diff --git a/src/evaluate/module.py b/src/evaluate/module.py index ca38b9b1..2116d026 100644 --- a/src/evaluate/module.py +++ b/src/evaluate/module.py @@ -982,6 +982,10 @@ def compute(self, predictions=None, references=None, **kwargs): def _merge_results(self, results): merged_results = {} + results = [ + result if isinstance(result, dict) else {module_name: result} + for module_name, result in zip(self.evaluation_module_names, results) + ] results_keys = list(itertools.chain.from_iterable([r.keys() for r in results])) duplicate_keys = {item for item, count in collections.Counter(results_keys).items() if count > 1} diff --git a/tests/test_metric.py b/tests/test_metric.py index 598b0f92..ef8edfae 100644 --- a/tests/test_metric.py +++ b/tests/test_metric.py @@ -757,3 +757,14 @@ def test_modules_from_string_poslabel(self): self.assertDictEqual( expected_result, combined_evaluation.compute(predictions=predictions, references=references, pos_label=0) ) + + def test_combined_evaluation_with_scalar_results(self): + predictions = ["this is the prediction", "there is an other sample"] + references = ["this is the reference", "there is another one"] + expected_result = {"wer": 0.5, "cer": 0.34146341463414637} + + combined_evaluation = combine(["wer", "cer"]) + + self.assertDictEqual( + expected_result, combined_evaluation.compute(predictions=predictions, references=references) + )