Repository navigation
[Refactor] Replace npu_ring_mla with FIA in MLA prefill - #5704
Conversation
|
👋 Hi! Thank you for contributing to the vLLM Ascend project. The following points will speed up your PR merge:
If CI fails, you can run linting and testing checks locally according Contributing and Testing. |
d9c21a8 to
ce906fa
Compare
There was a problem hiding this comment.
Code Review
This pull request refactors the MLA prefill attention path by replacing npu_ring_mla with the more performant npu_fused_infer_attention_score (FIA) for current token self-attention. The changes are well-structured and improve code clarity by unifying the attention backend and adding explanatory comments. I have one suggestion to further improve performance by replacing a Python loop with a more efficient torch.cumsum operation.
| query_lens_list = prefill_meta.query_lens.tolist() | ||
| actual_seq_lengths_q = [] | ||
| cumsum = 0 | ||
| for qlen in query_lens_list: | ||
| cumsum += qlen | ||
| actual_seq_lengths_q.append(cumsum) |
There was a problem hiding this comment.
The calculation of actual_seq_lengths_q using a Python loop is inefficient, especially for a large number of requests. This can be simplified and made more performant by using torch.cumsum directly on the tensor.
| query_lens_list = prefill_meta.query_lens.tolist() | |
| actual_seq_lengths_q = [] | |
| cumsum = 0 | |
| for qlen in query_lens_list: | |
| cumsum += qlen | |
| actual_seq_lengths_q.append(cumsum) | |
| actual_seq_lengths_q = torch.cumsum(prefill_meta.query_lens, dim=0).tolist() |
fb4026f to
b8dc6b4
Compare
|
How much performance improvement does this change bring? If it's not significant, I suggest waiting until FIA supports chunk prefill before making this modification. Otherwise, this PR won't achieve your goal and will only add unnecessary code complexity. |
be27e3e to
29ddf38
Compare
d1ad8fe to
dd087fd
Compare
ee075a8 to
f3b1f61
Compare
| query_lens = prefill_metadata.query_lens | ||
| if isinstance(query_lens, list): | ||
| query_lens = torch.tensor(query_lens) | ||
| actual_seq_lengths_q = torch.cumsum(query_lens, dim=0).tolist() |
There was a problem hiding this comment.
Why add this? This parameter can be added in AscendMlaMetadataBuilder like decode_metdata.
|
This pull request has conflicts, please resolve those before we can evaluate the pull request. |
ed1768b to
c6e7132
Compare
52ecb32 to
e07ea6d
Compare
|
This pull request has conflicts, please resolve those before we can evaluate the pull request. |
|
@LICO1314 Please rebase your code |
|
This pull request has conflicts, please resolve those before we can evaluate the pull request. |
Use FIA to replace npu_ring_mla in MLA prefill path and keep CP-related behavior aligned after rebasing on upstream/main. Precision test results: DeepSeek GSM8K 96.51, AIME24 93.3. Made-with: Cursor Signed-off-by: lico67373 <918688502@qq.com> Made-with: Cursor
What this PR does / why we need it?
Refactor: Replace npu_ring_mla with FIA in MLA prefill
This PR refactors the MLA (Multi-Layer Attention) prefill implementation by replacing
npu_ring_mlawithnpu_fused_infer_attention_score(FIA) operator, unifying the attention backend with the standard attention implementation.Key changes:
Core prefill refactoring (
mla_v1.py)npu_ring_mlawithnpu_fused_infer_attention_scorein_forward_prefilland_compute_prefill_contextsoftmax_lse_flag=Truefor prefill attentionnpu_attention_updateto merge multiple chunk outputs with LSE (Log-Sum-Exp)attn_maskfromget_final_mla_mask()toget_splitfuse_attn_mask()for FIA compatibilityData type handling
Metadata optimization
actual_seq_lengths_qinAscendMLAPrefillMetadatachunk_actual_seq_lengths_kv_listinChunkedContextMetadatatorch.cumsumoperations from forward pass to metadata building phaseCP compatibility (
mla_cp.py)_ring_mla_mask_builderto getnpu_ring_mla-compatible masks for Context Parallel scenarioschunk_actual_seq_lengths_kv_listfield toCPChunkedContextMetadataWhy we need it:
attention_v1.py)npu_attention_updateprovides native LSE-based output mergingnpu_ring_mlaremoval across the codebaseDoes this PR introduce any user-facing change?
No. This is a pure refactoring with no functional changes - same behavior, unified backend.