File size: 3,422 Bytes
d91766b | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 | 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",
]
|