[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, create_sliding_window_causal_mask,
) )
from transformers.modeling_outputs import BaseModelOutputWithPast from transformers.modeling_outputs import BaseModelOutputWithPast
from transformers.modeling_utils import ALL_ATTENTION_FUNCTIONS
from transformers.models.mistral3.modeling_mistral3 import ( from transformers.models.mistral3.modeling_mistral3 import (
Mistral3CausalLMOutputWithPast, Mistral3CausalLMOutputWithPast,
Mistral3ModelOutputWithPast, Mistral3ModelOutputWithPast,
@@ -37,13 +36,13 @@ from transformers.models.mistral.modeling_mistral import (
MistralRMSNorm, MistralRMSNorm,
MistralRotaryEmbedding, MistralRotaryEmbedding,
apply_rotary_pos_emb, apply_rotary_pos_emb,
eager_attention_forward,
) )
from sglang.multimodal_gen.runtime.distributed import ( from sglang.multimodal_gen.runtime.distributed import (
get_tp_world_size, get_tp_world_size,
model_parallel_is_initialized, model_parallel_is_initialized,
) )
from sglang.multimodal_gen.runtime.layers.attention import LocalAttention
from sglang.multimodal_gen.runtime.layers.linear import ( from sglang.multimodal_gen.runtime.layers.linear import (
ColumnParallelLinear, ColumnParallelLinear,
RowParallelLinear, 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 ( 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.platforms import (
AttentionBackendEnum,
current_platform,
)
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__)
@@ -107,6 +109,23 @@ def _make_row_linear(
return nn.Linear(in_features, out_features, bias=bias) 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: 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). 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, use_tensor_parallel=self.use_tensor_parallel,
) )
self.is_causal = True 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( def forward(
self, self,
@@ -224,32 +255,15 @@ class MistralAttention(nn.Module):
key_states, value_states, self.layer_idx, cache_kwargs key_states, value_states, self.layer_idx, cache_kwargs
) )
attn_implementation = getattr(self.config, "_attn_implementation", None) attn_output = self.attn(
attention_interface = eager_attention_forward query_states.transpose(1, 2),
if attn_implementation and attn_implementation != "eager": key_states.transpose(1, 2),
if hasattr(ALL_ATTENTION_FUNCTIONS, "get_interface"): value_states.transpose(1, 2),
attention_interface = ALL_ATTENTION_FUNCTIONS.get_interface( attn_mask=attention_mask,
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 = attn_output.reshape(*input_shape, -1).contiguous()
attn_output = _linear_output(self.o_proj, attn_output) attn_output = _linear_output(self.o_proj, attn_output)
return attn_output, attn_weights return attn_output, None
class MistralTPMLP(nn.Module): class MistralTPMLP(nn.Module):
@@ -388,20 +402,25 @@ class MistralModel(MistralPreTrainedModel):
if position_ids is None: if position_ids is None:
position_ids = cache_position.unsqueeze(0) position_ids = cache_position.unsqueeze(0)
mask_function = ( if _can_use_unmasked_causal_attention(
create_causal_mask attention_mask, self.config, past_key_values
if getattr(self.config, "sliding_window", None) is None ):
else create_sliding_window_causal_mask causal_mask = None
) else:
mask_kwargs = { mask_function = (
"config": self.config, create_causal_mask
_CREATE_CAUSAL_MASK_ARG: inputs_embeds, if getattr(self.config, "sliding_window", None) is None
"attention_mask": attention_mask, else create_sliding_window_causal_mask
"cache_position": cache_position, )
"past_key_values": past_key_values, mask_kwargs = {
"position_ids": position_ids, "config": self.config,
} _CREATE_CAUSAL_MASK_ARG: inputs_embeds,
causal_mask = mask_function(**mask_kwargs) "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 hidden_states = inputs_embeds
position_embeddings = self.rotary_emb(hidden_states, position_ids) position_embeddings = self.rotary_emb(hidden_states, position_ids)
@@ -506,7 +525,7 @@ class Mistral3ForConditionalGeneration(nn.Module, LayerwiseOffloadableModuleMixi
"^language_model.lm_head": "lm_head", "^language_model.lm_head": "lm_head",
} }
_tied_weights_keys = ["lm_head.weight"] _tied_weights_keys = ["lm_head.weight"]
uses_sglang_forward_context = False uses_sglang_forward_context = True
layerwise_offload_dit_group_enabled = False layerwise_offload_dit_group_enabled = False
layer_names = ["model.language_model.layers"] layer_names = ["model.language_model.layers"]