Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
21 changes: 14 additions & 7 deletions src/pyrecest/evaluation/group_results_by_filter.py
Original file line number Diff line number Diff line change
Expand Up @@ -20,13 +20,20 @@ def group_results_by_filter(data):
# Remove the 'name' key-value pair from the entry
entry_values = {k: v for k, v in entry.items() if k != "name"}

# Check if the name already exists in the output_dict
if name in output_dict:
for key, value in entry_values.items():
# Append values to the existing lists
output_dict[name][key].append(value)
else:
# Initialize the entry in the output_dict with lists for each value
if name not in output_dict:
output_dict[name] = {k: [v] for k, v in entry_values.items()}
continue

grouped_values = output_dict[name]
existing_row_count = len(next(iter(grouped_values.values()), []))

# Backfill columns that first appear in a later row so every column remains aligned.
for key in entry_values:
if key not in grouped_values:
grouped_values[key] = [None] * existing_row_count

# Append one value per known column, using None for fields omitted by this row.
for key in grouped_values:
grouped_values[key].append(entry_values.get(key))

return output_dict
13 changes: 13 additions & 0 deletions tests/test_group_results_by_filter.py
Original file line number Diff line number Diff line change
Expand Up @@ -16,6 +16,19 @@ def test_mixed_parameters_do_not_crash_sorting(self):
self.assertEqual(grouped["pf"]["parameter"], [None, 1, "b"])
self.assertEqual(grouped["pf"]["score"], [0.0, 1.0, 3.0])

def test_heterogeneous_metric_keys_stay_aligned(self):
rows = [
{"name": "pf", "parameter": 3, "std": 0.3},
{"name": "pf", "parameter": 1, "score": 1.0},
{"name": "pf", "parameter": 2, "score": 2.0, "std": 0.2},
]

grouped = group_results_by_filter(rows)

self.assertEqual(grouped["pf"]["parameter"], [1, 2, 3])
self.assertEqual(grouped["pf"]["score"], [1.0, 2.0, None])
self.assertEqual(grouped["pf"]["std"], [None, 0.2, 0.3])


if __name__ == "__main__":
unittest.main()
Loading