"""Decoder block: x += attn(norm(x)); x += ffn(norm(x)). attn/ffn are any modules with forward: [B,T,D] -> [B,T,D]. """ from __future__ import annotations from torch import nn from .kda_attn import KDAAttention from .latent_moe import LatentMoE from .mla import GatedMLA from .rmsnorm import RMSNorm from .swiglu import SwiGLUMLP def build_attn(config, kind: str) -> nn.Module: if kind == "kda": return KDAAttention.from_config(config) if kind == "mla": return GatedMLA.from_config(config) raise ValueError(f"unknown attn kind: {kind}") def build_ffn(config, kind: str) -> nn.Module: if kind == "swiglu": return SwiGLUMLP.from_config(config) if kind == "moe": return LatentMoE.from_config(config) raise ValueError(f"unknown ffn kind: {kind}") class DecoderBlock(nn.Module): def __init__(self, hidden_size: int, norm_eps: float, attn: nn.Module, ffn: nn.Module): super().__init__() self.attn_norm = RMSNorm(hidden_size, norm_eps) self.attn = attn self.ffn_norm = RMSNorm(hidden_size, norm_eps) self.ffn = ffn @classmethod def from_spec(cls, config, attn_kind: str, ffn_kind: str) -> DecoderBlock: return cls( config.hidden_size, config.norm_eps, build_attn(config, attn_kind), build_ffn(config, ffn_kind), ) def forward(self, x): x = x + self.attn(self.attn_norm(x)) return x + self.ffn(self.ffn_norm(x))