"""ForgePlex-M2 causal LM for Hugging Face Transformers. Preserves training-time Qwen3.5-style attention output gates and GPT-S2-style refresh gates (inject layers). RoPE uses NeoX even/odd interleaving (same as training) — no Llama half-rotate remapping. """ from __future__ import annotations from typing import Optional import torch import torch.nn as nn import torch.nn.functional as F from transformers import PreTrainedModel from transformers.cache_utils import DynamicCache from transformers.generation.utils import GenerationMixin from transformers.modeling_outputs import CausalLMOutputWithPast from .configuration_forgeplex_m2 import ForgePlexM2Config class RMSNorm(nn.Module): def __init__(self, dim: int, eps: float = 1e-6): super().__init__() self.eps = eps self.weight = nn.Parameter(torch.ones(dim)) def forward(self, x: torch.Tensor) -> torch.Tensor: rms = torch.rsqrt(x.float().pow(2).mean(-1, keepdim=True) + self.eps) return (x.float() * rms).type_as(x) * self.weight def precompute_rope_cos_sin( head_dim: int, seq_len: int, theta: float = 5000.0, device=None, ) -> tuple[torch.Tensor, torch.Tensor]: freqs = 1.0 / ( theta ** (torch.arange(0, head_dim, 2, dtype=torch.float32, device=device) / head_dim) ) positions = torch.arange(seq_len, dtype=torch.float32, device=device) freqs = torch.outer(positions, freqs) return freqs.cos(), freqs.sin() def apply_rotary_emb( q: torch.Tensor, k: torch.Tensor, rope_cos: torch.Tensor, rope_sin: torch.Tensor, ) -> tuple[torch.Tensor, torch.Tensor]: cos = rope_cos.unsqueeze(0).unsqueeze(0) sin = rope_sin.unsqueeze(0).unsqueeze(0) q_float = q.float().reshape(*q.shape[:-1], -1, 2) k_float = k.float().reshape(*k.shape[:-1], -1, 2) q_even, q_odd = q_float.unbind(-1) k_even, k_odd = k_float.unbind(-1) q_out = torch.stack( (q_even * cos - q_odd * sin, q_even * sin + q_odd * cos), dim=-1 ).flatten(-2) k_out = torch.stack( (k_even * cos - k_odd * sin, k_even * sin + k_odd * cos), dim=-1 ).flatten(-2) return q_out.type_as(q), k_out.type_as(k) class CausalSelfAttention(nn.Module): def __init__(self, config: ForgePlexM2Config, layer_idx: int): super().__init__() self.layer_idx = layer_idx self.n_head = config.num_attention_heads self.n_kv_heads = config.num_key_value_heads self.head_dim = config.head_dim self.n_rep = self.n_head // self.n_kv_heads self.use_xsa_projection = config.use_xsa_projection self.use_attn_output_gate = config.use_attn_output_gate self.q_proj = nn.Linear( config.hidden_size, self.n_head * self.head_dim, bias=False ) self.k_proj = nn.Linear( config.hidden_size, self.n_kv_heads * self.head_dim, bias=False ) self.v_proj = nn.Linear( config.hidden_size, self.n_kv_heads * self.head_dim, bias=False ) self.o_proj = nn.Linear( self.n_head * self.head_dim, config.hidden_size, bias=False ) if self.use_attn_output_gate: self.attn_gate = nn.Linear( config.hidden_size, self.n_head * self.head_dim, bias=False ) def forward( self, x: torch.Tensor, rope_cos: torch.Tensor, rope_sin: torch.Tensor, past_key_value: Optional[DynamicCache] = None, attention_mask: Optional[torch.Tensor] = None, ) -> torch.Tensor: batch_size, query_length, _ = x.size() q = self.q_proj(x).view( batch_size, query_length, self.n_head, self.head_dim ).transpose(1, 2) k = self.k_proj(x).view( batch_size, query_length, self.n_kv_heads, self.head_dim ).transpose(1, 2) v = self.v_proj(x).view( batch_size, query_length, self.n_kv_heads, self.head_dim ).transpose(1, 2) q, k = apply_rotary_emb(q, k, rope_cos, rope_sin) current_v = v if past_key_value is not None: k, v = past_key_value.update(k, v, self.layer_idx) key_length = k.size(2) # Prefer native GQA when available (training path); fall back to repeat. use_native_gqa = ( past_key_value is None and attention_mask is None and query_length == key_length and query_length > 1 ) if use_native_gqa: y = F.scaled_dot_product_attention( q, k, v, is_causal=True, enable_gqa=True ) else: k_repeated = k.repeat_interleave(self.n_rep, dim=1) v_repeated = v.repeat_interleave(self.n_rep, dim=1) past_length = key_length - query_length is_causal = query_length > 1 and past_length == 0 attn_mask = None if query_length > 1 and (past_length > 0 or attention_mask is not None): causal = torch.ones( query_length, key_length, dtype=torch.bool, device=x.device ).tril(diagonal=past_length) attn_mask = causal[None, None, :, :] if attention_mask is not None: key_padding = attention_mask[:, None, None, :key_length].to(torch.bool) attn_mask = key_padding if attn_mask is None else (key_padding & attn_mask) is_causal = False y = F.scaled_dot_product_attention( q, k_repeated, v_repeated, attn_mask=attn_mask, is_causal=is_causal ) if self.use_xsa_projection: y = y.view( batch_size, self.n_kv_heads, self.n_rep, query_length, self.head_dim, ) v_grouped = current_v.unsqueeze(2) denominator = v_grouped.pow(2).sum(dim=-1, keepdim=True).clamp_min(1e-6) y = y - ((y * v_grouped).sum(dim=-1, keepdim=True) / denominator) * v_grouped y = y.view(batch_size, self.n_head, query_length, self.head_dim) y = y.transpose(1, 2).contiguous().view( batch_size, query_length, self.n_head * self.head_dim ) if self.use_attn_output_gate: y = y * torch.sigmoid(self.attn_gate(x)) return self.o_proj(y) class SwiGLUMLP(nn.Module): def __init__(self, config: ForgePlexM2Config): super().__init__() hidden_dim = config.intermediate_size self.w_gate = nn.Linear(config.hidden_size, hidden_dim, bias=False) self.w_up = nn.Linear(config.hidden_size, hidden_dim, bias=False) self.w_down = nn.Linear(hidden_dim, config.hidden_size, bias=False) def forward(self, x: torch.Tensor) -> torch.Tensor: return self.w_down(F.silu(self.w_gate(x)) * self.w_up(x)) class RefreshGate(nn.Module): """Re-inject original token embeddings into the residual stream.""" def __init__(self, d_model: int, kernel: int = 9, eps: float = 1e-6): super().__init__() if kernel < 1: raise ValueError("refresh_kernel must be positive") self.kernel = kernel self.na = RMSNorm(d_model, eps=eps) self.ne = RMSNorm(d_model, eps=eps) self.gate_proj = nn.Linear(d_model, d_model, bias=False) self.gate_conv = nn.Conv1d( d_model, d_model, kernel, groups=d_model, bias=False, padding=kernel - 1, ) self.value_proj = nn.Linear(d_model, d_model, bias=False) self.out_proj = nn.Linear(d_model, d_model, bias=False) self.nz = RMSNorm(d_model, eps=eps) self.alpha = nn.Parameter(torch.tensor(0.0)) def forward( self, h: torch.Tensor, attn_out: torch.Tensor, e0: torch.Tensor, conv_state: dict | None = None, layer_idx: int | None = None, ) -> torch.Tensor: a = self.na(attn_out.detach()) e = self.ne(e0) batch_size, seq_len, channels = a.shape if conv_state is not None: prev = conv_state.get(layer_idx) if prev is None or prev.size(0) != batch_size: prev = a.new_zeros(batch_size, self.kernel - 1, channels) a_ext = torch.cat([prev, a], dim=1) conv_state[layer_idx] = a_ext[:, -(self.kernel - 1) :, :].detach() conv = F.conv1d( a_ext.transpose(1, 2), self.gate_conv.weight, bias=None, padding=0, groups=channels, ).transpose(1, 2) else: conv = self.gate_conv(a.transpose(1, 2)) conv = conv[:, :, :seq_len].transpose(1, 2) gate = self.gate_proj(a) + conv value = self.value_proj(e) z = self.nz(self.out_proj(F.silu(gate) * value)) return h + self.alpha * z class Block(nn.Module): def __init__(self, config: ForgePlexM2Config, layer_idx: int): super().__init__() self.layer_idx = layer_idx self.ln_1 = RMSNorm(config.hidden_size, eps=config.rms_norm_eps) self.attn = CausalSelfAttention(config, layer_idx) inject = config.use_refresh_gate and layer_idx in config.inject_layers self.refresh = ( RefreshGate( config.hidden_size, kernel=config.refresh_kernel, eps=config.rms_norm_eps, ) if inject else None ) self.ln_2 = RMSNorm(config.hidden_size, eps=config.rms_norm_eps) self.mlp = SwiGLUMLP(config) def forward( self, x: torch.Tensor, e0: torch.Tensor, rope_cos: torch.Tensor, rope_sin: torch.Tensor, past_key_value: Optional[DynamicCache] = None, attention_mask: Optional[torch.Tensor] = None, conv_state: dict | None = None, ) -> torch.Tensor: attn_out = self.attn( self.ln_1(x), rope_cos, rope_sin, past_key_value, attention_mask ) x = x + attn_out if self.refresh is not None: x = self.refresh( x, attn_out, e0, conv_state=conv_state, layer_idx=self.layer_idx ) return x + self.mlp(self.ln_2(x)) class ForgePlexM2PreTrainedModel(PreTrainedModel): config_class = ForgePlexM2Config base_model_prefix = "transformer" supports_gradient_checkpointing = False _supports_cache_class = True def _init_weights(self, module: nn.Module) -> None: std = 0.02 if isinstance(module, nn.Linear): nn.init.normal_(module.weight, mean=0.0, std=std) elif isinstance(module, nn.Conv1d): nn.init.normal_(module.weight, mean=0.0, std=std) elif isinstance(module, nn.Embedding): nn.init.normal_(module.weight, mean=0.0, std=0.02) class ForgePlexM2ForCausalLM(ForgePlexM2PreTrainedModel, GenerationMixin): _tied_weights_keys = {"lm_head.weight": "transformer.wte.weight"} def __init__(self, config: ForgePlexM2Config): super().__init__(config) self.transformer = nn.ModuleDict( { "wte": nn.Embedding(config.vocab_size, config.hidden_size), "h": nn.ModuleList( [Block(config, i) for i in range(config.num_hidden_layers)] ), "ln_f": RMSNorm(config.hidden_size, eps=config.rms_norm_eps), } ) self.lm_head = nn.Linear(config.hidden_size, config.vocab_size, bias=False) if config.tie_word_embeddings: self.lm_head.weight = self.transformer["wte"].weight self._rope_cache = None self.post_init() def get_input_embeddings(self): return self.transformer["wte"] def set_input_embeddings(self, value): self.transformer["wte"] = value if self.config.tie_word_embeddings: self.lm_head.weight = self.transformer["wte"].weight def get_output_embeddings(self): return self.lm_head def set_output_embeddings(self, value): self.lm_head = value def prepare_inputs_for_generation( self, input_ids, past_key_values=None, attention_mask=None, **kwargs ): if past_key_values is not None and past_key_values.get_seq_length() > 0: input_ids = input_ids[:, -1:] return { "input_ids": input_ids, "attention_mask": attention_mask, "past_key_values": past_key_values, "use_cache": kwargs.get("use_cache", True), } def _get_rope(self, seq_len: int, device): cache = self._rope_cache if cache is None or cache[0].device != device or cache[0].size(0) < seq_len: cache = precompute_rope_cos_sin( self.config.head_dim, seq_len, self.config.rope_theta, device=device, ) self._rope_cache = cache return cache[0][:seq_len], cache[1][:seq_len] def forward( self, input_ids: Optional[torch.LongTensor] = None, attention_mask: Optional[torch.Tensor] = None, labels: Optional[torch.LongTensor] = None, past_key_values: Optional[DynamicCache] = None, use_cache: Optional[bool] = None, **kwargs, ): if input_ids is None: raise ValueError("input_ids is required") _, query_length = input_ids.size() if use_cache and past_key_values is None: past_key_values = DynamicCache() past_length = ( past_key_values.get_seq_length() if past_key_values is not None else 0 ) total_length = past_length + query_length if total_length > self.config.max_position_embeddings: raise ValueError( f"Sequence length {total_length} exceeds " f"max_position_embeddings={self.config.max_position_embeddings}" ) x = self.transformer["wte"](input_ids) e0 = x rope_cos, rope_sin = self._get_rope(total_length, input_ids.device) rope_cos = rope_cos[past_length:] rope_sin = rope_sin[past_length:] # Refresh conv state lives on the module so generate can carry it # without ModelOutput plumbing. Reset when starting a new sequence. conv_state = None if use_cache: if past_length == 0: self._refresh_conv_state = {} conv_state = getattr(self, "_refresh_conv_state", None) if conv_state is None: self._refresh_conv_state = {} conv_state = self._refresh_conv_state cache = past_key_values if use_cache else None for block in self.transformer["h"]: x = block( x, e0, rope_cos, rope_sin, past_key_value=cache, attention_mask=attention_mask, conv_state=conv_state, ) logits = self.lm_head(self.transformer["ln_f"](x)) loss = None if labels is not None: shift_logits = logits[..., :-1, :].contiguous() shift_labels = labels[..., 1:].contiguous() loss = F.cross_entropy( shift_logits.view(-1, shift_logits.size(-1)), shift_labels.view(-1), ) return CausalLMOutputWithPast( loss=loss, logits=logits, past_key_values=past_key_values if use_cache else None, )