Skip to content

[ttx/npu] Optimize paged attention decode GQA with cube and split-KV - #395

Closed
lyujheng wants to merge 344 commits into
XPU-Forces:masterfrom
lyujheng:ttx-optimize-paged-attention-decode-gqa
Closed

lyujheng wants to merge 344 commits into
XPU-Forces:masterfrom
lyujheng:ttx-optimize-paged-attention-decode-gqa

Conversation

@lyujheng

Copy link
Copy Markdown
Contributor

Summary

Optimize MojoPagedDecodeGQA kernel for TTX NPU backend by exploiting GQA structure to reduce redundant KV cache loads and improve core utilization.

Changes

  • Refactor original paged_decode_kernel into paged_decode_vector_kernel, batching multiple Q heads per KV head into BLOCK_M dimension to share KV loads across the GQA group
  • Add paged_decode_cube_kernel leveraging cube cores for GQA-ratio batched decode when BLOCK_N >= 64 and gqa_ratio > 1
  • Add paged_decode_splitkv_kernel with TLE-based workspace for long-sequence parallel reduction when cube cores are underutilized
  • Auto-select kernel strategy at runtime based on gqa_ratio, sequence length, and core availability

Performance (bf16, device latency us)

Platform: 910B, CANN 9.0
Note: TLE depends on CANN 9.0 and FlagTree triton 3.2.x branch.

Shape (B, Q_H, KV_H, D, SeqLen, PageSize) Layout Triton-Ascend Triton-Ascend-Opt TLE Opt Speedup TLE Speedup
(8, 16, 4, 128, 1024, 32) ABAB 158.95 81.10 - 1.96x -
(8, 16, 4, 128, 1024, 32) AABB 158.40 81.10 - 1.95x -
(8, 16, 4, 96, 1024, 128) ABAB 163.22 57.12 - 2.86x -
(8, 16, 4, 96, 1024, 128) AABB 160.60 56.24 - 2.85x -
(8, 8, 1, 128, 8192, 128) ABAB 403.43 168.98 107.21 2.39x 3.76x
(8, 8, 1, 128, 8192, 128) AABB 404.04 168.46 106.60 2.40x 3.79x
(8, 32, 4, 128, 1024, 16) ABAB 443.56 179.74 - 2.47x -
(8, 32, 4, 128, 1024, 16) AABB 447.25 181.00 - 2.47x -
(8, 32, 4, 128, 512, 64) ABAB 107.55 57.18 - 1.88x -
(8, 32, 4, 128, 512, 64) AABB 106.26 55.61 - 1.91x -
(1, 8, 1, 128, 16384, 128) ABAB 402.30 310.33 44.37 1.30x 9.07x
(1, 8, 1, 128, 16384, 128) AABB 403.86 312.22 44.72 1.29x 9.03x
(4, 64, 8, 128, 2048, 128) ABAB 355.81 107.55 - 3.31x -
(4, 64, 8, 128, 2048, 128) AABB 356.42 106.70 - 3.34x -
(32, 8, 1, 128, 4096, 128) ABAB 730.54 202.99 - 3.60x -
(32, 8, 1, 128, 4096, 128) AABB 733.85 203.57 - 3.61x -
(8, 16, 4, 80, 1024, 128) ABAB 180.62 56.34 - 3.21x -
(8, 16, 4, 80, 1024, 128) AABB 179.32 57.24 - 3.13x -
(8, 8, 8, 128, 1024, 128)MHA ABAB 56.12 57.71 - 0.97x -
(8, 8, 8, 128, 1024, 128)MHA AABB 55.58 56.84 - 0.98x -

Accuracy (bf16/fp32, bf16: atol=2e-2 rtol=2e-2, fp32: atol=1e-5 rtol=1e-6)

Shape (B, Q_Heads, KV_Heads, D, SeqLen, PageSize) Result
(8, 16, 4, 128, 1024, 32) ✅ PASS
(8, 16, 4, 96, 1024, 128) ✅ PASS
(8, 8, 1, 128, 8192, 1024) ✅ PASS
(8, 8, 1, 128, 2048, 1024) ✅ PASS
(8, 8, 1, 128, 0, 1024) ✅ PASS
(8, 32, 4, 128, 1024, 16) ✅ PASS
(4, 16, 4, 128, 512, 32) ✅ PASS
(8, 32, 4, 128, 512, 64) ✅ PASS
(1, 8, 1, 128, 16384, 128) ✅ PASS
(2, 8, 1, 128, 8192, 64) ✅ PASS
(8, 8, 8, 128, 1024, 128) ✅ PASS
(4, 4, 4, 128, 2048, 64) ✅ PASS
(4, 64, 8, 128, 2048, 128) ✅ PASS
(2, 32, 1, 128, 4096, 128) ✅ PASS
(32, 8, 1, 128, 4096, 128) ✅ PASS
(8, 16, 4, 80, 1024, 128) ✅ PASS
(8, 16, 4, 80, 1024, 32) ✅ PASS
(1, 16, 4, 128, 256, 128) ✅ PASS
(4, 16, 4, 128, 1024, 128) fp32 ✅ PASS

YYYYimo and others added 30 commits January 27, 2026 18:27
Co-authored-by: lizichong <756066299@qq.com>
…t_tensor_factory_args

feat(mojo_opset/core/operator.py): support torch.empty factory args.
…Forces#129)

* fix: wrong logic in platform.py, add register fallback warning.

* fix: change warning to debug.
* Refine ttx activation kernels and tests

* Refactor normalization api and support patching rmsnorm

* chore: specify backend for ttx tests
添加forward_diff_with的百分比对比功能

.

框架报错人性化提示

添加forward_diff_with的百分比对比功能,并提取函数优化代码

格式化代码

Update operator.py

修改精度校验方法

.

.
…orces#113)

* Optimize causal_conv1d_update_kernel_bdt_fwd ttx kernel

* Remove useless static assert

* Fix accuracy problem

* Fix ut test error
* refine some core api

* Update README

* Add MoEGating torch impl
* Add MoE interface

* fix some
…computed tokens in kv_cache (XPU-Forces#130)

* feat(paged_prefill_attention): refactor paged_prefill_gqa to support computed tokens in kv_cache
* with computed tokens in kvcache, kv should be longer than q
* refactor reference impl and tests to support inequal q_len and kv_len
* refactor ttx impl of paged_prefill_gqa into persistent-kernel style

* fix(typo): fix typo in attention.py

* fix: fix gqa_layout for paged_attention_prefill

* fix: lift dtype of qk before scale

* feat: make seqlens_kv optional for paged_attention_prefill

* feat: add optional mask for paged_prefill_gqa
…f_weight_loading

chore(utils/hf_utils): add weight loading helper functions for hf mod…
* fix: change default gqa layout of Paged Attention

* fix: change paged_decode APIs as well

* ci: fix tests of MojoPagedDecodeGQA

* ci: lease the tolerance in ci of attention
…rces#149)

* fix: change default gqa layout of Paged Attention

* fix: change paged_decode APIs as well

* ci: fix tests of MojoPagedDecodeGQA

* ci: lease the tolerance in ci of attention

* feat: fix paged_decode_attn to support different gqa layout
- Change linear into gemm api.
- Remove MojoLinear because we can use torch.nn.Linear totally.
- Update readme
- Fix hf_utils
* fix: ci network.

* remove some testcase

---------

Co-authored-by: zhangjihang <zhangjihang@bytedance.com>
* add store_lowrank.py

* Fix short sequence boundary issues, add test cases

---------

Co-authored-by: lizichong <756066299@qq.com>
- Hides ttx npu-related impl within ttx.kernels.npu
- Make ttx api independent with specific target
- Format code with ruff
…-Forces#134)

* feat: refine kv_cache, moe impl, rope support varlen & rope_cut.

* fix: torch register.

* fix: moe interface.

* fix: comment.

* fix: comment.
Neuromancer42 and others added 7 commits July 1, 2026 15:15
Co-authored-by: chenyifan.42 <chenyifan.42@bytedance.com>
* [lvzheng/ttx] Optimize _layernorm_fwd_kernel performance

Refactor kernel into two stages: Stage 1 computes mean and variance
in one pass; Stage 2 performs normalization and writes output.

* [KMCompiler] Optimize layernorm kernels via single pass and nomask

* [KMCompiler] Optimize Layernorm by tune NPU LayerNorm forward tiling and skip infer stats stores

* [KMCompiler] Address LayerNorm review feedback

---------

Co-authored-by: lvzheng <lyujheng@gmail.com>
…es#365)

* [KMCompiler] Optimize silu with rowwise nomask kernels

* [KMCompiler] Address SiLU review feedback
* [KMCompiler]opt for rms_norm

* review update
…U-Forces#385)

`embedding_nf4_dequant_impl` used torch.empty to allocate the output
tensor. When an input id was out of vocabulary (>= vocab_size or < 0),
the kernel's token_mask went False and `tl.store` was masked out —
so the corresponding output row was never written. The row silently
kept whatever bytes torch.empty had handed us from the allocator.

On some runners those bytes happened to be near zero (test passed),
on others they weren't (test failed with Max absolute difference of
several thousand at exactly one row). This made
`test_embedding_nf4_dequant_impl` flake in a way that only reproduced
on specific CI hosts.

Switch to torch.zeros so OOB rows deterministically read back as
zero — matching the contract the reference implementation exercises
via `expected = torch.zeros(...); expected[valid_mask] = ...`.

Non-OOB behavior is unchanged: the kernel overwrites the zeros with
the real dequantized values.
…es#364)

* [KMCompiler] Optimize gelu with rowwise nomask kernels

* [KMCompiler] Address GeLU review feedback

* [KMCompiler] Fix GeLU libdevice import fallback
* fix bug of a: tl.constexpr=tl.cdiv(b, c)

* Apply suggestions from code review

Co-authored-by: gemini-code-assist[bot] <176961590+gemini-code-assist[bot]@users.noreply.github.com>

---------

Co-authored-by: zhouronghai <zhouronghai@cambricon.com>
Co-authored-by: gemini-code-assist[bot] <176961590+gemini-code-assist[bot]@users.noreply.github.com>

@gemini-code-assist gemini-code-assist Bot left a comment

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Code Review

This pull request refactors the paged attention decode implementation by introducing specialized vector, cube, and split-KV Triton kernels to optimize performance across different GQA ratios and sequence lengths. The review feedback highlights critical numerical stability improvements, specifically recommending safe division patterns during softmax normalization to prevent division by zero, and handling potential NaN propagation during split-KV merging when subtracting negative infinities.

Important

The consumer version of Gemini Code Assist on GitHub is being sunset. Starting June 18, 2026, new organization installations will be blocked, and all code review activity will officially cease on July 17, 2026.
For more details on the timeline and next steps, please review the Help Documentation.

Comment thread mojo_opset/backends/ttx/kernels/npu/flash_attention.py Outdated
Comment thread mojo_opset/backends/ttx/kernels/npu/a2/flash_attention.py
Comment thread mojo_opset/backends/ttx/kernels/npu/a2/flash_attention.py
Comment thread mojo_opset/backends/ttx/kernels/npu/a2/flash_attention.py
@lyujheng
lyujheng force-pushed the ttx-optimize-paged-attention-decode-gqa branch from 0a6f0c0 to 6e4f59d Compare July 13, 2026 06:58
kevin-hongkai and others added 13 commits July 15, 2026 11:31
* modify compiler_hint path to tl.extra.cann.extension; tl.max add propagate_nan=tl.PropagateNan.ALL;

* swa add enable_ubuf_saving to solve ub overflow

* normalization use casting_mode gemma to avoid F.rms_norm fp32 cast precision

* remove print

* flash attention use need_mask to do if else

* Revert "flash attention use need_mask to do if else"

This reverts commit d705131.

* modify test_over_encoding to switch to triton-ascend

* modify n_gram mask to solve random precision problem

* switch to triton-ascend 3.2.1, use wget temprorarily for checking CI is right

* switch to triton-ascend 3.2.1, use wget temprorarily for checking CI is right

* switch to triton-ascend 3.2.1, use wget temprorarily for checking CI is right

* add print to debug CI

* modify ccec path

* add CI debug

* add CI debug

* add CI debug

* rollback ci and switch to cann8.5.0 image

* add triton-ascend on CI

* fix inder quant para error

* fix perf test case

* modify extract_slice _compute_vision_rope to adapter triton-ascend

* add sync_solver=False to avoid groupgemm perf descend on triton-ascend

* switch to byted-triton-x 3.2.1

* add --index-url to switch to byted-triton-x 3.2.1

* add --index-url to switch to byted-triton-x 3.2.1

* Update mojo_opset/backends/ttx/kernels/npu/over_encoding/n_gram.py

Co-authored-by: gemini-code-assist[bot] <176961590+gemini-code-assist[bot]@users.noreply.github.com>

* rollback to rmsnorm_fwd llama mode, triton-ascend 3.2.1 precision is OK

* functions rms backward precision is not OK, change to gemma

* no need to add propagate_nan to all

* groupgemm: delete dot_pad_only_k hint and tl.multibuffer

* modify ci and pyproject

* use pyproject.toml to pip install

---------

Co-authored-by: gemini-code-assist[bot] <176961590+gemini-code-assist[bot]@users.noreply.github.com>
Co-authored-by: gemini-code-assist[bot] <176961590+gemini-code-assist[bot]@users.noreply.github.com>
[bug fix] fix some bug of 950 npu on master branch
本地测试通过,合入
[bugfix][NPU] update int8 gemm kernel to make autotune can adatper to triton 3.6.0
…st-dot scaling (XPU-Forces#317)

Move k_scale from pre-scaling K to post-dot application to better utilize AIC/AIV core memory bandwidth.
…PU-Forces#423)

All existing kernels under ttx/kernels/npu/ were written for the Ascend
910 (A2) generation. This refactor prepares the tree for adding Ascend
950 (A5) implementations without touching call sites.

- Move all top-level kernel modules into `a2/`; keep `utils.py` shared.
- Add empty `a5/` subpackage as the placeholder for future A5 kernels.
- Add `_dispatch.py` that detects the active SoC generation
  (env `MOJO_ASCEND_ARCH`, else triton `get_current_target().arch`) and
  registers each kernel under its public path `npu.<name>` in sys.modules.
- Fallback semantics: same-named `.py` files in `a5/` and `a2/` are
  wrapped in `_MergedModule` — attribute lookup prefers A5, falls back
  to A2, so A5 can override individual symbols without copying whole
  files. Subpackages use whole-package substitution.
- Public API unchanged: `from ...npu import X` and
  `from ...npu.<submod> import Y` continue to work.
- Detection outcome and source are logged at INFO; unrecognized arch or
  triton detection failure are logged at WARNING (silent DEBUG fallback
  was too easy to miss).
@Neuromancer42

Copy link
Copy Markdown
Collaborator

Please help resolve the conflicts.

And I have some questions:

  1. how is the num_splits and minimal split size decided
  2. would the updated kernel keep deterministic?

zhangjihang-BD and others added 3 commits July 29, 2026 16:14
…PU-Forces#423)

All existing kernels under ttx/kernels/npu/ were written for the Ascend
910 (A2) generation. This refactor prepares the tree for adding Ascend
950 (A5) implementations without touching call sites.

- Move all top-level kernel modules into `a2/`; keep `utils.py` shared.
- Add empty `a5/` subpackage as the placeholder for future A5 kernels.
- Add `_dispatch.py` that detects the active SoC generation
  (env `MOJO_ASCEND_ARCH`, else triton `get_current_target().arch`) and
  registers each kernel under its public path `npu.<name>` in sys.modules.
- Fallback semantics: same-named `.py` files in `a5/` and `a2/` are
  wrapped in `_MergedModule` — attribute lookup prefers A5, falls back
  to A2, so A5 can override individual symbols without copying whole
  files. Subpackages use whole-package substitution.
- Public API unchanged: `from ...npu import X` and
  `from ...npu.<submod> import Y` continue to work.
- Detection outcome and source are logged at INFO; unrecognized arch or
  triton detection failure are logged at WARNING (silent DEBUG fallback
  was too easy to miss).
@lyujheng

Copy link
Copy Markdown
Contributor Author

Please help resolve the conflicts.

And I have some questions:

  1. how is the num_splits and minimal split size decided
  2. would the updated kernel keep deterministic?

Resolved conflicts.

  1. num_splits is the minimum of: idle cores per sequence (total_cores / num_sequences) and sequence splits with minimum size (num_kv_blocks / 4)
  2. minimal split size is the experimental threshold, at least 4 paged blocks per split (each ~64-128 tokens)
  3. given the same input, the output is reproducible: split writes to its own workspace slot with no atomics or race conditions, and merge stage uses TLE global synchronization to ensure all splits complete before reduction
  4. The split-kv path requires TLE (FlagTree Triton 3.5.x).

Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

None yet

Projects

None yet

Development

Successfully merging this pull request may close these issues.