refactor: emit strategy metrics as floats

- Converts detached strategy metrics before returning loss output
- Removes redundant item conversion from the trainer loop
- Updates the documented contract and regression tests
This commit is contained in:
2026-08-02 06:38:28 +08:00
parent 1c7369f293
commit 020e2eff4e
4 changed files with 10 additions and 14 deletions
+3 -3
View File
@@ -15,7 +15,7 @@ from astrai.trainer.rollout import RolloutResult
class LossOutput(TypedDict):
loss: Tensor
metrics: Dict[str, Tensor]
metrics: Dict[str, float]
class LogprobsOutput(TypedDict):
@@ -148,14 +148,14 @@ class BaseStrategy(ABC):
metrics["loss"] = total_loss
return {
"loss": total_loss,
"metrics": {name: value.detach() for name, value in metrics.items()},
"metrics": {name: value.detach().item() for name, value in metrics.items()},
}
@staticmethod
def _normalize_output(output: Union[LossOutput, Tensor]) -> LossOutput:
if isinstance(output, dict):
return output
return {"loss": output, "metrics": {"loss": output.detach()}}
return {"loss": output, "metrics": {"loss": output.detach().item()}}
def supports_online(self) -> bool:
"""Whether this strategy can operate with a rollout runner.
+1 -4
View File
@@ -84,10 +84,7 @@ class Trainer:
self._call_callbacks("on_batch_begin", context)
loss_output = context.strategy(batch)
context.loss = loss_output["loss"].item()
context.metrics = {
name: value.item()
for name, value in loss_output["metrics"].items()
}
context.metrics = loss_output["metrics"]
stand_loss = loss_output["loss"] / executor.grad_accum_steps
executor.backward(stand_loss)
context.consumed_samples += (