ForgePlex-M2-9M / modeling_forgeplex_m2.py
Jsandero's picture
Add ForgePlex-M2-9M weights (peak Index 8.51 @ 452800, ~9.95M)
93e48a8
Raw History Blame Contribute Delete
15.7 kB
"""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,
)