| 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", |
| ] |
|
|