From 020e2eff4ed292ee9496e9276a25af6b70c2a961 Mon Sep 17 00:00:00 2001 From: ViperEkura <3081035982@qq.com> Date: Sun, 2 Aug 2026 06:38:28 +0800 Subject: [PATCH] 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 --- astrai/trainer/strategy.py | 6 +++--- astrai/trainer/trainer.py | 5 +---- docs/guides/training.md | 2 +- tests/trainer/test_loss_output.py | 11 +++++------ 4 files changed, 10 insertions(+), 14 deletions(-) diff --git a/astrai/trainer/strategy.py b/astrai/trainer/strategy.py index 07d0d53..61573f4 100644 --- a/astrai/trainer/strategy.py +++ b/astrai/trainer/strategy.py @@ -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. diff --git a/astrai/trainer/trainer.py b/astrai/trainer/trainer.py index a454f98..cb1d5eb 100644 --- a/astrai/trainer/trainer.py +++ b/astrai/trainer/trainer.py @@ -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 += ( diff --git a/docs/guides/training.md b/docs/guides/training.md index ddbe770..fcc3117 100644 --- a/docs/guides/training.md +++ b/docs/guides/training.md @@ -89,7 +89,7 @@ on_train_end Default callbacks (in order): `gradient_checkpointing` (activation checkpointing, optional), `checkpoint` (safetensors, rank-0), `metric` (JSONL + validation, rank-0), `progress_bar` (tqdm), `gradient_clipping` (always registered; computes grad norm, clips only when `max_grad_norm` is not `None`). -Strategies return `{"loss": Tensor, "metrics": Dict[str, Tensor]}` when called by the trainer. Built-in metrics include the task-specific loss and, for MoE models, `moe_aux_loss` plus `moe_aux_loss_weighted`. Direct `compute_loss(batch)` calls continue to return a single loss tensor. +Strategies return `{"loss": Tensor, "metrics": Dict[str, float]}` when called by the trainer. Built-in metrics include the task-specific loss and, for MoE models, `moe_aux_loss` plus `moe_aux_loss_weighted`. Direct `compute_loss(batch)` calls continue to return a single loss tensor. ## Strategies diff --git a/tests/trainer/test_loss_output.py b/tests/trainer/test_loss_output.py index c1372a6..5a2601a 100644 --- a/tests/trainer/test_loss_output.py +++ b/tests/trainer/test_loss_output.py @@ -1,5 +1,6 @@ from types import SimpleNamespace +import pytest import torch from astrai.model.transformer import AutoRegressiveLM @@ -33,16 +34,14 @@ def test_seq_strategy_combines_and_reports_moe_aux_loss(device): "moe_aux_loss", "moe_aux_loss_weighted", } - torch.testing.assert_close( - output["loss"], + assert output["loss"].item() == pytest.approx( output["metrics"]["task_loss"] + output["metrics"]["moe_aux_loss_weighted"], ) - torch.testing.assert_close( - output["metrics"]["moe_aux_loss_weighted"], + assert output["metrics"]["moe_aux_loss_weighted"] == pytest.approx( 0.25 * output["metrics"]["moe_aux_loss"], ) assert output["loss"].requires_grad - assert all(not metric.requires_grad for metric in output["metrics"].values()) + assert all(isinstance(metric, float) for metric in output["metrics"].values()) def test_metric_callback_includes_dynamic_strategy_metrics(tmp_path): @@ -103,4 +102,4 @@ def test_legacy_strategy_tensor_loss_is_normalized(): output = strategy({}) assert output["loss"].item() == 2.0 - assert output["metrics"]["loss"].item() == 2.0 + assert output["metrics"]["loss"] == 2.0