Text Generation
Transformers
Safetensors
English
forgeplex_m2
language-model
forgeplex
forgeworks
rope
swiglu
gqa
attn-output-gate
refresh-gate
custom_code
Instructions to use ForgeWorks/ForgePlex-M2-9M with libraries, inference providers, notebooks, and local apps. Follow these links to get started.
- Libraries
- Transformers
How to use ForgeWorks/ForgePlex-M2-9M with Transformers:
# Use a pipeline as a high-level helper from transformers import pipeline pipe = pipeline("text-generation", model="ForgeWorks/ForgePlex-M2-9M", trust_remote_code=True)# Load model directly from transformers import AutoModelForCausalLM model = AutoModelForCausalLM.from_pretrained("ForgeWorks/ForgePlex-M2-9M", trust_remote_code=True, device_map="auto") - Notebooks
- Google Colab
- Kaggle
- Local Apps Settings
- vLLM
How to use ForgeWorks/ForgePlex-M2-9M with vLLM:
Install from pip and serve model
# Install vLLM from pip: pip install vllm # Start the vLLM server: vllm serve "ForgeWorks/ForgePlex-M2-9M" # Call the server using curl (OpenAI-compatible API): curl -X POST "http://localhost:8000/v1/completions" \ -H "Content-Type: application/json" \ --data '{ "model": "ForgeWorks/ForgePlex-M2-9M", "prompt": "Once upon a time,", "max_tokens": 512, "temperature": 0.5 }'Use Docker
docker model run hf.co/ForgeWorks/ForgePlex-M2-9M
- SGLang
How to use ForgeWorks/ForgePlex-M2-9M with SGLang:
Install from pip and serve model
# Install SGLang from pip: pip install sglang # Start the SGLang server: python3 -m sglang.launch_server \ --model-path "ForgeWorks/ForgePlex-M2-9M" \ --host 0.0.0.0 \ --port 30000 # Call the server using curl (OpenAI-compatible API): curl -X POST "http://localhost:30000/v1/completions" \ -H "Content-Type: application/json" \ --data '{ "model": "ForgeWorks/ForgePlex-M2-9M", "prompt": "Once upon a time,", "max_tokens": 512, "temperature": 0.5 }'Use Docker images
docker run --gpus all \ --shm-size 32g \ -p 30000:30000 \ -v ~/.cache/huggingface:/root/.cache/huggingface \ --env "HF_TOKEN=<secret>" \ --ipc=host \ lmsysorg/sglang:latest \ python3 -m sglang.launch_server \ --model-path "ForgeWorks/ForgePlex-M2-9M" \ --host 0.0.0.0 \ --port 30000 # Call the server using curl (OpenAI-compatible API): curl -X POST "http://localhost:30000/v1/completions" \ -H "Content-Type: application/json" \ --data '{ "model": "ForgeWorks/ForgePlex-M2-9M", "prompt": "Once upon a time,", "max_tokens": 512, "temperature": 0.5 }' - Docker Model Runner
How to use ForgeWorks/ForgePlex-M2-9M with Docker Model Runner:
docker model run hf.co/ForgeWorks/ForgePlex-M2-9M
Download modeling_forgeplex_m2.py from ForgeWorks/ForgePlex-M2-9M: direct link, hf CLI and curl.
- Browser
- Download file 15.7 kB
-
https://huggingface.co/ForgeWorks/ForgePlex-M2-9M/resolve/main/modeling_forgeplex_m2.py
- Command line
-
hf download hf://ForgeWorks/ForgePlex-M2-9M/modeling_forgeplex_m2.py
-
curl -L -o modeling_forgeplex_m2.py https://huggingface.co/ForgeWorks/ForgePlex-M2-9M/resolve/main/modeling_forgeplex_m2.py
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, | |
| ) | |