"""Causal LM stem: embed -> DecoderBlock* -> norm -> lm_head. KDA-only and K3-like both use this class. Config.layer_specs() chooses attn/ffn per layer: ("kda"|"mla", "swiglu"|"moe"). ``config.attnres`` selects the depth mixer: off — standard residual inside each DecoderBlock (default) full — Full AttnRes over attn|ffn sublayers block — Block AttnRes (K3); block size from ``attnres_block_size`` """ from __future__ import annotations import torch import torch.nn.functional as F from torch import nn from torch.utils.checkpoint import checkpoint as activation_checkpoint from ..layers.attn_res import ( BlockAttnResStack, BorrowedSubLayer, FullAttnResStack, atomic_block_size, ) from ..layers.block import DecoderBlock from ..layers.rmsnorm import RMSNorm def _build_mixer(config, blocks: nn.ModuleList): mode = getattr(config, "attnres", "off") if mode == "off": return None atomics = [] for block in blocks: atomics.append(BorrowedSubLayer(block.attn_norm, block.attn)) atomics.append(BorrowedSubLayer(block.ffn_norm, block.ffn)) kwargs = dict( eps=config.norm_eps, zero_init_queries=getattr(config, "attnres_zero_init_queries", True), is_final_aggregate=getattr(config, "attnres_final_aggregate", True), ) if mode == "full": return FullAttnResStack(config.hidden_size, atomics, **kwargs) if mode == "block": return BlockAttnResStack( config.hidden_size, atomics, block_size=atomic_block_size( config.num_hidden_layers, getattr(config, "attnres_block_size", None) ), **kwargs, ) raise ValueError(f"unknown attnres mode: {mode!r}") class CausalLM(nn.Module): def __init__(self, config): super().__init__() self.config = config self.attnres = getattr(config, "attnres", "off") self.embedding = nn.Embedding(config.vocab_size, config.hidden_size) self.blocks = nn.ModuleList( [ DecoderBlock.from_spec(config, attn, ffn) for attn, ffn in config.layer_specs() ] ) self.mixer = _build_mixer(config, self.blocks) self.gradient_checkpointing = bool( getattr(config, "gradient_checkpointing", False) ) self.norm = RMSNorm(config.hidden_size, config.norm_eps) self.lm_head = nn.Linear(config.hidden_size, config.vocab_size, bias=False) nn.init.normal_(self.embedding.weight, std=config.initializer_range) nn.init.normal_(self.lm_head.weight, std=config.initializer_range) if config.tie_word_embeddings: self.lm_head.weight = self.embedding.weight def forward( self, input_ids: torch.Tensor, labels: torch.Tensor | None = None, ignore_index: int = -100, ): x = self.embedding(input_ids) if self.mixer is None: for block in self.blocks: if self.gradient_checkpointing and self.training: x = activation_checkpoint(block, x, use_reentrant=False) else: x = block(x) elif self.gradient_checkpointing and self.training: x = activation_checkpoint(self.mixer, x, use_reentrant=False) else: x = self.mixer(x) logits = self.lm_head(self.norm(x)) if labels is None: return logits return F.cross_entropy( logits[:, :-1].reshape(-1, logits.size(-1)), labels[:, 1:].reshape(-1), ignore_index=ignore_index, ) @torch.inference_mode() def generate( self, input_ids: torch.Tensor, max_new_tokens: int, temperature: float = 0.0, eos_token_id: int | None = None, ): for _ in range(max_new_tokens): logits = self(input_ids)[:, -1] if temperature > 0: probs = F.softmax(logits / temperature, dim=-1) next_token = torch.multinomial(probs, 1) else: next_token = logits.argmax(-1, keepdim=True) input_ids = torch.cat((input_ids, next_token), dim=1) if eos_token_id is not None and (next_token.squeeze(-1) == eos_token_id).all(): break return input_ids