DMax-Math-MMD / modeling_llada2_moe.py
free001style's picture
Add files using upload-large-folder tool
0b46285 verified
Raw History Blame Contribute Delete
16.9 kB
# coding=utf-8
# Copyright 2025 Antgroup and The HuggingFace Inc. team. All rights reserved.
#
# This code is based on EleutherAI's GPT-NeoX library and the GPT-NeoX
# and OPT implementations in this library. It has been modified from its
# original forms to accommodate minor architectural differences compared
# to GPT-NeoX and OPT used by the Meta AI team that trained the model.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
"""DMax LLaDA2 training layers; see NOTICE.md for upstream revision and adaptations.
Explicit 4D attention masks are passed unchanged to SDPA (True means attend).
"""
import math
import torch
import torch.nn.functional as F
from torch import nn
from torch.utils.checkpoint import checkpoint
from transformers.activations import ACT2FN
from transformers.modeling_outputs import MoeCausalLMOutputWithPast
from transformers.modeling_utils import PreTrainedModel
from .configuration_llada2_moe import LLaDA2MoeConfig
class LLaDA2MoeRMSNorm(nn.Module):
def __init__(self, hidden_size, eps=1e-6):
"""
LLaDA2MoeRMSNorm is equivalent to T5LayerNorm
"""
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 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)
# Copied from transformers.models.llama.modeling_llama.apply_rotary_pos_emb
def apply_rotary_pos_emb(q, k, cos, sin):
"""Rotate the configured head dimensions at the supplied logical positions."""
# RoPE values are [batch, token, rotary_dim]; broadcast over attention heads.
cos = cos.unsqueeze(1)
sin = sin.unsqueeze(1)
# Keep half or full tensor for later concatenation
rotary_dim = cos.shape[-1]
q_rot, q_pass = q[..., :rotary_dim], q[..., rotary_dim:]
k_rot, k_pass = k[..., :rotary_dim], k[..., rotary_dim:]
# Apply rotary embeddings on the first half or full tensor
q_embed = (q_rot * cos) + (rotate_half(q_rot) * sin)
k_embed = (k_rot * cos) + (rotate_half(k_rot) * sin)
# Concatenate back to full shape
q_embed = torch.cat([q_embed, q_pass], dim=-1)
k_embed = torch.cat([k_embed, k_pass], dim=-1)
return q_embed, k_embed
class LLaDA2MoeMLP(nn.Module):
def __init__(self, config: LLaDA2MoeConfig, intermediate_size: int):
super().__init__()
self.config = config
self.hidden_size = config.hidden_size
self.intermediate_size = 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):
return self.down_proj(self.act_fn(self.gate_proj(x)) * self.up_proj(x))
class LLaDA2MoeGate(nn.Module):
def __init__(self, config):
super().__init__()
self.config = config
self.top_k = config.num_experts_per_tok
self.num_experts = config.num_experts
self.n_group = config.n_group
self.topk_group = config.topk_group
# topk selection algorithm
self.gating_dim = config.hidden_size
self.weight = nn.Parameter(torch.empty((self.num_experts, self.gating_dim)))
self.routed_scaling_factor = config.routed_scaling_factor
self.register_buffer("expert_bias", torch.zeros((self.num_experts)))
self.reset_parameters()
def reset_parameters(self) -> None:
import torch.nn.init as init
init.kaiming_uniform_(self.weight, a=math.sqrt(5))
def group_limited_topk(
self,
scores: torch.Tensor,
):
num_tokens, _ = scores.size()
# Organize the experts into groups
group_scores = scores.view(num_tokens, self.n_group, -1).topk(2, dim=-1)[0].sum(dim=-1)
group_idx = torch.topk(group_scores, k=self.topk_group, dim=-1, sorted=False)[1]
group_mask = torch.zeros_like(group_scores)
group_mask.scatter_(1, group_idx, 1)
# Mask the experts based on selection groups
score_mask = (
group_mask.unsqueeze(-1)
.expand(num_tokens, self.n_group, self.num_experts // self.n_group)
.reshape(num_tokens, -1)
)
masked_scores = scores.masked_fill(~score_mask.bool(), float('-inf'))
probs, top_indices = torch.topk(masked_scores, k=self.top_k, dim=-1)
return probs, top_indices
def forward(self, hidden_states):
# compute gating score
hidden_states = hidden_states.view(-1, hidden_states.shape[-1])
logits = F.linear(hidden_states.type(torch.float32), self.weight.type(torch.float32))
scores = torch.sigmoid(logits.float()).type_as(logits)
scores_for_routing = scores + self.expert_bias
_, topk_idx = self.group_limited_topk(scores_for_routing)
scores = torch.gather(scores, dim=1, index=topk_idx).type_as(logits)
topk_weight = scores / (scores.sum(dim=-1, keepdim=True) + 1e-20) if self.top_k > 1 else scores
topk_weight = topk_weight * self.routed_scaling_factor
return topk_idx, topk_weight, logits
class LLaDA2MoeRotaryEmbedding(nn.Module):
def __init__(self, config):
super().__init__()
self.config = config
scaling = config.rope_scaling or {}
if scaling.get("rope_type", scaling.get("type", "default")) != "default":
raise ValueError("yrDMax supports the released LLaDA2 default RoPE only")
self.rebuild(torch.device("cpu"))
def rebuild(self, device):
dim = int(self.config.head_dim * self.config.partial_rotary_factor)
inv_freq = 1.0 / (self.config.rope_theta ** (
torch.arange(0, dim, 2, dtype=torch.float32, device=device) / dim))
self.register_buffer("inv_freq", inv_freq, persistent=False)
@torch.no_grad()
def forward(self, x, position_ids):
inv = self.inv_freq[None, :, None].float().expand(position_ids.shape[0], -1, 1)
with torch.autocast(device_type=x.device.type, enabled=False):
freqs = (inv @ position_ids[:, None, :].float()).transpose(1, 2)
emb = torch.cat((freqs, freqs), dim=-1)
return emb.cos().to(x.dtype), emb.sin().to(x.dtype)
class LLaDA2MoeSparseMoeBlock(nn.Module):
def __init__(self, config):
super().__init__()
self.config = config
self.num_experts_per_tok = config.num_experts_per_tok
self.experts = nn.ModuleList([
LLaDA2MoeMLP(config, config.moe_intermediate_size)
for _ in range(config.num_experts)])
self.gate = LLaDA2MoeGate(config)
self.shared_experts = (LLaDA2MoeMLP(
config, config.moe_intermediate_size * config.num_shared_experts)
if config.num_shared_experts else None)
def forward(self, hidden_states):
# Same expert and combine equations as upstream _forward; autograd also
# works in eval mode, so temporary mode changes cannot freeze experts.
shape = hidden_states.shape
indices, weights, router_logits = self.gate(hidden_states)
flat = hidden_states.reshape(-1, shape[-1])
expanded = flat.repeat_interleave(self.num_experts_per_tok, dim=0)
result = torch.empty_like(expanded)
for i, expert in enumerate(self.experts):
selected = indices.reshape(-1) == i
# Unwrapped/DDP FP32 masters may produce BF16 expert activations
# under autocast. Indexed assignment does not promote dtypes.
result[selected] = expert(expanded[selected]).to(result.dtype)
result = (result.view(*weights.shape, -1) * weights.unsqueeze(-1)).sum(1)
result = result.to(hidden_states.dtype).view(shape)
if self.shared_experts is not None:
result = result + self.shared_experts(hidden_states)
return result
def repeat_kv(hidden_states: torch.Tensor, n_rep: int) -> torch.Tensor:
"""
This is the equivalent of torch.repeat_interleave(x, dim=1, repeats=n_rep). The hidden states go from (batch,
num_key_value_heads, seqlen, head_dim) to (batch, num_attention_heads, seqlen, head_dim)
"""
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)
# Copied from transformers.models.llama.modeling_llama.LlamaAttention with Llama->LLaDA2Moe
class LLaDA2MoeAttention(nn.Module):
"""Multi-headed attention from 'Attention Is All You Need' paper"""
def __init__(self, config: LLaDA2MoeConfig, layer_idx=None):
super().__init__()
self.config = config
self.layer_idx = layer_idx
self.attention_dropout = config.attention_dropout
self.hidden_size = config.hidden_size
self.num_heads = config.num_attention_heads
self.head_dim = config.head_dim or self.hidden_size // self.num_heads
partial_rotary_factor = config.partial_rotary_factor if hasattr(config, "partial_rotary_factor") else 1.0
self.rope_dim = int(self.head_dim * partial_rotary_factor)
self.num_key_value_heads = config.num_key_value_heads
self.num_key_value_groups = self.num_heads // self.num_key_value_heads
self.max_position_embeddings = config.max_position_embeddings
self.rope_theta = config.rope_theta
self.is_causal = False
self.query_key_value = nn.Linear(
self.hidden_size,
(self.num_heads + 2 * self.num_key_value_heads) * self.head_dim,
bias=config.use_qkv_bias,
)
self.query_layernorm = LLaDA2MoeRMSNorm(self.head_dim, eps=config.rms_norm_eps)
self.key_layernorm = LLaDA2MoeRMSNorm(self.head_dim, eps=config.rms_norm_eps)
self.dense = nn.Linear(self.num_heads * self.head_dim, self.hidden_size, bias=config.use_bias)
def forward(self, hidden_states, attention_mask, position_embeddings):
bsz, q_len, _ = hidden_states.shape
qkv = self.query_key_value(hidden_states).view(
bsz, q_len, self.num_heads + 2 * self.num_key_value_heads, self.head_dim)
query, key, value = qkv.split(
[self.num_heads, self.num_key_value_heads, self.num_key_value_heads], dim=-2)
query = self.query_layernorm(query.transpose(1, 2))
key = self.key_layernorm(key.transpose(1, 2))
value = value.transpose(1, 2)
query, key = apply_rotary_pos_emb(query, key, *position_embeddings)
key, value = repeat_kv(key, self.num_key_value_groups), repeat_kv(value, self.num_key_value_groups)
out = F.scaled_dot_product_attention(
query.contiguous(), key.contiguous(), value.contiguous(),
attn_mask=attention_mask,
dropout_p=self.attention_dropout if self.training else 0.0,
is_causal=False)
out = self.dense(out.transpose(1, 2).contiguous().reshape(bsz, q_len, -1))
return out
class LLaDA2MoeDecoderLayer(nn.Module):
def __init__(self, config, layer_idx):
super().__init__()
self.attention = LLaDA2MoeAttention(config, layer_idx)
self.mlp = (LLaDA2MoeSparseMoeBlock(config)
if config.num_experts is not None and layer_idx >= config.first_k_dense_replace
else LLaDA2MoeMLP(config, config.intermediate_size))
self.input_layernorm = LLaDA2MoeRMSNorm(config.hidden_size, config.rms_norm_eps)
self.post_attention_layernorm = LLaDA2MoeRMSNorm(config.hidden_size, config.rms_norm_eps)
def forward(self, hidden_states, attention_mask, position_embeddings):
attention_output = self.attention(self.input_layernorm(hidden_states), attention_mask,
position_embeddings)
hidden_states = hidden_states + attention_output
hidden_states = hidden_states + self.mlp(self.post_attention_layernorm(hidden_states))
return hidden_states
class LLaDA2MoeModel(nn.Module):
def __init__(self, config):
super().__init__()
self.word_embeddings = nn.Embedding(config.vocab_size, config.hidden_size, config.pad_token_id)
self.layers = nn.ModuleList([LLaDA2MoeDecoderLayer(config, i) for i in range(config.num_hidden_layers)])
self.norm = LLaDA2MoeRMSNorm(config.hidden_size, config.rms_norm_eps)
self.rotary_emb = LLaDA2MoeRotaryEmbedding(config)
self.gradient_checkpointing = False
def forward(self, input_ids, attention_mask, position_ids, inputs_embeds=None):
hidden = self.word_embeddings(input_ids) if inputs_embeds is None else inputs_embeds
positions = self.rotary_emb(hidden, position_ids)
for layer in self.layers:
if self.gradient_checkpointing and self.training:
hidden = checkpoint(layer, hidden, attention_mask, positions, use_reentrant=False)
else:
hidden = layer(hidden, attention_mask, positions)
return self.norm(hidden)
class LLaDA2MoeModelLM(PreTrainedModel):
config_class = LLaDA2MoeConfig
base_model_prefix = "model"
_no_split_modules = ["LLaDA2MoeDecoderLayer"]
_supports_sdpa = True
supports_gradient_checkpointing = True
_tied_weights_keys = {"lm_head.weight": "model.word_embeddings.weight"}
def __init__(self, config):
super().__init__(config)
self.model = LLaDA2MoeModel(config)
self.lm_head = nn.Linear(config.hidden_size, config.vocab_size, bias=False)
self.post_init()
def _init_weights(self, module):
if isinstance(module, (nn.Linear, nn.Embedding)):
module.weight.data.normal_(mean=0.0, std=self.config.initializer_range)
if getattr(module, "bias", None) is not None:
module.bias.data.zero_()
if getattr(module, "padding_idx", None) is not None:
module.weight.data[module.padding_idx].zero_()
elif isinstance(module, LLaDA2MoeGate):
module.reset_parameters()
elif isinstance(module, LLaDA2MoeRMSNorm):
module.weight.data.fill_(1.0)
def get_input_embeddings(self):
return self.model.word_embeddings
def get_output_embeddings(self):
return self.lm_head
def gradient_checkpointing_enable(self, gradient_checkpointing_kwargs=None):
self.model.gradient_checkpointing = True
def forward(self, input_ids=None, attention_mask=None, position_ids=None,
inputs_embeds=None, logits_to_keep=None):
"""Apply the explicit OPUT grid and project only the requested loss positions."""
if (input_ids is None) == (inputs_embeds is None):
raise ValueError("Supply exactly one of input_ids and inputs_embeds")
shape = input_ids.shape if input_ids is not None else inputs_embeds.shape[:2]
if attention_mask is None or attention_mask.ndim != 4:
raise ValueError("Supply an explicit 4D block attention mask")
if attention_mask.shape[-2:] != (shape[1], shape[1]):
raise ValueError("Attention mask does not match the input grid")
if position_ids is None:
raise ValueError("Supply original DMax position_ids explicitly")
hidden = self.model(input_ids, attention_mask, position_ids, inputs_embeds)
if logits_to_keep is not None:
# The student needs only candidate logits; the frozen teacher needs
# no vocabulary projection when its decoder hook collects features.
hidden = hidden[:, logits_to_keep if isinstance(logits_to_keep, slice) else slice(*logits_to_keep)]
return MoeCausalLMOutputWithPast(logits=self.lm_head(hidden))