From 5623f6de9068af35ab7cd15d6e92224f100204c8 Mon Sep 17 00:00:00 2001 From: Madison May Date: Fri, 5 Jun 2015 04:06:14 -0400 Subject: [PATCH] Updated format for batch --- indicoio/utils/multi.py | 4 ---- tests/test_remote.py | 12 ++++++------ 2 files changed, 6 insertions(+), 10 deletions(-) diff --git a/indicoio/utils/multi.py b/indicoio/utils/multi.py index a947e05..c023e62 100644 --- a/indicoio/utils/multi.py +++ b/indicoio/utils/multi.py @@ -28,10 +28,6 @@ def multi(data, type, apis, available, batch=False, **kwargs): # Convert client api names to server names before sending request apis = map(CLIENT_SERVER_MAP.get, apis) result = api_handler(data, apis=apis, batch=batch, **kwargs) - - if batch: - return [handle_response(each) for each in result] - return handle_response(result) def handle_response(result): diff --git a/tests/test_remote.py b/tests/test_remote.py index 3b50c55..62ae094 100644 --- a/tests/test_remote.py +++ b/tests/test_remote.py @@ -126,22 +126,22 @@ class BatchAPIRun(unittest.TestCase): test_data = [generate_array((48,48))] response = batch_predict_image(test_data, apis=config.IMAGE_APIS, api_key=self.api_key) - self.assertTrue(isinstance(response, list)) - self.assertTrue(set(response[0].keys()) == set(config.IMAGE_APIS)) + self.assertTrue(isinstance(response, dict)) + self.assertTrue(set(response.keys()) == set(config.IMAGE_APIS)) def test_batch_multi_api_text(self): test_data = ['clearly an english sentence'] response = batch_predict_text(test_data, apis=config.TEXT_APIS, api_key=self.api_key) - self.assertTrue(isinstance(response, list)) - self.assertTrue(set(response[0].keys()) == set(config.TEXT_APIS)) + self.assertTrue(isinstance(response, dict)) + self.assertTrue(set(response.keys()) == set(config.TEXT_APIS)) def test_default_multi_api_text(self): test_data = ['clearly an english sentence'] response = batch_predict_text(test_data, api_key=self.api_key) - self.assertTrue(isinstance(response, list)) - self.assertTrue(set(response[0].keys()) == set(config.TEXT_APIS)) + self.assertTrue(isinstance(response, dict)) + self.assertTrue(set(response.keys()) == set(config.TEXT_APIS)) def test_multi_api_bad_api(self): self.assertRaises(ValueError,