diff --git a/python/sglang/multimodal_gen/configs/models/encoders/flux_2.py b/python/sglang/multimodal_gen/configs/models/encoders/flux_2.py index 40ce6e3f1..6c1c47c50 100644 --- a/python/sglang/multimodal_gen/configs/models/encoders/flux_2.py +++ b/python/sglang/multimodal_gen/configs/models/encoders/flux_2.py @@ -7,7 +7,9 @@ from sglang.multimodal_gen.configs.models.encoders.base import ( TextEncoderArchConfig, TextEncoderConfig, ) -from sglang.multimodal_gen.configs.models.fsdp import is_layer +from sglang.multimodal_gen.configs.models.encoders.mistral3 import ( + Mistral3EncoderArchConfig, +) FLUX_2_SYSTEM_MESSAGE = ( "You are an AI that reasons about image descriptions. You give structured responses focusing on object relationships, object\n" @@ -30,15 +32,15 @@ def build_flux2_text_messages(prompts: list[str]) -> list[list[dict]]: @dataclass -class Flux2MistralTextArchConfig(TextEncoderArchConfig): - stacked_params_mapping: list[tuple[str, str, str]] = field( - default_factory=lambda: [ - ("qkv_proj", "q_proj", "q"), - ("qkv_proj", "k_proj", "k"), - ("qkv_proj", "v_proj", "v"), - ] - ) - _fsdp_shard_conditions: list = field(default_factory=lambda: [is_layer]) +class Flux2MistralTextArchConfig(Mistral3EncoderArchConfig): + """FLUX.2 text-encoder arch config. + + Inherits Mistral3 defaults (hidden_size, num_attention_heads, head_dim, + rms_norm_eps, rope_parameters, ...) so the TP-parallel runtime has every + field it needs even when the checkpoint config doesn't override them. + Only the tokenizer behavior (max_length=512, padding=max_length) is + flux2-specific. + """ def __post_init__(self) -> None: self.tokenizer_kwargs = { diff --git a/python/sglang/multimodal_gen/configs/models/encoders/mistral3.py b/python/sglang/multimodal_gen/configs/models/encoders/mistral3.py index 7c8f1c969..6cc728f19 100644 --- a/python/sglang/multimodal_gen/configs/models/encoders/mistral3.py +++ b/python/sglang/multimodal_gen/configs/models/encoders/mistral3.py @@ -2,16 +2,12 @@ """Mistral3 text encoder configuration for SGLang diffusion models.""" from dataclasses import dataclass, field +from typing import Any from sglang.multimodal_gen.configs.models.encoders.base import ( TextEncoderArchConfig, TextEncoderConfig, ) -from sglang.multimodal_gen.configs.models.fsdp import ( - is_embed_tokens, - is_final_norm, - is_layer, -) @dataclass @@ -38,6 +34,12 @@ class Mistral3EncoderArchConfig(TextEncoderArchConfig): head_dim: int = 128 hidden_state_skip_layer: int = 2 # Use second-to-last hidden state text_len: int = 0 + # Mistral 3.x uses rope_theta=1e9 (yarn-free). + rope_parameters: dict[str, Any] = field( + default_factory=lambda: {"rope_theta": 1_000_000_000.0} + ) + attention_bias: bool = False + mlp_bias: bool = False stacked_params_mapping: list[tuple[str, str, str]] = field( default_factory=lambda: [ @@ -49,9 +51,8 @@ class Mistral3EncoderArchConfig(TextEncoderArchConfig): ] ) - _fsdp_shard_conditions: list = field( - default_factory=lambda: [is_layer, is_embed_tokens, is_final_norm] - ) + # TP-parallel runtime shards weights along TP dim; no FSDP needed. + _fsdp_shard_conditions: list = field(default_factory=list) def __post_init__(self): # Let the parent populate tokenizer_kwargs["max_length"] = self.text_len diff --git a/python/sglang/multimodal_gen/runtime/layers/custom_op.py b/python/sglang/multimodal_gen/runtime/layers/custom_op.py index 30b94beef..eb3aca58b 100644 --- a/python/sglang/multimodal_gen/runtime/layers/custom_op.py +++ b/python/sglang/multimodal_gen/runtime/layers/custom_op.py @@ -69,6 +69,11 @@ class CustomOp(nn.Module): # PyTorch-native implementation. return self.forward_native(*args, **kwargs) + def forward_xpu(self, *args, **kwargs) -> Any: + # By default, we assume that XPU ops are compatible with the + # PyTorch-native implementation. + return self.forward_native(*args, **kwargs) + def dispatch_forward(self) -> Callable: if _is_cuda: return self.forward_cuda diff --git a/python/sglang/multimodal_gen/runtime/models/encoders/mistral_3.py b/python/sglang/multimodal_gen/runtime/models/encoders/mistral_3.py index 72a7d9587..41e4bdfa5 100644 --- a/python/sglang/multimodal_gen/runtime/models/encoders/mistral_3.py +++ b/python/sglang/multimodal_gen/runtime/models/encoders/mistral_3.py @@ -13,455 +13,416 @@ # 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. -import inspect -from contextlib import nullcontext -from typing import Iterable, Optional, Union + +from collections.abc import Iterable +from typing import Any import torch from torch import nn -from torch.nn.attention import SDPBackend, sdpa_kernel -from transformers import Cache, DynamicCache, LlavaConfig, Mistral3Config, MistralConfig -from transformers.masking_utils import ( - create_causal_mask, - create_sliding_window_causal_mask, -) -from transformers.modeling_outputs import BaseModelOutputWithPast -from transformers.modeling_utils import ALL_ATTENTION_FUNCTIONS -from transformers.models.mistral3.modeling_mistral3 import ( - Mistral3CausalLMOutputWithPast, - Mistral3ModelOutputWithPast, -) -from transformers.models.mistral.modeling_mistral import ( - MistralMLP, - MistralPreTrainedModel, - MistralRMSNorm, - MistralRotaryEmbedding, - apply_rotary_pos_emb, - eager_attention_forward, -) -from sglang.multimodal_gen.runtime.loader.weight_utils import default_weight_loader +from sglang.multimodal_gen.configs.models.encoders import BaseEncoderOutput +from sglang.multimodal_gen.configs.models.encoders.mistral3 import ( + Mistral3EncoderConfig, +) +from sglang.multimodal_gen.runtime.distributed import get_tp_world_size +from sglang.multimodal_gen.runtime.layers.activation import SiluAndMul +from sglang.multimodal_gen.runtime.layers.attention import LocalAttention +from sglang.multimodal_gen.runtime.layers.layernorm import RMSNorm +from sglang.multimodal_gen.runtime.layers.linear import ( + MergedColumnParallelLinear, + QKVParallelLinear, + RowParallelLinear, +) +from sglang.multimodal_gen.runtime.layers.quantization import QuantizationConfig +from sglang.multimodal_gen.runtime.layers.rotary_embedding import get_rope +from sglang.multimodal_gen.runtime.layers.vocab_parallel_embedding import ( + VocabParallelEmbedding, +) +from sglang.multimodal_gen.runtime.loader.weight_utils import ( + default_weight_loader, + maybe_remap_kv_scale_name, +) from sglang.multimodal_gen.runtime.managers.memory_managers.layerwise_offload import ( LayerwiseOffloadableModuleMixin, ) -from sglang.multimodal_gen.runtime.platforms import current_platform +from sglang.multimodal_gen.runtime.models.encoders.base import TextEncoder from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger logger = init_logger(__name__) -_CREATE_CAUSAL_MASK_ARG = ( - "inputs_embeds" - if "inputs_embeds" in inspect.signature(create_causal_mask).parameters - else "input_embeds" -) -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) +class MistralMLP(nn.Module): + def __init__( + self, + hidden_size: int, + intermediate_size: int, + hidden_act: str, + quant_config: QuantizationConfig | None = None, + bias: bool = False, + prefix: str = "", + ) -> None: + super().__init__() + self.gate_up_proj = MergedColumnParallelLinear( + input_size=hidden_size, + output_sizes=[intermediate_size] * 2, + bias=bias, + quant_config=quant_config, + prefix=f"{prefix}.gate_up_proj", + ) + self.down_proj = RowParallelLinear( + input_size=intermediate_size, + output_size=hidden_size, + bias=bias, + quant_config=quant_config, + prefix=f"{prefix}.down_proj", + ) + if hidden_act != "silu": + raise ValueError( + f"Unsupported activation: {hidden_act}. " + "Only silu is supported for Mistral." + ) + self.act_fn = SiluAndMul() + + def forward(self, x: torch.Tensor) -> torch.Tensor: + x, _ = self.gate_up_proj(x) + x = self.act_fn(x) + x, _ = self.down_proj(x) + return x class MistralAttention(nn.Module): - """Multi-headed attention from 'Attention Is All You Need' paper""" - - def __init__(self, config: MistralConfig, layer_idx: int): + def __init__( + self, + config, + hidden_size: int, + num_heads: int, + num_kv_heads: int, + rope_theta: float = 10000.0, + rope_scaling: dict[str, Any] | None = None, + max_position_embeddings: int = 8192, + quant_config: QuantizationConfig | None = None, + bias: bool = False, + bias_o_proj: bool = False, + prefix: str = "", + ) -> None: super().__init__() - self.config = config - self.layer_idx = layer_idx - self.head_dim = ( - getattr(config, "head_dim", None) - or config.hidden_size // config.num_attention_heads - ) - self.num_key_value_groups = ( - config.num_attention_heads // config.num_key_value_heads + self.hidden_size = hidden_size + tp_size = get_tp_world_size() + self.total_num_heads = num_heads + assert self.total_num_heads % tp_size == 0 + self.num_heads = self.total_num_heads // tp_size + self.total_num_kv_heads = num_kv_heads + if self.total_num_kv_heads >= tp_size: + assert self.total_num_kv_heads % tp_size == 0 + else: + assert tp_size % self.total_num_kv_heads == 0 + self.num_kv_heads = max(1, self.total_num_kv_heads // tp_size) + # Mistral exposes an explicit head_dim (introduced by Mistral-Nemo). + self.head_dim = getattr( + config, "head_dim", self.hidden_size // self.total_num_heads ) + self.q_size = self.num_heads * self.head_dim + self.kv_size = self.num_kv_heads * self.head_dim self.scaling = self.head_dim**-0.5 - self.attention_dropout = config.attention_dropout - self.q_proj = nn.Linear( - config.hidden_size, config.num_attention_heads * self.head_dim, bias=False + self.rope_theta = rope_theta + self.max_position_embeddings = max_position_embeddings + + self.qkv_proj = QKVParallelLinear( + hidden_size=hidden_size, + head_size=self.head_dim, + total_num_heads=self.total_num_heads, + total_num_kv_heads=self.total_num_kv_heads, + bias=bias, + quant_config=quant_config, + prefix=f"{prefix}.qkv_proj", ) - self.k_proj = nn.Linear( - config.hidden_size, config.num_key_value_heads * self.head_dim, bias=False + self.o_proj = RowParallelLinear( + input_size=self.total_num_heads * self.head_dim, + output_size=hidden_size, + bias=bias_o_proj, + quant_config=quant_config, + prefix=f"{prefix}.o_proj", ) - self.v_proj = nn.Linear( - config.hidden_size, config.num_key_value_heads * self.head_dim, bias=False + + self.rotary_emb = get_rope( + self.head_dim, + rotary_dim=self.head_dim, + max_position=max_position_embeddings, + base=int(rope_theta), + rope_scaling=rope_scaling, + is_neox_style=True, ) - self.o_proj = nn.Linear( - config.num_attention_heads * self.head_dim, config.hidden_size, bias=False + + self.attn = LocalAttention( + self.num_heads, + self.head_dim, + self.num_kv_heads, + softmax_scale=self.scaling, + causal=True, + supported_attention_backends=config._supported_attention_backends, ) - self.is_causal = True - self.num_heads = config.num_attention_heads - self.num_key_value_heads = config.num_key_value_heads def forward( self, + positions: torch.Tensor, hidden_states: torch.Tensor, - position_embeddings: tuple[torch.Tensor, torch.Tensor], - attention_mask: Optional[torch.Tensor], - past_key_values: Optional[Cache] = None, - cache_position: Optional[torch.LongTensor] = None, - **kwargs, - ) -> tuple[torch.Tensor, Optional[torch.Tensor]]: - input_shape = hidden_states.shape[:-1] - hidden_shape = (*input_shape, -1, self.head_dim) + ) -> torch.Tensor: + qkv, _ = self.qkv_proj(hidden_states) + q, k, v = qkv.split([self.q_size, self.kv_size, self.kv_size], dim=-1) + q, k = self.rotary_emb(positions, q, k) - query_states = self.q_proj(hidden_states).view(hidden_shape).transpose(1, 2) - key_states = self.k_proj(hidden_states).view(hidden_shape).transpose(1, 2) - value_states = self.v_proj(hidden_states).view(hidden_shape).transpose(1, 2) - - cos, sin = position_embeddings - query_states, key_states = apply_rotary_pos_emb( - query_states, key_states, cos, sin + batch_size = q.shape[0] + seq_len = q.shape[1] + q = q.reshape(batch_size, seq_len, self.num_heads, self.head_dim) + k = k.reshape(batch_size, seq_len, self.num_kv_heads, self.head_dim) + v = v.reshape(batch_size, seq_len, self.num_kv_heads, self.head_dim) + attn_output = self.attn(q, k, v) + attn_output = attn_output.reshape( + batch_size, seq_len, self.num_heads * self.head_dim ) - - if past_key_values is not None: - # sin and cos are specific to RoPE models; cache_position needed for the static cache - cache_kwargs = {"sin": sin, "cos": cos, "cache_position": cache_position} - key_states, value_states = past_key_values.update( - key_states, value_states, self.layer_idx, cache_kwargs - ) - - attn_implementation = getattr(self.config, "_attn_implementation", None) - attention_interface = eager_attention_forward - if attn_implementation and attn_implementation != "eager": - if hasattr(ALL_ATTENTION_FUNCTIONS, "get_interface"): - attention_interface = ALL_ATTENTION_FUNCTIONS.get_interface( - attn_implementation, eager_attention_forward - ) - else: - attention_interface = ALL_ATTENTION_FUNCTIONS[attn_implementation] - attn_output, attn_weights = attention_interface( - self, - query_states, - key_states, - value_states, - attention_mask, - dropout=0.0, - scaling=self.scaling, - sliding_window=getattr( - self.config, "sliding_window", None - ), # main diff with Llama - **kwargs, - ) - - attn_output = attn_output.reshape(*input_shape, -1).contiguous() - attn_output = self.o_proj(attn_output) - return attn_output, attn_weights + output, _ = self.o_proj(attn_output) + return output class MistralDecoderLayer(nn.Module): - def __init__(self, config: MistralConfig, layer_idx: int): + def __init__( + self, + config, + quant_config: QuantizationConfig | None = None, + prefix: str = "", + ) -> None: super().__init__() self.hidden_size = config.hidden_size - self.self_attn = MistralAttention(config=config, layer_idx=layer_idx) - self.mlp = MistralMLP(config) - self.input_layernorm = MistralRMSNorm( - config.hidden_size, eps=config.rms_norm_eps + rope_parameters = getattr(config, "rope_parameters", None) or {} + rope_theta = rope_parameters.get("rope_theta", 10000.0) + rope_scaling = rope_parameters or None + max_position_embeddings = getattr(config, "max_position_embeddings", 8192) + attention_bias = getattr(config, "attention_bias", False) + bias_o_proj = attention_bias + + self.self_attn = MistralAttention( + config=config, + hidden_size=self.hidden_size, + num_heads=config.num_attention_heads, + num_kv_heads=getattr( + config, "num_key_value_heads", config.num_attention_heads + ), + rope_theta=rope_theta, + rope_scaling=rope_scaling, + max_position_embeddings=max_position_embeddings, + quant_config=quant_config, + bias=attention_bias, + bias_o_proj=bias_o_proj, + prefix=f"{prefix}.self_attn", ) - self.post_attention_layernorm = MistralRMSNorm( + self.mlp = MistralMLP( + hidden_size=self.hidden_size, + intermediate_size=config.intermediate_size, + hidden_act=config.hidden_act, + quant_config=quant_config, + bias=getattr(config, "mlp_bias", False), + prefix=f"{prefix}.mlp", + ) + 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, - attention_mask: Optional[torch.Tensor] = None, - position_ids: Optional[torch.LongTensor] = None, - past_key_values: Optional[Cache] = None, - use_cache: Optional[bool] = False, - cache_position: Optional[torch.LongTensor] = None, - position_embeddings: Optional[ - tuple[torch.Tensor, torch.Tensor] - ] = None, # necessary, but kept here for BC - **kwargs, - ) -> torch.Tensor: - residual = hidden_states - hidden_states = self.input_layernorm(hidden_states) - # Self Attention - hidden_states, _ = self.self_attn( - hidden_states=hidden_states, - attention_mask=attention_mask, - position_ids=position_ids, - past_key_values=past_key_values, - use_cache=use_cache, - cache_position=cache_position, - position_embeddings=position_embeddings, - **kwargs, - ) - hidden_states = residual + hidden_states - - # Fully Connected - residual = hidden_states - hidden_states = self.post_attention_layernorm(hidden_states) + residual: torch.Tensor | 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=positions, hidden_states=hidden_states) + hidden_states, residual = self.post_attention_layernorm(hidden_states, residual) hidden_states = self.mlp(hidden_states) - hidden_states = residual + hidden_states - return hidden_states + return hidden_states, residual -class MistralModel(MistralPreTrainedModel): - def __init__(self, config: MistralConfig): - super().__init__(config) - self.padding_idx = config.pad_token_id +class MistralModel(nn.Module): + """TP-parallel Mistral decoder stack used as a text encoder.""" + + def __init__(self, config: Mistral3EncoderConfig, prefix: str = "") -> None: + super().__init__() + self.config = config self.vocab_size = config.vocab_size - self.embed_tokens = nn.Embedding( - config.vocab_size, config.hidden_size, self.padding_idx + self.embed_tokens = VocabParallelEmbedding( + config.vocab_size, + config.hidden_size, + org_num_embeddings=config.vocab_size, + quant_config=getattr(config, "quant_config", None), ) + self.layers = nn.ModuleList( [ - MistralDecoderLayer(config, layer_idx) - for layer_idx in range(config.num_hidden_layers) + MistralDecoderLayer( + config=config, + quant_config=getattr(config, "quant_config", None), + prefix=f"{prefix}.layers.{i}", + ) + for i in range(config.num_hidden_layers) ] ) - self.norm = MistralRMSNorm(config.hidden_size, eps=config.rms_norm_eps) - self.rotary_emb = MistralRotaryEmbedding(config=config) - self.gradient_checkpointing = False - self.post_init() + self.norm = RMSNorm(config.hidden_size, eps=config.rms_norm_eps) + + def get_input_embeddings(self, input_ids: torch.Tensor) -> torch.Tensor: + return self.embed_tokens(input_ids) def forward( self, - input_ids: Optional[torch.LongTensor] = None, - attention_mask: Optional[torch.Tensor] = None, - position_ids: Optional[torch.LongTensor] = None, - past_key_values: Optional[Cache] = None, - inputs_embeds: Optional[torch.FloatTensor] = None, - use_cache: Optional[bool] = None, - cache_position: Optional[torch.LongTensor] = None, - output_hidden_states: Optional[bool] = None, + input_ids: torch.Tensor | None, + position_ids: torch.Tensor | None = None, + attention_mask: torch.Tensor | None = None, + inputs_embeds: torch.Tensor | None = None, + output_hidden_states: bool | None = True, **kwargs, - ) -> BaseModelOutputWithPast: - if (input_ids is None) ^ (inputs_embeds is not None): - raise ValueError( - "You must specify exactly one of input_ids or inputs_embeds" - ) - - if inputs_embeds is None: - inputs_embeds = self.embed_tokens(input_ids) - - if use_cache and past_key_values is None: - past_key_values = DynamicCache(config=self.config) - - if cache_position is None: - past_seen_tokens = ( - past_key_values.get_seq_length() if past_key_values is not None else 0 - ) - cache_position = torch.arange( - past_seen_tokens, - past_seen_tokens + inputs_embeds.shape[1], - device=inputs_embeds.device, - ) + ) -> BaseEncoderOutput: + if inputs_embeds is not None: + hidden_states = inputs_embeds + else: + hidden_states = self.get_input_embeddings(input_ids) if position_ids is None: - position_ids = cache_position.unsqueeze(0) - mask_function = ( - create_causal_mask - if getattr(self.config, "sliding_window", None) is None - else create_sliding_window_causal_mask + position_ids = torch.arange( + 0, hidden_states.shape[1], device=hidden_states.device + ).unsqueeze(0) + + residual: torch.Tensor | None = None + collect_hidden = bool(output_hidden_states) + all_hidden_states: tuple[torch.Tensor, ...] | None = ( + () if collect_hidden else None ) - mask_kwargs = { - "config": self.config, - _CREATE_CAUSAL_MASK_ARG: inputs_embeds, - "attention_mask": attention_mask, - "cache_position": cache_position, - "past_key_values": past_key_values, - "position_ids": position_ids, - } - causal_mask = mask_function(**mask_kwargs) + for layer in self.layers: + if all_hidden_states is not None: + all_hidden_states += ( + (hidden_states,) + if residual is None + else (hidden_states + residual,) + ) + hidden_states, residual = layer(position_ids, hidden_states, residual) - hidden_states = inputs_embeds - position_embeddings = self.rotary_emb(hidden_states, position_ids) + hidden_states, _ = self.norm(hidden_states, residual) + if all_hidden_states is not None: + all_hidden_states += (hidden_states,) - hidden_states_pool = [] if output_hidden_states else None - for decoder_layer in self.layers[: self.config.num_hidden_layers]: - if output_hidden_states: - hidden_states_pool.append(hidden_states) - hidden_states = decoder_layer( - hidden_states, - attention_mask=causal_mask, - position_ids=position_ids, - past_key_values=past_key_values, - use_cache=use_cache, - cache_position=cache_position, - position_embeddings=position_embeddings, - **kwargs, - ) - - hidden_states = self.norm(hidden_states) - if output_hidden_states: - hidden_states_pool.append(hidden_states) - - return BaseModelOutputWithPast( - hidden_states=hidden_states_pool, + return BaseEncoderOutput( last_hidden_state=hidden_states, - past_key_values=past_key_values if use_cache else None, + hidden_states=all_hidden_states, ) class Mistral3Model(nn.Module): - _checkpoint_conversion_mapping = {"language_model.model": "language_model"} + """Module-layout wrapper: exposes `language_model.*` under `model.*`.""" - def __init__(self, config: Mistral3Config): + def __init__(self, config: Mistral3EncoderConfig, prefix: str = "") -> None: super().__init__() - self.language_model = MistralModel(config.text_config) - self.config = config + self.language_model = MistralModel(config, prefix=f"{prefix}.language_model") - def get_input_embeddings(self): - return self.language_model.embed_tokens + def get_input_embeddings(self, input_ids: torch.Tensor) -> torch.Tensor: + return self.language_model.get_input_embeddings(input_ids) - def set_decoder(self, decoder): - self.language_model = decoder - - def get_decoder(self): - return self.language_model - - def forward( - self, - input_ids: Optional[torch.LongTensor] = None, - pixel_values: Optional[torch.FloatTensor] = None, - attention_mask: Optional[torch.Tensor] = None, - position_ids: Optional[torch.LongTensor] = None, - past_key_values: Optional[Cache] = None, - inputs_embeds: Optional[torch.FloatTensor] = None, - vision_feature_layer: Optional[Union[int, list[int]]] = None, - use_cache: Optional[bool] = None, - output_attentions: Optional[bool] = None, - output_hidden_states: Optional[bool] = None, - return_dict: Optional[bool] = None, - cache_position: Optional[torch.LongTensor] = None, - image_sizes: Optional[torch.Tensor] = None, - **kwargs, - ) -> Union[tuple, Mistral3ModelOutputWithPast]: - del pixel_values, vision_feature_layer, return_dict - output_attentions = False - output_hidden_states = True - - if (input_ids is None) ^ (inputs_embeds is not None): - raise ValueError( - "You must specify exactly one of input_ids or inputs_embeds" - ) - - if inputs_embeds is None: - inputs_embeds = self.get_input_embeddings()(input_ids) - - outputs: BaseModelOutputWithPast = self.language_model( - attention_mask=attention_mask, - position_ids=position_ids, - past_key_values=past_key_values, - inputs_embeds=inputs_embeds, - use_cache=use_cache, - output_attentions=output_attentions, - output_hidden_states=output_hidden_states, - return_dict=True, - cache_position=cache_position, - **kwargs, - ) - - return Mistral3ModelOutputWithPast( - last_hidden_state=outputs.last_hidden_state, - past_key_values=outputs.past_key_values, - hidden_states=outputs.hidden_states, - attentions=outputs.attentions, - ) + def forward(self, *args, **kwargs) -> BaseEncoderOutput: + return self.language_model(*args, **kwargs) -class Mistral3ForConditionalGeneration(nn.Module, LayerwiseOffloadableModuleMixin): +# Scalars that may live under Mistral3Config.text_config and must be hoisted. +_HOISTED_TEXT_CONFIG_FIELDS = ( + "vocab_size", + "hidden_size", + "intermediate_size", + "num_hidden_layers", + "num_attention_heads", + "num_key_value_heads", + "head_dim", + "hidden_act", + "max_position_embeddings", + "rms_norm_eps", + "tie_word_embeddings", + "pad_token_id", + "bos_token_id", + "eos_token_id", + "attention_bias", + "mlp_bias", + "sliding_window", +) + + +def _hoist_text_config(arch_config) -> None: + """Lift nested HF Mistral3Config.text_config scalars onto arch_config.""" + text_config = getattr(arch_config, "text_config", None) + if text_config is None: + return + for field_name in _HOISTED_TEXT_CONFIG_FIELDS: + if hasattr(text_config, field_name): + value = getattr(text_config, field_name) + if value is not None: + setattr(arch_config, field_name, value) + rope_theta = getattr(text_config, "rope_theta", None) + if rope_theta is not None: + rope_params = dict(getattr(arch_config, "rope_parameters", None) or {}) + rope_params["rope_theta"] = float(rope_theta) + rope_scaling = getattr(text_config, "rope_scaling", None) + if rope_scaling: + rope_params.update(rope_scaling) + arch_config.rope_parameters = rope_params + + +class Mistral3ForConditionalGeneration(TextEncoder, LayerwiseOffloadableModuleMixin): _checkpoint_conversion_mapping = { "^language_model.model": "model.language_model", "^multi_modal_projector": "model.multi_modal_projector", "^language_model.lm_head": "lm_head", } - _tied_weights_keys = ["lm_head.weight"] - uses_sglang_forward_context = False + uses_sglang_forward_context = True layerwise_offload_dit_group_enabled = False layer_names = ["model.language_model.layers"] + _supported_attention_backends = ( + Mistral3EncoderConfig()._supported_attention_backends + ) - def __init__(self, config: LlavaConfig): - super().__init__() - self.model = Mistral3Model(config.arch_config) + def __init__(self, config: Mistral3EncoderConfig) -> None: + super().__init__(config) + _hoist_text_config(config.arch_config) + self.model = Mistral3Model(config, prefix="model") - def get_input_embeddings(self): - return self.model.get_input_embeddings() - - def set_decoder(self, decoder): - self.model.set_decoder(decoder) - - def get_decoder(self): - return self.model.get_decoder() - - # Make modules available through conditional class for BC @property - def language_model(self): + def language_model(self) -> MistralModel: return self.model.language_model + def get_input_embeddings(self, input_ids: torch.Tensor) -> torch.Tensor: + return self.model.get_input_embeddings(input_ids) + def forward( self, - input_ids: Optional[torch.LongTensor] = None, - attention_mask: Optional[torch.Tensor] = None, - position_ids: Optional[torch.LongTensor] = None, - past_key_values: Optional[Cache] = None, - output_hidden_states: Optional[bool] = None, - inputs_embeds: Optional[torch.FloatTensor] = None, - labels: Optional[torch.LongTensor] = None, - use_cache: Optional[bool] = None, - return_dict: Optional[bool] = None, - cache_position: Optional[torch.LongTensor] = None, - logits_to_keep: Union[int, torch.Tensor] = 0, - image_sizes: Optional[torch.Tensor] = None, + input_ids: torch.Tensor | None = None, + position_ids: torch.Tensor | None = None, + attention_mask: torch.Tensor | None = None, + inputs_embeds: torch.Tensor | None = None, + output_hidden_states: bool | None = True, **kwargs, - ) -> Union[tuple, Mistral3CausalLMOutputWithPast]: - r""" - labels (`torch.LongTensor` of shape `(batch_size, sequence_length)`, *optional*): - Labels for computing the masked language modeling loss. Indices should either be in `[0, ..., - config.vocab_size]` or -100 (see `input_ids` docstring). Tokens with indices set to `-100` are ignored - (masked), the loss is only computed for the tokens with labels in `[0, ..., config.vocab_size]`. - - Example: - - """ - output_hidden_states = True - - execution_tensor = input_ids if input_ids is not None else inputs_embeds - sdpa_context = ( - sdpa_kernel(SDPBackend.CUDNN_ATTENTION) - if execution_tensor is not None - and execution_tensor.device.type == "cuda" - and current_platform.is_cuda() - else nullcontext() - ) - with sdpa_context: - # FLUX.2 uses the text-only Mistral3 path but still expects the - # same local SDPA kernel choice as the official HF implementation. - outputs = self.model( - input_ids=input_ids, - attention_mask=attention_mask, - position_ids=position_ids, - past_key_values=past_key_values, - inputs_embeds=inputs_embeds, - use_cache=use_cache, - output_hidden_states=output_hidden_states, - return_dict=True, - cache_position=cache_position, - image_sizes=image_sizes, - **kwargs, - ) - - return Mistral3CausalLMOutputWithPast( - hidden_states=outputs.hidden_states, + ) -> BaseEncoderOutput: + return self.model( + input_ids=input_ids, + position_ids=position_ids, + attention_mask=attention_mask, + inputs_embeds=inputs_embeds, + output_hidden_states=( + True if output_hidden_states is None else output_hidden_states + ), + **kwargs, ) def load_weights(self, weights: Iterable[tuple[str, torch.Tensor]]) -> set[str]: - # Define mapping for stacked parameters params_dict = dict(self.named_parameters()) loaded_params: set[str] = set() + stacked_params_mapping = self.config.arch_config.stacked_params_mapping for name, loaded_weight in weights: name_lower = name.lower() if ( @@ -470,16 +431,38 @@ class Mistral3ForConditionalGeneration(nn.Module, LayerwiseOffloadableModuleMixi or "lm_head" in name_lower ): continue - final_name = name.replace("language_model.model.", "model.language_model.") + name = name.replace("language_model.model.", "model.language_model.") + if "rotary_emb.inv_freq" in name: + continue + if "scale" in name: + kv_scale_name = maybe_remap_kv_scale_name(name, params_dict) + if kv_scale_name is None: + continue + name = kv_scale_name - if final_name in params_dict: - param = params_dict[final_name] + for param_name, weight_name, shard_id in stacked_params_mapping: + if weight_name not in name: + continue + merged_name = name.replace(weight_name, param_name) + if merged_name.endswith(".bias") and merged_name not in params_dict: + break + if merged_name not in params_dict: + break + param = params_dict[merged_name] + weight_loader = param.weight_loader + weight_loader(param, loaded_weight, shard_id) + loaded_params.add(merged_name) + break + else: + if name.endswith(".bias") and name not in params_dict: + continue + if name not in params_dict: + logger.warning("Param %s from weight is not loaded", name) + continue + param = params_dict[name] weight_loader = getattr(param, "weight_loader", default_weight_loader) weight_loader(param, loaded_weight) - loaded_params.add(final_name) - else: - logger.warning(f"Param {name=} {final_name=} from weight is not loaded") - + loaded_params.add(name) return loaded_params