Skip to content

[executorch][cuda] Run INT8 decode linears on autotuned QuantizedGemmFamily Triton kernels and delete the int8_plain_mm C shim - #23515

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

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

Conversation

@Gasoonjia

@Gasoonjia Gasoonjia commented Oct 6, 2026 •

Copy link
Copy Markdown
Contributor

Stack from ghstack (oldest at bottom):

Decode-sized (M <= 4) weight-only INT8 linears (torchao IntxUnpackedToInt8Tensor, target_dtype int8, no activation quantization) move onto QuantizedGemmFamily as triton::int8_quantized_gemm_m{1,2,3,4}, like INT4 and INT6 below. This removes the int8_plain_mm C shim everywhere.

Kernels, triton/kernels/int8_quantized_gemm.py. W8A8 DP4A, generated with KernelAgent.

  • qdata [N, K] int8 in natural K order; per-group bf16 scale and int8 zero.
  • Activations are quantized per K32 block in natural order. This uses a new natural_order option of the shared quantizer in quantized_gemm_utils.py, the only shared change.
  • Signed weight bytes are biased to unsigned for the shared U8xS8 DP4A helper and corrected in FP32 with (zero + 128) * x_sum.
  • 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.
  • _unsupported_reason rules: bf16 activation; int8 qdata/zero; bf16 scale; group size a power of two >= 32; static K % 256 == 0; consistent static shapes; contiguous; 4-byte-aligned qdata; same device, CUDA or fake; the shared M rule. K that is a multiple of 32 but not of 256 (which the shim served) now takes the dequant fallback.

Dispatch, quantize_op_dispatch/int8_dispatch.py. The type routing is kept: other IntxUnpackedToInt8Tensor configurations use their own dequantize. The weight-only INT8 path uses the shared quantized_linear and chunked_dequant_linear. Unsupported inputs fall back without raising. A keyword bias= is now honored (it was dropped before). The int8_plain_mm schema and its impls are removed.

C shim removal (fbcode + xplat):

  • runtime/shims/int8_plain_mm.{h,cu,cuh};
  • the entries in runtime/targets.bzl and CMakeLists.txt;
  • the fallback-kernel and custom_ops_to_c_shims entries in cuda_backend.py, and the test_sort_shim expectations;
  • the OSS wheel symbol list.

INT8 had no gtest or benchmark.

Op level (A100; full sequence = activation quantization + GEMM (+ reduce); the shim built from the pre-diff sources). No INT8 model is on disk, so these are representative Llama-style decode shapes (N x K): q_proj/o_proj 4096x4096, kv_proj 1024x4096, gate_up 14336x4096, down 4096x14336, lm_head 128256x4096, gs 32, plus gs 128 for gate_up and down. Each is run at M = 1..4 and with the 4-row bucket at a dynamic M = 2, 3, for 48 cases.

  • All 48 cases are faster than the shim, by 1.065x-3.575x. M >= 2 gains are large, because the shim's cost grows linearly with M (e.g. lm_head M=4: 397 us vs 1420 us).
  • The 12 closest cases (< 1.15x) were re-measured independently, alternating shim and Triton for 7 rounds and comparing medians: 1.060x-1.242x.

Model A/B: not run. There is no INT8 model on disk; this diff is validated at op level only.

Differential Revision: D123666321

[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/23515

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

❌ 1 New Failure, 1 Cancelled Job, 1 Unrelated Failure

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

NEW FAILURE - The following job has failed:

CANCELLED JOB - The following job was 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
Gasoonjia deployed to upload-benchmark-results October 7, 2026 00:24 — with GitHub Actions Active

This branch was successfully deployed

2 active deployments
upload-benchmark-results — 7b666d56 Deployed Oct 7, 2026 by Gasoonjia via upload-benchmark-results #20416
cadence — 7b666d56 Deployed Oct 6, 2026 by Gasoonjia via hifi-op-test / hifi4 #32075
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