fugu-lite / src /fugu_lite /evaluate.py
tahsinsoyak's picture
Upload private Fugu-Lite V1 snapshot
88e15cd verified
Raw
History Blame
876 Bytes
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