f8722e347e
CI / Lint (ruff check) (push) Successful in 30s
CI / Format (ruff format) (push) Successful in 29s
CI / Sync project version with tag (push) Has been skipped
CI / Lint (ruff check) (pull_request) Successful in 39s
CI / Type check (ty) (push) Successful in 44s
CI / Format (ruff format) (pull_request) Successful in 44s
CI / Sync project version with tag (pull_request) Has been skipped
CI / Type check (ty) (pull_request) Successful in 45s
CI / Tests (pull_request) Successful in 3m56s
CI / Tests (push) Successful in 3m58s
generator ∈ {"flow", "ddpm", "wgan"} was tested as a bare string in ~45
sites across models.py, sample.py, builders.py, trainers.py, and
stage2_inputs.py, each independently re-deriving one of five consequences
of the choice (needs a time embedding? what does the trunk take as input?
is the type slice folded into the trunk output? which sampler? which
loss?). giant/model/objectives.py adds an Objective ABC + OBJECTIVE_REGISTRY
+ build_objective factory, mirroring routers.py's Router pattern, and every
bare-string site now goes through it (needs_time, is_adversarial,
folds_type_slice, trunk_in_dim, build_schedule, stage1_loss/stage2_loss).
Per discussion: FlowDDPMStageTrainer and WGANStageTrainer stay separate
classes rather than merging into one StageTrainer as the issue's sketch
proposed — their training loops are genuinely different shapes (single loss
vs. dual G/D step with gradient penalty/n_critic/ST-Gumbel), and trainers.py
is the least-covered-by-fast-tests part of the codebase, so a full merge
was judged out of proportion to this issue's risk budget.
FlowDDPMStageTrainer's own loss dispatch (flow vs ddpm, one-shot vs AR) does
move onto the objective, so a future non-adversarial objective (rectified
flow, consistency distillation) is still a one-file, zero-trainer-edits
addition.
No config-schema change — stage{1,2}_model.generator stays the persisted
string, just looked up in the registry instead of string-compared. An
unrecognized generator value now fails fast with a clear ValueError instead
of silently falling through some bare-string checks and not others (same
behavior build_router/build_trunk already have for their own type keys).
Co-Authored-By: Claude Opus 5 <noreply@anthropic.com>
203 lines
7.4 KiB
Python
203 lines
7.4 KiB
Python
"""Generative objectives (flow/ddpm/wgan): `Objective` base + registry,
|
|
mirroring `giant.model.routers`'s `Router` pattern (gitea #32). Each objective
|
|
answers, in one place, the handful of questions every stage model/sampler/
|
|
trainer used to re-derive independently from a bare `generator` string: does
|
|
this stage need a time embedding, is it adversarial, does it fold the
|
|
secondary type slice into its own trunk output, what does the trunk take as
|
|
input, which stage-1/stage-2 loss does it train against.
|
|
|
|
Self-contained (no dependency on `giant.model.models`, unlike `Router` which
|
|
`giant.model.trunks` depends on) — `Objective` never needs to construct a
|
|
stage model or critic itself, only describe one. This also sidesteps a
|
|
`models.py` <-> `objectives.py` import cycle, since `models.py` calls
|
|
`build_objective`.
|
|
"""
|
|
|
|
import inspect
|
|
|
|
import torch
|
|
|
|
from giant.model.schedule import (
|
|
CosineSchedule,
|
|
flow_matching_loss,
|
|
flow_matching_loss_secondary,
|
|
flow_matching_loss_secondary_ar,
|
|
)
|
|
|
|
# ---------------------------------------------------------------------------
|
|
# Objective contract
|
|
# ---------------------------------------------------------------------------
|
|
|
|
|
|
class Objective:
|
|
"""Contract for a pluggable generative objective. Not an `nn.Module` —
|
|
unlike `Router`, no objective owns learnable parameters, so a plain
|
|
strategy object is the honest fit.
|
|
|
|
`needs_time`/`is_adversarial`/`folds_type_slice`/`supports_stage2_decoder`
|
|
are set by each concrete subclass (no defaults here — a new objective
|
|
should have to state all four, not silently inherit one that happens to
|
|
be wrong for it). See `FlowObjective`/`DdpmObjective`/`WganObjective`.
|
|
"""
|
|
|
|
needs_time: bool
|
|
is_adversarial: bool
|
|
folds_type_slice: bool
|
|
supports_stage2_decoder: bool = True
|
|
|
|
def trunk_in_dim(self, out_dim: int, noise_dim: int) -> int:
|
|
"""Width of the trunk's own input — `out_dim` (denoising/flow-matching
|
|
a same-shape vector) for every non-adversarial objective;
|
|
`WganObjective` overrides to `noise_dim` (a single-pass noise-to-output
|
|
generator)."""
|
|
return out_dim
|
|
|
|
def build_schedule(self, n_steps: int, device: torch.device) -> CosineSchedule | None:
|
|
"""Objective-owned auxiliary state a stage trainer must build once
|
|
and hold onto (device-placed) across its training loop. `None` for
|
|
every objective except `DdpmObjective` (its noise schedule)."""
|
|
return None
|
|
|
|
def stage1_loss(
|
|
self,
|
|
model: torch.nn.Module,
|
|
x1: torch.Tensor,
|
|
cond_cont: torch.Tensor,
|
|
cond_cat: torch.Tensor,
|
|
*,
|
|
schedule: object | None = None,
|
|
) -> torch.Tensor:
|
|
"""Stage-1 training loss. Only implemented by non-adversarial
|
|
objectives — `WganObjective` is unused here, `WGANStageTrainer` has
|
|
its own G/D step instead."""
|
|
raise NotImplementedError(f"{type(self).__name__} has no stage1_loss")
|
|
|
|
def stage2_loss(
|
|
self,
|
|
model: torch.nn.Module,
|
|
x1_s2: torch.Tensor,
|
|
cond_cont: torch.Tensor,
|
|
cond_cat: torch.Tensor,
|
|
stage1_ctx: torch.Tensor,
|
|
sec_mask: torch.Tensor,
|
|
*,
|
|
type_dim: int | None,
|
|
ar_inputs: dict[str, torch.Tensor] | None = None,
|
|
) -> torch.Tensor:
|
|
"""Stage-2 secondary-decoder training loss, one-shot or
|
|
autoregressive depending on whether `ar_inputs` is given. Same
|
|
adversarial caveat as `stage1_loss`."""
|
|
raise NotImplementedError(f"{type(self).__name__} has no stage2_loss")
|
|
|
|
|
|
OBJECTIVE_REGISTRY: dict[str, type[Objective]] = {}
|
|
|
|
|
|
def register_objective(name: str):
|
|
def decorator(cls: type[Objective]) -> type[Objective]:
|
|
OBJECTIVE_REGISTRY[name] = cls
|
|
return cls
|
|
|
|
return decorator
|
|
|
|
|
|
def build_objective(name: str, **kwargs) -> Objective:
|
|
"""Factory: look up an `Objective` subclass by name (a `generator`
|
|
config value) from the registry.
|
|
|
|
Every registered objective is fed the same kwargs; kwargs not declared by
|
|
that type's constructor are silently dropped, so per-type hyperparameters
|
|
(e.g. `DdpmObjective`'s `n_steps`) can coexist in one call without
|
|
special-casing — same convention as `giant.model.routers.build_router`.
|
|
"""
|
|
if name not in OBJECTIVE_REGISTRY:
|
|
raise ValueError(f"unknown generator/objective {name!r}; available: {sorted(OBJECTIVE_REGISTRY)}")
|
|
cls = OBJECTIVE_REGISTRY[name]
|
|
accepted = set(inspect.signature(cls.__init__).parameters) - {"self"}
|
|
filtered = {k: v for k, v in kwargs.items() if k in accepted}
|
|
return cls(**filtered)
|
|
|
|
|
|
# ---------------------------------------------------------------------------
|
|
# Concrete objectives
|
|
# ---------------------------------------------------------------------------
|
|
|
|
|
|
@register_objective("flow")
|
|
class FlowObjective(Objective):
|
|
"""Conditional flow matching (Lipman et al. 2022) — the primary
|
|
objective. ~10 ODE steps at inference (`giant.sample.sample_flow`)."""
|
|
|
|
needs_time = True
|
|
is_adversarial = False
|
|
folds_type_slice = False
|
|
|
|
def stage1_loss(self, model, x1, cond_cont, cond_cat, *, schedule=None) -> torch.Tensor:
|
|
return flow_matching_loss(model, x1, cond_cont, cond_cat)
|
|
|
|
def stage2_loss(
|
|
self,
|
|
model,
|
|
x1_s2,
|
|
cond_cont,
|
|
cond_cat,
|
|
stage1_ctx,
|
|
sec_mask,
|
|
*,
|
|
type_dim=None,
|
|
ar_inputs=None,
|
|
) -> torch.Tensor:
|
|
if ar_inputs is not None:
|
|
return flow_matching_loss_secondary_ar(
|
|
model,
|
|
x1_s2,
|
|
cond_cont,
|
|
cond_cat,
|
|
stage1_ctx,
|
|
ar_inputs["history_feat"],
|
|
ar_inputs["has_prev"],
|
|
ar_inputs["remaining_frac"],
|
|
ar_inputs["slot_idx"],
|
|
sec_mask,
|
|
type_dim=type_dim,
|
|
)
|
|
return flow_matching_loss_secondary(model, x1_s2, cond_cont, cond_cat, stage1_ctx, sec_mask, type_dim=type_dim)
|
|
|
|
|
|
@register_objective("ddpm")
|
|
class DdpmObjective(Objective):
|
|
"""Full DDPM ancestral sampling (Nichol & Dhariwal 2021 cosine schedule)
|
|
— the throwaway baseline. Stage-1 only: no `Stage2*` class has ever been
|
|
trained with `generator="ddpm"` in practice, so there's no stage-2 ddpm
|
|
loss to dispatch to (matches `FlowDDPMStageTrainer`'s pre-existing
|
|
stage-2 guard)."""
|
|
|
|
needs_time = True
|
|
is_adversarial = False
|
|
folds_type_slice = False
|
|
supports_stage2_decoder = False
|
|
|
|
def __init__(self, n_steps: int = 1000) -> None:
|
|
self.n_steps = n_steps
|
|
|
|
def build_schedule(self, n_steps: int, device: torch.device) -> CosineSchedule:
|
|
return CosineSchedule(T=n_steps).to(device)
|
|
|
|
def stage1_loss(self, model, x1, cond_cont, cond_cat, *, schedule=None) -> torch.Tensor:
|
|
assert schedule is not None, "DdpmObjective.stage1_loss needs a schedule (see build_schedule)"
|
|
return schedule.loss(model, x1, cond_cont, cond_cat)
|
|
|
|
|
|
@register_objective("wgan")
|
|
class WganObjective(Objective):
|
|
"""WGAN-GP (Gulrajani et al. 2017) — single forward pass instead of an
|
|
ODE loop. `stage1_loss`/`stage2_loss` are unused: `WGANStageTrainer` owns
|
|
its own dual generator/critic step instead of a single scalar loss."""
|
|
|
|
needs_time = False
|
|
is_adversarial = True
|
|
folds_type_slice = True
|
|
|
|
def trunk_in_dim(self, out_dim: int, noise_dim: int) -> int:
|
|
return noise_dim
|