Deduplicate giant/training/trainers.py shared per-stage logic
CI / Format (ruff format) (push) Successful in 29s
CI / Lint (ruff check) (push) Successful in 30s
CI / Sync project version with tag (push) Has been skipped
CI / Lint (ruff check) (pull_request) Successful in 40s
CI / Type check (ty) (push) Successful in 43s
CI / Format (ruff format) (pull_request) Successful in 32s
CI / Sync project version with tag (pull_request) Has been skipped
CI / Type check (ty) (pull_request) Successful in 31s
CI / Tests (pull_request) Successful in 3m12s
CI / Tests (push) Successful in 3m17s

Lift repeated per-batch operations into StageTrainer base-class helpers so
each is written once instead of being copy-pasted between FlowDDPMStageTrainer
and WGANStageTrainer:

- _n_sec_loss: the multiplicity classifier (stage1/stage2 predict_n_sec split
  + cross-entropy + accuracy), previously written three times. Gated on
  n_sec_head presence, not n_sec.mode, so a future stop_token model trains its
  EOS signal elsewhere and this stays zero.
- _sec_mask: the arange < n_sec prefix mask, previously in two places.
- _step_optimizer: the zero_grad/backward/clip_grad_norm_(1.0)/step quad,
  previously written three times; now the single home of the clip constant.
- _sec_target: collapses the byte-identical _ar_target/_real wrappers into one
  flatten-parameterized method (they differed only by .flatten(1)).

Also trim StageSpec.from_config to read DEFAULT_CONFIG-guaranteed train.* keys
directly instead of re-defaulting them.

The three particle-type targets (onehot CE, physical/embedding regression) and
_type_loss are intentionally left as separate paths — genuinely different
objectives, not duplication.

stage2_inputs.py: extract the shared _ar_meta helper for the has_prev/
remaining_frac/slot_idx trio used by both AR-input assemblers.

Behavior-preserving: same losses, optimizer order, and RNG draw order. Full
test suite (699) green; ruff + ty clean.

Co-Authored-By: Claude Opus 4.8 <noreply@anthropic.com>
This commit is contained in:
2026-08-10 10:32:26 +02:00
parent d3271bc798
commit fff61ebd61
2 changed files with 181 additions and 146 deletions
+22 -13
View File
@@ -138,6 +138,25 @@ def _ar_has_prev(k_max: int, device: torch.device) -> torch.Tensor:
return (torch.arange(k_max, device=device) >= 1).unsqueeze(0)
def _ar_meta(
k_max: int, batch: int, device: torch.device, fraction: torch.Tensor
) -> dict[str, torch.Tensor]:
"""`has_prev`/`remaining_frac`/`slot_idx` — the three per-token AR
conditioning tensors that don't depend on *which* history representation
(ground truth vs. the scheduled-sampling mix) produced `fraction`.
Shared by `_assemble_stage2_ar_inputs` and
`_assemble_stage2_ar_inputs_scheduled`, which differ only in
`history_feat`."""
slot_idx = (
torch.arange(k_max, device=device).float() / max(k_max - 1, 1)
).unsqueeze(0)
return {
"has_prev": _ar_has_prev(k_max, device).expand(batch, -1),
"remaining_frac": _remaining_energy_fraction(fraction),
"slot_idx": slot_idx.expand(batch, -1),
}
def _assemble_stage2_ar_inputs(
sec_cont: torch.Tensor,
sec_type_idx: torch.Tensor,
@@ -161,13 +180,7 @@ def _assemble_stage2_ar_inputs(
],
dim=-1,
)
slot_idx = (torch.arange(K, device=device).float() / max(K - 1, 1)).unsqueeze(0)
return {
"history_feat": history_feat,
"has_prev": _ar_has_prev(K, device).expand(B, -1),
"remaining_frac": _remaining_energy_fraction(fraction),
"slot_idx": slot_idx.expand(B, -1),
}
return {"history_feat": history_feat, **_ar_meta(K, B, device, fraction)}
def _stage2_tf_prob(
@@ -259,7 +272,7 @@ def _assemble_stage2_ar_inputs_scheduled(
model.train()
fraction_gt = _stick_fraction(sec_cont)
dir_gt = sec_cont[..., 1:4]
dir_gt = sec_cont[..., 1:CONT_SLOT_DIM]
type_repr_gt = _type_repr(
sec_type_idx, sec_cont, particle_type_cfg, cond_enc, emb_dim
)
@@ -275,11 +288,7 @@ def _assemble_stage2_ar_inputs_scheduled(
own_feat = torch.cat([fraction.unsqueeze(-1), direction, type_repr], dim=-1)
return {
"history_feat": _shift_prev(own_feat),
"has_prev": _ar_has_prev(K, device).expand(B, -1),
"remaining_frac": _remaining_energy_fraction(fraction),
"slot_idx": (torch.arange(K, device=device).float() / max(K - 1, 1))
.unsqueeze(0)
.expand(B, -1),
**_ar_meta(K, B, device, fraction),
}
+159 -133
View File
@@ -34,7 +34,6 @@ from giant.training.metrics import MetricSpec, stage_metric, train_metric, val_m
from giant.training.stage2_inputs import (
_assemble_stage2_ar_inputs_scheduled,
_assemble_stage2_ar_target,
_assemble_stage2_real,
_gumbel_tau,
_relax_onehot_type_slice,
_stage2_tf_prob,
@@ -148,9 +147,12 @@ class StageSpec:
particle_type=cfg["stage2_model"].get("particle_type")
or {"target": "physical"},
particle_type_emb_dim=cfg["conditioning"]["particle"]["emb_dim"],
# train.* keys are all guaranteed by DEFAULT_CONFIG's deep-merge
# (giant/config.py), so they read directly; the field defaults
# below exist only for tests that construct StageSpec by hand.
lr=t["lr"],
weight_decay=t.get("weight_decay", 0.01),
ema_decay=t.get("ema_decay", 0.9999),
weight_decay=t["weight_decay"],
ema_decay=t["ema_decay"],
warmup_epochs=t["warmup_epochs"],
epochs=t["epochs"],
steps_per_epoch=max(steps_per_epoch, 1),
@@ -166,7 +168,7 @@ class StageSpec:
# train.validate_steps as its flow-matching ODE step count — no
# dedicated config key for this (docs/v0.3.0-design.md §3.3 lists
# tf_p_start/tf_p_end/attn_n_heads/attn_n_layers only).
ar_sample_steps=t.get("validate_steps", 10),
ar_sample_steps=t["validate_steps"],
ddpm_n_steps=stage_cfg.get("ddpm", {}).get("n_steps", 1000),
n_critic=wgan_cfg.get("n_critic", 5),
gp_weight=wgan_cfg.get("gp_weight", 10.0),
@@ -288,6 +290,125 @@ class StageTrainer:
for module in self._modules:
module.eval()
# --- stage-2 secondary assembly (shared by both trainer subclasses) ---
def _ar_inputs(
self,
cond_cont: torch.Tensor,
cond_cat: torch.Tensor,
stage1_ctx: torch.Tensor,
sec_cont: torch.Tensor,
sec_type_idx: torch.Tensor,
n_sec: torch.Tensor,
epoch: int | None,
) -> dict[str, torch.Tensor]:
"""Per-token AR conditioning for this stage's secondary decoder.
`epoch=None` means full teacher forcing (`p_tf=1.0`) regardless of
`spec.teacher_forcing` — the val-loss convention, kept in this one
place so both trainer subclasses honor it identically.
"""
p_tf = (
1.0
if epoch is None
else _stage2_tf_prob(
self.spec.teacher_forcing,
self.spec.tf_p_start,
self.spec.tf_p_end,
epoch,
self.spec.epochs,
)
)
return _assemble_stage2_ar_inputs_scheduled(
self.model,
cond_cont,
cond_cat,
stage1_ctx,
sec_cont,
sec_type_idx,
n_sec,
self.particle_type_cfg,
self.model.cond_enc,
self.particle_type_emb_dim,
p_tf,
self.spec.ar_sample_steps,
)
def _sec_target(
self,
sec_cont: torch.Tensor,
sec_type_idx: torch.Tensor,
generator: str,
*,
flatten: bool,
) -> torch.Tensor:
"""Ground-truth stage-2 target for this stage's secondary decoder,
per the (particle-type target, generator) width rules in
`_assemble_stage2_ar_target`. `flatten=True` gives `Stage2OneShot`'s
flattened `(B, K*token_dim)` form (the old `_real`); `flatten=False`
gives `Stage2Autoregressive`'s per-token `(B, K, token_dim)` form (the
old `_ar_target`) — the two are the same tensor modulo `.flatten(1)`,
so the width rules live in one place (`stage2_inputs.py`)."""
target = _assemble_stage2_ar_target(
sec_cont,
sec_type_idx,
self.particle_type_cfg,
generator,
self.model.cond_enc,
self.particle_type_emb_dim,
)
return target.flatten(1) if flatten else target
@staticmethod
def _sec_mask(
n_sec: torch.Tensor, k_max: int, device: torch.device
) -> torch.Tensor:
"""`(B, K_MAX)` bool prefix mask: slot k is valid iff `k < n_sec`."""
return torch.arange(k_max, device=device).unsqueeze(0) < n_sec.unsqueeze(1)
def _n_sec_loss(
self,
cond_cont: torch.Tensor,
cond_cat: torch.Tensor,
stage1_ctx: torch.Tensor,
n_sec: torch.Tensor,
device: torch.device,
) -> tuple[torch.Tensor, torch.Tensor]:
"""`(l_nsec, nsec_acc)` for this stage's multiplicity classifier —
zeros when the stage owns no `n_sec_head` (stage 1 now that n_sec
defaults to stage 2, or any stage without the head). Owns the only
stage1-vs-stage2 `predict_n_sec` signature split, shared by the
non-adversarial and WGAN trainers.
Gated on `n_sec_head is None`, not on `n_sec.mode`: a future
`mode="stop_token"` model (design doc §11.2, currently rejected in
`validate_config`) carries no head and would train its EOS signal in
the generator/AR loss path instead, so this correctly stays zero.
"""
if self.model.n_sec_head is None:
zero = torch.zeros((), device=device)
return zero, zero
logits = (
self.model.predict_n_sec(cond_cont, cond_cat, stage1_ctx)
if self.is_stage2
else self.model.predict_n_sec(cond_cont, cond_cat)
)
l_nsec = F.cross_entropy(logits, n_sec)
nsec_acc = (logits.argmax(dim=-1) == n_sec).float().mean()
return l_nsec, nsec_acc
@staticmethod
def _step_optimizer(
optimizer: optim.Optimizer, loss: torch.Tensor, params: list
) -> float:
"""`zero_grad -> backward -> clip_grad_norm_(1.0) -> step`, returning
the pre-clip grad norm. The one place the grad-clip constant lives."""
optimizer.zero_grad()
loss.backward()
grad_norm = torch.nn.utils.clip_grad_norm_(params, 1.0)
optimizer.step()
return grad_norm.item()
def _extra_state(self) -> dict:
"""Subclass state beyond model/optimizer/lr_sched/EMA."""
return {}
@@ -342,11 +463,11 @@ class FlowDDPMStageTrainer(StageTrainer):
)
super().__init__(spec, model, device)
self.particle_type_lambda = self.particle_type_cfg.get("lambda", 1.0)
# Width of the type slice actually folded into x1_s2 by
# _assemble_stage2_real, under this trainer's generator (flow/ddpm
# only — see the NotImplementedError above): "physical" keeps it
# folded in (PARTICLE_PHYS_DIM wide, unchanged from v0.2); "onehot"/
# "embedding" pull it out into model.type_head instead (0 here).
# Width of the type slice actually folded into x1_s2 by _sec_target,
# under this trainer's generator (flow/ddpm only — see the
# NotImplementedError above): "physical" keeps it folded in
# (PARTICLE_PHYS_DIM wide, unchanged from v0.2); "onehot"/"embedding"
# pull it out into model.type_head instead (0 here).
self._flow_type_dim = (
None
if self.particle_type_cfg.get("target", "physical") == "physical"
@@ -396,13 +517,6 @@ class FlowDDPMStageTrainer(StageTrainer):
]
self.stage_metrics = [stage_metric("lr")]
def _predict_n_sec(self, cond_cont, cond_cat, stage1_ctx):
if self.model.n_sec_head is None:
return None
if self.is_stage2:
return self.model.predict_n_sec(cond_cont, cond_cat, stage1_ctx)
return self.model.predict_n_sec(cond_cont, cond_cat)
def _generator_loss(
self, cond_cont, cond_cat, x1_s1, x1_s2, sec_mask, stage1_ctx, ar_inputs=None
):
@@ -501,66 +615,29 @@ class FlowDDPMStageTrainer(StageTrainer):
proc_idx,
sec_type_idx,
) = _batch_to_device(batch, device)
sec_mask = torch.arange(sec_cont.size(1), device=device).unsqueeze(
0
) < n_sec.unsqueeze(1)
sec_mask = self._sec_mask(n_sec, sec_cont.size(1), device)
stage1_ctx = x1_s1.detach()
x1_s2 = None
ar_inputs = None
if self.is_stage2 and self.decoder == "autoregressive":
p_tf = (
1.0
if epoch is None
else _stage2_tf_prob(
self.spec.teacher_forcing,
self.spec.tf_p_start,
self.spec.tf_p_end,
epoch,
self.spec.epochs,
)
ar_inputs = self._ar_inputs(
cond_cont, cond_cat, stage1_ctx, sec_cont, sec_type_idx, n_sec, epoch
)
ar_inputs = _assemble_stage2_ar_inputs_scheduled(
self.model,
cond_cont,
cond_cat,
stage1_ctx,
sec_cont,
sec_type_idx,
n_sec,
self.particle_type_cfg,
self.model.cond_enc,
self.particle_type_emb_dim,
p_tf,
self.spec.ar_sample_steps,
)
x1_s2 = _assemble_stage2_ar_target(
sec_cont,
sec_type_idx,
self.particle_type_cfg,
self.generator,
self.model.cond_enc,
self.particle_type_emb_dim,
x1_s2 = self._sec_target(
sec_cont, sec_type_idx, self.generator, flatten=False
)
elif self.is_stage2:
x1_s2 = _assemble_stage2_real(
sec_cont,
sec_type_idx,
self.particle_type_cfg,
self.generator,
self.model.cond_enc,
self.particle_type_emb_dim,
x1_s2 = self._sec_target(
sec_cont, sec_type_idx, self.generator, flatten=True
)
l_gen = self._generator_loss(
cond_cont, cond_cat, x1_s1, x1_s2, sec_mask, stage1_ctx, ar_inputs=ar_inputs
)
n_sec_logits = self._predict_n_sec(cond_cont, cond_cat, stage1_ctx)
l_nsec = torch.zeros((), device=device)
nsec_acc = torch.zeros((), device=device)
if n_sec_logits is not None:
l_nsec = F.cross_entropy(n_sec_logits, n_sec)
nsec_acc = (n_sec_logits.argmax(dim=-1) == n_sec).float().mean()
l_nsec, nsec_acc = self._n_sec_loss(
cond_cont, cond_cat, stage1_ctx, n_sec, device
)
l_type, type_acc = self._type_loss(
cond_cont,
@@ -612,15 +689,12 @@ class FlowDDPMStageTrainer(StageTrainer):
)
epoch = global_step // self.spec.steps_per_epoch
out = self._compute(batch, device, epoch=epoch)
self.optimizer.zero_grad()
out["loss"].backward()
grad_norm = torch.nn.utils.clip_grad_norm_(self.params, 1.0)
self.optimizer.step()
grad_norm = self._step_optimizer(self.optimizer, out["loss"], self.params)
self.lr_sched.step()
if self.ema_model is not None:
_update_ema(self.ema_model, self.model, self.ema_decay)
stats = {key: value.item() for key, value in out.items()}
stats["grad_norm"] = grad_norm.item()
stats["grad_norm"] = grad_norm
stats["lr"] = self.optimizer.param_groups[0]["lr"]
return stats
@@ -719,7 +793,7 @@ class WGANStageTrainer(StageTrainer):
slot_width = CONT_SLOT_DIM + type_dim
k_max = sec_cont.size(1)
sec_mask = torch.arange(k_max, device=device).unsqueeze(0) < n_sec.unsqueeze(1)
sec_mask = self._sec_mask(n_sec, k_max, device)
mask = sec_mask.unsqueeze(-1).expand(-1, -1, slot_width).reshape(B, -1).float()
def critic_fn(x):
@@ -727,36 +801,13 @@ class WGANStageTrainer(StageTrainer):
if self.decoder == "autoregressive":
epoch = global_step // self.spec.steps_per_epoch
p_tf = _stage2_tf_prob(
self.spec.teacher_forcing,
self.spec.tf_p_start,
self.spec.tf_p_end,
epoch,
self.spec.epochs,
)
ar = _assemble_stage2_ar_inputs_scheduled(
self.model,
cond_cont,
cond_cat,
stage1_ctx,
sec_cont,
sec_type_idx,
n_sec,
self.particle_type_cfg,
self.model.cond_enc,
self.particle_type_emb_dim,
p_tf,
self.spec.ar_sample_steps,
ar = self._ar_inputs(
cond_cont, cond_cat, stage1_ctx, sec_cont, sec_type_idx, n_sec, epoch
)
real = (
_assemble_stage2_ar_target(
sec_cont,
sec_type_idx,
self.particle_type_cfg,
"wgan",
self.model.cond_enc,
self.particle_type_emb_dim,
).reshape(B, -1)
self._sec_target(sec_cont, sec_type_idx, "wgan", flatten=False).reshape(
B, -1
)
* mask
)
z = torch.randn(B, k_max, self.model.noise_dim, device=device)
@@ -771,17 +822,7 @@ class WGANStageTrainer(StageTrainer):
ar["slot_idx"],
).reshape(B, -1)
else:
real = (
_assemble_stage2_real(
sec_cont,
sec_type_idx,
self.particle_type_cfg,
"wgan",
self.model.cond_enc,
self.particle_type_emb_dim,
)
* mask
)
real = self._sec_target(sec_cont, sec_type_idx, "wgan", flatten=True) * mask
z = torch.randn(B, self.model.noise_dim, device=device)
fake_raw = self.model(z, cond_cont, cond_cat, stage1_ctx)
@@ -848,31 +889,19 @@ class WGANStageTrainer(StageTrainer):
d_loss = fake_score.mean() - real_score.mean() + self.gp_weight * gp
wasserstein = (real_score.mean() - fake_score.mean()).detach()
self.optimizer_d.zero_grad()
d_loss.backward()
grad_norm_d = torch.nn.utils.clip_grad_norm_(self.d_params, 1.0)
self.optimizer_d.step()
grad_norm_d = self._step_optimizer(self.optimizer_d, d_loss, self.d_params)
# --- generator (+ n_sec) step ---
did_g_step = global_step % self.n_critic == 0
n_sec_logits = None
if self.model.n_sec_head is not None:
n_sec_logits = (
self.model.predict_n_sec(cond_cont, cond_cat)
if not self.is_stage2
else self.model.predict_n_sec(cond_cont, cond_cat, stage1_ctx)
)
l_nsec = torch.zeros((), device=device)
nsec_acc = torch.zeros((), device=device)
if n_sec_logits is not None:
l_nsec = F.cross_entropy(n_sec_logits, n_sec)
nsec_acc = (n_sec_logits.argmax(dim=-1) == n_sec).float().mean()
l_nsec, nsec_acc = self._n_sec_loss(
cond_cont, cond_cat, stage1_ctx, n_sec, device
)
# On a non-generator-step batch with no n_sec_head on this stage
# (n_sec now defaults to stage 2, decision 1), there's nothing for
# the generator optimizer to do this batch — g_loss would otherwise
# be a graph-less zero tensor, which .backward() rejects outright.
skip_g_step = not did_g_step and n_sec_logits is None
skip_g_step = not did_g_step and self.model.n_sec_head is None
if did_g_step:
g_loss_adv = generator_loss(critic_fn, fake)
g_loss = (
@@ -882,12 +911,9 @@ class WGANStageTrainer(StageTrainer):
g_loss_adv = torch.zeros((), device=device)
g_loss = self.spec.n_sec_lambda * l_nsec
if skip_g_step:
grad_norm_g = torch.zeros(())
grad_norm_g = 0.0
else:
self.optimizer.zero_grad()
g_loss.backward()
grad_norm_g = torch.nn.utils.clip_grad_norm_(self.g_params, 1.0)
self.optimizer.step()
grad_norm_g = self._step_optimizer(self.optimizer, g_loss, self.g_params)
if did_g_step:
self.lr_sched.step()
@@ -902,9 +928,9 @@ class WGANStageTrainer(StageTrainer):
"loss_nsec": l_nsec.item(),
"nsec_acc": nsec_acc.item(),
"did_g_step": did_g_step,
"grad_norm": grad_norm_d.item() + grad_norm_g.item(),
"grad_norm_d": grad_norm_d.item(),
"grad_norm_g": grad_norm_g.item(),
"grad_norm": grad_norm_d + grad_norm_g,
"grad_norm_d": grad_norm_d,
"grad_norm_g": grad_norm_g,
"grad_norm_type_slice": grad_probe.get("type", 0.0),
"grad_norm_cont_slice": grad_probe.get("cont", 0.0),
"lr": self.optimizer.param_groups[0]["lr"],