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
+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 += (