mirror of
https://github.com/wassname/IndicoIo-python.git
synced 2026-08-17 11:13:07 +08:00
Merge pull request #107 from IndicoDataSolutions/Chris/fer-detect-no-resize
FIX: FER No resize for detect
This commit is contained in:
@@ -27,6 +27,10 @@ def fer(image, cloud=None, batch=False, api_key=None, **kwargs):
|
|||||||
:type image: list of lists
|
:type image: list of lists
|
||||||
:rtype: Dictionary containing emotion probability pairs
|
:rtype: Dictionary containing emotion probability pairs
|
||||||
"""
|
"""
|
||||||
image = image_preprocess(image, batch=batch)
|
|
||||||
|
image = image_preprocess(image, batch=batch,
|
||||||
|
size=None if kwargs.get("detect") else (48, 48)
|
||||||
|
)
|
||||||
|
|
||||||
url_params = {"batch": batch, "api_key": api_key}
|
url_params = {"batch": batch, "api_key": api_key}
|
||||||
return api_handler(image, cloud=cloud, api="fer", url_params=url_params, **kwargs)
|
return api_handler(image, cloud=cloud, api="fer", url_params=url_params, **kwargs)
|
||||||
|
|||||||
@@ -41,7 +41,8 @@ def image_preprocess(image, size=(48,48), min_axis=None, batch=False):
|
|||||||
else:
|
else:
|
||||||
raise IndicoError("Image must be a filepath, base64 encoded string, or a numpy array")
|
raise IndicoError("Image must be a filepath, base64 encoded string, or a numpy array")
|
||||||
|
|
||||||
out_image = resize_image(out_image, size, min_axis)
|
if size or min_axis:
|
||||||
|
out_image = resize_image(out_image, size, min_axis)
|
||||||
|
|
||||||
# convert to base64
|
# convert to base64
|
||||||
temp_output = StringIO.StringIO()
|
temp_output = StringIO.StringIO()
|
||||||
@@ -60,7 +61,7 @@ def resize_image(image, size, min_axis):
|
|||||||
warnings.warn(
|
warnings.warn(
|
||||||
"An aspect ratio greater than 10:1 is not recommended",
|
"An aspect ratio greater than 10:1 is not recommended",
|
||||||
Warning
|
Warning
|
||||||
)
|
)
|
||||||
size_arr = [0,0]
|
size_arr = [0,0]
|
||||||
size_arr[min_idx] = min_axis
|
size_arr[min_idx] = min_axis
|
||||||
size_arr[other_idx] = int(min_axis * aspect)
|
size_arr[other_idx] = int(min_axis * aspect)
|
||||||
@@ -89,4 +90,3 @@ def get_element_type(_list, dimens):
|
|||||||
for _ in xrange(len(dimens)):
|
for _ in xrange(len(dimens)):
|
||||||
elem = elem[0]
|
elem = elem[0]
|
||||||
return type(elem)
|
return type(elem)
|
||||||
|
|
||||||
|
|||||||
@@ -87,6 +87,13 @@ class BatchAPIRun(unittest.TestCase):
|
|||||||
self.assertTrue(isinstance(response, list))
|
self.assertTrue(isinstance(response, list))
|
||||||
self.assertTrue(isinstance(response[0], dict))
|
self.assertTrue(isinstance(response[0], dict))
|
||||||
|
|
||||||
|
def test_fer_detect(self):
|
||||||
|
test_data = os.path.normpath(os.path.join(DIR, "data/fear.png"))
|
||||||
|
response = fer(test_data, api_key=self.api_key, detect=True)
|
||||||
|
self.assertIsInstance(response, list)
|
||||||
|
self.assertEqual(len(response), 1)
|
||||||
|
self.assertIn("location", response[0])
|
||||||
|
|
||||||
def test_batch_fer_pil_image(self):
|
def test_batch_fer_pil_image(self):
|
||||||
test_data = [Image.open(os.path.normpath(os.path.join(DIR, "data/fear.png")))]
|
test_data = [Image.open(os.path.normpath(os.path.join(DIR, "data/fear.png")))]
|
||||||
response = fer(test_data, api_key=self.api_key)
|
response = fer(test_data, api_key=self.api_key)
|
||||||
|
|||||||
Reference in New Issue
Block a user