refactor: cleanup vision attention related codes (#13228)

Signed-off-by: Xinyuan Tong <xinyuantong.cs@gmail.com>
Co-authored-by: alisonshao <54658187+alisonshao@users.noreply.github.com>
Co-authored-by: Mick <mickjagger19@icloud.com>
Co-authored-by: Baizhou Zhang <sobereddiezhang@gmail.com>
Co-authored-by: Kangyan-Zhou <zky314343421@gmail.com>
This commit is contained in:
Xinyuan Tong
2025-11-16 06:44:40 -08:00
committed by GitHub
parent 8e3663d4e8
commit b1c688fba2
15 changed files with 4 additions and 142 deletions

View File

@@ -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),
)

View File

@@ -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,
)

View File

@@ -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)
]
)

View File

@@ -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,

View File

@@ -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),

View File

@@ -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,

View File

@@ -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),
)

View File

@@ -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,

View File

@@ -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",
)

View File

@@ -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,

View File

@@ -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),
)

View File

@@ -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),

View File

@@ -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),
)

View File

@@ -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),
)

View File

@@ -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),
)