Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
1 change: 0 additions & 1 deletion .ci/scripts/wheel/test_shared_libraries.py
Original file line number Diff line number Diff line change
Expand Up @@ -96,7 +96,6 @@
# kernel shim is listed rather than a sample, because a partially built library is exactly the
# failure this row exists to catch.
"aoti_torch_cuda__weight_int4pack_mm",
"aoti_torch_cuda_int4_plain_mm",
"aoti_torch_cuda_int5_plain_mm",
"aoti_torch_cuda_int6_plain_mm",
"aoti_torch_cuda_int8_plain_mm",
Expand Down
3 changes: 3 additions & 0 deletions backends/cuda/BUCK
Original file line number Diff line number Diff line change
Expand Up @@ -67,6 +67,7 @@ fbcode_target(
":coalesced_int4_tensor",
":dp4a_planar_int5_tensor",
":dp4a_planar_int6_tensor",
":triton_kernels",
"//caffe2:torch",
"//pytorch/ao:torchao",
],
Expand Down Expand Up @@ -165,8 +166,10 @@ fbcode_target(
srcs = [
"triton/kernels/__init__.py",
"triton/kernels/fused_moe.py",
"triton/kernels/int4_quantized_gemm.py",
"triton/kernels/int4_matmul.py",
"triton/kernels/quantized_gemm_family.py",
"triton/kernels/quantized_gemm_utils.py",
"triton/kernels/sdpa.py",
"triton/kernels/topk.py",
"triton/kernels/tq4_sdpa.py",
Expand Down
1 change: 0 additions & 1 deletion backends/cuda/CMakeLists.txt
Original file line number Diff line number Diff line change
Expand Up @@ -250,7 +250,6 @@ if(NOT EXECUTORCH_BUILD_ROCM AND CMAKE_CUDA_COMPILER)
APPEND
_aoti_cuda_shim_sources
runtime/shims/int4mm.cu
runtime/shims/int4_plain_mm.cu
runtime/shims/int5_plain_mm.cu
runtime/shims/int6_plain_mm.cu
runtime/shims/int8_plain_mm.cu
Expand Down
4 changes: 2 additions & 2 deletions backends/cuda/coalesced_int4_tensor.py
Original file line number Diff line number Diff line change
Expand Up @@ -43,8 +43,8 @@
Bits-per-weight: 4.0 (qdata) + 8/32 (scale codes) + 16/256 (fp16 scale step) +
8/32 (uint8 zero codes) + 16/256 (fp16 zero step) = 4.625 bpw.

The coalesced [N, n_groups] layout is exactly what the W4A8 dp4a matvec kernel
(``executorch_cuda::int4_plain_mm`` / ``int4_plain_mm.cuh``) reads row-for-row
The coalesced [N, n_groups] layout is exactly what the W4A8 DP4A decode kernels
(``triton::int4_quantized_gemm_m{M}``, triton/kernels/int4_quantized_gemm.py) read row-for-row
with qdata, so the exported decode graph carries no per-step transpose. The
is owned by :meth:`from_exportable_int4_tensor` so
it is baked into the serialized weight constant once at pack time.
Expand Down
8 changes: 0 additions & 8 deletions backends/cuda/cuda_backend.py
Original file line number Diff line number Diff line change
Expand Up @@ -771,8 +771,6 @@ def get_supported_fallback_kernels(cls) -> Dict[str, Any]:
"at::_ops::sort_stable::call": None,
"aoti_torch_cuda_sort_stable": None,
"aoti_torch_cuda_randint_low_out": None,
"executorch_cuda::int4_plain_mm": None,
"aoti_torch_cuda_int4_plain_mm": None,
"executorch_cuda::int5_plain_mm": None,
"aoti_torch_cuda_int5_plain_mm": None,
"executorch_cuda::int6_plain_mm": None,
Expand All @@ -788,12 +786,6 @@ def _get_custom_ops_to_c_shim_options() -> Dict[str, Any]:
try:
return {
"aot_inductor.custom_ops_to_c_shims": {
torch.ops.executorch_cuda.int4_plain_mm.default: [
"AOTITorchError aoti_torch_cuda_int4_plain_mm("
"AtenTensorHandle, AtenTensorHandle, AtenTensorHandle, "
"AtenTensorHandle, AtenTensorHandle, AtenTensorHandle, "
"int64_t, AtenTensorHandle*)"
],
torch.ops.executorch_cuda.int5_plain_mm.default: [
"AOTITorchError aoti_torch_cuda_int5_plain_mm("
"AtenTensorHandle, AtenTensorHandle, AtenTensorHandle, "
Expand Down
2 changes: 1 addition & 1 deletion backends/cuda/quantize_op_dispatch/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -10,7 +10,7 @@
weight tensors so that torch.export traces through ExecuTorch's custom ops and
dequant logic instead of torchao's defaults. It registers:

* INT4 (``CudaCoalescedInt4Tensor``) → ``executorch_cuda::int4_plain_mm``
* INT4 (``CudaCoalescedInt4Tensor``) → ``triton::int4_quantized_gemm_m{1,2,3,4}``
* INT5 (``CudaDp4aPlanarInt5Tensor``) → ``executorch_cuda::int5_plain_mm``
* INT6 (``CudaDp4aPlanarInt6Tensor``) → ``executorch_cuda::int6_plain_mm``
* INT8 (``IntxUnpackedToInt8Tensor``) → ``executorch_cuda::int8_plain_mm``
Expand Down
164 changes: 52 additions & 112 deletions backends/cuda/quantize_op_dispatch/int4_dispatch.py
Original file line number Diff line number Diff line change
Expand Up @@ -8,23 +8,27 @@

This module registers an F.linear dispatch on ``CudaCoalescedInt4Tensor`` (an
ExecuTorch-internal subclass, see ``coalesced_int4_tensor.py``) so that
torch.export traces through our custom op and dequant logic. Routing is by
torch.export traces through our Triton ops and dequant logic. Routing is by
*type*: stock torchao ``Int4Tensor`` weights are left untouched and keep using
torchao's default (mslk/tinygemm) path. The code here executes during eager
inference and during AOTI export tracing — it does NOT run at .pte runtime.

At .pte runtime, the captured graph is executed by the AOTI-generated .so:
- The custom op ``executorch_cuda::int4_plain_mm`` maps to a C shim that
runs the W4A8 dp4a matvec kernel (backends/cuda/runtime/shims/).
- ``triton::int4_quantized_gemm_m{M}`` is a Triton W4A8 DP4A kernel compiled
into it (see triton/kernels/int4_quantized_gemm.py).
- The inline dequant + F.linear is compiled by inductor into fused Triton
dequant + cuBLAS matmul kernels.
dequant + matmul kernels.

Dispatch strategy (determines what gets captured in the export graph):
Decode (M<=4): Custom op ``executorch_cuda::int4_plain_mm``
Prefill (M>4): Inline dequant + F.linear (standard PyTorch ops)
Dispatch (``_gemm_family_dispatch.quantized_linear``): when a Triton kernel can
run, the smallest bucket of ``INT4_QUANTIZED_GEMM`` that supports the inputs
(static M <= 4 takes its own bucket; a dynamic M provably within [1, 4], e.g. a
speculative block length in [2, 4], takes the smallest bucket that holds it).
Everything else (prefill, an unbounded dynamic M, a group size other than 32,
K not a multiple of 256, other dtypes, CPU eager) uses inline dequant +
F.linear, never an error.

Importing the parent ``quantize_op_dispatch`` package registers this dispatch
override (along with the INT8 one) before using nn.Linear with
override (along with the other formats) before using nn.Linear with
CudaCoalescedInt4Tensor weights::

import executorch.backends.cuda.quantize_op_dispatch # noqa: F401
Expand All @@ -33,51 +37,12 @@
import torch
import torch.nn.functional as F
from executorch.backends.cuda.coalesced_int4_tensor import CudaCoalescedInt4Tensor
from executorch.backends.cuda.quantize_op_dispatch._library import lib as _lib
from torch.library import impl

# ---------------------------------------------------------------------------
# Custom op for decode (M=1): dp4a matvec in C shim, dequant+F.linear in eager
# ---------------------------------------------------------------------------

_lib.define(
"int4_plain_mm(Tensor self, Tensor qdata, Tensor scale, Tensor scale_step, Tensor zero, Tensor zero_point_step, int group_size) -> Tensor"
from executorch.backends.cuda.quantize_op_dispatch._gemm_family_dispatch import (
chunked_dequant_linear,
quantized_linear,
)


@impl(_lib, "int4_plain_mm", "Meta")
def _meta(self, qdata, scale, scale_step, zero, zero_point_step, group_size):
return torch.empty(
self.shape[0], qdata.shape[0], dtype=self.dtype, device=self.device
)


@impl(_lib, "int4_plain_mm", "CUDA")
def _cuda(self, qdata, scale, scale_step, zero, zero_point_step, group_size):
# Metadata is stored in the coalesced [N, n_groups] layout (transposed at
# pack time, see pack_cuda.pack_linear_for_cuda). The scale is a uint8 code
# with a per-256 fp16 scale_step; the zero is a uint8 code with a per-256
# fp16 zero_point_step. _dequant_matmul reconstructs scale =
# code*scale_step[g//8], zero = code*zero_point_step[g//8].
return _dequant_matmul(
self, qdata, scale, scale_step, zero, zero_point_step, group_size
)


# Chunked dequant for the export GPU budget. The lm_head dequant (N = vocab_size,
# e.g. 262144) runs through the int4_plain_mm custom op (M=1); AOTI executes that
# op's CUDA impl during autotune / cpp_wrapper codegen, where it transiently holds
# ~5 full-size bf16 temporaries (low/high/data/data-z/w_deq) — ~10 GiB for a
# 262144-row weight even though the final w_deq is only ~2.6 GiB. Chunking along N
# caps that at ~chunk rows. It is numerically identical (F.linear output rows are
# independent), and because only the lm_head (custom-op) path crosses the N
# threshold — never the M>4 prefill inline path — it never enters the runtime
# graph: ZERO runtime / accuracy impact. Applied unconditionally to any weight
# whose row count exceeds the threshold.
_DEQUANT_N_THRESHOLD = 65536
_DEQUANT_N_CHUNK = 32768


def _dequant_matmul(x, qdata, scale, scale_step, zero, zero_point_step, group_size):
"""Dequant INT4 weights to input dtype and call F.linear.

Expand All @@ -87,10 +52,6 @@ def _dequant_matmul(x, qdata, scale, scale_step, zero, zero_point_step, group_si
real per-group scale is ``scale_code * scale_step[:, g // 8]``. The zero is a
uint8 code with a per-256-super-block fp16 ``zero_point_step`` ([N, K/256]);
the real per-group zero is ``zero_code * zero_point_step[:, g // 8]``.

Large weights (N > threshold, i.e. the lm_head) are chunked along N to bound
the dequant intermediate (see note above); smaller weights take the original
single-shot dequant.
"""
N, K_half = qdata.shape
K = K_half * 2
Expand All @@ -100,39 +61,26 @@ def _dequant_matmul(x, qdata, scale, scale_step, zero, zero_point_step, group_si
groups_per_super = n_groups // n_super
dtype = x.dtype

def _unit_dq_mm(qd, sc, s_step, ze, z_step, rows):
p = qd.to(torch.uint8).reshape(rows, n_groups, gs_half)
def dequant_linear_rows(i, j):
rows = j - i
p = qdata[i:j].to(torch.uint8).reshape(rows, n_groups, gs_half)
low = (p & 0x0F).to(dtype)
high = ((p >> 4) & 0x0F).to(dtype)
data = torch.stack([low, high], dim=-1).reshape(rows, n_groups, group_size)
# Scale: uint8 code * per-256 fp16 step (broadcast over the 8 groups in
# each super-block).
scale_step_g = s_step.to(dtype).repeat_interleave(groups_per_super, dim=1)
s = (sc.to(dtype) * scale_step_g).unsqueeze(-1)
# Zero: uint8 code * per-256 fp16 step (broadcast over the 8 groups in
# each super-block).
zero_point_step_g = z_step.to(dtype).repeat_interleave(groups_per_super, dim=1)
z = (ze.to(dtype) * zero_point_step_g).unsqueeze(-1)
# Scale and zero: uint8 code * per-256 fp16 step (broadcast over the
# groups in each super-block).
s = (
scale[i:j].to(dtype)
* scale_step[i:j].to(dtype).repeat_interleave(groups_per_super, dim=1)
).unsqueeze(-1)
z = (
zero[i:j].to(dtype)
* zero_point_step[i:j].to(dtype).repeat_interleave(groups_per_super, dim=1)
).unsqueeze(-1)
w_deq = ((data - z) * s).reshape(rows, K)
return F.linear(x, w_deq)

if N <= _DEQUANT_N_THRESHOLD:
return _unit_dq_mm(qdata, scale, scale_step, zero, zero_point_step, N)

outs = []
for i in range(0, N, _DEQUANT_N_CHUNK):
j = min(i + _DEQUANT_N_CHUNK, N)
outs.append(
_unit_dq_mm(
qdata[i:j],
scale[i:j],
scale_step[i:j],
zero[i:j],
zero_point_step[i:j],
j - i,
)
)
return torch.cat(outs, dim=-1)
return chunked_dequant_linear(x, N, dequant_linear_rows)


# ---------------------------------------------------------------------------
Expand All @@ -147,35 +95,27 @@ def _unit_dq_mm(qd, sc, s_step, ze, z_step, rows):
@_implements([aten.linear.default])
@_implements_torch_function([F.linear])
def _(func, types, args, kwargs):
from executorch.backends.cuda.triton.kernels.int4_quantized_gemm import (
INT4_QUANTIZED_GEMM,
)

input_tensor = args[0]
weight_tensor = args[1]
bias = args[2] if len(args) > 2 else None

orig_shape = input_tensor.shape
x_2d = input_tensor.reshape(-1, orig_shape[-1])

qdata = weight_tensor.qdata
scale = weight_tensor.scale
scale_step = weight_tensor.scale_step
zero = weight_tensor.zero_point
zero_point_step = weight_tensor.zero_point_step
gs = weight_tensor.block_size[-1]

M = x_2d.shape[0]
if M <= 4:
# The metadata is already in the coalesced [N, n_groups] layout the
# decode kernel reads directly (baked into the weight constant at pack
# time): scale as uint8 codes + per-256 fp16 scale_step; zero as uint8
# codes + per-256 fp16 zero_point_step. Passing them straight through
# keeps the export graph free of any per-step transpose/clone, so the
# coalesced layout is realized without recomputing it every decode step.
out = torch.ops.executorch_cuda.int4_plain_mm(
x_2d, qdata, scale, scale_step, zero, zero_point_step, gs
)
else:
out = _dequant_matmul(x_2d, qdata, scale, scale_step, zero, zero_point_step, gs)

out = out.reshape(*orig_shape[:-1], -1)
if bias is not None:
out = out + bias
return out
weight = args[1]
bias = args[2] if len(args) > 2 else kwargs.get("bias", None)
# The metadata is already in the coalesced [N, n_groups] layout the kernels
# read, so it passes straight through.
weight_args = (
weight.qdata,
weight.scale,
weight.scale_step,
weight.zero_point,
weight.zero_point_step,
weight.block_size[-1],
)
return quantized_linear(
INT4_QUANTIZED_GEMM,
input_tensor,
weight_args,
bias,
lambda x_2d: _dequant_matmul(x_2d, *weight_args),
)
Loading
Loading