Mistral3 add tensor parallel support for diffusion text encoder (#25950)

This commit is contained in:
sushil Dubey
2026-06-10 09:43:21 +08:00
committed by GitHub
parent af55025644
commit 5809bbe35d
4 changed files with 375 additions and 384 deletions
@@ -7,7 +7,9 @@ from sglang.multimodal_gen.configs.models.encoders.base import (
TextEncoderArchConfig, TextEncoderArchConfig,
TextEncoderConfig, 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 = ( FLUX_2_SYSTEM_MESSAGE = (
"You are an AI that reasons about image descriptions. You give structured responses focusing on object relationships, object\n" "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 @dataclass
class Flux2MistralTextArchConfig(TextEncoderArchConfig): class Flux2MistralTextArchConfig(Mistral3EncoderArchConfig):
stacked_params_mapping: list[tuple[str, str, str]] = field( """FLUX.2 text-encoder arch config.
default_factory=lambda: [
("qkv_proj", "q_proj", "q"), Inherits Mistral3 defaults (hidden_size, num_attention_heads, head_dim,
("qkv_proj", "k_proj", "k"), rms_norm_eps, rope_parameters, ...) so the TP-parallel runtime has every
("qkv_proj", "v_proj", "v"), 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.
_fsdp_shard_conditions: list = field(default_factory=lambda: [is_layer]) """
def __post_init__(self) -> None: def __post_init__(self) -> None:
self.tokenizer_kwargs = { self.tokenizer_kwargs = {
@@ -2,16 +2,12 @@
"""Mistral3 text encoder configuration for SGLang diffusion models.""" """Mistral3 text encoder configuration for SGLang diffusion models."""
from dataclasses import dataclass, field from dataclasses import dataclass, field
from typing import Any
from sglang.multimodal_gen.configs.models.encoders.base import ( from sglang.multimodal_gen.configs.models.encoders.base import (
TextEncoderArchConfig, TextEncoderArchConfig,
TextEncoderConfig, TextEncoderConfig,
) )
from sglang.multimodal_gen.configs.models.fsdp import (
is_embed_tokens,
is_final_norm,
is_layer,
)
@dataclass @dataclass
@@ -38,6 +34,12 @@ class Mistral3EncoderArchConfig(TextEncoderArchConfig):
head_dim: int = 128 head_dim: int = 128
hidden_state_skip_layer: int = 2 # Use second-to-last hidden state hidden_state_skip_layer: int = 2 # Use second-to-last hidden state
text_len: int = 0 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( stacked_params_mapping: list[tuple[str, str, str]] = field(
default_factory=lambda: [ default_factory=lambda: [
@@ -49,9 +51,8 @@ class Mistral3EncoderArchConfig(TextEncoderArchConfig):
] ]
) )
_fsdp_shard_conditions: list = field( # TP-parallel runtime shards weights along TP dim; no FSDP needed.
default_factory=lambda: [is_layer, is_embed_tokens, is_final_norm] _fsdp_shard_conditions: list = field(default_factory=list)
)
def __post_init__(self): def __post_init__(self):
# Let the parent populate tokenizer_kwargs["max_length"] = self.text_len # Let the parent populate tokenizer_kwargs["max_length"] = self.text_len
@@ -69,6 +69,11 @@ class CustomOp(nn.Module):
# PyTorch-native implementation. # PyTorch-native implementation.
return self.forward_native(*args, **kwargs) 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: def dispatch_forward(self) -> Callable:
if _is_cuda: if _is_cuda:
return self.forward_cuda return self.forward_cuda
@@ -13,455 +13,416 @@
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. # WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and # See the License for the specific language governing permissions and
# limitations under the License. # limitations under the License.
import inspect
from contextlib import nullcontext from collections.abc import Iterable
from typing import Iterable, Optional, Union from typing import Any
import torch import torch
from torch import nn 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 ( from sglang.multimodal_gen.runtime.managers.memory_managers.layerwise_offload import (
LayerwiseOffloadableModuleMixin, 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 from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
logger = init_logger(__name__) 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: class MistralMLP(nn.Module):
""" def __init__(
This is the equivalent of torch.repeat_interleave(x, dim=1, repeats=n_rep). self,
The hidden states go from (batch, num_key_value_heads, seqlen, head_dim) to hidden_size: int,
(batch, num_attention_heads, seqlen, head_dim) intermediate_size: int,
""" hidden_act: str,
batch, num_key_value_heads, slen, head_dim = hidden_states.shape quant_config: QuantizationConfig | None = None,
if n_rep == 1: bias: bool = False,
return hidden_states prefix: str = "",
hidden_states = hidden_states[:, :, None, :, :].expand( ) -> None:
batch, num_key_value_heads, n_rep, slen, head_dim super().__init__()
) self.gate_up_proj = MergedColumnParallelLinear(
return hidden_states.reshape(batch, num_key_value_heads * n_rep, slen, head_dim) 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): class MistralAttention(nn.Module):
"""Multi-headed attention from 'Attention Is All You Need' paper""" def __init__(
self,
def __init__(self, config: MistralConfig, layer_idx: int): 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__() super().__init__()
self.config = config self.hidden_size = hidden_size
self.layer_idx = layer_idx tp_size = get_tp_world_size()
self.head_dim = ( self.total_num_heads = num_heads
getattr(config, "head_dim", None) assert self.total_num_heads % tp_size == 0
or config.hidden_size // config.num_attention_heads self.num_heads = self.total_num_heads // tp_size
) self.total_num_kv_heads = num_kv_heads
self.num_key_value_groups = ( if self.total_num_kv_heads >= tp_size:
config.num_attention_heads // config.num_key_value_heads 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.scaling = self.head_dim**-0.5
self.attention_dropout = config.attention_dropout self.rope_theta = rope_theta
self.q_proj = nn.Linear( self.max_position_embeddings = max_position_embeddings
config.hidden_size, config.num_attention_heads * self.head_dim, bias=False
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( self.o_proj = RowParallelLinear(
config.hidden_size, config.num_key_value_heads * self.head_dim, bias=False 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( def forward(
self, self,
positions: torch.Tensor,
hidden_states: torch.Tensor, hidden_states: torch.Tensor,
position_embeddings: tuple[torch.Tensor, torch.Tensor], ) -> torch.Tensor:
attention_mask: Optional[torch.Tensor], qkv, _ = self.qkv_proj(hidden_states)
past_key_values: Optional[Cache] = None, q, k, v = qkv.split([self.q_size, self.kv_size, self.kv_size], dim=-1)
cache_position: Optional[torch.LongTensor] = None, q, k = self.rotary_emb(positions, q, k)
**kwargs,
) -> tuple[torch.Tensor, Optional[torch.Tensor]]:
input_shape = hidden_states.shape[:-1]
hidden_shape = (*input_shape, -1, self.head_dim)
query_states = self.q_proj(hidden_states).view(hidden_shape).transpose(1, 2) batch_size = q.shape[0]
key_states = self.k_proj(hidden_states).view(hidden_shape).transpose(1, 2) seq_len = q.shape[1]
value_states = self.v_proj(hidden_states).view(hidden_shape).transpose(1, 2) 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)
cos, sin = position_embeddings v = v.reshape(batch_size, seq_len, self.num_kv_heads, self.head_dim)
query_states, key_states = apply_rotary_pos_emb( attn_output = self.attn(q, k, v)
query_states, key_states, cos, sin attn_output = attn_output.reshape(
batch_size, seq_len, self.num_heads * self.head_dim
) )
output, _ = self.o_proj(attn_output)
if past_key_values is not None: return output
# 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
class MistralDecoderLayer(nn.Module): 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__() super().__init__()
self.hidden_size = config.hidden_size self.hidden_size = config.hidden_size
self.self_attn = MistralAttention(config=config, layer_idx=layer_idx) rope_parameters = getattr(config, "rope_parameters", None) or {}
self.mlp = MistralMLP(config) rope_theta = rope_parameters.get("rope_theta", 10000.0)
self.input_layernorm = MistralRMSNorm( rope_scaling = rope_parameters or None
config.hidden_size, eps=config.rms_norm_eps 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 config.hidden_size, eps=config.rms_norm_eps
) )
def forward( def forward(
self, self,
positions: torch.Tensor,
hidden_states: torch.Tensor, hidden_states: torch.Tensor,
attention_mask: Optional[torch.Tensor] = None, residual: torch.Tensor | None,
position_ids: Optional[torch.LongTensor] = None, ) -> tuple[torch.Tensor, torch.Tensor]:
past_key_values: Optional[Cache] = None, if residual is None:
use_cache: Optional[bool] = False, residual = hidden_states
cache_position: Optional[torch.LongTensor] = None, hidden_states = self.input_layernorm(hidden_states)
position_embeddings: Optional[ else:
tuple[torch.Tensor, torch.Tensor] hidden_states, residual = self.input_layernorm(hidden_states, residual)
] = None, # necessary, but kept here for BC hidden_states = self.self_attn(positions=positions, hidden_states=hidden_states)
**kwargs, hidden_states, residual = self.post_attention_layernorm(hidden_states, residual)
) -> 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)
hidden_states = self.mlp(hidden_states) hidden_states = self.mlp(hidden_states)
hidden_states = residual + hidden_states return hidden_states, residual
return hidden_states
class MistralModel(MistralPreTrainedModel): class MistralModel(nn.Module):
def __init__(self, config: MistralConfig): """TP-parallel Mistral decoder stack used as a text encoder."""
super().__init__(config)
self.padding_idx = config.pad_token_id def __init__(self, config: Mistral3EncoderConfig, prefix: str = "") -> None:
super().__init__()
self.config = config
self.vocab_size = config.vocab_size self.vocab_size = config.vocab_size
self.embed_tokens = nn.Embedding( self.embed_tokens = VocabParallelEmbedding(
config.vocab_size, config.hidden_size, self.padding_idx config.vocab_size,
config.hidden_size,
org_num_embeddings=config.vocab_size,
quant_config=getattr(config, "quant_config", None),
) )
self.layers = nn.ModuleList( self.layers = nn.ModuleList(
[ [
MistralDecoderLayer(config, layer_idx) MistralDecoderLayer(
for layer_idx in range(config.num_hidden_layers) 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.norm = RMSNorm(config.hidden_size, eps=config.rms_norm_eps)
self.rotary_emb = MistralRotaryEmbedding(config=config)
self.gradient_checkpointing = False def get_input_embeddings(self, input_ids: torch.Tensor) -> torch.Tensor:
self.post_init() return self.embed_tokens(input_ids)
def forward( def forward(
self, self,
input_ids: Optional[torch.LongTensor] = None, input_ids: torch.Tensor | None,
attention_mask: Optional[torch.Tensor] = None, position_ids: torch.Tensor | None = None,
position_ids: Optional[torch.LongTensor] = None, attention_mask: torch.Tensor | None = None,
past_key_values: Optional[Cache] = None, inputs_embeds: torch.Tensor | None = None,
inputs_embeds: Optional[torch.FloatTensor] = None, output_hidden_states: bool | None = True,
use_cache: Optional[bool] = None,
cache_position: Optional[torch.LongTensor] = None,
output_hidden_states: Optional[bool] = None,
**kwargs, **kwargs,
) -> BaseModelOutputWithPast: ) -> BaseEncoderOutput:
if (input_ids is None) ^ (inputs_embeds is not None): if inputs_embeds is not None:
raise ValueError( hidden_states = inputs_embeds
"You must specify exactly one of input_ids or inputs_embeds" else:
) hidden_states = self.get_input_embeddings(input_ids)
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,
)
if position_ids is None: if position_ids is None:
position_ids = cache_position.unsqueeze(0) position_ids = torch.arange(
mask_function = ( 0, hidden_states.shape[1], device=hidden_states.device
create_causal_mask ).unsqueeze(0)
if getattr(self.config, "sliding_window", None) is None
else create_sliding_window_causal_mask 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 = { for layer in self.layers:
"config": self.config, if all_hidden_states is not None:
_CREATE_CAUSAL_MASK_ARG: inputs_embeds, all_hidden_states += (
"attention_mask": attention_mask, (hidden_states,)
"cache_position": cache_position, if residual is None
"past_key_values": past_key_values, else (hidden_states + residual,)
"position_ids": position_ids, )
} hidden_states, residual = layer(position_ids, hidden_states, residual)
causal_mask = mask_function(**mask_kwargs)
hidden_states = inputs_embeds hidden_states, _ = self.norm(hidden_states, residual)
position_embeddings = self.rotary_emb(hidden_states, position_ids) if all_hidden_states is not None:
all_hidden_states += (hidden_states,)
hidden_states_pool = [] if output_hidden_states else None return BaseEncoderOutput(
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,
last_hidden_state=hidden_states, 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): 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__() super().__init__()
self.language_model = MistralModel(config.text_config) self.language_model = MistralModel(config, prefix=f"{prefix}.language_model")
self.config = config
def get_input_embeddings(self): def get_input_embeddings(self, input_ids: torch.Tensor) -> torch.Tensor:
return self.language_model.embed_tokens return self.language_model.get_input_embeddings(input_ids)
def set_decoder(self, decoder): def forward(self, *args, **kwargs) -> BaseEncoderOutput:
self.language_model = decoder return self.language_model(*args, **kwargs)
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,
)
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 = { _checkpoint_conversion_mapping = {
"^language_model.model": "model.language_model", "^language_model.model": "model.language_model",
"^multi_modal_projector": "model.multi_modal_projector", "^multi_modal_projector": "model.multi_modal_projector",
"^language_model.lm_head": "lm_head", "^language_model.lm_head": "lm_head",
} }
_tied_weights_keys = ["lm_head.weight"] uses_sglang_forward_context = True
uses_sglang_forward_context = False
layerwise_offload_dit_group_enabled = False layerwise_offload_dit_group_enabled = False
layer_names = ["model.language_model.layers"] layer_names = ["model.language_model.layers"]
_supported_attention_backends = (
Mistral3EncoderConfig()._supported_attention_backends
)
def __init__(self, config: LlavaConfig): def __init__(self, config: Mistral3EncoderConfig) -> None:
super().__init__() super().__init__(config)
self.model = Mistral3Model(config.arch_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 @property
def language_model(self): def language_model(self) -> MistralModel:
return self.model.language_model 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( def forward(
self, self,
input_ids: Optional[torch.LongTensor] = None, input_ids: torch.Tensor | None = None,
attention_mask: Optional[torch.Tensor] = None, position_ids: torch.Tensor | None = None,
position_ids: Optional[torch.LongTensor] = None, attention_mask: torch.Tensor | None = None,
past_key_values: Optional[Cache] = None, inputs_embeds: torch.Tensor | None = None,
output_hidden_states: Optional[bool] = None, output_hidden_states: bool | None = True,
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,
**kwargs, **kwargs,
) -> Union[tuple, Mistral3CausalLMOutputWithPast]: ) -> BaseEncoderOutput:
r""" return self.model(
labels (`torch.LongTensor` of shape `(batch_size, sequence_length)`, *optional*): input_ids=input_ids,
Labels for computing the masked language modeling loss. Indices should either be in `[0, ..., position_ids=position_ids,
config.vocab_size]` or -100 (see `input_ids` docstring). Tokens with indices set to `-100` are ignored attention_mask=attention_mask,
(masked), the loss is only computed for the tokens with labels in `[0, ..., config.vocab_size]`. inputs_embeds=inputs_embeds,
output_hidden_states=(
Example: True if output_hidden_states is None else output_hidden_states
),
""" **kwargs,
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,
) )
def load_weights(self, weights: Iterable[tuple[str, torch.Tensor]]) -> set[str]: def load_weights(self, weights: Iterable[tuple[str, torch.Tensor]]) -> set[str]:
# Define mapping for stacked parameters
params_dict = dict(self.named_parameters()) params_dict = dict(self.named_parameters())
loaded_params: set[str] = set() loaded_params: set[str] = set()
stacked_params_mapping = self.config.arch_config.stacked_params_mapping
for name, loaded_weight in weights: for name, loaded_weight in weights:
name_lower = name.lower() name_lower = name.lower()
if ( if (
@@ -470,16 +431,38 @@ class Mistral3ForConditionalGeneration(nn.Module, LayerwiseOffloadableModuleMixi
or "lm_head" in name_lower or "lm_head" in name_lower
): ):
continue 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: for param_name, weight_name, shard_id in stacked_params_mapping:
param = params_dict[final_name] 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 = getattr(param, "weight_loader", default_weight_loader)
weight_loader(param, loaded_weight) weight_loader(param, loaded_weight)
loaded_params.add(final_name) loaded_params.add(name)
else:
logger.warning(f"Param {name=} {final_name=} from weight is not loaded")
return loaded_params return loaded_params