# Scatterbrain - A single MoE layer looped multiple times # # Based on Qwen3MoeForCausalLM architecture with looping logic from Loopstral import copy from typing import Callable, Optional, Union import torch import torch.nn.functional as F from torch import nn from transformers.activations import ACT2FN from transformers.cache_utils import Cache, DynamicCache, DynamicLayer # Try to import Liger kernels for optimized operations try: from liger_kernel.transformers import LigerRMSNorm, LigerSwiGLUMLP LIGER_AVAILABLE = True except ImportError: LIGER_AVAILABLE = False # Try to import ScatterMoE for optimized MoE computation try: from scattermoe import flatten_sort_count, parallel_linear SCATTERMOE_AVAILABLE = True except ImportError: SCATTERMOE_AVAILABLE = False from transformers.generation import GenerationMixin from transformers.integrations import use_kernel_forward_from_hub from transformers.masking_utils import create_causal_mask, create_sliding_window_causal_mask from transformers.modeling_flash_attention_utils import FlashAttentionKwargs from transformers.modeling_layers import GradientCheckpointingLayer from transformers.modeling_outputs import MoeCausalLMOutputWithPast, MoeModelOutputWithPast from transformers.modeling_rope_utils import ROPE_INIT_FUNCTIONS, dynamic_rope_update from transformers.modeling_utils import ALL_ATTENTION_FUNCTIONS, PreTrainedModel from transformers.processing_utils import Unpack from transformers.utils import TransformersKwargs, auto_docstring, can_return_tuple from transformers.utils.deprecation import deprecate_kwarg from .configuration_scatterbrain import ScatterbrainConfig def rotate_half(x): """Rotates half the hidden dims of the input.""" x1 = x[..., : x.shape[-1] // 2] x2 = x[..., x.shape[-1] // 2 :] return torch.cat((-x2, x1), dim=-1) def apply_rotary_pos_emb(q, k, cos, sin, position_ids=None, unsqueeze_dim=1): """Applies Rotary Position Embedding to the query and key tensors.""" cos = cos.unsqueeze(unsqueeze_dim) sin = sin.unsqueeze(unsqueeze_dim) q_embed = (q * cos) + (rotate_half(q) * sin) k_embed = (k * cos) + (rotate_half(k) * sin) return q_embed, k_embed def repeat_kv(hidden_states: torch.Tensor, n_rep: int) -> torch.Tensor: """Repeat KV heads for GQA.""" batch, num_key_value_heads, slen, head_dim = hidden_states.shape if n_rep == 1: return hidden_states hidden_states = hidden_states[:, :, None, :, :].expand(batch, num_key_value_heads, n_rep, slen, head_dim) return hidden_states.reshape(batch, num_key_value_heads * n_rep, slen, head_dim) def eager_attention_forward( module: nn.Module, query: torch.Tensor, key: torch.Tensor, value: torch.Tensor, attention_mask: Optional[torch.Tensor], scaling: float, dropout: float = 0.0, **kwargs: Unpack[TransformersKwargs], ): key_states = repeat_kv(key, module.num_key_value_groups) value_states = repeat_kv(value, module.num_key_value_groups) attn_weights = torch.matmul(query, key_states.transpose(2, 3)) * scaling if attention_mask is not None: causal_mask = attention_mask[:, :, :, : key_states.shape[-2]] attn_weights = attn_weights + causal_mask attn_weights = nn.functional.softmax(attn_weights, dim=-1, dtype=torch.float32).to(query.dtype) attn_weights = nn.functional.dropout(attn_weights, p=dropout, training=module.training) attn_output = torch.matmul(attn_weights, value_states) attn_output = attn_output.transpose(1, 2).contiguous() return attn_output, attn_weights class ScatterbrainAttention(nn.Module): """Multi-headed attention with QK normalization.""" def __init__(self, config: ScatterbrainConfig, layer_idx: int): super().__init__() self.config = config self.layer_idx = layer_idx self.head_dim = getattr(config, "head_dim", config.hidden_size // config.num_attention_heads) self.num_key_value_groups = config.num_attention_heads // config.num_key_value_heads self.scaling = self.head_dim**-0.5 self.attention_dropout = config.attention_dropout self.is_causal = True self.q_proj = nn.Linear( config.hidden_size, config.num_attention_heads * self.head_dim, bias=config.attention_bias ) self.k_proj = nn.Linear( config.hidden_size, config.num_key_value_heads * self.head_dim, bias=config.attention_bias ) self.v_proj = nn.Linear( config.hidden_size, config.num_key_value_heads * self.head_dim, bias=config.attention_bias ) self.o_proj = nn.Linear( config.num_attention_heads * self.head_dim, config.hidden_size, bias=config.attention_bias ) self.q_norm = ScatterbrainRMSNorm(self.head_dim, eps=config.rms_norm_eps) self.k_norm = ScatterbrainRMSNorm(self.head_dim, eps=config.rms_norm_eps) self.sliding_window = getattr(config, "sliding_window", None) @deprecate_kwarg("past_key_value", new_name="past_key_values", version="4.58") def forward( self, hidden_states: torch.Tensor, position_embeddings: tuple[torch.Tensor, torch.Tensor], attention_mask: Optional[torch.Tensor], past_key_values: Optional[Cache] = None, cache_position: Optional[torch.LongTensor] = None, cache_slot_idx: Optional[int] = None, **kwargs: Unpack[FlashAttentionKwargs], ) -> tuple[torch.Tensor, Optional[torch.Tensor]]: input_shape = hidden_states.shape[:-1] hidden_shape = (*input_shape, -1, self.head_dim) query_states = self.q_norm(self.q_proj(hidden_states).view(hidden_shape)).transpose(1, 2) key_states = self.k_norm(self.k_proj(hidden_states).view(hidden_shape)).transpose(1, 2) value_states = self.v_proj(hidden_states).view(hidden_shape).transpose(1, 2) cos, sin = position_embeddings query_states, key_states = apply_rotary_pos_emb(query_states, key_states, cos, sin) if past_key_values is not None: cache_kwargs = {"sin": sin, "cos": cos, "cache_position": cache_position} # Use cache_slot_idx for looped layers - each iteration gets its own cache slot slot_idx = cache_slot_idx if cache_slot_idx is not None else self.layer_idx key_states, value_states = past_key_values.update(key_states, value_states, slot_idx, cache_kwargs) attention_interface: Callable = eager_attention_forward if self.config._attn_implementation != "eager": attention_interface = ALL_ATTENTION_FUNCTIONS[self.config._attn_implementation] attn_output, attn_weights = attention_interface( self, query_states, key_states, value_states, attention_mask, dropout=0.0 if not self.training else self.attention_dropout, scaling=self.scaling, sliding_window=self.sliding_window, **kwargs, ) attn_output = attn_output.reshape(*input_shape, -1).contiguous() attn_output = self.o_proj(attn_output) return attn_output, attn_weights class ScatterbrainMLP(nn.Module): def __init__(self, config, intermediate_size=None): super().__init__() self.config = config self.hidden_size = config.hidden_size self.intermediate_size = intermediate_size if intermediate_size is not None else config.intermediate_size self.gate_proj = nn.Linear(self.hidden_size, self.intermediate_size, bias=False) self.up_proj = nn.Linear(self.hidden_size, self.intermediate_size, bias=False) self.down_proj = nn.Linear(self.intermediate_size, self.hidden_size, bias=False) self.act_fn = ACT2FN[config.hidden_act] def forward(self, x): # Fused SwiGLU: down(act(gate(x)) * up(x)) return self.down_proj(self.act_fn(self.gate_proj(x)) * self.up_proj(x)) class ScatterbrainSparseMoeBlock(nn.Module): """ MoE block with optimized expert computation using ScatterMoE Triton kernels. When ScatterMoE is available, uses fused Triton kernels for massive speedup. Falls back to Python loop implementation when ScatterMoE is not available. Key optimizations with ScatterMoE: 1. Fused scatter-gather operations in Triton 2. All experts processed in parallel on GPU 3. Efficient memory access patterns """ def __init__(self, config): super().__init__() self.num_experts = config.num_experts self.top_k = config.num_experts_per_tok self.norm_topk_prob = config.norm_topk_prob self.hidden_size = config.hidden_size self.moe_intermediate_size = config.moe_intermediate_size self.gate = nn.Linear(config.hidden_size, config.num_experts, bias=False) # Stacked expert weights as Parameters # Shape: [num_experts, out_features, in_features] self.expert_gate_proj = nn.Parameter( torch.empty(config.num_experts, config.moe_intermediate_size, config.hidden_size) ) self.expert_up_proj = nn.Parameter( torch.empty(config.num_experts, config.moe_intermediate_size, config.hidden_size) ) self.expert_down_proj = nn.Parameter( torch.empty(config.num_experts, config.hidden_size, config.moe_intermediate_size) ) # Initialize weights for param in [self.expert_gate_proj, self.expert_up_proj, self.expert_down_proj]: nn.init.kaiming_uniform_(param, a=5**0.5) def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: batch_size, sequence_length, hidden_dim = hidden_states.shape hidden_states_flat = hidden_states.view(-1, hidden_dim) num_tokens = hidden_states_flat.shape[0] # Router router_logits = self.gate(hidden_states_flat) routing_weights = F.softmax(router_logits, dim=1, dtype=torch.float) routing_weights, selected_experts = torch.topk(routing_weights, self.top_k, dim=-1) if self.norm_topk_prob: routing_weights = routing_weights / routing_weights.sum(dim=-1, keepdim=True) routing_weights = routing_weights.to(hidden_states_flat.dtype) if SCATTERMOE_AVAILABLE: # Use ScatterMoE's optimized Triton kernels final_hidden_states = self._forward_scattermoe( hidden_states_flat, selected_experts, routing_weights ) else: # Fallback to Python loop final_hidden_states = self._forward_loop( hidden_states_flat, selected_experts, routing_weights ) final_hidden_states = final_hidden_states.reshape(batch_size, sequence_length, hidden_dim) return final_hidden_states, router_logits def _forward_scattermoe( self, hidden_states: torch.Tensor, selected_experts: torch.Tensor, routing_weights: torch.Tensor, ) -> torch.Tensor: """Forward pass using ScatterMoE Triton kernels.""" # Get sorting indices and expert offsets from ScatterMoE sorted_expert_idxs, sorted_scattered_idxs, expert_offsets = flatten_sort_count( selected_experts, num_experts=self.num_experts ) # ScatterMoE parallel_linear expects weights in [E, in, out] format # Our weights are stored as [E, out, in], so we need to permute them # This matches what ParallelExperts.forward() does internally # First pass: gate and up projections (hidden -> intermediate) # We need to do gate and up separately, then combine with SiLU gate_out = parallel_linear( inputs=hidden_states, expert_weights=self.expert_gate_proj.permute(0, 2, 1), # [E, hidden, intermediate] k=self.top_k, sorted_expert_idxs=sorted_expert_idxs, sorted_scattered_idxs=sorted_scattered_idxs, expert_offsets=expert_offsets, grouped_out=True, # Keep output grouped by expert (size: num_tokens * k) ) up_out = parallel_linear( inputs=hidden_states, expert_weights=self.expert_up_proj.permute(0, 2, 1), # [E, hidden, intermediate] k=self.top_k, sorted_expert_idxs=sorted_expert_idxs, sorted_scattered_idxs=sorted_scattered_idxs, expert_offsets=expert_offsets, grouped_out=True, # Keep output grouped by expert (size: num_tokens * k) ) # Apply SwiGLU activation activated = F.silu(gate_out) * up_out # Second pass: down projection (intermediate -> hidden) with gates for weighted sum # When grouped_in=True, use k=1 since input is already expanded to num_tokens * top_k output = parallel_linear( inputs=activated, expert_weights=self.expert_down_proj.permute(0, 2, 1), # [E, intermediate, hidden] k=1, # k=1 because input is already grouped/expanded sorted_expert_idxs=sorted_expert_idxs, sorted_scattered_idxs=sorted_scattered_idxs, expert_offsets=expert_offsets, gates=routing_weights, # Apply routing weights during scatter grouped_in=True, # Input is grouped by expert grouped_out=False, # Output ungrouped (back to token order) ) return output def _forward_loop( self, hidden_states: torch.Tensor, selected_experts: torch.Tensor, routing_weights: torch.Tensor, ) -> torch.Tensor: """Fallback forward pass using Python loop.""" num_tokens = hidden_states.shape[0] # Flatten token-expert assignments flat_expert_indices = selected_experts.view(-1) flat_token_indices = torch.arange(num_tokens, device=hidden_states.device).unsqueeze(1).expand(-1, self.top_k).reshape(-1) flat_routing_weights = routing_weights.view(-1) # Sort by expert for contiguous memory access sorted_indices = torch.argsort(flat_expert_indices, stable=True) sorted_expert_indices = flat_expert_indices[sorted_indices] sorted_token_indices = flat_token_indices[sorted_indices] sorted_routing_weights = flat_routing_weights[sorted_indices] sorted_hidden = hidden_states[sorted_token_indices] # Compute expert boundaries expert_counts = torch.bincount(sorted_expert_indices, minlength=self.num_experts) expert_offsets = torch.zeros(self.num_experts + 1, dtype=torch.long, device=hidden_states.device) expert_offsets[1:] = torch.cumsum(expert_counts, dim=0) final_hidden_states = torch.zeros_like(hidden_states) # Pre-fetch to lists to avoid repeated tensor->scalar conversions counts_list = expert_counts.tolist() offsets_list = expert_offsets.tolist() # Process each expert for expert_idx in range(self.num_experts): count = counts_list[expert_idx] if count == 0: continue start_idx = offsets_list[expert_idx] end_idx = offsets_list[expert_idx + 1] # Get tokens for this expert expert_tokens = sorted_hidden[start_idx:end_idx] expert_weights = sorted_routing_weights[start_idx:end_idx] token_indices = sorted_token_indices[start_idx:end_idx] # Get this expert's weight matrices gate_w = self.expert_gate_proj[expert_idx] up_w = self.expert_up_proj[expert_idx] down_w = self.expert_down_proj[expert_idx] # SwiGLU computation: down(silu(gate(x)) * up(x)) gate_out = F.linear(expert_tokens, gate_w) up_out = F.linear(expert_tokens, up_w) activated = F.silu(gate_out) * up_out expert_out = F.linear(activated, down_w) # Weight and accumulate weighted_out = expert_out * expert_weights.unsqueeze(-1) final_hidden_states.index_add_(0, token_indices, weighted_out) return final_hidden_states class _ScatterbrainRMSNormFallback(nn.Module): """Fallback RMSNorm when Liger is not available.""" def __init__(self, hidden_size, eps=1e-6): super().__init__() self.weight = nn.Parameter(torch.ones(hidden_size)) self.variance_epsilon = eps def forward(self, hidden_states): input_dtype = hidden_states.dtype hidden_states = hidden_states.to(torch.float32) variance = hidden_states.pow(2).mean(-1, keepdim=True) hidden_states = hidden_states * torch.rsqrt(variance + self.variance_epsilon) return self.weight * hidden_states.to(input_dtype) def extra_repr(self): return f"{tuple(self.weight.shape)}, eps={self.variance_epsilon}" # Use Liger RMSNorm if available, otherwise fallback ScatterbrainRMSNorm = LigerRMSNorm if LIGER_AVAILABLE else _ScatterbrainRMSNormFallback class ScatterbrainDecoderLayer(GradientCheckpointingLayer): def __init__(self, config: ScatterbrainConfig, layer_idx: int): super().__init__() self.hidden_size = config.hidden_size self.self_attn = ScatterbrainAttention(config, layer_idx) # MoE layer if config.num_experts > 0: self.mlp = ScatterbrainSparseMoeBlock(config) else: self.mlp = ScatterbrainMLP(config, intermediate_size=config.intermediate_size) self.input_layernorm = ScatterbrainRMSNorm(config.hidden_size, eps=config.rms_norm_eps) self.post_attention_layernorm = ScatterbrainRMSNorm(config.hidden_size, eps=config.rms_norm_eps) @deprecate_kwarg("past_key_value", new_name="past_key_values", version="4.58") def forward( self, hidden_states: torch.Tensor, position_embeddings: tuple[torch.Tensor, torch.Tensor], attention_mask: Optional[torch.Tensor] = None, position_ids: Optional[torch.LongTensor] = None, past_key_values: Optional[Cache] = None, cache_position: Optional[torch.LongTensor] = None, cache_slot_idx: Optional[int] = None, **kwargs: Unpack[FlashAttentionKwargs], ) -> torch.FloatTensor: residual = hidden_states hidden_states = self.input_layernorm(hidden_states) # Self Attention hidden_states, _ = self.self_attn( hidden_states=hidden_states, position_embeddings=position_embeddings, attention_mask=attention_mask, position_ids=position_ids, past_key_values=past_key_values, cache_position=cache_position, cache_slot_idx=cache_slot_idx, **kwargs, ) hidden_states = residual + hidden_states # MLP / MoE residual = hidden_states hidden_states = self.post_attention_layernorm(hidden_states) hidden_states = self.mlp(hidden_states) # Unpack MoE output if isinstance(hidden_states, tuple): hidden_states, _ = hidden_states hidden_states = residual + hidden_states return hidden_states class ScatterbrainRotaryEmbedding(nn.Module): inv_freq: torch.Tensor def __init__(self, config: ScatterbrainConfig, device=None): super().__init__() if hasattr(config, "rope_scaling") and isinstance(config.rope_scaling, dict): self.rope_type = config.rope_scaling.get("rope_type", config.rope_scaling.get("type")) else: self.rope_type = "default" self.max_seq_len_cached = config.max_position_embeddings self.original_max_seq_len = config.max_position_embeddings self.config = config self.rope_init_fn = ROPE_INIT_FUNCTIONS[self.rope_type] inv_freq, self.attention_scaling = self.rope_init_fn(self.config, device) self.register_buffer("inv_freq", inv_freq, persistent=False) self.original_inv_freq = self.inv_freq @torch.no_grad() @dynamic_rope_update def forward(self, x, position_ids): inv_freq_expanded = self.inv_freq[None, :, None].float().expand(position_ids.shape[0], -1, 1).to(x.device) position_ids_expanded = position_ids[:, None, :].float() device_type = x.device.type if isinstance(x.device.type, str) and x.device.type != "mps" else "cpu" with torch.autocast(device_type=device_type, enabled=False): freqs = (inv_freq_expanded.float() @ position_ids_expanded.float()).transpose(1, 2) emb = torch.cat((freqs, freqs), dim=-1) cos = emb.cos() * self.attention_scaling sin = emb.sin() * self.attention_scaling return cos.to(dtype=x.dtype), sin.to(dtype=x.dtype) @auto_docstring class ScatterbrainPreTrainedModel(PreTrainedModel): config_class = ScatterbrainConfig base_model_prefix = "model" supports_gradient_checkpointing = True _no_split_modules = ["ScatterbrainDecoderLayer"] _skip_keys_device_placement = ["past_key_values"] _supports_flash_attn = True _supports_sdpa = True _supports_flex_attn = True _can_compile_fullgraph = False # MoE models don't work with torch.compile _supports_attention_backend = True @auto_docstring class ScatterbrainModel(ScatterbrainPreTrainedModel): def __init__(self, config: ScatterbrainConfig): super().__init__(config) self.padding_idx = config.pad_token_id self.vocab_size = config.vocab_size self.embed_tokens = nn.Embedding(config.vocab_size, config.hidden_size, self.padding_idx) # Only one physical layer self.layers = nn.ModuleList( [ScatterbrainDecoderLayer(config, layer_idx) for layer_idx in range(config.num_hidden_layers)] ) self.norm = ScatterbrainRMSNorm(config.hidden_size, eps=config.rms_norm_eps) self.rotary_emb = ScatterbrainRotaryEmbedding(config=config) self.gradient_checkpointing = False # Number of loop iterations and cache slots self._num_loop_iterations = config.num_loop_iterations self._num_cache_slots = config.num_loop_iterations # Initialize weights and apply final processing self.post_init() @auto_docstring def forward( self, input_ids: Optional[torch.LongTensor] = None, attention_mask: Optional[torch.Tensor] = None, position_ids: Optional[torch.LongTensor] = None, past_key_values: Optional[Cache] = None, inputs_embeds: Optional[torch.FloatTensor] = None, use_cache: Optional[bool] = None, cache_position: Optional[torch.LongTensor] = None, **kwargs: Unpack[TransformersKwargs], ) -> MoeModelOutputWithPast: if (input_ids is None) ^ (inputs_embeds is not None): raise ValueError("You must specify exactly one of input_ids or inputs_embeds") if inputs_embeds is None: inputs_embeds = self.embed_tokens(input_ids) # Handle cache creation/extension for looping if use_cache: if past_key_values is None: # Create cache with enough slots for all loop iterations cache_config = copy.copy(self.config) cache_config.num_hidden_layers = self._num_cache_slots past_key_values = DynamicCache(config=cache_config) elif isinstance(past_key_values, DynamicCache) and len(past_key_values.layers) < self._num_cache_slots: # Extend cache if created externally with fewer slots while len(past_key_values.layers) < self._num_cache_slots: past_key_values.layers.append(DynamicLayer()) if cache_position is None: past_seen_tokens = past_key_values.get_seq_length() if past_key_values is not None else 0 cache_position = torch.arange( past_seen_tokens, past_seen_tokens + inputs_embeds.shape[1], device=inputs_embeds.device ) if position_ids is None: position_ids = cache_position.unsqueeze(0) mask_function = create_causal_mask if self.config.sliding_window is None else create_sliding_window_causal_mask causal_mask = mask_function( config=self.config, input_embeds=inputs_embeds, attention_mask=attention_mask, cache_position=cache_position, past_key_values=past_key_values, position_ids=position_ids, ) hidden_states = inputs_embeds position_embeddings = self.rotary_emb(hidden_states, position_ids) # Loop through the single layer multiple times # Each iteration gets its own cache slot decoder_layer = self.layers[0] for loop_idx in range(self._num_loop_iterations): hidden_states = decoder_layer( hidden_states, position_embeddings=position_embeddings, attention_mask=causal_mask, position_ids=position_ids, past_key_values=past_key_values, use_cache=use_cache, cache_position=cache_position, cache_slot_idx=loop_idx, **kwargs, ) hidden_states = self.norm(hidden_states) return MoeModelOutputWithPast( last_hidden_state=hidden_states, past_key_values=past_key_values, ) @auto_docstring class ScatterbrainForCausalLM(ScatterbrainPreTrainedModel, GenerationMixin): _tied_weights_keys = ["lm_head.weight"] def __init__(self, config): super().__init__(config) self.model = ScatterbrainModel(config) self.vocab_size = config.vocab_size self.lm_head = nn.Linear(config.hidden_size, config.vocab_size, bias=False) self.router_aux_loss_coef = config.router_aux_loss_coef self.num_experts = config.num_experts self.num_experts_per_tok = config.num_experts_per_tok # Initialize weights and apply final processing self.post_init() def get_input_embeddings(self): return self.model.embed_tokens def set_input_embeddings(self, value): self.model.embed_tokens = value def get_output_embeddings(self): return self.lm_head def set_output_embeddings(self, new_embeddings): self.lm_head = new_embeddings @can_return_tuple @auto_docstring def forward( self, input_ids: Optional[torch.LongTensor] = None, attention_mask: Optional[torch.Tensor] = None, position_ids: Optional[torch.LongTensor] = None, past_key_values: Optional[Cache] = None, inputs_embeds: Optional[torch.FloatTensor] = None, labels: Optional[torch.LongTensor] = None, use_cache: Optional[bool] = None, output_router_logits: Optional[bool] = None, cache_position: Optional[torch.LongTensor] = None, logits_to_keep: Union[int, torch.Tensor] = 0, **kwargs: Unpack[TransformersKwargs], ) -> MoeCausalLMOutputWithPast: output_router_logits = ( output_router_logits if output_router_logits is not None else self.config.output_router_logits ) outputs: MoeModelOutputWithPast = self.model( input_ids=input_ids, attention_mask=attention_mask, position_ids=position_ids, past_key_values=past_key_values, inputs_embeds=inputs_embeds, use_cache=use_cache, output_router_logits=output_router_logits, cache_position=cache_position, **kwargs, ) hidden_states = outputs.last_hidden_state slice_indices = slice(-logits_to_keep, None) if isinstance(logits_to_keep, int) else logits_to_keep logits = self.lm_head(hidden_states[:, slice_indices, :]) loss = None if labels is not None: loss = self.loss_function(logits, labels, self.vocab_size, **kwargs) return MoeCausalLMOutputWithPast( loss=loss, logits=logits, past_key_values=outputs.past_key_values, hidden_states=outputs.hidden_states, attentions=outputs.attentions, router_logits=outputs.router_logits, ) __all__ = [ "ScatterbrainConfig", "ScatterbrainForCausalLM", "ScatterbrainModel", "ScatterbrainPreTrainedModel", ]