feat: add MoE auxiliary loss metrics
- Propagates MoE load-balancing loss through model outputs - Logs task, auxiliary, and weighted losses across strategies - Computes only explicitly requested callback metrics - Preserves tensor compute_loss API and adds regression tests
This commit is contained in:
@@ -38,6 +38,7 @@ class TrainContext:
|
||||
epoch: int = field(default=0)
|
||||
consumed_samples: int = field(default=0)
|
||||
loss: float = field(default=0.0)
|
||||
metrics: Dict[str, float] = field(default_factory=dict)
|
||||
grad_norm: Optional[float] = field(default=None)
|
||||
grad_snr_tracker: GradSNRTracker = field(default_factory=GradSNRTracker)
|
||||
val_dataloader: Optional[DataLoader] = field(default=None)
|
||||
@@ -221,6 +222,7 @@ class TrainContextBuilder:
|
||||
obj.load_state_dict(extra[name])
|
||||
|
||||
strategy_kwargs = dict(cfg.extra_kwargs)
|
||||
strategy_kwargs.setdefault("moe_aux_loss_coef", cfg.moe_aux_loss_coef)
|
||||
|
||||
needs_ref = cfg.strategy in (
|
||||
"dpo",
|
||||
|
||||
Reference in New Issue
Block a user