From b1c688fba21da85a53271691fdab224b08d67bef Mon Sep 17 00:00:00 2001 From: Xinyuan Tong <115166877+JustinTong0323@users.noreply.github.com> Date: Sun, 16 Nov 2025 06:44:40 -0800 Subject: [PATCH] refactor: cleanup vision attention related codes (#13228) Signed-off-by: Xinyuan Tong Co-authored-by: alisonshao <54658187+alisonshao@users.noreply.github.com> Co-authored-by: Mick Co-authored-by: Baizhou Zhang Co-authored-by: Kangyan-Zhou --- python/sglang/srt/models/clip.py | 13 --------- .../sglang/srt/models/deepseek_janus_pro.py | 2 -- python/sglang/srt/models/dots_vlm_vit.py | 12 +-------- python/sglang/srt/models/glm4v.py | 27 +------------------ python/sglang/srt/models/idefics2.py | 1 - python/sglang/srt/models/internvl.py | 1 - python/sglang/srt/models/mllama.py | 3 --- python/sglang/srt/models/mllama4.py | 3 --- python/sglang/srt/models/pixtral.py | 1 - python/sglang/srt/models/qwen2_5_vl.py | 27 +------------------ python/sglang/srt/models/qwen2_vl.py | 13 --------- python/sglang/srt/models/qwen3_omni_moe.py | 3 --- python/sglang/srt/models/qwen3_vl.py | 24 +---------------- python/sglang/srt/models/siglip.py | 13 --------- python/sglang/srt/models/step3_vl.py | 3 --- 15 files changed, 4 insertions(+), 142 deletions(-) diff --git a/python/sglang/srt/models/clip.py b/python/sglang/srt/models/clip.py index ea9fee9ac..9294e6f88 100644 --- a/python/sglang/srt/models/clip.py +++ b/python/sglang/srt/models/clip.py @@ -141,7 +141,6 @@ class CLIPEncoderLayer(nn.Module): config: CLIPVisionConfig, act_layer: Type[nn.Module] = QuickGELU, norm_layer: Type[nn.Module] = None, - attn_implementation: Optional[str] = "sdpa", quant_config: Optional[QuantizationConfig] = None, prefix: str = "", ) -> None: @@ -150,22 +149,11 @@ class CLIPEncoderLayer(nn.Module): norm_layer = partial(nn.LayerNorm, eps=config.layer_norm_eps) self.layer_norm1 = norm_layer(config.hidden_size) self.layer_norm2 = norm_layer(config.hidden_size) - if attn_implementation == "sdpa": - qkv_backend = "sdpa" - softmax_in_single_precision = False - elif attn_implementation == "flash_attention_2": - qkv_backend = "triton_attn" - softmax_in_single_precision = False - elif attn_implementation == "eager": - qkv_backend = "sdpa" - softmax_in_single_precision = True self.self_attn = VisionAttention( embed_dim=config.hidden_size, num_heads=config.num_attention_heads, projection_size=config.hidden_size, use_qkv_parallel=True, - qkv_backend=qkv_backend, - softmax_in_single_precision=softmax_in_single_precision, flatten_batch=True, quant_config=quant_config, prefix=add_prefix("self_attn", prefix), @@ -233,7 +221,6 @@ class CLIPEncoder(nn.Module): CLIPEncoderLayer( config=config, norm_layer=norm_layer, - attn_implementation="sdpa", quant_config=quant_config, prefix=add_prefix(f"layers.{layer_idx}", prefix), ) diff --git a/python/sglang/srt/models/deepseek_janus_pro.py b/python/sglang/srt/models/deepseek_janus_pro.py index fe1c833f7..2167c4824 100644 --- a/python/sglang/srt/models/deepseek_janus_pro.py +++ b/python/sglang/srt/models/deepseek_janus_pro.py @@ -532,8 +532,6 @@ class VisionTransformerBlock(nn.Module): num_heads=num_heads, projection_size=dim, use_qkv_parallel=True, - qkv_backend="sdpa", - softmax_in_single_precision=False, dropout=attn_drop, ) diff --git a/python/sglang/srt/models/dots_vlm_vit.py b/python/sglang/srt/models/dots_vlm_vit.py index d094e1deb..ca82ddd5e 100644 --- a/python/sglang/srt/models/dots_vlm_vit.py +++ b/python/sglang/srt/models/dots_vlm_vit.py @@ -154,21 +154,13 @@ class DotsVisionBlock(nn.Module): config: DotsVisionConfig, quant_config: Optional[QuantizationConfig] = None, prefix: str = "", - attn_implementation: str = "flash_attention_2", ): super().__init__() - if attn_implementation == "flash_attention_2": - qkv_backend = "fa3" - softmax_in_single_precision = False - else: - raise RuntimeError("Unimplemented") self.attn = VisionAttention( embed_dim=config.embed_dim, num_heads=config.num_attention_heads, projection_size=config.embed_dim, use_qkv_parallel=True, - qkv_backend=qkv_backend, - softmax_in_single_precision=softmax_in_single_precision, flatten_batch=True, quant_config=quant_config, prefix=add_prefix("attn", prefix), @@ -211,9 +203,7 @@ class DotsVisionTransformer(PreTrainedModel): _num_hidden_layers = config.num_hidden_layers self.blocks = nn.ModuleList( [ - DotsVisionBlock( - config, quant_config, f"blocks.{i}", config.attn_implementation - ) + DotsVisionBlock(config, quant_config, f"blocks.{i}") for i in range(_num_hidden_layers) ] ) diff --git a/python/sglang/srt/models/glm4v.py b/python/sglang/srt/models/glm4v.py index 515018c45..b1ce0cc71 100644 --- a/python/sglang/srt/models/glm4v.py +++ b/python/sglang/srt/models/glm4v.py @@ -104,7 +104,6 @@ class Glm4vVisionBlock(nn.Module): dim: int, intermediate_dim: int, num_heads: int, - attn_implementation: Optional[str] = None, quant_config: Optional[QuantizationConfig] = None, prefix: str = "", num_dummy_heads: int = 0, @@ -114,37 +113,13 @@ class Glm4vVisionBlock(nn.Module): self.norm1 = RMSNorm(dim, eps=rms_norm_eps) self.norm2 = RMSNorm(dim, eps=rms_norm_eps) - if attn_implementation is None: - softmax_in_single_precision = False - qkv_backend = None - flatten_batch = True - elif attn_implementation == "sdpa": - softmax_in_single_precision = False - qkv_backend = "sdpa" - flatten_batch = True - elif attn_implementation == "flash_attention_2": - softmax_in_single_precision = False - qkv_backend = "triton_attn" - flatten_batch = True - elif attn_implementation == "eager": - softmax_in_single_precision = True - qkv_backend = "sdpa" - flatten_batch = True - elif attn_implementation == "flash_attention_3": - softmax_in_single_precision = False - qkv_backend = "fa3" - flatten_batch = True - self.attn = VisionAttention( embed_dim=dim, num_heads=num_heads, projection_size=dim, use_qkv_parallel=True, - rotary_embed="normal", proj_bias=True, - qkv_backend=qkv_backend, - softmax_in_single_precision=softmax_in_single_precision, - flatten_batch=flatten_batch, + flatten_batch=True, quant_config=quant_config, prefix=add_prefix("attn", prefix), num_dummy_heads=num_dummy_heads, diff --git a/python/sglang/srt/models/idefics2.py b/python/sglang/srt/models/idefics2.py index 75922d05c..02f8a2497 100644 --- a/python/sglang/srt/models/idefics2.py +++ b/python/sglang/srt/models/idefics2.py @@ -82,7 +82,6 @@ class Idefics2EncoderLayer(nn.Module): use_qkv_parallel=True, quant_config=quant_config, dropout=config.attention_dropout, - qkv_backend="sdpa", softmax_in_single_precision=True, flatten_batch=False, prefix=add_prefix("self_attn", prefix), diff --git a/python/sglang/srt/models/internvl.py b/python/sglang/srt/models/internvl.py index e19d73f2f..389c2a884 100644 --- a/python/sglang/srt/models/internvl.py +++ b/python/sglang/srt/models/internvl.py @@ -48,7 +48,6 @@ class InternAttention(nn.Module): self.scale = self.head_dim**-0.5 self.attn = VisionAttention( - qkv_backend="fa3", embed_dim=self.embed_dim, num_heads=self.num_heads, projection_size=self.embed_dim, diff --git a/python/sglang/srt/models/mllama.py b/python/sglang/srt/models/mllama.py index 8f89c32f1..5be7cda58 100644 --- a/python/sglang/srt/models/mllama.py +++ b/python/sglang/srt/models/mllama.py @@ -202,9 +202,6 @@ class MllamaVisionEncoderLayer(nn.Module): self.hidden_size, use_qkv_parallel=True, quant_config=quant_config, - dropout=0.0, - qkv_backend="sdpa", - softmax_in_single_precision=False, flatten_batch=False, prefix=add_prefix("self_attn", prefix), ) diff --git a/python/sglang/srt/models/mllama4.py b/python/sglang/srt/models/mllama4.py index c68c394e2..0913f9adf 100644 --- a/python/sglang/srt/models/mllama4.py +++ b/python/sglang/srt/models/mllama4.py @@ -173,9 +173,6 @@ class Llama4VisionEncoderLayer(nn.Module): use_qkv_parallel=True, # vision_model is explicitly ignored in Maverick-17B-128E-Instruct-FP8 quant_config=None, - dropout=0.0, - qkv_backend="sdpa", - softmax_in_single_precision=False, flatten_batch=False, prefix=add_prefix("self_attn", prefix), qkv_bias=True, diff --git a/python/sglang/srt/models/pixtral.py b/python/sglang/srt/models/pixtral.py index 209b40645..249a5ce81 100644 --- a/python/sglang/srt/models/pixtral.py +++ b/python/sglang/srt/models/pixtral.py @@ -106,7 +106,6 @@ class PixtralHFTransformerBlock(nn.Module): quant_config=quant_config, dropout=0.0, use_context_forward=False, - softmax_in_single_precision=False, flatten_batch=False, prefix=f"{prefix}.attention", ) diff --git a/python/sglang/srt/models/qwen2_5_vl.py b/python/sglang/srt/models/qwen2_5_vl.py index 75660a1fe..edcad6683 100644 --- a/python/sglang/srt/models/qwen2_5_vl.py +++ b/python/sglang/srt/models/qwen2_5_vl.py @@ -111,7 +111,6 @@ class Qwen2_5_VisionBlock(nn.Module): num_heads: int, hidden_act="silu", norm_layer: Type[nn.Module] = None, - attn_implementation: Optional[str] = None, quant_config: Optional[QuantizationConfig] = None, prefix: str = "", num_dummy_heads: int = 0, @@ -121,37 +120,13 @@ class Qwen2_5_VisionBlock(nn.Module): self.norm1 = RMSNorm(dim, eps=rms_norm_eps) self.norm2 = RMSNorm(dim, eps=rms_norm_eps) - if attn_implementation is None: - softmax_in_single_precision = False - qkv_backend = None - flatten_batch = True - elif attn_implementation == "sdpa": - softmax_in_single_precision = False - qkv_backend = "sdpa" - flatten_batch = True - elif attn_implementation == "flash_attention_2": - softmax_in_single_precision = False - qkv_backend = "triton_attn" - flatten_batch = True - elif attn_implementation == "eager": - softmax_in_single_precision = True - qkv_backend = "sdpa" - flatten_batch = True - elif attn_implementation == "flash_attention_3": - softmax_in_single_precision = False - qkv_backend = "fa3" - flatten_batch = True - self.attn = VisionAttention( embed_dim=dim, num_heads=num_heads, projection_size=dim, use_qkv_parallel=True, - rotary_embed="normal", proj_bias=True, - qkv_backend=qkv_backend, - softmax_in_single_precision=softmax_in_single_precision, - flatten_batch=flatten_batch, + flatten_batch=True, quant_config=quant_config, prefix=add_prefix("attn", prefix), num_dummy_heads=num_dummy_heads, diff --git a/python/sglang/srt/models/qwen2_vl.py b/python/sglang/srt/models/qwen2_vl.py index 943846210..88d6d5bc9 100644 --- a/python/sglang/srt/models/qwen2_vl.py +++ b/python/sglang/srt/models/qwen2_vl.py @@ -127,7 +127,6 @@ class Qwen2VisionBlock(nn.Module): mlp_ratio: float, act_layer: Type[nn.Module] = QuickGELU, norm_layer: Type[nn.Module] = None, - attn_implementation: Optional[str] = "sdpa", quant_config: Optional[QuantizationConfig] = None, prefix: str = "", ) -> None: @@ -137,23 +136,12 @@ class Qwen2VisionBlock(nn.Module): self.norm1 = norm_layer(dim) self.norm2 = norm_layer(dim) mlp_hidden_dim = int(dim * mlp_ratio) - if attn_implementation == "sdpa": - qkv_backend = "sdpa" - softmax_in_single_precision = False - elif attn_implementation == "flash_attention_2": - qkv_backend = "triton_attn" - softmax_in_single_precision = False - elif attn_implementation == "eager": - qkv_backend = "sdpa" - softmax_in_single_precision = True self.attn = VisionAttention( embed_dim=dim, num_heads=num_heads, projection_size=dim, use_qkv_parallel=True, - qkv_backend=qkv_backend, - softmax_in_single_precision=softmax_in_single_precision, flatten_batch=True, quant_config=quant_config, prefix=add_prefix("attn", prefix), @@ -333,7 +321,6 @@ class Qwen2VisionTransformer(nn.Module): num_heads=num_heads, mlp_ratio=mlp_ratio, norm_layer=norm_layer, - attn_implementation="sdpa", quant_config=quant_config, prefix=add_prefix(f"blocks.{i}", prefix), ) diff --git a/python/sglang/srt/models/qwen3_omni_moe.py b/python/sglang/srt/models/qwen3_omni_moe.py index 805e5d7a2..8663e5ac5 100644 --- a/python/sglang/srt/models/qwen3_omni_moe.py +++ b/python/sglang/srt/models/qwen3_omni_moe.py @@ -61,10 +61,7 @@ class Qwen3OmniMoeAudioEncoderLayer(nn.Module): num_heads=config.encoder_attention_heads, projection_size=embed_dim, use_qkv_parallel=True, - rotary_embed="normal", proj_bias=True, - qkv_backend="fa3", - softmax_in_single_precision=False, flatten_batch=True, quant_config=quant_config, prefix=add_prefix("attn", prefix), diff --git a/python/sglang/srt/models/qwen3_vl.py b/python/sglang/srt/models/qwen3_vl.py index b471829f4..c4d9456bc 100644 --- a/python/sglang/srt/models/qwen3_vl.py +++ b/python/sglang/srt/models/qwen3_vl.py @@ -130,7 +130,6 @@ class Qwen3_VisionBlock(nn.Module): intermediate_dim: int, hidden_act="silu", norm_layer: Optional[Callable[[int], nn.Module]] = None, - attn_implementation: Optional[str] = "sdpa", quant_config: Optional[QuantizationConfig] = None, prefix: str = "", ) -> None: @@ -140,33 +139,13 @@ class Qwen3_VisionBlock(nn.Module): self.norm1 = norm_layer(dim) self.norm2 = norm_layer(dim) - if attn_implementation == "sdpa": - softmax_in_single_precision = False - qkv_backend = "sdpa" - flatten_batch = True - elif attn_implementation == "flash_attention_2": - softmax_in_single_precision = False - qkv_backend = "triton_attn" - flatten_batch = True - elif attn_implementation == "eager": - softmax_in_single_precision = True - qkv_backend = "sdpa" - flatten_batch = True - elif attn_implementation == "flash_attention_3": - softmax_in_single_precision = False - qkv_backend = "fa3" - flatten_batch = True - self.attn = VisionAttention( embed_dim=dim, num_heads=num_heads, projection_size=dim, use_qkv_parallel=True, - rotary_embed="normal", proj_bias=True, - qkv_backend=qkv_backend, - softmax_in_single_precision=softmax_in_single_precision, - flatten_batch=flatten_batch, + flatten_batch=True, quant_config=quant_config, prefix=add_prefix("attn", prefix), ) @@ -283,7 +262,6 @@ class Qwen3VLMoeVisionModel(nn.Module): intermediate_dim=vision_config.intermediate_size, hidden_act=vision_config.hidden_act, norm_layer=norm_layer, - attn_implementation="flash_attention_3", quant_config=quant_config, prefix=add_prefix(f"blocks.{layer_idx}", prefix), ) diff --git a/python/sglang/srt/models/siglip.py b/python/sglang/srt/models/siglip.py index 2a76dc286..34afe07f8 100644 --- a/python/sglang/srt/models/siglip.py +++ b/python/sglang/srt/models/siglip.py @@ -97,7 +97,6 @@ class SiglipEncoderLayer(nn.Module): config: SiglipVisionConfig, act_layer: Type[nn.Module] = QuickGELU, norm_layer: Type[nn.Module] = None, - attn_implementation: Optional[str] = "sdpa", quant_config: Optional[QuantizationConfig] = None, prefix: str = "", ) -> None: @@ -106,22 +105,11 @@ class SiglipEncoderLayer(nn.Module): norm_layer = partial(nn.LayerNorm, eps=config.layer_norm_eps) self.layer_norm1 = norm_layer(config.hidden_size) self.layer_norm2 = norm_layer(config.hidden_size) - if attn_implementation == "sdpa": - qkv_backend = "sdpa" - softmax_in_single_precision = False - elif attn_implementation == "flash_attention_2": - qkv_backend = "triton_attn" - softmax_in_single_precision = False - elif attn_implementation == "eager": - qkv_backend = "sdpa" - softmax_in_single_precision = True self.self_attn = VisionAttention( embed_dim=config.hidden_size, num_heads=config.num_attention_heads, projection_size=config.hidden_size, use_qkv_parallel=True, - qkv_backend=qkv_backend, - softmax_in_single_precision=softmax_in_single_precision, flatten_batch=True, quant_config=quant_config, prefix=add_prefix("self_attn", prefix), @@ -190,7 +178,6 @@ class SiglipEncoder(nn.Module): SiglipEncoderLayer( config=config, norm_layer=norm_layer, - attn_implementation="sdpa", quant_config=quant_config, prefix=add_prefix(f"layers.{layer_idx}", prefix), ) diff --git a/python/sglang/srt/models/step3_vl.py b/python/sglang/srt/models/step3_vl.py index 5a9e74ab6..4474f62d5 100644 --- a/python/sglang/srt/models/step3_vl.py +++ b/python/sglang/srt/models/step3_vl.py @@ -571,7 +571,6 @@ class Step3VisionAttention(nn.Module): self, dim: int, num_heads: int = 16, - qkv_backend="fa3", quant_config=None, prefix: str = "", ) -> None: @@ -593,9 +592,7 @@ class Step3VisionAttention(nn.Module): num_heads=num_heads, projection_size=dim, use_qkv_parallel=True, - rotary_embed="normal", proj_bias=True, - qkv_backend=qkv_backend, quant_config=quant_config, prefix=add_prefix("attn", prefix), )