Compare commits
8 Commits
| Author | SHA1 | Date | |
|---|---|---|---|
| 97f5bbf9f0 | |||
| 060353ea4a | |||
| b3f28e98af | |||
| 4b2e0ba98e | |||
| 1675052ecd | |||
| eb9d331bea | |||
| 12689cf5b6 | |||
| 732d5f1cd2 |
+1
-1
@@ -1,5 +1,5 @@
|
||||
[tool.bumpversion]
|
||||
current_version = "0.3.4"
|
||||
current_version = "0.3.6"
|
||||
parse = "(?P<major>\\d+)\\.(?P<minor>\\d+)\\.(?P<patch>\\d+)"
|
||||
serialize = ["{major}.{minor}.{patch}"]
|
||||
search = "{current_version}"
|
||||
|
||||
@@ -1,5 +1,17 @@
|
||||
# Changelog
|
||||
|
||||
## [0.3.6] - 2026-08-24
|
||||
|
||||
### Changed
|
||||
|
||||
- Give CriticModel a registry-built trunk and StageModel base [gitea #57](https://git.larsbogner.de/lars/giant/issues/57)
|
||||
|
||||
## [0.3.5] - 2026-08-24
|
||||
|
||||
### Added
|
||||
|
||||
- Add "none" variants for router, history, and trunk [gitea #45](https://git.larsbogner.de/lars/giant/issues/45)
|
||||
|
||||
## [0.3.4] - 2026-08-23
|
||||
|
||||
### Added
|
||||
|
||||
@@ -202,6 +202,8 @@ def build_critics(model_config: dict) -> dict[str, nn.Module | None]:
|
||||
cond_out_dim=cond_out_dim,
|
||||
dropout=s1_spec.dropout,
|
||||
stage="stage1",
|
||||
trunk_type=s1_spec.trunk.type,
|
||||
block_conditioning=s1_spec.trunk.block_conditioning,
|
||||
)
|
||||
|
||||
if s2_spec.active and build_objective(s2_spec.generator).is_adversarial:
|
||||
@@ -225,6 +227,8 @@ def build_critics(model_config: dict) -> dict[str, nn.Module | None]:
|
||||
dropout=s2_spec.dropout,
|
||||
stage="stage2",
|
||||
context_dim=s2_spec.context_dim,
|
||||
trunk_type=s2_spec.trunk.type,
|
||||
block_conditioning=s2_spec.trunk.block_conditioning,
|
||||
)
|
||||
|
||||
return result
|
||||
|
||||
@@ -63,6 +63,23 @@ def build_history(name: str, in_dim: int, out_dim: int, **kwargs) -> HistoryEnco
|
||||
return cls(in_dim, out_dim, **filtered)
|
||||
|
||||
|
||||
@register_history("none")
|
||||
class NoHistory(HistoryEncoder):
|
||||
"""No history signal at all — ignores feat/has_prev entirely and always
|
||||
returns zeros. Ablates whether the AR decoder's history conditioning is
|
||||
earning its parameters. `init_cache`/`step` use the base class's O(1)
|
||||
defaults unmodified (this encoder's own `forward` is already O(1) per
|
||||
call regardless of prefix length)."""
|
||||
|
||||
def __init__(self, in_dim: int, out_dim: int) -> None:
|
||||
super().__init__()
|
||||
self.out_dim = out_dim
|
||||
|
||||
def forward(self, feat: torch.Tensor, has_prev: torch.Tensor) -> torch.Tensor:
|
||||
B, K, _ = feat.shape
|
||||
return torch.zeros(B, K, self.out_dim, device=feat.device, dtype=feat.dtype)
|
||||
|
||||
|
||||
@register_history("markov")
|
||||
class MarkovHistory(HistoryEncoder):
|
||||
"""Summarizes the previous secondary's own `(energy_fraction, direction,
|
||||
|
||||
+57
-32
@@ -8,7 +8,7 @@ from giant.config import ConditioningAxisConfig, HeadConfig, ParticleTypeConfig
|
||||
from giant.constants import CONT_SLOT_DIM, K_MAX, PARTICLE_PHYS_DIM, SEC_DIM, SEC_SLOT_DIM, X_DIM
|
||||
from giant.model.encoders import ConditionEncoder
|
||||
from giant.model.history import HistoryEncoder, build_history
|
||||
from giant.model.layers import ContextAdapter, ResBlock, SinusoidalEmbedding, build_mlp_head
|
||||
from giant.model.layers import ContextAdapter, SinusoidalEmbedding, build_mlp_head
|
||||
from giant.model.objectives import build_objective
|
||||
from giant.model.routers import Router
|
||||
from giant.model.trunks import build_trunk
|
||||
@@ -188,6 +188,26 @@ class StageModel(nn.Module):
|
||||
hidden = max(1, round(hidden_dim * head_cfg.hidden_ratio))
|
||||
self.stop_head = build_mlp_head(cond_out_dim, 1, hidden, head_cfg.depth)
|
||||
|
||||
def _build_context_fusion(self, x_dim: int, context_dim: int, cond_out_dim: int) -> None:
|
||||
"""Builds `self.context_adapter`/`self.fuse` — the stage-2-style
|
||||
context-fusion pattern (project the previous stage's outcome down to
|
||||
`context_dim` via `ContextAdapter`, concat onto the base conditioning,
|
||||
project back to `cond_out_dim`) shared by `Stage2OneShot` and a
|
||||
`stage="stage2"` `CriticModel` (gitea #57). Call from a subclass's
|
||||
`__init__` before using `_cond_embed`."""
|
||||
self.context_adapter = ContextAdapter(x_dim, context_dim)
|
||||
self.fuse = nn.Sequential(
|
||||
nn.Linear(cond_out_dim + context_dim, cond_out_dim),
|
||||
nn.SiLU(),
|
||||
)
|
||||
|
||||
def _cond_embed(self, cond_cont: torch.Tensor, cond_cat: torch.Tensor, stage1_out: torch.Tensor) -> torch.Tensor:
|
||||
"""Fuses base conditioning with the previous stage's outcome — pairs
|
||||
with `_build_context_fusion`."""
|
||||
base = self.cond_enc(cond_cont, cond_cat)
|
||||
ctx = self.context_adapter(stage1_out)
|
||||
return self.fuse(torch.cat([base, ctx], dim=-1))
|
||||
|
||||
def _require_n_sec_head(self) -> None:
|
||||
if self.n_sec_head is None:
|
||||
raise RuntimeError(
|
||||
@@ -363,11 +383,7 @@ class Stage2OneShot(StageModel):
|
||||
particle_type_cfg=particle_type_cfg,
|
||||
cond_enc=cond_enc,
|
||||
)
|
||||
self.context_adapter = ContextAdapter(x_dim, context_dim)
|
||||
self.fuse = nn.Sequential(
|
||||
nn.Linear(cond_out_dim + context_dim, cond_out_dim),
|
||||
nn.SiLU(),
|
||||
)
|
||||
self._build_context_fusion(x_dim, context_dim, cond_out_dim)
|
||||
target = self.particle_type_cfg.target
|
||||
type_head_out_dim = None if target == "physical" else k_max * self.type_dim
|
||||
self._build_trunk_and_heads(
|
||||
@@ -386,11 +402,6 @@ class Stage2OneShot(StageModel):
|
||||
type_head_cfg=type_head_cfg,
|
||||
)
|
||||
|
||||
def _cond_embed(self, cond_cont: torch.Tensor, cond_cat: torch.Tensor, stage1_out: torch.Tensor) -> torch.Tensor:
|
||||
base = self.cond_enc(cond_cont, cond_cat)
|
||||
ctx = self.context_adapter(stage1_out)
|
||||
return self.fuse(torch.cat([base, ctx], dim=-1))
|
||||
|
||||
def forward(
|
||||
self,
|
||||
x_t: torch.Tensor,
|
||||
@@ -702,11 +713,25 @@ class Stage2Autoregressive(StageModel):
|
||||
return self.stop_head(c_emb.reshape(B * K, -1)).view(B, K)
|
||||
|
||||
|
||||
class CriticModel(nn.Module):
|
||||
class CriticModel(StageModel):
|
||||
"""Generator-agnostic WGAN-GP critic body: a scalar realism score, for
|
||||
either stage (`stage="stage1"` mirrors v0.2 `Critic`; `stage="stage2"`
|
||||
mirrors v0.2 `SecondaryCritic`, adding the same context-fusion path as
|
||||
`Stage2OneShot`). Used only when that stage's `generator == "wgan"`."""
|
||||
`Stage2OneShot`, via `StageModel._build_context_fusion`/`_cond_embed`).
|
||||
Used only when that stage's `generator == "wgan"`.
|
||||
|
||||
Subclasses `StageModel` for the `cond_enc` construction and (stage 2)
|
||||
context-fusion scaffolding only — its trunk is built directly via
|
||||
`build_trunk` (output width 1) rather than through
|
||||
`_build_trunk_and_heads`, since that helper is shaped around a
|
||||
generator's `Objective`/time-embedding/flow-matching concerns
|
||||
(`forward`'s `(x_t, cond) -> vector` shape) that don't apply to a critic's
|
||||
`(x, cond) -> scalar` (gitea #57). `generator="wgan"` is passed to the
|
||||
base purely because that's factually when a critic exists; nothing here
|
||||
ever calls `_build_trunk_and_heads`, so no head/time-embedding machinery
|
||||
is built from it. Never routed (MoE) — that's a separate, unrequested
|
||||
axis of scope; see gitea #57's proposal, which covers only the trunk/
|
||||
block registries."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
@@ -722,22 +747,26 @@ class CriticModel(nn.Module):
|
||||
stage: str = "stage1",
|
||||
context_dim: int = 64,
|
||||
context_in_dim: int = X_DIM,
|
||||
trunk_type: str = "resmlp",
|
||||
block_conditioning: str = "add",
|
||||
) -> None:
|
||||
super().__init__()
|
||||
super().__init__(
|
||||
pdg_vocab,
|
||||
mat_vocab,
|
||||
particle_cfg,
|
||||
material_cfg,
|
||||
cond_out_dim=cond_out_dim,
|
||||
generator="wgan",
|
||||
noise_dim=0,
|
||||
)
|
||||
if stage not in ("stage1", "stage2"):
|
||||
raise ValueError(f"stage must be 'stage1' or 'stage2', got {stage!r}")
|
||||
self.stage = stage
|
||||
self.cond_enc = ConditionEncoder(pdg_vocab, mat_vocab, particle_cfg, material_cfg, out_dim=cond_out_dim)
|
||||
if stage == "stage2":
|
||||
self.context_adapter = ContextAdapter(context_in_dim, context_dim)
|
||||
self.fuse = nn.Sequential(
|
||||
nn.Linear(cond_out_dim + context_dim, cond_out_dim),
|
||||
nn.SiLU(),
|
||||
)
|
||||
self.input_proj = nn.Linear(in_dim, hidden_dim)
|
||||
self.blocks = nn.ModuleList([ResBlock(hidden_dim, cond_out_dim, dropout=dropout) for _ in range(n_res_blocks)])
|
||||
self.out_norm = nn.LayerNorm(hidden_dim)
|
||||
self.out_proj = nn.Linear(hidden_dim, 1)
|
||||
self._build_context_fusion(context_in_dim, context_dim, cond_out_dim)
|
||||
self.trunk = build_trunk(
|
||||
None, trunk_type, in_dim, 1, hidden_dim, n_res_blocks, cond_out_dim, dropout, block_conditioning
|
||||
)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
@@ -746,13 +775,9 @@ class CriticModel(nn.Module):
|
||||
cond_cat: torch.Tensor,
|
||||
stage1_out: torch.Tensor | None = None,
|
||||
) -> torch.Tensor:
|
||||
base = self.cond_enc(cond_cont, cond_cat)
|
||||
if self.stage == "stage2":
|
||||
ctx = self.context_adapter(stage1_out)
|
||||
cond = self.fuse(torch.cat([base, ctx], dim=-1))
|
||||
assert stage1_out is not None, "stage='stage2' CriticModel requires stage1_out"
|
||||
cond = self._cond_embed(cond_cont, cond_cat, stage1_out)
|
||||
else:
|
||||
cond = base
|
||||
h = self.input_proj(x)
|
||||
for block in self.blocks:
|
||||
h = block(h, cond)
|
||||
return self.out_proj(self.out_norm(h)).squeeze(-1)
|
||||
cond = self.cond_enc(cond_cont, cond_cat)
|
||||
return self.trunk(x, cond, cond_cont, cond_cat).squeeze(-1)
|
||||
|
||||
@@ -15,6 +15,7 @@ from giant.model.history import (
|
||||
AttentionHistory,
|
||||
HistoryEncoder,
|
||||
MarkovHistory,
|
||||
NoHistory,
|
||||
_CausalAttnBlock,
|
||||
build_history,
|
||||
register_history,
|
||||
@@ -54,6 +55,7 @@ from giant.model.routers import (
|
||||
ROUTER_REGISTRY,
|
||||
ComposedRouter,
|
||||
EnergyRouter,
|
||||
NoneRouter,
|
||||
PdgRouter,
|
||||
ProcessRouter,
|
||||
Router,
|
||||
@@ -67,6 +69,7 @@ from giant.model.routers import (
|
||||
from giant.model.trunks import (
|
||||
TRUNK_REGISTRY,
|
||||
ExpertTrunk,
|
||||
LinearTrunk,
|
||||
RoutedTrunk,
|
||||
Trunk,
|
||||
_route_forward,
|
||||
@@ -90,7 +93,10 @@ __all__ = [
|
||||
"FlowObjective",
|
||||
"HISTORY_REGISTRY",
|
||||
"HistoryEncoder",
|
||||
"LinearTrunk",
|
||||
"MarkovHistory",
|
||||
"NoHistory",
|
||||
"NoneRouter",
|
||||
"OBJECTIVE_REGISTRY",
|
||||
"Objective",
|
||||
"PdgRouter",
|
||||
|
||||
@@ -144,6 +144,24 @@ def _inverse_bounded_interp(value: float, lo: float, hi: float) -> float:
|
||||
return math.log(p / (1 - p))
|
||||
|
||||
|
||||
@register_router("none")
|
||||
class NoneRouter(Router):
|
||||
"""Uniform 1/n_experts gate — no learned routing signal at all.
|
||||
|
||||
Still builds n_experts expert trunks via RoutedTrunk (same parameter
|
||||
budget as a real router), but every row gets an identical weight
|
||||
regardless of conditioning. Ablates whether the *learned routing
|
||||
signal* — as opposed to simply having multiple experts — is earning
|
||||
its parameters. `top1()` (the base class default) always dispatches to
|
||||
expert 0 (argmax of a uniform vector), which still exercises
|
||||
RoutedTrunk's real per-expert grouped-dispatch code path at eval time.
|
||||
"""
|
||||
|
||||
def gate(self, cond_cont: torch.Tensor, cond_cat: torch.Tensor) -> torch.Tensor:
|
||||
B = cond_cont.shape[0]
|
||||
return torch.full((B, self.n_experts), 1.0 / self.n_experts, device=cond_cont.device)
|
||||
|
||||
|
||||
@register_router("energy")
|
||||
class EnergyRouter(Router):
|
||||
"""Soft turn-on gate over normalized pre-step log-energy.
|
||||
|
||||
@@ -94,6 +94,49 @@ class ExpertTrunk(nn.Module):
|
||||
return self.out_proj(x)
|
||||
|
||||
|
||||
@register_trunk("linear")
|
||||
class LinearTrunk(nn.Module):
|
||||
"""`nn.Linear(in_dim + cond_dim, out_dim)` over `concat([x, cond])` —
|
||||
the trivial trunk body: no hidden layer, no ResBlock stack, no
|
||||
nonlinearity. Ablates whether trunk depth/nonlinearity is earning its
|
||||
parameters, holding everything else (heads, ConditionEncoder,
|
||||
generator, ...) fixed. Composes for free with `router.enabled = true`
|
||||
(gitea #33): a RoutedTrunk of n_experts linear bodies is "mixture of
|
||||
trivial linear experts". `hidden_dim`/`n_blocks`/`dropout`/
|
||||
`block_conditioning` are accepted and ignored, matching
|
||||
`build_expert_body`'s shared factory signature.
|
||||
|
||||
`x` — the trunk's own input (e.g. the noised primary vector for flow
|
||||
matching) — does not already carry conditioning; that's fused in
|
||||
per-body via `cond`. So this concatenates `x` and `cond` itself to
|
||||
remain a valid, conditioning-dependent model.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
in_dim: int,
|
||||
out_dim: int,
|
||||
hidden_dim: int,
|
||||
n_blocks: int,
|
||||
cond_dim: int,
|
||||
dropout: float = 0.0,
|
||||
block_conditioning: str = "add",
|
||||
) -> None:
|
||||
super().__init__()
|
||||
self.in_dim = in_dim
|
||||
self.out_dim = out_dim
|
||||
self.linear = nn.Linear(in_dim + cond_dim, out_dim)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
x: torch.Tensor,
|
||||
cond: torch.Tensor,
|
||||
cond_cont: torch.Tensor | None = None,
|
||||
cond_cat: torch.Tensor | None = None,
|
||||
) -> torch.Tensor:
|
||||
return self.linear(torch.cat([x, cond], dim=-1))
|
||||
|
||||
|
||||
def _route_forward(
|
||||
experts: nn.ModuleList,
|
||||
router: Router,
|
||||
|
||||
+1
-1
@@ -1,6 +1,6 @@
|
||||
[project]
|
||||
name = "giant"
|
||||
version = "0.3.4"
|
||||
version = "0.3.6"
|
||||
description = "Geant4 step-function surrogate via conditional flow matching"
|
||||
readme = "README.md"
|
||||
requires-python = ">=3.12"
|
||||
|
||||
+85
-14
@@ -8,8 +8,12 @@ from giant.model.network import (
|
||||
HISTORY_REGISTRY,
|
||||
AttentionHistory,
|
||||
ConditionEncoder,
|
||||
CriticModel,
|
||||
FilmResBlock,
|
||||
HistoryEncoder,
|
||||
LinearTrunk,
|
||||
MarkovHistory,
|
||||
NoHistory,
|
||||
SinusoidalEmbedding,
|
||||
Stage1Model,
|
||||
Stage2Autoregressive,
|
||||
@@ -465,16 +469,52 @@ def test_attention_history_step_matches_forward():
|
||||
assert torch.allclose(stepped, expected, atol=1e-5)
|
||||
|
||||
|
||||
# --- NoHistory (gitea #45) ----------------------------------------------------
|
||||
|
||||
|
||||
def test_no_history_shape():
|
||||
hist = NoHistory(in_dim=7, out_dim=12)
|
||||
B, K = 3, 5
|
||||
feat = torch.randn(B, K, 7)
|
||||
has_prev = (torch.arange(K) >= 1).unsqueeze(0).expand(B, -1)
|
||||
out = hist(feat, has_prev)
|
||||
assert out.shape == (B, K, 12)
|
||||
|
||||
|
||||
def test_no_history_ignores_feat_and_has_prev():
|
||||
hist = NoHistory(in_dim=4, out_dim=6)
|
||||
B, K = 2, 3
|
||||
has_prev_a = (torch.arange(K) >= 1).unsqueeze(0).expand(B, -1)
|
||||
has_prev_b = torch.zeros(B, K, dtype=torch.bool)
|
||||
feat_a = torch.randn(B, K, 4)
|
||||
feat_b = torch.randn(B, K, 4) * 100
|
||||
out_a = hist(feat_a, has_prev_a)
|
||||
out_b = hist(feat_b, has_prev_b)
|
||||
assert torch.equal(out_a, torch.zeros(B, K, 6))
|
||||
assert torch.equal(out_a, out_b)
|
||||
|
||||
|
||||
def test_no_history_uses_base_class_o1_defaults():
|
||||
hist = NoHistory(in_dim=4, out_dim=6)
|
||||
assert hist.init_cache() is None
|
||||
feat = torch.randn(2, 1, 4)
|
||||
has_prev = torch.ones(2, 1, dtype=torch.bool)
|
||||
out, cache = hist.step(feat, has_prev, "unused-cache")
|
||||
assert torch.equal(out, torch.zeros(2, 1, 6))
|
||||
assert cache == "unused-cache"
|
||||
|
||||
|
||||
# --- HISTORY_REGISTRY / build_history (gitea #35) ----------------------------
|
||||
|
||||
|
||||
def test_history_registry_has_exactly_the_two_known_histories():
|
||||
assert set(HISTORY_REGISTRY) == {"markov", "attention"}
|
||||
def test_history_registry_has_exactly_the_known_histories():
|
||||
assert set(HISTORY_REGISTRY) == {"markov", "attention", "none"}
|
||||
|
||||
|
||||
def test_build_history_returns_correct_concrete_type():
|
||||
assert isinstance(build_history("markov", 4, 6), MarkovHistory)
|
||||
assert isinstance(build_history("attention", 4, 8), AttentionHistory)
|
||||
assert isinstance(build_history("none", 4, 6), NoHistory)
|
||||
|
||||
|
||||
def test_build_history_unknown_name_raises():
|
||||
@@ -605,7 +645,7 @@ def test_stage2_autoregressive_n_sec_head_and_type_head_cfg_control_hidden_width
|
||||
|
||||
@pytest.mark.parametrize("target", ["physical", "onehot", "embedding"])
|
||||
@pytest.mark.parametrize("generator", ["wgan", "flow"])
|
||||
@pytest.mark.parametrize("history", ["markov", "attention"])
|
||||
@pytest.mark.parametrize("history", ["markov", "attention", "none"])
|
||||
def test_stage2_autoregressive_forward_shape(target, generator, history):
|
||||
B, K, emb_dim = 4, 5, 6
|
||||
model = _build_stage2_ar(target, generator, emb_dim=emb_dim, k_max=K, history=history)
|
||||
@@ -907,7 +947,7 @@ def test_build_critics_particle_type_n_classes_overrides_conditioning_emb_dim():
|
||||
assert wider_critic is not None
|
||||
# k_max=3 slots, each CONT_SLOT_DIM + n_classes wide under wgan folding —
|
||||
# widening n_classes alone (emb_dim stays 4) must widen the critic input.
|
||||
assert wider_critic.input_proj.in_features > default_n_classes_critic.input_proj.in_features
|
||||
assert wider_critic.trunk.input_proj.in_features > default_n_classes_critic.trunk.input_proj.in_features
|
||||
|
||||
|
||||
# ── build_models/build_critics: DEFAULT_CONFIG fallback drift (issues.md #1) ─
|
||||
@@ -984,12 +1024,12 @@ def test_build_critics_omitted_particle_type_matches_default_config():
|
||||
cfg["stage2_model"]["generator"] = "wgan"
|
||||
onehot_critic = build_critics(cfg)["stage2"]
|
||||
assert onehot_critic is not None
|
||||
onehot_in_dim = onehot_critic.input_proj.in_features
|
||||
onehot_in_dim = onehot_critic.trunk.input_proj.in_features
|
||||
|
||||
cfg["stage2_model"]["particle_type"] = {"target": "physical"}
|
||||
physical_critic = build_critics(cfg)["stage2"]
|
||||
assert physical_critic is not None
|
||||
physical_in_dim = physical_critic.input_proj.in_features
|
||||
physical_in_dim = physical_critic.trunk.input_proj.in_features
|
||||
|
||||
# onehot's per-slot type width is emb_dim classes vs. physical's fixed
|
||||
# (log-mass, charge) pair — different unless emb_dim happens to be 2, so
|
||||
@@ -1009,15 +1049,15 @@ def test_build_critics_stage1_critic_hidden_dim_and_n_res_blocks_override_genera
|
||||
|
||||
inherited = build_critics(cfg)["stage1"]
|
||||
assert inherited is not None
|
||||
assert inherited.input_proj.out_features == 8
|
||||
assert len(inherited.blocks) == 1
|
||||
assert inherited.trunk.input_proj.out_features == 8
|
||||
assert len(inherited.trunk.blocks) == 1
|
||||
|
||||
cfg["stage1_model"]["wgan"]["critic_hidden_dim"] = 16
|
||||
cfg["stage1_model"]["wgan"]["critic_n_res_blocks"] = 3
|
||||
overridden = build_critics(cfg)["stage1"]
|
||||
assert overridden is not None
|
||||
assert overridden.input_proj.out_features == 16
|
||||
assert len(overridden.blocks) == 3
|
||||
assert overridden.trunk.input_proj.out_features == 16
|
||||
assert len(overridden.trunk.blocks) == 3
|
||||
|
||||
|
||||
def test_build_critics_stage2_critic_hidden_dim_and_n_res_blocks_override_generator_size():
|
||||
@@ -1028,15 +1068,15 @@ def test_build_critics_stage2_critic_hidden_dim_and_n_res_blocks_override_genera
|
||||
|
||||
inherited = build_critics(cfg)["stage2"]
|
||||
assert inherited is not None
|
||||
assert inherited.input_proj.out_features == 8
|
||||
assert len(inherited.blocks) == 1
|
||||
assert inherited.trunk.input_proj.out_features == 8
|
||||
assert len(inherited.trunk.blocks) == 1
|
||||
|
||||
cfg["stage2_model"]["wgan"]["critic_hidden_dim"] = 16
|
||||
cfg["stage2_model"]["wgan"]["critic_n_res_blocks"] = 3
|
||||
overridden = build_critics(cfg)["stage2"]
|
||||
assert overridden is not None
|
||||
assert overridden.input_proj.out_features == 16
|
||||
assert len(overridden.blocks) == 3
|
||||
assert overridden.trunk.input_proj.out_features == 16
|
||||
assert len(overridden.trunk.blocks) == 3
|
||||
|
||||
|
||||
# ── StageModel base (gitea #39): Stage1Model/Stage2OneShot/Stage2Autoregressive
|
||||
@@ -1197,3 +1237,34 @@ def test_stagemodel_time_emb_matches_objective_needs_time(cls, generator):
|
||||
assert model.generator_kind == generator
|
||||
assert model.noise_dim == 8
|
||||
assert (model.time_emb is not None) == build_objective(generator).needs_time
|
||||
|
||||
|
||||
# ── CriticModel uses the trunk/block registries + StageModel base (gitea #57) ─
|
||||
|
||||
|
||||
def test_critic_model_is_stagemodel_subclass():
|
||||
assert issubclass(CriticModel, StageModel)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("stage", ["stage1", "stage2"])
|
||||
def test_build_critics_threads_trunk_type_from_generator_config(stage):
|
||||
cfg = _minimal_model_config(share_stages=False)
|
||||
cfg["stage1_model"]["generator"] = "wgan"
|
||||
cfg["stage2_model"]["generator"] = "wgan"
|
||||
cfg[f"{stage}_model"]["trunk"] = {"type": "linear"}
|
||||
|
||||
critic = build_critics(cfg)[stage]
|
||||
assert critic is not None
|
||||
assert isinstance(critic.trunk, LinearTrunk)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("stage", ["stage1", "stage2"])
|
||||
def test_build_critics_threads_block_conditioning_from_generator_config(stage):
|
||||
cfg = _minimal_model_config(share_stages=False)
|
||||
cfg["stage1_model"]["generator"] = "wgan"
|
||||
cfg["stage2_model"]["generator"] = "wgan"
|
||||
cfg[f"{stage}_model"]["trunk"] = {"block_conditioning": "film"}
|
||||
|
||||
critic = build_critics(cfg)[stage]
|
||||
assert critic is not None
|
||||
assert all(isinstance(block, FilmResBlock) for block in critic.trunk.blocks)
|
||||
|
||||
@@ -13,6 +13,8 @@ from giant.model.network import (
|
||||
EnergyRouter,
|
||||
ExpertTrunk,
|
||||
FilmResBlock,
|
||||
LinearTrunk,
|
||||
NoneRouter,
|
||||
PdgRouter,
|
||||
ProcessRouter,
|
||||
ROUTER_REGISTRY,
|
||||
@@ -73,6 +75,34 @@ def test_energy_router_registered():
|
||||
assert ROUTER_REGISTRY["energy"] is EnergyRouter
|
||||
|
||||
|
||||
# ── NoneRouter (gitea #45) ──────────────────────────────────────────────────
|
||||
|
||||
|
||||
def test_none_router_registered():
|
||||
assert ROUTER_REGISTRY["none"] is NoneRouter
|
||||
|
||||
|
||||
def test_none_router_gate_is_uniform():
|
||||
router = NoneRouter(n_experts=4)
|
||||
cond_cont, cond_cat = _cond(16)
|
||||
g = router.gate(cond_cont, cond_cat)
|
||||
assert g.shape == (16, 4)
|
||||
torch.testing.assert_close(g, torch.full((16, 4), 0.25))
|
||||
|
||||
|
||||
def test_none_router_gate_ignores_conditioning():
|
||||
router = NoneRouter(n_experts=3)
|
||||
cond_cont_a, cond_cat_a = _cond(8)
|
||||
cond_cont_b, cond_cat_b = _cond(8)
|
||||
torch.testing.assert_close(router.gate(cond_cont_a, cond_cat_a), router.gate(cond_cont_b, cond_cat_b))
|
||||
|
||||
|
||||
def test_none_router_top1_always_expert_zero():
|
||||
router = NoneRouter(n_experts=4)
|
||||
cond_cont, cond_cat = _cond(16)
|
||||
assert torch.equal(router.top1(cond_cont, cond_cat), torch.zeros(16, dtype=torch.long))
|
||||
|
||||
|
||||
def test_energy_router_gate_partition_of_unity():
|
||||
router = EnergyRouter(n_experts=4)
|
||||
cond_cont, cond_cat = _cond(16)
|
||||
@@ -168,6 +198,47 @@ def test_build_expert_body_unknown_type_raises():
|
||||
raise AssertionError("expected ValueError for unknown trunk type")
|
||||
|
||||
|
||||
# ── LinearTrunk (gitea #45) ──────────────────────────────────────────────────
|
||||
|
||||
|
||||
def test_trunk_registry_has_linear():
|
||||
assert "linear" in TRUNK_REGISTRY
|
||||
assert TRUNK_REGISTRY["linear"] is LinearTrunk
|
||||
|
||||
|
||||
def test_linear_trunk_forward_shape():
|
||||
trunk = build_expert_body("linear", in_dim=9, out_dim=9, hidden_dim=64, n_blocks=6, cond_dim=12)
|
||||
assert trunk.in_dim == 9
|
||||
assert trunk.out_dim == 9
|
||||
x = torch.randn(5, 9)
|
||||
cond = torch.randn(5, 12)
|
||||
out = trunk(x, cond)
|
||||
assert out.shape == (5, 9)
|
||||
|
||||
|
||||
def test_linear_trunk_depends_on_x_and_cond():
|
||||
trunk = build_expert_body("linear", in_dim=9, out_dim=9, hidden_dim=64, n_blocks=6, cond_dim=12)
|
||||
x = torch.randn(5, 9)
|
||||
cond_a = torch.randn(5, 12)
|
||||
cond_b = torch.randn(5, 12)
|
||||
assert not torch.allclose(trunk(x, cond_a), trunk(x, cond_b))
|
||||
|
||||
|
||||
def test_routed_linear_trunk_is_mixture_of_trivial_experts():
|
||||
"""trunk.type = 'linear' composes for free with router.enabled = true
|
||||
(gitea #33's comment on this issue) — a RoutedTrunk of n_experts linear
|
||||
bodies."""
|
||||
router = build_router("energy", n_experts=3)
|
||||
trunk = RoutedTrunk(router, "linear", in_dim=9, out_dim=9, hidden_dim=64, n_res_blocks=6, cond_dim=12)
|
||||
assert len(trunk.experts) == 3
|
||||
assert all(isinstance(e, LinearTrunk) for e in trunk.experts)
|
||||
x = torch.randn(5, 9)
|
||||
cond = torch.randn(5, 12)
|
||||
cond_cont, cond_cat = _cond(5)
|
||||
out = trunk(x, cond, cond_cont, cond_cat)
|
||||
assert out.shape == (5, 9)
|
||||
|
||||
|
||||
# ── BLOCK_REGISTRY / build_block (gitea #34) ────────────────────────────────
|
||||
|
||||
|
||||
|
||||
+28
-1
@@ -2,7 +2,7 @@ import torch
|
||||
|
||||
from giant.config import ConditioningAxisConfig
|
||||
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.network import CriticModel, LinearTrunk, Stage1Model, Stage2OneShot
|
||||
from giant.model.wgan import critic_loss, generator_loss, gradient_penalty
|
||||
from giant.sample import sample_secondaries_wgan, sample_wgan
|
||||
|
||||
@@ -115,6 +115,33 @@ def test_critic_output_shape():
|
||||
assert out.shape == (B,)
|
||||
|
||||
|
||||
def test_critic_model_honours_trunk_type_and_block_conditioning():
|
||||
"""gitea #57: CriticModel routes its body through build_trunk/build_block
|
||||
like every generator stage model, instead of hand-rolling a plain
|
||||
ResBlock stack."""
|
||||
B = 8
|
||||
critic = 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",
|
||||
trunk_type="linear",
|
||||
block_conditioning="adaln",
|
||||
)
|
||||
assert isinstance(critic.trunk, LinearTrunk)
|
||||
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)
|
||||
loss.backward()
|
||||
for name, p in critic.named_parameters():
|
||||
assert p.grad is not None, f"no grad for {name}"
|
||||
|
||||
|
||||
def test_sample_wgan_shape():
|
||||
B = 6
|
||||
model = _small_generator()
|
||||
|
||||
Reference in New Issue
Block a user