72f5a891bf
CI / Format (ruff format) (push) Successful in 27s
CI / Lint (ruff check) (push) Successful in 29s
CI / Sync project version with tag (push) Has been skipped
CI / Lint (ruff check) (pull_request) Successful in 31s
CI / Type check (ty) (push) Successful in 35s
CI / Format (ruff format) (pull_request) Successful in 42s
CI / Sync project version with tag (pull_request) Has been skipped
CI / Type check (ty) (pull_request) Successful in 42s
CI / Tests (pull_request) Successful in 3m42s
CI / Tests (push) Successful in 3m54s
Pure file-move refactor: network.py's 1742 lines held six distinct concerns (layers, condition encoder, routers, trunks, history encoders, stage models, legacy migration, builders) that the v0.3.0 composable-parts refactor already separated at the class level but not the file level. Split along those seams into layers.py/encoders.py/routers.py/trunks.py/ history.py/models.py/_legacy.py/builders.py; network.py is now an 83-line re-export shim so no external import site needed to change. No logic, signature, or behavior changes.
156 lines
4.7 KiB
Python
156 lines
4.7 KiB
Python
"""Trunks: everything downstream of the fused conditioning vector — monolithic
|
|
or expert-routed (issues.md Issue 8)."""
|
|
|
|
import torch
|
|
import torch.nn as nn
|
|
|
|
from giant.model.layers import ResBlock
|
|
from giant.model.routers import Router
|
|
|
|
|
|
class ExpertTrunk(nn.Module):
|
|
"""One small expert: `input_proj -> ResBlock stack -> out_proj`.
|
|
|
|
Unlike v0.2, `out_dim` is independent of `in_dim` — needed by stage-2 AR
|
|
tokens later (`noise_dim` in, `4 + type_dim` out), even though every
|
|
step-2/3 caller still has `in_dim == out_dim`.
|
|
"""
|
|
|
|
def __init__(
|
|
self,
|
|
in_dim: int,
|
|
out_dim: int,
|
|
hidden_dim: int,
|
|
n_blocks: int,
|
|
cond_dim: int,
|
|
dropout: float = 0.0,
|
|
) -> None:
|
|
super().__init__()
|
|
self.input_proj = nn.Linear(in_dim, hidden_dim)
|
|
self.blocks = nn.ModuleList([ResBlock(hidden_dim, cond_dim, dropout=dropout) for _ in range(n_blocks)])
|
|
self.out_proj = nn.Linear(hidden_dim, out_dim)
|
|
|
|
def forward(self, x: torch.Tensor, cond: torch.Tensor) -> torch.Tensor:
|
|
x = self.input_proj(x)
|
|
for block in self.blocks:
|
|
x = block(x, cond)
|
|
return self.out_proj(x)
|
|
|
|
|
|
def _route_forward(
|
|
experts: nn.ModuleList,
|
|
router: Router,
|
|
x: torch.Tensor,
|
|
cond: torch.Tensor,
|
|
cond_cont: torch.Tensor,
|
|
cond_cat: torch.Tensor,
|
|
training: bool,
|
|
) -> torch.Tensor:
|
|
"""Shared dispatch for `RoutedTrunk`.
|
|
|
|
Train mode: full mixture `sum_i weight_i * expert_i(x)` — always
|
|
N-expert dense compute, fully differentiable (`weight` is
|
|
`router.combine_weights`). Eval mode: grouped top-1 dispatch — each row
|
|
runs exactly one expert, the actual source of the per-call speedup.
|
|
"""
|
|
if training:
|
|
weights = router.combine_weights(cond_cont, cond_cat) # (B, n_experts)
|
|
out = torch.zeros(x.shape[0], experts[0].out_proj.out_features, device=x.device)
|
|
for i, expert in enumerate(experts):
|
|
out = out + weights[:, i : i + 1] * expert(x, cond)
|
|
return out
|
|
|
|
idx = router.top1(cond_cont, cond_cat) # (B,)
|
|
out_dim = experts[0].out_proj.out_features
|
|
out = torch.zeros(x.shape[0], out_dim, device=x.device)
|
|
for i, expert in enumerate(experts):
|
|
mask = idx == i
|
|
if mask.any():
|
|
out[mask] = expert(x[mask], cond[mask])
|
|
return out
|
|
|
|
|
|
class Trunk(nn.Module):
|
|
"""Interface implemented by `MonolithicTrunk`/`RoutedTrunk`: everything
|
|
downstream of the fused conditioning vector, i.e. the actual generative
|
|
trunk of a stage (`input_proj -> blocks -> out_proj`, monolithic or
|
|
expert-routed)."""
|
|
|
|
def forward(
|
|
self,
|
|
x: torch.Tensor,
|
|
cond: torch.Tensor,
|
|
cond_cont: torch.Tensor,
|
|
cond_cat: torch.Tensor,
|
|
) -> torch.Tensor:
|
|
raise NotImplementedError
|
|
|
|
|
|
class MonolithicTrunk(Trunk):
|
|
def __init__(
|
|
self,
|
|
in_dim: int,
|
|
out_dim: int,
|
|
hidden_dim: int,
|
|
n_res_blocks: int,
|
|
cond_dim: int,
|
|
dropout: float = 0.0,
|
|
) -> None:
|
|
super().__init__()
|
|
self.input_proj = nn.Linear(in_dim, hidden_dim)
|
|
self.blocks = nn.ModuleList([ResBlock(hidden_dim, cond_dim, dropout=dropout) for _ in range(n_res_blocks)])
|
|
self.out_proj = nn.Linear(hidden_dim, out_dim)
|
|
|
|
def forward(
|
|
self,
|
|
x: torch.Tensor,
|
|
cond: torch.Tensor,
|
|
cond_cont: torch.Tensor,
|
|
cond_cat: torch.Tensor,
|
|
) -> torch.Tensor:
|
|
x = self.input_proj(x)
|
|
for block in self.blocks:
|
|
x = block(x, cond)
|
|
return self.out_proj(x)
|
|
|
|
|
|
class RoutedTrunk(Trunk):
|
|
def __init__(
|
|
self,
|
|
router: Router,
|
|
in_dim: int,
|
|
out_dim: int,
|
|
hidden_dim: int,
|
|
n_res_blocks: int,
|
|
cond_dim: int,
|
|
dropout: float = 0.0,
|
|
) -> None:
|
|
super().__init__()
|
|
self.router = router
|
|
self.experts = nn.ModuleList(
|
|
[ExpertTrunk(in_dim, out_dim, hidden_dim, n_res_blocks, cond_dim, dropout) for _ in range(router.n_experts)]
|
|
)
|
|
|
|
def forward(
|
|
self,
|
|
x: torch.Tensor,
|
|
cond: torch.Tensor,
|
|
cond_cont: torch.Tensor,
|
|
cond_cat: torch.Tensor,
|
|
) -> torch.Tensor:
|
|
return _route_forward(self.experts, self.router, x, cond, cond_cont, cond_cat, self.training)
|
|
|
|
|
|
def build_trunk(
|
|
router: Router | None,
|
|
in_dim: int,
|
|
out_dim: int,
|
|
hidden_dim: int,
|
|
n_res_blocks: int,
|
|
cond_dim: int,
|
|
dropout: float = 0.0,
|
|
) -> Trunk:
|
|
if router is not None:
|
|
return RoutedTrunk(router, in_dim, out_dim, hidden_dim, n_res_blocks, cond_dim, dropout)
|
|
return MonolithicTrunk(in_dim, out_dim, hidden_dim, n_res_blocks, cond_dim, dropout)
|