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:
@@ -15,7 +15,7 @@ from astrai.trainer.rollout import RolloutResult
|
|||||||
|
|
||||||
class LossOutput(TypedDict):
|
class LossOutput(TypedDict):
|
||||||
loss: Tensor
|
loss: Tensor
|
||||||
metrics: Dict[str, Tensor]
|
metrics: Dict[str, float]
|
||||||
|
|
||||||
|
|
||||||
class LogprobsOutput(TypedDict):
|
class LogprobsOutput(TypedDict):
|
||||||
@@ -148,14 +148,14 @@ class BaseStrategy(ABC):
|
|||||||
metrics["loss"] = total_loss
|
metrics["loss"] = total_loss
|
||||||
return {
|
return {
|
||||||
"loss": total_loss,
|
"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
|
@staticmethod
|
||||||
def _normalize_output(output: Union[LossOutput, Tensor]) -> LossOutput:
|
def _normalize_output(output: Union[LossOutput, Tensor]) -> LossOutput:
|
||||||
if isinstance(output, dict):
|
if isinstance(output, dict):
|
||||||
return output
|
return output
|
||||||
return {"loss": output, "metrics": {"loss": output.detach()}}
|
return {"loss": output, "metrics": {"loss": output.detach().item()}}
|
||||||
|
|
||||||
def supports_online(self) -> bool:
|
def supports_online(self) -> bool:
|
||||||
"""Whether this strategy can operate with a rollout runner.
|
"""Whether this strategy can operate with a rollout runner.
|
||||||
|
|||||||
@@ -84,10 +84,7 @@ class Trainer:
|
|||||||
self._call_callbacks("on_batch_begin", context)
|
self._call_callbacks("on_batch_begin", context)
|
||||||
loss_output = context.strategy(batch)
|
loss_output = context.strategy(batch)
|
||||||
context.loss = loss_output["loss"].item()
|
context.loss = loss_output["loss"].item()
|
||||||
context.metrics = {
|
context.metrics = loss_output["metrics"]
|
||||||
name: value.item()
|
|
||||||
for name, value in loss_output["metrics"].items()
|
|
||||||
}
|
|
||||||
stand_loss = loss_output["loss"] / executor.grad_accum_steps
|
stand_loss = loss_output["loss"] / executor.grad_accum_steps
|
||||||
executor.backward(stand_loss)
|
executor.backward(stand_loss)
|
||||||
context.consumed_samples += (
|
context.consumed_samples += (
|
||||||
|
|||||||
@@ -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`).
|
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
|
## Strategies
|
||||||
|
|
||||||
|
|||||||
@@ -1,5 +1,6 @@
|
|||||||
from types import SimpleNamespace
|
from types import SimpleNamespace
|
||||||
|
|
||||||
|
import pytest
|
||||||
import torch
|
import torch
|
||||||
|
|
||||||
from astrai.model.transformer import AutoRegressiveLM
|
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",
|
||||||
"moe_aux_loss_weighted",
|
"moe_aux_loss_weighted",
|
||||||
}
|
}
|
||||||
torch.testing.assert_close(
|
assert output["loss"].item() == pytest.approx(
|
||||||
output["loss"],
|
|
||||||
output["metrics"]["task_loss"] + output["metrics"]["moe_aux_loss_weighted"],
|
output["metrics"]["task_loss"] + output["metrics"]["moe_aux_loss_weighted"],
|
||||||
)
|
)
|
||||||
torch.testing.assert_close(
|
assert output["metrics"]["moe_aux_loss_weighted"] == pytest.approx(
|
||||||
output["metrics"]["moe_aux_loss_weighted"],
|
|
||||||
0.25 * output["metrics"]["moe_aux_loss"],
|
0.25 * output["metrics"]["moe_aux_loss"],
|
||||||
)
|
)
|
||||||
assert output["loss"].requires_grad
|
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):
|
def test_metric_callback_includes_dynamic_strategy_metrics(tmp_path):
|
||||||
@@ -103,4 +102,4 @@ def test_legacy_strategy_tensor_loss_is_normalized():
|
|||||||
output = strategy({})
|
output = strategy({})
|
||||||
|
|
||||||
assert output["loss"].item() == 2.0
|
assert output["loss"].item() == 2.0
|
||||||
assert output["metrics"]["loss"].item() == 2.0
|
assert output["metrics"]["loss"] == 2.0
|
||||||
|
|||||||
Reference in New Issue
Block a user