| from __future__ import annotations | |
| from pathlib import Path | |
| from .io import write_json | |
| from .model import RouterModel | |
| from .schemas import RewardRecord | |
| from .training_common import evaluate_loader, make_loader, resolve_device | |
| def evaluate_checkpoint( | |
| records: list[RewardRecord], | |
| checkpoint: str | Path, | |
| split: str = "test", | |
| output_path: str | Path | None = None, | |
| ) -> dict: | |
| model = RouterModel.from_checkpoint(checkpoint, device=resolve_device()) | |
| if model.worker_ids != records[0].worker_ids: | |
| raise ValueError("Checkpoint worker order does not match reward dataset") | |
| loader = make_loader(records, model, split, batch_size=32, shuffle=False) | |
| report = evaluate_loader(model, loader) | |
| report["split"] = split | |
| report["checkpoint"] = str(checkpoint) | |
| if output_path: | |
| write_json(output_path, report) | |
| return report | |