[NPU]Strip padding before FIA kernel for vision encoder padded sequences (#36329)
This commit is contained in:
@@ -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
|
||||
|
||||
|
||||
|
||||
Reference in New Issue
Block a user