diff --git a/python/sglang/srt/layers/attention/vision.py b/python/sglang/srt/layers/attention/vision.py index 591ea5c99..fbc0f4efd 100644 --- a/python/sglang/srt/layers/attention/vision.py +++ b/python/sglang/srt/layers/attention/vision.py @@ -878,6 +878,7 @@ class VisionAscendAttention(nn.Module): else: cu_seqlens = resolve_seqlens(cu_seqlens, bsz, seq_len, device="cpu") seq_len_arg = cu_seqlens[1:].to(torch.int32) + output = torch.empty_like(q) _, num_heads, head_size = q.shape num_kv_heads = k.shape[1] @@ -885,7 +886,42 @@ class VisionAscendAttention(nn.Module): scale_value = softmax_scale if softmax_scale is not None else head_size**-0.5 seq_len_arg = seq_len_arg.tolist() - output = torch_npu.npu_fused_infer_attention_score( + + total_tokens = int(seq_len_arg[-1]) + + # Vision encoders may pad each image's token sequence up to a common + # length (e.g. MiniCPM-V2.6 pads every image to the batch's max patch + # count). FIA requires queryT == last(actual_seq_lengths), so drop the + # padding before the kernel and put the result back into the padded + # layout afterwards. + valid_indices = None + if q.shape[0] != total_tokens: + # Recover each image's real token count from the cumulative seqlens. + per_seq_lens = [seq_len_arg[0] - 0] + [ + seq_len_arg[i] - seq_len_arg[i - 1] for i in range(1, len(seq_len_arg)) + ] + if len(per_seq_lens) == 1: + q = q[:total_tokens] + k = k[:total_tokens] + v = v[:total_tokens] + else: + padded_seq_len = q.shape[0] // len(per_seq_lens) + valid_indices = torch.cat( + [ + torch.arange( + b * padded_seq_len, + b * padded_seq_len + per_seq_lens[b], + dtype=torch.long, + device=q.device, + ) + for b in range(len(per_seq_lens)) + ] + ) + q = q[valid_indices] + k = k[valid_indices] + v = v[valid_indices] + + attn_output = torch_npu.npu_fused_infer_attention_score( query=q, key=k, value=v, @@ -897,6 +933,15 @@ class VisionAscendAttention(nn.Module): sparse_mode=0, input_layout="TND", )[0] + + if valid_indices is not None: + output.index_copy_(0, valid_indices, attn_output) + elif output.shape[0] != total_tokens: + # Single image: write the compact result into the front of the + # pre-allocated padded output, leaving the padding rows as-is. + output[:total_tokens].copy_(attn_output) + else: + output = attn_output return output