Conversation
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.
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>
There was a problem hiding this comment.
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.
0a6f0c0 to
6e4f59d
Compare
* 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).
|
Please help resolve the conflicts. And I have some questions:
|
…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).
Resolved conflicts.
|
7defc82 to
df2f3da
Compare
Summary
Optimize
MojoPagedDecodeGQAkernel for TTX NPU backend by exploiting GQA structure to reduce redundant KV cache loads and improve core utilization.Changes
paged_decode_kernelintopaged_decode_vector_kernel, batching multiple Q heads per KV head into BLOCK_M dimension to share KV loads across the GQA grouppaged_decode_cube_kernelleveraging cube cores for GQA-ratio batched decode when BLOCK_N >= 64 and gqa_ratio > 1paged_decode_splitkv_kernelwith TLE-based workspace for long-sequence parallel reduction when cube cores are underutilizedPerformance (bf16, device latency us)
Platform: 910B, CANN 9.0
Note: TLE depends on CANN 9.0 and FlagTree triton 3.2.x branch.
Accuracy (bf16/fp32, bf16: atol=2e-2 rtol=2e-2, fp32: atol=1e-5 rtol=1e-6)