diff --git a/giant/training/trainers.py b/giant/training/trainers.py index 770980d..a998e77 100644 --- a/giant/training/trainers.py +++ b/giant/training/trainers.py @@ -618,9 +618,12 @@ class FlowDDPMStageTrainer(StageTrainer): l_balance = l_proc = l_entropy = torch.zeros((), device=device) if self.router is not None: - l_balance = self.router.balance_loss(cond_cont, cond_cat) - l_proc = self.router.classify_loss(cond_cont, cond_cat, proc_idx) - l_entropy = self.router.entropy_loss(cond_cont, cond_cat) + if self.spec.lambda_balance > 0: + l_balance = self.router.balance_loss(cond_cont, cond_cat) + if self.spec.lambda_proc > 0: + l_proc = self.router.classify_loss(cond_cont, cond_cat, proc_idx) + if self.spec.lambda_entropy > 0: + l_entropy = self.router.entropy_loss(cond_cont, cond_cat) total = self.spec.lambda_weight * l_gen + self.spec.n_sec_lambda * l_nsec + self.particle_type_lambda * l_type if self.spec.lambda_balance > 0: diff --git a/tests/test_train.py b/tests/test_train.py index ef0787f..6815309 100644 --- a/tests/test_train.py +++ b/tests/test_train.py @@ -5,6 +5,7 @@ import csv import math import tempfile from pathlib import Path +from unittest.mock import MagicMock import pytest import torch @@ -587,6 +588,60 @@ def test_stage_spec_from_config_omitted_decoder_and_particle_type_match_default_ assert spec.particle_type.target == "onehot" +def _routed_stage1_trainer(lambda_balance, lambda_proc, lambda_entropy): + cfg = _base_cfg() + cfg["stage1_model"]["router"] = { + "enabled": True, + "type": "energy", + "n_experts": 3, + "temperature": 0.5, + "learn_centers": True, + "lambda_balance": lambda_balance, + "lambda_proc": lambda_proc, + "lambda_entropy": lambda_entropy, + } + model_config = _model_config(cfg) + models = build_models(model_config) + critics = build_critics(model_config) + trainers = build_stage_trainers(cfg, models, critics, torch.device("cpu"), total_train_batches=4) + return trainers["stage1"] + + +def test_router_aux_losses_skipped_when_lambda_zero_but_run_when_positive(): + """Gitea #31: FlowDDPMStageTrainer._compute must not call + router.balance_loss/classify_loss/entropy_loss when the corresponding + lambda is 0 (the default) -- those calls do their own router.gate(...) + forward pass that is wasted once the term is masked out of the total + loss anyway. Checked both ways: zero lambdas must skip all three calls, + positive lambdas must still make them (the guard must not accidentally + suppress the real path).""" + batch = _fake_batches(1, 4)[0] + device = torch.device("cpu") + + trainer_zero = _routed_stage1_trainer(0.0, 0.0, 0.0) + router_zero = trainer_zero.router + router_zero.balance_loss = MagicMock(wraps=router_zero.balance_loss) + router_zero.classify_loss = MagicMock(wraps=router_zero.classify_loss) + router_zero.entropy_loss = MagicMock(wraps=router_zero.entropy_loss) + stats_zero = trainer_zero.step(batch, device, global_step=1) + assert router_zero.balance_loss.call_count == 0 + assert router_zero.classify_loss.call_count == 0 + assert router_zero.entropy_loss.call_count == 0 + assert stats_zero["loss_balance"] == 0.0 + assert stats_zero["loss_proc"] == 0.0 + assert stats_zero["loss_entropy"] == 0.0 + + trainer_pos = _routed_stage1_trainer(0.1, 0.1, 0.01) + router_pos = trainer_pos.router + router_pos.balance_loss = MagicMock(wraps=router_pos.balance_loss) + router_pos.classify_loss = MagicMock(wraps=router_pos.classify_loss) + router_pos.entropy_loss = MagicMock(wraps=router_pos.entropy_loss) + trainer_pos.step(batch, device, global_step=1) + assert router_pos.balance_loss.call_count == 1 + assert router_pos.classify_loss.call_count == 1 + assert router_pos.entropy_loss.call_count == 1 + + # --- AR trainer wiring (v0.3.0 step 5) --------------------------------------