# Copyright Lightning AI. Licensed under the Apache License 2.0, see LICENSE file. """Full definition of a decoder-only transformer-based language model, all of it in this single file. Based on the nanoGPT implementation: https://github.com/karpathy/nanoGPT and https://github.com/EleutherAI/gpt-neox/tree/main/megatron/model. """ import math from copy import copy from functools import partial from typing import Any import torch import torch.nn as nn import torch.nn.functional as F from torch.utils.checkpoint import checkpoint from torch.nn.attention.flex_attention import BlockMask, create_block_mask, flex_attention, flex_attention from typing_extensions import Self from litgpt.config import Config from litgpt.scripts.convert_hf_checkpoint import qkv_reassemble def build_attention(config: Config, attention_idx: int) -> nn.Module: """Construct the configured mixer for one attention-layer ordinal.""" if config.kda_enabled: period = config.kda_mla_ratio + 1 if (attention_idx + 1) % period: return KimiDeltaAttentionAdapter(config, attention_idx) return MultiheadLatentAttention(config, attention_idx) if config.latent_attention: return MultiheadLatentAttention(config, attention_idx) if config.attention_variant == "shared_diff": return SharedDiffCausalSelfAttention(config, attention_idx) return CausalSelfAttention(config, attention_idx) class GPT(nn.Module): def __init__(self, config: Config) -> None: super().__init__() assert config.padded_vocab_size is not None self.config = config self.lm_head = ( nn.Identity() if config.mamba3_hierarchical_vocab else nn.Linear(config.n_embd, config.padded_vocab_size, bias=config.lm_head_bias) ) if config.mamba3_hierarchical_vocab: mixtures = config.mamba3_hierarchical_mixtures self.mamba3_cluster_head = nn.Linear( config.n_embd, mixtures * 64, bias=False ) self.mamba3_slot_head = nn.Linear( config.n_embd, mixtures * 64, bias=False ) self.mamba3_mixture_head = ( nn.Linear(config.n_embd, mixtures, bias=False) if mixtures > 1 else None ) multipliers = torch.tensor( [1, 2053, 1597, 3135, 817, 3327, 1231, 2871], dtype=torch.long, )[:mixtures] offsets = torch.tensor( [0, 911, 1877, 283, 3491, 1453, 2579, 3677], dtype=torch.long, )[:mixtures] vocabulary = torch.arange(4096, dtype=torch.long) self.register_buffer( "mamba3_vocab_multipliers", multipliers, persistent=False ) self.register_buffer( "mamba3_vocab_offsets", offsets, persistent=False ) self.register_buffer( "mamba3_vocab_permutations", (multipliers[:, None] * vocabulary[None] + offsets[:, None]) .bitwise_and(4095), persistent=False, ) self.mamba3_position_embedding = ( nn.Embedding(config.mamba3_position_table_size, config.n_embd) if config.mamba3_position_table_size else None ) self.mamba3_phase_embedding = ( nn.Embedding(config.mamba3_phase_table_size, config.n_embd) if config.mamba3_phase_table_size else None ) self.mamba3_token_scalar = ( nn.Embedding(config.mamba3_token_scalar_params, 1) if config.mamba3_token_scalar_params else None ) block_class = RMSBudgetedBlock if config.rms_budgeted_block else Block if config.mamba3_enabled: from litgpt.mamba3_hybrid import Mamba3MultiscreenBlock blocks = nn.ModuleList(Mamba3MultiscreenBlock(config, block_idx) for block_idx in range(config.n_layer)) elif config.rwkv7_enabled: from litgpt.rwkv7 import RWKV7HybridBlock if not config.rwkv7_hybrid_multiscreen: raise NotImplementedError("The prototype currently exposes the RWKV-7 + Multiscreen hybrid.") blocks = nn.ModuleList(RWKV7HybridBlock(config, block_idx) for block_idx in range(config.n_layer)) elif config.multiscreen_enabled: blocks = nn.ModuleList(MultiscreenBlock(config, block_idx) for block_idx in range(config.n_layer)) elif config.multiscreen_layer_interval: blocks = nn.ModuleList( SelectiveRecallBlock(config, block_idx) for block_idx in range(config.n_layer) ) elif config.sandwich_coefficient: k = config.sandwich_coefficient architecture = "s" * k + "sf" * (config.n_layer - k) + "f" * k blocks = nn.ModuleList( SandwichSublayer(config, sublayer_idx, kind) for sublayer_idx, kind in enumerate(architecture) ) else: blocks = nn.ModuleList(block_class(config, block_idx) for block_idx in range(config.n_layer)) depth_memory = nn.ModuleDict() if config.depth_memory_mode != "none": for sublayer_idx in range(len(blocks)): depth = sublayer_idx + 1 if depth >= config.depth_memory_start and depth % config.depth_memory_interval == 0: depth_memory[str(sublayer_idx)] = AdaptiveDepthMemory(config) self.transformer = nn.ModuleDict( dict( wte=nn.Embedding( config.padded_vocab_size * config.mamba3_embedding_banks, config.n_embd, ), h=blocks, depth_memory=depth_memory, ln_f=( nn.Identity() if config.multiscreen_enabled else nn.LayerNorm( config.n_embd, eps=config.norm_eps, bias=not config.rwkv7_exact_10m_norms, ) if config.rwkv7_enabled and config.rwkv7_exact_10m_norms else config.norm_class(config.n_embd, eps=config.norm_eps) ), ) ) mtp_config = config if config.mtp_num_layers and config.sliding_window_indices is not None: mtp_config = copy(config) period = config.sliding_window_indices mtp_config.sliding_window_indices = [ *period, *(period[index % len(period)] for index in range(config.mtp_num_layers)), ] mtp_block_class = MultiscreenBlock if config.multiscreen_enabled else Block self.mtp_blocks = nn.ModuleList( mtp_block_class(mtp_config, config.n_layer + index) for index in range(config.mtp_num_layers) ) self.mtp_enorm = ( config.norm_class(config.n_embd, eps=config.norm_eps) if config.mtp_num_layers else None ) self.mtp_hnorm = ( config.norm_class(config.n_embd, eps=config.norm_eps) if config.mtp_num_layers else None ) self.mtp_eh_proj = ( nn.Linear(2 * config.n_embd, config.n_embd, bias=False) if config.mtp_num_layers else None ) self.mtp_eagle3_feature_proj = ( nn.Linear(3 * config.n_embd, config.n_embd, bias=False) if config.mtp_eagle3 else None ) if config.attn_residual_num_blocks or config.attn_residual_block_size: num_sublayers = len(blocks) num_reads = num_sublayers + 1 self.attn_residual_norms = nn.ModuleList( config.norm_class(config.n_embd, eps=config.norm_eps) for _ in range(num_reads) ) self.attn_residual_queries = nn.ModuleList( nn.Linear(config.n_embd, 1, bias=False) for _ in range(num_reads) ) self.attn_residual_block_size = ( config.attn_residual_block_size or num_sublayers // config.attn_residual_num_blocks ) else: self.attn_residual_norms = nn.ModuleList() self.attn_residual_queries = nn.ModuleList() self.attn_residual_block_size = 0 if config.multiscreen_enabled: self.multiscreen_embedding_scale = nn.Parameter(torch.zeros(())) # The Multiscreen residual initially remains strongly aligned # with the current normalized token embedding. Because the readout # is tied to that same normalized embedding table, initializing # its norm to sqrt(n_embd) gives the current token a logit near # sqrt(n_embd), even though the target is the next token. # # A fixed initial norm of 4 keeps initial CE near log(vocab_size) # at every residual width. This is still a learnable log-scale. self.multiscreen_readout_scale = nn.Parameter( torch.tensor(math.log(4.0)) ) self.register_buffer("depth_memory_scale", torch.tensor(1.0), persistent=False) self.model_match_scale = ( nn.Parameter(torch.zeros(config.model_match_adapter_params)) if config.model_match_adapter_params else None ) self.mask_cache: torch.Tensor | None = None self.max_seq_length = self.config.block_size def _apply_model_match_scale(self, x: torch.Tensor) -> torch.Tensor: if self.model_match_scale is not None: scale = self.model_match_scale if scale.numel() < self.config.n_embd: scale = F.pad(scale, (0, self.config.n_embd - scale.numel())) elif scale.numel() > self.config.n_embd: rows = math.ceil(scale.numel() / self.config.n_embd) scale = F.pad( scale, (0, rows * self.config.n_embd - scale.numel()), ).view(rows, self.config.n_embd).mean(dim=0) x = x * (1.0 + 0.01 * torch.tanh(scale)) return x def _matched_final_norm(self, x: torch.Tensor) -> torch.Tensor: return self._apply_model_match_scale(self.transformer.ln_f(x)) def _attention_residual_read( self, blocks: list[torch.Tensor], partial_block: torch.Tensor, norm: nn.Module, query: nn.Module, ) -> torch.Tensor: """Kimi Block AttnRes inter-block softmax aggregation.""" values = torch.stack([*blocks, partial_block], dim=0) keys = norm(values) logits = query(keys).squeeze(-1) weights = F.softmax(logits, dim=0, dtype=torch.float).to(values.dtype) return torch.einsum("nbt,nbtd->btd", weights, values) def prepare_sliding_window_masks(self, sequence_length: int, device: torch.device) -> None: """Build block-sparse causal masks before ``torch.compile`` traces the model.""" for module in self.modules(): if isinstance(module, CausalSelfAttention): module.prepare_sliding_window_mask(sequence_length, device) def set_sliding_window_enabled(self, enabled: bool) -> int: """Toggle configured local-attention layers without changing parameters.""" changed = 0 for module in self.modules(): if not isinstance(module, CausalSelfAttention): continue next_value = module.configured_sliding_window_attention and enabled if module.apply_sliding_window_attention != next_value: module.apply_sliding_window_attention = next_value changed += 1 return changed @property def max_seq_length(self) -> int: return self._max_seq_length @max_seq_length.setter def max_seq_length(self, value: int) -> None: """ When doing inference, the sequences used might be shorter than the model's context length. This allows setting a smaller number to avoid allocating unused memory """ if value > self.config.block_size: raise ValueError( f"Cannot attend to {value}, block size is only {self.config.block_size}." " This is likely because the input text exceeds the supported context length of this model." ) self._max_seq_length = value if not hasattr(self, "cos"): # first call cos, sin = self.rope_cache() self.register_buffer("cos", cos, persistent=False) self.register_buffer("sin", sin, persistent=False) # override elif value != self.cos.size(0): self.cos, self.sin = self.rope_cache(device=self.cos.device) # the mask and kv cache size will get updated on `set_kv_cache`. we cannot update it here because we don't know # if the kv cache is expected if self.mask_cache is not None and self.mask_cache.shape[-1] < value: print( f"Warning: KV cache has length {self.mask_cache.shape[-1]} < {value} = max_seq_length. Call 'set_kv_cache' before doing any forwards!" ) def reset_parameters(self) -> None: # Trigger resetting the rope-cache self.cos, self.sin = self.rope_cache(device=self.cos.device) def _init_weights(self, module: nn.Module) -> None: """Meant to be used with `gpt.apply(gpt._init_weights)`.""" if module is self.mtp_eagle3_feature_proj: # Kimi K3 initializes W_E3=[0 0 I], so enabling EAGLE-3 starts # exactly from the high-level feature used for MTP pretraining. with torch.no_grad(): module.weight.zero_() module.weight[:, -self.config.n_embd :].copy_( torch.eye( self.config.n_embd, device=module.weight.device, dtype=module.weight.dtype, ) ) elif isinstance(module, GroupedTopkRouter): torch.nn.init.normal_(module.weight.data, mean=0.0, std=0.02) elif isinstance(module, nn.Linear): torch.nn.init.normal_(module.weight, mean=0.0, std=0.02) if module.bias is not None: torch.nn.init.zeros_(module.bias) elif isinstance(module, nn.Embedding): torch.nn.init.normal_(module.weight, mean=0.0, std=0.02) def forward( self, idx: torch.Tensor, input_pos: torch.Tensor | None = None, return_hidden: bool = False, return_mtp_features: bool = False, input_pos_maxp1: int | None = None, lm_head_chunk_size: int = 0, targets: torch.Tensor | None = None, ) -> torch.Tensor | list[torch.Tensor]: """ If `input_pos` is provided, the KV cache uses K and V vectors for positions smaller than entries in `input_pos`. For efficiency, pass `input_pos_maxp1` as `max(input_pos) + 1` if already available from your forward algorithm. This slices the KV cache buffers and speeds up multi-head attention. Without `input_pos_maxp1`, the computation uses the full KV cache (`max_seq_length`) with masking applied. Note that inferring `input_pos_maxp1` from `input_pos` causes graph breaks and prevents compilation. Args: idx: Token indices of input sequences, shape `(B, T)`, where `B` is batch size. input_pos: Optional. Positions of input tokens. The default is `arange(T)`. Can have shape `(T,)` or `(B, T)` (batched index). input_pos_maxp1: Optional. See above. lm_head_chunk_size: Optional. If `lm_head_chunk_size > 0`, the final `lm_head` computation is done in chunks of this size. Returns: Logit outputs, shape `(B, T, config.padded_vocab_size)`. If `lm_head_chunk_size > 0`, this is a list of chunks of shape `(B, lm_head_chunk_size, config.padded_vocab_size)`, the final entry can be shorter. """ if return_hidden and return_mtp_features: raise ValueError("Request either the final hidden state or EAGLE-3 features, not both.") if return_mtp_features and not self.config.mtp_eagle3: raise ValueError("return_mtp_features requires mtp_eagle3=True.") T = idx.size(1) if self.max_seq_length < T: raise ValueError(f"Cannot forward sequence of length {T}, max seq length is only {self.max_seq_length}.") if input_pos is not None: # use the kv cache if input_pos.dim() > 2: # otherwise, things go wrong in `apply_rope` raise ValueError(f"input_pos must have 1 or 2 dimensions, input_pos.shape = {input_pos.shape}") if input_pos.shape[-1] != T: raise ValueError(f"input_pos.shape[-1] = {input_pos.shape[-1]} != {T} = idx.shape[1], must be the same") cos = batched_index_select(self.cos, 0, input_pos) sin = batched_index_select(self.sin, 0, input_pos) if input_pos.dim() == 1: cos = cos.unsqueeze(0) sin = sin.unsqueeze(0) if self.mask_cache is None: raise TypeError("You need to call `gpt.set_kv_cache()`") mask = batched_index_select(self.mask_cache, 2, input_pos) if mask.dim() > 4: # the mask cache has a batch dim of 1 in addition to the one # we get if input_pos has a batch dimension mask = mask.view(*(mask.shape[0:1] + mask.shape[2:])) if input_pos_maxp1 is not None: # Shorten final dimension so it just covers all `input_pos` entries if input_pos_maxp1 > self.max_seq_length: raise ValueError(f"Positions in 'input_pos' must be in [0,{self.max_seq_length})") mask = mask[..., :input_pos_maxp1] else: # unsqueeze to have a batch dimension cos = self.cos[:T].unsqueeze(0) sin = self.sin[:T].unsqueeze(0) # `cos`, `sin` have shape (1, T, config.rope_n_elem) mask = None # defaults to causal mask input_pos_maxp1 = None if self.config.multiscreen_enabled and input_pos is not None: raise NotImplementedError("Multiscreen KV-cache decode is not implemented; use full-sequence evaluation.") if self.config.multiscreen_enabled: embedding = F.normalize(self.transformer.wte.weight, dim=-1, eps=1e-6) x = F.embedding(idx, embedding) * self.multiscreen_embedding_scale.exp() elif self.config.mamba3_hierarchical_vocab: positions = torch.arange(T, device=idx.device) bank = positions.remainder(self.config.mamba3_embedding_banks) striped_idx = idx + bank[None] * self.config.padded_vocab_size x = self.transformer.wte(striped_idx) if self.mamba3_position_embedding is not None: x = x + self.mamba3_position_embedding( positions.remainder(self.config.mamba3_position_table_size) )[None] if self.mamba3_phase_embedding is not None: x = x + self.mamba3_phase_embedding( positions.remainder(self.config.mamba3_phase_table_size) )[None] if self.mamba3_token_scalar is not None: x = x + self.mamba3_token_scalar( idx.remainder(self.config.mamba3_token_scalar_params) ) else: x = self.transformer.wte(idx) # token embeddings of shape (B, T, n_embd) if self.config.scale_embeddings: x = x * torch.tensor(self.config.n_embd**0.5, dtype=x.dtype) first_layer_values: torch.Tensor | None = None depth_memory_states = [x] if ( self.config.multiscreen_enabled and self.config.multiscreen_activation_checkpointing and self.training and torch.is_grad_enabled() ): if self.config.rope_indices is not None or self.config.value_residual_mix > 0.0: raise NotImplementedError( "Segmented Multiscreen checkpointing does not support per-layer RoPE " "indices or value residual mixing." ) group_size = max(1, self.config.multiscreen_checkpoint_group_size) blocks = tuple(self.transformer.h) for start in range(0, len(blocks), group_size): segment_blocks = blocks[start : start + group_size] def run_segment( segment_input: torch.Tensor, segment: tuple[nn.Module, ...] = segment_blocks, ) -> torch.Tensor: for segment_block in segment: segment_input = segment_block( segment_input, cos, sin, None, None, None ) return segment_input x = checkpoint( run_segment, x, # Reentrant checkpointing runs the original segment # forward under no_grad. For Multiscreen this avoids # retaining the compiled inner autograd graph, which is # much larger than the single segment input at B=32. use_reentrant=True, preserve_rng_state=False, ) x = F.normalize(x, dim=-1, eps=1e-6) * self.multiscreen_readout_scale.exp() x = self._apply_model_match_scale(x) weight = F.normalize(self.transformer.wte.weight, dim=-1, eps=1e-6) if targets is not None and self.config.multiscreen_chunked_lm_loss: from litgpt.multiscreen_loss import chunked_linear_cross_entropy loss = chunked_linear_cross_entropy( x, weight, targets, self.config.multiscreen_lm_loss_chunk_tokens, ) return (loss, x) if return_hidden else loss if lm_head_chunk_size > 0: return [F.linear(x_i, weight) for x_i in x.split(lm_head_chunk_size, dim=1)] logits = F.linear(x, weight) return (logits, x) if return_hidden else logits mamba3_checkpointed = ( self.config.mamba3_enabled and self.config.mamba3_activation_checkpointing and self.training and torch.is_grad_enabled() ) if mamba3_checkpointed: group_size = self.config.mamba3_checkpoint_group_size blocks = tuple(self.transformer.h) for start in range(0, len(blocks), group_size): segment_blocks = blocks[start : start + group_size] def run_mamba3_segment( segment_input: torch.Tensor, segment: tuple[nn.Module, ...] = segment_blocks, ) -> torch.Tensor: for segment_block in segment: segment_input = segment_block( segment_input, cos, sin, None, None, None ) return segment_input x = checkpoint( run_mamba3_segment, x, # Mamba-3's custom Triton autograd reads saved_tensors more # than once, which is incompatible with non-reentrant # checkpoint hooks. use_reentrant=True, preserve_rng_state=False, ) rwkv_first_value: torch.Tensor | None = None attn_residual_blocks: list[torch.Tensor] | None = None attn_residual_partial: torch.Tensor | None = None if self.config.attn_residual_num_blocks or self.config.attn_residual_block_size: # The paper treats the token embedding as the first completed # source and the current hidden state as the initial partial sum. attn_residual_blocks = [x] attn_residual_partial = x active_blocks = self.transformer.h if not mamba3_checkpointed else () attn_residual_norms = ( self.attn_residual_norms[:-1] if attn_residual_blocks is not None else (None,) * len(active_blocks) ) attn_residual_queries = ( self.attn_residual_queries[:-1] if attn_residual_blocks is not None else (None,) * len(active_blocks) ) for block_idx, (block, attn_residual_norm, attn_residual_query) in enumerate( zip(active_blocks, attn_residual_norms, attn_residual_queries) ): if self.config.rope_indices is not None: block_cos = cos[..., self.config.rope_indices[block_idx]] block_sin = sin[..., self.config.rope_indices[block_idx]] else: block_cos = cos block_sin = sin if attn_residual_blocks is not None: assert attn_residual_partial is not None # Keep the complete depth-attention calculation in the parent # forward graph. Calling a shared helper here makes Dynamo # specialize that helper on every different ``len(blocks)`` # and fall back after eight AttnRes reads. attn_residual_values = torch.stack( [*attn_residual_blocks, attn_residual_partial], dim=0, ) attn_residual_keys = attn_residual_norm(attn_residual_values) attn_residual_logits = attn_residual_query( attn_residual_keys ).squeeze(-1) attn_residual_weights = F.softmax( attn_residual_logits, dim=0, dtype=torch.float, ).to(attn_residual_values.dtype) h = torch.einsum( "nbt,nbtd->btd", attn_residual_weights, attn_residual_values, ) # Block AttnRes commits the completed partial block after its # final read and starts the next block from the next update. if block_idx and block_idx % self.attn_residual_block_size == 0: attn_residual_blocks.append(attn_residual_partial) attn_residual_partial = None update = block( h, block_cos, block_sin, mask, input_pos, input_pos_maxp1, ) - h attn_residual_partial = ( update if attn_residual_partial is None else attn_residual_partial + update ) x = attn_residual_partial elif self.config.rwkv7_enabled: x, rwkv_first_value, _ = block( x, block_cos, block_sin, mask, input_pos, input_pos_maxp1, first_value=rwkv_first_value, ) elif self.config.value_residual_mix > 0.0: x, current_values = block( x, block_cos, block_sin, mask, input_pos, input_pos_maxp1, value_residual=first_layer_values, ) if first_layer_values is None: first_layer_values = current_values else: x = block(x, block_cos, block_sin, mask, input_pos, input_pos_maxp1) depth_memory_key = str(block_idx) if depth_memory_key in self.transformer.depth_memory: if self.config.depth_memory_mode == "adaptive": x = self.transformer.depth_memory[depth_memory_key]( x, depth_memory_states, self.depth_memory_scale ) if self.config.depth_memory_score_scale == "recurrent": # ADeM-R keeps only the previous committed depth state. # This avoids retaining and stacking four full residuals. depth_memory_states = [x] else: depth_memory_states.append(x) mtp_features: tuple[torch.Tensor, torch.Tensor, torch.Tensor] | None = None if attn_residual_blocks is not None: assert attn_residual_partial is not None completed_attnres_blocks = [ *attn_residual_blocks[1:], attn_residual_partial, ] if return_mtp_features: if len(completed_attnres_blocks) < 4: raise RuntimeError("EAGLE-3 requires at least four completed AttnRes blocks.") # Kimi K3 uses outputs of the 1st, 4th, and final blocks. mtp_features = ( completed_attnres_blocks[0], completed_attnres_blocks[3], completed_attnres_blocks[-1], ) attn_residual_values = torch.stack( [*attn_residual_blocks, attn_residual_partial], dim=0, ) attn_residual_keys = self.attn_residual_norms[-1]( attn_residual_values ) attn_residual_logits = self.attn_residual_queries[-1]( attn_residual_keys ).squeeze(-1) attn_residual_weights = F.softmax( attn_residual_logits, dim=0, dtype=torch.float, ).to(attn_residual_values.dtype) x = torch.einsum( "nbt,nbtd->btd", attn_residual_weights, attn_residual_values, ) if self.config.multiscreen_enabled: x = F.normalize(x, dim=-1, eps=1e-6) * self.multiscreen_readout_scale.exp() x = self._apply_model_match_scale(x) weight = F.normalize(self.transformer.wte.weight, dim=-1, eps=1e-6) if targets is not None and self.config.multiscreen_chunked_lm_loss: from litgpt.multiscreen_loss import chunked_linear_cross_entropy loss = chunked_linear_cross_entropy( x, weight, targets, self.config.multiscreen_lm_loss_chunk_tokens, ) return (loss, x) if return_hidden else loss if lm_head_chunk_size > 0: return [F.linear(x_i, weight) for x_i in x.split(lm_head_chunk_size, dim=1)] logits = F.linear(x, weight) return (logits, x) if return_hidden else logits mtp_hidden = x x = self._matched_final_norm(x) if self.config.mamba3_hierarchical_vocab: logits_or_loss = self._hierarchical_vocab_output(x, targets) if targets is not None: return logits_or_loss logits = logits_or_loss return (logits, mtp_hidden) if return_hidden else logits if ( targets is not None and self.config.multiscreen_chunked_lm_loss ): if self.config.final_logit_softcapping is not None: raise NotImplementedError( "Chunked linear CE does not support logit softcapping." ) from litgpt.multiscreen_loss import chunked_linear_cross_entropy loss = chunked_linear_cross_entropy( x, self.lm_head.weight, targets, self.config.multiscreen_lm_loss_chunk_tokens, ) return (loss, mtp_hidden) if return_hidden else loss clamp_head = ( partial(do_softcapping, thresh=self.config.final_logit_softcapping) if self.config.final_logit_softcapping is not None else nn.Identity() ) if lm_head_chunk_size > 0: return [clamp_head(self.lm_head(x_i)) for x_i in x.split(lm_head_chunk_size, dim=1)] else: logits = clamp_head(self.lm_head(x)) if return_mtp_features: assert mtp_features is not None return logits, mtp_features return (logits, mtp_hidden) if return_hidden else logits def forward_rwkv7_step( self, idx: torch.Tensor, states: list[Any] | None = None, ) -> tuple[torch.Tensor, list[Any]]: """Decode one token with fixed-size RWKV and local-screening state.""" if not self.config.rwkv7_enabled: raise RuntimeError("forward_rwkv7_step requires rwkv7_enabled.") if idx.ndim != 2 or idx.size(1) != 1: raise ValueError("forward_rwkv7_step expects token indices shaped (batch, 1).") if states is not None and len(states) != len(self.transformer.h): raise ValueError("One recurrent state is required per hybrid block.") x = self.transformer.wte(idx) if self.config.scale_embeddings: x = x * torch.tensor(self.config.n_embd**0.5, dtype=x.dtype, device=x.device) next_states = [] first_value = None for index, block in enumerate(self.transformer.h): layer_state = None if states is None else states[index] x, first_value, layer_state = block.forward_step(x, first_value, layer_state) next_states.append(layer_state) logits = self.lm_head(self._matched_final_norm(x)) return logits, next_states def forward_multiscreen_step( self, idx: torch.Tensor, states: list[Any] | None = None, ) -> tuple[torch.Tensor, list[Any]]: """Decode one token with bounded per-layer Multiscreen key/value caches.""" if not self.config.multiscreen_enabled or self.config.mamba3_enabled: raise RuntimeError("forward_multiscreen_step requires pure Multiscreen.") if idx.ndim != 2 or idx.size(1) != 1: raise ValueError("forward_multiscreen_step expects token indices shaped (batch, 1).") if states is not None and len(states) != len(self.transformer.h): raise ValueError("One screening state is required per Multiscreen block.") embedding = F.normalize(self.transformer.wte.weight, dim=-1, eps=1e-6) x = F.embedding(idx, embedding) * self.multiscreen_embedding_scale.exp() if self.config.scale_embeddings: x = x * torch.tensor(self.config.n_embd**0.5, dtype=x.dtype, device=x.device) next_states = [] for index, block in enumerate(self.transformer.h): layer_state = None if states is None else states[index] x, layer_state = block.forward_step(x, layer_state) next_states.append(layer_state) x = F.normalize(x, dim=-1, eps=1e-6) * self.multiscreen_readout_scale.exp() x = self._apply_model_match_scale(x) weight = F.normalize(self.transformer.wte.weight, dim=-1, eps=1e-6) return F.linear(x, weight), next_states def forward_mamba3_step( self, idx: torch.Tensor, states: list[Any] | None = None, ) -> tuple[torch.Tensor, list[Any]]: """Decode one token with fixed-size Mamba-3 and local-screening state.""" if not self.config.mamba3_enabled: raise RuntimeError("forward_mamba3_step requires mamba3_enabled.") if idx.ndim != 2 or idx.size(1) != 1: raise ValueError("forward_mamba3_step expects token indices shaped (batch, 1).") if states is not None and len(states) != len(self.transformer.h): raise ValueError("One recurrent state is required per hybrid block.") if self.config.mamba3_hierarchical_vocab: position = 0 if states is None else states[0].next_position bank = position % self.config.mamba3_embedding_banks x = self.transformer.wte(idx + bank * self.config.padded_vocab_size) if self.mamba3_position_embedding is not None: position_idx = torch.tensor( [position % self.config.mamba3_position_table_size], device=idx.device, ) x = x + self.mamba3_position_embedding(position_idx)[None] if self.mamba3_phase_embedding is not None: phase_idx = torch.tensor( [position % self.config.mamba3_phase_table_size], device=idx.device, ) x = x + self.mamba3_phase_embedding(phase_idx)[None] if self.mamba3_token_scalar is not None: x = x + self.mamba3_token_scalar( idx.remainder(self.config.mamba3_token_scalar_params) ) else: x = self.transformer.wte(idx) next_states = [] for index, block in enumerate(self.transformer.h): layer_state = None if states is None else states[index] x, layer_state = block.forward_step(x, layer_state) next_states.append(layer_state) x = self._matched_final_norm(x) if self.config.mamba3_hierarchical_vocab: logits = self._hierarchical_vocab_output(x, None) return logits, next_states return self.lm_head(x), next_states def _hierarchical_vocab_output( self, hidden: torch.Tensor, targets: torch.Tensor | None, ) -> torch.Tensor: mixtures = self.config.mamba3_hierarchical_mixtures cluster = self.mamba3_cluster_head(hidden).view( *hidden.shape[:-1], mixtures, 64 ) slot = self.mamba3_slot_head(hidden).view( *hidden.shape[:-1], mixtures, 64 ) if mixtures == 1: if targets is not None: cluster_targets = torch.div(targets, 64, rounding_mode="floor") slot_targets = targets.remainder(64) return F.cross_entropy( cluster.flatten(0, -2).float(), cluster_targets.flatten(), ) + F.cross_entropy( slot.flatten(0, -2).float(), slot_targets.flatten(), ) cluster = cluster.squeeze(-2) slot = slot.squeeze(-2) return (cluster.unsqueeze(-1) + slot.unsqueeze(-2)).flatten(-2) cluster_log_probability = F.log_softmax(cluster.float(), dim=-1) slot_log_probability = F.log_softmax(slot.float(), dim=-1) mixture_log_probability = F.log_softmax( self.mamba3_mixture_head(hidden).float(), dim=-1 ) if targets is not None: mapped = ( targets[..., None] * self.mamba3_vocab_multipliers + self.mamba3_vocab_offsets ).bitwise_and(4095) cluster_target = torch.div(mapped, 64, rounding_mode="floor") slot_target = mapped.remainder(64) cluster_target_logp = cluster_log_probability.gather( -1, cluster_target.unsqueeze(-1) ).squeeze(-1) slot_target_logp = slot_log_probability.gather( -1, slot_target.unsqueeze(-1) ).squeeze(-1) token_log_probability = torch.logsumexp( mixture_log_probability + cluster_target_logp + slot_target_logp, dim=-1, ) return -token_log_probability.mean() component_log_probability = ( cluster_log_probability.unsqueeze(-1) + slot_log_probability.unsqueeze(-2) ).flatten(-2) gather_index = self.mamba3_vocab_permutations.view( *((1,) * (component_log_probability.ndim - 2)), mixtures, 4096, ).expand_as(component_log_probability) component_log_probability = component_log_probability.gather( -1, gather_index ) return torch.logsumexp( mixture_log_probability.unsqueeze(-1) + component_log_probability, dim=-2, ) def mtp_forward( self, hidden, next_tokens, return_hidden: bool = False, targets: torch.Tensor | None = None, ): T = hidden.size(1) if self.config.multiscreen_enabled: embedding_weight = F.normalize( self.transformer.wte.weight, dim=-1, eps=1e-6, ) embedding = ( F.embedding(next_tokens, embedding_weight) * self.multiscreen_embedding_scale.exp() ) else: embedding = self.transformer.wte(next_tokens) # DeepSeek-V3 Eq. 21: M_k[RMSNorm(h); RMSNorm(Emb(t_{i+k}))]. x = self.mtp_eh_proj( torch.cat((self.mtp_hnorm(hidden), self.mtp_enorm(embedding)), dim=-1) ) cos = self.cos[:T].unsqueeze(0) sin = self.sin[:T].unsqueeze(0) for block in self.mtp_blocks: x = block(x, cos, sin) if self.config.multiscreen_enabled: x = F.normalize(x, dim=-1, eps=1e-6) * self.multiscreen_readout_scale.exp() x = self._apply_model_match_scale(x) weight = F.normalize(self.transformer.wte.weight, dim=-1, eps=1e-6) if targets is not None and self.config.multiscreen_chunked_lm_loss: from litgpt.multiscreen_loss import chunked_linear_cross_entropy return chunked_linear_cross_entropy( x, weight, targets, self.config.multiscreen_lm_loss_chunk_tokens, ) logits = F.linear(x, weight) else: x = self._matched_final_norm(x) if targets is not None and self.config.multiscreen_chunked_lm_loss: from litgpt.multiscreen_loss import chunked_linear_cross_entropy return chunked_linear_cross_entropy( x, self.lm_head.weight, targets, self.config.multiscreen_lm_loss_chunk_tokens, ) logits = self.lm_head(x) return (logits, x) if return_hidden else logits def mtp_forward_multiscreen_step( self, hidden: torch.Tensor, next_tokens: torch.Tensor, states: list[Any] | None = None, return_hidden: bool = False, ) -> tuple[torch.Tensor, list[Any]] | tuple[ torch.Tensor, torch.Tensor, list[Any], ]: """Advance the MTP drafter with bounded Multiscreen recurrent state.""" if not self.config.multiscreen_enabled or self.config.mamba3_enabled: raise RuntimeError( "mtp_forward_multiscreen_step requires pure Multiscreen." ) if len(self.mtp_blocks) != 1: raise RuntimeError( "mtp_forward_multiscreen_step requires exactly one MTP layer." ) if hidden.shape[:2] != next_tokens.shape or next_tokens.size(1) != 1: raise ValueError( "MTP hidden states and tokens must both have shape (batch, 1, ...)." ) if states is not None and len(states) != len(self.mtp_blocks): raise ValueError("One recurrent state is required per MTP block.") embedding_weight = F.normalize( self.transformer.wte.weight, dim=-1, eps=1e-6, ) embedding = ( F.embedding(next_tokens, embedding_weight) * self.multiscreen_embedding_scale.exp() ) x = self.mtp_eh_proj( torch.cat( (self.mtp_hnorm(hidden), self.mtp_enorm(embedding)), dim=-1, ) ) next_states = [] for index, block in enumerate(self.mtp_blocks): layer_state = None if states is None else states[index] x, layer_state = block.forward_step(x, layer_state) next_states.append(layer_state) draft_hidden = x x = F.normalize(x, dim=-1, eps=1e-6) * self.multiscreen_readout_scale.exp() x = self._apply_model_match_scale(x) logits = F.linear(x, embedding_weight) if return_hidden: return logits, draft_hidden, next_states return logits, next_states def mtp_eagle3_forward( self, features: tuple[torch.Tensor, torch.Tensor, torch.Tensor], next_tokens: torch.Tensor, return_hidden: bool = False, ): """Run the Kimi K3 EAGLE-3 drafter from three AttnRes block features.""" if self.mtp_eagle3_feature_proj is None: raise RuntimeError("mtp_eagle3_forward requires mtp_eagle3=True.") if len(features) != 3: raise ValueError("EAGLE-3 requires exactly low-, mid-, and high-level features.") hidden = self.mtp_eagle3_feature_proj(torch.cat(features, dim=-1)) return self.mtp_forward(hidden, next_tokens, return_hidden=return_hidden) def mtp_eagle3_forward_step( self, features: tuple[torch.Tensor, torch.Tensor, torch.Tensor], next_tokens: torch.Tensor, input_pos: torch.Tensor, input_pos_maxp1: int, return_hidden: bool = False, ) -> torch.Tensor | tuple[torch.Tensor, torch.Tensor]: """Advance EAGLE-3 using the cached pretrained MTP module.""" if self.mtp_eagle3_feature_proj is None: raise RuntimeError("mtp_eagle3_forward_step requires mtp_eagle3=True.") if len(features) != 3: raise ValueError("EAGLE-3 requires exactly low-, mid-, and high-level features.") hidden = self.mtp_eagle3_feature_proj(torch.cat(features, dim=-1)) return self.mtp_forward_step( hidden, next_tokens, input_pos, input_pos_maxp1, return_hidden=return_hidden, ) def mtp_forward_step( self, hidden: torch.Tensor, next_tokens: torch.Tensor, input_pos: torch.Tensor, input_pos_maxp1: int, return_hidden: bool = False, ) -> torch.Tensor | tuple[torch.Tensor, torch.Tensor]: """Advance the one-depth DeepSeek MTP module with its own KV cache.""" if len(self.mtp_blocks) != 1: raise RuntimeError("mtp_forward_step requires exactly one MTP module.") if hidden.shape[:2] != next_tokens.shape or next_tokens.size(1) != input_pos.shape[-1]: raise ValueError("MTP hidden states, tokens, and positions must have matching sequence lengths.") if self.mask_cache is None: raise TypeError("Call set_kv_cache() before cached MTP decoding.") embedding = self.transformer.wte(next_tokens) x = self.mtp_eh_proj( torch.cat((self.mtp_hnorm(hidden), self.mtp_enorm(embedding)), dim=-1) ) cos = batched_index_select(self.cos, 0, input_pos) sin = batched_index_select(self.sin, 0, input_pos) if input_pos.dim() == 1: cos = cos.unsqueeze(0) sin = sin.unsqueeze(0) mask = batched_index_select(self.mask_cache, 2, input_pos) if mask.dim() > 4: mask = mask.view(*(mask.shape[0:1] + mask.shape[2:])) mask = mask[..., :input_pos_maxp1] for block in self.mtp_blocks: x = block( x, cos, sin, mask, input_pos, input_pos_maxp1, ) logits = self.lm_head(self._matched_final_norm(x)) return (logits, x) if return_hidden else logits @classmethod def from_name(cls, name: str, **kwargs: Any) -> Self: return cls(Config.from_name(name, **kwargs)) def rope_cache(self, device: torch.device | None = None) -> tuple[torch.Tensor, torch.Tensor]: if self.config.rope_adjustments is None: extra_config = None else: # Check for mutually exclusive parameter sets llama3_params = ["low_freq_factor", "high_freq_factor"] yarn_params = ["beta_fast", "beta_slow"] has_llama3 = any(param in self.config.rope_adjustments for param in llama3_params) has_yarn = any(param in self.config.rope_adjustments for param in yarn_params) if has_llama3 and has_yarn: raise ValueError( "RoPE adjustments cannot contain both Llama3 parameters (low_freq_factor, high_freq_factor) " "and YaRN parameters (beta_fast, beta_slow). These are mutually exclusive." ) # Llama3-style RoPE if has_llama3: adjusted_params_required = ["factor", "low_freq_factor", "high_freq_factor", "original_max_seq_len"] params_present = [param in self.config.rope_adjustments for param in adjusted_params_required] if all(params_present): extra_config = {name: self.config.rope_adjustments[name] for name in adjusted_params_required} else: missing_params = [ param for param, present in zip(adjusted_params_required, params_present) if not present ] raise ValueError( f"The following Llama3 RoPE parameters are missing in rope_adjustments: {', '.join(missing_params)}. " "All Llama3 parameters must be specified together." ) # YaRN-style RoPE elif has_yarn: # Required: factor, beta_fast, beta_slow, original_max_seq_len # Optional: mscale, mscale_all_dim yarn_required_params = ["factor", "beta_fast", "beta_slow", "original_max_seq_len"] params_present = [param in self.config.rope_adjustments for param in yarn_required_params] if not all(params_present): missing_params = [ param for param, present in zip(yarn_required_params, params_present) if not present ] raise ValueError( f"The following YaRN RoPE parameters are missing in rope_adjustments: {', '.join(missing_params)}. " "All YaRN required parameters must be specified together." ) extra_config = {name: self.config.rope_adjustments[name] for name in yarn_required_params} # Add optional YaRN parameters for param in ["mscale", "mscale_all_dim"]: if param in self.config.rope_adjustments: extra_config[param] = self.config.rope_adjustments[param] # Linear or standard RoPE elif "factor" in self.config.rope_adjustments: # linear RoPE adjusted_params_required = ["factor"] extra_config = {name: self.config.rope_adjustments[name] for name in adjusted_params_required} else: extra_config = None # uses standard RoPE return build_rope_cache( seq_len=self.max_seq_length, n_elem=self.config.rope_n_elem, device=device, condense_ratio=self.config.rope_condense_ratio, base=self.config.rope_base, extra_config=extra_config, rope_local_base_freq=self.config.rope_local_base_freq, ) def rope_cache_length(self) -> int: """ Extract the head dimension (n_elem) from RoPE cache regardless of shape. The RoPE cache can have different shapes depending on model configuration: - Standard RoPE: (seq_len, n_elem) - 2D tensor - Dual RoPE (local/global): (seq_len, n_elem, 2) - 3D tensor Returns: int: n_elem (head dimension for RoPE) """ return self.cos.size(1) def set_kv_cache( self, batch_size: int, max_seq_length: int | None = None, rope_cache_length: int | None = None, device: torch.device | None = None, dtype: torch.dtype | None = None, ) -> None: if rope_cache_length is None: rope_cache_length = self.rope_cache_length() if max_seq_length is None: max_seq_length = self.max_seq_length # initialize the kv cache for all blocks for block in self.transformer.h: if block.attn is None: continue block.attn.kv_cache = block.attn.build_kv_cache( batch_size, max_seq_length, rope_cache_length, device, dtype, ) for block in self.mtp_blocks: block.attn.kv_cache = block.attn.build_kv_cache( batch_size, max_seq_length, rope_cache_length, device, dtype, ) if self.mask_cache is None or self.mask_cache.size(3) != max_seq_length: # passing `attn_mask` to SDPA disables the flash implementation. since we only need the mask # for the kv-cache support (only during inference), we only create it in that situation self.mask_cache = build_mask_cache(max_seq_length, device) def clear_kv_cache(self) -> None: self.mask_cache = None for block in self.transformer.h: if block.attn is not None: block.attn.kv_cache = None for block in self.mtp_blocks: block.attn.kv_cache = None class AdaptiveDepthMemory(nn.Module): """Bounded, token-adaptive retrieval from earlier residual states. Retrieval weights are content-addressed independently for every token. A zero-initialized channel gate makes the adaptive path exactly identical to the control at initialization while allowing a live first-step gradient. Earlier states are RMS-matched to the current stream before interpolation, preventing depth-dependent residual scale from dominating retrieval. """ def __init__(self, config: Config) -> None: super().__init__() self.max_sources = config.depth_memory_max_sources self.max_commit = config.depth_memory_max_commit self.score_scale = config.depth_memory_score_scale # Direct zeros consume no RNG, preserving every shared parent tensor. self.channel_gate = nn.Parameter(torch.zeros(config.n_embd)) self.source_bias = nn.Parameter(torch.zeros(self.max_sources)) self.log_temperature = nn.Parameter(torch.zeros(())) def forward( self, x: torch.Tensor, states: list[torch.Tensor], scale: torch.Tensor | float = 1.0, ) -> torch.Tensor: sources = torch.stack(states[-self.max_sources :], dim=-2) x_float = x.float() sources_float = sources.float() source_count = sources.size(-2) bias = self.source_bias[-source_count:].float() if self.score_scale in {"factorized", "recurrent"}: # Factorized ADeM: global depth selection plus a local # token/channel agreement modulation. This removes the costly # per-token source softmax while retaining input-dependent gates. if self.score_scale == "recurrent": retrieved = sources_float[..., -1, :] else: weights = torch.softmax(bias, dim=-1) retrieved = (weights.unsqueeze(-1) * sources_float).sum(dim=-2) current_rms = x_float.pow(2).mean(dim=-1, keepdim=True).sqrt().clamp_min(1e-6) retrieved_rms = retrieved.pow(2).mean(dim=-1, keepdim=True).sqrt().clamp_min(1e-6) retrieved = retrieved * (current_rms / retrieved_rms) correction = (retrieved - x_float).to(dtype=x.dtype) agreement = torch.tanh( x_float * retrieved / (current_rms * current_rms).clamp_min(1e-6) ).to(dtype=x.dtype) modulation = 1.0 + 0.5 * torch.tanh(self.log_temperature).to(dtype=x.dtype) * agreement gate = self.max_commit * torch.tanh(self.channel_gate).to(dtype=x.dtype) * modulation return x + torch.as_tensor(scale, device=x.device, dtype=x.dtype) * gate * correction query_unit = F.normalize(x_float, dim=-1, eps=1e-6) source_unit = F.normalize(sources_float, dim=-1, eps=1e-6) scores = (source_unit * query_unit.unsqueeze(-2)).sum(dim=-1) if self.score_scale == "inverse_sqrt": # Retain exact ADeM-v1 checkpoint semantics. ADeM-v2 uses the # cosine directly because both operands are already unit norm. scores = scores / math.sqrt(x.size(-1)) temperature = F.softplus(self.log_temperature.float()) + 0.5 weights = torch.softmax(scores * temperature + bias, dim=-1) retrieved = (weights.unsqueeze(-1) * sources_float).sum(dim=-2) current_rms = x_float.pow(2).mean(dim=-1, keepdim=True).sqrt().clamp_min(1e-6) retrieved_rms = retrieved.pow(2).mean(dim=-1, keepdim=True).sqrt().clamp_min(1e-6) retrieved = retrieved * (current_rms / retrieved_rms) correction = (retrieved - x_float).to(dtype=x.dtype) gate = self.max_commit * torch.tanh(self.channel_gate).to(dtype=x.dtype) return x + torch.as_tensor(scale, device=x.device, dtype=x.dtype) * gate * correction class _PackedMultiscreenProjection(torch.autograd.Function): """One GEMM over four original parameters, with gradients returned separately.""" @staticmethod def forward(ctx, x, packed_weight, *weights): ctx.save_for_backward(x, packed_weight) ctx.shapes = tuple(weight.shape for weight in weights) return F.linear(x, packed_weight) @staticmethod def backward(ctx, grad_output): x, packed_weight = ctx.saved_tensors x_flat = x.reshape(-1, x.size(-1)) grad_flat = grad_output.reshape(-1, grad_output.size(-1)) grad_x = torch.mm(grad_flat, packed_weight).reshape_as(x) grad_packed = torch.mm(grad_flat.transpose(0, 1), x_flat) gradients = [] offset = 0 for heads, hidden, features in ctx.shapes: width = heads * features gradient = ( grad_packed[offset : offset + width] .view(heads, features, hidden) .permute(0, 2, 1) .contiguous() ) gradients.append(gradient) offset += width return grad_x, None, *gradients class ExactMultiscreenResidualAdapter(nn.Module): """Compiler-fusible active calibration for a small exact budget gap.""" def __init__(self, width: int, parameter_budget: int) -> None: super().__init__() scale_size = min(width, parameter_budget) bias_size = parameter_budget - scale_size self.width = width self.scale = nn.Parameter(torch.zeros(scale_size)) self.bias = nn.Parameter(torch.zeros(bias_size)) if sum(parameter.numel() for parameter in self.parameters()) != parameter_budget: raise AssertionError("Multiscreen adapter did not consume its exact budget.") def forward(self, x: torch.Tensor) -> torch.Tensor: scale = F.pad(self.scale, (0, self.width - self.scale.numel())) bias = F.pad(self.bias, (0, self.width - self.bias.numel())) return x * (0.01 * torch.tanh(scale)) + bias class BudgetedMultiscreenResidualAdapter(nn.Module): """Useful low-rank residual that consumes an arbitrary exact budget.""" def __init__(self, width: int, parameter_budget: int) -> None: super().__init__() rank, remainder = divmod(parameter_budget, 2 * width) if rank <= 0: raise ValueError("Budgeted adapter requires at least 2 * width parameters.") self.width = width self.up = nn.Parameter(torch.empty(width, rank)) self.down = nn.Parameter(torch.zeros(rank, width)) scale_size = min(width, remainder) self.scale = nn.Parameter(torch.zeros(scale_size)) self.bias = nn.Parameter(torch.zeros(remainder - scale_size)) nn.init.orthogonal_(self.up) if sum(parameter.numel() for parameter in self.parameters()) != parameter_budget: raise AssertionError("Budgeted adapter did not consume its exact budget.") def forward(self, x: torch.Tensor) -> torch.Tensor: low_rank = F.silu(x @ self.up) @ self.down scale = F.pad(self.scale, (0, self.width - self.scale.numel())) bias = F.pad(self.bias, (0, self.width - self.bias.numel())) return low_rank + 0.01 * x * torch.tanh(scale) + bias class MultiscreenBlock(nn.Module): """A paper-faithful gated screening layer from arXiv:2604.01178. The query loop is only a memory-bounding implementation detail. It computes the same dense causal screening equation as the unchunked reference. """ def __init__(self, config: Config, block_idx: int) -> None: super().__init__() d = config.n_embd h = config.multiscreen_num_heads dk = config.multiscreen_key_dim dv = config.multiscreen_value_dim self.num_heads = h self.num_layers = config.n_layer self.key_dim = dk self.value_dim = dv self.window_threshold = float(config.multiscreen_window_threshold) self.query_chunk_size = config.multiscreen_query_chunk_size self.landmark_stride = config.multiscreen_landmark_stride self.triton_inference = config.multiscreen_triton_inference self.fused_projections = config.multiscreen_fused_projections self.swe_context_length = config.multiscreen_swe_context_length self.hard_max_window = config.multiscreen_hard_max_window self.match_adapter = None if config.multiscreen_match_adapter_params and block_idx == config.n_layer - 1: adapter_budget = config.multiscreen_match_adapter_params self.match_adapter = ( ExactMultiscreenResidualAdapter(d, adapter_budget) if adapter_budget <= 2 * d else BudgetedMultiscreenResidualAdapter(d, adapter_budget) ) self.query_weight = nn.Parameter(torch.empty(h, d, dk)) self.key_weight = nn.Parameter(torch.empty(h, d, dk)) self.value_weight = nn.Parameter(torch.empty(h, d, dv)) self.gate_weight = nn.Parameter(torch.empty(h, d, dv)) self.output_weight = nn.Parameter(torch.empty(h, dv, d)) self.log_window_minus_one = nn.Parameter(torch.empty(h)) self.acceptance_logit = nn.Parameter(torch.zeros(h)) self.log_output_scale = nn.Parameter(torch.empty(h)) self.mlp_norm = ( config.norm_class(d, eps=config.norm_eps) if config.multiscreen_mlp_enabled else None ) self.mlp = config.mlp_class(config) if config.multiscreen_mlp_enabled else None self.reset_parameters() @torch.no_grad() def reset_parameters(self) -> None: nn.init.normal_(self.query_weight, std=0.1 / math.sqrt(self.key_dim)) nn.init.normal_(self.key_weight, std=0.1 / math.sqrt(self.key_dim)) nn.init.normal_(self.value_weight, std=0.1 / math.sqrt(self.value_dim)) nn.init.normal_(self.gate_weight, std=0.1) nn.init.normal_(self.output_weight, std=0.1 / math.sqrt(self.output_weight.shape[-1])) self.log_window_minus_one.copy_( torch.linspace(0.0, math.log(self.window_threshold), self.num_heads) ) self.acceptance_logit.zero_() self.log_output_scale.fill_(-0.5 * math.log(self.num_heads * self.num_layers)) @staticmethod def _mipe(x: torch.Tensor, angle: torch.Tensor) -> torch.Tensor: first, second = x[..., 0], x[..., 1] cos, sin = angle.cos(), angle.sin() rotated_pair = torch.stack( (first * cos - second * sin, first * sin + second * cos), dim=-1, ) # Avoid clone + two in-place CopySlices nodes. Those nodes retain the # full pre-rotation tensor for backward and become expensive at large # microbatches. return torch.cat((rotated_pair, x[..., 2:]), dim=-1) def forward( self, x: torch.Tensor, cos: torch.Tensor, sin: torch.Tensor, mask: torch.Tensor | None = None, input_pos: torch.Tensor | None = None, input_pos_maxp1: int | None = None, ) -> torch.Tensor: del cos, sin, mask, input_pos, input_pos_maxp1 _, sequence_length, _ = x.shape if self.fused_projections: from litgpt.multiscreen_projection_triton import direct_multiscreen_projection h, dk, dv = self.num_heads, self.key_dim, self.value_dim projected = direct_multiscreen_projection( x, self.query_weight, self.key_weight, self.value_weight, self.gate_weight, ) query, key, value, gate = projected.split( (h * dk, h * dk, h * dv, h * dv), dim=-1 ) query = query.view(*query.shape[:2], h, dk).transpose(1, 2) key = key.view(*key.shape[:2], h, dk).transpose(1, 2) value = value.view(*value.shape[:2], h, dv).transpose(1, 2) gate = gate.view(*gate.shape[:2], h, dv).transpose(1, 2) else: query = torch.einsum("btd,hdk->bhtk", x, self.query_weight) key = torch.einsum("btd,hdk->bhtk", x, self.key_weight) value = torch.einsum("btd,hdv->bhtv", x, self.value_weight) gate = torch.einsum("btd,hdv->bhtv", x, self.gate_weight) query = F.normalize(query, dim=-1, eps=1e-6) key = F.normalize(key, dim=-1, eps=1e-6) value = F.normalize(value, dim=-1, eps=1e-6) learned_window = self.log_window_minus_one.exp() + 1.0 if self.hard_max_window > 0: learned_window = learned_window.clamp_max(float(self.hard_max_window)) use_swe = not torch.is_grad_enabled() and self.swe_context_length > 0 if use_swe: # A negative value is the internal infinity sentinel understood # by both the reference and Triton screening implementations. window = torch.where( learned_window > self.swe_context_length, torch.full_like(learned_window, -1.0), learned_window, ) else: window = learned_window finite_window = window > 0.0 safe_window = torch.where(finite_window, window, torch.ones_like(window)) gamma = torch.where( finite_window & (window < self.window_threshold), 0.5 * (torch.cos(math.pi * safe_window / self.window_threshold) + 1.0), torch.zeros_like(window), ) positions = torch.arange(sequence_length, device=x.device, dtype=torch.float32) query_angle = (math.pi * gamma / safe_window)[:, None] * positions[None, :] query = self._mipe(query, query_angle[None].to(query.dtype)) key_positions_1d = positions[:: self.landmark_stride] key = key[:, :, :: self.landmark_stride] value = value[:, :, :: self.landmark_stride] key_angle = ( (math.pi * gamma / safe_window)[:, None] * key_positions_1d[None, :] ) key = self._mipe(key, key_angle[None].to(key.dtype)) acceptance_width = torch.sigmoid(self.acceptance_logit).clamp_min(1e-6) use_triton = ( self.triton_inference and query.is_cuda and query.dtype in {torch.bfloat16, torch.float32} and self.key_dim == 16 and self.value_dim in {32, 64, 128} and self.landmark_stride == 1 ) if use_triton: from litgpt.multiscreen_triton import screening_aggregate_triton aggregated = screening_aggregate_triton( query.contiguous(), key.contiguous(), value.contiguous(), acceptance_width, window ) else: chunks = [] key_positions = key_positions_1d[None, None, None, :] for start in range(0, sequence_length, self.query_chunk_size): end = min(start + self.query_chunk_size, sequence_length) similarity = torch.einsum("bhqd,bhkd->bhqk", query[:, :, start:end], key) relevance = F.relu( 1.0 - (1.0 - similarity.float()) / acceptance_width.float()[None, :, None, None] ).square() query_positions = positions[start:end][None, None, :, None] distance = key_positions - query_positions infinite = ~finite_window[None, :, None, None] valid = (distance <= 0.0) & ( infinite | (distance > -safe_window[None, :, None, None]) ) finite_softmask = 0.5 * ( torch.cos( math.pi * distance / safe_window[None, :, None, None] ) + 1.0 ) softmask = torch.where(infinite, torch.ones_like(finite_softmask), finite_softmask) relevance = relevance * torch.where(valid, softmask, torch.zeros_like(softmask)) chunks.append( torch.einsum( "bhqk,bhkv->bhqv", relevance.to(value.dtype), value, ) ) aggregated = torch.cat(chunks, dim=2) norm = aggregated.float().norm(dim=-1, keepdim=True) aggregated = aggregated * (norm.tanh() / norm.clamp_min(1e-6)).to(aggregated.dtype) gated = aggregated * torch.tanh(F.silu(gate)) gated = gated * self.log_output_scale.exp()[None, :, None, None] update = torch.einsum("bhtv,hvd->btd", gated, self.output_weight) x = x + update if self.mlp is not None: x = x + self.mlp(self.mlp_norm(x)) if self.match_adapter is not None: x = x + self.match_adapter(x) return x def forward_step( self, x: torch.Tensor, state: tuple[torch.Tensor, torch.Tensor, int] | None = None, ) -> tuple[torch.Tensor, tuple[torch.Tensor, torch.Tensor, int]]: """Decode one token while retaining only the largest learned finite window.""" if x.ndim != 3 or x.size(1) != 1: raise ValueError("MultiscreenBlock.forward_step expects shape (batch, 1, width).") if self.swe_context_length > 0: raise NotImplementedError( "Bounded cached decode does not support screening-window expansion." ) if self.fused_projections: from litgpt.multiscreen_projection_triton import direct_multiscreen_projection h, dk, dv = self.num_heads, self.key_dim, self.value_dim projected = direct_multiscreen_projection( x, self.query_weight, self.key_weight, self.value_weight, self.gate_weight, ) query, key, value, gate = projected.split( (h * dk, h * dk, h * dv, h * dv), dim=-1 ) query = query.view(*query.shape[:2], h, dk).transpose(1, 2) key = key.view(*key.shape[:2], h, dk).transpose(1, 2) value = value.view(*value.shape[:2], h, dv).transpose(1, 2) gate = gate.view(*gate.shape[:2], h, dv).transpose(1, 2) else: query = torch.einsum("btd,hdk->bhtk", x, self.query_weight) key = torch.einsum("btd,hdk->bhtk", x, self.key_weight) value = torch.einsum("btd,hdv->bhtv", x, self.value_weight) gate = torch.einsum("btd,hdv->bhtv", x, self.gate_weight) query = F.normalize(query, dim=-1, eps=1e-6) key = F.normalize(key, dim=-1, eps=1e-6) value = F.normalize(value, dim=-1, eps=1e-6) learned_window = self.log_window_minus_one.exp() + 1.0 if self.hard_max_window > 0: learned_window = learned_window.clamp_max(float(self.hard_max_window)) safe_window = learned_window.clamp_min(1.0) gamma = torch.where( learned_window < self.window_threshold, 0.5 * (torch.cos(math.pi * safe_window / self.window_threshold) + 1.0), torch.zeros_like(learned_window), ) position = 0 if state is None else state[2] angle = (math.pi * gamma / safe_window * position)[None, :, None].to(query.dtype) query = self._mipe(query, angle) key = self._mipe(key, angle) if state is not None: key = torch.cat((state[0], key), dim=2) value = torch.cat((state[1], value), dim=2) cache_tokens = max(1, int(torch.ceil(safe_window.max()).item())) if key.size(2) > cache_tokens: key = key[:, :, -cache_tokens:] value = value[:, :, -cache_tokens:] history = key.size(2) distance = torch.arange( 1 - history, 1, device=x.device, dtype=torch.float32, )[None, None, None, :] similarity = torch.einsum("bhqd,bhkd->bhqk", query, key) acceptance_width = torch.sigmoid(self.acceptance_logit).clamp_min(1e-6) relevance = F.relu( 1.0 - (1.0 - similarity.float()) / acceptance_width.float()[None, :, None, None] ).square() valid = distance > -safe_window[None, :, None, None] softmask = 0.5 * ( torch.cos(math.pi * distance / safe_window[None, :, None, None]) + 1.0 ) relevance = relevance * torch.where(valid, softmask, torch.zeros_like(softmask)) aggregated = torch.einsum("bhqk,bhkv->bhqv", relevance, value) norm = aggregated.float().norm(dim=-1, keepdim=True) aggregated = aggregated * (norm.tanh() / norm.clamp_min(1e-6)).to(aggregated.dtype) gated = aggregated * torch.tanh(F.silu(gate)) gated = gated * self.log_output_scale.exp()[None, :, None, None] update = torch.einsum("bhtv,hvd->btd", gated, self.output_weight) x = x + update if self.mlp is not None: x = x + self.mlp(self.mlp_norm(x)) if self.match_adapter is not None: x = x + self.match_adapter(x) return x, (key, value, position + 1) class KimiDeltaAttentionAdapter(nn.Module): """LitGPT adapter around Moonshot/FLA's official KDA implementation.""" def __init__(self, config: Config, block_idx: int) -> None: super().__init__() try: from fla.layers import KimiDeltaAttention except ImportError as error: raise ImportError( "KDA requires the official `fla-core`/`flash-linear-attention` package." ) from error self.attn = KimiDeltaAttention( hidden_size=config.n_embd, expand_v=1, head_dim=config.kda_head_dim, num_heads=config.kda_num_heads, num_v_heads=config.kda_num_heads, mode="chunk", use_short_conv=True, conv_size=config.kda_short_conv_kernel_size, conv_bias=False, safe_gate=config.kda_safe_gate, lower_bound=config.kda_lower_bound if config.kda_safe_gate else None, layer_idx=block_idx, norm_eps=config.norm_eps, ) if config.kda_full_rank_output_gate: # Kimi K3 Eq. 6: y = Wo[sigmoid(Wg x) * RMSNorm(o)]. # The current FLA layer defaults to Kimi Linear's low-rank Wg; # replacing it with one bias-free projection preserves the # official FLA KDA kernel while matching K3's parameterization. self.attn.g_proj = nn.Linear( config.n_embd, self.attn.value_dim, bias=False, ) self.kv_cache = None def build_kv_cache( self, batch_size: int, max_seq_length: int, rope_cache_length: int, device: torch.device | None, dtype: torch.dtype | None, ): """Create the official FLA recurrent cache used by KDA decoding.""" del batch_size, max_seq_length, rope_cache_length, device, dtype from fla.models.utils import Cache return Cache() def forward( self, x: torch.Tensor, cos: torch.Tensor, sin: torch.Tensor, mask: torch.Tensor | None = None, input_pos: torch.Tensor | None = None, input_pos_maxp1: int | None = None, ) -> torch.Tensor: use_cache = input_pos is not None output, _, cache = self.attn( hidden_states=x, attention_mask=None, past_key_values=self.kv_cache if use_cache else None, use_cache=use_cache, output_attentions=False, ) if use_cache: self.kv_cache = cache return output class Block(nn.Module): def __init__( self, config: Config, block_idx: int, ) -> None: super().__init__() if not config.parallel_residual and config.shared_attention_norm: raise NotImplementedError( "No checkpoint amongst the ones we support uses this configuration" " (non-parallel residual and shared attention norm)." ) self.norm_1 = nn.Identity() if not config.norm_1 else config.norm_class(config.n_embd, eps=config.norm_eps) self.attn = build_attention(config, block_idx) self.post_attention_norm = ( config.norm_class(config.n_embd, eps=config.norm_eps) if config.post_attention_norm else nn.Identity() ) self.norm_2 = ( nn.Identity() if not config.norm_2 else (None if config.shared_attention_norm else config.norm_class(config.n_embd, eps=config.norm_eps)) ) mlp_config = config if config.mlp_taper_schedule != "none": mlp_config = copy(config) mlp_config.intermediate_size = config.tapered_mlp_intermediate_sizes()[block_idx] self.mlp = mlp_config.mlp_class(mlp_config) self.mlp_intermediate_size = mlp_config.intermediate_size if config.first_k_dense_replace is not None and block_idx < config.first_k_dense_replace: self.mlp = LLaMAMLP(config) if hasattr(self.mlp, "set_block_index"): self.mlp.set_block_index(block_idx, config.n_layer) self.post_mlp_norm = ( config.norm_class(config.n_embd, eps=config.norm_eps) if config.post_mlp_norm else nn.Identity() ) self.grouped_mlp_rotation_shift = 0 is_grouped_mlp = ( config.mlp_class_name.startswith("TileRouted") or config.mlp_class_name == "HiddenBlockDSwiGLUMLP" ) and not isinstance(self.mlp, LLaMAMLP) if is_grouped_mlp and config.grouped_mlp_channel_rotation: stride = config.grouped_mlp_channel_rotation_stride if stride is None: stride = max(1, config.n_embd // max(1, config.sparse_mlp_num_groups)) self.grouped_mlp_rotation_shift = (block_idx * stride) % config.n_embd self.grouped_mlp_mixer: nn.Module | None = None if ( is_grouped_mlp and config.grouped_mlp_mix_every_n_layers > 0 and (block_idx + 1) % config.grouped_mlp_mix_every_n_layers == 0 ): if config.grouped_mlp_mix_rank > 0: self.grouped_mlp_mixer = nn.Sequential( nn.Linear(config.n_embd, config.grouped_mlp_mix_rank, bias=False), nn.Linear(config.grouped_mlp_mix_rank, config.n_embd, bias=False), ) else: self.grouped_mlp_mixer = nn.Linear(config.n_embd, config.n_embd, bias=False) mix_alpha = torch.tensor(float(config.grouped_mlp_mix_alpha)) if config.grouped_mlp_mix_alpha_learnable: self.grouped_mlp_mix_alpha = nn.Parameter(mix_alpha) else: self.register_buffer("grouped_mlp_mix_alpha", mix_alpha, persistent=False) self.config = config def forward( self, x: torch.Tensor, cos: torch.Tensor, sin: torch.Tensor, mask: torch.Tensor | None = None, input_pos: torch.Tensor | None = None, input_pos_maxp1: int | None = None, value_residual: torch.Tensor | None = None, ) -> torch.Tensor | tuple[torch.Tensor, torch.Tensor]: """ Non-parallel residual Parallel residual ┌─ x ┌─ x ──────────────────┐ Note: if `shared_attention_norm` is True, │ ↓ │ ↓ ↓ the output from `norm_1` is reused │ norm_1 │ norm_1 ───────► norm_2 │ ↓ │ ↓ ↓ │ attn │ attn MLP │ ↓ │ ↓ ↓ | post_attn_norm | post_attn_norm post_mlp_norm | ↓ | ↓ ↓ ┌─ └► + └► + ◄─────────────────┘ | ↓ │ norm_2 │ ↓ │ MLP │ ↓ | post_mlp_norm | ↓ └───► + """ x_normed = self.norm_1(x) if self.config.value_residual_mix > 0.0: attention_output, current_values = self.attn( x_normed, cos, sin, mask, input_pos, input_pos_maxp1, value_residual=value_residual, ) else: attention_output = self.attn(x_normed, cos, sin, mask, input_pos, input_pos_maxp1) attention_output = self.post_attention_norm(attention_output) if self.config.parallel_residual: if not self.config.shared_attention_norm: x_normed = self.norm_2(x) x = attention_output + x else: x = attention_output + x x_normed = self.norm_2(x) mlp_input = x_normed if self.grouped_mlp_rotation_shift: mlp_input = torch.roll(mlp_input, shifts=self.grouped_mlp_rotation_shift, dims=-1) mlp_output = self.mlp(mlp_input) if self.grouped_mlp_rotation_shift: mlp_output = torch.roll(mlp_output, shifts=-self.grouped_mlp_rotation_shift, dims=-1) if self.grouped_mlp_mixer is not None: mix_alpha = self.grouped_mlp_mix_alpha.to(device=mlp_output.device, dtype=mlp_output.dtype) mlp_output = mlp_output + mix_alpha * self.grouped_mlp_mixer(mlp_output) output = self.post_mlp_norm(mlp_output) + x if self.config.value_residual_mix > 0.0: return output, current_values return output class SelectiveRecallBlock(nn.Module): """Fast Transformer block with a periodic causal landmark screen.""" def __init__(self, config: Config, block_idx: int) -> None: super().__init__() self.base = Block(config, block_idx) self.screen_enabled = ( (block_idx + 1) % config.multiscreen_layer_interval == 0 ) self.screen_norm = ( config.norm_class(config.n_embd, eps=config.norm_eps) if self.screen_enabled else None ) if self.screen_enabled: screen_config = copy(config) screen_config.multiscreen_mlp_enabled = False self.screen = MultiscreenBlock(screen_config, block_idx) else: self.screen = None def forward( self, x: torch.Tensor, cos: torch.Tensor, sin: torch.Tensor, mask: torch.Tensor | None = None, input_pos: torch.Tensor | None = None, input_pos_maxp1: int | None = None, value_residual: torch.Tensor | None = None, ) -> torch.Tensor | tuple[torch.Tensor, torch.Tensor]: base_output = self.base( x, cos, sin, mask, input_pos, input_pos_maxp1, value_residual, ) if isinstance(base_output, tuple): x, current_values = base_output else: x, current_values = base_output, None if self.screen is not None: normalized = self.screen_norm(x) x = x + ( self.screen( normalized, cos, sin, mask, input_pos, input_pos_maxp1, ) - normalized ) if current_values is not None: return x, current_values return x class SandwichSublayer(nn.Module): """One attention or FFN sublayer in the official Sandwich ordering. The source architecture is ``s^k (sf)^(L-k) f^k`` from https://github.com/ofirpress/sandwich_transformer. This class only adapts that ordering to LitGPT's pre-norm residual convention. """ def __init__(self, config: Config, sublayer_idx: int, kind: str) -> None: super().__init__() if kind not in {"s", "f"}: raise ValueError(f"Unknown Sandwich sublayer kind: {kind!r}") self.kind = kind self.config = config self.sublayer_idx = sublayer_idx logical_depth = min(sublayer_idx // 2, config.n_layer - 1) if kind == "s": # Sandwich starts with ``k`` consecutive attention sublayers, so # ``sublayer_idx // 2`` maps the first two attention modules to the # same layer. Use the true attention ordinal for per-layer # attention schedules (sliding windows, NoPE, and similar flags). # This is identical to ``logical_depth`` for the ordinary # interleaved architecture and changes no parameters. k = config.sandwich_coefficient attention_idx = sublayer_idx if sublayer_idx < k else k + (sublayer_idx - k) // 2 self.norm = config.norm_class(config.n_embd, eps=config.norm_eps) self.attn = build_attention(config, attention_idx) self.mlp = None self.post_norm = ( config.norm_class(config.n_embd, eps=config.norm_eps) if config.post_attention_norm else nn.Identity() ) else: self.norm = config.norm_class(config.n_embd, eps=config.norm_eps) self.attn = None self.mlp = config.mlp_class(config) if hasattr(self.mlp, "set_block_index"): self.mlp.set_block_index(logical_depth, config.n_layer) self.post_norm = ( config.norm_class(config.n_embd, eps=config.norm_eps) if config.post_mlp_norm else nn.Identity() ) adaptive_start = 2 * config.n_layer - config.sandwich_adaptive_commit if kind == "f" and sublayer_idx >= adaptive_start: # Direct zero initialization consumes no RNG and keeps every shared # parent parameter and logit bit-identical to the matched control. self.adaptive_commit_router = nn.Parameter(torch.zeros(config.n_embd)) else: self.register_parameter("adaptive_commit_router", None) def forward( self, x: torch.Tensor, cos: torch.Tensor, sin: torch.Tensor, mask: torch.Tensor | None = None, input_pos: torch.Tensor | None = None, input_pos_maxp1: int | None = None, value_residual: torch.Tensor | None = None, ) -> torch.Tensor: x_normed = self.norm(x) if self.kind == "s": update = self.attn(x_normed, cos, sin, mask, input_pos, input_pos_maxp1) else: update = self.mlp(x_normed) update = self.post_norm(update) if self.adaptive_commit_router is not None: router_logit = F.linear( x_normed.float(), self.adaptive_commit_router.float().unsqueeze(0), ) commit = (2.0 * torch.sigmoid(router_logit)).to(dtype=update.dtype) update = commit * update return x + update class RMSBudgetedBlock(Block): """Transformer block with a protected residual stream and bounded commits. Strong attention/MLP computation may happen in the branch state, but the write back to the main residual stream is clamped by an update/input RMS budget. The block is opt-in through ``Config.rms_budgeted_block`` so the proven dense and grouped baselines keep their exact execution path. """ def __init__(self, config: Config, block_idx: int) -> None: super().__init__(config, block_idx) self.block_idx = block_idx self.rms_budget_dual_stream = bool(config.rms_budget_dual_stream) self.rms_budget_attn_mlp = bool(config.rms_budget_attn_mlp) self.rms_budget_direction_aware = bool(config.rms_budget_direction_aware) self.rms_budget_work_scale = float(config.rms_budget_work_scale) self.rms_budget_min_scale = float(config.rms_budget_min_scale) self.rms_budget_exact_target = bool(config.rms_budget_exact_target) self.rms_budget_target_min = float(config.rms_budget_target_min) self.rms_budget_target_max = float(config.rms_budget_target_max) target = float(config.rms_budget_target_init) target = min(max(target, self.rms_budget_target_min + 1e-6), self.rms_budget_target_max - 1e-6) if config.rms_budget_target_learnable: ratio = (target - self.rms_budget_target_min) / (self.rms_budget_target_max - self.rms_budget_target_min) self.rms_budget_target_raw = nn.Parameter(torch.tensor(math.log(ratio / (1.0 - ratio)))) else: self.register_buffer("rms_budget_target", torch.tensor(target), persistent=False) if self.rms_budget_attn_mlp: self.budget_gate_weight = nn.Parameter(torch.zeros(2, config.n_embd)) self.budget_gate_bias = nn.Parameter(torch.zeros(2)) self.last_residual_ratio: torch.Tensor | None = None self.last_residual_scale: torch.Tensor | None = None self.last_attn_budget_mean: torch.Tensor | None = None self.last_mlp_budget_mean: torch.Tensor | None = None def _target_ratio(self, x: torch.Tensor) -> torch.Tensor: if hasattr(self, "rms_budget_target_raw"): raw = self.rms_budget_target_raw.to(device=x.device, dtype=torch.float32) target = self.rms_budget_target_min + (self.rms_budget_target_max - self.rms_budget_target_min) * torch.sigmoid(raw) return target.to(dtype=x.dtype) return self.rms_budget_target.to(device=x.device, dtype=x.dtype) def _direction_filter(self, x: torch.Tensor, update: torch.Tensor) -> torch.Tensor: if not self.rms_budget_direction_aware: return update x_float = x.float() update_float = update.float() denom = x_float.pow(2).sum(dim=-1, keepdim=True).clamp_min(1e-8) coeff = (update_float * x_float).sum(dim=-1, keepdim=True) / denom anti_coeff = coeff.clamp(max=0.0).to(dtype=update.dtype) return update - anti_coeff * x def _combine_updates(self, x: torch.Tensor, attention_output: torch.Tensor, mlp_output: torch.Tensor) -> torch.Tensor: if not self.rms_budget_attn_mlp: if not torch.is_grad_enabled(): one = x.new_tensor(1.0) self.last_attn_budget_mean = one self.last_mlp_budget_mean = one return attention_output + mlp_output x_rms = x.float().pow(2).mean(dim=-1, keepdim=True).sqrt().clamp_min(1e-8) x_normed = (x.float() / x_rms).to(dtype=x.dtype) logits = F.linear( x_normed, self.budget_gate_weight.to(device=x.device, dtype=x.dtype), self.budget_gate_bias.to(device=x.device, dtype=x.dtype), ) weights = torch.softmax(logits.float(), dim=-1).to(dtype=x.dtype) * 2.0 if not torch.is_grad_enabled(): self.last_attn_budget_mean = weights[..., 0].detach().mean() self.last_mlp_budget_mean = weights[..., 1].detach().mean() return attention_output * weights[..., 0:1] + mlp_output * weights[..., 1:2] def _budget_update(self, x: torch.Tensor, update: torch.Tensor) -> torch.Tensor: update = self._direction_filter(x, update) input_rms = x.float().pow(2).mean(dim=-1, keepdim=True).sqrt() update_rms = update.float().pow(2).mean(dim=-1, keepdim=True).sqrt() target = self._target_ratio(x).float() desired = target * input_rms scale = desired / update_rms.clamp_min(1e-8) if not self.rms_budget_exact_target: scale = scale.clamp(max=1.0) if self.rms_budget_min_scale > 0.0: scale = scale.clamp(min=self.rms_budget_min_scale) if not torch.is_grad_enabled(): ratio = update_rms / input_rms.clamp_min(1e-8) self.last_residual_ratio = ratio.detach().mean() self.last_residual_scale = scale.detach().mean() return update * scale.to(device=update.device, dtype=update.dtype) def _mlp_forward_with_hooks(self, mlp_input: torch.Tensor) -> torch.Tensor: if self.grouped_mlp_rotation_shift: mlp_input = torch.roll(mlp_input, shifts=self.grouped_mlp_rotation_shift, dims=-1) mlp_output = self.mlp(mlp_input) if self.grouped_mlp_rotation_shift: mlp_output = torch.roll(mlp_output, shifts=-self.grouped_mlp_rotation_shift, dims=-1) if self.grouped_mlp_mixer is not None: mix_alpha = self.grouped_mlp_mix_alpha.to(device=mlp_output.device, dtype=mlp_output.dtype) mlp_output = mlp_output + mix_alpha * self.grouped_mlp_mixer(mlp_output) return self.post_mlp_norm(mlp_output) def forward( self, x: torch.Tensor, cos: torch.Tensor, sin: torch.Tensor, mask: torch.Tensor | None = None, input_pos: torch.Tensor | None = None, input_pos_maxp1: int | None = None, ) -> torch.Tensor: x_lm = x x_normed = self.norm_1(x_lm) attention_output = self.attn(x_normed, cos, sin, mask, input_pos, input_pos_maxp1) attention_output = self.post_attention_norm(attention_output) if self.config.parallel_residual: if self.rms_budget_dual_stream: mlp_source = x_lm + self.rms_budget_work_scale * attention_output else: mlp_source = x_lm mlp_input = x_normed if self.config.shared_attention_norm else self.norm_2(mlp_source) else: work_scale = self.rms_budget_work_scale if self.rms_budget_dual_stream else 1.0 mlp_input = self.norm_2(x_lm + work_scale * attention_output) mlp_output = self._mlp_forward_with_hooks(mlp_input) update = self._combine_updates(x_lm, attention_output, mlp_output) return x_lm + self._budget_update(x_lm, update) class _CausalLinearMemory(torch.autograd.Function): """Normalized causal linear memory with an explicit reverse-scan backward.""" @staticmethod def forward( ctx, query: torch.Tensor, key: torch.Tensor, value: torch.Tensor, ) -> torch.Tensor: kv_state = torch.cumsum( key.unsqueeze(-1) * value.unsqueeze(-2), dim=1, ) key_state = torch.cumsum(key, dim=1) numerator = torch.einsum("btr,btrh->bth", query, kv_state) denominator_raw = torch.einsum("btr,btr->bt", query, key_state) denominator = denominator_raw.clamp_min(1e-6) ctx.save_for_backward( query, key, value, kv_state, key_state, numerator, denominator, denominator_raw, ) return numerator / denominator.unsqueeze(-1) @staticmethod def backward(ctx, grad_output: torch.Tensor): ( query, key, value, kv_state, key_state, numerator, denominator, denominator_raw, ) = ctx.saved_tensors grad_numerator = grad_output / denominator.unsqueeze(-1) grad_denominator = -( grad_output * numerator ).sum(dim=-1) / denominator.square() grad_denominator = grad_denominator * (denominator_raw > 1e-6) grad_query = torch.einsum( "btrh,bth->btr", kv_state, grad_numerator ) grad_query = grad_query + key_state * grad_denominator.unsqueeze(-1) state_contribution = ( query.unsqueeze(-1) * grad_numerator.unsqueeze(-2) ) grad_state = torch.flip( torch.cumsum(torch.flip(state_contribution, dims=(1,)), dim=1), dims=(1,), ) key_contribution = query * grad_denominator.unsqueeze(-1) grad_key_state = torch.flip( torch.cumsum(torch.flip(key_contribution, dims=(1,)), dim=1), dims=(1,), ) grad_key = torch.einsum( "btrh,bth->btr", grad_state, value ) + grad_key_state grad_value = torch.einsum("btrh,btr->bth", grad_state, key) return grad_query, grad_key, grad_value class CausalSelfAttention(nn.Module): def __init__(self, config: Config, block_idx: int) -> None: super().__init__() self.attention_output_gate = config.attention_output_gate if self.attention_output_gate == "headwise": self.gate_size = config.n_head elif self.attention_output_gate == "elementwise": self.gate_size = config.n_head * config.head_size else: self.gate_size = 0 # key, query and value projections for all heads, but in a batch self.qkv = nn.Linear( config.n_embd, (config.n_head + 2 * config.n_query_groups) * config.head_size + self.gate_size, bias=config.bias or config.attn_bias, ) # output projection self.proj = nn.Linear(config.head_size * config.n_head, config.n_embd, bias=config.bias) self.dual_path_linear_enabled = ( config.dual_path_linear_enabled and (block_idx + 1) % config.dual_path_linear_interval == 0 ) self.dual_path_linear_rank = config.dual_path_linear_rank self.dual_path_linear_explicit_backward = ( config.dual_path_linear_explicit_backward ) self.dual_path_gate = ( nn.Parameter( torch.full( (config.n_head,), float(config.dual_path_linear_gate_init), ) ) if self.dual_path_linear_enabled else None ) # disabled by default self.kv_cache: KVCache | MLACompressedKVCache | None = None self.apply_sliding_window_attention = False self.configured_sliding_window_attention = False self.sliding_window_block_mask: BlockMask | None = None if config.sliding_window_size is not None and config.sliding_window_indices is not None: self.configured_sliding_window_attention = bool(config.sliding_window_indices[block_idx]) self.apply_sliding_window_attention = self.configured_sliding_window_attention if config.norm_qk: norm_q_size = config.n_head * config.head_size if config.norm_qk_type == "olmo2" else config.head_size norm_k_size = ( config.n_query_groups * config.head_size if config.norm_qk_type == "olmo2" else config.head_size ) self.norm_q = config.norm_class(norm_q_size, eps=config.norm_eps) self.norm_k = config.norm_class(norm_k_size, eps=config.norm_eps) else: self.norm_q = self.norm_k = None if config.rope_adjustments is not None: mscale_all_dim = config.rope_adjustments.get("mscale_all_dim", None) scaling_factor = config.rope_adjustments.get("factor", None) if mscale_all_dim and scaling_factor: # YaRN self.mscale = yarn_get_mscale(scaling_factor, mscale_all_dim) else: self.mscale = 1.0 else: self.mscale = 1.0 self.config = config self.block_idx = block_idx def prepare_sliding_window_mask(self, sequence_length: int, device: torch.device) -> None: if not self.configured_sliding_window_attention: return window_size = int(self.config.sliding_window_size) def sliding_causal_mask( batch_index: torch.Tensor, head_index: torch.Tensor, query_index: torch.Tensor, key_index: torch.Tensor, ) -> torch.Tensor: del batch_index, head_index distance = query_index - key_index return (distance >= 0) & (distance < window_size) self.sliding_window_block_mask = create_block_mask( sliding_causal_mask, B=None, H=None, Q_LEN=sequence_length, KV_LEN=sequence_length, device=device, ) def forward( self, x: torch.Tensor, cos: torch.Tensor, sin: torch.Tensor, mask: torch.Tensor | None = None, input_pos: torch.Tensor | None = None, input_pos_maxp1: int | None = None, value_residual: torch.Tensor | None = None, ) -> torch.Tensor | tuple[torch.Tensor, torch.Tensor]: # Notation: # - B | batch size # - T | time-step (sequence length) # - C | model's embeddings size (n_embd) # - C* | attentions's embeddings size # - hs | head size # - nh_(q,k,v) | number of heads for query, key and value # - n_query_groups = nh_k = nh_v | number of query groups sharing key and value heads # alternative notation: num_kv_groups = n_query_groups # ┌───┐┌───┐┌───┐┌───┐ ┌───┐ ┌───┐ ┌───┐ # │ v ││ v ││ v ││ v │ │ v │ │ v │ │ v │ # └───┘└───┘└───┘└───┘ └───┘ └───┘ └───┘ # │ │ │ │ │ │ │ # ┌───┐┌───┐┌───┐┌───┐ ┌───┐ ┌───┐ ┌───┐ # │ k ││ k ││ k ││ k │ │ k │ │ k │ │ k │ # └───┘└───┘└───┘└───┘ └───┘ └───┘ └───┘ # │ │ │ │ ┌──┴──┐ ┌──┴──┐ ┌────┬──┴─┬────┐ # ┌───┐┌───┐┌───┐┌───┐ ┌───┐┌───┐┌───┐┌───┐ ┌───┐┌───┐┌───┐┌───┐ # │ q ││ q ││ q ││ q │ │ q ││ q ││ q ││ q │ │ q ││ q ││ q ││ q │ # └───┘└───┘└───┘└───┘ └───┘└───┘└───┘└───┘ └───┘└───┘└───┘└───┘ # ◀──────────────────▶ ◀──────────────────▶ ◀──────────────────▶ # MHA GQA MQA # n_query_groups=4 n_query_groups=2 n_query_groups=1 # # credit https://arxiv.org/pdf/2305.13245.pdf head_size = self.config.head_size n_head = self.config.n_head n_query_groups = self.config.n_query_groups rope_n_elem = self.config.rope_n_elem B, T, C = x.size() # batch size, sequence length, embedding dimensionality (n_embd) # Perform a single multiplication operation using a combined QKV matrix to calculate `query`, `key`, and `value` # instead of individually multiplying the input `x` with the respective weight matrices. qkv = self.qkv(x) # (B, T, 3xC*) # Define query, key and value sizes. # If grouped/multi query is enabled, these sizes are not equal (see the diagram above). query_size = n_head * head_size query_projection_size = query_size + self.gate_size key_size = value_size = n_query_groups * head_size # Gated Attention appends gate logits to each KV-group's query projection. q_projection, k, v = qkv.split((query_projection_size, key_size, value_size), dim=-1) gate_score: torch.Tensor | None = None if self.attention_output_gate != "none": queries_per_group = n_head // n_query_groups q_projection = q_projection.view(B, T, n_query_groups, -1) gate_features = ( queries_per_group if self.attention_output_gate == "headwise" else queries_per_group * head_size ) q, gate_score = q_projection.split((queries_per_group * head_size, gate_features), dim=-1) q = q.reshape(B, T, n_head, head_size) if self.attention_output_gate == "headwise": gate_score = gate_score.reshape(B, T, n_head, 1) else: gate_score = gate_score.reshape(B, T, n_head, head_size) else: q = q_projection current_values = v if value_residual is not None: mix = self.config.value_residual_mix v = mix * value_residual + (1.0 - mix) * v if self.config.norm_qk and self.config.norm_qk_type == "olmo2": q = self.norm_q(q) k = self.norm_k(k) # To place the num_heads (nh) dimension right after the batch (B) dimension, the first step is to decouple the # embedding size (C) into num_heads (nh) and head_size (hs). # The original GQA paper is followed here and the term query groups is used. # alternative notation: Query groups are also referred to as KV groups. if self.attention_output_gate == "none": q = q.view(B, T, n_head, head_size) # (B, T, nh_q, hs) k = k.view(B, T, n_query_groups, head_size) # (B, T, n_query_groups, hs) v = v.view(B, T, n_query_groups, head_size) # (B, T, n_query_groups, hs) # The tensors `query`, `key`, and `value` are now accurately structured: within each batch element (B), there are # multiple heads (nh), and within each head, there is a sequence of elements (T), each represented by a vector # of size `hs`. q = q.transpose(1, 2) # (B, nh_q, T, hs) k = k.transpose(1, 2) # (B, nh_k, T, hs) v = v.transpose(1, 2) # (B, nh_v, T, hs) v_current = v if self.config.norm_qk and self.config.norm_qk_type == "default": q = self.norm_q(q) k = self.norm_k(k) # SmolLM3 uses a 3:1 RoPE/NoPE pattern by skipping every fourth layer. use_rope = ( rope_n_elem > 0 and ( self.config.no_rope_layer_interval == 0 or (self.block_idx + 1) % self.config.no_rope_layer_interval != 0 ) ) if use_rope: if self.config.rope_interleave: q_roped = apply_rope_interleave(q[..., :rope_n_elem], cos, sin) k_roped = apply_rope_interleave(k[..., :rope_n_elem], cos, sin) else: q_roped = apply_rope(q[..., :rope_n_elem], cos, sin) k_roped = apply_rope(k[..., :rope_n_elem], cos, sin) q = torch.cat((q_roped, q[..., rope_n_elem:]), dim=-1) # (B, nh_q, T, hs) k = torch.cat((k_roped, k[..., rope_n_elem:]), dim=-1) # (B, nh_k, T, hs) # Apply kv-cache during inference. if input_pos is not None: if not isinstance(self.kv_cache, KVCache): raise TypeError("You need to call `gpt.set_kv_cache()`") k, v = self.kv_cache(input_pos, k, v) if input_pos_maxp1 is not None: # Subselect along sequence dimension k = k[..., :input_pos_maxp1, :] v = v[..., :input_pos_maxp1, :] if self.apply_sliding_window_attention: # The cache keeps absolute positions so cached prefill and # decode can use the same semantics as training, including # prompts longer than one local window. A compact circular # cache needs chronological reordering and a position-aware # mask; until that fast path exists, keep correctness exact. key_position = torch.arange(k.size(2), device=input_pos.device) window_size = int(self.config.sliding_window_size) if input_pos.dim() == 1: distance = input_pos[:, None] - key_position[None, :] mask = ((distance >= 0) & (distance < window_size))[None, None, :, :] else: distance = input_pos[:, :, None] - key_position[None, None, :] mask = ((distance >= 0) & (distance < window_size))[:, None, :, :] # k, v: (B, nh_k, input_pos_maxp1, hs) # If input_pos_maxp1 is None -> max_seq_length use_flex_window = ( self.apply_sliding_window_attention and input_pos is None and self.sliding_window_block_mask is not None and self.sliding_window_block_mask.shape[-1] == T and self.config.attention_logit_softcapping is None ) # Grouped queries: balance the number of heads across all three matrices. # NOTE: flash attention requires it in training mode. # Multi-query: this step can be skipped since there is only 1 head, allowing us to use broadcasting. if n_query_groups != n_head and (input_pos is None or n_query_groups != 1) and not use_flex_window: q_per_kv = n_head // n_query_groups k = k.repeat_interleave(q_per_kv, dim=1) # (B, nh_q, T, hs) v = v.repeat_interleave(q_per_kv, dim=1) # (B, nh_q, T, hs) if self.apply_sliding_window_attention and not use_flex_window: """ Global Window Sliding window Sliding window attention mask + bias = attention mask ┌────────────────────────┐ ┌───────────────────────┐ ┌─────────────────────────┐ │ True False False False │ │ True True True True │ │ True False False False │ │ True True False False │ │ True True True True │ │ True True False False │ │ True True True False │ │ False True True True │ │ False True True False │ │ True True True True │ │ False False True True │ │ False False True True │ └────────────────────────┘ └───────────────────────┘ └─────────────────────────┘ """ if input_pos is None: if mask is None: mask = torch.ones(T, T, dtype=q.dtype, device=q.device).triu(diagonal=1) mask.masked_fill_(mask.bool(), float("-inf")) mask = mask.view(1, 1, *mask.shape) sliding_window_mask = torch.full((T, T), float("-inf"), dtype=q.dtype, device=q.device) for i in range(T): window_start = max(0, i - self.config.sliding_window_size + 1) sliding_window_mask[i, window_start : i + 1] = 0.0 sliding_window_mask = sliding_window_mask.view(1, 1, T, T) mask = sliding_window_mask # Efficient attention using Flash Attention CUDA kernels. # NOTE: efficient implementation is disabled if `mask` is not None or softcapping is enabled. # ↓ (B, nh, T, hs) @ (B, nh, T, hs).mT --> (B, nh, T, T) @ (B, nh, T, hs) --> (B, nh, T, hs) if use_flex_window: scale = 1.0 / math.sqrt(self.config.attention_scores_scalar or self.config.head_size) scale = scale * self.mscale * self.mscale y = flex_attention( q, k, v, block_mask=self.sliding_window_block_mask, scale=scale, enable_gqa=n_query_groups != n_head, ).transpose(1, 2) else: y = self.scaled_dot_product_attention(q, k, v, mask) if self.dual_path_linear_enabled: # Reuse the already projected MQA features. Positive feature maps # produce a normalized causal linear-attention recurrence: # S_t = S_{t-1} + phi(k_t) outer v_t. rank = self.dual_path_linear_rank q_feature = F.elu(q[..., :rank].mean(dim=1).float()) + 1.0 k_feature = F.elu(k[:, 0, :, :rank].float()) + 1.0 value_feature = v_current[:, 0].float() if self.dual_path_linear_explicit_backward: global_value = _CausalLinearMemory.apply( q_feature, k_feature, value_feature, ) else: kv_state = torch.cumsum( k_feature.unsqueeze(-1) * value_feature.unsqueeze(-2), dim=1, ) key_state = torch.cumsum(k_feature, dim=1) numerator = torch.einsum( "btr,btrh->bth", q_feature, kv_state ) denominator = torch.einsum( "btr,btr->bt", q_feature, key_state ) global_value = numerator / denominator.clamp_min( 1e-6 ).unsqueeze(-1) global_value = global_value.to(dtype=y.dtype) gate = torch.tanh(self.dual_path_gate).to(dtype=y.dtype) y = y + global_value[:, :, None, :] * gate[None, None, :, None] if gate_score is not None: y = y * torch.sigmoid(gate_score) if self.config.xsa_projection: if v.size(2) == T: v_projection_base = v.transpose(1, 2) else: q_per_kv = n_head // n_query_groups v_projection_base = v_current.repeat_interleave(q_per_kv, dim=1).transpose(1, 2) v_projection_base = F.normalize(v_projection_base, dim=-1) y = y - (y * v_projection_base).sum(dim=-1, keepdim=True) * v_projection_base # Re-assemble all head outputs side by side. y = y.reshape(B, T, head_size * n_head) # Output projection. output = self.proj(y) # (B, T, C) if self.config.value_residual_mix > 0.0: return output, current_values return output def scaled_dot_product_attention( self, q: torch.Tensor, k: torch.Tensor, v: torch.Tensor, mask: torch.Tensor | None = None ) -> torch.Tensor: scale = 1.0 / math.sqrt(self.config.attention_scores_scalar or self.config.head_size) scale = scale * self.mscale * self.mscale # with softcapping we cannot use SDPA if self.config.attention_logit_softcapping is not None: scores = q @ k.mT * scale scores = do_softcapping(scores, self.config.attention_logit_softcapping) if mask is None: mask = torch.ones(q.size(2), q.size(2), dtype=q.dtype, device=q.device).triu(diagonal=1) mask.masked_fill_(mask.bool(), torch.finfo(q.dtype).min) scores = scores + mask scores = F.softmax(scores, dim=-1, dtype=torch.float).to(dtype=q.dtype) y = scores @ v else: y = F.scaled_dot_product_attention( q, k, v, attn_mask=mask, dropout_p=0.0, scale=scale, is_causal=mask is None ) return y.transpose(1, 2) def build_kv_cache( self, batch_size: int, max_seq_length: int, rope_cache_length: int | None = None, device: torch.device | None = None, dtype: torch.dtype | None = None, ) -> "KVCache": # Keep absolute-position caches for sliding layers. This deliberately # gives up the potential cache-memory saving until circular-buffer # decode can preserve chronological order and exact causal masking. effective_cache_size = max_seq_length v_shape = (batch_size, self.config.n_query_groups, effective_cache_size, self.config.head_size) if rope_cache_length is None: if self.config.rotary_percentage != 1.0: raise TypeError( "Please pass the `rope_cache_length` parameter. " "Use `rope_cache_length=model.rope_cache_length()` to extract it automatically." ) k_shape = v_shape else: k_shape = ( batch_size, self.config.n_query_groups, effective_cache_size, rope_cache_length + self.config.head_size - self.config.rope_n_elem, ) return KVCache( k_shape, v_shape, device=device, dtype=dtype, is_sliding_window=self.apply_sliding_window_attention, sliding_window_size=self.config.sliding_window_size if self.apply_sliding_window_attention else None, ) def _load_from_state_dict(self, state_dict: dict, prefix: str, *args: Any, **kwargs: Any) -> None: """For compatibility with legacy checkpoints.""" for attr in ("weight", "bias"): legacy_key = f"{prefix}attn.{attr}" current_key = f"{prefix}qkv.{attr}" if legacy_key in state_dict: state_dict[current_key] = qkv_reassemble(state_dict.pop(legacy_key), self.config) super()._load_from_state_dict(state_dict, prefix, *args, **kwargs) class SharedDiffCausalSelfAttention(nn.Module): """Shared DIFF attention from arXiv:2501.17900. The differential attention and lambda parameterization follow Microsoft's reference DIFF implementation. Q and K use the Shared DIFF paper's shared base projection plus two trainable low-rank updates. """ def __init__(self, config: Config, block_idx: int) -> None: super().__init__() self.config = config self.block_idx = block_idx self.num_heads = config.n_head // 2 # Config head counts describe the matched baseline. The official DIFF # geometry halves both query and KV heads while preserving head_dim. self.num_kv_heads = config.n_query_groups // 2 self.head_dim = config.head_size self.q_size = self.num_heads * self.head_dim self.k_size = self.num_kv_heads * self.head_dim self.v_size = self.num_kv_heads * 2 * self.head_dim rank = config.shared_diff_rank self.q_base = nn.Linear(config.n_embd, self.q_size, bias=config.attn_bias) self.k_base = nn.Linear(config.n_embd, self.k_size, bias=config.attn_bias) self.v_proj = nn.Linear(config.n_embd, self.v_size, bias=config.attn_bias) self.proj = nn.Linear(config.n_embd, config.n_embd, bias=config.bias) self.q_lora_a = nn.Parameter(torch.empty(2, self.num_heads, config.n_embd, rank)) self.q_lora_b = nn.Parameter(torch.empty(2, self.num_heads, rank, self.head_dim)) self.k_lora_a = nn.Parameter(torch.empty(2, self.num_kv_heads, config.n_embd, rank)) self.k_lora_b = nn.Parameter(torch.empty(2, self.num_kv_heads, rank, self.head_dim)) for parameter in (self.q_lora_a, self.q_lora_b, self.k_lora_a, self.k_lora_b): nn.init.normal_(parameter, mean=0.0, std=0.02) self.lambda_init = 0.8 - 0.6 * math.exp(-0.3 * block_idx) self.lambda_q1 = nn.Parameter(torch.empty(self.head_dim).normal_(mean=0.0, std=0.1)) self.lambda_k1 = nn.Parameter(torch.empty(self.head_dim).normal_(mean=0.0, std=0.1)) self.lambda_q2 = nn.Parameter(torch.empty(self.head_dim).normal_(mean=0.0, std=0.1)) self.lambda_k2 = nn.Parameter(torch.empty(self.head_dim).normal_(mean=0.0, std=0.1)) self.subln = RMSNorm(2 * self.head_dim, eps=config.norm_eps) self.kv_cache: KVCache | None = None if config.rope_adjustments is not None: mscale_all_dim = config.rope_adjustments.get("mscale_all_dim") scaling_factor = config.rope_adjustments.get("factor") self.mscale = yarn_get_mscale(scaling_factor, mscale_all_dim) if mscale_all_dim and scaling_factor else 1.0 else: self.mscale = 1.0 @staticmethod def _low_rank_pair(x: torch.Tensor, a: torch.Tensor, b: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]: hidden = torch.einsum("btc,shcr->btshr", x, a) updates = torch.einsum("btshr,shrd->btshd", hidden, b) return updates[:, :, 0], updates[:, :, 1] def _apply_rope(self, tensor: torch.Tensor, cos: torch.Tensor, sin: torch.Tensor) -> torch.Tensor: rope_n_elem = self.config.rope_n_elem if self.config.rope_interleave: roped = apply_rope_interleave(tensor[..., :rope_n_elem], cos, sin) else: roped = apply_rope(tensor[..., :rope_n_elem], cos, sin) return torch.cat((roped, tensor[..., rope_n_elem:]), dim=-1) def forward( self, x: torch.Tensor, cos: torch.Tensor, sin: torch.Tensor, mask: torch.Tensor | None = None, input_pos: torch.Tensor | None = None, input_pos_maxp1: int | None = None, value_residual: torch.Tensor | None = None, ) -> torch.Tensor: if value_residual is not None: raise ValueError("Shared DIFF does not support value residuals.") batch_size, sequence_length, _ = x.shape q_delta1, q_delta2 = self._low_rank_pair(x, self.q_lora_a, self.q_lora_b) k_delta1, k_delta2 = self._low_rank_pair(x, self.k_lora_a, self.k_lora_b) q_base = self.q_base(x) k_base = self.k_base(x) q_base = q_base.view(batch_size, sequence_length, self.num_heads, self.head_dim) k_base = k_base.view(batch_size, sequence_length, self.num_kv_heads, self.head_dim) q1 = q_base + q_delta1 q2 = q_base + q_delta2 k1 = k_base + k_delta1 k2 = k_base + k_delta2 v = self.v_proj(x).view(batch_size, sequence_length, self.num_kv_heads, 2 * self.head_dim) q1 = self._apply_rope(q1.transpose(1, 2), cos, sin) q2 = self._apply_rope(q2.transpose(1, 2), cos, sin) k1 = self._apply_rope(k1.transpose(1, 2), cos, sin) k2 = self._apply_rope(k2.transpose(1, 2), cos, sin) v = v.transpose(1, 2) if input_pos is not None: if not isinstance(self.kv_cache, KVCache): raise TypeError("You need to call `gpt.set_kv_cache()`") cached_k, v = self.kv_cache(input_pos, torch.cat((k1, k2), dim=1), v) k1, k2 = cached_k.split(self.num_kv_heads, dim=1) if input_pos_maxp1 is not None: k1 = k1[..., :input_pos_maxp1, :] k2 = k2[..., :input_pos_maxp1, :] v = v[..., :input_pos_maxp1, :] if self.num_kv_heads != self.num_heads: repeats = self.num_heads // self.num_kv_heads k1 = k1.repeat_interleave(repeats, dim=1) k2 = k2.repeat_interleave(repeats, dim=1) v = v.repeat_interleave(repeats, dim=1) scale = self.mscale * self.mscale / math.sqrt(self.head_dim) is_causal = mask is None attn1 = F.scaled_dot_product_attention(q1, k1, v, attn_mask=mask, dropout_p=0.0, scale=scale, is_causal=is_causal) attn2 = F.scaled_dot_product_attention(q2, k2, v, attn_mask=mask, dropout_p=0.0, scale=scale, is_causal=is_causal) lambda1 = torch.exp(torch.sum(self.lambda_q1 * self.lambda_k1).float()).to(dtype=x.dtype) lambda2 = torch.exp(torch.sum(self.lambda_q2 * self.lambda_k2).float()).to(dtype=x.dtype) lambda_full = lambda1 - lambda2 + self.lambda_init output = self.subln(attn1 - lambda_full * attn2) * (1.0 - self.lambda_init) output = output.transpose(1, 2).reshape(batch_size, sequence_length, self.config.n_embd) return self.proj(output) def build_kv_cache( self, batch_size: int, max_seq_length: int, rope_cache_length: int | None = None, device: torch.device | None = None, dtype: torch.dtype | None = None, ) -> "KVCache": if rope_cache_length is None: if self.config.rotary_percentage != 1.0: raise TypeError("Please pass rope_cache_length when rotary_percentage != 1.0") key_dim = self.head_dim else: key_dim = rope_cache_length + self.head_dim - self.config.rope_n_elem k_shape = (batch_size, 2 * self.num_kv_heads, max_seq_length, key_dim) v_shape = (batch_size, self.num_kv_heads, max_seq_length, 2 * self.head_dim) return KVCache(k_shape, v_shape, device=device, dtype=dtype) class MultiheadLatentAttention(nn.Module): def __init__(self, config: Config, block_idx: int) -> None: super().__init__() self.q_a_proj = nn.Linear(config.n_embd, config.q_lora_rank, bias=config.attn_bias) self.q_a_norm = RMSNorm(config.q_lora_rank, eps=config.norm_eps) self.q_b_proj = nn.Linear(config.q_lora_rank, config.n_head * config.qk_head_dim, bias=config.bias) self.kv_a_proj_with_mqa = nn.Linear( config.n_embd, config.kv_lora_rank + config.qk_rope_head_dim, bias=config.attn_bias ) self.kv_a_norm = RMSNorm(config.kv_lora_rank, eps=config.norm_eps) self.kv_b_proj = nn.Linear( config.kv_lora_rank, config.n_query_groups * (config.qk_nope_head_dim + config.v_head_dim), bias=config.bias, ) # output projection self.proj = nn.Linear(config.n_head * config.v_head_dim, config.n_embd, bias=config.bias) self.output_gate = ( nn.Linear( config.n_embd, config.n_head * config.v_head_dim, bias=False, ) if config.mla_use_output_gate else None ) # disabled by default self.kv_cache: KVCache | None = None if config.rope_adjustments is not None: mscale_all_dim = config.rope_adjustments.get("mscale_all_dim", None) scaling_factor = config.rope_adjustments.get("factor", None) if mscale_all_dim and scaling_factor: # YaRN self.mscale = yarn_get_mscale(scaling_factor, mscale_all_dim) else: self.mscale = 1.0 else: self.mscale = 1.0 self.config = config self.block_idx = block_idx def forward( self, x: torch.Tensor, cos: torch.Tensor, sin: torch.Tensor, mask: torch.Tensor | None = None, input_pos: torch.Tensor | None = None, input_pos_maxp1: int | None = None, ) -> torch.Tensor: # Notation: # - B | batch size # - T | time-step (sequence length) # - C | model's embeddings size (n_embd) # - C* | attentions's embeddings size # - hs | head size # - nh_(q,k,v) | number of heads for query, key and value # - n_query_groups = nh_k = nh_v | number of query groups sharing key and value heads # alternative notation: num_kv_groups = n_query_groups B, T, C = x.size() # batch size, sequence length, embedding dimensionality (n_embd) q = self.q_b_proj(self.q_a_norm(self.q_a_proj(x))) # (B, T, n_head * qk_head_dim) q = q.view(B, T, -1, self.config.qk_head_dim) # (B, T, n_head, qk_head_dim) q = q.transpose(1, 2) # (B, n_head, T, qk_head_dim) q_pass, q_rot = torch.split(q, [self.config.qk_nope_head_dim, self.config.qk_rope_head_dim], dim=-1) compressed_kv = self.kv_a_proj_with_mqa(x) # (B, T, kv_lora_rank + qk_rope_head_dim) compressed_latent, k_rot = torch.split( compressed_kv, [self.config.kv_lora_rank, self.config.qk_rope_head_dim], dim=-1 ) compressed_latent = self.kv_a_norm(compressed_latent) k_rot = k_rot.view(B, 1, T, self.config.qk_rope_head_dim) if self.config.mla_use_nope: q_roped = q_rot k_roped = k_rot elif self.config.rope_interleave: q_roped = apply_rope_interleave(q_rot, cos, sin) k_roped = apply_rope_interleave(k_rot, cos, sin) else: q_roped = apply_rope(q_rot, cos, sin) k_roped = apply_rope(k_rot, cos, sin) if input_pos is not None and isinstance(self.kv_cache, MLACompressedKVCache): # The absorbed compressed-cache path is valid for both one-token # decode and batched speculative verification. Supporting T > 1 # avoids materializing expanded K/V and preserves prior context. return self.compressed_cache_attention( q_pass, q_roped, compressed_latent, k_roped, input_pos, input_pos_maxp1, mask, x, ) k_pass = self.kv_b_proj(compressed_latent) k_pass = k_pass.view(B, T, self.config.n_query_groups, -1) k_pass = k_pass.transpose(1, 2) k_pass, v = torch.split(k_pass, [self.config.qk_nope_head_dim, self.config.v_head_dim], dim=-1) k_roped = k_roped.expand(*k_pass.shape[:-1], -1) # (B, n_head, T, qk_rope_head_dim) q = torch.cat((q_pass, q_roped), dim=-1) k = torch.cat((k_pass, k_roped), dim=-1) # Apply kv-cache during inference. if input_pos is not None and isinstance(self.kv_cache, KVCache): if not isinstance(self.kv_cache, KVCache): raise TypeError("You need to call `gpt.set_kv_cache()`") k, v = self.kv_cache(input_pos, k, v) if input_pos_maxp1 is not None: # Subselect along sequence dimension k = k[..., :input_pos_maxp1, :] v = v[..., :input_pos_maxp1, :] # k, v: (B, nh_k, input_pos_maxp1, hs) # If input_pos_maxp1 is None -> max_seq_length # Grouped queries: balance the number of heads across all three matrices. # NOTE: flash attention requires it in training mode. # Multi-query: this step can be skipped since there is only 1 head, allowing us to use broadcasting. if self.config.n_query_groups != self.config.n_head and ( input_pos is None or self.config.n_query_groups != 1 or T > 1 ): q_per_kv = self.config.n_head // self.config.n_query_groups k = k.repeat_interleave(q_per_kv, dim=1) # (B, nh_q, T, hs) v = v.repeat_interleave(q_per_kv, dim=1) # (B, nh_q, T, hs) # Efficient attention using Flash Attention CUDA kernels. # NOTE: efficient implementation is disabled if `mask` is not None or softcapping is enabled. # ↓ (B, nh, T, hs) @ (B, nh, T, hs).mT --> (B, nh, T, T) @ (B, nh, T, hs) --> (B, nh, T, hs) y = self.scaled_dot_product_attention(q, k, v, mask) # Re-assemble all head outputs side by side. y = y.reshape(B, T, self.config.n_head * self.config.v_head_dim) if self.output_gate is not None: y = torch.sigmoid(self.output_gate(x)) * y # Output projection. return self.proj(y) # (B, T, C) def compressed_cache_attention( self, q_pass: torch.Tensor, q_roped: torch.Tensor, compressed_latent: torch.Tensor, k_roped: torch.Tensor, input_pos: torch.Tensor, input_pos_maxp1: int | None, mask: torch.Tensor | None, gate_input: torch.Tensor, ) -> torch.Tensor: """Decode from compressed MLA state without materializing per-head K/V.""" if self.kv_b_proj.bias is not None or self.proj.bias is not None: raise NotImplementedError("Compressed MLA cache currently requires bias=False.") latent_cache, rope_cache = self.kv_cache(input_pos, compressed_latent, k_roped) if input_pos_maxp1 is not None: latent_cache = latent_cache[:, :input_pos_maxp1, :] rope_cache = rope_cache[..., :input_pos_maxp1, :] if mask is not None: mask = mask[..., :input_pos_maxp1] h = self.config.n_head g = self.config.n_query_groups q_per_group = h // g r = self.config.kv_lora_rank k_dim = self.config.qk_nope_head_dim v_dim = self.config.v_head_dim kv_weight = self.kv_b_proj.weight.view(g, k_dim + v_dim, r) k_up = kv_weight[:, :k_dim, :].repeat_interleave(q_per_group, dim=0) v_up = kv_weight[:, k_dim:, :].repeat_interleave(q_per_group, dim=0) latent_q = torch.einsum("bhtd,hdr->bhtr", q_pass, k_up) scores = torch.einsum("bhtr,bsr->bhts", latent_q, latent_cache) scores = scores + torch.einsum("bhtd,bnsd->bhts", q_roped, rope_cache) scale = self.mscale * self.mscale / math.sqrt( self.config.attention_scores_scalar or self.config.qk_head_dim ) scores = scores * scale if mask is not None: scores = scores.masked_fill(~mask, torch.finfo(scores.dtype).min) if mask.dtype == torch.bool else scores + mask weights = F.softmax(scores, dim=-1, dtype=torch.float).to(dtype=scores.dtype) latent_context = torch.einsum("bhts,bsr->bhtr", weights, latent_cache) if self.output_gate is not None: context = torch.einsum("bhtr,hvr->bhtv", latent_context, v_up) context = context.transpose(1, 2).reshape( gate_input.size(0), gate_input.size(1), h * v_dim, ) context = torch.sigmoid(self.output_gate(gate_input)) * context return self.proj(context) out_weight = self.proj.weight.view(self.config.n_embd, h, v_dim) absorbed_v_o = torch.einsum("ohv,hvr->hor", out_weight, v_up) return torch.einsum("bhtr,hor->bto", latent_context, absorbed_v_o) def scaled_dot_product_attention( self, q: torch.Tensor, k: torch.Tensor, v: torch.Tensor, mask: torch.Tensor | None = None ) -> torch.Tensor: scale = 1.0 / math.sqrt(self.config.attention_scores_scalar or self.config.qk_head_dim) scale = scale * self.mscale * self.mscale # with softcapping we cannot use SDPA if self.config.attention_logit_softcapping is not None: scores = q @ k.mT * scale scores = do_softcapping(scores, self.config.attention_logit_softcapping) if mask is None: mask = torch.ones(q.size(2), q.size(2), dtype=q.dtype, device=q.device).triu(diagonal=1) mask.masked_fill_(mask.bool(), torch.finfo(q.dtype).min) scores = scores + mask scores = F.softmax(scores, dim=-1, dtype=torch.float).to(dtype=q.dtype) y = scores @ v else: y = F.scaled_dot_product_attention( q, k, v, attn_mask=mask, dropout_p=0.0, scale=scale, is_causal=mask is None ) return y.transpose(1, 2) def build_kv_cache( self, batch_size: int, max_seq_length: int, rope_cache_length: int | None = None, device: torch.device | None = None, dtype: torch.dtype | None = None, ) -> "MLACompressedKVCache": latent_shape = (batch_size, max_seq_length, self.config.kv_lora_rank) rope_shape = (batch_size, 1, max_seq_length, self.config.qk_rope_head_dim) if rope_cache_length is not None: print("Warning: `rope_cache_length` has no effect on MultiheadLatentAttention!") if self.config.rotary_percentage != 1.0: print("Warning: `rotary_percentage` has no effect on MultiheadLatentAttention!") return MLACompressedKVCache(latent_shape, rope_shape, device=device, dtype=dtype) class GptNeoxMLP(nn.Module): def __init__(self, config: Config, intermediate_size: int | None = None) -> None: super().__init__() self.intermediate_size = intermediate_size or config.intermediate_size self.fc = nn.Linear(config.n_embd, self.intermediate_size, bias=config.bias) self.proj = nn.Linear(self.intermediate_size, config.n_embd, bias=config.bias) self.config = config def forward(self, x: torch.Tensor) -> torch.Tensor: x = self.fc(x) x = F.gelu(x, approximate=self.config.gelu_approximate) return self.proj(x) def powlu_gate(x: torch.Tensor, m: float) -> torch.Tensor: positive_mask = x > 0 positive = torch.where(positive_mask, x, torch.ones_like(x)) exponent = float(m) / (torch.sqrt(positive) + 1.0) positive_gate = torch.pow(positive, exponent) * torch.sigmoid(x) return torch.where(positive_mask, positive_gate, F.silu(x)) class LLaMAMLP(nn.Module): def __init__(self, config: Config, intermediate_size: int | None = None) -> None: super().__init__() self.intermediate_size = intermediate_size or config.intermediate_size self.fc_1 = nn.Linear(config.n_embd, self.intermediate_size, bias=config.bias) self.fc_2 = nn.Linear(config.n_embd, self.intermediate_size, bias=config.bias) self.proj = nn.Linear(self.intermediate_size, config.n_embd, bias=config.bias) self.config = config def forward(self, x: torch.Tensor) -> torch.Tensor: x_fc_1 = self.fc_1(x) x_fc_2 = self.fc_2(x) x = F.silu(x_fc_1) * x_fc_2 return self.proj(x) class KimiSiTUGLUMLP(LLaMAMLP): """Kimi K3 Eq. 12 dense SiTU-GLU. Both multiplicative branches are smoothly capped in FP32 using the constants released by Moonshot (beta_gate=4, beta_up=25). """ def forward(self, x: torch.Tensor) -> torch.Tensor: gate = self.fc_1(x).float() up = self.fc_2(x).float() beta_gate = self.config.kimi_situ_beta beta_up = self.config.kimi_situ_linear_beta gate = beta_gate * torch.tanh(gate / beta_gate) * torch.sigmoid(gate) up = beta_up * torch.tanh(up / beta_up) return self.proj((gate * up).to(x.dtype)) class _KimiSiTUExpert(nn.Module): def __init__(self, input_size: int, intermediate_size: int, config: Config) -> None: super().__init__() self.gate = nn.Linear(input_size, intermediate_size, bias=False) self.up = nn.Linear(input_size, intermediate_size, bias=False) self.down = nn.Linear(intermediate_size, input_size, bias=False) self.beta_gate = config.kimi_situ_beta self.beta_up = config.kimi_situ_linear_beta def forward(self, x: torch.Tensor) -> torch.Tensor: gate = self.gate(x).float() up = self.up(x).float() gate = self.beta_gate * torch.tanh(gate / self.beta_gate) * torch.sigmoid(gate) up = self.beta_up * torch.tanh(up / self.beta_up) return self.down((gate * up).to(x.dtype)) class KimiStableLatentMoEMLP(nn.Module): """Training-capable Kimi K3 Stable LatentMoE (Eqs. 11-14). The routed experts operate at latent width, their weighted sum is RMS-normalized before the full-width up projection, two full-width shared SiTU experts are always active, and the next batch's dispatch bias is updated with exact single-device Quantile Balancing. """ def __init__(self, config: Config) -> None: super().__init__() d = config.n_embd latent = config.kimi_latent_moe_latent_size n_experts = config.kimi_latent_moe_num_experts routed_hidden = config.kimi_latent_moe_expert_intermediate_size shared_hidden = ( config.kimi_latent_moe_shared_intermediate_size * config.kimi_latent_moe_num_shared_experts ) self.config = config self.router = nn.Linear(d, n_experts, bias=False) self.routed_down = nn.Linear(d, latent, bias=False) self.routed_up = nn.Linear(latent, d, bias=False) self.routed_norm = config.norm_class(latent, eps=config.norm_eps) # Pack all routed experts so selected experts execute as batched # matmuls. This is algebraically identical to a ModuleList of # independent SiTU experts, but avoids one Python/CUDA launch chain # per expert. self.routed_gate_weight = nn.Parameter( torch.empty(n_experts, routed_hidden, latent) ) self.routed_up_weight = nn.Parameter( torch.empty(n_experts, routed_hidden, latent) ) self.routed_down_weight = nn.Parameter( torch.empty(n_experts, latent, routed_hidden) ) for parameter in ( self.routed_gate_weight, self.routed_up_weight, self.routed_down_weight, ): nn.init.normal_(parameter, mean=0.0, std=0.02) # Concatenating the two experts' hidden axes and using one down # projection is algebraically identical to summing two independent # full-width GLUs, while issuing three large GEMMs instead of six. self.shared_experts = _KimiSiTUExpert(d, shared_hidden, config) self.register_buffer("quantile_bias", torch.zeros(n_experts)) @torch.no_grad() def _next_quantile_bias( self, scores: torch.Tensor, biased_scores: torch.Tensor, ) -> torch.Tensor: k = self.config.kimi_latent_moe_top_k experts = self.config.kimi_latent_moe_num_experts bins = self.config.kimi_latent_moe_qb_bins cutoff = biased_scores.topk(k + 1, dim=-1, sorted=True).values[:, k] # Appendix D: histogram the required bias r_ij = alpha_i - s_ij # over [b_min - 1, b_max + 1], then recover its k/n quantile. required_bias = cutoff[:, None] - scores lower = self.quantile_bias.min().float() - 1.0 upper = self.quantile_bias.max().float() + 1.0 bin_width = (upper - lower) / bins bin_index = torch.floor( (required_bias - lower) / bin_width ).to(torch.long).clamp_(0, bins - 1) expert_offset = ( torch.arange(experts, device=scores.device, dtype=torch.long) * bins ) flat_index = (bin_index + expert_offset).reshape(-1) histogram = torch.zeros( experts * bins, device=scores.device, dtype=torch.int32, ) histogram.scatter_add_( 0, flat_index, torch.ones_like(flat_index, dtype=torch.int32), ) histogram = histogram.view(experts, bins) cumulative = histogram.cumsum(dim=1) target_load = (scores.size(0) * k + experts - 1) // experts selected_bin = (cumulative >= target_load).to(torch.int32).argmax(dim=1) selected_count = histogram.gather(1, selected_bin[:, None]).squeeze(1) prior_bin = (selected_bin - 1).clamp_min(0) prior_count = cumulative.gather(1, prior_bin[:, None]).squeeze(1) prior_count = torch.where( selected_bin == 0, torch.zeros_like(prior_count), prior_count, ) fraction = ( (target_load - prior_count).float() / selected_count.clamp_min(1).float() ).clamp_(0.0, 1.0) next_bias = lower + (selected_bin.float() + fraction) * bin_width return next_bias - next_bias.mean() def forward(self, x: torch.Tensor) -> torch.Tensor: shape = x.shape flat = x.reshape(-1, shape[-1]) scores = torch.sigmoid(self.router(flat).float()) biased_scores = scores + self.quantile_bias.float() k = self.config.kimi_latent_moe_top_k indices = biased_scores.topk(k, dim=-1, sorted=False).indices weights = scores.gather(1, indices) weights = weights / weights.sum(dim=-1, keepdim=True).clamp_min(1e-20) latent = self.routed_down(flat) token_count, route_count = indices.shape route_input = latent[:, None, :].expand( token_count, route_count, latent.size(-1), ).reshape(-1, latent.size(-1)) flat_indices = indices.reshape(-1) gate = torch.bmm( self.routed_gate_weight[flat_indices], route_input.unsqueeze(-1), ).squeeze(-1).float() up = torch.bmm( self.routed_up_weight[flat_indices], route_input.unsqueeze(-1), ).squeeze(-1).float() gate = ( self.config.kimi_situ_beta * torch.tanh(gate / self.config.kimi_situ_beta) * torch.sigmoid(gate) ) up = self.config.kimi_situ_linear_beta * torch.tanh( up / self.config.kimi_situ_linear_beta ) expert_hidden = (gate * up).to(x.dtype) expert_output = torch.bmm( self.routed_down_weight[flat_indices], expert_hidden.unsqueeze(-1), ).squeeze(-1) expert_output = expert_output * weights.reshape(-1, 1).to(x.dtype) routed = torch.zeros_like(latent) routed.index_add_( 0, torch.arange(token_count, device=x.device).repeat_interleave( route_count ), expert_output, ) if self.training and self.config.kimi_latent_moe_quantile_balancing: self.quantile_bias.copy_( self._next_quantile_bias(scores.detach(), biased_scores.detach()) ) routed = self.routed_up(self.routed_norm(routed)) shared = self.shared_experts(flat) return (shared + routed).view(shape) class LLaMAPowLUMLP(LLaMAMLP): def forward(self, x: torch.Tensor) -> torch.Tensor: x_fc_1 = self.fc_1(x) x_fc_2 = self.fc_2(x) x = powlu_gate(x_fc_1, self.config.powlu_m) * x_fc_2 return self.proj(x) class DSwiGLUMLP(nn.Module): def __init__(self, config: Config, intermediate_size: int | None = None) -> None: super().__init__() self.intermediate_size = intermediate_size or config.intermediate_size self.fc_1 = nn.Linear(config.n_embd, self.intermediate_size, bias=config.bias) self.fc_2 = nn.Linear(config.n_embd, self.intermediate_size, bias=config.bias) self.fc_3 = nn.Linear(config.n_embd, self.intermediate_size, bias=config.bias) self.proj = nn.Linear(self.intermediate_size, config.n_embd, bias=config.bias) self.config = config self.last_gate_sparsity: torch.Tensor | None = None def _hard_topk_gate(self, logits: torch.Tensor) -> torch.Tensor: active = max(1, min(self.intermediate_size, int(round(self.intermediate_size * self.config.drelu_target_active)))) threshold = torch.topk(logits.float(), k=active, dim=-1).values[..., -1:].detach() centered = logits.float() - threshold soft = torch.sigmoid(self.config.drelu_gate_sharpness * centered).to(dtype=logits.dtype) hard = (centered >= 0).to(dtype=logits.dtype) return hard + soft - soft.detach() def forward(self, x: torch.Tensor) -> torch.Tensor: representation = F.silu(self.fc_1(x)) * self.fc_2(x) gate = self._hard_topk_gate(self.fc_3(x)) self.last_gate_sparsity = (gate.detach() == 0).float().mean() return self.proj(representation * gate) class AdaptiveDSwiGLUMLP(nn.Module): def __init__(self, config: Config, intermediate_size: int | None = None) -> None: super().__init__() self.intermediate_size = intermediate_size or config.intermediate_size self.fc_1 = nn.Linear(config.n_embd, self.intermediate_size, bias=config.bias) self.fc_2 = nn.Linear(config.n_embd, self.intermediate_size, bias=config.bias) self.fc_3 = nn.Linear(config.n_embd, self.intermediate_size, bias=config.bias) self.threshold = nn.Linear(config.n_embd, 1, bias=True) self.proj = nn.Linear(self.intermediate_size, config.n_embd, bias=config.bias) self.config = config self.last_gate_sparsity: torch.Tensor | None = None self.last_threshold_mean: torch.Tensor | None = None init = min(max(config.drelu_threshold_init, 1e-4), 1 - 1e-4) torch.nn.init.zeros_(self.threshold.weight) torch.nn.init.constant_(self.threshold.bias, math.log(init / (1 - init))) def _adaptive_gate(self, logits: torch.Tensor, x: torch.Tensor) -> torch.Tensor: probabilities = torch.sigmoid(logits.float()) threshold = torch.sigmoid(self.threshold(x).float()) centered = probabilities - threshold soft = torch.sigmoid(self.config.drelu_gate_sharpness * centered).to(dtype=logits.dtype) hard = (centered >= 0).to(dtype=logits.dtype) self.last_threshold_mean = threshold.detach().mean() return hard + soft - soft.detach() def forward(self, x: torch.Tensor) -> torch.Tensor: representation = F.silu(self.fc_1(x)) * self.fc_2(x) gate = self._adaptive_gate(self.fc_3(x), x) self.last_gate_sparsity = (gate.detach() == 0).float().mean() return self.proj(representation * gate) class BlockSparseAdaptiveDSwiGLUMLP(nn.Module): """Adaptive dSwiGLU with structured conditional compute. The dense dSwiGLU variants compute every SwiGLU channel and then zero out activations. This variant splits the intermediate MLP into groups and only runs the selected groups for each token. The Python implementation is meant for laptop-scale research and correctness; production speedups need a fused grouped kernel. """ def __init__(self, config: Config, intermediate_size: int | None = None) -> None: super().__init__() self.intermediate_size = intermediate_size or config.intermediate_size self.num_groups = config.sparse_mlp_num_groups self.min_active_groups = config.sparse_mlp_min_active_groups self.max_active_groups = config.sparse_mlp_max_active_groups if self.intermediate_size % self.num_groups != 0: raise ValueError("intermediate_size must be divisible by sparse_mlp_num_groups") self.group_size = self.intermediate_size // self.num_groups if config.bias: raise ValueError("BlockSparseAdaptiveDSwiGLUMLP currently expects bias=False for grouped GEMM.") self.up_weight = nn.Parameter(torch.empty(self.num_groups, config.n_embd, self.group_size)) self.value_weight = nn.Parameter(torch.empty(self.num_groups, config.n_embd, self.group_size)) self.down_weight = nn.Parameter(torch.empty(self.num_groups, self.group_size, config.n_embd)) self.router = nn.Linear(config.n_embd, self.num_groups, bias=True) self.threshold = nn.Linear(config.n_embd, 1, bias=True) self.config = config self.last_gate_sparsity: torch.Tensor | None = None self.last_threshold_mean: torch.Tensor | None = None self.last_active_groups_mean: torch.Tensor | None = None init = min(max(config.drelu_threshold_init, 1e-4), 1 - 1e-4) self.reset_parameters() def reset_parameters(self) -> None: std = math.sqrt(2.0 / 5 / self.config.n_embd) nn.init.normal_(self.up_weight, mean=0.0, std=std) nn.init.normal_(self.value_weight, mean=0.0, std=std) nn.init.normal_(self.down_weight, mean=0.0, std=std) self.router.reset_parameters() torch.nn.init.zeros_(self.threshold.weight) init = min(max(self.config.drelu_threshold_init, 1e-4), 1 - 1e-4) torch.nn.init.constant_(self.threshold.bias, math.log(init / (1 - init))) def _active_group_mask(self, x: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]: probabilities = torch.sigmoid(self.router(x).float()) threshold = torch.sigmoid(self.threshold(x).float()) _, top_indices = torch.topk(probabilities, k=self.max_active_groups, dim=-1) top_mask = torch.zeros_like(probabilities, dtype=torch.bool) top_mask.scatter_(-1, top_indices, True) threshold_mask = probabilities >= threshold active_mask = threshold_mask & top_mask active_count = active_mask.sum(dim=-1) if self.min_active_groups > 0: _, min_indices = torch.topk(probabilities, k=self.min_active_groups, dim=-1) min_mask = torch.zeros_like(active_mask) min_mask.scatter_(-1, min_indices, True) active_mask = active_mask | (active_count.unsqueeze(-1) < self.min_active_groups) & min_mask self.last_gate_sparsity = 1.0 - active_mask.detach().float().mean() self.last_threshold_mean = threshold.detach().mean() self.last_active_groups_mean = active_mask.detach().float().sum(dim=-1).mean() return active_mask, probabilities.to(dtype=x.dtype), threshold.to(dtype=x.dtype) def _exact_topk_pairs(self, x: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]: probabilities = torch.sigmoid(self.router(x).float()) top_values, top_indices = torch.topk(probabilities, k=self.max_active_groups, dim=-1) self.last_gate_sparsity = x.new_tensor(1.0 - self.max_active_groups / self.num_groups) self.last_threshold_mean = torch.sigmoid(self.threshold(x).float()).detach().mean() self.last_active_groups_mean = x.new_tensor(float(self.max_active_groups)) token_idx = torch.arange(x.shape[0], device=x.device).repeat_interleave(self.max_active_groups) return token_idx, top_indices.reshape(-1), top_values.reshape(-1).to(dtype=x.dtype) @torch._dynamo.disable def forward(self, x: torch.Tensor) -> torch.Tensor: original_shape = x.shape x_flat = x.reshape(-1, original_shape[-1]) output = x_flat.new_zeros(x_flat.shape) if self.min_active_groups == self.max_active_groups: token_idx, group_idx, scale_values = self._exact_topk_pairs(x_flat) else: active_mask, probabilities, _ = self._active_group_mask(x_flat) active_pairs = torch.nonzero(active_mask, as_tuple=False) if active_pairs.numel() == 0: return output.reshape(original_shape) token_idx = active_pairs[:, 0] group_idx = active_pairs[:, 1] scale_values = probabilities[token_idx, group_idx] order = torch.argsort(group_idx) token_idx = token_idx.index_select(0, order) group_idx = group_idx.index_select(0, order) scale_values = scale_values.index_select(0, order) selected_x = x_flat.index_select(0, token_idx).contiguous() counts = torch.bincount(group_idx, minlength=self.num_groups) offsets = torch.cumsum(counts, dim=0).to(dtype=torch.int32) if not selected_x.is_cuda: for group_idx_int in range(self.num_groups): pair_idx = torch.nonzero(group_idx == group_idx_int, as_tuple=False).flatten() if pair_idx.numel() == 0: continue group_token_idx = token_idx.index_select(0, pair_idx) x_group = x_flat.index_select(0, group_token_idx) hidden = F.silu(x_group @ self.up_weight[group_idx_int]) * (x_group @ self.value_weight[group_idx_int]) projected = hidden @ self.down_weight[group_idx_int] scale = scale_values.index_select(0, pair_idx).unsqueeze(-1) output.index_add_(0, group_token_idx, projected * scale.to(dtype=projected.dtype)) return output.reshape(original_shape) grouped_dtype = torch.bfloat16 if selected_x.is_cuda else selected_x.dtype selected_x_gemm = selected_x.to(dtype=grouped_dtype) up_weight = self.up_weight.to(dtype=grouped_dtype) value_weight = self.value_weight.to(dtype=grouped_dtype) down_weight = self.down_weight.to(dtype=grouped_dtype) up = torch._grouped_mm(selected_x_gemm, up_weight, offsets) value = torch._grouped_mm(selected_x_gemm, value_weight, offsets) hidden = F.silu(up) * value projected = torch._grouped_mm(hidden, down_weight, offsets).to(dtype=output.dtype) scale = scale_values.unsqueeze(-1).to(dtype=projected.dtype) output.index_add_(0, token_idx, projected * scale) return output.reshape(original_shape) class HiddenBlockDSwiGLUMLP(nn.Module): """Parameter-matched hidden-channel block SwiGLU. Unlike ``TileRoutedDSwiGLUMLP``, which gives every group the full residual stream and partitions the intermediate channels, this module partitions the residual stream itself into blocks. Each block owns a local SwiGLU with the full configured intermediate width, so total parameter count and matmul work are close to a dense dSwiGLU at the same ``intermediate_size``. ``grouped_mlp_stack_alpha`` turns the parallel blocks into a small cascade: each later block can see a scaled residual update from the previous block. The existing block-level channel rotation hook can rotate hidden channels across transformer layers, preventing the same channels from being trapped in the same local block for the whole network. """ def __init__(self, config: Config, intermediate_size: int | None = None) -> None: super().__init__() self.intermediate_size = intermediate_size or config.intermediate_size self.num_blocks = config.sparse_mlp_num_groups self.active_groups = config.sparse_mlp_max_active_groups if self.active_groups != self.num_blocks: raise ValueError("HiddenBlockDSwiGLUMLP is a full-active block variant.") if config.n_embd % self.num_blocks != 0: raise ValueError("n_embd must be divisible by sparse_mlp_num_groups") if config.bias: raise ValueError("HiddenBlockDSwiGLUMLP currently expects bias=False.") self.block_hidden_size = config.n_embd // self.num_blocks self.up_weight = nn.Parameter(torch.empty(self.num_blocks, self.block_hidden_size, self.intermediate_size)) self.value_weight = nn.Parameter(torch.empty(self.num_blocks, self.block_hidden_size, self.intermediate_size)) self.down_weight = nn.Parameter(torch.empty(self.num_blocks, self.intermediate_size, self.block_hidden_size)) stack_alpha = torch.tensor(float(config.grouped_mlp_stack_alpha)) if config.grouped_mlp_stack_alpha_learnable: self.stack_alpha = nn.Parameter(stack_alpha) else: self.register_buffer("stack_alpha", stack_alpha, persistent=False) output_scale = torch.tensor(float(config.grouped_mlp_stack_output_scale)) if config.grouped_mlp_stack_output_scale_learnable: self.stack_output_scale = nn.Parameter(output_scale) else: self.register_buffer("stack_output_scale", output_scale, persistent=False) self.config = config self.last_gate_sparsity: torch.Tensor | None = None self.last_threshold_mean: torch.Tensor | None = None self.last_active_groups_mean: torch.Tensor | None = None self.reset_parameters() def reset_parameters(self) -> None: std = math.sqrt(2.0 / 5 / self.config.n_embd) std *= float(getattr(self.config, "grouped_mlp_stack_init_scale", 1.0)) nn.init.normal_(self.up_weight, mean=0.0, std=std) nn.init.normal_(self.value_weight, mean=0.0, std=std) nn.init.normal_(self.down_weight, mean=0.0, std=std) def forward(self, x: torch.Tensor) -> torch.Tensor: original_shape = x.shape x_flat = x.reshape(-1, original_shape[-1]) x_blocks = x_flat.reshape(x_flat.shape[0], self.num_blocks, self.block_hidden_size) stack_alpha = self.stack_alpha.to(device=x_flat.device, dtype=x_flat.dtype) deltas: list[torch.Tensor] = [] carry: torch.Tensor | None = None for block_idx in range(self.num_blocks): state = x_blocks[:, block_idx, :] if carry is not None: state = state + stack_alpha * carry hidden = F.silu(state @ self.up_weight[block_idx]) * (state @ self.value_weight[block_idx]) carry = hidden @ self.down_weight[block_idx] deltas.append(carry) output = torch.stack(deltas, dim=1).reshape(original_shape) if not torch.is_grad_enabled(): self.last_gate_sparsity = x_flat.new_tensor(0.0) self.last_threshold_mean = x_flat.new_tensor(0.0) self.last_active_groups_mean = x_flat.new_tensor(float(self.active_groups)) output_scale = self.stack_output_scale.to(device=x_flat.device, dtype=x_flat.dtype) return output * output_scale class TileRoutedDSwiGLUMLP(nn.Module): """Tile-routed grouped dSwiGLU. This is the hardware-aligned sparse MLP experiment. It keeps one active MLP group per token, but routes contiguous token tiles instead of individual tokens. The forward is expressed as batched GEMMs so backward uses efficient dense kernels instead of per-token atomic accumulation. """ def __init__(self, config: Config, intermediate_size: int | None = None) -> None: super().__init__() self.intermediate_size = intermediate_size or config.intermediate_size self.num_groups = config.sparse_mlp_num_groups if self.intermediate_size % self.num_groups != 0: raise ValueError("intermediate_size must be divisible by sparse_mlp_num_groups") if config.bias: raise ValueError("TileRoutedDSwiGLUMLP currently expects bias=False.") self.group_size = self.intermediate_size // self.num_groups self.active_groups = config.sparse_mlp_max_active_groups self.up_weight = nn.Parameter(torch.empty(self.num_groups, config.n_embd, self.group_size)) self.value_weight = nn.Parameter(torch.empty(self.num_groups, config.n_embd, self.group_size)) self.down_weight = nn.Parameter(torch.empty(self.num_groups, self.group_size, config.n_embd)) self.shared_path: LLaMAMLP | None = None if config.grouped_mlp_shared_intermediate_size > 0: self.shared_path = LLaMAMLP(config, intermediate_size=config.grouped_mlp_shared_intermediate_size) shared_alpha = torch.tensor(float(config.grouped_mlp_shared_alpha)) if config.grouped_mlp_shared_alpha_learnable: self.shared_alpha = nn.Parameter(shared_alpha) else: self.register_buffer("shared_alpha", shared_alpha, persistent=False) self.config = config self.last_gate_sparsity: torch.Tensor | None = None self.last_threshold_mean: torch.Tensor | None = None self.last_active_groups_mean: torch.Tensor | None = None self.reset_parameters() def reset_parameters(self) -> None: std = math.sqrt(2.0 / 5 / self.config.n_embd) nn.init.normal_(self.up_weight, mean=0.0, std=std) nn.init.normal_(self.value_weight, mean=0.0, std=std) nn.init.normal_(self.down_weight, mean=0.0, std=std) def _apply_shared_path(self, x: torch.Tensor, grouped_output: torch.Tensor) -> torch.Tensor: if self.shared_path is None: return grouped_output shared_alpha = self.shared_alpha.to(device=grouped_output.device, dtype=grouped_output.dtype) return grouped_output + shared_alpha * self.shared_path(x) def _gate(self, up: torch.Tensor) -> torch.Tensor: return F.silu(up) def _forward_full_active(self, x: torch.Tensor, original_shape: torch.Size, x_flat: torch.Tensor) -> torch.Tensor: # Full-active grouped MLP is mathematically a dense SwiGLU whose # intermediate channels are partitioned into groups. Collapse grouped # parameter tensors into dense GEMM views for the 4/4 quality path. hidden_size = original_shape[-1] up_weight = self.up_weight.permute(1, 0, 2).reshape(hidden_size, self.intermediate_size) value_weight = self.value_weight.permute(1, 0, 2).reshape(hidden_size, self.intermediate_size) down_weight = self.down_weight.reshape(self.intermediate_size, hidden_size) output = (self._gate(x_flat @ up_weight) * (x_flat @ value_weight)) @ down_weight return output.reshape(original_shape) def forward(self, x: torch.Tensor) -> torch.Tensor: original_shape = x.shape x_flat = x.reshape(-1, original_shape[-1]) if self.active_groups == self.num_groups: output = self._forward_full_active(x, original_shape, x_flat) if not torch.is_grad_enabled(): self.last_gate_sparsity = x_flat.new_tensor(0.0) self.last_threshold_mean = x_flat.new_tensor(0.0) self.last_active_groups_mean = x_flat.new_tensor(float(self.active_groups)) return self._apply_shared_path(x, output) usable_tokens = x_flat.shape[0] - (x_flat.shape[0] % self.num_groups) output = x_flat.new_empty(x_flat.shape) if usable_tokens: x_tiles = x_flat[:usable_tokens].reshape(self.num_groups, usable_tokens // self.num_groups, x_flat.shape[-1]) out_tiles = torch.zeros_like(x_tiles) for route_offset in range(self.active_groups): if route_offset == 0: up_weight = self.up_weight value_weight = self.value_weight down_weight = self.down_weight else: up_weight = torch.roll(self.up_weight, shifts=-route_offset, dims=0) value_weight = torch.roll(self.value_weight, shifts=-route_offset, dims=0) down_weight = torch.roll(self.down_weight, shifts=-route_offset, dims=0) up = torch.bmm(x_tiles, up_weight) value = torch.bmm(x_tiles, value_weight) out_tiles = out_tiles + torch.bmm(self._gate(up) * value, down_weight) output[:usable_tokens] = out_tiles.reshape(usable_tokens, x_flat.shape[-1]) if usable_tokens < x_flat.shape[0]: tail = x_flat[usable_tokens:] tail_out = torch.zeros_like(tail) for route_offset in range(self.active_groups): hidden = self._gate(tail @ self.up_weight[route_offset]) * (tail @ self.value_weight[route_offset]) tail_out = tail_out + hidden @ self.down_weight[route_offset] output[usable_tokens:] = tail_out if not torch.is_grad_enabled(): self.last_gate_sparsity = x_flat.new_tensor(1.0 - self.active_groups / self.num_groups) self.last_threshold_mean = x_flat.new_tensor(0.0) self.last_active_groups_mean = x_flat.new_tensor(float(self.active_groups)) return self._apply_shared_path(x, output.reshape(original_shape)) class TileRoutedActivationDSwiGLUMLP(TileRoutedDSwiGLUMLP): """Full-active grouped GLU with a configurable gate activation.""" gate_kind = "silu" def _gate(self, up: torch.Tensor) -> torch.Tensor: if self.gate_kind == "gelu": return F.gelu(up, approximate=self.config.gelu_approximate) if self.gate_kind == "relu": return F.relu(up) if self.gate_kind == "squared_relu": relu = F.relu(up) return relu * relu if self.gate_kind == "sigmoid": return torch.sigmoid(up) if self.gate_kind == "powlu": return powlu_gate(up, self.config.powlu_m) return F.silu(up) def _forward_full_active(self, x: torch.Tensor, original_shape: torch.Size, x_flat: torch.Tensor) -> torch.Tensor: hidden_size = original_shape[-1] up_weight = self.up_weight.permute(1, 0, 2).reshape(hidden_size, self.intermediate_size) value_weight = self.value_weight.permute(1, 0, 2).reshape(hidden_size, self.intermediate_size) down_weight = self.down_weight.reshape(self.intermediate_size, hidden_size) hidden = self._gate(x_flat @ up_weight) * (x_flat @ value_weight) return (hidden @ down_weight).reshape(original_shape) class TileRoutedGEGLUMLP(TileRoutedActivationDSwiGLUMLP): gate_kind = "gelu" class TileRoutedReGLUMLP(TileRoutedActivationDSwiGLUMLP): gate_kind = "relu" class TileRoutedSquaredReGLUMLP(TileRoutedActivationDSwiGLUMLP): gate_kind = "squared_relu" class TileRoutedSigmoidGLUMLP(TileRoutedActivationDSwiGLUMLP): gate_kind = "sigmoid" class TileRoutedPowLUGLUMLP(TileRoutedActivationDSwiGLUMLP): gate_kind = "powlu" class TileRoutedEulerGLUMLP(TileRoutedActivationDSwiGLUMLP): """Grouped SwiGLU with a normalized Euler maximum envelope. For ``p = softplus(u)``, ``p ** (1 / p)`` reaches its unique global maximum at ``p = e``. Normalizing by ``e ** (1 / e)`` bounds the positive modulation by one while SiLU retains the signed gate behavior. """ def _gate(self, up: torch.Tensor) -> torch.Tensor: p = F.softplus(up.float()).clamp_min(torch.finfo(torch.float32).tiny) envelope = torch.exp(torch.log(p) / p - (1.0 / math.e)).to(dtype=up.dtype) return F.silu(up) * envelope class TileRoutedEulerMishBlendGLUMLP(TileRoutedEulerGLUMLP): """Euler gate with a zero-start, per-group Mish residual. The gate begins as the fixed Euler gate exactly. Each contiguous 4/4 group can learn a bounded interpolation toward Mish, preserving the proven positive Euler envelope while exposing Mish's smooth negative curvature. """ def __init__(self, config: Config, intermediate_size: int | None = None) -> None: super().__init__(config, intermediate_size) if self.active_groups != self.num_groups: raise ValueError("Euler-Mish blend requires all groups to be active.") self.mish_residual = nn.Parameter(torch.zeros(self.num_groups)) def _gate(self, up: torch.Tensor) -> torch.Tensor: p = F.softplus(up.float()).clamp_min(torch.finfo(torch.float32).tiny) envelope = torch.exp(torch.log(p) / p - (1.0 / math.e)).to(dtype=up.dtype) euler_gate = F.silu(up) * envelope mish_gate = up * torch.tanh(p).to(dtype=up.dtype) mix = torch.tanh(self.mish_residual).repeat_interleave(self.group_size).to(dtype=up.dtype) return euler_gate + mix * (mish_gate - euler_gate) class TileRoutedContextualValueRotationEulerGLUMLP(TileRoutedEulerGLUMLP): """Euler 4/4 MLP with contextual orthogonal value rotations. The 1664-channel intermediate is interpreted as 13 aligned tiles, each containing four 32-channel branches. Four response-conditioned butterfly rotations mix the branches before the existing dense down projection. All controller coefficients start at zero, making this exactly the fixed Euler gate at initialization while retaining nonzero controller gradients. """ _tile_size = 32 _rotation_pairs = ((0, 1), (2, 3), (0, 2), (1, 3)) _max_half_angle = math.tan(math.pi / 24.0) def __init__(self, config: Config, intermediate_size: int | None = None) -> None: super().__init__(config, intermediate_size) if self.active_groups != self.num_groups or self.num_groups != 4: raise ValueError("Contextual value rotation requires full-active grouped 4/4.") if self.group_size % self._tile_size: raise ValueError("Contextual value rotation requires a group size divisible by 32.") self.num_value_tiles = self.group_size // self._tile_size self.cvr_energy_coeff = nn.Parameter(torch.zeros(len(self._rotation_pairs))) self.cvr_agreement_coeff = nn.Parameter(torch.zeros(len(self._rotation_pairs))) def _rotate_pair( self, left: torch.Tensor, right: torch.Tensor, pair_index: int, ) -> tuple[torch.Tensor, torch.Tensor]: left_float = left.float() right_float = right.float() epsilon = torch.finfo(torch.float32).eps left_energy = left_float.square().mean(dim=-1).clamp_min(epsilon) right_energy = right_float.square().mean(dim=-1).clamp_min(epsilon) energy_contrast = torch.tanh(0.5 * (torch.log(left_energy) - torch.log(right_energy))) agreement = (left_float * right_float).mean(dim=-1) * torch.rsqrt(left_energy * right_energy) agreement = agreement.clamp(-1.0, 1.0) controller = ( self.cvr_energy_coeff[pair_index].float() * energy_contrast + self.cvr_agreement_coeff[pair_index].float() * agreement ) half_angle = self._max_half_angle * torch.tanh(controller) half_angle_squared = half_angle.square() denominator = 1.0 + half_angle_squared cosine = ((1.0 - half_angle_squared) / denominator).to(dtype=left.dtype).unsqueeze(-1) sine = ((2.0 * half_angle) / denominator).to(dtype=left.dtype).unsqueeze(-1) return cosine * left - sine * right, sine * left + cosine * right def _forward_full_active(self, x: torch.Tensor, original_shape: torch.Size, x_flat: torch.Tensor) -> torch.Tensor: hidden_size = original_shape[-1] up_weight = self.up_weight.permute(1, 0, 2).reshape(hidden_size, self.intermediate_size) value_weight = self.value_weight.permute(1, 0, 2).reshape(hidden_size, self.intermediate_size) down_weight = self.down_weight.reshape(self.intermediate_size, hidden_size) up = x_flat @ up_weight value = x_flat @ value_weight p = F.softplus(up.float()).clamp_min(torch.finfo(torch.float32).tiny) envelope = torch.exp(torch.log(p) / p - (1.0 / math.e)).to(dtype=up.dtype) hidden = (F.silu(up) * envelope * value).reshape( x_flat.shape[0], self.num_groups, self.num_value_tiles, self._tile_size ) groups = list(hidden.unbind(dim=1)) for pair_index, (left_index, right_index) in enumerate(self._rotation_pairs): groups[left_index], groups[right_index] = self._rotate_pair( groups[left_index], groups[right_index], pair_index ) hidden = torch.stack(groups, dim=1).reshape(x_flat.shape[0], self.intermediate_size) return (hidden @ down_weight).reshape(original_shape) class TileRoutedContextAdaptiveEulerTemperatureGLUMLP(TileRoutedEulerGLUMLP): """Euler 4/4 gate with a bounded, response-conditioned temperature. The normalized Euler log-envelope is non-positive and has its unique zero at ``softplus(up) = e``. Each group learns how strongly to sharpen one side of that maximum while relaxing the other, using the bounded signed displacement from ``e`` as local context. Zero parameters reproduce the fixed Euler gate exactly, and the temperature always stays in [0.5, 1.5], so the envelope remains bounded by one with the same global maximum. """ _max_temperature_delta = 0.5 def __init__(self, config: Config, intermediate_size: int | None = None) -> None: super().__init__(config, intermediate_size) if self.active_groups != self.num_groups or self.num_groups != 4: raise ValueError("Context-adaptive Euler temperature requires full-active grouped 4/4.") self.caet_temperature_raw = nn.Parameter(torch.zeros(self.num_groups)) def _gate(self, up: torch.Tensor) -> torch.Tensor: if up.shape[-1] != self.intermediate_size: raise RuntimeError("Context-adaptive Euler expects the full active intermediate dimension.") up_float = up.float() p = F.softplus(up_float).clamp_min(torch.finfo(torch.float32).tiny) log_envelope = torch.log(p) / p - (1.0 / math.e) response = (p - math.e) / (p + math.e) group_delta = self._max_temperature_delta * torch.tanh(self.caet_temperature_raw.float()) channel_delta = group_delta.repeat_interleave(self.group_size) temperature = 1.0 + channel_delta * response envelope = torch.exp(log_envelope * temperature).to(dtype=up.dtype) return F.silu(up) * envelope class TileRoutedAdaptiveEulerGLUMLP(TileRoutedDSwiGLUMLP): """Full-active grouped SwiGLU with a learned SiLU/Euler choice per group. The 4/4 path stores channels in four contiguous groups. Each group receives an independent differentiable selector between ordinary SiLU and the normalized Euler envelope. This adds four scalars per layer, not another MLP projection, so it can express activation specialization without changing the model's MLP width or matrix-multiply count. """ def __init__(self, config: Config, intermediate_size: int | None = None) -> None: super().__init__(config, intermediate_size) if self.active_groups != self.num_groups: raise ValueError("TileRoutedAdaptiveEulerGLUMLP requires all groups to be active.") initial_euler_probability = 0.75 initial_logit = math.log(initial_euler_probability / (1.0 - initial_euler_probability)) self.euler_mix_logits = nn.Parameter(torch.full((self.num_groups,), initial_logit)) def _gate(self, up: torch.Tensor) -> torch.Tensor: if up.shape[-1] != self.intermediate_size: raise RuntimeError("Adaptive Euler gating expects the full active intermediate dimension.") p = F.softplus(up.float()).clamp_min(torch.finfo(torch.float32).tiny) envelope = torch.exp(torch.log(p) / p - (1.0 / math.e)).to(dtype=up.dtype) euler_weight = torch.sigmoid(self.euler_mix_logits).repeat_interleave(self.group_size) return F.silu(up) * (1.0 + euler_weight.to(dtype=up.dtype) * (envelope - 1.0)) class TileRoutedDiverseEulerGLUMLP(TileRoutedAdaptiveEulerGLUMLP): """Adaptive Euler gating with deliberately diverse group initializations.""" def __init__(self, config: Config, intermediate_size: int | None = None) -> None: super().__init__(config, intermediate_size) probabilities = torch.linspace(0.125, 0.875, self.num_groups) with torch.no_grad(): self.euler_mix_logits.copy_(torch.logit(probabilities)) class TileRoutedGroupTokenAttentionEulerGLUMLP(TileRoutedDSwiGLUMLP): """Euler-gated 4/4 MLP with token-conditioned attention between groups.""" def __init__(self, config: Config, intermediate_size: int | None = None) -> None: super().__init__(config, intermediate_size) if self.active_groups != self.num_groups: raise ValueError("Group-token attention requires all groups to be active.") self.group_query_scale = nn.Parameter(torch.ones(self.num_groups)) self.group_key_scale = nn.Parameter(torch.ones(self.num_groups)) self.group_value_scale = nn.Parameter(torch.ones(self.num_groups)) self.group_attention_bias = nn.Parameter(torch.zeros(self.num_groups, self.num_groups)) self.group_attention_gain = nn.Parameter(torch.full((self.num_groups,), 0.1)) def _forward_full_active(self, x: torch.Tensor, original_shape: torch.Size, x_flat: torch.Tensor) -> torch.Tensor: hidden_size = original_shape[-1] up_weight = self.up_weight.permute(1, 0, 2).reshape(hidden_size, self.intermediate_size) value_weight = self.value_weight.permute(1, 0, 2).reshape(hidden_size, self.intermediate_size) down_weight = self.down_weight.reshape(self.intermediate_size, hidden_size) up = x_flat @ up_weight value = x_flat @ value_weight p = F.softplus(up.float()).clamp_min(torch.finfo(torch.float32).tiny) envelope = torch.exp(torch.log(p) / p - (1.0 / math.e)).to(dtype=up.dtype) hidden = (F.silu(up) * envelope * value).reshape(-1, self.num_groups, self.group_size) summaries = hidden.float().mean(dim=-1) query = summaries * self.group_query_scale.float() key = summaries * self.group_key_scale.float() scores = query.unsqueeze(-1) * key.unsqueeze(-2) + self.group_attention_bias.float() attention = F.softmax(scores, dim=-1) context = (attention @ (summaries * self.group_value_scale.float()).unsqueeze(-1)).squeeze(-1) gains = 1.0 + torch.tanh(context * self.group_attention_gain.float()) hidden = hidden * gains.to(dtype=hidden.dtype).unsqueeze(-1) return (hidden.reshape(-1, self.intermediate_size) @ down_weight).reshape(original_shape) class TileRoutedHierarchicalEulerGLUMLP(TileRoutedDSwiGLUMLP): """Euler-gated 4/4 MLP with a learned 2x2 macro-group stage.""" def __init__(self, config: Config, intermediate_size: int | None = None) -> None: super().__init__(config, intermediate_size) if self.num_groups != 4 or self.active_groups != self.num_groups: raise ValueError("Hierarchical Euler MLP is specialized for full-active grouped 4/4.") self.macro_mix_delta = nn.Parameter(torch.zeros(2, 2)) self.macro_gain = nn.Parameter(torch.full((2, 2), 0.1)) def _forward_full_active(self, x: torch.Tensor, original_shape: torch.Size, x_flat: torch.Tensor) -> torch.Tensor: hidden_size = original_shape[-1] up_weight = self.up_weight.permute(1, 0, 2).reshape(hidden_size, self.intermediate_size) value_weight = self.value_weight.permute(1, 0, 2).reshape(hidden_size, self.intermediate_size) down_weight = self.down_weight.reshape(self.intermediate_size, hidden_size) up = x_flat @ up_weight value = x_flat @ value_weight p = F.softplus(up.float()).clamp_min(torch.finfo(torch.float32).tiny) envelope = torch.exp(torch.log(p) / p - (1.0 / math.e)).to(dtype=up.dtype) hidden = (F.silu(up) * envelope * value).reshape(-1, 2, 2, self.group_size) macro = hidden.float().mean(dim=2) mix = torch.eye(2, device=macro.device, dtype=macro.dtype) + self.macro_mix_delta.float() macro = torch.einsum('tpg,pq->tqg', macro, mix) macro_summary = macro.mean(dim=-1) macro_gains = 1.0 + torch.tanh( macro_summary.unsqueeze(-1) * self.macro_gain.float().unsqueeze(0) ) hidden = hidden * macro_gains.to(dtype=hidden.dtype).unsqueeze(-1) return (hidden.reshape(-1, self.intermediate_size) @ down_weight).reshape(original_shape) class TileRoutedTokenExpertEulerGLUMLP(TileRoutedDSwiGLUMLP): """Token- and group-routed mixture of SiLU, Euler, and GELU gates.""" def __init__(self, config: Config, intermediate_size: int | None = None) -> None: super().__init__(config, intermediate_size) if self.active_groups != self.num_groups: raise ValueError("Token expert routing requires all groups to be active.") self.expert_router_slope = nn.Parameter(torch.zeros(self.num_groups, 3)) initial_bias = torch.tensor([-2.2, 2.9, -2.2]) self.expert_router_bias = nn.Parameter(initial_bias.repeat(self.num_groups, 1)) def _gate(self, up: torch.Tensor) -> torch.Tensor: if up.shape[-1] != self.intermediate_size: raise RuntimeError("Token expert routing expects the full active intermediate dimension.") grouped_up = up.reshape(-1, self.num_groups, self.group_size) summaries = grouped_up.float().mean(dim=-1) logits = summaries.unsqueeze(-1) * self.expert_router_slope.float() + self.expert_router_bias.float() weights = F.softmax(logits, dim=-1) p = F.softplus(grouped_up.float()).clamp_min(torch.finfo(torch.float32).tiny) envelope = torch.exp(torch.log(p) / p - (1.0 / math.e)).to(dtype=up.dtype) silu_gate = F.silu(grouped_up) gates = torch.stack((silu_gate, silu_gate * envelope, F.gelu(grouped_up)), dim=-1) return (gates * weights.to(dtype=up.dtype).unsqueeze(-2)).sum(dim=-1).reshape_as(up) class TileRoutedDynamicGroupMixerEulerGLUMLP(TileRoutedDSwiGLUMLP): """Euler-gated 4/4 MLP with a token-conditioned 4x4 group mixer.""" def __init__(self, config: Config, intermediate_size: int | None = None) -> None: super().__init__(config, intermediate_size) if self.num_groups != 4 or self.active_groups != self.num_groups: raise ValueError("Dynamic group mixing is specialized for full-active grouped 4/4.") self.mixer_slope = nn.Parameter(torch.zeros(4, 4)) self.mixer_bias = nn.Parameter(torch.full((4, 4), -3.0)) with torch.no_grad(): self.mixer_bias.diagonal().fill_(3.0) def _forward_full_active(self, x: torch.Tensor, original_shape: torch.Size, x_flat: torch.Tensor) -> torch.Tensor: hidden_size = original_shape[-1] up_weight = self.up_weight.permute(1, 0, 2).reshape(hidden_size, self.intermediate_size) value_weight = self.value_weight.permute(1, 0, 2).reshape(hidden_size, self.intermediate_size) down_weight = self.down_weight.reshape(self.intermediate_size, hidden_size) up = x_flat @ up_weight value = x_flat @ value_weight p = F.softplus(up.float()).clamp_min(torch.finfo(torch.float32).tiny) envelope = torch.exp(torch.log(p) / p - (1.0 / math.e)).to(dtype=up.dtype) hidden = (F.silu(up) * envelope * value).reshape(-1, 4, self.group_size) summaries = hidden.float().mean(dim=-1) logits = self.mixer_bias.float().unsqueeze(0) + summaries.unsqueeze(-1) * self.mixer_slope.float() mixer = F.softmax(logits, dim=-1) hidden = torch.einsum('tij,tjg->tig', mixer.to(dtype=hidden.dtype), hidden) return (hidden.reshape(-1, self.intermediate_size) @ down_weight).reshape(original_shape) class TileRoutedGroupCompetitionEulerGLUMLP(TileRoutedDSwiGLUMLP): """Euler-gated 4/4 MLP with token-wise conserved group energy.""" def __init__(self, config: Config, intermediate_size: int | None = None) -> None: super().__init__(config, intermediate_size) if self.active_groups != self.num_groups: raise ValueError("Group competition requires all groups to be active.") self.competition_temperature_raw = nn.Parameter(torch.tensor(0.0)) self.competition_strength_raw = nn.Parameter(torch.tensor(-2.0)) def _forward_full_active(self, x: torch.Tensor, original_shape: torch.Size, x_flat: torch.Tensor) -> torch.Tensor: hidden_size = original_shape[-1] up_weight = self.up_weight.permute(1, 0, 2).reshape(hidden_size, self.intermediate_size) value_weight = self.value_weight.permute(1, 0, 2).reshape(hidden_size, self.intermediate_size) down_weight = self.down_weight.reshape(self.intermediate_size, hidden_size) up = x_flat @ up_weight value = x_flat @ value_weight p = F.softplus(up.float()).clamp_min(torch.finfo(torch.float32).tiny) envelope = torch.exp(torch.log(p) / p - (1.0 / math.e)).to(dtype=up.dtype) hidden = (F.silu(up) * envelope * value).reshape(-1, self.num_groups, self.group_size) energy = hidden.float().square().mean(dim=-1).add(1e-6).log() temperature = F.softplus(self.competition_temperature_raw.float()) allocation = F.softmax(energy * temperature, dim=-1) * self.num_groups strength = torch.sigmoid(self.competition_strength_raw.float()) gains = 1.0 + strength * (allocation - 1.0) hidden = hidden * gains.to(dtype=hidden.dtype).unsqueeze(-1) return (hidden.reshape(-1, self.intermediate_size) @ down_weight).reshape(original_shape) class TileRoutedEulerGroupStateMemoryGLUMLP(TileRoutedDSwiGLUMLP): """Euler-gated full-active grouped MLP with token-conditioned group state.""" def __init__(self, config: Config, intermediate_size: int | None = None) -> None: super().__init__(config, intermediate_size) self.group_state_memory_gain_max = float(config.group_state_memory_gain_max) self.group_state_memory_gain_raw = nn.Parameter(torch.zeros(self.num_groups, 1)) def _apply_group_state_memory(self, hidden: torch.Tensor) -> torch.Tensor: hidden_groups = hidden.reshape(hidden.shape[0], self.num_groups, self.group_size) group_mean = hidden_groups.float().mean(dim=-1, keepdim=True) group_rms = hidden_groups.float().square().mean(dim=-1, keepdim=True).sqrt().clamp_min(1e-6) state = torch.tanh(group_mean / group_rms).to(dtype=hidden_groups.dtype) gain = self.group_state_memory_gain_max * torch.tanh(self.group_state_memory_gain_raw).view(1, self.num_groups, 1) return (hidden_groups * (1.0 + gain.to(dtype=hidden_groups.dtype) * state)).reshape_as(hidden) def _forward_full_active(self, x: torch.Tensor, original_shape: torch.Size, x_flat: torch.Tensor) -> torch.Tensor: hidden_size = original_shape[-1] up_weight = self.up_weight.permute(1, 0, 2).reshape(hidden_size, self.intermediate_size) value_weight = self.value_weight.permute(1, 0, 2).reshape(hidden_size, self.intermediate_size) down_weight = self.down_weight.reshape(self.intermediate_size, hidden_size) up = x_flat @ up_weight p = F.softplus(up.float()).clamp_min(torch.finfo(torch.float32).tiny) envelope = torch.exp(torch.log(p) / p - (1.0 / math.e)).to(dtype=up.dtype) hidden = F.silu(up) * envelope * (x_flat @ value_weight) hidden = self._apply_group_state_memory(hidden) return (hidden @ down_weight).reshape(original_shape) class TileRoutedEulerFusedGroupMixGLUMLP(TileRoutedDSwiGLUMLP): """Fixed Euler 4/4 MLP with group mixing folded into down weights. Mixing hidden groups before the down projection is algebraically identical to mixing the four down-weight tiles. Folding it avoids a token-side einsum and retains the dense full-active GEMM path. """ def __init__(self, config: Config, intermediate_size: int | None = None) -> None: super().__init__(config, intermediate_size) logits = torch.full((self.num_groups, self.num_groups), -6.0) logits.fill_diagonal_(6.0) self.group_mix_logits = nn.Parameter(logits) def _forward_full_active(self, x: torch.Tensor, original_shape: torch.Size, x_flat: torch.Tensor) -> torch.Tensor: hidden_size = original_shape[-1] up_weight = self.up_weight.permute(1, 0, 2).reshape(hidden_size, self.intermediate_size) value_weight = self.value_weight.permute(1, 0, 2).reshape(hidden_size, self.intermediate_size) up = x_flat @ up_weight value = x_flat @ value_weight p = F.softplus(up.float()).clamp_min(torch.finfo(torch.float32).tiny) envelope = torch.exp(torch.log(p) / p - (1.0 / math.e)).to(dtype=up.dtype) hidden = F.silu(up) * envelope * value mix = F.softmax(self.group_mix_logits.float(), dim=-1).to(dtype=hidden.dtype) effective_down = torch.einsum('og,osh->gsh', mix, self.down_weight) return (hidden @ effective_down.reshape(self.intermediate_size, hidden_size)).reshape(original_shape) class TileRoutedAdaptiveOperatorFieldEulerGLUMLP(TileRoutedDSwiGLUMLP): """Euler 4/4 MLP with a token-conditioned group operator field. A small controller reads the full residual token and selects a convex combination of 4x4 operator bases. Because the operator changes per token, it cannot be folded into the down projection like a static group mixer. """ def __init__(self, config: Config, intermediate_size: int | None = None) -> None: super().__init__(config, intermediate_size) if self.num_groups != 4 or self.active_groups != self.num_groups: raise ValueError("Adaptive operator field is specialized for full-active grouped 4/4.") controller_dim = 16 num_bases = 4 self.operator_controller = nn.Linear(config.n_embd, controller_dim, bias=False) self.operator_router = nn.Linear(controller_dim, num_bases, bias=False) self.operator_bases = nn.Parameter(torch.empty(num_bases, self.num_groups, self.num_groups)) self.operator_strength_raw = nn.Parameter(torch.tensor(-2.2)) nn.init.normal_(self.operator_controller.weight, mean=0.0, std=0.02) nn.init.normal_(self.operator_router.weight, mean=0.0, std=0.02) nn.init.normal_(self.operator_bases, mean=0.0, std=0.02) def _forward_full_active(self, x: torch.Tensor, original_shape: torch.Size, x_flat: torch.Tensor) -> torch.Tensor: hidden_size = original_shape[-1] up_weight = self.up_weight.permute(1, 0, 2).reshape(hidden_size, self.intermediate_size) value_weight = self.value_weight.permute(1, 0, 2).reshape(hidden_size, self.intermediate_size) down_weight = self.down_weight.reshape(self.intermediate_size, hidden_size) up = x_flat @ up_weight value = x_flat @ value_weight p = F.softplus(up.float()).clamp_min(torch.finfo(torch.float32).tiny) envelope = torch.exp(torch.log(p) / p - (1.0 / math.e)).to(dtype=up.dtype) hidden = (F.silu(up) * envelope * value).reshape(-1, self.num_groups, self.group_size) controller = F.rms_norm(x_flat.float(), (x_flat.shape[-1],)).to(dtype=x_flat.dtype) router = F.softmax(self.operator_router(F.silu(self.operator_controller(controller))).float(), dim=-1) delta = torch.einsum('tk,kij->tij', router, self.operator_bases.float()) strength = torch.sigmoid(self.operator_strength_raw.float()) identity = torch.eye(self.num_groups, device=hidden.device, dtype=torch.float32) operator = identity.unsqueeze(0) + strength * torch.tanh(delta) hidden = torch.bmm(operator.to(dtype=hidden.dtype), hidden) return (hidden.reshape(-1, self.intermediate_size) @ down_weight).reshape(original_shape) class TileRoutedAdaptiveBasisEulerGLUMLP(TileRoutedDSwiGLUMLP): """Euler 4/4 MLP with token-dependent low-rank feature construction. A controller selects one of two low-rank bases that update the up and value features before the Euler gate. This changes the token's feature basis and is not equivalent to a post-activation group mixer. """ def __init__(self, config: Config, intermediate_size: int | None = None) -> None: super().__init__(config, intermediate_size) if self.active_groups != self.num_groups: raise ValueError("Adaptive basis Euler MLP requires all groups to be active.") rank = 4 num_bases = 2 self.basis_input = nn.Parameter(torch.empty(config.n_embd, rank)) self.basis_up = nn.Parameter(torch.empty(num_bases, rank, self.intermediate_size)) self.basis_value = nn.Parameter(torch.empty(num_bases, rank, self.intermediate_size)) self.basis_router = nn.Linear(config.n_embd, num_bases, bias=False) self.basis_strength_raw = nn.Parameter(torch.tensor(-1.4)) nn.init.normal_(self.basis_input, mean=0.0, std=0.02) nn.init.normal_(self.basis_up, mean=0.0, std=0.02) nn.init.normal_(self.basis_value, mean=0.0, std=0.02) nn.init.zeros_(self.basis_router.weight) def _forward_full_active(self, x: torch.Tensor, original_shape: torch.Size, x_flat: torch.Tensor) -> torch.Tensor: hidden_size = original_shape[-1] up_weight = self.up_weight.permute(1, 0, 2).reshape(hidden_size, self.intermediate_size) value_weight = self.value_weight.permute(1, 0, 2).reshape(hidden_size, self.intermediate_size) down_weight = self.down_weight.reshape(self.intermediate_size, hidden_size) normalized_x = F.rms_norm(x_flat.float(), (x_flat.shape[-1],)).to(dtype=x_flat.dtype) router = F.softmax(self.basis_router(normalized_x).float(), dim=-1) latent = x_flat @ self.basis_input.to(dtype=x_flat.dtype) basis_up = torch.einsum('tk,kri->tri', router, self.basis_up.to(dtype=x_flat.dtype)) basis_value = torch.einsum('tk,kri->tri', router, self.basis_value.to(dtype=x_flat.dtype)) strength = torch.sigmoid(self.basis_strength_raw.float()).to(dtype=x_flat.dtype) up = x_flat @ up_weight + strength * torch.bmm(latent.unsqueeze(1), basis_up).squeeze(1) value = x_flat @ value_weight + strength * torch.bmm(latent.unsqueeze(1), basis_value).squeeze(1) p = F.softplus(up.float()).clamp_min(torch.finfo(torch.float32).tiny) envelope = torch.exp(torch.log(p) / p - (1.0 / math.e)).to(dtype=up.dtype) return ((F.silu(up) * envelope * value) @ down_weight).reshape(original_shape) class TileRoutedFusedAdaptiveBasisEulerGLUMLP(TileRoutedAdaptiveBasisEulerGLUMLP): """Adaptive basis Euler MLP without materializing token-specific bases.""" def _forward_full_active(self, x: torch.Tensor, original_shape: torch.Size, x_flat: torch.Tensor) -> torch.Tensor: hidden_size = original_shape[-1] up_weight = self.up_weight.permute(1, 0, 2).reshape(hidden_size, self.intermediate_size) value_weight = self.value_weight.permute(1, 0, 2).reshape(hidden_size, self.intermediate_size) down_weight = self.down_weight.reshape(self.intermediate_size, hidden_size) normalized_x = F.rms_norm(x_flat.float(), (x_flat.shape[-1],)).to(dtype=x_flat.dtype) router = F.softmax(self.basis_router(normalized_x).float(), dim=-1).to(dtype=x_flat.dtype) latent = x_flat @ self.basis_input.to(dtype=x_flat.dtype) delta_up = torch.zeros(x_flat.shape[0], self.intermediate_size, device=x_flat.device, dtype=x_flat.dtype) delta_value = torch.zeros_like(delta_up) for basis_idx in range(self.basis_up.shape[0]): weighted_latent = latent * router[:, basis_idx : basis_idx + 1] delta_up = delta_up + weighted_latent @ self.basis_up[basis_idx].to(dtype=x_flat.dtype) delta_value = delta_value + weighted_latent @ self.basis_value[basis_idx].to(dtype=x_flat.dtype) strength = torch.sigmoid(self.basis_strength_raw.float()).to(dtype=x_flat.dtype) up = x_flat @ up_weight + strength * delta_up value = x_flat @ value_weight + strength * delta_value p = F.softplus(up.float()).clamp_min(torch.finfo(torch.float32).tiny) envelope = torch.exp(torch.log(p) / p - (1.0 / math.e)).to(dtype=up.dtype) return ((F.silu(up) * envelope * value) @ down_weight).reshape(original_shape) class TileRoutedAdaptiveDirectionEulerGLUMLP(TileRoutedDSwiGLUMLP): """Euler 4/4 MLP with a token-selected rank-one feature direction. Two learned directions are mixed by a controller, then multiplied by a token-specific scalar projection. This is a conditional low-rank update to the up/value features without materializing token-specific matrices. """ def __init__(self, config: Config, intermediate_size: int | None = None) -> None: super().__init__(config, intermediate_size) if self.active_groups != self.num_groups: raise ValueError("Adaptive direction Euler MLP requires all groups to be active.") num_directions = 2 self.direction_input = nn.Parameter(torch.empty(config.n_embd)) self.direction_up = nn.Parameter(torch.empty(num_directions, self.intermediate_size)) self.direction_value = nn.Parameter(torch.empty(num_directions, self.intermediate_size)) self.direction_router = nn.Linear(config.n_embd, num_directions, bias=False) self.direction_strength_raw = nn.Parameter(torch.tensor(-2.0)) nn.init.normal_(self.direction_input, mean=0.0, std=0.02) nn.init.normal_(self.direction_up, mean=0.0, std=0.02) nn.init.normal_(self.direction_value, mean=0.0, std=0.02) nn.init.zeros_(self.direction_router.weight) def _forward_full_active(self, x: torch.Tensor, original_shape: torch.Size, x_flat: torch.Tensor) -> torch.Tensor: hidden_size = original_shape[-1] up_weight = self.up_weight.permute(1, 0, 2).reshape(hidden_size, self.intermediate_size) value_weight = self.value_weight.permute(1, 0, 2).reshape(hidden_size, self.intermediate_size) down_weight = self.down_weight.reshape(self.intermediate_size, hidden_size) normalized_x = F.rms_norm(x_flat.float(), (x_flat.shape[-1],)).to(dtype=x_flat.dtype) router = F.softmax(self.direction_router(normalized_x).float(), dim=-1).to(dtype=x_flat.dtype) token_scalar = (x_flat @ self.direction_input.to(dtype=x_flat.dtype)).unsqueeze(-1) up_direction = router @ self.direction_up.to(dtype=x_flat.dtype) value_direction = router @ self.direction_value.to(dtype=x_flat.dtype) strength = torch.sigmoid(self.direction_strength_raw.float()).to(dtype=x_flat.dtype) up = x_flat @ up_weight + strength * token_scalar * up_direction value = x_flat @ value_weight + strength * token_scalar * value_direction p = F.softplus(up.float()).clamp_min(torch.finfo(torch.float32).tiny) envelope = torch.exp(torch.log(p) / p - (1.0 / math.e)).to(dtype=up.dtype) return ((F.silu(up) * envelope * value) @ down_weight).reshape(original_shape) class TileRoutedAdaptiveDirectionEulerFusedGLUMLP(TileRoutedAdaptiveDirectionEulerGLUMLP): """Checkpoint-compatible two-route fused form of Adaptive Direction.""" def _forward_full_active(self, x: torch.Tensor, original_shape: torch.Size, x_flat: torch.Tensor) -> torch.Tensor: hidden_size = original_shape[-1] up_weight = self.up_weight.permute(1, 0, 2).reshape(hidden_size, self.intermediate_size) value_weight = self.value_weight.permute(1, 0, 2).reshape(hidden_size, self.intermediate_size) down_weight = self.down_weight.reshape(self.intermediate_size, hidden_size) normalized_x = F.rms_norm(x_flat.float(), (x_flat.shape[-1],)).to(dtype=x_flat.dtype) router_weight = self.direction_router.weight.to(dtype=x_flat.dtype) route_first = torch.sigmoid(normalized_x @ (router_weight[0] - router_weight[1])) token_scalar = (x_flat @ self.direction_input.to(dtype=x_flat.dtype)).unsqueeze(-1) up_base = self.direction_up[1].to(dtype=x_flat.dtype) value_base = self.direction_value[1].to(dtype=x_flat.dtype) up_direction = up_base + route_first.unsqueeze(-1) * (self.direction_up[0].to(dtype=x_flat.dtype) - up_base) value_direction = value_base + route_first.unsqueeze(-1) * (self.direction_value[0].to(dtype=x_flat.dtype) - value_base) strength = torch.sigmoid(self.direction_strength_raw.float()).to(dtype=x_flat.dtype) up = x_flat @ up_weight + strength * token_scalar * up_direction value = x_flat @ value_weight + strength * token_scalar * value_direction p = F.softplus(up.float()).clamp_min(torch.finfo(torch.float32).tiny) envelope = torch.exp(torch.log(p) / p - (1.0 / math.e)).to(dtype=up.dtype) return ((F.silu(up) * envelope * value) @ down_weight).reshape(original_shape) class TileRoutedGateLawDSwiGLUMLP(TileRoutedDSwiGLUMLP): """Vectorized grouped4 GLU with a compile-friendly gate law. The full-active grouped 4/4 path stays as the same three dense GEMMs as ``TileRoutedDSwiGLUMLP``. Subclasses only add cheap elementwise transforms between the up/value GEMMs and the down GEMM. """ use_powrat_gate = False use_sin_gate = False use_spon_shift = False default_spon_value = False def __init__(self, config: Config, intermediate_size: int | None = None) -> None: super().__init__(config, intermediate_size) self.block_idx = 0 self.n_layer = config.n_layer self.gate_alpha_max = float(config.grouped_gate_alpha_max) self.gate_sin_eps = float(config.grouped_gate_sin_eps) self.gate_sin_freq = float(config.grouped_gate_sin_freq) self.gate_spon_max = float(config.grouped_gate_spon_max) self.gate_layer_schedule = str(config.grouped_gate_layer_schedule) self.spon_value_enabled = bool(config.grouped_gate_spon_value or self.default_spon_value) self._gate_schedule_initialized = False if self.use_powrat_gate: self.gate_alpha_raw = nn.Parameter(torch.empty(self.num_groups, 1, 1)) self.gate_beta = nn.Parameter(torch.full((self.num_groups, 1, 1), float(config.grouped_gate_beta_init))) self._fill_alpha(float(config.grouped_gate_alpha_init)) if self.use_spon_shift: self.spon_up_shift = nn.Parameter(torch.zeros(self.intermediate_size)) self.spon_value_shift = nn.Parameter(torch.zeros(self.intermediate_size)) self.last_gate_alpha_mean: torch.Tensor | None = None self.last_gate_beta_mean: torch.Tensor | None = None self.last_spon_up_absmean: torch.Tensor | None = None self.last_spon_value_absmean: torch.Tensor | None = None self.last_gate_outlier_rate: torch.Tensor | None = None self.last_hidden_outlier_rate: torch.Tensor | None = None def _fill_alpha(self, alpha: float) -> None: alpha = max(-0.999 * self.gate_alpha_max, min(0.999 * self.gate_alpha_max, alpha)) raw = math.atanh(alpha / self.gate_alpha_max) with torch.no_grad(): self.gate_alpha_raw.fill_(raw) def set_block_index(self, block_idx: int, n_layer: int) -> None: self.block_idx = block_idx self.n_layer = n_layer if self.use_powrat_gate and not self._gate_schedule_initialized: self._fill_alpha(self._scheduled_alpha(block_idx, n_layer)) self._gate_schedule_initialized = True def _scheduled_alpha(self, block_idx: int, n_layer: int) -> float: if self.gate_layer_schedule == "constant": return float(self.config.grouped_gate_alpha_init) if self.gate_layer_schedule == "late_positive": if n_layer <= 1: return self.gate_alpha_max pos = block_idx / max(1, n_layer - 1) return 0.0 if pos < 0.67 else self.gate_alpha_max * (pos - 0.67) / 0.33 observed = [-0.074, -0.090, -0.089, -0.090, -0.084, -0.084, -0.060, -0.043, -0.020, 0.107, 0.235, 0.351] if n_layer <= 1: value = observed[-1] else: pos = block_idx * (len(observed) - 1) / max(1, n_layer - 1) lo = int(math.floor(pos)) hi = min(len(observed) - 1, lo + 1) frac = pos - lo value = observed[lo] * (1.0 - frac) + observed[hi] * frac scale = self.gate_alpha_max / max(abs(v) for v in observed) return max(-self.gate_alpha_max, min(self.gate_alpha_max, value * scale)) def _dense_weights(self, hidden_size: int) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]: up_weight = self.up_weight.permute(1, 0, 2).reshape(hidden_size, self.intermediate_size) value_weight = self.value_weight.permute(1, 0, 2).reshape(hidden_size, self.intermediate_size) down_weight = self.down_weight.reshape(self.intermediate_size, hidden_size) return up_weight, value_weight, down_weight def _spon_biases(self, dtype: torch.dtype) -> tuple[torch.Tensor | None, torch.Tensor | None]: if not self.use_spon_shift or self.gate_spon_max == 0.0: return None, None up_bias = self.spon_up_shift.to(dtype=dtype) value_bias = self.spon_value_shift.to(dtype=dtype) if self.spon_value_enabled else None return up_bias, value_bias def _gate_law(self, up: torch.Tensor) -> torch.Tensor: if not self.use_powrat_gate: return F.silu(up) up_groups = up.reshape(up.shape[0], self.num_groups, self.group_size) alpha = self.gate_alpha_max * torch.tanh( self.gate_alpha_raw.to(device=up.device, dtype=up.dtype) ).view(1, self.num_groups, 1) beta = self.gate_beta.to(device=up.device, dtype=up.dtype).abs().clamp_min(1e-4).view(1, self.num_groups, 1) modifier = 1.0 + alpha * up_groups * torch.rsqrt(1.0 + beta * up_groups.square()) gate = F.silu(up_groups) * modifier if self.use_sin_gate and self.gate_sin_eps > 0.0: gate = gate * (1.0 + self.gate_sin_eps * torch.sin(self.gate_sin_freq * up_groups)) return gate.reshape_as(up) def _forward_full_active(self, x: torch.Tensor, original_shape: torch.Size, x_flat: torch.Tensor) -> torch.Tensor: hidden_size = original_shape[-1] up_weight, value_weight, down_weight = self._dense_weights(hidden_size) up = x_flat @ up_weight value = x_flat @ value_weight up_bias, value_bias = self._spon_biases(x_flat.dtype) if up_bias is not None: up.add_(up_bias) if value_bias is not None: value.add_(value_bias) gate = self._gate_law(up) hidden = gate * value output = hidden @ down_weight if not torch.is_grad_enabled(): if self.use_powrat_gate: alpha = self.gate_alpha_max * torch.tanh(self.gate_alpha_raw.detach().float()) self.last_gate_alpha_mean = alpha.mean().to(device=x_flat.device, dtype=x_flat.dtype) self.last_gate_beta_mean = self.gate_beta.detach().float().abs().mean().to(device=x_flat.device, dtype=x_flat.dtype) if self.use_spon_shift: self.last_spon_up_absmean = self.spon_up_shift.detach().float().abs().mean().to(device=x_flat.device, dtype=x_flat.dtype) self.last_spon_value_absmean = self.spon_value_shift.detach().float().abs().mean().to(device=x_flat.device, dtype=x_flat.dtype) self.last_gate_outlier_rate = (gate.detach().float().abs() > 8.0).float().mean().to(device=x_flat.device, dtype=x_flat.dtype) self.last_hidden_outlier_rate = (hidden.detach().float().abs() > 8.0).float().mean().to(device=x_flat.device, dtype=x_flat.dtype) return output.reshape(original_shape) class TileRoutedPowRatGLUMLP(TileRoutedGateLawDSwiGLUMLP): use_powrat_gate = True class TileRoutedSinPowRatGLUMLP(TileRoutedGateLawDSwiGLUMLP): use_powrat_gate = True use_sin_gate = True class TileRoutedSPONGLUMLP(TileRoutedGateLawDSwiGLUMLP): use_spon_shift = True class TileRoutedTrainOnlySPONGLUMLP(TileRoutedSPONGLUMLP): """Use SPON as a training-time activation perturbation with free inference.""" def _forward_full_active(self, x: torch.Tensor, original_shape: torch.Size, x_flat: torch.Tensor) -> torch.Tensor: if not torch.is_grad_enabled(): return TileRoutedDSwiGLUMLP._forward_full_active(self, x, original_shape, x_flat) return super()._forward_full_active(x, original_shape, x_flat) class TileRoutedSPONPowRatGLUMLP(TileRoutedGateLawDSwiGLUMLP): use_powrat_gate = True use_spon_shift = True class TileRoutedAttentionMemoryGateDSwiGLUMLP(TileRoutedDSwiGLUMLP): """Grouped4 GLU whose gate is steered by a small learned memory attention. The main compute path keeps grouped4's full-active dense GEMMs: up/value/down are still single matmuls. The fixed SwiGLU nonlinearity is replaced by ``up * sigmoid(up + memory_delta)`` where ``memory_delta`` is a per-group token-conditioned lookup over learned hidden memory slots. """ def __init__(self, config: Config, intermediate_size: int | None = None) -> None: super().__init__(config, intermediate_size) self.attn_gate_rank = int(config.attn_gate_rank) self.attn_gate_slots = int(config.attn_gate_slots) self.attn_gate_alpha_max = float(config.attn_gate_alpha_max) self.attn_gate_temperature = float(config.attn_gate_temperature) self.attn_q_weight = nn.Parameter(torch.empty(self.num_groups, self.group_size, self.attn_gate_rank)) self.attn_k = nn.Parameter(torch.empty(self.num_groups, self.attn_gate_slots, self.attn_gate_rank)) self.attn_v = nn.Parameter(torch.empty(self.num_groups, self.attn_gate_slots, self.group_size)) self.attn_alpha_raw = nn.Parameter(torch.empty(self.num_groups, 1, 1)) self._reset_attention_memory(float(config.attn_gate_alpha_init)) self.last_attn_gate_alpha_mean: torch.Tensor | None = None self.last_attn_gate_entropy: torch.Tensor | None = None self.last_attn_gate_delta_rms: torch.Tensor | None = None self.last_attn_gate_up_rms: torch.Tensor | None = None def _reset_attention_memory(self, alpha_init: float) -> None: nn.init.normal_(self.attn_q_weight, mean=0.0, std=self.group_size**-0.5) nn.init.normal_(self.attn_k, mean=0.0, std=self.attn_gate_rank**-0.5) nn.init.normal_(self.attn_v, mean=0.0, std=0.01) alpha = max(-0.999 * self.attn_gate_alpha_max, min(0.999 * self.attn_gate_alpha_max, alpha_init)) raw = math.atanh(alpha / self.attn_gate_alpha_max) if self.attn_gate_alpha_max > 0.0 else 0.0 with torch.no_grad(): self.attn_alpha_raw.fill_(raw) def _dense_weights(self, hidden_size: int) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]: up_weight = self.up_weight.permute(1, 0, 2).reshape(hidden_size, self.intermediate_size) value_weight = self.value_weight.permute(1, 0, 2).reshape(hidden_size, self.intermediate_size) down_weight = self.down_weight.reshape(self.intermediate_size, hidden_size) return up_weight, value_weight, down_weight def _memory_delta(self, up_groups: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]: dtype = up_groups.dtype up_norm = up_groups * torch.rsqrt(up_groups.float().square().mean(dim=-1, keepdim=True).to(dtype=dtype) + 1e-6) q = torch.einsum("ngc,gcr->ngr", up_norm, self.attn_q_weight.to(device=up_groups.device, dtype=dtype)) k = self.attn_k.to(device=up_groups.device, dtype=dtype) logits = torch.einsum("ngr,gsr->ngs", q, k) logits = logits * (self.attn_gate_rank**-0.5) / max(self.attn_gate_temperature, 1e-4) attn = torch.softmax(logits.float(), dim=-1).to(dtype=dtype) delta = torch.einsum("ngs,gsc->ngc", attn, self.attn_v.to(device=up_groups.device, dtype=dtype)) return delta, attn def _forward_full_active(self, x: torch.Tensor, original_shape: torch.Size, x_flat: torch.Tensor) -> torch.Tensor: hidden_size = original_shape[-1] up_weight, value_weight, down_weight = self._dense_weights(hidden_size) up = x_flat @ up_weight value = x_flat @ value_weight up_groups = up.reshape(up.shape[0], self.num_groups, self.group_size) delta, attn = self._memory_delta(up_groups) alpha = self.attn_gate_alpha_max * torch.tanh( self.attn_alpha_raw.to(device=up.device, dtype=up.dtype) ).view(1, self.num_groups, 1) gate_groups = up_groups * torch.sigmoid(up_groups + alpha * delta) hidden = gate_groups.reshape_as(up) * value output = hidden @ down_weight if not torch.is_grad_enabled(): entropy = -(attn.float().clamp_min(1e-8) * attn.float().clamp_min(1e-8).log()).sum(dim=-1).mean() self.last_attn_gate_alpha_mean = alpha.detach().float().mean().to(device=x_flat.device, dtype=x_flat.dtype) self.last_attn_gate_entropy = entropy.to(device=x_flat.device, dtype=x_flat.dtype) self.last_attn_gate_delta_rms = delta.detach().float().square().mean().sqrt().to(device=x_flat.device, dtype=x_flat.dtype) self.last_attn_gate_up_rms = up_groups.detach().float().square().mean().sqrt().to(device=x_flat.device, dtype=x_flat.dtype) return output.reshape(original_shape) class TileRoutedInnerAttentionFFNMLP(TileRoutedDSwiGLUMLP): """Activation-free grouped FFN with an attention block inside the FFN. This removes pointwise gate activations from the full-active path. The hidden groups become four tiny tokens: ``up`` builds q/k, ``value`` builds values, group attention mixes values, and the result is down-projected. """ def __init__(self, config: Config, intermediate_size: int | None = None) -> None: super().__init__(config, intermediate_size) self.inner_attn_rank = int(config.inner_attn_rank) self.inner_attn_alpha_max = float(config.inner_attn_alpha_max) self.inner_attn_temperature = float(config.inner_attn_temperature) self.inner_q_weight = nn.Parameter(torch.empty(self.num_groups, self.group_size, self.inner_attn_rank)) self.inner_k_weight = nn.Parameter(torch.empty(self.num_groups, self.group_size, self.inner_attn_rank)) self.inner_attn_bias = nn.Parameter(torch.empty(self.num_groups, self.num_groups)) self.inner_alpha_raw = nn.Parameter(torch.empty(self.num_groups, 1, 1)) self._reset_inner_attention(float(config.inner_attn_alpha_init)) self.last_inner_attn_alpha_mean: torch.Tensor | None = None self.last_inner_attn_entropy: torch.Tensor | None = None self.last_inner_attn_offdiag: torch.Tensor | None = None self.last_inner_attn_update_rms: torch.Tensor | None = None self.last_inner_attn_value_rms: torch.Tensor | None = None def _reset_inner_attention(self, alpha_init: float) -> None: nn.init.normal_(self.inner_q_weight, mean=0.0, std=self.group_size**-0.5) nn.init.normal_(self.inner_k_weight, mean=0.0, std=self.group_size**-0.5) with torch.no_grad(): self.inner_attn_bias.fill_(-2.0) self.inner_attn_bias.diagonal().fill_(2.0) alpha = max(-0.999 * self.inner_attn_alpha_max, min(0.999 * self.inner_attn_alpha_max, alpha_init)) raw = math.atanh(alpha / self.inner_attn_alpha_max) if self.inner_attn_alpha_max > 0.0 else 0.0 with torch.no_grad(): self.inner_alpha_raw.fill_(raw) def _dense_weights(self, hidden_size: int) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]: up_weight = self.up_weight.permute(1, 0, 2).reshape(hidden_size, self.intermediate_size) value_weight = self.value_weight.permute(1, 0, 2).reshape(hidden_size, self.intermediate_size) down_weight = self.down_weight.reshape(self.intermediate_size, hidden_size) return up_weight, value_weight, down_weight def _forward_full_active(self, x: torch.Tensor, original_shape: torch.Size, x_flat: torch.Tensor) -> torch.Tensor: hidden_size = original_shape[-1] up_weight, value_weight, down_weight = self._dense_weights(hidden_size) up = x_flat @ up_weight value = x_flat @ value_weight up_groups = up.reshape(up.shape[0], self.num_groups, self.group_size) value_groups = value.reshape(value.shape[0], self.num_groups, self.group_size) dtype = up_groups.dtype up_norm = up_groups * torch.rsqrt(up_groups.float().square().mean(dim=-1, keepdim=True).to(dtype=dtype) + 1e-6) q = torch.einsum("ngc,gcr->ngr", up_norm, self.inner_q_weight.to(device=up.device, dtype=dtype)) k = torch.einsum("ngc,gcr->ngr", up_norm, self.inner_k_weight.to(device=up.device, dtype=dtype)) logits = torch.einsum("ngr,nhr->ngh", q, k) logits = logits * (self.inner_attn_rank**-0.5) / max(self.inner_attn_temperature, 1e-4) logits = logits + self.inner_attn_bias.to(device=up.device, dtype=dtype).view(1, self.num_groups, self.num_groups) attn = torch.softmax(logits.float(), dim=-1).to(dtype=dtype) mixed = torch.einsum("ngh,nhc->ngc", attn, value_groups) alpha = self.inner_attn_alpha_max * torch.tanh( self.inner_alpha_raw.to(device=up.device, dtype=dtype) ).view(1, self.num_groups, 1) hidden_groups = value_groups + alpha * (mixed - value_groups) output = hidden_groups.reshape_as(value) @ down_weight if not torch.is_grad_enabled(): probs = attn.detach().float().clamp_min(1e-8) entropy = -(probs * probs.log()).sum(dim=-1).mean() diag = torch.diagonal(attn.detach().float(), dim1=-2, dim2=-1).mean() update = (mixed - value_groups).detach().float() value_float = value_groups.detach().float() self.last_inner_attn_alpha_mean = alpha.detach().float().mean().to(device=x_flat.device, dtype=x_flat.dtype) self.last_inner_attn_entropy = entropy.to(device=x_flat.device, dtype=x_flat.dtype) self.last_inner_attn_offdiag = (1.0 - diag).to(device=x_flat.device, dtype=x_flat.dtype) self.last_inner_attn_update_rms = update.square().mean().sqrt().to(device=x_flat.device, dtype=x_flat.dtype) self.last_inner_attn_value_rms = value_float.square().mean().sqrt().to(device=x_flat.device, dtype=x_flat.dtype) return output.reshape(original_shape) class TileRoutedBilinearMemoryStepFFNMLP(TileRoutedDSwiGLUMLP): """Grouped FFN with a small bounded inner update loop instead of a gate. This is not a pointwise activation replacement. The FFN makes group hidden states, repeatedly applies a bilinear product correction, RMS-bounds that correction, shares a small mean-bus between groups, and only then writes through the down projection. """ def __init__(self, config: Config, intermediate_size: int | None = None) -> None: super().__init__(config, intermediate_size) self.bilinear_memory_depth = int(config.bilinear_memory_depth) self.bilinear_memory_alpha_max = float(config.bilinear_memory_alpha_max) self.bilinear_memory_target_ratio = float(config.bilinear_memory_target_ratio) self.bilinear_memory_bus_mix = float(config.bilinear_memory_bus_mix) self.bilinear_memory_state_update = float(config.bilinear_memory_state_update) self.memory_scale = nn.Parameter(torch.empty(1, self.num_groups, self.group_size)) self.inner_alpha_raw = nn.Parameter(torch.empty(self.num_groups, 1, 1)) self._reset_bilinear_memory(float(config.bilinear_memory_alpha_init)) self.last_bilinear_alpha_mean: torch.Tensor | None = None self.last_bilinear_update_ratio: torch.Tensor | None = None self.last_bilinear_hidden_rms: torch.Tensor | None = None self.last_bilinear_state_rms: torch.Tensor | None = None def _reset_bilinear_memory(self, alpha_init: float) -> None: nn.init.normal_(self.memory_scale, mean=1.0, std=0.02) alpha = max(-0.999 * self.bilinear_memory_alpha_max, min(0.999 * self.bilinear_memory_alpha_max, alpha_init)) raw = math.atanh(alpha / self.bilinear_memory_alpha_max) if self.bilinear_memory_alpha_max > 0.0 else 0.0 with torch.no_grad(): self.inner_alpha_raw.fill_(raw) def _dense_weights(self, hidden_size: int) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]: up_weight = self.up_weight.permute(1, 0, 2).reshape(hidden_size, self.intermediate_size) value_weight = self.value_weight.permute(1, 0, 2).reshape(hidden_size, self.intermediate_size) down_weight = self.down_weight.reshape(self.intermediate_size, hidden_size) return up_weight, value_weight, down_weight @staticmethod def _rms_norm(x: torch.Tensor) -> torch.Tensor: return x * torch.rsqrt(x.float().square().mean(dim=-1, keepdim=True).to(dtype=x.dtype) + 1e-6) def _forward_full_active(self, x: torch.Tensor, original_shape: torch.Size, x_flat: torch.Tensor) -> torch.Tensor: hidden_size = original_shape[-1] up_weight, value_weight, down_weight = self._dense_weights(hidden_size) up = x_flat @ up_weight value = x_flat @ value_weight state = self._rms_norm(up.reshape(up.shape[0], self.num_groups, self.group_size)) hidden = value.reshape(value.shape[0], self.num_groups, self.group_size) alpha = self.bilinear_memory_alpha_max * torch.tanh( self.inner_alpha_raw.to(device=up.device, dtype=up.dtype) ).view(1, self.num_groups, 1) memory = self.memory_scale.to(device=up.device, dtype=up.dtype) last_ratio = hidden.new_tensor(0.0) for _ in range(self.bilinear_memory_depth): correction = state * hidden * memory if self.bilinear_memory_bus_mix > 0.0: bus = correction.mean(dim=1, keepdim=True) correction = correction + self.bilinear_memory_bus_mix * (bus - correction) hidden_rms = hidden.float().square().mean(dim=-1, keepdim=True).sqrt().to(dtype=hidden.dtype).clamp_min(1e-6) correction_rms = correction.float().square().mean(dim=-1, keepdim=True).sqrt().to(dtype=hidden.dtype).clamp_min(1e-6) scale = (self.bilinear_memory_target_ratio * hidden_rms / correction_rms).detach() correction = correction * scale hidden = hidden + alpha * correction if self.bilinear_memory_state_update > 0.0: state = self._rms_norm(state + self.bilinear_memory_state_update * correction) if not torch.is_grad_enabled(): last_ratio = (correction.detach().float().square().mean().sqrt() / hidden.detach().float().square().mean().sqrt().clamp_min(1e-6)).to(dtype=hidden.dtype) output = hidden.reshape_as(value) @ down_weight if not torch.is_grad_enabled(): self.last_bilinear_alpha_mean = alpha.detach().float().mean().to(device=x_flat.device, dtype=x_flat.dtype) self.last_bilinear_update_ratio = last_ratio.to(device=x_flat.device, dtype=x_flat.dtype) self.last_bilinear_hidden_rms = hidden.detach().float().square().mean().sqrt().to(device=x_flat.device, dtype=x_flat.dtype) self.last_bilinear_state_rms = state.detach().float().square().mean().sqrt().to(device=x_flat.device, dtype=x_flat.dtype) return output.reshape(original_shape) class TileRoutedRationalGateDSwiGLUMLP(TileRoutedDSwiGLUMLP): """SwiGLU plus a tiny learned rational gate correction. Inspired by KAT/KAN-style learnable activations, but initialized as exactly SwiGLU so the model starts from the known-good grouped path. """ def __init__(self, config: Config, intermediate_size: int | None = None) -> None: super().__init__(config, intermediate_size) self.rational_num = nn.Parameter(torch.zeros(self.num_groups, 1, 1)) self.rational_den = nn.Parameter(torch.ones(self.num_groups, 1, 1)) def _forward_full_active(self, x: torch.Tensor, original_shape: torch.Size, x_flat: torch.Tensor) -> torch.Tensor: output = torch.zeros_like(x_flat) for group_idx in range(self.num_groups): up = x_flat @ self.up_weight[group_idx] value = x_flat @ self.value_weight[group_idx] correction = self.rational_num[group_idx] * up.square() / (1.0 + self.rational_den[group_idx].abs() * up.abs()) hidden = (F.silu(up) + correction) * value output = output + hidden @ self.down_weight[group_idx] return output.reshape(original_shape) class TileRoutedExpGateDSwiGLUMLP(TileRoutedDSwiGLUMLP): """SwiGLU with a learnable stabilized exponential gate multiplier.""" def __init__(self, config: Config, intermediate_size: int | None = None) -> None: super().__init__(config, intermediate_size) self.exp_gate_alpha = nn.Parameter(torch.zeros(self.num_groups, 1, 1)) def _forward_full_active(self, x: torch.Tensor, original_shape: torch.Size, x_flat: torch.Tensor) -> torch.Tensor: output = torch.zeros_like(x_flat) for group_idx in range(self.num_groups): up = x_flat @ self.up_weight[group_idx] value = x_flat @ self.value_weight[group_idx] multiplier = torch.exp(self.exp_gate_alpha[group_idx].to(dtype=up.dtype) * torch.tanh(up)) hidden = F.silu(up) * multiplier * value output = output + hidden @ self.down_weight[group_idx] return output.reshape(original_shape) class TileRoutedHiddenRMSNormDSwiGLUMLP(TileRoutedDSwiGLUMLP): """Zero-init hidden RMS correction before the down projection.""" def __init__(self, config: Config, intermediate_size: int | None = None) -> None: super().__init__(config, intermediate_size) self.hidden_norm_alpha = nn.Parameter(torch.zeros(self.num_groups, 1, 1)) def _forward_full_active(self, x: torch.Tensor, original_shape: torch.Size, x_flat: torch.Tensor) -> torch.Tensor: output = torch.zeros_like(x_flat) for group_idx in range(self.num_groups): hidden = F.silu(x_flat @ self.up_weight[group_idx]) * (x_flat @ self.value_weight[group_idx]) rms = hidden.float().pow(2).mean(dim=-1, keepdim=True).add(self.config.norm_eps).rsqrt().to(dtype=hidden.dtype) hidden_normed = hidden * rms alpha = self.hidden_norm_alpha[group_idx].to(device=hidden.device, dtype=hidden.dtype).clamp(-1.0, 1.0) hidden = hidden + alpha * (hidden_normed - hidden) output = output + hidden @ self.down_weight[group_idx] return output.reshape(original_shape) class TileRoutedLateTokenHiddenRMSPreserveDSwiGLUMLP(TileRoutedDSwiGLUMLP): """Late, token-gated hidden RMS correction with the grouped fastpath preserved.""" def __init__(self, config: Config, intermediate_size: int | None = None) -> None: super().__init__(config, intermediate_size) self.block_idx = 0 self.n_layer = config.n_layer self.late_layer_start = int(config.grouped_mlp_hidden_rms_late_layer_start) self.alpha_max = float(config.grouped_mlp_hidden_rms_alpha_max) alpha_init = max( -0.999 * self.alpha_max, min(0.999 * self.alpha_max, float(config.grouped_mlp_hidden_rms_alpha_init)), ) raw_alpha = math.atanh(alpha_init / self.alpha_max) self.hidden_norm_alpha_raw = nn.Parameter(torch.full((self.num_groups, 1, 1), raw_alpha)) self.hidden_norm_gate = nn.Linear(config.n_embd, self.num_groups, bias=True) nn.init.zeros_(self.hidden_norm_gate.weight) gate_init = max(1e-4, min(1.0 - 1e-4, float(config.grouped_mlp_hidden_rms_gate_init))) nn.init.constant_(self.hidden_norm_gate.bias, math.log(gate_init / (1.0 - gate_init))) self.hidden_norm_output_scale = float(config.grouped_mlp_hidden_rms_output_scale) self.hidden_norm_token_gate = bool(config.grouped_mlp_hidden_rms_token_gate) def set_block_index(self, block_idx: int, n_layer: int) -> None: self.block_idx = block_idx self.n_layer = n_layer def _forward_full_active(self, x: torch.Tensor, original_shape: torch.Size, x_flat: torch.Tensor) -> torch.Tensor: hidden_size = original_shape[-1] up_weight = self.up_weight.permute(1, 0, 2).reshape(hidden_size, self.intermediate_size) value_weight = self.value_weight.permute(1, 0, 2).reshape(hidden_size, self.intermediate_size) down_weight = self.down_weight.reshape(self.intermediate_size, hidden_size) hidden = F.silu(x_flat @ up_weight) * (x_flat @ value_weight) if self.block_idx < self.late_layer_start or self.hidden_norm_output_scale == 0.0: return (hidden @ down_weight).reshape(original_shape) hidden_groups = hidden.reshape(hidden.shape[0], self.num_groups, self.group_size) rms = hidden_groups.float().pow(2).mean(dim=-1, keepdim=True).add(self.config.norm_eps).rsqrt() hidden_normed = hidden_groups * rms.to(dtype=hidden_groups.dtype) alpha = self.alpha_max * torch.tanh( self.hidden_norm_alpha_raw.to(device=hidden.device, dtype=hidden.dtype) ).view(1, self.num_groups, 1) if self.hidden_norm_token_gate: gate = torch.sigmoid(self.hidden_norm_gate(x_flat).float()).to(dtype=hidden.dtype).unsqueeze(-1) else: gate = 1.0 correction = (hidden_normed - hidden_groups) * (alpha * gate * self.hidden_norm_output_scale) corrected_hidden = (hidden_groups + correction).reshape(hidden.shape[0], self.intermediate_size) return (corrected_hidden @ down_weight).reshape(original_shape) class TileRoutedRMSMemoryDSwiGLUMLP(TileRoutedDSwiGLUMLP): """Full-active grouped SwiGLU with normalized hidden group memory exchange. The base 4/4 grouped SwiGLU is preserved. Hidden group activations are RMS normalized, softly mixed across groups, then blended back before the down projection. This moves communication into hidden scratch space instead of adding an unbounded late residual correction. """ def __init__(self, config: Config, intermediate_size: int | None = None) -> None: super().__init__(config, intermediate_size) logits = torch.full((self.num_groups, self.num_groups), float(config.rms_memory_mix_offdiag_logit)) logits.fill_diagonal_(float(config.rms_memory_mix_identity_logit)) self.group_memory_mix_logits = nn.Parameter(logits) self.rms_memory_alpha_max = float(config.rms_memory_alpha_max) self.hidden_memory_alpha_raw = nn.Parameter(torch.zeros(self.num_groups, 1, 1)) target = float(config.rms_memory_target_init) target_min = float(config.rms_memory_target_min) target_max = float(config.rms_memory_target_max) target = min(max(target, target_min + 1e-6), target_max - 1e-6) ratio = (target - target_min) / (target_max - target_min) self.hidden_rms_target_raw = nn.Parameter(torch.full((self.num_groups, 1, 1), math.log(ratio / (1.0 - ratio)))) self.rms_memory_target_min = target_min self.rms_memory_target_max = target_max self.rms_memory_token_gate = bool(config.rms_memory_token_gate) if self.rms_memory_token_gate: self.memory_gate_weight = nn.Parameter(torch.zeros(self.num_groups, config.n_embd)) gate_init = min(max(float(config.rms_memory_gate_init), 1e-4), 1.0 - 1e-4) self.memory_gate_bias = nn.Parameter(torch.full((self.num_groups,), math.log(gate_init / (1.0 - gate_init)))) self.rms_weight_reparam = bool(config.rms_weight_reparam) if self.rms_weight_reparam: self.up_weight_gain = nn.Parameter(torch.ones(self.num_groups, 1, 1)) self.value_weight_gain = nn.Parameter(torch.ones(self.num_groups, 1, 1)) self.down_weight_gain = nn.Parameter(torch.ones(self.num_groups, 1, 1)) self.block_idx = 0 self.n_layer = config.n_layer self._scheduled_alpha_initialized = False self.last_hidden_rms_before: torch.Tensor | None = None self.last_hidden_rms_after: torch.Tensor | None = None self.last_memory_alpha_mean: torch.Tensor | None = None self.last_memory_mix_offdiag: torch.Tensor | None = None def set_block_index(self, block_idx: int, n_layer: int) -> None: self.block_idx = block_idx self.n_layer = n_layer if not self._scheduled_alpha_initialized: alpha = self._scheduled_alpha(block_idx, n_layer) alpha = max(-0.999 * self.rms_memory_alpha_max, min(0.999 * self.rms_memory_alpha_max, alpha)) with torch.no_grad(): self.hidden_memory_alpha_raw.fill_(math.atanh(alpha / self.rms_memory_alpha_max)) self._scheduled_alpha_initialized = True def _scheduled_alpha(self, block_idx: int, n_layer: int) -> float: observed = [-0.074, -0.090, -0.089, -0.090, -0.084, -0.084, -0.060, -0.043, -0.020, 0.107, 0.235, 0.351] if n_layer <= 1: return observed[-1] pos = block_idx * (len(observed) - 1) / max(1, n_layer - 1) lo = int(math.floor(pos)) hi = min(len(observed) - 1, lo + 1) frac = pos - lo return observed[lo] * (1.0 - frac) + observed[hi] * frac def _target_rms(self, hidden: torch.Tensor) -> torch.Tensor: raw = self.hidden_rms_target_raw.to(device=hidden.device, dtype=torch.float32) target = self.rms_memory_target_min + (self.rms_memory_target_max - self.rms_memory_target_min) * torch.sigmoid(raw) return target.to(dtype=hidden.dtype).view(1, self.num_groups, 1) def _maybe_reparam_weight(self, weight: torch.Tensor, gain: torch.Tensor | None, target_std: float) -> torch.Tensor: if not self.rms_weight_reparam or gain is None: return weight rms = weight.float().pow(2).mean(dim=(-2, -1), keepdim=True).sqrt().clamp_min(1e-8) scale = (target_std / rms).to(device=weight.device, dtype=weight.dtype) return weight * scale * gain.to(device=weight.device, dtype=weight.dtype) def _weights(self) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]: std = math.sqrt(2.0 / 5 / self.config.n_embd) down_std = 1.0 / math.sqrt(self.config.n_embd) / self.config.n_layer up = self._maybe_reparam_weight(self.up_weight, getattr(self, "up_weight_gain", None), std) value = self._maybe_reparam_weight(self.value_weight, getattr(self, "value_weight_gain", None), std) down = self._maybe_reparam_weight(self.down_weight, getattr(self, "down_weight_gain", None), down_std) return up, value, down def _token_gate(self, x_flat: torch.Tensor) -> torch.Tensor | float: if not self.rms_memory_token_gate: return 1.0 x_rms = x_flat.float().pow(2).mean(dim=-1, keepdim=True).sqrt().clamp_min(1e-8) x_normed = (x_flat.float() / x_rms).to(dtype=x_flat.dtype) logits = F.linear( x_normed, self.memory_gate_weight.to(device=x_flat.device, dtype=x_flat.dtype), self.memory_gate_bias.to(device=x_flat.device, dtype=x_flat.dtype), ) return torch.sigmoid(logits.float()).to(dtype=x_flat.dtype).unsqueeze(-1) def _forward_full_active(self, x: torch.Tensor, original_shape: torch.Size, x_flat: torch.Tensor) -> torch.Tensor: hidden_size = original_shape[-1] up_weight, value_weight, down_weight = self._weights() up_dense = up_weight.permute(1, 0, 2).reshape(hidden_size, self.intermediate_size) value_dense = value_weight.permute(1, 0, 2).reshape(hidden_size, self.intermediate_size) hidden = F.silu(x_flat @ up_dense) * (x_flat @ value_dense) hidden_groups = hidden.reshape(hidden.shape[0], self.num_groups, self.group_size) hidden_rms = hidden_groups.float().pow(2).mean(dim=-1, keepdim=True).sqrt() hidden_normed = hidden_groups * torch.rsqrt( hidden_groups.float().pow(2).mean(dim=-1, keepdim=True) + self.config.norm_eps ).to(dtype=hidden_groups.dtype) hidden_normed = hidden_normed * self._target_rms(hidden_groups) mix = torch.softmax(self.group_memory_mix_logits.float(), dim=-1).to(device=hidden.device, dtype=hidden.dtype) mixed = torch.einsum("og,ngs->nos", mix, hidden_normed) alpha = self.rms_memory_alpha_max * torch.tanh( self.hidden_memory_alpha_raw.to(device=hidden.device, dtype=hidden.dtype) ).view(1, self.num_groups, 1) gate = self._token_gate(x_flat) hidden_groups = hidden_groups + alpha * gate * (mixed - hidden_groups) down_dense = down_weight.reshape(self.intermediate_size, hidden_size) output = hidden_groups.reshape(hidden.shape[0], self.intermediate_size) @ down_dense if not torch.is_grad_enabled(): self.last_hidden_rms_before = hidden_rms.detach().mean() self.last_hidden_rms_after = hidden_groups.float().pow(2).mean(dim=-1, keepdim=True).sqrt().detach().mean() self.last_memory_alpha_mean = alpha.detach().mean() eye = torch.eye(self.num_groups, device=mix.device, dtype=torch.bool) self.last_memory_mix_offdiag = mix.masked_select(~eye).detach().mean() return output.reshape(original_shape) class TileRoutedFastRMSMemoryDSwiGLUMLP(TileRoutedDSwiGLUMLP): """Speed-first hidden RMS memory for full-active grouped SwiGLU. The early layers run the exact grouped4 dense-GEMM fastpath. Enabled late layers keep the same up/value/down GEMMs, but insert a cheap hidden-space mean-bus exchange before the down projection. The exchange preserves each group's hidden RMS, so it changes direction more than magnitude. """ def __init__(self, config: Config, intermediate_size: int | None = None) -> None: super().__init__(config, intermediate_size) self.fast_rms_memory_schedule = str(config.fast_rms_memory_schedule) self.fast_rms_memory_alpha_max = float(config.fast_rms_memory_alpha_max) self.fast_rms_memory_use_rms_norm = bool(config.fast_rms_memory_use_rms_norm) self.hidden_memory_alpha_raw = nn.Parameter(torch.zeros(self.num_groups, 1, 1)) target = float(config.rms_memory_target_init) target_min = float(config.rms_memory_target_min) target_max = float(config.rms_memory_target_max) target = min(max(target, target_min + 1e-6), target_max - 1e-6) ratio = (target - target_min) / (target_max - target_min) self.hidden_rms_target_raw = nn.Parameter(torch.full((self.num_groups, 1, 1), math.log(ratio / (1.0 - ratio)))) self.rms_memory_target_min = target_min self.rms_memory_target_max = target_max self.block_idx = 0 self.n_layer = config.n_layer self.memory_enabled = False self._scheduled_alpha_initialized = False self.last_hidden_rms_before: torch.Tensor | None = None self.last_hidden_rms_after: torch.Tensor | None = None self.last_memory_alpha_mean: torch.Tensor | None = None self.last_memory_bus_rms: torch.Tensor | None = None def set_block_index(self, block_idx: int, n_layer: int) -> None: self.block_idx = block_idx self.n_layer = n_layer if not self._scheduled_alpha_initialized: alpha = self._scheduled_alpha(block_idx, n_layer) self.memory_enabled = abs(alpha) > 0.0 alpha = max(-0.999 * self.fast_rms_memory_alpha_max, min(0.999 * self.fast_rms_memory_alpha_max, alpha)) with torch.no_grad(): self.hidden_memory_alpha_raw.fill_(math.atanh(alpha / self.fast_rms_memory_alpha_max)) self._scheduled_alpha_initialized = True def _scheduled_alpha(self, block_idx: int, n_layer: int) -> float: if self.fast_rms_memory_schedule == "late8_meanbus": if n_layer == 12: schedule = [0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.05, 0.10, 0.18, 0.25] return schedule[block_idx] start = max(0, n_layer - 4) late = [0.05, 0.10, 0.18, 0.25] return late[block_idx - start] if block_idx >= start else 0.0 if self.fast_rms_memory_schedule == "late9_meanbus": if n_layer == 12: schedule = [0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.08, 0.16, 0.25] return schedule[block_idx] start = max(0, n_layer - 3) late = [0.08, 0.16, 0.25] return late[block_idx - start] if block_idx >= start else 0.0 if self.fast_rms_memory_schedule == "all_tiny_meanbus": observed = [-0.074, -0.090, -0.089, -0.090, -0.084, -0.084, -0.060, -0.043, -0.020, 0.107, 0.235, 0.351] if n_layer <= 1: value = observed[-1] else: pos = block_idx * (len(observed) - 1) / max(1, n_layer - 1) lo = int(math.floor(pos)) hi = min(len(observed) - 1, lo + 1) frac = pos - lo value = observed[lo] * (1.0 - frac) + observed[hi] * frac return max(-0.03, min(0.08, value)) raise ValueError(f"unknown fast_rms_memory_schedule={self.fast_rms_memory_schedule!r}") def _target_scale(self, hidden: torch.Tensor) -> torch.Tensor: raw = self.hidden_rms_target_raw.to(device=hidden.device, dtype=torch.float32) target = self.rms_memory_target_min + (self.rms_memory_target_max - self.rms_memory_target_min) * torch.sigmoid(raw) return target.to(dtype=hidden.dtype).view(1, self.num_groups, 1) def _forward_full_active(self, x: torch.Tensor, original_shape: torch.Size, x_flat: torch.Tensor) -> torch.Tensor: if not self.memory_enabled: return super()._forward_full_active(x, original_shape, x_flat) hidden_size = original_shape[-1] up_weight = self.up_weight.permute(1, 0, 2).reshape(hidden_size, self.intermediate_size) value_weight = self.value_weight.permute(1, 0, 2).reshape(hidden_size, self.intermediate_size) hidden = F.silu(x_flat @ up_weight) * (x_flat @ value_weight) alpha = self.fast_rms_memory_alpha_max * torch.tanh( self.hidden_memory_alpha_raw.to(device=hidden.device, dtype=hidden.dtype) ) if not self.fast_rms_memory_use_rms_norm: # Linear mean-bus exchange can be folded into the down projection: # mixed_o = (1-a_o) h_o + a_o * mean_g(h_g). alpha_weight = alpha.to(device=self.down_weight.device, dtype=self.down_weight.dtype) shared_down = (alpha_weight * self.down_weight).mean(dim=0, keepdim=True) effective_down = (1.0 - alpha_weight) * self.down_weight + shared_down output = hidden @ effective_down.reshape(self.intermediate_size, hidden_size) if not torch.is_grad_enabled(): self.last_hidden_rms_before = hidden.float().pow(2).mean(dim=-1).sqrt().detach().mean() self.last_hidden_rms_after = self.last_hidden_rms_before self.last_memory_alpha_mean = alpha.detach().mean() self.last_memory_bus_rms = hidden.new_tensor(0.0) return output.reshape(original_shape) hidden_groups = hidden.reshape(hidden.shape[0], self.num_groups, self.group_size) hidden_var = hidden_groups.float().pow(2).mean(dim=-1, keepdim=True) hidden_rms = torch.sqrt(hidden_var + self.config.norm_eps) hidden_normed = hidden_groups * torch.rsqrt(hidden_var + self.config.norm_eps).to(dtype=hidden_groups.dtype) bus = hidden_normed.mean(dim=1, keepdim=True) bus_var = bus.float().pow(2).mean(dim=-1, keepdim=True) bus = bus * torch.rsqrt(bus_var + self.config.norm_eps).to(dtype=bus.dtype) memory = bus * hidden_rms.to(dtype=hidden_groups.dtype) * self._target_scale(hidden_groups) alpha = alpha.view(1, self.num_groups, 1) hidden_groups = hidden_groups + alpha * (memory - hidden_groups) down_weight = self.down_weight.reshape(self.intermediate_size, hidden_size) output = hidden_groups.reshape(hidden.shape[0], self.intermediate_size) @ down_weight if not torch.is_grad_enabled(): self.last_hidden_rms_before = hidden_rms.detach().mean() self.last_hidden_rms_after = hidden_groups.float().pow(2).mean(dim=-1, keepdim=True).sqrt().detach().mean() self.last_memory_alpha_mean = alpha.detach().mean() self.last_memory_bus_rms = bus.float().pow(2).mean(dim=-1, keepdim=True).sqrt().detach().mean() return output.reshape(original_shape) class TileRoutedCoAdaptHiddenControllerDSwiGLUMLP(TileRoutedDSwiGLUMLP): """Late-only train-time hidden controller for grouped4. This keeps grouped4's full-active up/value/down GEMM path. In late layers, a zero-init low-rank controller can add a tiny RMS-bounded hidden-space correction before the down projection. The design is the train-time, co-adapted version of the oracle-positive hidden write probe; at initialization it is exactly the grouped4 fastpath. """ def __init__(self, config: Config, intermediate_size: int | None = None) -> None: super().__init__(config, intermediate_size) self.block_idx = 0 self.n_layer = config.n_layer self.controller_rank = int(config.hidden_controller_rank) self.controller_rho_frac = float(config.hidden_controller_rho_frac) self.controller_late_layer_start = int(config.hidden_controller_late_layer_start) self.controller_down = nn.Parameter(torch.empty(self.num_groups, self.group_size, self.controller_rank)) self.controller_up = nn.Parameter(torch.zeros(self.num_groups, self.controller_rank, self.group_size)) nn.init.normal_(self.controller_down, mean=0.0, std=self.group_size**-0.5) self.last_hidden_controller_ratio: torch.Tensor | None = None self.last_hidden_controller_raw_rms: torch.Tensor | None = None def set_block_index(self, block_idx: int, n_layer: int) -> None: self.block_idx = block_idx self.n_layer = n_layer enabled = self._controller_enabled() self.controller_down.requires_grad_(enabled) self.controller_up.requires_grad_(enabled) def _controller_enabled(self) -> bool: return self.controller_rho_frac > 0.0 and self.block_idx >= self.controller_late_layer_start def _hidden_controller(self, hidden_groups: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]: dtype = hidden_groups.dtype hidden_rms = hidden_groups.float().square().mean(dim=-1, keepdim=True).sqrt().to(dtype=dtype) h_norm = hidden_groups * torch.rsqrt( hidden_groups.float().square().mean(dim=-1, keepdim=True).to(dtype=dtype) + self.config.norm_eps ) low = torch.einsum("ngs,gsr->ngr", h_norm, self.controller_down.to(dtype=dtype)) raw = torch.einsum("ngr,grs->ngs", torch.tanh(low), self.controller_up.to(dtype=dtype)) raw_rms = torch.sqrt(raw.float().square().mean(dim=-1, keepdim=True) + self.config.norm_eps) raw_rms = raw_rms.to(dtype=dtype) target = self.controller_rho_frac * hidden_rms correction = raw * (target / raw_rms).clamp(max=1.0) ratio = correction.float().square().mean().sqrt() / hidden_groups.float().square().mean().sqrt().clamp_min(1e-8) return hidden_groups + correction, ratio, raw_rms.detach().mean() def _forward_full_active(self, x: torch.Tensor, original_shape: torch.Size, x_flat: torch.Tensor) -> torch.Tensor: if not self._controller_enabled(): return super()._forward_full_active(x, original_shape, x_flat) hidden_size = original_shape[-1] up_weight = self.up_weight.permute(1, 0, 2).reshape(hidden_size, self.intermediate_size) value_weight = self.value_weight.permute(1, 0, 2).reshape(hidden_size, self.intermediate_size) down_weight = self.down_weight.reshape(self.intermediate_size, hidden_size) hidden = F.silu(x_flat @ up_weight) * (x_flat @ value_weight) hidden_groups = hidden.reshape(hidden.shape[0], self.num_groups, self.group_size) hidden_groups, ratio, raw_rms = self._hidden_controller(hidden_groups) if not torch.is_grad_enabled(): self.last_hidden_controller_ratio = ratio.detach() self.last_hidden_controller_raw_rms = raw_rms.detach() output = hidden_groups.reshape(hidden.shape[0], self.intermediate_size) @ down_weight return output.reshape(original_shape) class TileRoutedHiddenDirectionMemoryDSwiGLUMLP(TileRoutedDSwiGLUMLP): """Late-only hidden key/value direction memory for grouped4. Each group owns a tiny bank of correction directions. Hidden activations query the bank, select a value direction with softmax, and commit only an RMS-bounded correction before the down projection. Values are zero-init, so the model starts exactly as grouped4 and must learn useful directions during pretraining instead of receiving an unbounded residual update. """ def __init__(self, config: Config, intermediate_size: int | None = None) -> None: super().__init__(config, intermediate_size) self.block_idx = 0 self.n_layer = config.n_layer self.memory_slots = int(config.hidden_direction_memory_slots) self.memory_rho_frac = float(config.hidden_direction_memory_rho_frac) self.memory_late_layer_start = int(config.hidden_direction_memory_late_layer_start) self.memory_temperature = float(config.hidden_direction_memory_temperature) self.memory_key = nn.Parameter(torch.empty(self.num_groups, self.group_size, self.memory_slots)) self.memory_value = nn.Parameter(torch.zeros(self.num_groups, self.memory_slots, self.group_size)) nn.init.normal_(self.memory_key, mean=0.0, std=self.group_size**-0.5) self.last_direction_memory_entropy: torch.Tensor | None = None self.last_direction_memory_ratio: torch.Tensor | None = None def set_block_index(self, block_idx: int, n_layer: int) -> None: self.block_idx = block_idx self.n_layer = n_layer enabled = self._memory_enabled() self.memory_key.requires_grad_(enabled) self.memory_value.requires_grad_(enabled) def _memory_enabled(self) -> bool: return self.memory_rho_frac > 0.0 and self.block_idx >= self.memory_late_layer_start def _direction_memory(self, hidden_groups: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]: dtype = hidden_groups.dtype hidden_rms = hidden_groups.float().square().mean(dim=-1, keepdim=True).sqrt().to(dtype=dtype) h_norm = hidden_groups * torch.rsqrt( hidden_groups.float().square().mean(dim=-1, keepdim=True).to(dtype=dtype) + self.config.norm_eps ) key = self.memory_key.to(dtype=dtype) value = self.memory_value.to(dtype=dtype) logits = torch.einsum("ngs,gsm->ngm", h_norm, key) logits = logits / (math.sqrt(self.group_size) * self.memory_temperature) weights = torch.softmax(logits.float(), dim=-1).to(dtype=dtype) raw = torch.einsum("ngm,gms->ngs", weights, value) raw_rms = torch.sqrt(raw.float().square().mean(dim=-1, keepdim=True) + self.config.norm_eps).to(dtype=dtype) target = self.memory_rho_frac * hidden_rms correction = raw * (target / raw_rms).clamp(max=1.0) ratio = correction.float().square().mean().sqrt() / hidden_groups.float().square().mean().sqrt().clamp_min(1e-8) entropy = -(weights.float() * weights.float().clamp_min(1e-8).log()).sum(dim=-1).mean() return hidden_groups + correction, ratio, entropy def _forward_full_active(self, x: torch.Tensor, original_shape: torch.Size, x_flat: torch.Tensor) -> torch.Tensor: if not self._memory_enabled(): return super()._forward_full_active(x, original_shape, x_flat) hidden_size = original_shape[-1] up_weight = self.up_weight.permute(1, 0, 2).reshape(hidden_size, self.intermediate_size) value_weight = self.value_weight.permute(1, 0, 2).reshape(hidden_size, self.intermediate_size) down_weight = self.down_weight.reshape(self.intermediate_size, hidden_size) hidden = F.silu(x_flat @ up_weight) * (x_flat @ value_weight) hidden_groups = hidden.reshape(hidden.shape[0], self.num_groups, self.group_size) hidden_groups, ratio, entropy = self._direction_memory(hidden_groups) if not torch.is_grad_enabled(): self.last_direction_memory_ratio = ratio.detach() self.last_direction_memory_entropy = entropy.detach() output = hidden_groups.reshape(hidden.shape[0], self.intermediate_size) @ down_weight return output.reshape(original_shape) class TileRoutedScalarDirectionMemoryDSwiGLUMLP(TileRoutedDSwiGLUMLP): """Speed-first hidden direction memory with no extra hidden matmul. Each group stores one learned direction. A token reads that memory by cosine-like alignment with its normalized hidden state, then writes a tiny RMS-capped signed correction before the down projection. This is less expressive than a key/value bank but keeps the extra work to reductions and elementwise ops in late layers. """ def __init__(self, config: Config, intermediate_size: int | None = None) -> None: super().__init__(config, intermediate_size) self.block_idx = 0 self.n_layer = config.n_layer self.scalar_memory_rho_frac = float(config.scalar_direction_memory_rho_frac) self.scalar_memory_late_layer_start = int(config.scalar_direction_memory_late_layer_start) self.scalar_memory_direction = nn.Parameter(torch.empty(self.num_groups, self.group_size)) self.scalar_memory_alpha_raw = nn.Parameter(torch.zeros(self.num_groups, 1, 1)) nn.init.normal_(self.scalar_memory_direction, mean=0.0, std=1.0) self.last_scalar_memory_alpha: torch.Tensor | None = None self.last_scalar_memory_gate_abs: torch.Tensor | None = None self.last_scalar_memory_ratio: torch.Tensor | None = None def set_block_index(self, block_idx: int, n_layer: int) -> None: self.block_idx = block_idx self.n_layer = n_layer enabled = self._memory_enabled() self.scalar_memory_direction.requires_grad_(enabled) self.scalar_memory_alpha_raw.requires_grad_(enabled) def _memory_enabled(self) -> bool: return self.scalar_memory_rho_frac > 0.0 and self.block_idx >= self.scalar_memory_late_layer_start def _scalar_memory(self, hidden_groups: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]: dtype = hidden_groups.dtype hidden_rms = hidden_groups.float().square().mean(dim=-1, keepdim=True).sqrt().to(dtype=dtype) h_norm = hidden_groups * torch.rsqrt( hidden_groups.float().square().mean(dim=-1, keepdim=True).to(dtype=dtype) + self.config.norm_eps ) direction = self.scalar_memory_direction.to(dtype=dtype) direction = direction * torch.rsqrt(direction.float().square().mean(dim=-1, keepdim=True).to(dtype=dtype) + self.config.norm_eps) direction = direction.unsqueeze(0) gate = torch.tanh((h_norm * direction).mean(dim=-1, keepdim=True)) alpha = torch.tanh(self.scalar_memory_alpha_raw.to(dtype=dtype)).view(1, self.num_groups, 1) correction = alpha * gate * direction * (self.scalar_memory_rho_frac * hidden_rms) ratio = correction.float().square().mean().sqrt() / hidden_groups.float().square().mean().sqrt().clamp_min(1e-8) return hidden_groups + correction, ratio, alpha.detach().abs().mean(), gate.detach().abs().mean() def _forward_full_active(self, x: torch.Tensor, original_shape: torch.Size, x_flat: torch.Tensor) -> torch.Tensor: if not self._memory_enabled(): return super()._forward_full_active(x, original_shape, x_flat) hidden_size = original_shape[-1] up_weight = self.up_weight.permute(1, 0, 2).reshape(hidden_size, self.intermediate_size) value_weight = self.value_weight.permute(1, 0, 2).reshape(hidden_size, self.intermediate_size) down_weight = self.down_weight.reshape(self.intermediate_size, hidden_size) hidden = F.silu(x_flat @ up_weight) * (x_flat @ value_weight) hidden_groups = hidden.reshape(hidden.shape[0], self.num_groups, self.group_size) hidden_groups, ratio, alpha_abs, gate_abs = self._scalar_memory(hidden_groups) if not torch.is_grad_enabled(): self.last_scalar_memory_ratio = ratio.detach() self.last_scalar_memory_alpha = alpha_abs.detach() self.last_scalar_memory_gate_abs = gate_abs.detach() output = hidden_groups.reshape(hidden.shape[0], self.intermediate_size) @ down_weight return output.reshape(original_shape) class TileRoutedChannelMemoryDSwiGLUMLP(TileRoutedDSwiGLUMLP): """Foldable FFN memory-value gain for grouped4. FFN hidden channels act like keys; down-projection rows are the memory values. This class learns a bounded per-channel value strength. It starts exactly as grouped4 and can be folded into ``down_weight`` for inference, so it changes what the FFN remembers without adding a new token-time reasoning path. """ def __init__(self, config: Config, intermediate_size: int | None = None) -> None: super().__init__(config, intermediate_size) self.block_idx = 0 self.n_layer = config.n_layer self.channel_memory_gain_max = float(config.channel_memory_gain_max) self.channel_memory_late_layer_start = int(config.channel_memory_late_layer_start) self.channel_memory_gain_raw = nn.Parameter(torch.zeros(self.num_groups, self.group_size, 1)) self.last_channel_memory_gain_abs: torch.Tensor | None = None def set_block_index(self, block_idx: int, n_layer: int) -> None: self.block_idx = block_idx self.n_layer = n_layer self.channel_memory_gain_raw.requires_grad_(self._memory_enabled()) def _memory_enabled(self) -> bool: return self.channel_memory_gain_max > 0.0 and self.block_idx >= self.channel_memory_late_layer_start def _effective_down_weight(self) -> torch.Tensor: if not self._memory_enabled(): return self.down_weight gain = self.channel_memory_gain_max * torch.tanh( self.channel_memory_gain_raw.to(device=self.down_weight.device, dtype=self.down_weight.dtype) ) if not torch.is_grad_enabled(): self.last_channel_memory_gain_abs = gain.detach().abs().mean() return self.down_weight * (1.0 + gain) def _forward_full_active(self, x: torch.Tensor, original_shape: torch.Size, x_flat: torch.Tensor) -> torch.Tensor: hidden_size = original_shape[-1] up_weight = self.up_weight.permute(1, 0, 2).reshape(hidden_size, self.intermediate_size) value_weight = self.value_weight.permute(1, 0, 2).reshape(hidden_size, self.intermediate_size) down_weight = self._effective_down_weight().reshape(self.intermediate_size, hidden_size) hidden = F.silu(x_flat @ up_weight) * (x_flat @ value_weight) return (hidden @ down_weight).reshape(original_shape) class TileRoutedChannelStateMemoryDSwiGLUMLP(TileRoutedDSwiGLUMLP): """Per-channel token-conditioned memory for grouped4. The previous scalar group-state memory could move argmaxes but only had one learned gain per group. This variant keeps the same cheap group-state signal but gives each hidden channel its own bounded value-control coefficient. Zero gain is exactly grouped4. """ def __init__(self, config: Config, intermediate_size: int | None = None) -> None: super().__init__(config, intermediate_size) self.block_idx = 0 self.n_layer = config.n_layer self.channel_memory_gain_max = float(config.channel_memory_gain_max) self.channel_memory_late_layer_start = int(config.channel_memory_late_layer_start) self.channel_memory_gain_raw = nn.Parameter(torch.zeros(self.num_groups, self.group_size)) self.last_channel_state_gain_abs: torch.Tensor | None = None self.last_channel_state_abs: torch.Tensor | None = None def set_block_index(self, block_idx: int, n_layer: int) -> None: self.block_idx = block_idx self.n_layer = n_layer self.channel_memory_gain_raw.requires_grad_(self._memory_enabled()) def _memory_enabled(self) -> bool: return self.channel_memory_gain_max > 0.0 and self.block_idx >= self.channel_memory_late_layer_start def _apply_channel_state_memory(self, hidden: torch.Tensor) -> torch.Tensor: if not self._memory_enabled(): return hidden hidden_groups = hidden.reshape(hidden.shape[0], self.num_groups, self.group_size) hidden_float = hidden_groups.float() group_mean = hidden_float.mean(dim=-1, keepdim=True) group_rms = hidden_float.square().mean(dim=-1, keepdim=True).sqrt().clamp_min(1e-6) group_state = torch.tanh(group_mean / group_rms).to(dtype=hidden_groups.dtype) gain = self.channel_memory_gain_max * torch.tanh( self.channel_memory_gain_raw.to(device=hidden.device, dtype=hidden_groups.dtype) ).unsqueeze(0) if not torch.is_grad_enabled(): self.last_channel_state_gain_abs = gain.detach().abs().mean() self.last_channel_state_abs = group_state.detach().abs().mean() hidden_groups = hidden_groups * (1.0 + gain * group_state) return hidden_groups.reshape(hidden.shape) def _forward_full_active(self, x: torch.Tensor, original_shape: torch.Size, x_flat: torch.Tensor) -> torch.Tensor: hidden_size = original_shape[-1] up_weight = self.up_weight.permute(1, 0, 2).reshape(hidden_size, self.intermediate_size) value_weight = self.value_weight.permute(1, 0, 2).reshape(hidden_size, self.intermediate_size) down_weight = self.down_weight.reshape(self.intermediate_size, hidden_size) hidden = F.silu(x_flat @ up_weight) * (x_flat @ value_weight) hidden = self._apply_channel_state_memory(hidden) return (hidden @ down_weight).reshape(original_shape) class TileRoutedCenteredChannelStateMemoryDSwiGLUMLP(TileRoutedDSwiGLUMLP): """Mean-preserving per-channel token-conditioned memory for grouped4.""" def __init__(self, config: Config, intermediate_size: int | None = None) -> None: super().__init__(config, intermediate_size) self.block_idx = 0 self.n_layer = config.n_layer self.channel_memory_gain_max = float(config.channel_memory_gain_max) self.channel_memory_late_layer_start = int(config.channel_memory_late_layer_start) self.channel_memory_gain_raw = nn.Parameter(torch.zeros(self.num_groups, self.group_size)) self.last_channel_state_gain_abs: torch.Tensor | None = None self.last_channel_state_abs: torch.Tensor | None = None def set_block_index(self, block_idx: int, n_layer: int) -> None: self.block_idx = block_idx self.n_layer = n_layer self.channel_memory_gain_raw.requires_grad_(self._memory_enabled()) def _memory_enabled(self) -> bool: return self.channel_memory_gain_max > 0.0 and self.block_idx >= self.channel_memory_late_layer_start def _apply_centered_channel_state_memory(self, hidden: torch.Tensor) -> torch.Tensor: if not self._memory_enabled(): return hidden hidden_groups = hidden.reshape(hidden.shape[0], self.num_groups, self.group_size) hidden_float = hidden_groups.float() group_mean = hidden_float.mean(dim=-1, keepdim=True) centered = hidden_groups - group_mean.to(dtype=hidden_groups.dtype) group_rms = hidden_float.square().mean(dim=-1, keepdim=True).sqrt().clamp_min(1e-6) group_state = torch.tanh(group_mean / group_rms).to(dtype=hidden_groups.dtype) gain = self.channel_memory_gain_max * torch.tanh( self.channel_memory_gain_raw.to(device=hidden.device, dtype=hidden_groups.dtype) ).unsqueeze(0) if not torch.is_grad_enabled(): self.last_channel_state_gain_abs = gain.detach().abs().mean() self.last_channel_state_abs = group_state.detach().abs().mean() hidden_groups = hidden_groups + centered * (gain * group_state) return hidden_groups.reshape(hidden.shape) def _forward_full_active(self, x: torch.Tensor, original_shape: torch.Size, x_flat: torch.Tensor) -> torch.Tensor: hidden_size = original_shape[-1] up_weight = self.up_weight.permute(1, 0, 2).reshape(hidden_size, self.intermediate_size) value_weight = self.value_weight.permute(1, 0, 2).reshape(hidden_size, self.intermediate_size) down_weight = self.down_weight.reshape(self.intermediate_size, hidden_size) hidden = F.silu(x_flat @ up_weight) * (x_flat @ value_weight) hidden = self._apply_centered_channel_state_memory(hidden) return (hidden @ down_weight).reshape(original_shape) class TileRoutedCenteredChannelGroupStateMemoryDSwiGLUMLP(TileRoutedDSwiGLUMLP): """Mean-preserving channel memory plus scalar group-state memory.""" def __init__(self, config: Config, intermediate_size: int | None = None) -> None: super().__init__(config, intermediate_size) self.block_idx = 0 self.n_layer = config.n_layer self.channel_memory_gain_max = float(config.channel_memory_gain_max) self.channel_memory_late_layer_start = int(config.channel_memory_late_layer_start) self.group_state_memory_gain_max = float(config.group_state_memory_gain_max) self.group_state_memory_late_layer_start = int(config.group_state_memory_late_layer_start) self.channel_memory_gain_raw = nn.Parameter(torch.zeros(self.num_groups, self.group_size)) self.group_state_memory_gain_raw = nn.Parameter(torch.zeros(self.num_groups, 1)) self.last_channel_state_gain_abs: torch.Tensor | None = None self.last_group_state_gain_abs: torch.Tensor | None = None self.last_group_state_abs: torch.Tensor | None = None def set_block_index(self, block_idx: int, n_layer: int) -> None: self.block_idx = block_idx self.n_layer = n_layer self.channel_memory_gain_raw.requires_grad_(self._channel_memory_enabled()) self.group_state_memory_gain_raw.requires_grad_(self._group_state_memory_enabled()) def _channel_memory_enabled(self) -> bool: return self.channel_memory_gain_max > 0.0 and self.block_idx >= self.channel_memory_late_layer_start def _group_state_memory_enabled(self) -> bool: return self.group_state_memory_gain_max > 0.0 and self.block_idx >= self.group_state_memory_late_layer_start def _apply_memory(self, hidden: torch.Tensor) -> torch.Tensor: if not self._channel_memory_enabled() and not self._group_state_memory_enabled(): return hidden hidden_groups = hidden.reshape(hidden.shape[0], self.num_groups, self.group_size) hidden_float = hidden_groups.float() group_mean = hidden_float.mean(dim=-1, keepdim=True) group_rms = hidden_float.square().mean(dim=-1, keepdim=True).sqrt().clamp_min(1e-6) group_state = torch.tanh(group_mean / group_rms).to(dtype=hidden_groups.dtype) if self._channel_memory_enabled(): centered = hidden_groups - group_mean.to(dtype=hidden_groups.dtype) channel_gain = self.channel_memory_gain_max * torch.tanh( self.channel_memory_gain_raw.to(device=hidden.device, dtype=hidden_groups.dtype) ).unsqueeze(0) if not torch.is_grad_enabled(): self.last_channel_state_gain_abs = channel_gain.detach().abs().mean() hidden_groups = hidden_groups + centered * (channel_gain * group_state) if self._group_state_memory_enabled(): group_gain = self.group_state_memory_gain_max * torch.tanh( self.group_state_memory_gain_raw.to(device=hidden.device, dtype=hidden_groups.dtype) ).view(1, self.num_groups, 1) if not torch.is_grad_enabled(): self.last_group_state_gain_abs = group_gain.detach().abs().mean() self.last_group_state_abs = group_state.detach().abs().mean() hidden_groups = hidden_groups * (1.0 + group_gain * group_state) return hidden_groups.reshape(hidden.shape) def _forward_full_active(self, x: torch.Tensor, original_shape: torch.Size, x_flat: torch.Tensor) -> torch.Tensor: hidden_size = original_shape[-1] up_weight = self.up_weight.permute(1, 0, 2).reshape(hidden_size, self.intermediate_size) value_weight = self.value_weight.permute(1, 0, 2).reshape(hidden_size, self.intermediate_size) down_weight = self.down_weight.reshape(self.intermediate_size, hidden_size) hidden = F.silu(x_flat @ up_weight) * (x_flat @ value_weight) hidden = self._apply_memory(hidden) return (hidden @ down_weight).reshape(original_shape) class TileRoutedKeyedChannelMemoryDSwiGLUMLP(TileRoutedDSwiGLUMLP): """Prototype-keyed token-conditioned memory for grouped4. This reads a richer token signal than group mean/RMS: each group compares its normalized centered hidden state to learned prototype keys, then uses zero-init bounded value vectors to modulate centered channels. With zero values the function is exactly grouped4. """ def __init__(self, config: Config, intermediate_size: int | None = None) -> None: super().__init__(config, intermediate_size) self.block_idx = 0 self.n_layer = config.n_layer self.channel_memory_gain_max = float(config.channel_memory_gain_max) self.channel_memory_late_layer_start = int(config.channel_memory_late_layer_start) self.keyed_channel_memory_rank = int(config.keyed_channel_memory_rank) key = torch.randn(self.num_groups, self.keyed_channel_memory_rank, self.group_size) key = key * (self.group_size ** -0.5) self.keyed_memory_key = nn.Parameter(key) self.keyed_memory_value_raw = nn.Parameter( torch.zeros(self.num_groups, self.keyed_channel_memory_rank, self.group_size) ) self.last_keyed_memory_gain_abs: torch.Tensor | None = None self.last_keyed_memory_state_abs: torch.Tensor | None = None def set_block_index(self, block_idx: int, n_layer: int) -> None: self.block_idx = block_idx self.n_layer = n_layer enabled = self._memory_enabled() self.keyed_memory_key.requires_grad_(enabled) self.keyed_memory_value_raw.requires_grad_(enabled) def _memory_enabled(self) -> bool: return self.channel_memory_gain_max > 0.0 and self.block_idx >= self.channel_memory_late_layer_start def _apply_keyed_memory(self, hidden: torch.Tensor) -> torch.Tensor: if not self._memory_enabled(): return hidden hidden_groups = hidden.reshape(hidden.shape[0], self.num_groups, self.group_size) hidden_float = hidden_groups.float() group_mean = hidden_float.mean(dim=-1, keepdim=True) centered = hidden_groups - group_mean.to(dtype=hidden_groups.dtype) centered_float = centered.float() centered_rms = centered_float.square().mean(dim=-1, keepdim=True).sqrt().clamp_min(1e-6) centered_norm = centered_float / centered_rms key = self.keyed_memory_key.to(device=hidden.device, dtype=torch.float32) key = key / key.square().mean(dim=-1, keepdim=True).sqrt().clamp_min(1e-6) state = torch.tanh((centered_norm.unsqueeze(2) * key.unsqueeze(0)).mean(dim=-1)) value = self.channel_memory_gain_max * torch.tanh( self.keyed_memory_value_raw.to(device=hidden.device, dtype=hidden_groups.dtype) ) channel_gain = (state.to(dtype=hidden_groups.dtype).unsqueeze(-1) * value.unsqueeze(0)).sum(dim=2) channel_gain = channel_gain * (self.keyed_channel_memory_rank ** -0.5) if not torch.is_grad_enabled(): self.last_keyed_memory_gain_abs = channel_gain.detach().abs().mean() self.last_keyed_memory_state_abs = state.detach().abs().mean() hidden_groups = hidden_groups + centered * channel_gain return hidden_groups.reshape(hidden.shape) def _forward_full_active(self, x: torch.Tensor, original_shape: torch.Size, x_flat: torch.Tensor) -> torch.Tensor: hidden_size = original_shape[-1] up_weight = self.up_weight.permute(1, 0, 2).reshape(hidden_size, self.intermediate_size) value_weight = self.value_weight.permute(1, 0, 2).reshape(hidden_size, self.intermediate_size) down_weight = self.down_weight.reshape(self.intermediate_size, hidden_size) hidden = F.silu(x_flat @ up_weight) * (x_flat @ value_weight) hidden = self._apply_keyed_memory(hidden) return (hidden @ down_weight).reshape(original_shape) class TileRoutedGroupStateMemoryDSwiGLUMLP(TileRoutedDSwiGLUMLP): """Token-conditioned bounded group memory for grouped4. This keeps grouped4's full-active GEMM path, then lets each hidden group scale its own channels from a signed normalized group state. Zero gain is exactly grouped4, while nonzero gain gives a cheap token-dependent memory write before the down projection. """ def __init__(self, config: Config, intermediate_size: int | None = None) -> None: super().__init__(config, intermediate_size) self.block_idx = 0 self.n_layer = config.n_layer self.group_state_memory_gain_max = float(config.group_state_memory_gain_max) self.group_state_memory_late_layer_start = int(config.group_state_memory_late_layer_start) self.group_state_memory_gain_raw = nn.Parameter(torch.zeros(self.num_groups, 1)) self.last_group_state_gain_abs: torch.Tensor | None = None self.last_group_state_abs: torch.Tensor | None = None def set_block_index(self, block_idx: int, n_layer: int) -> None: self.block_idx = block_idx self.n_layer = n_layer self.group_state_memory_gain_raw.requires_grad_(self._memory_enabled()) def _memory_enabled(self) -> bool: return self.group_state_memory_gain_max > 0.0 and self.block_idx >= self.group_state_memory_late_layer_start def _apply_group_state_memory(self, hidden: torch.Tensor) -> torch.Tensor: if not self._memory_enabled(): return hidden hidden_groups = hidden.reshape(hidden.shape[0], self.num_groups, self.group_size) hidden_float = hidden_groups.float() group_mean = hidden_float.mean(dim=-1, keepdim=True) group_rms = hidden_float.square().mean(dim=-1, keepdim=True).sqrt().clamp_min(1e-6) group_state = torch.tanh(group_mean / group_rms).to(dtype=hidden_groups.dtype) gain = self.group_state_memory_gain_max * torch.tanh( self.group_state_memory_gain_raw.to(device=hidden.device, dtype=hidden_groups.dtype) ).view(1, self.num_groups, 1) if not torch.is_grad_enabled(): self.last_group_state_gain_abs = gain.detach().abs().mean() self.last_group_state_abs = group_state.detach().abs().mean() hidden_groups = hidden_groups * (1.0 + gain * group_state) return hidden_groups.reshape(hidden.shape) def _forward_full_active(self, x: torch.Tensor, original_shape: torch.Size, x_flat: torch.Tensor) -> torch.Tensor: hidden_size = original_shape[-1] up_weight = self.up_weight.permute(1, 0, 2).reshape(hidden_size, self.intermediate_size) value_weight = self.value_weight.permute(1, 0, 2).reshape(hidden_size, self.intermediate_size) down_weight = self.down_weight.reshape(self.intermediate_size, hidden_size) hidden = F.silu(x_flat @ up_weight) * (x_flat @ value_weight) hidden = self._apply_group_state_memory(hidden) return (hidden @ down_weight).reshape(original_shape) class TileRoutedChannelGroupStateMemoryDSwiGLUMLP(TileRoutedDSwiGLUMLP): """Combined foldable channel values plus token-conditioned group memory.""" def __init__(self, config: Config, intermediate_size: int | None = None) -> None: super().__init__(config, intermediate_size) self.block_idx = 0 self.n_layer = config.n_layer self.channel_memory_gain_max = float(config.channel_memory_gain_max) self.channel_memory_late_layer_start = int(config.channel_memory_late_layer_start) self.group_state_memory_gain_max = float(config.group_state_memory_gain_max) self.group_state_memory_late_layer_start = int(config.group_state_memory_late_layer_start) self.channel_memory_gain_raw = nn.Parameter(torch.zeros(self.num_groups, self.group_size, 1)) self.group_state_memory_gain_raw = nn.Parameter(torch.zeros(self.num_groups, 1)) self.last_channel_memory_gain_abs: torch.Tensor | None = None self.last_group_state_gain_abs: torch.Tensor | None = None self.last_group_state_abs: torch.Tensor | None = None def set_block_index(self, block_idx: int, n_layer: int) -> None: self.block_idx = block_idx self.n_layer = n_layer self.channel_memory_gain_raw.requires_grad_(self._channel_memory_enabled()) self.group_state_memory_gain_raw.requires_grad_(self._group_state_memory_enabled()) def _channel_memory_enabled(self) -> bool: return self.channel_memory_gain_max > 0.0 and self.block_idx >= self.channel_memory_late_layer_start def _group_state_memory_enabled(self) -> bool: return self.group_state_memory_gain_max > 0.0 and self.block_idx >= self.group_state_memory_late_layer_start def _effective_down_weight(self) -> torch.Tensor: if not self._channel_memory_enabled(): return self.down_weight gain = self.channel_memory_gain_max * torch.tanh( self.channel_memory_gain_raw.to(device=self.down_weight.device, dtype=self.down_weight.dtype) ) if not torch.is_grad_enabled(): self.last_channel_memory_gain_abs = gain.detach().abs().mean() return self.down_weight * (1.0 + gain) def _apply_group_state_memory(self, hidden: torch.Tensor) -> torch.Tensor: if not self._group_state_memory_enabled(): return hidden hidden_groups = hidden.reshape(hidden.shape[0], self.num_groups, self.group_size) hidden_float = hidden_groups.float() group_mean = hidden_float.mean(dim=-1, keepdim=True) group_rms = hidden_float.square().mean(dim=-1, keepdim=True).sqrt().clamp_min(1e-6) group_state = torch.tanh(group_mean / group_rms).to(dtype=hidden_groups.dtype) gain = self.group_state_memory_gain_max * torch.tanh( self.group_state_memory_gain_raw.to(device=hidden.device, dtype=hidden_groups.dtype) ).view(1, self.num_groups, 1) if not torch.is_grad_enabled(): self.last_group_state_gain_abs = gain.detach().abs().mean() self.last_group_state_abs = group_state.detach().abs().mean() hidden_groups = hidden_groups * (1.0 + gain * group_state) return hidden_groups.reshape(hidden.shape) def _forward_full_active(self, x: torch.Tensor, original_shape: torch.Size, x_flat: torch.Tensor) -> torch.Tensor: hidden_size = original_shape[-1] up_weight = self.up_weight.permute(1, 0, 2).reshape(hidden_size, self.intermediate_size) value_weight = self.value_weight.permute(1, 0, 2).reshape(hidden_size, self.intermediate_size) down_weight = self._effective_down_weight().reshape(self.intermediate_size, hidden_size) hidden = F.silu(x_flat @ up_weight) * (x_flat @ value_weight) hidden = self._apply_group_state_memory(hidden) return (hidden @ down_weight).reshape(original_shape) class TileRoutedGRNDSwiGLUMLP(TileRoutedDSwiGLUMLP): """Grouped hidden Global Response Normalization before the down projection.""" def __init__(self, config: Config, intermediate_size: int | None = None) -> None: super().__init__(config, intermediate_size) self.grn_gamma = nn.Parameter(torch.zeros(self.num_groups, self.group_size)) self.grn_beta = nn.Parameter(torch.zeros(self.num_groups, self.group_size)) def _forward_full_active(self, x: torch.Tensor, original_shape: torch.Size, x_flat: torch.Tensor) -> torch.Tensor: hidden_size = original_shape[-1] up_weight = self.up_weight.permute(1, 0, 2).reshape(hidden_size, self.intermediate_size) value_weight = self.value_weight.permute(1, 0, 2).reshape(hidden_size, self.intermediate_size) hidden = F.silu(x_flat @ up_weight) * (x_flat @ value_weight) hidden_groups = hidden.reshape(hidden.shape[0], self.num_groups, self.group_size) group_norm = hidden_groups.float().pow(2).mean(dim=-1, keepdim=True).sqrt() response = group_norm / group_norm.mean(dim=1, keepdim=True).clamp_min(1e-6) gamma = self.grn_gamma.to(device=hidden.device, dtype=hidden.dtype).unsqueeze(0) beta = self.grn_beta.to(device=hidden.device, dtype=hidden.dtype).unsqueeze(0) hidden_groups = hidden_groups + gamma * (hidden_groups * response.to(dtype=hidden.dtype)) + beta down_weight = self.down_weight.reshape(self.intermediate_size, hidden_size) return (hidden_groups.reshape_as(hidden) @ down_weight).reshape(original_shape) class TileRoutedGroupMixDSwiGLUMLP(TileRoutedDSwiGLUMLP): """Learn a tiny soft group-to-group hidden mixer before the down weights.""" def __init__(self, config: Config, intermediate_size: int | None = None) -> None: super().__init__(config, intermediate_size) logits = torch.full((self.num_groups, self.num_groups), -6.0) logits.fill_diagonal_(6.0) self.group_mix_logits = nn.Parameter(logits) def _forward_full_active(self, x: torch.Tensor, original_shape: torch.Size, x_flat: torch.Tensor) -> torch.Tensor: hidden_size = original_shape[-1] up_weight = self.up_weight.permute(1, 0, 2).reshape(hidden_size, self.intermediate_size) value_weight = self.value_weight.permute(1, 0, 2).reshape(hidden_size, self.intermediate_size) hidden = F.silu(x_flat @ up_weight) * (x_flat @ value_weight) hidden_groups = hidden.reshape(hidden.shape[0], self.num_groups, self.group_size) mix = torch.softmax(self.group_mix_logits.to(device=hidden.device, dtype=torch.float32), dim=-1).to(dtype=hidden.dtype) hidden_groups = torch.einsum("og,ngs->nos", mix, hidden_groups) down_weight = self.down_weight.reshape(self.intermediate_size, hidden_size) return (hidden_groups.reshape_as(hidden) @ down_weight).reshape(original_shape) class TileRoutedTokenGroupGateDSwiGLUMLP(TileRoutedDSwiGLUMLP): """Token-dependent competition between full-active groups.""" def __init__(self, config: Config, intermediate_size: int | None = None) -> None: super().__init__(config, intermediate_size) self.group_gate_alpha = nn.Parameter(torch.tensor(0.05)) def _forward_full_active(self, x: torch.Tensor, original_shape: torch.Size, x_flat: torch.Tensor) -> torch.Tensor: hidden_size = original_shape[-1] up_weight = self.up_weight.permute(1, 0, 2).reshape(hidden_size, self.intermediate_size) value_weight = self.value_weight.permute(1, 0, 2).reshape(hidden_size, self.intermediate_size) hidden = F.silu(x_flat @ up_weight) * (x_flat @ value_weight) hidden_groups = hidden.reshape(hidden.shape[0], self.num_groups, self.group_size) scores = hidden_groups.float().pow(2).mean(dim=-1).sqrt() gates = torch.softmax(scores, dim=-1).to(dtype=hidden.dtype).unsqueeze(-1) * self.num_groups alpha = self.group_gate_alpha.to(device=hidden.device, dtype=hidden.dtype).clamp(-0.5, 0.5) hidden_groups = hidden_groups * (1.0 + alpha * (gates - 1.0)) down_weight = self.down_weight.reshape(self.intermediate_size, hidden_size) return (hidden_groups.reshape_as(hidden) @ down_weight).reshape(original_shape) class TileRoutedBlockNormDSwiGLUMLP(TileRoutedDSwiGLUMLP): """Full-active grouped MLP with block-local RMS normalization per group. Hard hidden blocks damaged LM quality because each FFN block lost access to most of the residual stream. This softer variant keeps the standard full-hidden grouped MLP: every intermediate group still sees all hidden channels. The non-absorbed part is an input-dependent local RMS normalization over hidden-channel blocks before each group path. """ def __init__(self, config: Config, intermediate_size: int | None = None) -> None: super().__init__(config, intermediate_size) if config.n_embd % self.num_groups != 0: raise ValueError("n_embd must be divisible by sparse_mlp_num_groups") self.block_hidden_size = config.n_embd // self.num_groups self.group_norm_weight = nn.Parameter(torch.ones(self.num_groups, config.n_embd)) def _block_norm(self, x_flat: torch.Tensor, group_idx: int) -> torch.Tensor: x_blocks = x_flat.reshape(x_flat.shape[0], self.num_groups, self.block_hidden_size) scale = torch.rsqrt(x_blocks.pow(2).mean(dim=-1, keepdim=True) + self.config.norm_eps) x_normed = (x_blocks * scale).reshape_as(x_flat) weight = self.group_norm_weight[group_idx].to(device=x_flat.device, dtype=x_flat.dtype) return x_normed * weight def _forward_full_active(self, x: torch.Tensor, original_shape: torch.Size, x_flat: torch.Tensor) -> torch.Tensor: output = torch.zeros_like(x_flat) for group_idx in range(self.num_groups): group_input = self._block_norm(x_flat, group_idx) hidden = F.silu(group_input @ self.up_weight[group_idx]) * (group_input @ self.value_weight[group_idx]) output = output + hidden @ self.down_weight[group_idx] return output.reshape(original_shape) class TileRoutedStackedDSwiGLUMLP(TileRoutedDSwiGLUMLP): """Full-active grouped dSwiGLU with serial group communication. At ``grouped_mlp_stack_alpha == 0`` this is equivalent to full-active grouped 4/4. With positive alpha, each group sees the residual state updated by earlier group outputs, making the groups a small ordered sub-FFN instead of independent parallel subspaces. """ def __init__(self, config: Config, intermediate_size: int | None = None) -> None: super().__init__(config, intermediate_size) stack_alpha = torch.tensor(float(config.grouped_mlp_stack_alpha)) if config.grouped_mlp_stack_alpha_learnable: self.stack_alpha = nn.Parameter(stack_alpha) else: self.register_buffer("stack_alpha", stack_alpha, persistent=False) def _forward_full_active(self, x: torch.Tensor, original_shape: torch.Size, x_flat: torch.Tensor) -> torch.Tensor: output = torch.zeros_like(x_flat) state = x_flat stack_alpha = self.stack_alpha.to(device=x_flat.device, dtype=x_flat.dtype) for group_idx in range(self.num_groups): hidden = F.silu(state @ self.up_weight[group_idx]) * (state @ self.value_weight[group_idx]) delta = hidden @ self.down_weight[group_idx] output = output + delta if group_idx + 1 < self.num_groups: state = state + stack_alpha * delta return output.reshape(original_shape) def forward(self, x: torch.Tensor) -> torch.Tensor: original_shape = x.shape x_flat = x.reshape(-1, original_shape[-1]) if self.active_groups != self.num_groups: return super().forward(x) output = self._forward_full_active(x, original_shape, x_flat) if not torch.is_grad_enabled(): self.last_gate_sparsity = x_flat.new_tensor(0.0) self.last_threshold_mean = x_flat.new_tensor(0.0) self.last_active_groups_mean = x_flat.new_tensor(float(self.active_groups)) return self._apply_shared_path(x, output) class TileRoutedChunkedStackedDSwiGLUMLP(TileRoutedDSwiGLUMLP): """Full-active grouped MLP with serial stage chunks inside existing groups. This keeps the exact ``TileRoutedDSwiGLUMLP`` parameter layout. The intermediate channels in each group are split into ``grouped_mlp_stack_depth`` chunks. At alpha=0, summing all chunks is exactly the normal grouped MLP. Non-zero alpha lets later chunks see a small residual correction from earlier chunks, testing the ARC margin effect without changing initialization layout. """ def __init__(self, config: Config, intermediate_size: int | None = None) -> None: super().__init__(config, intermediate_size) self.stack_depth = config.grouped_mlp_stack_depth if self.stack_depth < 1: raise ValueError("grouped_mlp_stack_depth must be >= 1") if self.group_size % self.stack_depth != 0: raise ValueError("group_size must be divisible by grouped_mlp_stack_depth") self.stage_group_size = self.group_size // self.stack_depth self.stage_intermediate_size = self.stage_group_size * self.num_groups stack_alpha = torch.tensor(float(config.grouped_mlp_stack_alpha)) if config.grouped_mlp_stack_alpha_learnable: self.stack_alpha = nn.Parameter(stack_alpha) else: self.register_buffer("stack_alpha", stack_alpha, persistent=False) output_scale = torch.tensor(float(config.grouped_mlp_stack_output_scale)) if config.grouped_mlp_stack_output_scale_learnable: self.stack_output_scale = nn.Parameter(output_scale) else: self.register_buffer("stack_output_scale", output_scale, persistent=False) def _stage_forward(self, stage_idx: int, state: torch.Tensor) -> torch.Tensor: hidden_size = state.shape[-1] start = stage_idx * self.stage_group_size stop = (stage_idx + 1) * self.stage_group_size up_weight = self.up_weight[:, :, start:stop].permute(1, 0, 2).reshape(hidden_size, self.stage_intermediate_size) value_weight = self.value_weight[:, :, start:stop].permute(1, 0, 2).reshape(hidden_size, self.stage_intermediate_size) down_weight = self.down_weight[:, start:stop, :].reshape(self.stage_intermediate_size, hidden_size) return (F.silu(state @ up_weight) * (state @ value_weight)) @ down_weight def _forward_full_active(self, x: torch.Tensor, original_shape: torch.Size, x_flat: torch.Tensor) -> torch.Tensor: state = x_flat output = torch.zeros_like(x_flat) stack_alpha = self.stack_alpha.to(device=x_flat.device, dtype=x_flat.dtype) for stage_idx in range(self.stack_depth): delta = self._stage_forward(stage_idx, state) output = output + delta if stage_idx + 1 < self.stack_depth: state = state + stack_alpha * delta output_scale = self.stack_output_scale.to(device=x_flat.device, dtype=x_flat.dtype) return (output * output_scale).reshape(original_shape) class TileRoutedGatedStackedDSwiGLUMLP(nn.Module): """Grouped MLP main path plus a gated independent split-stage branch. The total intermediate budget is conserved: - ``intermediate_size - grouped_mlp_contrastive_intermediate_size`` goes to a normal full-active grouped MLP main path. - ``grouped_mlp_contrastive_intermediate_size`` goes to a split-stage branch, the architecture that produced stronger ARC margins. This tests whether the split-stage branch is useful as a bounded contrastive/ranking perturbation instead of replacing the whole MLP. """ def __init__(self, config: Config, intermediate_size: int | None = None) -> None: super().__init__() self.intermediate_size = intermediate_size or config.intermediate_size self.num_groups = config.sparse_mlp_num_groups self.active_groups = config.sparse_mlp_max_active_groups self.stack_depth = config.grouped_mlp_stack_depth self.branch_intermediate_size = int(config.grouped_mlp_contrastive_intermediate_size) if self.active_groups != self.num_groups: raise ValueError("TileRoutedGatedStackedDSwiGLUMLP is a full-active grouped variant.") if self.stack_depth < 1: raise ValueError("grouped_mlp_stack_depth must be >= 1") if config.bias: raise ValueError("TileRoutedGatedStackedDSwiGLUMLP currently expects bias=False.") if self.branch_intermediate_size <= 0: raise ValueError("grouped_mlp_contrastive_intermediate_size must be > 0") if self.branch_intermediate_size >= self.intermediate_size: raise ValueError("contrastive branch must be smaller than intermediate_size") self.main_intermediate_size = self.intermediate_size - self.branch_intermediate_size if self.main_intermediate_size % self.num_groups != 0: raise ValueError("main intermediate must divide across sparse_mlp_num_groups") if self.branch_intermediate_size % (self.num_groups * self.stack_depth) != 0: raise ValueError("contrastive branch intermediate must divide across groups * stack depth") self.main_group_size = self.main_intermediate_size // self.num_groups self.branch_stage_intermediate_size = self.branch_intermediate_size // self.stack_depth self.branch_group_size = self.branch_stage_intermediate_size // self.num_groups self.main_up_weight = nn.Parameter(torch.empty(self.num_groups, config.n_embd, self.main_group_size)) self.main_value_weight = nn.Parameter(torch.empty(self.num_groups, config.n_embd, self.main_group_size)) self.main_down_weight = nn.Parameter(torch.empty(self.num_groups, self.main_group_size, config.n_embd)) self.branch_up_weight = nn.Parameter( torch.empty(self.stack_depth, self.num_groups, config.n_embd, self.branch_group_size) ) self.branch_value_weight = nn.Parameter( torch.empty(self.stack_depth, self.num_groups, config.n_embd, self.branch_group_size) ) self.branch_down_weight = nn.Parameter( torch.empty(self.stack_depth, self.num_groups, self.branch_group_size, config.n_embd) ) stack_alpha = torch.tensor(float(config.grouped_mlp_stack_alpha)) if config.grouped_mlp_stack_alpha_learnable: self.stack_alpha = nn.Parameter(stack_alpha) else: self.register_buffer("stack_alpha", stack_alpha, persistent=False) gate_max = float(config.grouped_mlp_contrastive_gate_max) if gate_max <= 0: raise ValueError("grouped_mlp_contrastive_gate_max must be > 0") self.contrastive_gate_max = gate_max gate = float(config.grouped_mlp_contrastive_gate) if config.grouped_mlp_contrastive_gate_learnable: ratio = max(-0.999, min(0.999, gate / gate_max)) self.raw_contrastive_gate = nn.Parameter(torch.tensor(math.atanh(ratio))) else: self.register_buffer("contrastive_gate", torch.tensor(gate), persistent=False) self.config = config self.last_gate_sparsity: torch.Tensor | None = None self.last_threshold_mean: torch.Tensor | None = None self.last_active_groups_mean: torch.Tensor | None = None self.reset_parameters() def reset_parameters(self) -> None: std = math.sqrt(2.0 / 5 / self.config.n_embd) nn.init.normal_(self.main_up_weight, mean=0.0, std=std) nn.init.normal_(self.main_value_weight, mean=0.0, std=std) nn.init.normal_(self.main_down_weight, mean=0.0, std=std) branch_std = std * float(getattr(self.config, "grouped_mlp_stack_init_scale", 1.0)) nn.init.normal_(self.branch_up_weight, mean=0.0, std=branch_std) nn.init.normal_(self.branch_value_weight, mean=0.0, std=branch_std) nn.init.normal_(self.branch_down_weight, mean=0.0, std=branch_std) def _main_forward(self, x_flat: torch.Tensor) -> torch.Tensor: hidden_size = x_flat.shape[-1] up_weight = self.main_up_weight.permute(1, 0, 2).reshape(hidden_size, self.main_intermediate_size) value_weight = self.main_value_weight.permute(1, 0, 2).reshape(hidden_size, self.main_intermediate_size) down_weight = self.main_down_weight.reshape(self.main_intermediate_size, hidden_size) return (F.silu(x_flat @ up_weight) * (x_flat @ value_weight)) @ down_weight def _branch_stage_forward(self, stage_idx: int, state: torch.Tensor) -> torch.Tensor: hidden_size = state.shape[-1] up_weight = self.branch_up_weight[stage_idx].permute(1, 0, 2).reshape( hidden_size, self.branch_stage_intermediate_size ) value_weight = self.branch_value_weight[stage_idx].permute(1, 0, 2).reshape( hidden_size, self.branch_stage_intermediate_size ) down_weight = self.branch_down_weight[stage_idx].reshape(self.branch_stage_intermediate_size, hidden_size) return (F.silu(state @ up_weight) * (state @ value_weight)) @ down_weight def _branch_forward(self, x_flat: torch.Tensor) -> torch.Tensor: state = x_flat output = torch.zeros_like(x_flat) stack_alpha = self.stack_alpha.to(device=x_flat.device, dtype=x_flat.dtype) for stage_idx in range(self.stack_depth): delta = self._branch_stage_forward(stage_idx, state) output = output + delta if stage_idx + 1 < self.stack_depth: state = state + stack_alpha * delta return output def _gate(self, x_flat: torch.Tensor) -> torch.Tensor: if hasattr(self, "raw_contrastive_gate"): gate = self.contrastive_gate_max * torch.tanh(self.raw_contrastive_gate) else: gate = self.contrastive_gate return gate.to(device=x_flat.device, dtype=x_flat.dtype) def forward(self, x: torch.Tensor) -> torch.Tensor: original_shape = x.shape x_flat = x.reshape(-1, original_shape[-1]) output = self._main_forward(x_flat) + self._gate(x_flat) * self._branch_forward(x_flat) if not torch.is_grad_enabled(): self.last_gate_sparsity = x_flat.new_tensor(0.0) self.last_threshold_mean = x_flat.new_tensor(0.0) self.last_active_groups_mean = x_flat.new_tensor(float(self.active_groups)) return output.reshape(original_shape) class TileRoutedFullMainGatedStackedDSwiGLUMLP(TileRoutedGatedStackedDSwiGLUMLP): """Full grouped MLP plus a ReZero-style split-stack residual branch. ``TileRoutedGatedStackedDSwiGLUMLP`` conserves total intermediate width by moving channels from the normal grouped MLP into the split-stack branch. That is useful as a fixed-budget control, but it weakens the LM path whenever the branch gate starts small. This variant keeps the full grouped MLP width and adds the split-stack branch on top behind the same bounded learnable gate. """ def __init__(self, config: Config, intermediate_size: int | None = None) -> None: super().__init__(config, intermediate_size) self.main_intermediate_size = self.intermediate_size self.main_group_size = self.intermediate_size // self.num_groups self.main_up_weight = nn.Parameter(torch.empty(self.num_groups, config.n_embd, self.main_group_size)) self.main_value_weight = nn.Parameter(torch.empty(self.num_groups, config.n_embd, self.main_group_size)) self.main_down_weight = nn.Parameter(torch.empty(self.num_groups, self.main_group_size, config.n_embd)) self.reset_parameters() def reset_parameters(self) -> None: std = math.sqrt(2.0 / 5 / self.config.n_embd) nn.init.normal_(self.main_up_weight, mean=0.0, std=std) nn.init.normal_(self.main_value_weight, mean=0.0, std=std) nn.init.normal_(self.main_down_weight, mean=0.0, std=std) branch_std = std * float(getattr(self.config, "grouped_mlp_stack_init_scale", 1.0)) nn.init.normal_(self.branch_up_weight, mean=0.0, std=branch_std) nn.init.normal_(self.branch_value_weight, mean=0.0, std=branch_std) nn.init.normal_(self.branch_down_weight, mean=0.0, std=branch_std) def forward(self, x: torch.Tensor) -> torch.Tensor: original_shape = x.shape x_flat = x.reshape(-1, original_shape[-1]) main_output = self._main_forward(x_flat) branch_output = self._branch_forward(x_flat) output = main_output + self._gate(x_flat) * branch_output if not torch.is_grad_enabled(): self.last_gate_sparsity = x_flat.new_tensor(0.0) self.last_threshold_mean = x_flat.new_tensor(0.0) self.last_active_groups_mean = x_flat.new_tensor(float(self.active_groups)) return output.reshape(original_shape) class TileRoutedFullMainRMSBoundedGatedStackedDSwiGLUMLP(TileRoutedFullMainGatedStackedDSwiGLUMLP): """Full-main gated split-stack with final RMS residual scaling. The full-main gated screen showed that a small branch gate alone does not prevent the whole MLP update from growing to split-stack scale. This variant interprets ``grouped_mlp_stack_output_scale`` as the target output/input RMS ratio for the complete MLP output. """ def __init__(self, config: Config, intermediate_size: int | None = None) -> None: super().__init__(config, intermediate_size) output_scale = torch.tensor(float(config.grouped_mlp_stack_output_scale)) if config.grouped_mlp_stack_output_scale_learnable: self.stack_output_scale = nn.Parameter(output_scale) else: self.register_buffer("stack_output_scale", output_scale, persistent=False) def forward(self, x: torch.Tensor) -> torch.Tensor: original_shape = x.shape x_flat = x.reshape(-1, original_shape[-1]) main_output = self._main_forward(x_flat) branch_output = self._branch_forward(x_flat) output = main_output + self._gate(x_flat) * branch_output target_ratio = self.stack_output_scale.to(device=x_flat.device, dtype=x_flat.dtype) input_rms = x_flat.float().pow(2).mean().sqrt().to(dtype=x_flat.dtype) output_rms = output.float().pow(2).mean().sqrt().to(dtype=x_flat.dtype) scale = target_ratio * input_rms / output_rms.clamp_min(torch.finfo(output.dtype).eps) output = output * scale if not torch.is_grad_enabled(): self.last_gate_sparsity = x_flat.new_tensor(0.0) self.last_threshold_mean = x_flat.new_tensor(0.0) self.last_active_groups_mean = x_flat.new_tensor(float(self.active_groups)) return output.reshape(original_shape) class TileRoutedFullMainDirectionalRMSGatedStackedDSwiGLUMLP( TileRoutedFullMainRMSBoundedGatedStackedDSwiGLUMLP ): def _rms_bound_output(self, x_flat: torch.Tensor, output: torch.Tensor) -> torch.Tensor: target_ratio = self.stack_output_scale.to(device=x_flat.device, dtype=x_flat.dtype) input_rms = x_flat.float().pow(2).mean().sqrt().to(dtype=x_flat.dtype) output_rms = output.float().pow(2).mean().sqrt().to(dtype=x_flat.dtype) scale = target_ratio * input_rms / output_rms.clamp_min(torch.finfo(output.dtype).eps) return output * scale def _match_token_rms(self, source: torch.Tensor, reference: torch.Tensor) -> torch.Tensor: source_rms = source.float().pow(2).mean(dim=-1, keepdim=True).sqrt() reference_rms = reference.float().pow(2).mean(dim=-1, keepdim=True).sqrt() scale = reference_rms / source_rms.clamp_min(1e-8) return source * scale.to(device=source.device, dtype=source.dtype) class TileRoutedFullMainOrthogonalRMSGatedStackedDSwiGLUMLP( TileRoutedFullMainDirectionalRMSGatedStackedDSwiGLUMLP ): """Add only the branch component that is orthogonal to the main grouped update.""" def forward(self, x: torch.Tensor) -> torch.Tensor: original_shape = x.shape x_flat = x.reshape(-1, original_shape[-1]) main_output = self._main_forward(x_flat) branch_output = self._branch_forward(x_flat) dot = (branch_output.float() * main_output.float()).sum(dim=-1, keepdim=True) denom = main_output.float().pow(2).sum(dim=-1, keepdim=True).clamp_min(1e-8) branch_output = branch_output - (dot / denom).to(dtype=branch_output.dtype) * main_output branch_output = self._match_token_rms(branch_output, main_output) output = main_output + self._gate(x_flat) * branch_output output = self._rms_bound_output(x_flat, output) if not torch.is_grad_enabled(): self.last_gate_sparsity = x_flat.new_tensor(0.0) self.last_threshold_mean = x_flat.new_tensor(0.0) self.last_active_groups_mean = x_flat.new_tensor(float(self.active_groups)) return output.reshape(original_shape) class TileRoutedFullMainAlignedRMSGatedStackedDSwiGLUMLP( TileRoutedFullMainDirectionalRMSGatedStackedDSwiGLUMLP ): """Convert the branch into a bounded token-wise modulation of the main update.""" def forward(self, x: torch.Tensor) -> torch.Tensor: original_shape = x.shape x_flat = x.reshape(-1, original_shape[-1]) main_output = self._main_forward(x_flat) branch_output = self._branch_forward(x_flat) dot = (branch_output.float() * main_output.float()).sum(dim=-1, keepdim=True) denom = main_output.float().pow(2).sum(dim=-1, keepdim=True).clamp_min(1e-8) coeff = (dot / denom).clamp(min=-1.0, max=1.0).to(dtype=main_output.dtype) output = main_output * (1.0 + self._gate(x_flat) * coeff) output = self._rms_bound_output(x_flat, output) if not torch.is_grad_enabled(): self.last_gate_sparsity = x_flat.new_tensor(0.0) self.last_threshold_mean = x_flat.new_tensor(0.0) self.last_active_groups_mean = x_flat.new_tensor(float(self.active_groups)) return output.reshape(original_shape) class TileRoutedFullMainNormedBranchRMSGatedStackedDSwiGLUMLP( TileRoutedFullMainDirectionalRMSGatedStackedDSwiGLUMLP ): """Normalize branch token RMS to the main update before the gated add.""" def forward(self, x: torch.Tensor) -> torch.Tensor: original_shape = x.shape x_flat = x.reshape(-1, original_shape[-1]) main_output = self._main_forward(x_flat) branch_output = self._match_token_rms(self._branch_forward(x_flat), main_output) output = main_output + self._gate(x_flat) * branch_output output = self._rms_bound_output(x_flat, output) if not torch.is_grad_enabled(): self.last_gate_sparsity = x_flat.new_tensor(0.0) self.last_threshold_mean = x_flat.new_tensor(0.0) self.last_active_groups_mean = x_flat.new_tensor(float(self.active_groups)) return output.reshape(original_shape) class TileRoutedFullMainCompressedRMSGatedStackedDSwiGLUMLP( TileRoutedFullMainDirectionalRMSGatedStackedDSwiGLUMLP ): """Compress the split branch through a small dense bottleneck before adding it.""" def __init__(self, config: Config, intermediate_size: int | None = None) -> None: super().__init__(config, intermediate_size) rank = int(config.grouped_mlp_mix_rank) if config.grouped_mlp_mix_rank > 0 else max(1, config.n_embd // 4) self.branch_compress = nn.Linear(config.n_embd, rank, bias=False) self.branch_expand = nn.Linear(rank, config.n_embd, bias=False) nn.init.normal_(self.branch_compress.weight, mean=0.0, std=0.02) nn.init.normal_(self.branch_expand.weight, mean=0.0, std=0.02) def forward(self, x: torch.Tensor) -> torch.Tensor: original_shape = x.shape x_flat = x.reshape(-1, original_shape[-1]) main_output = self._main_forward(x_flat) branch_output = self.branch_expand(torch.tanh(self.branch_compress(self._branch_forward(x_flat)))) branch_output = self._match_token_rms(branch_output, main_output) output = main_output + self._gate(x_flat) * branch_output output = self._rms_bound_output(x_flat, output) if not torch.is_grad_enabled(): self.last_gate_sparsity = x_flat.new_tensor(0.0) self.last_threshold_mean = x_flat.new_tensor(0.0) self.last_active_groups_mean = x_flat.new_tensor(float(self.active_groups)) return output.reshape(original_shape) class TileRoutedLateOnlyRMSGatedStackedDSwiGLUMLP(TileRoutedFullMainDirectionalRMSGatedStackedDSwiGLUMLP): """Use the branch only in later transformer blocks.""" def set_block_index(self, block_idx: int, n_layer: int) -> None: self.block_idx = block_idx self.n_layer = n_layer def forward(self, x: torch.Tensor) -> torch.Tensor: original_shape = x.shape x_flat = x.reshape(-1, original_shape[-1]) main_output = self._main_forward(x_flat) start_layer = int(getattr(self.config, "grouped_mlp_branch_start_layer", 0)) if getattr(self, "block_idx", 0) >= start_layer: branch_output = self._match_token_rms(self._branch_forward(x_flat), main_output) output = main_output + self._gate(x_flat) * branch_output else: output = main_output output = self._rms_bound_output(x_flat, output) if not torch.is_grad_enabled(): self.last_gate_sparsity = x_flat.new_tensor(0.0) self.last_threshold_mean = x_flat.new_tensor(0.0) self.last_active_groups_mean = x_flat.new_tensor(float(self.active_groups)) return output.reshape(original_shape) class TileRoutedAdditiveLateTokenRMSGatedDSwiGLUMLP(TileRoutedDSwiGLUMLP): """Full grouped 4/4 MLP plus a tiny late token-conditioned branch. This keeps the base grouped MLP parameter names and shapes unchanged so a trained ``TileRoutedDSwiGLUMLP`` checkpoint can initialize the full LM path exactly. Only the additive branch and its token gate are new. """ def __init__(self, config: Config, intermediate_size: int | None = None) -> None: super().__init__(config, intermediate_size) branch_intermediate_size = int(config.grouped_mlp_contrastive_intermediate_size) if branch_intermediate_size <= 0: raise ValueError("grouped_mlp_contrastive_intermediate_size must be > 0") if branch_intermediate_size % self.num_groups != 0: raise ValueError("contrastive branch intermediate must divide across sparse_mlp_num_groups") self.branch_intermediate_size = branch_intermediate_size self.branch_group_size = branch_intermediate_size // self.num_groups self.branch_up_weight = nn.Parameter(torch.empty(self.num_groups, config.n_embd, self.branch_group_size)) self.branch_value_weight = nn.Parameter(torch.empty(self.num_groups, config.n_embd, self.branch_group_size)) self.branch_down_weight = nn.Parameter(torch.empty(self.num_groups, self.branch_group_size, config.n_embd)) self.token_gate = nn.Linear(config.n_embd, 1, bias=True) gate_max = float(config.grouped_mlp_contrastive_gate_max) if gate_max <= 0: raise ValueError("grouped_mlp_contrastive_gate_max must be > 0") self.contrastive_gate_max = gate_max gate = float(config.grouped_mlp_contrastive_gate) ratio = max(-0.999, min(0.999, gate / gate_max)) self.initial_gate_raw = math.atanh(ratio) self.block_idx = 0 self.n_layer = 1 self.reset_branch_parameters() def reset_branch_parameters(self) -> None: std = math.sqrt(2.0 / 5 / self.config.n_embd) branch_std = std * float(getattr(self.config, "grouped_mlp_stack_init_scale", 1.0)) nn.init.normal_(self.branch_up_weight, mean=0.0, std=branch_std) nn.init.normal_(self.branch_value_weight, mean=0.0, std=branch_std) nn.init.normal_(self.branch_down_weight, mean=0.0, std=branch_std) nn.init.zeros_(self.token_gate.weight) nn.init.constant_(self.token_gate.bias, self.initial_gate_raw) def set_block_index(self, block_idx: int, n_layer: int) -> None: self.block_idx = block_idx self.n_layer = n_layer def _branch_forward(self, x_flat: torch.Tensor) -> torch.Tensor: hidden_size = x_flat.shape[-1] up_weight = self.branch_up_weight.permute(1, 0, 2).reshape(hidden_size, self.branch_intermediate_size) value_weight = self.branch_value_weight.permute(1, 0, 2).reshape(hidden_size, self.branch_intermediate_size) down_weight = self.branch_down_weight.reshape(self.branch_intermediate_size, hidden_size) return (F.silu(x_flat @ up_weight) * (x_flat @ value_weight)) @ down_weight def _match_token_rms(self, source: torch.Tensor, reference: torch.Tensor) -> torch.Tensor: source_rms = source.float().pow(2).mean(dim=-1, keepdim=True).sqrt() reference_rms = reference.float().pow(2).mean(dim=-1, keepdim=True).sqrt() scale = reference_rms / source_rms.clamp_min(1e-8) return source * scale.to(device=source.device, dtype=source.dtype) def _token_gate(self, x_flat: torch.Tensor) -> torch.Tensor: x_rms = x_flat.float().pow(2).mean(dim=-1, keepdim=True).sqrt().clamp_min(1e-8) x_normed = (x_flat.float() / x_rms).to(dtype=x_flat.dtype) raw_gate = self.token_gate(x_normed) return self.contrastive_gate_max * torch.tanh(raw_gate).to(dtype=x_flat.dtype) def forward(self, x: torch.Tensor) -> torch.Tensor: main_output = super().forward(x) start_layer = int(getattr(self.config, "grouped_mlp_branch_start_layer", 0)) if self.block_idx < start_layer: return main_output original_shape = x.shape x_flat = x.reshape(-1, original_shape[-1]) main_flat = main_output.reshape(-1, original_shape[-1]) branch_output = self._match_token_rms(self._branch_forward(x_flat), main_flat) output = main_flat + self._token_gate(x_flat) * branch_output return output.reshape(original_shape) class TileRoutedStackedLayersDSwiGLUMLP(nn.Module): """Compute-matched stack of full-active grouped 4/4 MLP stages. The total intermediate budget is split across ``grouped_mlp_stack_depth`` stages. Each stage is a full-active grouped 4/4 MLP. Earlier stages update the state consumed by later stages, giving group layers a serial path for communication without increasing total MLP parameters or matmul FLOPs. """ def __init__(self, config: Config, intermediate_size: int | None = None) -> None: super().__init__() self.intermediate_size = intermediate_size or config.intermediate_size self.num_groups = config.sparse_mlp_num_groups self.stack_depth = config.grouped_mlp_stack_depth if self.stack_depth < 1: raise ValueError("grouped_mlp_stack_depth must be >= 1") if config.sparse_mlp_max_active_groups != self.num_groups: raise ValueError("TileRoutedStackedLayersDSwiGLUMLP is a full-active 4/4 variant.") if self.intermediate_size % (self.num_groups * self.stack_depth) != 0: raise ValueError("intermediate_size must divide across stack_depth * sparse_mlp_num_groups") if config.bias: raise ValueError("TileRoutedStackedLayersDSwiGLUMLP currently expects bias=False.") self.stage_intermediate_size = self.intermediate_size // self.stack_depth self.group_size = self.stage_intermediate_size // self.num_groups self.active_groups = config.sparse_mlp_max_active_groups self.up_weight = nn.Parameter(torch.empty(self.stack_depth, self.num_groups, config.n_embd, self.group_size)) self.value_weight = nn.Parameter(torch.empty(self.stack_depth, self.num_groups, config.n_embd, self.group_size)) self.down_weight = nn.Parameter(torch.empty(self.stack_depth, self.num_groups, self.group_size, config.n_embd)) stack_alpha = torch.tensor(float(config.grouped_mlp_stack_alpha)) if config.grouped_mlp_stack_alpha_learnable: self.stack_alpha = nn.Parameter(stack_alpha) else: self.register_buffer("stack_alpha", stack_alpha, persistent=False) output_scale = torch.tensor(float(config.grouped_mlp_stack_output_scale)) if config.grouped_mlp_stack_output_scale_learnable: self.stack_output_scale = nn.Parameter(output_scale) else: self.register_buffer("stack_output_scale", output_scale, persistent=False) self.config = config self.last_gate_sparsity: torch.Tensor | None = None self.last_threshold_mean: torch.Tensor | None = None self.last_active_groups_mean: torch.Tensor | None = None self.reset_parameters() def reset_parameters(self) -> None: std = math.sqrt(2.0 / 5 / self.config.n_embd) std *= float(getattr(self.config, "grouped_mlp_stack_init_scale", 1.0)) nn.init.normal_(self.up_weight, mean=0.0, std=std) nn.init.normal_(self.value_weight, mean=0.0, std=std) nn.init.normal_(self.down_weight, mean=0.0, std=std) def _stage_forward(self, stage_idx: int, state: torch.Tensor) -> torch.Tensor: hidden_size = state.shape[-1] up_weight = self.up_weight[stage_idx].permute(1, 0, 2).reshape(hidden_size, self.stage_intermediate_size) value_weight = self.value_weight[stage_idx].permute(1, 0, 2).reshape(hidden_size, self.stage_intermediate_size) down_weight = self.down_weight[stage_idx].reshape(self.stage_intermediate_size, hidden_size) return (F.silu(state @ up_weight) * (state @ value_weight)) @ down_weight def forward(self, x: torch.Tensor) -> torch.Tensor: original_shape = x.shape x_flat = x.reshape(-1, original_shape[-1]) state = x_flat output = torch.zeros_like(x_flat) stack_alpha = self.stack_alpha.to(device=x_flat.device, dtype=x_flat.dtype) for stage_idx in range(self.stack_depth): delta = self._stage_forward(stage_idx, state) output = output + delta if stage_idx + 1 < self.stack_depth: state = state + stack_alpha * delta if not torch.is_grad_enabled(): self.last_gate_sparsity = x_flat.new_tensor(0.0) self.last_threshold_mean = x_flat.new_tensor(0.0) self.last_active_groups_mean = x_flat.new_tensor(float(self.active_groups)) output_scale = self.stack_output_scale.to(device=x_flat.device, dtype=x_flat.dtype) return (output * output_scale).reshape(original_shape) class TileRoutedRMSBoundedStackedLayersDSwiGLUMLP(TileRoutedStackedLayersDSwiGLUMLP): """Stacked grouped layers with bounded residual update RMS. The unbounded split-stack d2 alpha -0.03 improves some choice benchmarks but produces MLP updates about four times stronger than dense/grouped MLPs. Here ``grouped_mlp_stack_output_scale`` is interpreted as a target output/input RMS ratio, keeping the serial group computation while limiting how hard it can perturb the residual stream. """ def forward(self, x: torch.Tensor) -> torch.Tensor: original_shape = x.shape x_flat = x.reshape(-1, original_shape[-1]) state = x_flat output = torch.zeros_like(x_flat) stack_alpha = self.stack_alpha.to(device=x_flat.device, dtype=x_flat.dtype) for stage_idx in range(self.stack_depth): delta = self._stage_forward(stage_idx, state) output = output + delta if stage_idx + 1 < self.stack_depth: state = state + stack_alpha * delta target_ratio = self.stack_output_scale.to(device=x_flat.device, dtype=x_flat.dtype) input_rms = x_flat.float().pow(2).mean().sqrt().to(dtype=x_flat.dtype) output_rms = output.float().pow(2).mean().sqrt().to(dtype=x_flat.dtype) scale = target_ratio * input_rms / output_rms.clamp_min(torch.finfo(output.dtype).eps) output = output * scale if not torch.is_grad_enabled(): self.last_gate_sparsity = x_flat.new_tensor(0.0) self.last_threshold_mean = x_flat.new_tensor(0.0) self.last_active_groups_mean = x_flat.new_tensor(float(self.active_groups)) return output.reshape(original_shape) class TileRoutedBlockLayeredNormDSwiGLUMLP(nn.Module): """Deep MLP with block-local RMS normalization at each layer. Unlike split-stack which uses serial state updates, this has real layers where the output of one layer feeds into the next layer. Each layer has block-normalized grouped MLPs. This tests whether deep blocknorm layers improve LM quality without the serial state perturbation of split-stack. The depth is controlled by `grouped_mlp_stack_depth`, but each layer is a full feedforward pass, not a state update. """ def __init__(self, config: Config, intermediate_size: int | None = None) -> None: super().__init__() self.intermediate_size = intermediate_size or config.intermediate_size self.num_groups = config.sparse_mlp_num_groups self.num_layers = config.grouped_mlp_stack_depth self.config = config if config.n_embd % self.num_groups != 0: raise ValueError("n_embd must be divisible by sparse_mlp_num_groups") self.block_hidden_size = config.n_embd // self.num_groups self.group_size = self.intermediate_size // self.num_groups # Each layer has its own grouped weights and block norm weights self.layer_up_weight = nn.Parameter(torch.empty(self.num_layers, self.num_groups, config.n_embd, self.group_size)) self.layer_value_weight = nn.Parameter(torch.empty(self.num_layers, self.num_groups, config.n_embd, self.group_size)) self.layer_down_weight = nn.Parameter(torch.empty(self.num_layers, self.num_groups, self.group_size, config.n_embd)) self.layer_norm_weight = nn.Parameter(torch.ones(self.num_layers, self.num_groups, config.n_embd)) self.reset_parameters() def reset_parameters(self) -> None: std = math.sqrt(2.0 / 5 / self.config.n_embd) nn.init.normal_(self.layer_up_weight, mean=0.0, std=std) nn.init.normal_(self.layer_value_weight, mean=0.0, std=std) nn.init.normal_(self.layer_down_weight, mean=0.0, std=std) def _block_norm(self, x_flat: torch.Tensor, layer_idx: int, group_idx: int) -> torch.Tensor: x_blocks = x_flat.reshape(x_flat.shape[0], self.num_groups, self.block_hidden_size) scale = torch.rsqrt(x_blocks.pow(2).mean(dim=-1, keepdim=True) + self.config.norm_eps) x_normed = (x_blocks * scale).reshape_as(x_flat) weight = self.layer_norm_weight[layer_idx, group_idx].to(device=x_flat.device, dtype=x_flat.dtype) return x_normed * weight def _layer_forward(self, x_flat: torch.Tensor, layer_idx: int) -> torch.Tensor: hidden_size = x_flat.shape[-1] # Reshape layer weights for grouped computation up_weight = self.layer_up_weight[layer_idx].permute(1, 0, 2).reshape(hidden_size, self.intermediate_size) value_weight = self.layer_value_weight[layer_idx].permute(1, 0, 2).reshape(hidden_size, self.intermediate_size) down_weight = self.layer_down_weight[layer_idx].reshape(self.intermediate_size, hidden_size) # Apply block normalization and compute grouped output output = torch.zeros_like(x_flat) for group_idx in range(self.num_groups): group_input = self._block_norm(x_flat, layer_idx, group_idx) hidden = F.silu(group_input @ up_weight) * (group_input @ value_weight) output = output + hidden @ down_weight return output def forward(self, x: torch.Tensor) -> torch.Tensor: original_shape = x.shape x_flat = x.reshape(-1, original_shape[-1]) # Pass through each layer sequentially (real layers, not state updates) for layer_idx in range(self.num_layers): x_flat = self._layer_forward(x_flat, layer_idx) return x_flat.reshape(original_shape) class TileRoutedRegulatedStackedLayersDSwiGLUMLP(TileRoutedRMSBoundedStackedLayersDSwiGLUMLP): def _match_global_rms(self, source: torch.Tensor, reference: torch.Tensor) -> torch.Tensor: source_rms = source.float().pow(2).mean().sqrt() reference_rms = reference.float().pow(2).mean().sqrt() scale = reference_rms / source_rms.clamp_min(1e-8) return source * scale.to(device=source.device, dtype=source.dtype) def _rms_bound_output(self, x_flat: torch.Tensor, output: torch.Tensor) -> torch.Tensor: target_ratio = self.stack_output_scale.to(device=x_flat.device, dtype=x_flat.dtype) input_rms = x_flat.float().pow(2).mean().sqrt().to(dtype=x_flat.dtype) output_rms = output.float().pow(2).mean().sqrt().to(dtype=x_flat.dtype) scale = target_ratio * input_rms / output_rms.clamp_min(torch.finfo(output.dtype).eps) return output * scale class TileRoutedStageNormRMSBoundedStackedLayersDSwiGLUMLP(TileRoutedRegulatedStackedLayersDSwiGLUMLP): """Normalize each inter-stage state update before the next stage reads it.""" def forward(self, x: torch.Tensor) -> torch.Tensor: original_shape = x.shape x_flat = x.reshape(-1, original_shape[-1]) state = x_flat output = torch.zeros_like(x_flat) stack_alpha = self.stack_alpha.to(device=x_flat.device, dtype=x_flat.dtype) for stage_idx in range(self.stack_depth): delta = self._stage_forward(stage_idx, state) output = output + delta if stage_idx + 1 < self.stack_depth: state_delta = self._match_global_rms(delta, state) state = state + stack_alpha * state_delta output = self._rms_bound_output(x_flat, output) if not torch.is_grad_enabled(): self.last_gate_sparsity = x_flat.new_tensor(0.0) self.last_threshold_mean = x_flat.new_tensor(0.0) self.last_active_groups_mean = x_flat.new_tensor(float(self.active_groups)) return output.reshape(original_shape) class TileRoutedWeightedStageRMSBoundedStackedLayersDSwiGLUMLP(TileRoutedRegulatedStackedLayersDSwiGLUMLP): """Learn a soft blend over stage deltas instead of summing stages equally.""" def __init__(self, config: Config, intermediate_size: int | None = None) -> None: super().__init__(config, intermediate_size) self.raw_stage_weight = nn.Parameter(torch.zeros(self.stack_depth)) def forward(self, x: torch.Tensor) -> torch.Tensor: original_shape = x.shape x_flat = x.reshape(-1, original_shape[-1]) state = x_flat output = torch.zeros_like(x_flat) stack_alpha = self.stack_alpha.to(device=x_flat.device, dtype=x_flat.dtype) weights = F.softmax(self.raw_stage_weight.float(), dim=0).to(device=x_flat.device, dtype=x_flat.dtype) for stage_idx in range(self.stack_depth): delta = self._stage_forward(stage_idx, state) output = output + weights[stage_idx] * self.stack_depth * delta if stage_idx + 1 < self.stack_depth: state_delta = self._match_global_rms(delta, state) state = state + stack_alpha * state_delta output = self._rms_bound_output(x_flat, output) if not torch.is_grad_enabled(): self.last_gate_sparsity = x_flat.new_tensor(0.0) self.last_threshold_mean = x_flat.new_tensor(0.0) self.last_active_groups_mean = x_flat.new_tensor(float(self.active_groups)) return output.reshape(original_shape) class TileRoutedMomentumRMSBoundedStackedLayersDSwiGLUMLP(TileRoutedRegulatedStackedLayersDSwiGLUMLP): """Use a momentum stream over stage deltas before updating later stages.""" def __init__(self, config: Config, intermediate_size: int | None = None) -> None: super().__init__(config, intermediate_size) beta = min(max(float(config.grouped_mlp_mix_alpha), 1e-4), 1 - 1e-4) self.raw_momentum_beta = nn.Parameter(torch.tensor(math.log(beta / (1.0 - beta)))) def forward(self, x: torch.Tensor) -> torch.Tensor: original_shape = x.shape x_flat = x.reshape(-1, original_shape[-1]) state = x_flat output = torch.zeros_like(x_flat) momentum: torch.Tensor | None = None beta = torch.sigmoid(self.raw_momentum_beta).to(device=x_flat.device, dtype=x_flat.dtype) stack_alpha = self.stack_alpha.to(device=x_flat.device, dtype=x_flat.dtype) for stage_idx in range(self.stack_depth): delta = self._stage_forward(stage_idx, state) momentum = delta if momentum is None else beta * momentum + (1.0 - beta) * delta output = output + momentum if stage_idx + 1 < self.stack_depth: state_delta = self._match_global_rms(momentum, state) state = state + stack_alpha * state_delta output = self._rms_bound_output(x_flat, output) if not torch.is_grad_enabled(): self.last_gate_sparsity = x_flat.new_tensor(0.0) self.last_threshold_mean = x_flat.new_tensor(0.0) self.last_active_groups_mean = x_flat.new_tensor(float(self.active_groups)) return output.reshape(original_shape) class TileRoutedLateLayerStackedLayersDSwiGLUMLP(TileRoutedStackedLayersDSwiGLUMLP): """Split-stack only active in late layers where ranking matters most. Early layers use standard grouped 4/4 MLP for clean LM learning. Late layers (config.splitstack_late_layer_start and above) use split-stack for stronger answer-choice ranking signal. This follows the momentum merge result where late8 5% injection preserved LM quality while improving OpenBookQA. """ def __init__(self, config: Config, intermediate_size: int | None = None) -> None: super().__init__(config, intermediate_size) self.late_layer_start = getattr(config, "splitstack_late_layer_start", 0) self.block_idx = 0 self.is_late_layer = False def set_block_index(self, block_idx: int, n_layer: int) -> None: self.block_idx = block_idx self.is_late_layer = block_idx >= self.late_layer_start def forward(self, x: torch.Tensor) -> torch.Tensor: if not self.is_late_layer: # Early layers: use single-stage grouped MLP (standard grouped 4/4) original_shape = x.shape x_flat = x.reshape(-1, original_shape[-1]) hidden_size = x_flat.shape[-1] up_weight = self.up_weight[0].permute(1, 0, 2).reshape(hidden_size, self.stage_intermediate_size) value_weight = self.value_weight[0].permute(1, 0, 2).reshape(hidden_size, self.stage_intermediate_size) down_weight = self.down_weight[0].reshape(self.stage_intermediate_size, hidden_size) output = (F.silu(x_flat @ up_weight) * (x_flat @ value_weight)) @ down_weight output_scale = self.stack_output_scale.to(device=x_flat.device, dtype=x_flat.dtype) return (output * output_scale).reshape(original_shape) else: # Late layers: use full split-stack return super().forward(x) class TileRoutedGatedStackedLayersDSwiGLUMLP(TileRoutedRMSBoundedStackedLayersDSwiGLUMLP): """Split-stack with learned per-layer contribution gate. A scalar gate controls how much split-stack contributes vs a baseline grouped MLP. This allows the model to learn when strong residual perturbations are helpful (choice ranking) vs harmful (general LM modeling). """ def __init__(self, config: Config, intermediate_size: int | None = None) -> None: super().__init__(config, intermediate_size) initial_gate = getattr(config, "splitstack_initial_gate", 0.1) self.raw_gate = nn.Parameter(torch.tensor(math.log(initial_gate / (1.0 - initial_gate)))) self.block_idx = 0 def set_block_index(self, block_idx: int, n_layer: int) -> None: self.block_idx = block_idx def forward(self, x: torch.Tensor) -> torch.Tensor: original_shape = x.shape x_flat = x.reshape(-1, original_shape[-1]) # Baseline grouped MLP (single-stage) hidden_size = x_flat.shape[-1] up_weight = self.up_weight[0].permute(1, 0, 2).reshape(hidden_size, self.stage_intermediate_size) value_weight = self.value_weight[0].permute(1, 0, 2).reshape(hidden_size, self.stage_intermediate_size) down_weight = self.down_weight[0].reshape(self.stage_intermediate_size, hidden_size) baseline_output = (F.silu(x_flat @ up_weight) * (x_flat @ value_weight)) @ down_weight # Full split-stack output state = x_flat stack_output = torch.zeros_like(x_flat) stack_alpha = self.stack_alpha.to(device=x_flat.device, dtype=x_flat.dtype) for stage_idx in range(self.stack_depth): delta = self._stage_forward(stage_idx, state) stack_output = stack_output + delta if stage_idx + 1 < self.stack_depth: state = state + stack_alpha * delta # RMS-bound the split-stack output target_ratio = self.stack_output_scale.to(device=x_flat.device, dtype=x_flat.dtype) input_rms = x_flat.float().pow(2).mean().sqrt().to(dtype=x_flat.dtype) stack_rms = stack_output.float().pow(2).mean().sqrt().to(dtype=x_flat.dtype) scale = target_ratio * input_rms / stack_rms.clamp_min(torch.finfo(stack_output.dtype).eps) stack_output = stack_output * scale # Learned gate between baseline and split-stack gate = torch.sigmoid(self.raw_gate).to(device=x_flat.device, dtype=x_flat.dtype) output = (1.0 - gate) * baseline_output + gate * stack_output if not torch.is_grad_enabled(): self.last_gate_sparsity = x_flat.new_tensor(0.0) self.last_threshold_mean = x_flat.new_tensor(0.0) self.last_active_groups_mean = x_flat.new_tensor(float(self.active_groups)) return output.reshape(original_shape) class TileRoutedAttentionOutputStackedLayersDSwiGLUMLP(TileRoutedStackedLayersDSwiGLUMLP): """Apply split-stack to attention output instead of MLP input. The split-stack processes the attention residual stream, then the standard grouped MLP processes the combined stream. This moves the strong residual perturbation to the attention path where it may be less disruptive to general LM modeling while still providing ranking signal. """ def __init__(self, config: Config, intermediate_size: int | None = None) -> None: super().__init__(config, intermediate_size) self.attention_stack_alpha = getattr(config, "splitstack_attention_alpha", -0.03) self.attention_stack_alpha_tensor = torch.tensor(self.attention_stack_alpha) def forward(self, x: torch.Tensor) -> torch.Tensor: # This MLP processes the combined stream after attention + split-stack # The split-stack logic is handled at the Block level # For now, use standard grouped MLP behavior original_shape = x.shape x_flat = x.reshape(-1, original_shape[-1]) hidden_size = x_flat.shape[-1] up_weight = self.up_weight[0].permute(1, 0, 2).reshape(hidden_size, self.stage_intermediate_size) value_weight = self.value_weight[0].permute(1, 0, 2).reshape(hidden_size, self.stage_intermediate_size) down_weight = self.down_weight[0].reshape(self.stage_intermediate_size, hidden_size) output = (F.silu(x_flat @ up_weight) * (x_flat @ value_weight)) @ down_weight output_scale = self.stack_output_scale.to(device=x_flat.device, dtype=x_flat.dtype) return (output * output_scale).reshape(original_shape) class TileRoutedStaticDSwiGLUMLP(TileRoutedDSwiGLUMLP): """Static tile-routed dSwiGLU without the per-route Python math loop. The active route pattern is fixed: tile group `g` uses expert groups `g, g + 1, ...` modulo `num_groups`. This keeps the same structural routing as `TileRoutedDSwiGLUMLP`, but executes the active routes as one batched matmul stack instead of rolling weights and launching one MLP path per route. """ def __init__(self, config: Config, intermediate_size: int | None = None) -> None: super().__init__(config, intermediate_size) route_indices = [ (torch.arange(self.num_groups, dtype=torch.long) + route_offset) % self.num_groups for route_offset in range(self.active_groups) ] self.register_buffer("route_indices", torch.stack(route_indices, dim=0), persistent=False) def forward(self, x: torch.Tensor) -> torch.Tensor: original_shape = x.shape x_flat = x.reshape(-1, original_shape[-1]) usable_tokens = x_flat.shape[0] - (x_flat.shape[0] % self.num_groups) output = x_flat.new_empty(x_flat.shape) if usable_tokens: tokens_per_group = usable_tokens // self.num_groups x_tiles = x_flat[:usable_tokens].reshape(self.num_groups, tokens_per_group, x_flat.shape[-1]) route_idx = self.route_indices.to(device=x_flat.device) route_up = self.up_weight.index_select(0, route_idx.reshape(-1)).reshape( self.active_groups, self.num_groups, x_flat.shape[-1], self.group_size ) route_value = self.value_weight.index_select(0, route_idx.reshape(-1)).reshape( self.active_groups, self.num_groups, x_flat.shape[-1], self.group_size ) route_down = self.down_weight.index_select(0, route_idx.reshape(-1)).reshape( self.active_groups, self.num_groups, self.group_size, x_flat.shape[-1] ) x_routes = x_tiles.unsqueeze(0).expand(self.active_groups, -1, -1, -1) up = torch.matmul(x_routes, route_up) value = torch.matmul(x_routes, route_value) out_tiles = torch.matmul(F.silu(up) * value, route_down).sum(dim=0) output[:usable_tokens] = out_tiles.reshape(usable_tokens, x_flat.shape[-1]) if usable_tokens < x_flat.shape[0]: tail = x_flat[usable_tokens:] tail_out = torch.zeros_like(tail) for route_offset in range(self.active_groups): group_id = route_offset % self.num_groups hidden = F.silu(tail @ self.up_weight[group_id]) * (tail @ self.value_weight[group_id]) tail_out = tail_out + hidden @ self.down_weight[group_id] output[usable_tokens:] = tail_out if not torch.is_grad_enabled(): self.last_gate_sparsity = x_flat.new_tensor(1.0 - self.active_groups / self.num_groups) self.last_threshold_mean = x_flat.new_tensor(0.0) self.last_active_groups_mean = x_flat.new_tensor(float(self.active_groups)) return self._apply_shared_path(x, output.reshape(original_shape)) class TileRoutedDSwiGLUMLPStaticA2(nn.Module): """Hard-specialized fixed-route tile dSwiGLU for the 1M model. This variant intentionally trades flexibility for a simpler execution shape: - hidden size 128 - intermediate size 1024 - 8 fixed tile groups - exactly 2 active fixed routes per tile There is no router, threshold, argsort, gather/scatter, or torch.roll in the forward path. Route selection is encoded directly in the parameter layout: `[active_route, tile_group, hidden, group_size]`. """ def __init__(self, config: Config, intermediate_size: int | None = None) -> None: super().__init__() self.intermediate_size = intermediate_size or config.intermediate_size self.num_groups = config.sparse_mlp_num_groups self.active_groups = config.sparse_mlp_max_active_groups if config.bias: raise ValueError("TileRoutedDSwiGLUMLPStaticA2 expects bias=False.") if config.n_embd != 128: raise ValueError("TileRoutedDSwiGLUMLPStaticA2 is specialized for hidden size 128.") if self.intermediate_size != 1024: raise ValueError("TileRoutedDSwiGLUMLPStaticA2 is specialized for intermediate size 1024.") if self.num_groups != 8: raise ValueError("TileRoutedDSwiGLUMLPStaticA2 is specialized for 8 groups.") if self.active_groups != 2 or config.sparse_mlp_min_active_groups != 2: raise ValueError("TileRoutedDSwiGLUMLPStaticA2 is specialized for active groups 2.") self.group_size = self.intermediate_size // self.num_groups self.up_weight = nn.Parameter(torch.empty(2, 8, 128, self.group_size)) self.value_weight = nn.Parameter(torch.empty(2, 8, 128, self.group_size)) self.down_weight = nn.Parameter(torch.empty(2, 8, self.group_size, 128)) self.config = config self.last_gate_sparsity: torch.Tensor | None = None self.last_threshold_mean: torch.Tensor | None = None self.last_active_groups_mean: torch.Tensor | None = None self.reset_parameters() def reset_parameters(self) -> None: std = math.sqrt(2.0 / 5 / self.config.n_embd) nn.init.normal_(self.up_weight, mean=0.0, std=std) nn.init.normal_(self.value_weight, mean=0.0, std=std) nn.init.normal_(self.down_weight, mean=0.0, std=std) def forward(self, x: torch.Tensor) -> torch.Tensor: original_shape = x.shape x_flat = x.reshape(-1, 128) usable_tokens = x_flat.shape[0] - (x_flat.shape[0] % 8) output = x_flat.new_empty(x_flat.shape) if usable_tokens: tokens_per_group = usable_tokens // 8 x_tiles = x_flat[:usable_tokens].reshape(8, tokens_per_group, 128) x_routes = x_tiles.unsqueeze(0) up = torch.matmul(x_routes, self.up_weight) value = torch.matmul(x_routes, self.value_weight) out_tiles = torch.matmul(F.silu(up) * value, self.down_weight).sum(dim=0) output[:usable_tokens] = out_tiles.reshape(usable_tokens, 128) if usable_tokens < x_flat.shape[0]: tail = x_flat[usable_tokens:] tail_routes = tail.view(1, 1, -1, 128) up = torch.matmul(tail_routes, self.up_weight[:, :1]) value = torch.matmul(tail_routes, self.value_weight[:, :1]) tail_out = torch.matmul(F.silu(up) * value, self.down_weight[:, :1]).sum(dim=(0, 1)) output[usable_tokens:] = tail_out self.last_gate_sparsity = x_flat.new_tensor(0.75) self.last_threshold_mean = x_flat.new_tensor(0.0) self.last_active_groups_mean = x_flat.new_tensor(2.0) return output.reshape(original_shape) class TileRoutedDSwiGLUMLPStaticGPTS(nn.Module): """Hard-specialized fixed-route tile dSwiGLU for GPT-S-5M shape. Exact shape: - hidden size 192 - intermediate size 672 - 6 fixed tile groups - 2 active fixed routes per tile This intentionally uses a route-major parameter layout to avoid dynamic routing overhead during training: no router, threshold, argsort, gather/scatter, or torch.roll in the forward path. """ def __init__(self, config: Config, intermediate_size: int | None = None) -> None: super().__init__() self.intermediate_size = intermediate_size or config.intermediate_size self.num_groups = config.sparse_mlp_num_groups self.active_groups = config.sparse_mlp_max_active_groups if config.bias: raise ValueError("TileRoutedDSwiGLUMLPStaticGPTS expects bias=False.") if config.n_embd != 192: raise ValueError("TileRoutedDSwiGLUMLPStaticGPTS is specialized for hidden size 192.") if self.intermediate_size != 672: raise ValueError("TileRoutedDSwiGLUMLPStaticGPTS is specialized for intermediate size 672.") if self.num_groups != 6: raise ValueError("TileRoutedDSwiGLUMLPStaticGPTS is specialized for 6 groups.") if self.active_groups != 2 or config.sparse_mlp_min_active_groups != 2: raise ValueError("TileRoutedDSwiGLUMLPStaticGPTS is specialized for active groups 2.") self.group_size = self.intermediate_size // self.num_groups self.up_weight = nn.Parameter(torch.empty(2, 6, 192, self.group_size)) self.value_weight = nn.Parameter(torch.empty(2, 6, 192, self.group_size)) self.down_weight = nn.Parameter(torch.empty(2, 6, self.group_size, 192)) self.config = config self.last_gate_sparsity: torch.Tensor | None = None self.last_threshold_mean: torch.Tensor | None = None self.last_active_groups_mean: torch.Tensor | None = None self.reset_parameters() def reset_parameters(self) -> None: std = math.sqrt(2.0 / 5 / self.config.n_embd) nn.init.normal_(self.up_weight, mean=0.0, std=std) nn.init.normal_(self.value_weight, mean=0.0, std=std) nn.init.normal_(self.down_weight, mean=0.0, std=std) def forward(self, x: torch.Tensor) -> torch.Tensor: original_shape = x.shape x_flat = x.reshape(-1, 192) usable_tokens = x_flat.shape[0] - (x_flat.shape[0] % 6) output = x_flat.new_empty(x_flat.shape) if usable_tokens: tokens_per_group = usable_tokens // 6 x_tiles = x_flat[:usable_tokens].reshape(6, tokens_per_group, 192) x_routes = x_tiles.unsqueeze(0) up = torch.matmul(x_routes, self.up_weight) value = torch.matmul(x_routes, self.value_weight) out_tiles = torch.matmul(F.silu(up) * value, self.down_weight).sum(dim=0) output[:usable_tokens] = out_tiles.reshape(usable_tokens, 192) if usable_tokens < x_flat.shape[0]: tail = x_flat[usable_tokens:] tail_routes = tail.view(1, 1, -1, 192) up = torch.matmul(tail_routes, self.up_weight[:, :1]) value = torch.matmul(tail_routes, self.value_weight[:, :1]) tail_out = torch.matmul(F.silu(up) * value, self.down_weight[:, :1]).sum(dim=(0, 1)) output[usable_tokens:] = tail_out self.last_gate_sparsity = x_flat.new_tensor(2.0 / 3.0) self.last_threshold_mean = x_flat.new_tensor(0.0) self.last_active_groups_mean = x_flat.new_tensor(2.0) return output.reshape(original_shape) class GemmaMLP(LLaMAMLP): def forward(self, x: torch.Tensor) -> torch.Tensor: x_fc_1 = self.fc_1(x) x_fc_2 = self.fc_2(x) x = F.gelu(x_fc_1, approximate=self.config.gelu_approximate) * x_fc_2 return self.proj(x) class LLaMAMoE(nn.Module): def __init__(self, config: Config) -> None: super().__init__() self.gate = ( nn.Linear(config.n_embd, config.n_expert, bias=False) if not config.n_expert_groups else GroupedTopkRouter(config) ) self.experts = nn.ModuleList( LLaMAMLP(config, intermediate_size=config.moe_intermediate_size) for _ in range(config.n_expert) ) if config.n_shared_expert: self.shared_experts = LLaMAMLP( config, intermediate_size=config.moe_intermediate_size * config.n_shared_expert ) self.config = config def forward(self, x: torch.Tensor) -> torch.Tensor: """ Derived from: https://github.com/mistralai/mistral-src/blob/b46d6/moe_one_file_ref.py#L203-L219 See also figure 1 in https://arxiv.org/abs/2211.15841 """ B, T, C = x.size() # batch size, sequence length, embedding dimensionality (n_embd) residual_x = x.clone() x = x.view(-1, C) # (B*T, C) if not self.config.n_expert_groups: router = self.gate(x) # (B*T, n_expert) probs, indices = torch.topk(router, self.config.n_expert_per_token) # (B*T, n_expert_per_token) probs = probs.softmax(dim=1, dtype=torch.float).to(dtype=x.dtype) else: probs, indices = self.gate(x) if self.config.routed_scaling_factor != 1.0: probs = probs * self.config.routed_scaling_factor masks = indices.unsqueeze(-1) == torch.arange(self.config.n_expert, device=x.device) masks = masks.permute(2, 0, 1) # (n_expert, B*T, n_expert_per_token) y = torch.zeros_like(x) # (B*T, C) for mask, expert in zip(masks, self.experts): token_idx, expert_idx = torch.where(mask) y[token_idx] += probs[token_idx, expert_idx, None] * expert(x[token_idx]) y = y.view(B, T, C) if self.config.n_shared_expert: y = y + self.shared_experts(residual_x) return y class GroupedTopkRouter(nn.Module): """ Derived from: https://github.com/huggingface/transformers/blob/main/src/transformers/models/deepseek_v3/modeling_deepseek_v3.py. DeepseekV3TopkRouter class. """ def __init__(self, config: Config) -> None: super().__init__() self.config = config self.weight = nn.Parameter(torch.empty(config.n_expert, config.n_embd)) self.register_buffer("e_score_correction_bias", torch.zeros(config.n_expert)) @torch.no_grad() def get_topk_indices(self, scores: torch.Tensor) -> torch.Tensor: scores_for_choice = scores.view(-1, self.config.n_expert) + self.e_score_correction_bias.unsqueeze(0) group_scores = ( scores_for_choice.view(-1, self.config.n_expert_groups, self.config.n_expert // self.config.n_expert_groups) .topk(self.config.n_topk_scores_per_group, dim=-1)[0] # Top k scores for each group .sum(dim=-1) ) group_idx = torch.topk(group_scores, k=self.config.n_topk_groups, dim=-1, sorted=False)[1] group_mask = torch.zeros_like(group_scores) group_mask.scatter_(1, group_idx, 1) score_mask = ( group_mask.unsqueeze(-1) .expand(-1, self.config.n_expert_groups, self.config.n_expert // self.config.n_expert_groups) .reshape(-1, self.config.n_expert) ) scores_for_choice = scores_for_choice.masked_fill(~score_mask.bool(), 0.0) topk_indices = torch.topk(scores_for_choice, k=self.config.n_expert_per_token, dim=-1, sorted=False)[1] return topk_indices def forward(self, x: torch.Tensor) -> torch.Tensor: router_logits = F.linear(x.type(torch.float32), self.weight.type(torch.float32)) scores = router_logits.sigmoid() topk_indices = self.get_topk_indices(scores) topk_weights = scores.gather(1, topk_indices) if self.config.norm_topk_prob: denominator = topk_weights.sum(dim=-1, keepdim=True) + 1e-20 topk_weights /= denominator return topk_weights, topk_indices # ROPE: YaRN (Yet another RoPE extensioN) scaling function for extended context def yarn_get_mscale(scale=1, mscale=1): if scale <= 1: return 1.0 return 0.1 * mscale * math.log(scale) + 1.0 def build_rope_cache( seq_len: int, n_elem: int, device: torch.device | None = None, base: int = 10000, condense_ratio: int = 1, extra_config: dict | None = None, rope_local_base_freq: float | None = None, ) -> tuple[torch.Tensor, torch.Tensor]: """ Enhanced Transformer with Rotary Position Embedding. Args: seq_len (int): Sequence length. n_elem (int): Number of elements (head dimension). device (torch.device, optional): Device for tensor allocations. base (int, optional): Base for computing inverse frequencies. condense_ratio (int, optional): Ratio to condense the position indices. extra_config (dict, optional): Configuration parameters for frequency adjustments (used by Llama 3.1 and 3.2) Returns: Tuple[torch.Tensor, torch.Tensor]: Cosine and sine caches for RoPE. Shapes are `(seq_len, n_elem)`. """ # Compute the inverse frequencies theta theta = 1.0 / (base ** (torch.arange(0, n_elem, 2, device=device).float() / n_elem)) # Initialize attention scaling factor (modified for YaRN) attention_scaling = 1.0 if extra_config is not None: factor = extra_config["factor"] # Check YaRN first (has beta_fast/beta_slow) if "beta_fast" in extra_config or "beta_slow" in extra_config: # YaRN-style RoPE scaling beta_fast = extra_config["beta_fast"] beta_slow = extra_config["beta_slow"] original_max_seq_len = extra_config["original_max_seq_len"] # Calculate attention scaling factor based on mscale and mscale_all_dim mscale = extra_config.get("mscale") mscale_all_dim = extra_config.get("mscale_all_dim") if mscale and mscale_all_dim: attention_scaling = yarn_get_mscale(factor, mscale) / yarn_get_mscale(factor, mscale_all_dim) elif mscale_all_dim: attention_scaling = yarn_get_mscale(factor, mscale_all_dim) elif mscale: attention_scaling = yarn_get_mscale(factor, mscale) # else: attention_scaling remains 1.0 # Create two frequency sets: extrapolation (unscaled) and interpolation (scaled) pos_freqs = base ** (torch.arange(0, n_elem, 2, device=device).float() / n_elem) theta_extrapolation = 1.0 / pos_freqs theta_interpolation = 1.0 / (factor * pos_freqs) # Find correction range based on rotation counts # Inverse dimension formula to find dimension based on number of rotations def find_correction_dim(num_rotations, dim, base_val, max_pos): return (dim * math.log(max_pos / (num_rotations * 2 * math.pi))) / (2 * math.log(base_val)) low_dim = find_correction_dim(beta_fast, n_elem, base, original_max_seq_len) high_dim = find_correction_dim(beta_slow, n_elem, base, original_max_seq_len) # Apply truncation if specified if extra_config.get("truncate", True): low_dim = math.floor(low_dim) high_dim = math.ceil(high_dim) low_dim = max(low_dim, 0) high_dim = min(high_dim, n_elem // 2 - 1) # Create linear ramp factor for blending dim_range = torch.arange(n_elem // 2, device=device, dtype=torch.float32) if low_dim == high_dim: high_dim += 0.001 # Prevent singularity linear_func = (dim_range - low_dim) / (high_dim - low_dim) ramp_func = torch.clamp(linear_func, 0.0, 1.0) # Blend extrapolation and interpolation frequencies # ramp_func = 0 -> use interpolation (scaled), ramp_func = 1 -> use extrapolation (unscaled) theta_extrapolation_factor = ramp_func theta = ( theta_interpolation * (1 - theta_extrapolation_factor) + theta_extrapolation * theta_extrapolation_factor ) elif "original_max_seq_len" in extra_config: # Llama3-style RoPE scaling orig_context_len = extra_config["original_max_seq_len"] low_freq_factor = extra_config["low_freq_factor"] high_freq_factor = extra_config["high_freq_factor"] wavelen = 2 * torch.pi / theta ratio = orig_context_len / wavelen smooth_factor = (ratio - low_freq_factor) / (high_freq_factor - low_freq_factor) smooth_factor = torch.clamp(smooth_factor, min=0.0, max=1.0) # Compute adjusted_theta without masked indexing adjusted_theta = (1 - smooth_factor) * (theta / factor) + smooth_factor * theta theta = adjusted_theta else: # Linear scaling fallback theta = theta / factor # Create position indices `[0, 1, ..., seq_len - 1]` seq_idx = torch.arange(seq_len, device=device).float() / condense_ratio # Calculate the product of position index and $\theta_i$ idx_theta = torch.outer(seq_idx, theta).repeat(1, 2) # If `n_elem` is odd, the final dimension of `idx_theta` has size # `n_elem + 1`, so need to cut something off. # Due to a current bug in Hugging Face, in the case `n_elem == 1`, we leave # `idx_theta`, `cos`, `sin` as is. Things work out in `apply_rope` due to # broadcasting. If we shorten `idx_theta`, unit tests comparing to # Hugging Face fail. # https://github.com/huggingface/transformers/issues/35233 if idx_theta.shape[-1] > n_elem > 1: idx_theta = idx_theta[..., :n_elem] # if rope_local_base_freq is given, have a separate rope value for local embedding # For now, we use default RoPE for local embedding if rope_local_base_freq is not None: local_theta = 1.0 / (rope_local_base_freq ** (torch.arange(0, n_elem, 2, device=device).float() / n_elem)) local_idx_theta = torch.outer(seq_idx, local_theta) local_idx_theta = local_idx_theta.repeat(1, 2) if local_idx_theta.shape[-1] > n_elem > 1: local_idx_theta = local_idx_theta[..., :n_elem] idx_theta = torch.stack((idx_theta, local_idx_theta), dim=-1) cos = torch.cos(idx_theta) * attention_scaling sin = torch.sin(idx_theta) * attention_scaling return cos, sin def batched_index_select(t, dim, idx): """index_select for batched index and unbatched t""" if idx.dim() == 1: return torch.index_select(t, dim, idx) *batch_shape, idx_size = idx.shape res = torch.index_select(t, dim, idx.reshape(-1)) # flat index # split out single batch idx res = res.view(*t.shape[:dim], -1, idx_size, *t.shape[dim + 1 :]) if dim > 0: # move batch dim to front, this is np.rollaxis(res, dim, 0) for tensors dims = [dim] + list(range(res.dim())) del dims[dim + 1] res = res.permute(dims) # unflatten batch dims res = res.view(*batch_shape, *res.shape[1:]) return res def batched_index_copy_(t, dim, idx, val): """Index copy for batched t, idx, val""" if t.device.type == "mps": # Normalize negative dimensions if dim < 0: dim = t.dim() + dim if idx.dim() == 1: idx_shape = [1] * val.dim() idx_shape[dim] = -1 idx_expanded = idx.view(*idx_shape) idx_expanded = idx_expanded.expand_as(val) t.scatter_(dim, idx_expanded, val) return t elif idx.dim() == 2: assert dim != 0, "Cannot index the batch dimension" batch_size = idx.size(0) idx_size = idx.size(1) assert batch_size == t.size(0) == val.size(0) idx_shape = [batch_size] + [1] * (val.dim() - 1) idx_shape[dim] = idx_size idx_expanded = idx.view(*idx_shape) idx_expanded = idx_expanded.expand_as(val) t.scatter_(dim, idx_expanded, val) return t else: raise NotImplementedError(f"idx.dim() == {idx.dim()} not supported") else: if idx.dim() == 1: return t.index_copy_(dim, idx, val) assert idx.dim() == 2, f"multiple batch dims not yet {idx.shape=}" assert dim != 0, f"cannot index batch dim {dim=}" batch_size, idx_size = idx.shape assert batch_size == t.size(0) assert batch_size == val.size(0) # if we can view the batch and indexed dimensions together, we could # do index trickery. This is, sadly, not the case for kvcache so we # fall back to for loop for i in range(batch_size): unbatched_dim = dim if dim < 0 else dim - 1 t[i].index_copy_(unbatched_dim, idx[i], val[i]) return t def apply_rope(x: torch.Tensor, cos: torch.Tensor, sin: torch.Tensor) -> torch.Tensor: """ Applies RoPE transform to `x`. Note that `cos`, `sin` need to have a batch dimension. Args: x: Input tensor, `(B, ..., T, head_size)` cos: Cached cosines, `(B, T, head_size)` or `(1, T, head_size)` sin: Cached sines, `(B, T, head_size)` or `(1, T, head_size)` Returns: Encoded tensor, `(B, ..., T, head_size)` """ if cos.dim() != 3: raise ValueError(f"cos must be three-dimensional, but shape is {cos.shape}") if cos.shape != sin.shape: raise ValueError(f"cos, sin must have same shape, but cos.shape={cos.shape}, sin.shape={sin.shape}") head_size_half = x.size(-1) // 2 x1 = x[..., :head_size_half] # (B, ..., T, head_size/2) x2 = x[..., head_size_half:] # (B, ..., T, head_size/2) rotated = torch.cat((-x2, x1), dim=-1) # (B, ..., T, head_size) dims_diff = x.dim() - cos.dim() if dims_diff > 0: # Ensure that shapes of `x`, `cos`, `sin` align new_shape = cos.shape[0:1] + (1,) * dims_diff + cos.shape[1:] cos = cos.view(*new_shape) sin = sin.view(*new_shape) roped = (x * cos) + (rotated * sin) return roped.to(dtype=x.dtype) def apply_rope_interleave(x: torch.Tensor, cos: torch.Tensor, sin: torch.Tensor) -> torch.Tensor: """Apply rotary position embeddings with interleaved tensor layout. This version rearranges the input tensor to group even/odd indices separately before applying the standard RoPE rotation, matching HuggingFace's apply_rotary_pos_emb_interleave behavior. Args: x: Input tensor of shape (..., seq_len, head_dim) cos: Cosine component of shape (B, seq_len, head_dim) or (1, seq_len, head_dim) sin: Sine component of shape (B, seq_len, head_dim) or (1, seq_len, head_dim) Returns: Tensor with RoPE applied, same shape as input """ if cos.dim() != 3: raise ValueError(f"cos must be three-dimensional, but shape is {cos.shape}") if cos.shape != sin.shape: raise ValueError(f"cos, sin must have same shape, but cos.shape={cos.shape}, sin.shape={sin.shape}") # Rearrange tensor to group even/odd indices: [x0,x1,x2,x3,...] -> [x0,x2,x4,...,x1,x3,x5,...] *batch_dims, d = x.shape x = x.view(*batch_dims, d // 2, 2).transpose(-1, -2).reshape(*batch_dims, d) # Standard rotation logic (same as apply_rope) head_size_half = x.size(-1) // 2 x1 = x[..., :head_size_half] x2 = x[..., head_size_half:] rotated = torch.cat((-x2, x1), dim=-1) # Auto-detect dimension mismatch and reshape cos/sin dims_diff = x.dim() - cos.dim() if dims_diff > 0: new_shape = cos.shape[0:1] + (1,) * dims_diff + cos.shape[1:] cos = cos.view(*new_shape) sin = sin.view(*new_shape) roped = (x * cos) + (rotated * sin) return roped.to(dtype=x.dtype) def do_softcapping(x: torch.Tensor, thresh: float) -> torch.Tensor: return torch.tanh(x / thresh) * thresh class MLACompressedKVCache(nn.Module): """Cache normalized c^KV and the shared RoPE key instead of expanded K/V.""" def __init__( self, latent_shape: tuple[int, int, int], rope_shape: tuple[int, int, int, int], device: torch.device | None = None, dtype: torch.dtype | None = None, ) -> None: super().__init__() self.register_buffer("latent", torch.zeros(latent_shape, device=device, dtype=dtype), persistent=False) self.register_buffer("rope", torch.zeros(rope_shape, device=device, dtype=dtype), persistent=False) def forward( self, input_pos: torch.Tensor, latent: torch.Tensor, rope: torch.Tensor ) -> tuple[torch.Tensor, torch.Tensor]: if self.latent.dtype != latent.dtype: self.latent = self.latent.to(latent.dtype) if self.rope.dtype != rope.dtype: self.rope = self.rope.to(rope.dtype) bs = latent.size(0) if input_pos.dim() == 1: self.latent[:bs].index_copy_(1, input_pos, latent) self.rope[:bs].index_copy_(2, input_pos, rope) else: batch = torch.arange(bs, device=input_pos.device)[:, None] self.latent[batch, input_pos] = latent self.rope[:bs, 0][batch, input_pos] = rope[:, 0] return self.latent[:bs], self.rope[:bs] @property def bytes_per_token(self) -> int: return (self.latent.shape[-1] + self.rope.shape[-1]) * self.latent.element_size() class KVCache(nn.Module): """ Buffers `k`, `v` have shape `(batch_size, n_query_groups, max_seq_length, head_size)`. """ def __init__( self, k_shape: tuple[int, int, int, int], v_shape: tuple[int, int, int, int], device: torch.device | None = None, dtype: torch.dtype | None = None, is_sliding_window: bool = False, sliding_window_size: int | None = None, ) -> None: super().__init__() self.register_buffer("k", torch.zeros(k_shape, device=device, dtype=dtype), persistent=False) self.register_buffer("v", torch.zeros(v_shape, device=device, dtype=dtype), persistent=False) self.is_sliding_window = is_sliding_window self.sliding_window_size = sliding_window_size self.max_cache_len = k_shape[2] def forward(self, input_pos: torch.Tensor, k: torch.Tensor, v: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]: """ Writes new values `k` and `v` into the cache at the positions specified by `input_pos` along the sequence dimension (`max_seq_length`). The batch size of `k` and `v` (`bs`) must be smaller or equal to `KVCache` batch size. Returns the full buffers, adjusted to the batch size `bs`. Args: input_pos: Position index, `(bs, T)` or `(T,)` k: New values, `(bs, n_query_groups, T, head_size)` v: New values, `(bs, n_query_groups, T, head_size)` Returns: k_full, v_full, `(bs, n_query_groups, max_seq_length, head_size)` """ # move the buffer to the activation dtype for when AMP is used if self.k.dtype != k.dtype: self.k = self.k.to(k.dtype) if self.v.dtype != v.dtype: self.v = self.v.to(v.dtype) # update the cache bs = k.size(0) if self.is_sliding_window: # Circular buffer for sliding window prefill_len = input_pos.shape[-1] if prefill_len > self.max_cache_len: raise ValueError( f"Prefill length ({prefill_len}) exceeds the sliding window size ({self.max_cache_len}). " f"This causes the ring-buffer KV cache to overwrite entries, but the attention mask is not " f"rebuilt to reflect the true positions, which silently violates causality. " f"Please use chunked prefill with chunk size <= {self.max_cache_len} to avoid this issue." ) cache_positions = input_pos % self.max_cache_len k = batched_index_copy_(self.k[:bs, ...], -2, cache_positions, k) v = batched_index_copy_(self.v[:bs, ...], -2, cache_positions, v) max_pos = input_pos.max().item() if max_pos < self.max_cache_len: k = k[:, :, : max_pos + 1, :] v = v[:, :, : max_pos + 1, :] else: # Standard KV cache (global attention) k = batched_index_copy_(self.k[:bs, ...], -2, input_pos, k) v = batched_index_copy_(self.v[:bs, ...], -2, input_pos, v) return k, v def reset_parameters(self) -> None: torch.nn.init.zeros_(self.k) torch.nn.init.zeros_(self.v) def build_mask_cache(max_seq_length: int, device: torch.device | None = None) -> torch.Tensor: ones = torch.ones((max_seq_length, max_seq_length), device=device, dtype=torch.bool) return torch.tril(ones).unsqueeze(0).unsqueeze(0) class RMSNorm(torch.nn.Module): """Root Mean Square Layer Normalization. Derived from https://github.com/bzhangGo/rmsnorm/blob/master/rmsnorm_torch.py. BSD 3-Clause License: https://github.com/bzhangGo/rmsnorm/blob/master/LICENSE. """ def __init__(self, size: int, dim: int = -1, eps: float = 1e-6, add_unit_offset: bool = False) -> None: super().__init__() self.weight = torch.nn.Parameter(torch.ones(size)) self.eps = eps self.dim = dim self.add_unit_offset = add_unit_offset def forward(self, x: torch.Tensor) -> torch.Tensor: dtype = x.dtype x = x.float() # NOTE: the original RMSNorm paper implementation is not equivalent norm_x = torch.mean(x * x, dim=self.dim, keepdim=True) x_normed = x * torch.rsqrt(norm_x + self.eps) weight = (1 + self.weight) if self.add_unit_offset else self.weight return (x_normed * weight.float()).to(dtype=dtype) def reset_parameters(self) -> None: torch.nn.init.ones_(self.weight)