[diffusion] feat: use LocalAttention for mistral3 encoder (#28176)

This commit is contained in:
Mick
2026-06-17 21:18:41 +08:00
committed by GitHub
parent dad890fff1
commit 735a256f98
@@ -27,7 +27,6 @@ from transformers.masking_utils import (
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,
@@ -37,13 +36,13 @@ from transformers.models.mistral.modeling_mistral import (
MistralRMSNorm,
MistralRotaryEmbedding,
apply_rotary_pos_emb,
eager_attention_forward,
)
from sglang.multimodal_gen.runtime.distributed import (
get_tp_world_size,
model_parallel_is_initialized,
)
from sglang.multimodal_gen.runtime.layers.attention import LocalAttention
from sglang.multimodal_gen.runtime.layers.linear import (
ColumnParallelLinear,
RowParallelLinear,
@@ -52,7 +51,10 @@ from sglang.multimodal_gen.runtime.loader.weight_utils import default_weight_loa
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.platforms import (
AttentionBackendEnum,
current_platform,
)
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
logger = init_logger(__name__)
@@ -107,6 +109,23 @@ def _make_row_linear(
return nn.Linear(in_features, out_features, bias=bias)
def _can_use_unmasked_causal_attention(
attention_mask: Optional[torch.Tensor],
config: MistralConfig,
past_key_values: Optional[Cache],
) -> bool:
if (
getattr(config, "sliding_window", None) is not None
or past_key_values is not None
):
return False
if attention_mask is None:
return True
if attention_mask.dim() != 2:
return False
return bool(torch.all(attention_mask > 0).item())
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).
@@ -183,6 +202,18 @@ class MistralAttention(nn.Module):
use_tensor_parallel=self.use_tensor_parallel,
)
self.is_causal = True
self.attn = LocalAttention(
self.num_heads,
self.head_dim,
self.num_key_value_heads,
softmax_scale=self.scaling,
causal=True,
supported_attention_backends={
AttentionBackendEnum.FA,
AttentionBackendEnum.TORCH_SDPA,
},
allow_cudnn_sdp=True,
)
def forward(
self,
@@ -224,32 +255,15 @@ class MistralAttention(nn.Module):
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 = self.attn(
query_states.transpose(1, 2),
key_states.transpose(1, 2),
value_states.transpose(1, 2),
attn_mask=attention_mask,
)
attn_output = attn_output.reshape(*input_shape, -1).contiguous()
attn_output = _linear_output(self.o_proj, attn_output)
return attn_output, attn_weights
return attn_output, None
class MistralTPMLP(nn.Module):
@@ -388,20 +402,25 @@ class MistralModel(MistralPreTrainedModel):
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
)
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)
if _can_use_unmasked_causal_attention(
attention_mask, self.config, past_key_values
):
causal_mask = None
else:
mask_function = (
create_causal_mask
if getattr(self.config, "sliding_window", None) is None
else create_sliding_window_causal_mask
)
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)
hidden_states = inputs_embeds
position_embeddings = self.rotary_emb(hidden_states, position_ids)
@@ -506,7 +525,7 @@ class Mistral3ForConditionalGeneration(nn.Module, LayerwiseOffloadableModuleMixi
"^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"]