Files
giant/tests/test_wgan.py
T
lars 9ce55e5013
CI / Format (ruff format) (push) Failing after 25s
CI / Lint (ruff check) (push) Successful in 27s
CI / Sync project version with tag (push) Has been skipped
CI / Type check (ty) (push) Failing after 24s
CI / Tests (push) Has been skipped
v0.3.0 step 2: network.py refactor to composable stage models
Decomposes the ten permutation classes in giant/model/network.py into
the reusable parts from docs/v0.3.0-design.md §5: ConditionEncoder (now
independently configurable per particle/material axis), ContextAdapter,
Trunk/MonolithicTrunk/RoutedTrunk/ExpertTrunk, and the stage classes
Stage1Model/Stage2OneShot/CriticModel (Stage2Autoregressive stubbed,
raises NotImplementedError until step 4/5). build_models/build_critics
now return a dict keyed by stage and accept the new nested config shape,
with routed WGAN reachable for the first time (the old --mode wgan
--router rejection is gone) and stage2_model.router.tie_to_stage1
sharing a literal Router instance.

A v0.2 checkpoint's flat model_config auto-migrates via
_migrate_legacy_model_config + migrate_legacy_state_dict, preserving the
n_sec_head's attachment to Stage1Model (legacy_owner="stage1", design
doc §4.1). tests/test_migration_v02_v03.py proves this bit-identical
against a frozen v0.2 snapshot (tests/legacy/network_v02_snapshot.py)
for both flow and wgan, both conditioning modes.
scripts/check_migration_v02_v03.py is the real-checkpoint counterpart
for a portal machine with /ceph access.

giant/model/schedule.py's flow-matching/DDPM loss helpers are updated
to the new model-call convention (t as a keyword). giant/sample.py,
giant/rollout.py, and giant/validate.py are not yet updated (deferred
to design doc step 6) — their exercising tests are marked xfail with
that reasoning rather than silently broken.

Co-Authored-By: Claude Sonnet 5 <noreply@anthropic.com>
2026-08-06 10:55:29 +02:00

217 lines
5.8 KiB
Python

import torch
from giant.constants import COND_DIM, K_MAX, SEC_DIM, SEC_SLOT_DIM, X_DIM
from giant.model.network import CriticModel, Stage1Model, Stage2OneShot
from giant.model.wgan import critic_loss, generator_loss, gradient_penalty
from giant.sample import sample_secondaries_wgan, sample_wgan
PARTICLE_CFG = {"type": "physical", "emb_dim": 8, "n_layers": 1}
MATERIAL_CFG = {"type": "physical", "emb_dim": 8, "n_layers": 1}
def _cond(B=8):
cond_cont = torch.randn(B, COND_DIM)
cond_cat = torch.zeros(B, 2, dtype=torch.long)
return cond_cont, cond_cat
def _small_generator():
return Stage1Model(
pdg_vocab=3,
mat_vocab=2,
particle_cfg=PARTICLE_CFG,
material_cfg=MATERIAL_CFG,
hidden_dim=32,
n_res_blocks=2,
generator="wgan",
noise_dim=8,
n_sec_head_k_max=K_MAX,
)
def _small_critic():
return CriticModel(
pdg_vocab=3,
mat_vocab=2,
particle_cfg=PARTICLE_CFG,
material_cfg=MATERIAL_CFG,
in_dim=X_DIM,
hidden_dim=32,
n_res_blocks=2,
stage="stage1",
)
def _small_sec_generator():
return Stage2OneShot(
pdg_vocab=3,
mat_vocab=2,
particle_cfg=PARTICLE_CFG,
material_cfg=MATERIAL_CFG,
hidden_dim=32,
n_res_blocks=2,
generator="wgan",
noise_dim=8,
)
def _small_sec_critic():
return CriticModel(
pdg_vocab=3,
mat_vocab=2,
particle_cfg=PARTICLE_CFG,
material_cfg=MATERIAL_CFG,
in_dim=SEC_DIM,
hidden_dim=32,
n_res_blocks=2,
stage="stage2",
)
def _mask(B, n_sec):
sec_mask = torch.arange(K_MAX).unsqueeze(0) < n_sec.unsqueeze(1)
return sec_mask.unsqueeze(-1).expand(-1, -1, SEC_SLOT_DIM).reshape(B, -1).float()
# --- Stage-1 generator/critic ---
def test_wgan_generator_output_shape():
B = 8
model = _small_generator()
cond_cont, cond_cat = _cond(B)
z = torch.randn(B, model.noise_dim)
out = model(z, cond_cont, cond_cat)
assert out.shape == (B, X_DIM)
def test_wgan_generator_predict_n_sec_shape():
B = 6
model = _small_generator()
cond_cont, cond_cat = _cond(B)
logits = model.predict_n_sec(cond_cont, cond_cat)
assert logits.shape == (B, K_MAX + 1)
def test_wgan_generator_gradients_flow():
B = 4
model = _small_generator()
cond_cont, cond_cat = _cond(B)
z = torch.randn(B, model.noise_dim)
gen_loss = model(z, cond_cont, cond_cat).sum()
nsec_loss = model.predict_n_sec(cond_cont, cond_cat).sum()
(gen_loss + nsec_loss).backward()
for name, p in model.named_parameters():
assert p.grad is not None, f"no grad for {name}"
def test_critic_output_shape():
B = 8
critic = _small_critic()
cond_cont, cond_cat = _cond(B)
x = torch.randn(B, X_DIM)
out = critic(x, cond_cont, cond_cat)
assert out.shape == (B,)
def test_sample_wgan_shape():
B = 6
model = _small_generator()
cond_cont, cond_cat = _cond(B)
sample, n_sec = sample_wgan(model, cond_cont, cond_cat)
assert sample.shape == (B, X_DIM)
assert n_sec.shape == (B,)
# --- Stage-2 generator/critic ---
def test_wgan_secondary_generator_output_shape():
B = 8
model = _small_sec_generator()
cond_cont, cond_cat = _cond(B)
stage1_out = torch.randn(B, X_DIM)
z = torch.randn(B, model.noise_dim)
out = model(z, cond_cont, cond_cat, stage1_out)
assert out.shape == (B, SEC_DIM)
def test_secondary_critic_output_shape():
B = 8
critic = _small_sec_critic()
cond_cont, cond_cat = _cond(B)
stage1_out = torch.randn(B, X_DIM)
x = torch.randn(B, SEC_DIM)
out = critic(x, cond_cont, cond_cat, stage1_out)
assert out.shape == (B,)
def test_sample_secondaries_wgan_shape():
B = 5
model = _small_sec_generator()
cond_cont, cond_cat = _cond(B)
stage1_out = torch.randn(B, X_DIM)
n_sec_pred = torch.randint(0, K_MAX, (B,))
sec_cont, sec_phys, sec_valid = sample_secondaries_wgan(
model, cond_cont, cond_cat, stage1_out, n_sec_pred
)
assert sec_cont.shape == (B, K_MAX, 4)
assert sec_phys.shape == (B, K_MAX, 2)
assert sec_valid.shape == (B, K_MAX)
# --- Losses ---
def test_gradient_penalty_nonneg():
B = 8
critic = _small_critic()
cond_cont, cond_cat = _cond(B)
real = torch.randn(B, X_DIM)
fake = torch.randn(B, X_DIM)
gp = gradient_penalty(lambda x: critic(x, cond_cont, cond_cat), real, fake)
assert gp.item() >= 0.0
assert gp.shape == ()
def test_gradient_penalty_masked():
B = 8
sec_critic = _small_sec_critic()
cond_cont, cond_cat = _cond(B)
stage1_out = torch.randn(B, X_DIM)
n_sec = torch.randint(0, K_MAX, (B,))
mask = _mask(B, n_sec)
real = torch.randn(B, SEC_DIM) * mask
fake = torch.randn(B, SEC_DIM) * mask
gp = gradient_penalty(
lambda x: sec_critic(x, cond_cont, cond_cat, stage1_out), real, fake, mask=mask
)
assert gp.item() >= 0.0
def test_critic_loss_scalar_and_grad():
B = 8
critic = _small_critic()
cond_cont, cond_cat = _cond(B)
real = torch.randn(B, X_DIM)
fake = torch.randn(B, X_DIM)
loss = critic_loss(
lambda x: critic(x, cond_cont, cond_cat), real, fake.detach(), gp_weight=10.0
)
assert loss.shape == ()
loss.backward()
assert any(p.grad is not None for p in critic.parameters())
def test_generator_loss_scalar_and_grad():
B = 4
generator = _small_generator()
critic = _small_critic()
cond_cont, cond_cat = _cond(B)
z = torch.randn(B, generator.noise_dim)
fake = generator(z, cond_cont, cond_cat)
loss = generator_loss(lambda x: critic(x, cond_cont, cond_cat), fake)
assert loss.shape == ()
loss.backward()
assert any(p.grad is not None for p in generator.parameters())