Ouzhang's picture
Add files using upload-large-folder tool
d91766b verified
Raw
History Blame Contribute Delete
3.42 kB
from __future__ import annotations
import torch
import torch.nn as nn
from diffulex.layer.embed_head import ParallelLMHead, VocabParallelEmbedding
from diffulex.layer.layernorm import RMSNorm
from diffulex.model.auto_model import AutoModelForDiffusionLM
from diffulex.model.sdar import SDARAttention, SDARMLP
from diffulex.moe import build_mlp_or_moe
class SDARMoEDecoderLayer(nn.Module):
"""SDAR decoder layer with optional sparse MoE MLP."""
def __init__(self, config, layer_idx: int) -> None:
super().__init__()
self.self_attn = SDARAttention(config)
self.mlp = build_mlp_or_moe(config, layer_idx, lambda: SDARMLP(config))
self.input_layernorm = RMSNorm(config.hidden_size, eps=config.rms_norm_eps)
self.post_attention_layernorm = RMSNorm(config.hidden_size, eps=config.rms_norm_eps)
def forward(
self,
positions: torch.Tensor,
hidden_states: torch.Tensor,
residual: torch.Tensor | None,
mask: torch.Tensor | None = None,
) -> tuple[torch.Tensor, torch.Tensor]:
if residual is None:
residual = hidden_states
hidden_states = self.input_layernorm(hidden_states)
else:
hidden_states, residual = self.input_layernorm(hidden_states, residual)
hidden_states = self.self_attn(positions, hidden_states, mask)
hidden_states, residual = self.post_attention_layernorm(hidden_states, residual)
mlp_output = self.mlp(hidden_states)
if isinstance(mlp_output, tuple):
hidden_states, _router_logits = mlp_output
else:
hidden_states = mlp_output
return hidden_states, residual
class SDARMoEModel(nn.Module):
def __init__(self, config) -> None:
super().__init__()
self.embed_tokens = VocabParallelEmbedding(config.vocab_size, config.hidden_size)
self.layers = nn.ModuleList(
[SDARMoEDecoderLayer(config, layer_idx) for layer_idx in range(config.num_hidden_layers)]
)
self.norm = RMSNorm(config.hidden_size, eps=config.rms_norm_eps)
def forward(
self,
input_ids: torch.Tensor,
positions: torch.Tensor,
mask: torch.Tensor | None = None,
) -> torch.Tensor:
hidden_states = self.embed_tokens(input_ids)
residual = None
for layer in self.layers:
hidden_states, residual = layer(positions, hidden_states, residual, mask)
hidden_states, _ = self.norm(hidden_states, residual)
return hidden_states
@AutoModelForDiffusionLM.register("sdar_moe")
class SDARMoEForDiffusionLM(nn.Module):
packed_modules_mapping = {}
def __init__(self, config) -> None:
super().__init__()
self.model = SDARMoEModel(config)
self.lm_head = ParallelLMHead(config.vocab_size, config.hidden_size)
if getattr(config, "tie_word_embeddings", False):
self.lm_head.weight.data = self.model.embed_tokens.weight.data
def forward(
self,
input_ids: torch.Tensor,
positions: torch.Tensor,
mask: torch.Tensor | None = None,
) -> torch.Tensor:
return self.model(input_ids, positions, mask)
def compute_logits(self, hidden_states: torch.Tensor) -> torch.Tensor:
return self.lm_head(hidden_states)
__all__ = [
"SDARMoEDecoderLayer",
"SDARMoEModel",
"SDARMoEForDiffusionLM",
]