diff --git a/python/sglang/multimodal_gen/runtime/layers/attention/backends/aiter.py b/python/sglang/multimodal_gen/runtime/layers/attention/backends/aiter.py index efd69a1c6..457299d6e 100644 --- a/python/sglang/multimodal_gen/runtime/layers/attention/backends/aiter.py +++ b/python/sglang/multimodal_gen/runtime/layers/attention/backends/aiter.py @@ -71,16 +71,14 @@ class AITerImpl(AttentionImpl): Performs attention using aiter.flash_attn_func. Args: - query: Query tensor of shape [batch_size, num_heads, seq_len, head_dim] - key: Key tensor of shape [batch_size, num_heads, seq_len, head_dim] - value: Value tensor of shape [batch_size, num_heads, seq_len, head_dim] + query: Query tensor of shape [batch_size, seq_len, num_heads, head_dim] + key: Key tensor of shape [batch_size, seq_len, num_heads, head_dim] + value: Value tensor of shape [batch_size, seq_len, num_heads, head_dim] attn_metadata: Metadata for the attention operation (unused). Returns: - Output tensor of shape [batch_size, num_heads, seq_len, head_dim] + Output tensor of shape [batch_size, seq_len, num_heads, head_dim] """ - # aiter.flash_attn_func expects tensors in [B, H, S, D] layout, - # which is what ring_attn provides. output, _ = aiter.flash_attn_func( query, key, diff --git a/python/sglang/multimodal_gen/runtime/models/dits/flux_2.py b/python/sglang/multimodal_gen/runtime/models/dits/flux_2.py index 35ba8902e..d087c6228 100644 --- a/python/sglang/multimodal_gen/runtime/models/dits/flux_2.py +++ b/python/sglang/multimodal_gen/runtime/models/dits/flux_2.py @@ -865,6 +865,8 @@ class Flux2Transformer2DModel(CachableDiT, OffloadableDiTMixin): _supported_attention_backends = { AttentionBackendEnum.TORCH_SDPA, AttentionBackendEnum.FA, + AttentionBackendEnum.AITER, + AttentionBackendEnum.AITER_SAGE, } def post_load_weights(self) -> None: