0f95e0eaae
CI / Format (ruff format) (push) Successful in 29s
CI / Lint (ruff check) (push) Successful in 33s
CI / Sync project version with tag (push) Has been skipped
CI / Type check (ty) (push) Successful in 32s
CI / Tests (push) Successful in 1m58s
CI / Format (ruff format) (pull_request) Successful in 28s
CI / Lint (ruff check) (pull_request) Successful in 31s
CI / Sync project version with tag (pull_request) Has been skipped
CI / Type check (ty) (pull_request) Successful in 27s
CI / Tests (pull_request) Successful in 2m5s
ResBlock injected conditioning exactly one way — h = linear1(h) + cond_proj(cond), a conditional bias, the weakest standard option for a model whose entire job is to be conditional. Adds BLOCK_REGISTRY (giant/model/layers.py), mirroring the TRUNK_REGISTRY/ROUTER_REGISTRY registry+factory idiom (gitea #33), with two new drop-in alternatives: FilmResBlock (per-channel scale+shift modulating the norm output, zero-init so conditioning has no effect at construction) and AdaLNResBlock (DiT-style AdaLN-Zero — the norm's own affine is replaced by a conditioning-derived scale/shift, plus a zero-init gate on the residual branch, making the block the exact identity function at init). Selected per stage via a new stage{1,2}_model.trunk.block_conditioning config leaf ("add" | "film" | "adaln", default "add"), threaded through build_trunk/build_expert_body/RoutedTrunk and the three stage model constructors. Default stays "add" and ResBlock's body is unchanged, so existing configs/checkpoints are bit-identical to before this change. Decided during planning: the new field lives on the existing TrunkConfig rather than a new top-level block/blocks config section; the WGAN CriticModel (which builds its own ResBlock stack outside TRUNK_REGISTRY) and the issue's mentioned blocks.norm/blocks.activation axes are both left out of scope. Co-Authored-By: Claude Opus 5 <noreply@anthropic.com>
160 lines
6.1 KiB
Python
160 lines
6.1 KiB
Python
"""Small stateless-ish building blocks shared across encoders/trunks/models —
|
|
no dependency on any other `giant.model` submodule (issues.md Issue 8)."""
|
|
|
|
import math
|
|
|
|
import torch
|
|
import torch.nn as nn
|
|
|
|
|
|
class SinusoidalEmbedding(nn.Module):
|
|
def __init__(self, dim: int) -> None:
|
|
super().__init__()
|
|
assert dim % 2 == 0, "dim must be even"
|
|
half = dim // 2
|
|
freqs = torch.exp(-math.log(10000) * torch.arange(half, dtype=torch.float32) / max(half - 1, 1))
|
|
self.register_buffer("freqs", freqs)
|
|
|
|
def forward(self, t: torch.Tensor) -> torch.Tensor:
|
|
t = t.reshape(-1, 1).float()
|
|
args = t * self.freqs.unsqueeze(0) # (B, half)
|
|
return torch.cat([args.sin(), args.cos()], dim=-1) # (B, dim)
|
|
|
|
|
|
def _make_axis_mlp(in_dim: int, emb_dim: int, n_layers: int) -> nn.Sequential:
|
|
"""`n_layers`-deep MLP producing an `emb_dim`-wide vector from `in_dim`
|
|
physical properties (`conditioning.{particle,material}.n_layers`).
|
|
|
|
`n_layers=1` (the v0.3.0 default): a single `Linear`, no hidden
|
|
activation. `n_layers=2` reproduces v0.2's hardcoded depth exactly —
|
|
`Linear -> SiLU -> Linear` — which is why `migrate_config` back-fills
|
|
`n_layers=2` for migrated configs rather than the v0.3 default of 1 (see
|
|
its docstring).
|
|
"""
|
|
if n_layers < 1:
|
|
raise ValueError(f"n_layers must be >= 1, got {n_layers}")
|
|
if n_layers == 1:
|
|
return nn.Sequential(nn.Linear(in_dim, emb_dim))
|
|
layers: list[nn.Module] = [nn.Linear(in_dim, emb_dim), nn.SiLU()]
|
|
for _ in range(n_layers - 2):
|
|
layers += [nn.Linear(emb_dim, emb_dim), nn.SiLU()]
|
|
layers.append(nn.Linear(emb_dim, emb_dim))
|
|
return nn.Sequential(*layers)
|
|
|
|
|
|
class ContextAdapter(nn.Module):
|
|
"""Projects a stage's outcome (e.g. Stage 1's 9D target) down to a
|
|
fixed-width context vector for a downstream stage's conditioning —
|
|
`stage2_model.context_dim`. Was `SecondaryConditionEncoder.stage1_proj`
|
|
(+ its `tanh`) in v0.2; pulled out as its own module in v0.3.0 since
|
|
`SecondaryConditionEncoder` as a wrapper class disappears."""
|
|
|
|
def __init__(self, in_dim: int, context_dim: int) -> None:
|
|
super().__init__()
|
|
self.proj = nn.Linear(in_dim, context_dim)
|
|
|
|
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
|
return torch.tanh(self.proj(x))
|
|
|
|
|
|
BLOCK_REGISTRY: dict[str, type[nn.Module]] = {}
|
|
|
|
|
|
def register_block(name: str):
|
|
def decorator(cls: type[nn.Module]) -> type[nn.Module]:
|
|
BLOCK_REGISTRY[name] = cls
|
|
return cls
|
|
|
|
return decorator
|
|
|
|
|
|
def build_block(name: str, dim: int, cond_dim: int, dropout: float = 0.0) -> nn.Module:
|
|
"""Factory: look up a registered conditioning-injection block by name and
|
|
construct one instance — `trunk.block_conditioning` (gitea #34)."""
|
|
if name not in BLOCK_REGISTRY:
|
|
raise ValueError(f"unknown block conditioning type {name!r}; available: {sorted(BLOCK_REGISTRY)}")
|
|
return BLOCK_REGISTRY[name](dim, cond_dim, dropout)
|
|
|
|
|
|
@register_block("add")
|
|
class ResBlock(nn.Module):
|
|
def __init__(self, dim: int, cond_dim: int, dropout: float = 0.0) -> None:
|
|
super().__init__()
|
|
self.norm = nn.LayerNorm(dim)
|
|
self.linear1 = nn.Linear(dim, dim)
|
|
self.cond_proj = nn.Linear(cond_dim, dim, bias=False)
|
|
self.act = nn.SiLU()
|
|
self.dropout = nn.Dropout(dropout)
|
|
self.linear2 = nn.Linear(dim, dim)
|
|
|
|
def forward(self, x: torch.Tensor, cond: torch.Tensor) -> torch.Tensor:
|
|
h = self.norm(x)
|
|
h = self.linear1(h) + self.cond_proj(cond)
|
|
h = self.act(h)
|
|
h = self.dropout(h)
|
|
h = self.linear2(h)
|
|
return x + h
|
|
|
|
|
|
@register_block("film")
|
|
class FilmResBlock(nn.Module):
|
|
"""FiLM conditioning (Perez et al. 2018): a per-channel scale+shift
|
|
modulates the normalized features, on top of the norm's own affine —
|
|
an *additional* modulation, unlike `AdaLNResBlock` below, which replaces
|
|
the norm's affine outright. `film_proj` is zero-initialized so
|
|
`gamma=beta=0` at construction — conditioning has no effect on the
|
|
output until training moves it, a stable starting point (though not a
|
|
literal identity block, since `linear1`/`linear2` aren't zero-init)."""
|
|
|
|
def __init__(self, dim: int, cond_dim: int, dropout: float = 0.0) -> None:
|
|
super().__init__()
|
|
self.norm = nn.LayerNorm(dim)
|
|
self.linear1 = nn.Linear(dim, dim)
|
|
self.film_proj = nn.Linear(cond_dim, 2 * dim)
|
|
nn.init.zeros_(self.film_proj.weight)
|
|
nn.init.zeros_(self.film_proj.bias)
|
|
self.act = nn.SiLU()
|
|
self.dropout = nn.Dropout(dropout)
|
|
self.linear2 = nn.Linear(dim, dim)
|
|
|
|
def forward(self, x: torch.Tensor, cond: torch.Tensor) -> torch.Tensor:
|
|
h = self.norm(x)
|
|
gamma, beta = self.film_proj(cond).chunk(2, dim=-1)
|
|
h = h * (1 + gamma) + beta
|
|
h = self.linear1(h)
|
|
h = self.act(h)
|
|
h = self.dropout(h)
|
|
h = self.linear2(h)
|
|
return x + h
|
|
|
|
|
|
@register_block("adaln")
|
|
class AdaLNResBlock(nn.Module):
|
|
"""AdaLN-Zero conditioning (DiT, Peebles & Xie 2022): the norm's own
|
|
affine is replaced by a conditioning-derived scale/shift, and the
|
|
residual branch is scaled by a conditioning-derived gate. `adaln_proj`
|
|
is zero-initialized, so `scale=shift=gate=0` at construction — the block
|
|
is the exact identity function at init (`x + 0 * h' == x`), regardless
|
|
of `x`/`cond`."""
|
|
|
|
def __init__(self, dim: int, cond_dim: int, dropout: float = 0.0) -> None:
|
|
super().__init__()
|
|
self.norm = nn.LayerNorm(dim, elementwise_affine=False)
|
|
self.linear1 = nn.Linear(dim, dim)
|
|
self.adaln_proj = nn.Linear(cond_dim, 3 * dim)
|
|
nn.init.zeros_(self.adaln_proj.weight)
|
|
nn.init.zeros_(self.adaln_proj.bias)
|
|
self.act = nn.SiLU()
|
|
self.dropout = nn.Dropout(dropout)
|
|
self.linear2 = nn.Linear(dim, dim)
|
|
|
|
def forward(self, x: torch.Tensor, cond: torch.Tensor) -> torch.Tensor:
|
|
h = self.norm(x)
|
|
scale, shift, gate = self.adaln_proj(cond).chunk(3, dim=-1)
|
|
h = h * (1 + scale) + shift
|
|
h = self.linear1(h)
|
|
h = self.act(h)
|
|
h = self.dropout(h)
|
|
h = self.linear2(h)
|
|
return x + gate * h
|