|
1 | 1 | from pathlib import Path |
2 | | -from typing import Any, Dict, List, Optional, Union |
| 2 | +from typing import Any, Callable, Dict, List, Optional, Union |
3 | 3 | import asyncio |
4 | 4 |
|
5 | 5 | from tqdm.auto import tqdm |
@@ -28,6 +28,7 @@ def __init__( |
28 | 28 | show_progress: bool = True, |
29 | 29 | n_repetitions: int = 1, |
30 | 30 | save_dir: Optional[str] = None, |
| 31 | + on_model_done: Optional[Callable[[str, "RepeatedExperimentResults"], None]] = None, |
31 | 32 | ): |
32 | 33 | if not models or any("model" not in m for m in models): |
33 | 34 | raise ValueError("Models must be dicts with a 'model' key.") |
@@ -58,6 +59,7 @@ def __init__( |
58 | 59 | self.show_progress = show_progress |
59 | 60 | self.n_repetitions = n_repetitions |
60 | 61 | self.save_dir = Path(save_dir) if save_dir else None |
| 62 | + self.on_model_done = on_model_done |
61 | 63 |
|
62 | 64 | def _make_label(self, model_info: Dict[str, Any]) -> str: |
63 | 65 | return model_info.get("label") or model_info["model"] |
@@ -156,6 +158,9 @@ async def run_async( |
156 | 158 | pbar_reps.update(1) |
157 | 159 |
|
158 | 160 | runs_by_model[label] = [runs_ordered[i] for i in range(self.n_repetitions)] |
| 161 | + if self.on_model_done: |
| 162 | + partial = RepeatedExperimentResults({label: runs_by_model[label]}) |
| 163 | + self.on_model_done(label, partial) |
159 | 164 | pbar_models.update(1) |
160 | 165 |
|
161 | 166 | judge_info = { |
|
0 commit comments