Skip to content

[executorch][cuda] Run INT5 decode linears on autotuned QuantizedGemmFamily Triton kernels and delete the last C shim (int5_plain_mm) - #23516

Open
Gasoonjia wants to merge 2 commits into
gh/gasoonjia/198/basefrom
gh/gasoonjia/198/head
Open

Gasoonjia wants to merge 2 commits into
gh/gasoonjia/198/basefrom
gh/gasoonjia/198/head

Conversation

@Gasoonjia

@Gasoonjia Gasoonjia commented Oct 6, 2026 •

Copy link
Copy Markdown
Contributor

Stack from ghstack (oldest at bottom):

Decode-sized (M <= 4) INT5 linears (CudaDp4aPlanarInt5Tensor, GGUF Q5_K) move onto QuantizedGemmFamily as triton::int5_quantized_gemm_m{1,2,3,4}, like INT4/INT6/INT8 below. int5_plain_mm was the last decode C shim, so the custom-op-to-C-shim machinery goes with it.

Kernels, triton/kernels/int5_quantized_gemm.py. W5A8 DP4A, generated with KernelAgent.

  • Weights: ql [N, K/2] low nibbles (even/odd, as INT4); qh [N, K/8] high bits; u in [0, 31]; w = (scale_code * scale_step) * (u - zero_code * zero_point_step).
  • The activation is quantized once per K32 block with the shared quantize_activations_q8. One warp owns an output row. The 5-bit words are rebuilt with a multiply-and-mask high-bit spread, (nibble * 0x00204081) & 0x01010101, exhaustively checked for all 16 nibbles.
  • Each bucket autotunes the shared generic space over the generic kernel, the explicit kernel with the bucket's rows, and the next-larger one, with rows per CTA {1, 2, 4, 8} and stages {1, 2, 3}, pruned by trip count. Under a dynamic M, an explicit kernel branches to the one with exactly M rows. Split-K is the shared rule.
  • Shared helpers are imported, not copied; quantized_gemm_utils.py is unchanged.

Dispatch, quantize_op_dispatch/int5_dispatch.py. It uses the shared quantized_linear and chunked_dequant_linear. Unsupported inputs fall back to dequant + F.linear and never raise. The int5_plain_mm schema and its impls are removed.

C shim removal (fbcode + xplat):

  • runtime/shims/int5_plain_mm.{h,cu,cuh}, its gtest and its benchmark, and their CMake entries;
  • gen_plain_mm_test_vectors.py, whose last user was the INT5 gtest;
  • the entries in runtime/targets.bzl and CMakeLists.txt, and the OSS wheel symbol list;
  • in cuda_backend.py: the INT5 fallback-kernel entries and _get_custom_ops_to_c_shim_options, since no custom op maps to a C shim any more;
  • quantize_op_dispatch/_library.py: the executorch_cuda op namespace has no ops left;
  • test_sort_shim now asserts no custom-op C shim map is set. The INT4 dispatch tests drop their "no int4_plain_mm" assertions.

With this diff, _plain_mm no longer appears under fbcode/executorch, except in the Windows import library aoti_cuda_shims.lib, which is left for the OSS CI.

Op level (A100; full sequence = activation quantization + GEMM (+ reduce); the shim built from the pre-diff sources; each case runs the config the CUDA-graph autotune timing picks from the pruned space, then shim and Triton alternate for 7 rounds and medians are compared). The only Q5_K linear in the muse-glimmer GGUF is its lm_head; two representative shapes are added:

shape (N x K) M=1 M=2 M=3 M=4 M=2 on bucket 4 M=3 on bucket 4
muse lm_head 202048 x 6656 1.153x 2.073x 2.211x 1.254x 1.867x 2.000x
4096 x 6656 1.109x 1.521x 1.603x 1.099x 1.343x 1.312x
6656 x 19968 1.248x 1.904x 1.915x 1.299x 1.653x 1.660x

Geomean 1.487x for static M and 1.620x for dynamic M; no case is slower than the shim. Mean relative difference vs the shim is < 0.0005. lm_head M=1 runs at about 87% of HBM bandwidth.

Model A/B: pending (muse-glimmer, whose lm_head is Q5_K).

Differential Revision: D123707826

[ghstack-poisoned]
@pytorch-bot

pytorch-bot Bot commented Oct 6, 2026 •

Copy link
Copy Markdown

🔗 Helpful Links

🧪 See artifacts and rendered test results at hud.pytorch.org/pr/pytorch/executorch/23516

Note: Links to docs will display an error until the docs builds have been completed.

❌ 1 New Failure, 2 Cancelled Jobs, 1 Unrelated Failure

As of commit b7b909d with merge base 8789aa5 (image):

NEW FAILURE - The following job has failed:

CANCELLED JOBS - The following jobs were cancelled. Please retry:

UNSTABLE - The following job is marked as unstable, possibly due to flakiness on trunk:

This comment was automatically generated by Dr. CI and updates every 15 minutes.

@github-actions

github-actions Bot commented Oct 6, 2026

Copy link
Copy Markdown

This PR needs a release notes: label

If your change should be included in the release notes (i.e. would users of this library care about this change?), please use a label starting with release notes:. This helps us keep track and include your important work in the next release notes.

To add a label, you can comment to pytorchbot, for example
@pytorchbot label "release notes: none"

For more information, see
https://github.com/pytorch/pytorch/wiki/PyTorch-AutoLabel-Bot#why-categorize-for-release-notes-and-how-does-it-work.

@meta-cla meta-cla Bot added the CLA Signed This label is managed by the Facebook bot. Authors need to sign the CLA before a PR can be reviewed. label Oct 6, 2026
[ghstack-poisoned]
Gasoonjia added a commit that referenced this pull request Oct 6, 2026
…Family Triton kernels and delete the last C shim (int5_plain_mm)

Pull Request resolved: #23516

Decode-sized (M <= 4) INT5 linears (`CudaDp4aPlanarInt5Tensor`, GGUF Q5_K) move onto `QuantizedGemmFamily` as `triton::int5_quantized_gemm_m{1,2,3,4}`, like INT4/INT6/INT8 below. `int5_plain_mm` was the last decode C shim, so the custom-op-to-C-shim machinery goes with it.

**Kernels, `triton/kernels/int5_quantized_gemm.py`.** W5A8 DP4A, generated with KernelAgent.
- Weights: ql [N, K/2] low nibbles (even/odd, as INT4); qh [N, K/8] high bits; u in [0, 31]; `w = (scale_code * scale_step) * (u - zero_code * zero_point_step)`.
- The activation is quantized once per K32 block with the shared `quantize_activations_q8`. One warp owns an output row. The 5-bit words are rebuilt with a multiply-and-mask high-bit spread, `(nibble * 0x00204081) & 0x01010101`, exhaustively checked for all 16 nibbles.
- Each bucket autotunes the shared generic space over the generic kernel, the explicit kernel with the bucket's rows, and the next-larger one, with rows per CTA {1, 2, 4, 8} and stages {1, 2, 3}, pruned by trip count. Under a dynamic M, an explicit kernel branches to the one with exactly M rows. Split-K is the shared rule.
- Shared helpers are imported, not copied; `quantized_gemm_utils.py` is unchanged.

**Dispatch, `quantize_op_dispatch/int5_dispatch.py`.** It uses the shared `quantized_linear` and `chunked_dequant_linear`. Unsupported inputs fall back to dequant + `F.linear` and never raise. The `int5_plain_mm` schema and its impls are removed.

**C shim removal (fbcode + xplat):**
- `runtime/shims/int5_plain_mm.{h,cu,cuh}`, its gtest and its benchmark, and their CMake entries;
- `gen_plain_mm_test_vectors.py`, whose last user was the INT5 gtest;
- the entries in `runtime/targets.bzl` and `CMakeLists.txt`, and the OSS wheel symbol list;
- in `cuda_backend.py`: the INT5 fallback-kernel entries and `_get_custom_ops_to_c_shim_options`, since no custom op maps to a C shim any more;
- `quantize_op_dispatch/_library.py`: the `executorch_cuda` op namespace has no ops left;
- `test_sort_shim` now asserts no custom-op C shim map is set. The INT4 dispatch tests drop their "no int4_plain_mm" assertions.

With this diff, `_plain_mm` no longer appears under fbcode/executorch, except in the Windows import library `aoti_cuda_shims.lib`, which is left for the OSS CI.

Op level (A100; full sequence = activation quantization + GEMM (+ reduce); the shim built from the pre-diff sources; each case runs the config the CUDA-graph autotune timing picks from the pruned space, then shim and Triton alternate for 7 rounds and medians are compared). The only Q5_K linear in the muse-glimmer GGUF is its lm_head; two representative shapes are added:

| shape (N x K) | M=1 | M=2 | M=3 | M=4 | M=2 on bucket 4 | M=3 on bucket 4 |
|---|---|---|---|---|---|---|
| muse lm_head 202048 x 6656 | 1.153x | 2.073x | 2.211x | 1.254x | 1.867x | 2.000x |
| 4096 x 6656 | 1.109x | 1.521x | 1.603x | 1.099x | 1.343x | 1.312x |
| 6656 x 19968 | 1.248x | 1.904x | 1.915x | 1.299x | 1.653x | 1.660x |

Geomean 1.487x for static M and 1.620x for dynamic M; no case is slower than the shim. Mean relative difference vs the shim is < 0.0005. lm_head M=1 runs at about 87% of HBM bandwidth.

Model A/B: pending (muse-glimmer, whose lm_head is Q5_K).
ghstack-source-id: 443135252
@exported-using-ghexport

Differential Revision: [D123707826](https://our.internmc.facebook.com/intern/diff/D123707826/)
@Gasoonjia
Gasoonjia deployed to upload-benchmark-results October 7, 2026 00:17 — with GitHub Actions Active

This branch was successfully deployed

2 active deployments
upload-benchmark-results — b7b909d7 Deployed Oct 7, 2026 by Gasoonjia via upload-benchmark-results #20417
cadence — b7b909d7 Deployed Oct 6, 2026 by Gasoonjia via hifi-op-test / hifi4 #32076
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

CLA Signed This label is managed by the Facebook bot. Authors need to sign the CLA before a PR can be reviewed. meta-exported

Projects

None yet

Development

Successfully merging this pull request may close these issues.

1 participant