[Performance][DSA-CP] Optimize local token metadata computation with fused Triton kernel and caching - #12193
Conversation
Summary of ChangesHello, I'm Gemini Code Assist1! I'm currently reviewing this pull request and will post my feedback shortly. In the meantime, here's a summary to help you and other reviewers quickly get up to speed! This pull request optimizes the metadata building process for DSA context-parallel mode by introducing caching mechanisms for intermediate results that are identical across attention groups. By reusing these results and fusing multiple metadata operations into a single Triton kernel, the PR significantly lowers the computational overhead per step, leading to improved performance in high-concurrency scenarios. Highlights
New Features🧠 You can now enable Memory (public preview) to help Gemini Code Assist learn from your team's feedback. This makes future code reviews more consistent and personalized to your project's style. Click here to enable Memory in your admin console. Using Gemini Code AssistThe full guide for Gemini Code Assist can be found on our documentation page, here are some quick tips. Invoking Gemini You can request assistance from Gemini at any point by creating a comment using either
Customization To customize the Gemini Code Assist for GitHub experience, repository maintainers can create a configuration file and/or provide a custom code review style guide (such as PEP-8 for Python) by creating and adding files to a Limitations & Feedback Gemini Code Assist may make mistakes. Please leave feedback on any instances where its feedback is incorrect or counterproductive. You can react with 👍 and 👎 on @gemini-code-assist comments. If you're interested in giving your feedback about your experience with Gemini Code Assist for GitHub and other Google products, sign up here. Footnotes
|
|
👋 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. |
There was a problem hiding this comment.
Code Review
Suggested PR Title:
[Attention][Feature] Optimize local token metadata computation with fused Triton kernel and cachingSuggested PR Summary:
### What this PR does / why we need it?
This pull request optimizes the local token metadata computation in context-parallel DSA execution. It introduces a fused Triton kernel `_build_local_metadata_npu_kernel` to reduce kernel launch overhead on NPUs, and implements caching for GPU local metadata, CPU local metadata, and RoPE local slices across kv-cache groups. Additionally, a `pad_to` method is added to `RopeDataProxy` to support the optimized path.
Feedback on the changes:
1. An `AttributeError` will occur at import time if Triton is not installed because `@triton.jit` is evaluated when `triton` is `None`. A fallback decorator `triton_jit` should be defined.
2. In the Triton kernel, multiplying a boolean tensor with an integer tensor `(lql > 0) * (seq_len - offset)` may cause compilation issues on some backends. It is safer to use `tl.where`.
### Does this PR introduce _any_ user-facing change?
No.
### How was this patch tested?
A new equivalence test suite `tests/ut/ops/test_rope_proxy.py` has been added to verify that the optimized path (`pad_to` + slice) is semantically equivalent to the original path (pad positions + gather + slice).ef6aa99 to
0679d26
Compare
|
@pisceskkk PTAL Threre's a performance optimization for dsa_cp with dsv4. |
8f5c1d2 to
b7283e0
Compare
@pisceskkk, sorry to ping again. I have another fix for dsa_cp in the MTP scenario, which modifies the same function and will cause merge conflicts. Could you please take a look at the current performance optimization PR? Thanks! |
pisceskkk
left a comment
There was a problem hiding this comment.
Sorry for later reply. LGTM, leave a little suggestion.
|
This pull request has conflicts, please resolve those before we can evaluate the pull request. |
f506733 to
5a60fd1
Compare
5a60fd1 to
590b634
Compare
…fused Triton kernel and caching Signed-off-by: Csrayz <33659823+Csrayz@users.noreply.github.com>
Signed-off-by: Csrayz <33659823+Csrayz@users.noreply.github.com>
68f99a1 to
15409ed
Compare
|
@kunpengW-code, linfeng-yuan is currently unreachable. This PR's CI is all green and pisceskkk has already approved. Could you take a look and approve/merge when you get a chance? |
…fused Triton kernel and caching (vllm-project#12193) ### What this PR does / why we need it? In DSA context-parallel mode, `AscendDSACPMetadataBuilder.build()` runs once per attention group within a step. Under DeepSeek-V4's dsv4 attention, there can be 8 build() per step, and each build independently recomputes metadata from scratch. Across all groups, the following data are identical: - `query_start_loc`, `seq_lens`, `num_reqs`, `num_input_tokens` - Clamp + cumsum + offset + mask results - RoPE local slices - `start_pos` This PR caches these intermediate results so that only the first group computes them; subsequent groups reuse the cached values. Additionally, the fused `_build_local_metadata_npu_kernel` is introduced to reduce kernel launch overhead. The metadata building phase dropped from approximately 23ms to 10ms per step (measured on DeepSeek-V4 Flash, mtp=1). before (main) <img width="1618" height="485" alt="c5fa490ea22b40dbaafa4c2dab1f760d_IMAGE_1784202330638" src="https://github.com/user-attachments/assets/f7145c7f-f78a-4fd0-b367-c45ef95d66d3" /> after (this pr) <img width="889" height="479" alt="7281c93267ea4205b34526f9fb6ebd3a_IMAGE_1784202321384" src="https://github.com/user-attachments/assets/eeb12a34-17bd-4be7-8b8b-8f94c6c2ee1b" /> **Changes summary** - New `_ensure_gpu_local_metadata` method that computes GPU local metadata once and caches the result in `common_ratio_to_sas_metadata["_gpu_local"]`. - Cross-group caching for CPU local metadata (`"_cpu_local"`) and RoPE local slices (`"_rope_local"`). - New Triton kernel `_build_local_metadata_npu_kernel` that fuses clamp, cumsum, offset, mask, and start_pos into a single kernel launch. - Refactored `build_req_metadata` to use the cached results and inlined RoPE slicing. - Cleaned up `_build_local_token_metadata` by removing unused parameters (`input_positions`, `use_cache`, `seq_lens_q`). ### Does this PR introduce *any* user-facing change? No. Behavior is identical. ### How was this patch tested? - Added a unit test for `RopeDataProxy.pad_to`. - GMS8K evaluation results passed. ``` before (main) +---------+-----------+----------+----------+-------+---------+---------+ | Model | Dataset | Metric | Subset | Num | Score | Cat.0 | +=========+===========+==========+==========+=======+=========+=========+ | dsv4-f | gsm8k | mean_acc | main | 1319 | 0.9682 | default | +---------+-----------+----------+----------+-------+---------+---------+ ``` aftet (this pr) ``` +---------+-----------+----------+----------+-------+---------+---------+ | Model | Dataset | Metric | Subset | Num | Score | Cat.0 | +=========+===========+==========+==========+=======+=========+=========+ | dsv4-f | gsm8k | mean_acc | main | 1319 | 0.9712 | default | +---------+-----------+----------+----------+-------+---------+---------+ ``` - vLLM version: v0.25.1 - vLLM main: vllm-project/vllm@54503ec --------- Signed-off-by: Csrayz <33659823+Csrayz@users.noreply.github.com>
…fused Triton kernel and caching (vllm-project#12193) ### What this PR does / why we need it? In DSA context-parallel mode, `AscendDSACPMetadataBuilder.build()` runs once per attention group within a step. Under DeepSeek-V4's dsv4 attention, there can be 8 build() per step, and each build independently recomputes metadata from scratch. Across all groups, the following data are identical: - `query_start_loc`, `seq_lens`, `num_reqs`, `num_input_tokens` - Clamp + cumsum + offset + mask results - RoPE local slices - `start_pos` This PR caches these intermediate results so that only the first group computes them; subsequent groups reuse the cached values. Additionally, the fused `_build_local_metadata_npu_kernel` is introduced to reduce kernel launch overhead. The metadata building phase dropped from approximately 23ms to 10ms per step (measured on DeepSeek-V4 Flash, mtp=1). before (main) <img width="1618" height="485" alt="c5fa490ea22b40dbaafa4c2dab1f760d_IMAGE_1784202330638" src="https://github.com/user-attachments/assets/f7145c7f-f78a-4fd0-b367-c45ef95d66d3" /> after (this pr) <img width="889" height="479" alt="7281c93267ea4205b34526f9fb6ebd3a_IMAGE_1784202321384" src="https://github.com/user-attachments/assets/eeb12a34-17bd-4be7-8b8b-8f94c6c2ee1b" /> **Changes summary** - New `_ensure_gpu_local_metadata` method that computes GPU local metadata once and caches the result in `common_ratio_to_sas_metadata["_gpu_local"]`. - Cross-group caching for CPU local metadata (`"_cpu_local"`) and RoPE local slices (`"_rope_local"`). - New Triton kernel `_build_local_metadata_npu_kernel` that fuses clamp, cumsum, offset, mask, and start_pos into a single kernel launch. - Refactored `build_req_metadata` to use the cached results and inlined RoPE slicing. - Cleaned up `_build_local_token_metadata` by removing unused parameters (`input_positions`, `use_cache`, `seq_lens_q`). ### Does this PR introduce *any* user-facing change? No. Behavior is identical. ### How was this patch tested? - Added a unit test for `RopeDataProxy.pad_to`. - GMS8K evaluation results passed. ``` before (main) +---------+-----------+----------+----------+-------+---------+---------+ | Model | Dataset | Metric | Subset | Num | Score | Cat.0 | +=========+===========+==========+==========+=======+=========+=========+ | dsv4-f | gsm8k | mean_acc | main | 1319 | 0.9682 | default | +---------+-----------+----------+----------+-------+---------+---------+ ``` aftet (this pr) ``` +---------+-----------+----------+----------+-------+---------+---------+ | Model | Dataset | Metric | Subset | Num | Score | Cat.0 | +=========+===========+==========+==========+=======+=========+=========+ | dsv4-f | gsm8k | mean_acc | main | 1319 | 0.9712 | default | +---------+-----------+----------+----------+-------+---------+---------+ ``` - vLLM version: v0.25.1 - vLLM main: vllm-project/vllm@54503ec --------- Signed-off-by: Csrayz <33659823+Csrayz@users.noreply.github.com>
### What this PR does / why we need it? This PR enables MTP=3 speculative decoding in the DSA‑CP path for both eager and graph modes. It introduces full ACL graph support for the MTP draft steps, including pre‑allocated metadata buffers that allow deterministic graph capture and replay when enable_dsa_cp is true. ### Changes summary #### MTP > 1 support in DSA‑CP ~Previously build_for_drafting only handled MTP=1 and lacked the CPU‑side sequence length path, which caused tensor dimension errors for any MTP > 1.~ Supported by #13249. This PR updates build_for_drafting to populate CPU sequence lengths, so that all downstream metadata builders can correctly process multiple draft steps. #### Full ACL graph support for MTP draft steps ACL graph capture and replay require stable tensor addresses across invocations. To guarantee this, the PR pre‑allocates per‑step buffers for all draft‑step metadata during builder initialization, and pads seq_lens_cpu to match the graph‑dispatched batch size, ensuring deterministic addresses throughout replay. In addition, a static `update_graph_params` method was added, so that the graph dispatch can correctly invoke the DSA-CP attention backend during graph replay. ~#### Depends on PR #12193.~ ~#12193 also fixes metadata mismatch issues in the DSA-CP builder, and this PR is directly based on the refactored code. Cherry-picking the commits onto main without #12193 causes runtime errors under concurrent requests.~ ~Only the top commits (after Commits on Jul 21, 2026) belong to this PR. Please review by focusing on the top commits. Once it is merged, I'll rebase onto main quickly.`~ ### Does this PR introduce _any_ user-facing change? No. ### How was this patch tested? - Added an E2E accuracy test in `tests/e2e/pull_request/four_card/test_deepseek_v4.py` with exact `expected_token_ids` assertion to guard against regressions in the DSA-CP + MTP=3 + full graph path. - GMS8K accuracy evaluation for MTP+eager and MTP+graph is attached below. ``` mtp3+eager +---------+-----------+----------+----------+-------+---------+---------+ | Model | Dataset | Metric | Subset | Num | Score | Cat.0 | +=========+===========+==========+==========+=======+=========+=========+ | dsv4-f | gsm8k | mean_acc | main | 1319 | 0.9666 | default | +---------+-----------+----------+----------+-------+---------+---------+ mtp3+full graph +---------+-----------+----------+----------+-------+---------+---------+ | Model | Dataset | Metric | Subset | Num | Score | Cat.0 | +=========+===========+==========+==========+=======+=========+=========+ | dsv4-f | gsm8k | mean_acc | main | 1319 | 0.9659 | default | +---------+-----------+----------+----------+-------+---------+---------+ ``` - vLLM main: vllm-project/vllm@e6bfe03 --------- Signed-off-by: frankie <wangyongsheng686@gmail.com> Co-authored-by: frankie <wangyongsheng686@gmail.com>
### What this PR does / why we need it? This PR enables MTP=3 speculative decoding in the DSA‑CP path for both eager and graph modes. It introduces full ACL graph support for the MTP draft steps, including pre‑allocated metadata buffers that allow deterministic graph capture and replay when enable_dsa_cp is true. ### Changes summary #### MTP > 1 support in DSA‑CP ~Previously build_for_drafting only handled MTP=1 and lacked the CPU‑side sequence length path, which caused tensor dimension errors for any MTP > 1.~ Supported by vllm-project#13249. This PR updates build_for_drafting to populate CPU sequence lengths, so that all downstream metadata builders can correctly process multiple draft steps. #### Full ACL graph support for MTP draft steps ACL graph capture and replay require stable tensor addresses across invocations. To guarantee this, the PR pre‑allocates per‑step buffers for all draft‑step metadata during builder initialization, and pads seq_lens_cpu to match the graph‑dispatched batch size, ensuring deterministic addresses throughout replay. In addition, a static `update_graph_params` method was added, so that the graph dispatch can correctly invoke the DSA-CP attention backend during graph replay. ~#### Depends on PR vllm-project#12193.~ ~vllm-project#12193 also fixes metadata mismatch issues in the DSA-CP builder, and this PR is directly based on the refactored code. Cherry-picking the commits onto main without vllm-project#12193 causes runtime errors under concurrent requests.~ ~Only the top commits (after Commits on Jul 21, 2026) belong to this PR. Please review by focusing on the top commits. Once it is merged, I'll rebase onto main quickly.`~ ### Does this PR introduce _any_ user-facing change? No. ### How was this patch tested? - Added an E2E accuracy test in `tests/e2e/pull_request/four_card/test_deepseek_v4.py` with exact `expected_token_ids` assertion to guard against regressions in the DSA-CP + MTP=3 + full graph path. - GMS8K accuracy evaluation for MTP+eager and MTP+graph is attached below. ``` mtp3+eager +---------+-----------+----------+----------+-------+---------+---------+ | Model | Dataset | Metric | Subset | Num | Score | Cat.0 | +=========+===========+==========+==========+=======+=========+=========+ | dsv4-f | gsm8k | mean_acc | main | 1319 | 0.9666 | default | +---------+-----------+----------+----------+-------+---------+---------+ mtp3+full graph +---------+-----------+----------+----------+-------+---------+---------+ | Model | Dataset | Metric | Subset | Num | Score | Cat.0 | +=========+===========+==========+==========+=======+=========+=========+ | dsv4-f | gsm8k | mean_acc | main | 1319 | 0.9659 | default | +---------+-----------+----------+----------+-------+---------+---------+ ``` - vLLM main: vllm-project/vllm@e6bfe03 --------- Signed-off-by: frankie <wangyongsheng686@gmail.com> Co-authored-by: frankie <wangyongsheng686@gmail.com>
…fused Triton kernel and caching (vllm-project#12193) ### What this PR does / why we need it? In DSA context-parallel mode, `AscendDSACPMetadataBuilder.build()` runs once per attention group within a step. Under DeepSeek-V4's dsv4 attention, there can be 8 build() per step, and each build independently recomputes metadata from scratch. Across all groups, the following data are identical: - `query_start_loc`, `seq_lens`, `num_reqs`, `num_input_tokens` - Clamp + cumsum + offset + mask results - RoPE local slices - `start_pos` This PR caches these intermediate results so that only the first group computes them; subsequent groups reuse the cached values. Additionally, the fused `_build_local_metadata_npu_kernel` is introduced to reduce kernel launch overhead. The metadata building phase dropped from approximately 23ms to 10ms per step (measured on DeepSeek-V4 Flash, mtp=1). before (main) <img width="1618" height="485" alt="c5fa490ea22b40dbaafa4c2dab1f760d_IMAGE_1784202330638" src="https://github.com/user-attachments/assets/f7145c7f-f78a-4fd0-b367-c45ef95d66d3" /> after (this pr) <img width="889" height="479" alt="7281c93267ea4205b34526f9fb6ebd3a_IMAGE_1784202321384" src="https://github.com/user-attachments/assets/eeb12a34-17bd-4be7-8b8b-8f94c6c2ee1b" /> **Changes summary** - New `_ensure_gpu_local_metadata` method that computes GPU local metadata once and caches the result in `common_ratio_to_sas_metadata["_gpu_local"]`. - Cross-group caching for CPU local metadata (`"_cpu_local"`) and RoPE local slices (`"_rope_local"`). - New Triton kernel `_build_local_metadata_npu_kernel` that fuses clamp, cumsum, offset, mask, and start_pos into a single kernel launch. - Refactored `build_req_metadata` to use the cached results and inlined RoPE slicing. - Cleaned up `_build_local_token_metadata` by removing unused parameters (`input_positions`, `use_cache`, `seq_lens_q`). ### Does this PR introduce *any* user-facing change? No. Behavior is identical. ### How was this patch tested? - Added a unit test for `RopeDataProxy.pad_to`. - GMS8K evaluation results passed. ``` before (main) +---------+-----------+----------+----------+-------+---------+---------+ | Model | Dataset | Metric | Subset | Num | Score | Cat.0 | +=========+===========+==========+==========+=======+=========+=========+ | dsv4-f | gsm8k | mean_acc | main | 1319 | 0.9682 | default | +---------+-----------+----------+----------+-------+---------+---------+ ``` aftet (this pr) ``` +---------+-----------+----------+----------+-------+---------+---------+ | Model | Dataset | Metric | Subset | Num | Score | Cat.0 | +=========+===========+==========+==========+=======+=========+=========+ | dsv4-f | gsm8k | mean_acc | main | 1319 | 0.9712 | default | +---------+-----------+----------+----------+-------+---------+---------+ ``` - vLLM version: v0.25.1 - vLLM main: vllm-project/vllm@54503ec --------- Signed-off-by: Csrayz <33659823+Csrayz@users.noreply.github.com>
### What this PR does / why we need it? This PR enables MTP=3 speculative decoding in the DSA‑CP path for both eager and graph modes. It introduces full ACL graph support for the MTP draft steps, including pre‑allocated metadata buffers that allow deterministic graph capture and replay when enable_dsa_cp is true. ### Changes summary #### MTP > 1 support in DSA‑CP ~Previously build_for_drafting only handled MTP=1 and lacked the CPU‑side sequence length path, which caused tensor dimension errors for any MTP > 1.~ Supported by vllm-project#13249. This PR updates build_for_drafting to populate CPU sequence lengths, so that all downstream metadata builders can correctly process multiple draft steps. #### Full ACL graph support for MTP draft steps ACL graph capture and replay require stable tensor addresses across invocations. To guarantee this, the PR pre‑allocates per‑step buffers for all draft‑step metadata during builder initialization, and pads seq_lens_cpu to match the graph‑dispatched batch size, ensuring deterministic addresses throughout replay. In addition, a static `update_graph_params` method was added, so that the graph dispatch can correctly invoke the DSA-CP attention backend during graph replay. ~#### Depends on PR vllm-project#12193.~ ~vllm-project#12193 also fixes metadata mismatch issues in the DSA-CP builder, and this PR is directly based on the refactored code. Cherry-picking the commits onto main without vllm-project#12193 causes runtime errors under concurrent requests.~ ~Only the top commits (after Commits on Jul 21, 2026) belong to this PR. Please review by focusing on the top commits. Once it is merged, I'll rebase onto main quickly.`~ ### Does this PR introduce _any_ user-facing change? No. ### How was this patch tested? - Added an E2E accuracy test in `tests/e2e/pull_request/four_card/test_deepseek_v4.py` with exact `expected_token_ids` assertion to guard against regressions in the DSA-CP + MTP=3 + full graph path. - GMS8K accuracy evaluation for MTP+eager and MTP+graph is attached below. ``` mtp3+eager +---------+-----------+----------+----------+-------+---------+---------+ | Model | Dataset | Metric | Subset | Num | Score | Cat.0 | +=========+===========+==========+==========+=======+=========+=========+ | dsv4-f | gsm8k | mean_acc | main | 1319 | 0.9666 | default | +---------+-----------+----------+----------+-------+---------+---------+ mtp3+full graph +---------+-----------+----------+----------+-------+---------+---------+ | Model | Dataset | Metric | Subset | Num | Score | Cat.0 | +=========+===========+==========+==========+=======+=========+=========+ | dsv4-f | gsm8k | mean_acc | main | 1319 | 0.9659 | default | +---------+-----------+----------+----------+-------+---------+---------+ ``` - vLLM main: vllm-project/vllm@e6bfe03 --------- Signed-off-by: frankie <wangyongsheng686@gmail.com> Co-authored-by: frankie <wangyongsheng686@gmail.com>
### What this PR does / why we need it? This PR enables MTP=3 speculative decoding in the DSA‑CP path for both eager and graph modes. It introduces full ACL graph support for the MTP draft steps, including pre‑allocated metadata buffers that allow deterministic graph capture and replay when enable_dsa_cp is true. ### Changes summary #### MTP > 1 support in DSA‑CP ~Previously build_for_drafting only handled MTP=1 and lacked the CPU‑side sequence length path, which caused tensor dimension errors for any MTP > 1.~ Supported by vllm-project#13249. This PR updates build_for_drafting to populate CPU sequence lengths, so that all downstream metadata builders can correctly process multiple draft steps. #### Full ACL graph support for MTP draft steps ACL graph capture and replay require stable tensor addresses across invocations. To guarantee this, the PR pre‑allocates per‑step buffers for all draft‑step metadata during builder initialization, and pads seq_lens_cpu to match the graph‑dispatched batch size, ensuring deterministic addresses throughout replay. In addition, a static `update_graph_params` method was added, so that the graph dispatch can correctly invoke the DSA-CP attention backend during graph replay. ~#### Depends on PR vllm-project#12193.~ ~vllm-project#12193 also fixes metadata mismatch issues in the DSA-CP builder, and this PR is directly based on the refactored code. Cherry-picking the commits onto main without vllm-project#12193 causes runtime errors under concurrent requests.~ ~Only the top commits (after Commits on Jul 21, 2026) belong to this PR. Please review by focusing on the top commits. Once it is merged, I'll rebase onto main quickly.`~ ### Does this PR introduce _any_ user-facing change? No. ### How was this patch tested? - Added an E2E accuracy test in `tests/e2e/pull_request/four_card/test_deepseek_v4.py` with exact `expected_token_ids` assertion to guard against regressions in the DSA-CP + MTP=3 + full graph path. - GMS8K accuracy evaluation for MTP+eager and MTP+graph is attached below. ``` mtp3+eager +---------+-----------+----------+----------+-------+---------+---------+ | Model | Dataset | Metric | Subset | Num | Score | Cat.0 | +=========+===========+==========+==========+=======+=========+=========+ | dsv4-f | gsm8k | mean_acc | main | 1319 | 0.9666 | default | +---------+-----------+----------+----------+-------+---------+---------+ mtp3+full graph +---------+-----------+----------+----------+-------+---------+---------+ | Model | Dataset | Metric | Subset | Num | Score | Cat.0 | +=========+===========+==========+==========+=======+=========+=========+ | dsv4-f | gsm8k | mean_acc | main | 1319 | 0.9659 | default | +---------+-----------+----------+----------+-------+---------+---------+ ``` - vLLM main: vllm-project/vllm@e6bfe03 --------- Signed-off-by: frankie <wangyongsheng686@gmail.com> Co-authored-by: frankie <wangyongsheng686@gmail.com> Signed-off-by: like-0517 <ithwlike@126.com>
What this PR does / why we need it?
In DSA context-parallel mode,
AscendDSACPMetadataBuilder.build()runs once per attention group within a step. Under DeepSeek-V4's dsv4 attention, there can be 8 build() per step, and each build independently recomputes metadata from scratch.Across all groups, the following data are identical:
query_start_loc,seq_lens,num_reqs,num_input_tokensstart_posThis PR caches these intermediate results so that only the first group computes them; subsequent groups reuse the cached values. Additionally, the fused
_build_local_metadata_npu_kernelis introduced to reduce kernel launch overhead.The metadata building phase dropped from approximately 23ms to 10ms per step (measured on DeepSeek-V4 Flash, mtp=1).
before (main)
after (this pr)
Changes summary
_ensure_gpu_local_metadatamethod that computes GPU local metadata once and caches the result incommon_ratio_to_sas_metadata["_gpu_local"]."_cpu_local") and RoPE local slices ("_rope_local")._build_local_metadata_npu_kernelthat fuses clamp, cumsum, offset, mask, and start_pos into a single kernel launch.build_req_metadatato use the cached results and inlined RoPE slicing._build_local_token_metadataby removing unused parameters (input_positions,use_cache,seq_lens_q).Does this PR introduce any user-facing change?
No. Behavior is identical.
How was this patch tested?
RopeDataProxy.pad_to.aftet (this pr)