From f5151e230ebb640e9b8777ce45752a0de4eadbda Mon Sep 17 00:00:00 2001 From: gasoonjia Date: Tue, 6 Oct 2026 15:07:46 -0700 Subject: [PATCH] Update [ghstack-poisoned] --- .ci/scripts/wheel/test_shared_libraries.py | 1 - backends/cuda/BUCK | 1 + backends/cuda/CMakeLists.txt | 1 - backends/cuda/cuda_backend.py | 8 - backends/cuda/dp4a_planar_int6_tensor.py | 6 +- .../cuda/quantize_op_dispatch/__init__.py | 2 +- .../quantize_op_dispatch/int6_dispatch.py | 110 +- backends/cuda/runtime/shims/int6_plain_mm.cu | 87 -- backends/cuda/runtime/shims/int6_plain_mm.cuh | 761 ----------- backends/cuda/runtime/shims/int6_plain_mm.h | 68 - .../cuda/runtime/shims/tests/CMakeLists.txt | 31 - .../shims/tests/benchmark_int6_plain_mm.cu | 486 ------- .../shims/tests/gen_plain_mm_test_vectors.py | 74 +- .../test_aoti_torch_cuda_int6_plain_mm.cpp | 1032 -------------- backends/cuda/runtime/targets.bzl | 3 - backends/cuda/tests/targets.bzl | 24 + backends/cuda/tests/test_int6_dispatch.py | 183 ++- .../cuda/tests/test_int6_quantized_gemm.py | 383 ++++++ backends/cuda/tests/test_sort_shim.py | 4 +- .../triton/kernels/int6_quantized_gemm.py | 1205 +++++++++++++++++ .../gemma4_31b/quant/tests/test_pack_cuda.py | 14 +- .../gemma4_31b/tests/test_cuda_packers.py | 14 +- 22 files changed, 1837 insertions(+), 2661 deletions(-) delete mode 100644 backends/cuda/runtime/shims/int6_plain_mm.cu delete mode 100644 backends/cuda/runtime/shims/int6_plain_mm.cuh delete mode 100644 backends/cuda/runtime/shims/int6_plain_mm.h delete mode 100644 backends/cuda/runtime/shims/tests/benchmark_int6_plain_mm.cu delete mode 100644 backends/cuda/runtime/shims/tests/test_aoti_torch_cuda_int6_plain_mm.cpp create mode 100644 backends/cuda/tests/test_int6_quantized_gemm.py create mode 100644 backends/cuda/triton/kernels/int6_quantized_gemm.py diff --git a/.ci/scripts/wheel/test_shared_libraries.py b/.ci/scripts/wheel/test_shared_libraries.py index 61dc8f6d708..642714e9884 100644 --- a/.ci/scripts/wheel/test_shared_libraries.py +++ b/.ci/scripts/wheel/test_shared_libraries.py @@ -97,7 +97,6 @@ # failure this row exists to catch. "aoti_torch_cuda__weight_int4pack_mm", "aoti_torch_cuda_int5_plain_mm", - "aoti_torch_cuda_int6_plain_mm", "aoti_torch_cuda_int8_plain_mm", "aoti_torch_cuda_rand", "aoti_torch_cuda_randint_low_out", diff --git a/backends/cuda/BUCK b/backends/cuda/BUCK index 5b5a28b8299..49df52273a2 100644 --- a/backends/cuda/BUCK +++ b/backends/cuda/BUCK @@ -167,6 +167,7 @@ fbcode_target( "triton/kernels/__init__.py", "triton/kernels/fused_moe.py", "triton/kernels/int4_quantized_gemm.py", + "triton/kernels/int6_quantized_gemm.py", "triton/kernels/int4_matmul.py", "triton/kernels/quantized_gemm_family.py", "triton/kernels/quantized_gemm_utils.py", diff --git a/backends/cuda/CMakeLists.txt b/backends/cuda/CMakeLists.txt index d1670714014..cb4629bd0ba 100644 --- a/backends/cuda/CMakeLists.txt +++ b/backends/cuda/CMakeLists.txt @@ -251,7 +251,6 @@ if(NOT EXECUTORCH_BUILD_ROCM AND CMAKE_CUDA_COMPILER) _aoti_cuda_shim_sources runtime/shims/int4mm.cu runtime/shims/int5_plain_mm.cu - runtime/shims/int6_plain_mm.cu runtime/shims/int8_plain_mm.cu runtime/shims/sort.cu runtime/shims/rand.cu diff --git a/backends/cuda/cuda_backend.py b/backends/cuda/cuda_backend.py index dc47fa73bfe..dcf31e82ddf 100644 --- a/backends/cuda/cuda_backend.py +++ b/backends/cuda/cuda_backend.py @@ -773,8 +773,6 @@ def get_supported_fallback_kernels(cls) -> Dict[str, Any]: "aoti_torch_cuda_randint_low_out": None, "executorch_cuda::int5_plain_mm": None, "aoti_torch_cuda_int5_plain_mm": None, - "executorch_cuda::int6_plain_mm": None, - "aoti_torch_cuda_int6_plain_mm": None, "executorch_cuda::int8_plain_mm": None, "aoti_torch_cuda_int8_plain_mm": None, } @@ -792,12 +790,6 @@ def _get_custom_ops_to_c_shim_options() -> Dict[str, Any]: "AtenTensorHandle, AtenTensorHandle, AtenTensorHandle, " "AtenTensorHandle, int64_t, AtenTensorHandle*)" ], - torch.ops.executorch_cuda.int6_plain_mm.default: [ - "AOTITorchError aoti_torch_cuda_int6_plain_mm(" - "AtenTensorHandle, AtenTensorHandle, AtenTensorHandle, " - "AtenTensorHandle, AtenTensorHandle, int64_t, " - "AtenTensorHandle*)" - ], torch.ops.executorch_cuda.int8_plain_mm.default: [ "AOTITorchError aoti_torch_cuda_int8_plain_mm(" "AtenTensorHandle, AtenTensorHandle, AtenTensorHandle, " diff --git a/backends/cuda/dp4a_planar_int6_tensor.py b/backends/cuda/dp4a_planar_int6_tensor.py index 9b0c4dd7e9d..0f13fde2cc4 100644 --- a/backends/cuda/dp4a_planar_int6_tensor.py +++ b/backends/cuda/dp4a_planar_int6_tensor.py @@ -47,9 +47,9 @@ layout is UNCHANGED. The pack/unpack helpers (:func:`pack_int6`, :func:`unpack_int6`) must stay in -lockstep with ``int6_plain_mm.cuh`` (the decode kernel) — the per-32-weight -``hi_even``/``hi_odd`` byte order is the single most error-prone detail and is -covered by the pack round-trip and the C++ gtest. +lockstep with ``triton/kernels/int6_quantized_gemm.py`` (the decode kernels) — +the per-32-weight ``hi_even``/``hi_odd`` byte order is the single most +error-prone detail and is covered by the pack round-trip and the kernel tests. """ from typing import List, Optional, Tuple diff --git a/backends/cuda/quantize_op_dispatch/__init__.py b/backends/cuda/quantize_op_dispatch/__init__.py index e90e700a142..24dcf2f9111 100644 --- a/backends/cuda/quantize_op_dispatch/__init__.py +++ b/backends/cuda/quantize_op_dispatch/__init__.py @@ -12,7 +12,7 @@ * 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`` + * INT6 (``CudaDp4aPlanarInt6Tensor``) → ``triton::int6_quantized_gemm_m{1,2,3,4}`` * INT8 (``IntxUnpackedToInt8Tensor``) → ``executorch_cuda::int8_plain_mm`` See ``int4_dispatch``, ``int5_dispatch``, ``int6_dispatch`` and ``int8_dispatch`` diff --git a/backends/cuda/quantize_op_dispatch/int6_dispatch.py b/backends/cuda/quantize_op_dispatch/int6_dispatch.py index 1c86ccae522..958008c7bbe 100644 --- a/backends/cuda/quantize_op_dispatch/int6_dispatch.py +++ b/backends/cuda/quantize_op_dispatch/int6_dispatch.py @@ -8,28 +8,29 @@ This module registers an F.linear dispatch on ``CudaDp4aPlanarInt6Tensor`` (an ExecuTorch-internal subclass, see ``dp4a_planar_int6_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*: only GGUF Q6_K weights (converted to ``CudaDp4aPlanarInt6Tensor``) take the packed-int6 path; genuine INT8 weights stay on the int8 path. The code here runs during eager inference and 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::int6_plain_mm`` maps to a C shim that runs - the W6A8 dp4a matvec kernel (backends/cuda/runtime/shims/int6_plain_mm.*). + - ``triton::int6_quantized_gemm_m{M}`` is a Triton W6A8 DP4A kernel compiled + into it (see triton/kernels/int6_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::int6_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 ``INT6_QUANTIZED_GEMM`` that supports the inputs +(static M <= 4 takes its own bucket; a dynamic M provably within [1, 4] takes +the smallest bucket that holds it). Everything else (prefill, an unbounded +dynamic M, other group sizes or dtypes, CPU eager) uses inline dequant + +F.linear, never an error. The packed-int6 weight is symmetric (no zero point): ``w = q * scale`` with -``q`` in ``[-32, 31]`` stored as the ql/qh planes. The op signature mirrors -int4_plain_mm / int8_plain_mm but takes two weight planes (ql, qh) instead of -one, and no zero tensor. +``q`` in ``[-32, 31]`` stored as the ql/qh planes. Importing the parent ``quantize_op_dispatch`` package registers this dispatch -override (along with the INT4 / INT8 ones):: +override (along with the other formats):: import executorch.backends.cuda.quantize_op_dispatch # noqa: F401 """ @@ -40,36 +41,18 @@ CudaDp4aPlanarInt6Tensor, unpack_int6, ) -from executorch.backends.cuda.quantize_op_dispatch._library import lib as _lib -from torch.library import impl - -# --------------------------------------------------------------------------- -# Custom op for INT6 decode (M<=4): W6A8 dp4a matvec in C shim. -# --------------------------------------------------------------------------- - -_lib.define( - "int6_plain_mm(Tensor self, Tensor ql, Tensor qh, Tensor scale, Tensor steps, int group_size) -> Tensor" +from executorch.backends.cuda.quantize_op_dispatch._gemm_family_dispatch import ( + chunked_dequant_linear, + quantized_linear, ) -@impl(_lib, "int6_plain_mm", "Meta") -def _meta_int6(self, ql, qh, scale, steps, group_size): - return torch.empty(self.shape[0], ql.shape[0], dtype=self.dtype, device=self.device) - - -@impl(_lib, "int6_plain_mm", "CUDA") -def _cuda_int6(self, ql, qh, scale, steps, group_size): - # scale is int8 codes in the [N, n_groups] layout; steps is the per-256 - # super-block [N, K/256] fp16 scale step. _unit_dq_mm_int6 reconstructs - # scale = code * steps[:, g // (256 // gs)]. - return _unit_dq_mm_int6(self, ql, qh, scale, steps, group_size) - - def _unit_dq_mm_int6(x, ql, qh, scale, steps, group_size): """Dequant packed-INT6 weights to input dtype and call F.linear. ql [N, K/2] / qh [N, K/4] pack symmetric Q6_K values q in [-32, 31]. - scale [N, K//gs] is int8 codes; steps [N, K//256] fp16 is the per-256 + scale [N, K//gs] is signed 8-bit codes (raw uint8 storage is read as + int8, as the kernels do); steps [N, K//256] fp16 is the per-256 super-block scale step, so the real per-group scale is ``scale_code * steps[:, g // (256 // gs)]``. Dequant: w[n, k] = q[n, k] * (scale_code[n, k//gs] * steps[n, (k//gs) // gps]). @@ -80,19 +63,22 @@ def _unit_dq_mm_int6(x, ql, qh, scale, steps, group_size): n_super = steps.shape[1] groups_per_super = n_groups // n_super dtype = x.dtype + codes = scale.view(torch.int8) if scale.dtype == torch.uint8 else scale - q = unpack_int6(ql, qh, N, K).to(dtype).reshape(N, n_groups, group_size) - # Broadcast the per-256 step over the groups_per_super groups in each - # super-block, then multiply by the int8 code -> effective per-group scale. - step_g = steps.to(dtype).repeat_interleave(groups_per_super, dim=1) - s = (scale.to(dtype) * step_g).reshape(N, n_groups, 1) - w_deq = (q * s).reshape(N, K) + def dequant_linear_rows(i, j): + rows = j - i + q = unpack_int6(ql[i:j], qh[i:j], rows, K).to(dtype).reshape(rows, n_groups, group_size) + # Broadcast the per-256 step over the groups in each super-block, then + # multiply by the int8 code -> effective per-group scale. + step_g = steps[i:j].to(dtype).repeat_interleave(groups_per_super, dim=1) + s = (codes[i:j].to(dtype) * step_g).reshape(rows, n_groups, 1) + return F.linear(x, (q * s).reshape(rows, K)) - return F.linear(x, w_deq) + return chunked_dequant_linear(x, N, dequant_linear_rows) # --------------------------------------------------------------------------- -# CudaDp4aPlanarInt6Tensor F.linear dispatch (W6A8 dp4a for decode) +# CudaDp4aPlanarInt6Tensor F.linear dispatch # --------------------------------------------------------------------------- aten = torch.ops.aten @@ -103,26 +89,24 @@ def _unit_dq_mm_int6(x, ql, qh, scale, steps, group_size): @_implements_i6([aten.linear.default]) @_implements_torch_function_i6([F.linear]) def _(func, types, args, kwargs): + from executorch.backends.cuda.triton.kernels.int6_quantized_gemm import ( + INT6_QUANTIZED_GEMM, + ) + input_tensor = args[0] - weight_tensor = args[1] + weight = args[1] bias = args[2] if len(args) > 2 else kwargs.get("bias", None) - - orig_shape = input_tensor.shape - x_2d = input_tensor.reshape(-1, orig_shape[-1]) - - ql = weight_tensor.ql - qh = weight_tensor.qh - scale = weight_tensor.scale - steps = weight_tensor.steps - gs = weight_tensor.block_size[-1] - - M = x_2d.shape[0] - if M <= 4: - out = torch.ops.executorch_cuda.int6_plain_mm(x_2d, ql, qh, scale, steps, gs) - else: - out = _unit_dq_mm_int6(x_2d, ql, qh, scale, steps, gs) - - out = out.reshape(*orig_shape[:-1], -1) - if bias is not None: - out = out + bias - return out + weight_args = ( + weight.ql, + weight.qh, + weight.scale, + weight.steps, + weight.block_size[-1], + ) + return quantized_linear( + INT6_QUANTIZED_GEMM, + input_tensor, + weight_args, + bias, + lambda x_2d: _unit_dq_mm_int6(x_2d, *weight_args), + ) diff --git a/backends/cuda/runtime/shims/int6_plain_mm.cu b/backends/cuda/runtime/shims/int6_plain_mm.cu deleted file mode 100644 index 5a2c68b91e0..00000000000 --- a/backends/cuda/runtime/shims/int6_plain_mm.cu +++ /dev/null @@ -1,87 +0,0 @@ -/* - * Copyright (c) Meta Platforms, Inc. and affiliates. - * All rights reserved. - * - * This source code is licensed under the BSD-style license found in the - * LICENSE file in the root directory of this source tree. - */ - -#include -#include - -#include -#include -#include -#include -#include - -namespace executorch::backends::cuda { -#ifdef __cplusplus -extern "C" { -#endif - -AOTITorchError aoti_torch_cuda_int6_plain_mm( - Tensor* self, - Tensor* ql, - Tensor* qh, - Tensor* scale, - Tensor* steps, - int64_t group_size, - Tensor** ret0) { - ET_CHECK_OR_RETURN_ERROR( - self != nullptr, - InvalidArgument, - "aoti_torch_cuda_int6_plain_mm: self is null"); - - ET_CHECK_OR_RETURN_ERROR( - ql != nullptr, - InvalidArgument, - "aoti_torch_cuda_int6_plain_mm: ql is null"); - - ET_CHECK_OR_RETURN_ERROR( - qh != nullptr, - InvalidArgument, - "aoti_torch_cuda_int6_plain_mm: qh is null"); - - ET_CHECK_OR_RETURN_ERROR( - scale != nullptr, - InvalidArgument, - "aoti_torch_cuda_int6_plain_mm: scale is null"); - - ET_CHECK_OR_RETURN_ERROR( - steps != nullptr, - InvalidArgument, - "aoti_torch_cuda_int6_plain_mm: steps is null"); - - ET_CHECK_OR_RETURN_ERROR( - ret0 != nullptr, - InvalidArgument, - "aoti_torch_cuda_int6_plain_mm: ret0 is null"); - - int32_t M = self->size(0); - int32_t N = ql->size(0); - Tensor* C = nullptr; - std::array c_shape = {M, N}; - std::array c_stride = {N, 1}; - aoti_torch_empty_strided( - 2, - c_shape.data(), - c_stride.data(), - static_cast( - executorch::backends::aoti::slim::c10::ScalarType::BFloat16), - static_cast( - executorch::backends::aoti::slim::c10::DeviceType::CUDA), - 0, - &C); - - _int6_plain_mm_cuda(*self, *ql, *qh, *scale, *steps, group_size, C); - ET_CUDA_KERNEL_LAUNCH_CHECK_OR_RETURN_ERROR(); - - *ret0 = C; - return Error::Ok; -} - -#ifdef __cplusplus -} -#endif -} // namespace executorch::backends::cuda diff --git a/backends/cuda/runtime/shims/int6_plain_mm.cuh b/backends/cuda/runtime/shims/int6_plain_mm.cuh deleted file mode 100644 index 125f707d1f9..00000000000 --- a/backends/cuda/runtime/shims/int6_plain_mm.cuh +++ /dev/null @@ -1,761 +0,0 @@ -/* - * Copyright (c) Meta Platforms, Inc. and affiliates. - * All rights reserved. - * - * This source code is licensed under the BSD-style license found in the - * LICENSE file in the root directory of this source tree. - */ - -// W6A8 dp4a matvec for packed INT6 decode (M <= 4), used for GGUF Q6_K weights. -// -// Reads a genuine 6-bit packed weight (CudaDp4aPlanarInt6Tensor format), split -// into two planes: -// ql : [N, K/2] uint8 — low-nibble plane, nibble-packed even/odd exactly -// like the INT4 path (ql[:,j] = lo[:,2j] | (lo[:,2j+1] << 4)). -// qh : [N, K/4] uint8 — high-2-bit plane, 4 values/byte, arranged per -// 32-weight chunk as hi_even_packed[4] then hi_odd_packed[4] (each -// byte holds the four 2-bit highs of one dp4a word in even/odd -// order). -// scale : [N, K/gs] int8 — per-group signed scale *codes* (row-major, -// coalesced; no zero), decoded with a per-256-super-block fp16 step -// scale_step[N, K/256]: the group scale is scale_code * scale_step[b], -// b = super-block index = k >> 8. group_size is 16 (GGUF Q6_K), so a -// 256-weight super-block spans 16 groups. -// The stored 6-bit value is u = q + 32 in [0, 63] (q in [-32, 31]); the -// constant -32 offset is applied in the kernel, so Q6_K's symmetry means NO -// zero tensor. The finer per-256 fp16 scale step (vs a per-row step) mirrors -// GGUF Q6_K's own per-super-block fp16 d and lifts whole-weight dequant SNR. -// -// Metadata amortization (mirrors the code-load hoist int4_plain_mm.cuh landed): -// the per-group signed scale is an int8 code decoded with a per-256-super-block -// fp16 step scale_step[N, K/256]. Two cheap reuses keep decode perf-neutral vs a -// per-row step: (1) the fp16 step is loaded once per uint4 (its super-block is -// constant across the uint4's 4 dp4a words) — 8 lanes share each tiny -// L1-resident [N,K/256] address, so the read-only-cache broadcast is cheaper -// than a __shfl; (2) the int8 scale codes are loaded once per group per uint4 -// (gs=16 => 2 codes/uint4, held in registers across the 4 words) instead of the -// 4/uint4 the naive per-word load did. Register-only (no smem, no occupancy -// cliff). An earlier T3 warp-shuffle variant tied this on perf but was strictly -// more complex, since int6's per-row baseline had nothing to amortize (unlike -// int4's per-group step); the code-load hoist is the actual lever. -// -// Dynamically quantizes bf16 activations to INT8 (per-32-element blocks, -// even/odd order, identical to the INT4 path), reconstructs full 6-bit weight -// bytes per dp4a word (vfull = vi_lo | (spread2(hi_byte) << 4)), and uses dp4a -// for fused int6xint8 dot products with vectorized weight loads and -// warp-cooperative quantization. -// -// Symbol names are suffixed _i6 / distinct from int4_plain_mm.cuh and -// int8_plain_mm.cuh so all three translation units can be linked together -// without ODR conflicts. - -#pragma once - -#include -#include -#include -#include - -#include -#include -#include -#include -#include - -namespace executorch::backends::cuda { - -using executorch::backends::aoti::Tensor; -namespace c10 = executorch::backends::aoti::slim::c10; - -// --------------------------------------------------------------------------- -// Constants -// --------------------------------------------------------------------------- - -constexpr int32_t MV6_NWARPS = 8; -constexpr int32_t MV6_WARP_SIZE = 32; -constexpr int32_t MV6_THREADS = MV6_NWARPS * MV6_WARP_SIZE; -constexpr int32_t Q8_BLOCK_SIZE_I6 = 32; -// GGUF Q6_K super-block = 256 weights; the fp16 scale step is per-super-block. -constexpr int32_t SUPER_BLOCK_I6 = 256; -constexpr int32_t SUPER_BLOCK_SHIFT_I6 = 8; // log2(SUPER_BLOCK_I6) - -__host__ __forceinline__ int32_t log2_pow2_i6(int32_t v) { - int32_t r = 0; - while (v > 1) { - v >>= 1; - r++; - } - return r; -} - -// Expand a byte's four 2-bit fields into four byte lanes (each in bits 0-1): -// in : b = [.. b7 b6 | b5 b4 | b3 b2 | b1 b0] -// out : lane0=[b1 b0], lane1=[b3 b2], lane2=[b5 b4], lane3=[b7 b6] -// ~6 ALU ops; verified by truth-table. Used to place the high 2 bits of each -// weight into bits 4-5 of the corresponding dp4a byte lane. -__device__ __forceinline__ uint32_t spread2_i6(uint32_t b) { - uint32_t t = (b | (b << 12)) & 0x000F000F; - uint32_t r = (t | (t << 6)) & 0x03030303; - return r; -} - -// --------------------------------------------------------------------------- -// Activation quantization: bf16 -> int8 (warp-cooperative, per-32-element -// blocks, EVEN/ODD order — identical to the INT4 path's Q8Block). -// --------------------------------------------------------------------------- - -// alignas(16) pads sizeof(Q8Block_i6) to 48 so each block (and its -// qs_even/qs_odd 16-byte halves) is 16-byte aligned, allowing two vectorized -// uint4 loads of a block's int8 activations instead of eight scalar int32 -// loads. -struct alignas(16) Q8Block_i6 { - int8_t qs_even[Q8_BLOCK_SIZE_I6 / 2]; - int8_t qs_odd[Q8_BLOCK_SIZE_I6 / 2]; - int16_t sum8[4]; - float d; // scale -}; - -__global__ void quantize_activations_q8_i6_kernel( - const __nv_bfloat16* __restrict__ A, - Q8Block_i6* __restrict__ q8, - int32_t K) { - const int32_t m = blockIdx.y; - const int32_t block_id = blockIdx.x * blockDim.y + threadIdx.y; - const int32_t n_blocks = K / Q8_BLOCK_SIZE_I6; - if (block_id >= n_blocks) - return; - - const int32_t lane = threadIdx.x; - const __nv_bfloat16* src = - A + static_cast(m) * K + block_id * Q8_BLOCK_SIZE_I6; - Q8Block_i6* dst = q8 + static_cast(m) * n_blocks + block_id; - - float val = __bfloat162float(src[lane]); - - float amax = fabsf(val); - for (int offset = 16; offset > 0; offset >>= 1) - amax = fmaxf(amax, __shfl_xor_sync(0xffffffff, amax, offset)); - - float d = amax / 127.0f; - float id = (d > 0.0f) ? 1.0f / d : 0.0f; - int32_t q = __float2int_rn(val * id); - q = max(-128, min(127, q)); - - if (lane % 2 == 0) - dst->qs_even[lane / 2] = static_cast(q); - else - dst->qs_odd[lane / 2] = static_cast(q); - - int32_t sum8 = q; -#pragma unroll - for (int offset = 4; offset > 0; offset >>= 1) { - sum8 += __shfl_xor_sync(0xffffffff, sum8, offset); - } - if ((lane & 7) == 0) { - dst->sum8[lane >> 3] = static_cast(sum8); - } - - if (lane == 0) - dst->d = d; -} - -// --------------------------------------------------------------------------- -// W6A8 dp4a matvec kernel (per-256 fp16 step + code-load hoist) -// -// dp4a is linear, so reconstructing v = lo + (hi<<4) and dotting once is -// equivalent to two separate dp4a passes. We reconstruct the full 6-bit byte -// (vfull = vi_lo | (spread2(hi_byte) << 4)) so a single dp4a per even/odd half -// covers the whole weight. The per-group zero is the constant 32 (in u-space), -// applied as out += scale * a_scale * (dp - 32 * a_sum) — no zero load. -// -// The per-group signed scale is a coalesced int8 code decoded with a -// per-256-super-block fp16 step. Both are amortized per uint4 (see file header): -// the fp16 step is loaded once per uint4 (constant across its super-block; 8 -// lanes share the tiny L1-resident address), and the int8 codes once per group -// per uint4 (gs=16 => 2/uint4, register-held across the 4 words) rather than -// per word. Register-only (no smem => no occupancy cliff). -// --------------------------------------------------------------------------- - -__device__ __forceinline__ uint32_t uint4_at_i6(uint4 v, int32_t i) { - return i == 0 ? v.x : (i == 1 ? v.y : (i == 2 ? v.z : v.w)); -} - - -__device__ __forceinline__ void accum_i6( - int32_t vfull_even, - int32_t vfull_odd, - uint32_t a_even, - uint32_t a_odd, - float scale, - float& sum) { - int32_t dp = __dp4a(vfull_even, static_cast(a_even), 0); - dp = __dp4a(vfull_odd, static_cast(a_odd), dp); - int32_t a_sum = __dp4a(0x01010101, static_cast(a_even), 0); - a_sum = __dp4a(0x01010101, static_cast(a_odd), a_sum); - sum += scale * (static_cast(dp) - 32.0f * static_cast(a_sum)); -} - -__device__ __forceinline__ void accum_i6_sum( - int32_t vfull_even, - int32_t vfull_odd, - uint32_t a_even, - uint32_t a_odd, - float scale, - int32_t a_sum, - float& sum) { - int32_t dp = __dp4a(vfull_even, static_cast(a_even), 0); - dp = __dp4a(vfull_odd, static_cast(a_odd), dp); - sum += scale * (static_cast(dp) - 32.0f * static_cast(a_sum)); -} - -template -__device__ __forceinline__ void int6_w6a8_matvec_body( - const uint8_t* __restrict__ ql, - const uint8_t* __restrict__ qh, - const int8_t* __restrict__ w_scale, - const __half* __restrict__ w_scale_step, - const Q8Block_i6* __restrict__ q8, - __nv_bfloat16* __restrict__ out, - int32_t n, - int32_t N, - int32_t K, - int32_t gs_shift, - int32_t n_groups, - int32_t n_super) { - const int32_t K_half = K / 2; - const int32_t K_quarter = K / 4; - const int32_t lane_id = threadIdx.x; - const int32_t n_q8_blocks = K / Q8_BLOCK_SIZE_I6; - - const uint8_t* qlrow = ql + static_cast(n) * K_half; - const uint8_t* qhrow = qh + static_cast(n) * K_quarter; - const int8_t* scale_row = w_scale + static_cast(n) * n_groups; - const __half* scale_step_row = - w_scale_step + static_cast(n) * n_super; - - const uint4* qlrow16 = reinterpret_cast(qlrow); - const uint2* qhrow8 = reinterpret_cast(qhrow); - const int32_t K_half_16 = K_half / 16; - const int32_t sb_shift = SUPER_BLOCK_SHIFT_I6 - gs_shift; - const int32_t wpg_shift = gs_shift - 3; - float sums[ROWS] = {}; - - for (int32_t i = lane_id; i < K_half_16; i += MV6_WARP_SIZE) { - const uint4 packed16 = __ldg(&qlrow16[i]); - const uint2 qh_chunk = __ldg(&qhrow8[i]); - const int32_t k_base = i * 32; - const uint32_t words[4] = {packed16.x, packed16.y, packed16.z, packed16.w}; - const uint32_t hi_even_word = qh_chunk.x; - const uint32_t hi_odd_word = qh_chunk.y; - - const int32_t g_base = GS16 ? (k_base >> 4) : (k_base >> gs_shift); - const float scale_step = __half2float( - __ldg(&scale_step_row[GS16 ? (g_base >> 4) : (g_base >> sb_shift)])); - const float ws0 = - static_cast(__ldg(&scale_row[g_base])) * scale_step; - const float ws1 = - static_cast(__ldg(&scale_row[g_base + 1])) * scale_step; - - uint4 activations_even[ROWS]; - uint4 activations_odd[ROWS]; - float activation_scales[ROWS]; - int16_t activation_sums[ROWS][4]; -#pragma unroll - for (int32_t row = 0; row < ROWS; ++row) { - const Q8Block_i6* qb = q8 + static_cast(row) * n_q8_blocks + i; - activations_even[row] = *reinterpret_cast(qb->qs_even); - activations_odd[row] = *reinterpret_cast(qb->qs_odd); - activation_scales[row] = qb->d; - if constexpr (USE_SUM) { - activation_sums[row][0] = qb->sum8[0]; - activation_sums[row][1] = qb->sum8[1]; - activation_sums[row][2] = qb->sum8[2]; - activation_sums[row][3] = qb->sum8[3]; - } - } - -#pragma unroll - for (int32_t w = 0; w < 4; w++) { - const uint32_t packed = words[w]; - const int32_t vi_lo = static_cast(packed & 0x0F0F0F0F); - const int32_t vi_hi = static_cast((packed >> 4) & 0x0F0F0F0F); - const uint32_t hi_even_byte = (hi_even_word >> (w * 8)) & 0xFF; - const uint32_t hi_odd_byte = (hi_odd_word >> (w * 8)) & 0xFF; - const int32_t vfull_even = - vi_lo | static_cast(spread2_i6(hi_even_byte) << 4); - const int32_t vfull_odd = - vi_hi | static_cast(spread2_i6(hi_odd_byte) << 4); - const float ws = GS16 ? ((w >= 2) ? ws1 : ws0) - : ((w >> wpg_shift) ? ws1 : ws0); - -#pragma unroll - for (int32_t row = 0; row < ROWS; ++row) { - if constexpr (USE_SUM) { - accum_i6_sum( - vfull_even, - vfull_odd, - uint4_at_i6(activations_even[row], w), - uint4_at_i6(activations_odd[row], w), - ws * activation_scales[row], - static_cast(activation_sums[row][w]), - sums[row]); - } else { - accum_i6( - vfull_even, - vfull_odd, - uint4_at_i6(activations_even[row], w), - uint4_at_i6(activations_odd[row], w), - ws * activation_scales[row], - sums[row]); - } - } - } - } - - for (int offset = MV6_WARP_SIZE / 2; offset > 0; offset >>= 1) { -#pragma unroll - for (int32_t row = 0; row < ROWS; ++row) { - sums[row] += __shfl_xor_sync(0xffffffff, sums[row], offset); - } - } - if (lane_id == 0) { -#pragma unroll - for (int32_t row = 0; row < ROWS; ++row) { - out[static_cast(row) * N + n] = __float2bfloat16(sums[row]); - } - } -} - -__global__ void __launch_bounds__(MV6_THREADS) int6_w6a8_matvec_kernel( - const uint8_t* __restrict__ ql, // [N, K/2] - const uint8_t* __restrict__ qh, // [N, K/4] - const int8_t* __restrict__ w_scale, // [N, n_groups] int8 codes - const __half* __restrict__ w_scale_step, // [N, n_super] fp16 - const Q8Block_i6* __restrict__ q8, - __nv_bfloat16* __restrict__ out, - int32_t N, - int32_t K, - int32_t M, - int32_t gs_shift, - int32_t n_groups, - int32_t n_super) { - const int32_t n = blockIdx.x * MV6_NWARPS + threadIdx.y; - if (n >= N) { - return; - } - if (M == 1) { - int6_w6a8_matvec_body<1, false, false>( - ql, qh, w_scale, w_scale_step, q8, out, n, N, K, gs_shift, n_groups, n_super); - return; - } - if (M == 2) { - int6_w6a8_matvec_body<2, false, false>( - ql, qh, w_scale, w_scale_step, q8, out, n, N, K, gs_shift, n_groups, n_super); - return; - } - if (M == 3) { - int6_w6a8_matvec_body<3, false, false>( - ql, qh, w_scale, w_scale_step, q8, out, n, N, K, gs_shift, n_groups, n_super); - return; - } - int6_w6a8_matvec_body<4, false, false>( - ql, qh, w_scale, w_scale_step, q8, out, n, N, K, gs_shift, n_groups, n_super); -} - -#define DEFINE_INT6_GS16_KERNEL(ROWS) \ - __global__ void __launch_bounds__(MV6_THREADS) \ - int6_w6a8_matvec_m##ROWS##_gs16_kernel( \ - const uint8_t* __restrict__ ql, \ - const uint8_t* __restrict__ qh, \ - const int8_t* __restrict__ w_scale, \ - const __half* __restrict__ w_scale_step, \ - const Q8Block_i6* __restrict__ q8, \ - __nv_bfloat16* __restrict__ out, \ - int32_t N, \ - int32_t K, \ - int32_t n_groups, \ - int32_t n_super) { \ - const int32_t n = blockIdx.x * MV6_NWARPS + threadIdx.y; \ - if (n >= N) { \ - return; \ - } \ - int6_w6a8_matvec_body( \ - ql, qh, w_scale, w_scale_step, q8, out, n, N, K, 4, n_groups, n_super); \ - } - -DEFINE_INT6_GS16_KERNEL(1) -DEFINE_INT6_GS16_KERNEL(2) - -__global__ void __launch_bounds__(MV6_THREADS) int6_w6a8_matvec_m3_gs16_kernel( - const uint8_t* __restrict__ ql, - const uint8_t* __restrict__ qh, - const int8_t* __restrict__ w_scale, - const __half* __restrict__ w_scale_step, - const Q8Block_i6* __restrict__ q8, - __nv_bfloat16* __restrict__ out, - int32_t N, - int32_t K, - int32_t n_groups, - int32_t n_super) { - const int32_t n = blockIdx.x * MV6_NWARPS + threadIdx.y; - if (n >= N) - return; - - const int32_t K_half = K / 2; - const int32_t K_quarter = K / 4; - const int32_t lane_id = threadIdx.x; - const int32_t n_q8_blocks = K / Q8_BLOCK_SIZE_I6; - - const uint8_t* qlrow = ql + static_cast(n) * K_half; - const uint8_t* qhrow = qh + static_cast(n) * K_quarter; - const int8_t* scale_row = w_scale + static_cast(n) * n_groups; - const __half* scale_step_row = - w_scale_step + static_cast(n) * n_super; - - const uint4* qlrow16 = reinterpret_cast(qlrow); - const uint2* qhrow8 = reinterpret_cast(qhrow); - const int32_t K_half_16 = K_half / 16; - float s0 = 0.0f; - float s1 = 0.0f; - float s2 = 0.0f; - - for (int32_t i = lane_id; i < K_half_16; i += MV6_WARP_SIZE) { - uint4 packed16 = __ldg(&qlrow16[i]); - uint2 qh_chunk = __ldg(&qhrow8[i]); - int32_t k_base = i * 32; - uint32_t words[4] = {packed16.x, packed16.y, packed16.z, packed16.w}; - uint32_t hi_even_word = qh_chunk.x; - uint32_t hi_odd_word = qh_chunk.y; - - int32_t g_base = k_base >> 4; - float scale_step = - __half2float(__ldg(&scale_step_row[g_base >> 4])); - const float ws0 = - static_cast(__ldg(&scale_row[g_base])) * scale_step; - const float ws1 = - static_cast(__ldg(&scale_row[g_base + 1])) * scale_step; - - const Q8Block_i6* qb0 = q8 + i; - const Q8Block_i6* qb1 = q8 + n_q8_blocks + i; - const Q8Block_i6* qb2 = q8 + static_cast(2) * n_q8_blocks + i; - uint4 ae0 = *reinterpret_cast(qb0->qs_even); - uint4 ao0 = *reinterpret_cast(qb0->qs_odd); - uint4 ae1 = *reinterpret_cast(qb1->qs_even); - uint4 ao1 = *reinterpret_cast(qb1->qs_odd); - uint4 ae2 = *reinterpret_cast(qb2->qs_even); - uint4 ao2 = *reinterpret_cast(qb2->qs_odd); - float as0 = qb0->d; - float as1 = qb1->d; - float as2 = qb2->d; - -#pragma unroll - for (int32_t w = 0; w < 4; w++) { - uint32_t packed = words[w]; - int32_t vi_lo = static_cast(packed & 0x0F0F0F0F); - int32_t vi_hi = static_cast((packed >> 4) & 0x0F0F0F0F); - uint32_t hi_even_byte = (hi_even_word >> (w * 8)) & 0xFF; - uint32_t hi_odd_byte = (hi_odd_word >> (w * 8)) & 0xFF; - int32_t vfull_even = - vi_lo | static_cast(spread2_i6(hi_even_byte) << 4); - int32_t vfull_odd = - vi_hi | static_cast(spread2_i6(hi_odd_byte) << 4); - float ws = (w >= 2) ? ws1 : ws0; - accum_i6( - vfull_even, - vfull_odd, - uint4_at_i6(ae0, w), - uint4_at_i6(ao0, w), - ws * as0, - s0); - accum_i6( - vfull_even, - vfull_odd, - uint4_at_i6(ae1, w), - uint4_at_i6(ao1, w), - ws * as1, - s1); - accum_i6( - vfull_even, - vfull_odd, - uint4_at_i6(ae2, w), - uint4_at_i6(ao2, w), - ws * as2, - s2); - } - } - - for (int offset = MV6_WARP_SIZE / 2; offset > 0; offset >>= 1) { - s0 += __shfl_xor_sync(0xffffffff, s0, offset); - s1 += __shfl_xor_sync(0xffffffff, s1, offset); - s2 += __shfl_xor_sync(0xffffffff, s2, offset); - } - if (lane_id == 0) { - out[n] = __float2bfloat16(s0); - out[static_cast(N) + n] = __float2bfloat16(s1); - out[static_cast(2) * N + n] = __float2bfloat16(s2); - } -} - -DEFINE_INT6_GS16_KERNEL(4) - -#undef DEFINE_INT6_GS16_KERNEL - -#define DEFINE_INT6_SUM_GS16_KERNEL(ROWS) \ - __global__ void __launch_bounds__(MV6_THREADS) \ - int6_w6a8_matvec_m##ROWS##_sum_gs16_kernel( \ - const uint8_t* __restrict__ ql, \ - const uint8_t* __restrict__ qh, \ - const int8_t* __restrict__ w_scale, \ - const __half* __restrict__ w_scale_step, \ - const Q8Block_i6* __restrict__ q8, \ - __nv_bfloat16* __restrict__ out, \ - int32_t N, \ - int32_t K, \ - int32_t n_groups, \ - int32_t n_super) { \ - const int32_t n = blockIdx.x * MV6_NWARPS + threadIdx.y; \ - if (n >= N) { \ - return; \ - } \ - int6_w6a8_matvec_body( \ - ql, qh, w_scale, w_scale_step, q8, out, n, N, K, 4, n_groups, n_super); \ - } - -DEFINE_INT6_SUM_GS16_KERNEL(3) -DEFINE_INT6_SUM_GS16_KERNEL(4) - -#undef DEFINE_INT6_SUM_GS16_KERNEL - -__global__ void __launch_bounds__(MV6_THREADS) int6_w6a8_matvec_m4_kernel( - const uint8_t* __restrict__ ql, - const uint8_t* __restrict__ qh, - const int8_t* __restrict__ w_scale, - const __half* __restrict__ w_scale_step, - const Q8Block_i6* __restrict__ q8, - __nv_bfloat16* __restrict__ out, - int32_t N, - int32_t K, - int32_t gs_shift, - int32_t n_groups, - int32_t n_super) { - const int32_t n = blockIdx.x * MV6_NWARPS + threadIdx.y; - if (n >= N) { - return; - } - int6_w6a8_matvec_body<4, false, false>( - ql, qh, w_scale, w_scale_step, q8, out, n, N, K, gs_shift, n_groups, n_super); -} - -// --------------------------------------------------------------------------- -// Persistent Q8 buffer (lazy init, not thread-safe — single-stream only). -// Freed at process exit via a static guard so leak detectors stay quiet; the -// CUDA runtime would otherwise reclaim it on teardown anyway. -// --------------------------------------------------------------------------- - -static Q8Block_i6* g_q8_buf_i6 = nullptr; -static size_t g_q8_buf_i6_size = 0; - -namespace { -struct Q8BufferGuardI6 { - ~Q8BufferGuardI6() { - if (g_q8_buf_i6) { - // Ignore errors: during process teardown the CUDA context may already be - // gone (cudaErrorCudartUnloading), which is harmless here. - cudaFree(g_q8_buf_i6); - g_q8_buf_i6 = nullptr; - g_q8_buf_i6_size = 0; - } - } -}; -Q8BufferGuardI6 g_q8_buf_i6_guard; -} // namespace - -static Q8Block_i6* get_q8_buffer_i6(size_t needed) { - if (g_q8_buf_i6_size < needed) { - if (g_q8_buf_i6) - cudaFree(g_q8_buf_i6); - cudaError_t err = cudaMalloc(&g_q8_buf_i6, needed); - ET_CHECK_MSG( - err == cudaSuccess, - "cudaMalloc failed for Q8 buffer (int6): %s", - cudaGetErrorString(err)); - g_q8_buf_i6_size = needed; - } - return g_q8_buf_i6; -} - -// --------------------------------------------------------------------------- -// Main entry point -// --------------------------------------------------------------------------- - -inline void _int6_plain_mm_cuda( - const Tensor& A, // [M, K] bf16 - const Tensor& ql, // [N, K/2] uint8 - const Tensor& qh, // [N, K/4] uint8 - const Tensor& scale, // [N, K/gs] int8 codes - const Tensor& steps, // [N, K/256] fp16 per-256 scale_step - int64_t group_size, - Tensor* output) { // [M, N] bf16, pre-allocated - int32_t M = A.size(0); - int32_t K = A.size(1); - int32_t N = ql.size(0); - - ET_CHECK(A.dtype() == c10::ScalarType::BFloat16); - ET_CHECK_MSG(M >= 1 && M <= 4, "int6 short-query M=%d must be in [1, 4]", M); - ET_CHECK( - ql.dtype() == c10::ScalarType::Byte || - ql.dtype() == c10::ScalarType::Char); - ET_CHECK( - qh.dtype() == c10::ScalarType::Byte || - qh.dtype() == c10::ScalarType::Char); - ET_CHECK( - scale.dtype() == c10::ScalarType::Byte || - scale.dtype() == c10::ScalarType::Char); - ET_CHECK(steps.dtype() == c10::ScalarType::Half); - ET_CHECK(A.dim() == 2); - ET_CHECK(ql.dim() == 2); - ET_CHECK(ql.size(1) == K / 2); - ET_CHECK(qh.dim() == 2); - ET_CHECK(qh.size(1) == K / 4); - ET_CHECK(scale.dim() == 2); - ET_CHECK(scale.size(0) == N); - ET_CHECK(steps.dim() == 2); - ET_CHECK(steps.size(0) == N); - ET_CHECK(steps.size(1) == K / SUPER_BLOCK_I6); - - int32_t gs = static_cast(group_size); - ET_CHECK_MSG( - gs > 0 && (gs & (gs - 1)) == 0, "group_size=%d must be a power of 2", gs); - // group_size must be a multiple of 8 (the dp4a word stride) so a word never - // straddles a group boundary; gs=16 covers GGUF Q6_K. - ET_CHECK_MSG( - gs % 8 == 0, - "group_size=%d must be a multiple of 8 (e.g. 16 for GGUF Q6_K)", - gs); - ET_CHECK_MSG( - K >= Q8_BLOCK_SIZE_I6 && K % Q8_BLOCK_SIZE_I6 == 0, - "K=%d must be a positive multiple of %d for dp4a int6 kernel", - K, - Q8_BLOCK_SIZE_I6); - ET_CHECK_MSG( - K % SUPER_BLOCK_I6 == 0, - "K=%d must be a multiple of %d (super-block) for the per-256 scale step", - K, - SUPER_BLOCK_I6); - - auto stream_result = getCurrentCUDAStream(0); - ET_CHECK_MSG(stream_result.ok(), "Failed to get CUDA stream"); - cudaStream_t stream = stream_result.get(); - - int32_t gs_shift = log2_pow2_i6(gs); - - // Quantize activations to INT8 (even/odd order) - int32_t n_q8_blocks = K / Q8_BLOCK_SIZE_I6; - size_t q8_bytes = static_cast(M) * n_q8_blocks * sizeof(Q8Block_i6); - Q8Block_i6* q8_buf = get_q8_buffer_i6(q8_bytes); - - constexpr int32_t Q8_WARPS = 8; - int32_t blocks_per_m = (n_q8_blocks + Q8_WARPS - 1) / Q8_WARPS; - dim3 q8_grid(blocks_per_m, M); - dim3 q8_block(MV6_WARP_SIZE, Q8_WARPS); - quantize_activations_q8_i6_kernel<<>>( - reinterpret_cast(A.data_ptr()), q8_buf, K); - - // dp4a matvec - dim3 grid((N + MV6_NWARPS - 1) / MV6_NWARPS); - dim3 block(MV6_WARP_SIZE, MV6_NWARPS); - - int32_t n_groups = static_cast(scale.size(1)); - int32_t n_super = static_cast(steps.size(1)); - if (M == 4 && gs == 16) { - int6_w6a8_matvec_m4_sum_gs16_kernel<<>>( - reinterpret_cast(ql.data_ptr()), - reinterpret_cast(qh.data_ptr()), - reinterpret_cast(scale.data_ptr()), - reinterpret_cast(steps.data_ptr()), - q8_buf, - reinterpret_cast<__nv_bfloat16*>(output->data_ptr()), - N, - K, - n_groups, - n_super); - return; - } - - if (M == 3 && gs == 16 && N >= 1024) { - int6_w6a8_matvec_m3_sum_gs16_kernel<<>>( - reinterpret_cast(ql.data_ptr()), - reinterpret_cast(qh.data_ptr()), - reinterpret_cast(scale.data_ptr()), - reinterpret_cast(steps.data_ptr()), - q8_buf, - reinterpret_cast<__nv_bfloat16*>(output->data_ptr()), - N, - K, - n_groups, - n_super); - return; - } - - if (M == 2 && gs == 16 && N >= 1024) { - int6_w6a8_matvec_m2_gs16_kernel<<>>( - reinterpret_cast(ql.data_ptr()), - reinterpret_cast(qh.data_ptr()), - reinterpret_cast(scale.data_ptr()), - reinterpret_cast(steps.data_ptr()), - q8_buf, - reinterpret_cast<__nv_bfloat16*>(output->data_ptr()), - N, - K, - n_groups, - n_super); - return; - } - - if (M == 1 && gs == 16 && N >= 1024 && K >= 8192) { - int6_w6a8_matvec_m1_gs16_kernel<<>>( - reinterpret_cast(ql.data_ptr()), - reinterpret_cast(qh.data_ptr()), - reinterpret_cast(scale.data_ptr()), - reinterpret_cast(steps.data_ptr()), - q8_buf, - reinterpret_cast<__nv_bfloat16*>(output->data_ptr()), - N, - K, - n_groups, - n_super); - return; - } - - if (M == 4) { - int6_w6a8_matvec_m4_kernel<<>>( - reinterpret_cast(ql.data_ptr()), - reinterpret_cast(qh.data_ptr()), - reinterpret_cast(scale.data_ptr()), - reinterpret_cast(steps.data_ptr()), - q8_buf, - reinterpret_cast<__nv_bfloat16*>(output->data_ptr()), - N, - K, - gs_shift, - n_groups, - n_super); - return; - } - - int6_w6a8_matvec_kernel<<>>( - reinterpret_cast(ql.data_ptr()), - reinterpret_cast(qh.data_ptr()), - reinterpret_cast(scale.data_ptr()), - reinterpret_cast(steps.data_ptr()), - q8_buf, - reinterpret_cast<__nv_bfloat16*>(output->data_ptr()), - N, - K, - M, - gs_shift, - n_groups, - n_super); -} - -} // namespace executorch::backends::cuda diff --git a/backends/cuda/runtime/shims/int6_plain_mm.h b/backends/cuda/runtime/shims/int6_plain_mm.h deleted file mode 100644 index 26e7c42e974..00000000000 --- a/backends/cuda/runtime/shims/int6_plain_mm.h +++ /dev/null @@ -1,68 +0,0 @@ -/* - * Copyright (c) Meta Platforms, Inc. and affiliates. - * All rights reserved. - * - * This source code is licensed under the BSD-style license found in the - * LICENSE file in the root directory of this source tree. - */ - -#pragma once - -#include -#include -#include - -namespace executorch::backends::cuda { - -using executorch::backends::aoti::AOTITorchError; -using executorch::backends::aoti::Tensor; - -#ifdef __cplusplus -extern "C" { -#endif - -/** - * Packed INT6 matrix multiplication for GGUF Q6_K weights (symmetric). - * - * The 6-bit weight is split into two planes plus a per-group scale; there is - * NO zero tensor — Q6_K is symmetric and the stored value is u = q + 32 in - * [0, 63] (q in [-32, 31]), with the constant -32 offset applied in the kernel. - * - * Weight format: - * ql : [N, K/2] uint8 — low-nibble plane, nibble-packed even/odd - * (ql[:,j] = (u[:,2j] & 0xF) | ((u[:,2j+1] & 0xF) << 4)). - * qh : [N, K/4] uint8 — high-2-bit plane, 4 values/byte, arranged per - * 32-weight chunk as hi_even_packed[4] then hi_odd_packed[4]; each - * byte holds the four 2-bit highs of one dp4a word, bit field j - * (bits 2j..2j+1) = high 2 bits of that word's j-th even/odd weight. - * scale : [N, K//group_size] int8 per-group signed scale codes (row-major), - * decoded with a per-256-super-block [N, K//256] fp16 ``steps``: the - * group scale is ``scale_code * steps[:, g // (256 // group_size)]``. - * This mirrors GGUF Q6_K's own per-super-block fp16 ``d`` - * granularity. W6A8 dp4a matvec: dynamically quantizes activations to INT8, - * reconstructs full 6-bit weight bytes, then uses dp4a for fused int6xint8 dot - * products. - * - * @param self Input activation [M, K] bf16 - * @param ql Low-nibble plane [N, K/2] uint8 - * @param qh High-2-bit plane [N, K/4] uint8 - * @param scale Per-group scale codes [N, K//group_size] int8 - * @param steps Per-256-super-block scale step [N, K//256] fp16 (scale = - * code * steps[:, g // (256 // group_size)]) - * @param group_size Quantization group size (multiple of 8; e.g. 16 for Q6_K) - * @param ret0 Output [M, N] bf16 - */ -AOTI_SHIM_EXPORT AOTITorchError aoti_torch_cuda_int6_plain_mm( - Tensor* self, - Tensor* ql, - Tensor* qh, - Tensor* scale, - Tensor* steps, - int64_t group_size, - Tensor** ret0); - -#ifdef __cplusplus -} -#endif - -} // namespace executorch::backends::cuda diff --git a/backends/cuda/runtime/shims/tests/CMakeLists.txt b/backends/cuda/runtime/shims/tests/CMakeLists.txt index 5e848e23bff..ce6d0dea4d4 100644 --- a/backends/cuda/runtime/shims/tests/CMakeLists.txt +++ b/backends/cuda/runtime/shims/tests/CMakeLists.txt @@ -52,7 +52,6 @@ set(CUDA_SHIM_TESTS # CUDA-specific tests requiring GPU kernels set(CUDA_KERNEL_TESTS test_aoti_torch_cuda__weight_int4pack_mm test_aoti_torch_cuda_int5_plain_mm - test_aoti_torch_cuda_int6_plain_mm ) enable_testing() @@ -75,36 +74,6 @@ foreach(test_name ${CUDA_SHIM_TESTS}) add_test(NAME ${test_name} COMMAND ${test_name}) endforeach() -add_executable(benchmark_int6_plain_mm benchmark_int6_plain_mm.cu) -target_include_directories( - benchmark_int6_plain_mm PRIVATE ${EXECUTORCH_ROOT}/.. ${EXECUTORCH_ROOT} - ${CUDAToolkit_INCLUDE_DIRS} -) -target_compile_definitions(benchmark_int6_plain_mm PRIVATE CUDA_AVAILABLE=1) -set_property(TARGET benchmark_int6_plain_mm PROPERTY CUDA_ARCHITECTURES 80) -target_compile_options( - benchmark_int6_plain_mm PRIVATE $<$:-Xptxas=-v> -) -target_link_libraries( - benchmark_int6_plain_mm PRIVATE aoti_cuda_shims executorch_core CUDA::cudart -) -add_test(NAME benchmark_int6_plain_mm_m1 - COMMAND benchmark_int6_plain_mm --M=1 --N=1024 --K=8192 --gs=16 - --warmup=1 --iters=1 --seeds=3 -) -add_test(NAME benchmark_int6_plain_mm_m2 - COMMAND benchmark_int6_plain_mm --M=2 --N=1024 --K=256 --gs=16 - --warmup=1 --iters=1 --seeds=3 -) -add_test(NAME benchmark_int6_plain_mm_m3 - COMMAND benchmark_int6_plain_mm --M=3 --N=1024 --K=256 --gs=16 - --warmup=1 --iters=1 --seeds=3 -) -add_test(NAME benchmark_int6_plain_mm_m4 - COMMAND benchmark_int6_plain_mm --M=4 --N=64 --K=256 --gs=16 - --warmup=1 --iters=1 --seeds=3 -) - add_executable(benchmark_int5_plain_mm benchmark_int5_plain_mm.cu) target_include_directories( benchmark_int5_plain_mm PRIVATE ${EXECUTORCH_ROOT}/.. ${EXECUTORCH_ROOT} diff --git a/backends/cuda/runtime/shims/tests/benchmark_int6_plain_mm.cu b/backends/cuda/runtime/shims/tests/benchmark_int6_plain_mm.cu deleted file mode 100644 index d25476d6190..00000000000 --- a/backends/cuda/runtime/shims/tests/benchmark_int6_plain_mm.cu +++ /dev/null @@ -1,486 +0,0 @@ -/* Copyright (c) Meta Platforms, Inc. and affiliates. */ - -#include -#include -#include - -#include - -#include -#include -#include -#include -#include -#include -#include -#include - -namespace cuda_shims = executorch::backends::cuda; - -namespace { - -#define CUDA_CHECK(expr) \ - do { \ - cudaError_t err__ = (expr); \ - if (err__ != cudaSuccess) { \ - std::fprintf( \ - stderr, \ - "CUDA error %s:%d: %s\n", \ - __FILE__, \ - __LINE__, \ - cudaGetErrorString(err__)); \ - std::exit(1); \ - } \ - } while (0) - -__global__ void __launch_bounds__(cuda_shims::MV6_THREADS) - int6_w6a8_matvec_reference_kernel( - const uint8_t* __restrict__ ql, - const uint8_t* __restrict__ qh, - const int8_t* __restrict__ w_scale, - const __half* __restrict__ w_scale_step, - const cuda_shims::Q8Block_i6* __restrict__ q8, - __nv_bfloat16* __restrict__ out, - int32_t N, - int32_t K, - int32_t M, - int32_t gs_shift, - int32_t n_groups, - int32_t n_super) { - const int32_t n = blockIdx.x * cuda_shims::MV6_NWARPS + threadIdx.y; - const int32_t m = blockIdx.y; - if (n >= N || m >= M) { - return; - } - - const int32_t K_half = K / 2; - const int32_t K_quarter = K / 4; - const int32_t lane_id = threadIdx.x; - const int32_t n_q8_blocks = K / cuda_shims::Q8_BLOCK_SIZE_I6; - - const uint8_t* qlrow = ql + static_cast(n) * K_half; - const uint8_t* qhrow = qh + static_cast(n) * K_quarter; - const int8_t* scale_row = w_scale + static_cast(n) * n_groups; - const __half* scale_step_row = - w_scale_step + static_cast(n) * n_super; - const cuda_shims::Q8Block_i6* q8_row = - q8 + static_cast(m) * n_q8_blocks; - - const uint4* qlrow16 = reinterpret_cast(qlrow); - const uint2* qhrow8 = reinterpret_cast(qhrow); - const int32_t K_half_16 = K_half / 16; - const int32_t sb_shift = cuda_shims::SUPER_BLOCK_SHIFT_I6 - gs_shift; - const int32_t wpg_shift = gs_shift - 3; - - float sum = 0.0f; - for (int32_t i = lane_id; i < K_half_16; i += cuda_shims::MV6_WARP_SIZE) { - uint4 packed16 = __ldg(&qlrow16[i]); - uint2 qh_chunk = __ldg(&qhrow8[i]); - int32_t k_base = i * 32; - uint32_t words[4] = {packed16.x, packed16.y, packed16.z, packed16.w}; - uint32_t hi_even_word = qh_chunk.x; - uint32_t hi_odd_word = qh_chunk.y; - int32_t g_base = k_base >> gs_shift; - float scale_step = - __half2float(__ldg(&scale_step_row[g_base >> sb_shift])); - float ws0 = static_cast(__ldg(&scale_row[g_base])) * scale_step; - float ws1 = static_cast(__ldg(&scale_row[g_base + 1])) * scale_step; - const cuda_shims::Q8Block_i6* qb = &q8_row[i]; - uint4 ae = *reinterpret_cast(qb->qs_even); - uint4 ao = *reinterpret_cast(qb->qs_odd); - float a_scale = qb->d; - uint32_t a_even[4] = {ae.x, ae.y, ae.z, ae.w}; - uint32_t a_odd[4] = {ao.x, ao.y, ao.z, ao.w}; - -#pragma unroll - for (int32_t w = 0; w < 4; ++w) { - uint32_t packed = words[w]; - int32_t vi_lo = static_cast(packed & 0x0F0F0F0F); - int32_t vi_hi = static_cast((packed >> 4) & 0x0F0F0F0F); - uint32_t hi_even_byte = (hi_even_word >> (w * 8)) & 0xFF; - uint32_t hi_odd_byte = (hi_odd_word >> (w * 8)) & 0xFF; - int32_t vfull_even = - vi_lo | static_cast(cuda_shims::spread2_i6(hi_even_byte) << 4); - int32_t vfull_odd = - vi_hi | static_cast(cuda_shims::spread2_i6(hi_odd_byte) << 4); - int32_t dp = __dp4a(vfull_even, static_cast(a_even[w]), 0); - dp = __dp4a(vfull_odd, static_cast(a_odd[w]), dp); - int32_t a_sum = - __dp4a(0x01010101, static_cast(a_even[w]), 0); - a_sum = __dp4a(0x01010101, static_cast(a_odd[w]), a_sum); - float ws = (w >> wpg_shift) ? ws1 : ws0; - sum += ws * a_scale * - (static_cast(dp) - 32.0f * static_cast(a_sum)); - } - } - - for (int offset = cuda_shims::MV6_WARP_SIZE / 2; offset > 0; offset >>= 1) { - sum += __shfl_xor_sync(0xffffffff, sum, offset); - } - if (lane_id == 0) { - out[static_cast(m) * N + n] = __float2bfloat16(sum); - } -} - -int32_t log2_pow2_host(int32_t v) { - int32_t r = 0; - while (v > 1) { - v >>= 1; - ++r; - } - return r; -} - -uint16_t float_to_bf16(float x) { - uint32_t bits; - std::memcpy(&bits, &x, sizeof(bits)); - return static_cast(bits >> 16); -} - -void fill_case( - int64_t M, - int64_t N, - int64_t K, - int64_t gs, - uint32_t seed, - std::vector& A, - std::vector& ql, - std::vector& qh, - std::vector& scale, - std::vector& steps) { - std::mt19937 rng(seed); - std::uniform_real_distribution adist(-2.0f, 2.0f); - std::uniform_int_distribution qdist(0, 63); - std::uniform_int_distribution sdist(-8, 8); - std::uniform_real_distribution stepdist(0.003f, 0.035f); - - A.resize(M * K); - ql.assign(N * (K / 2), 0); - qh.assign(N * (K / 4), 0); - scale.resize(N * (K / gs)); - steps.resize(N * (K / 256)); - - for (auto& x : A) { - x = float_to_bf16(adist(rng)); - } - for (auto& x : scale) { - int v = sdist(rng); - x = static_cast(v == 0 ? 1 : v); - } - for (auto& x : steps) { - x = __half_as_ushort(__float2half(stepdist(rng))); - } - - std::vector u(K); - for (int64_t n = 0; n < N; ++n) { - uint8_t* qlrow = ql.data() + n * (K / 2); - uint8_t* qhrow = qh.data() + n * (K / 4); - for (int64_t k = 0; k < K; ++k) { - u[k] = static_cast(qdist(rng)); - } - for (int64_t k = 0; k < K; k += 2) { - qlrow[k / 2] = static_cast((u[k] & 0xF) | ((u[k + 1] & 0xF) << 4)); - } - for (int64_t k = 0; k < K; k += 32) { - uint8_t* chunk = qhrow + k / 4; - for (int w = 0; w < 4; ++w) { - uint8_t he = 0; - uint8_t ho = 0; - for (int j = 0; j < 4; ++j) { - he |= static_cast(((u[k + w * 8 + j * 2] >> 4) & 0x3) << (j * 2)); - ho |= static_cast(((u[k + w * 8 + j * 2 + 1] >> 4) & 0x3) << (j * 2)); - } - chunk[w] = he; - chunk[4 + w] = ho; - } - } - } -} - -template -float time_ms(Fn fn, int warmup, int iterations) { - for (int i = 0; i < warmup; ++i) { - fn(); - } - CUDA_CHECK(cudaDeviceSynchronize()); - cudaEvent_t start, stop; - CUDA_CHECK(cudaEventCreate(&start)); - CUDA_CHECK(cudaEventCreate(&stop)); - CUDA_CHECK(cudaEventRecord(start)); - for (int i = 0; i < iterations; ++i) { - fn(); - } - CUDA_CHECK(cudaEventRecord(stop)); - CUDA_CHECK(cudaEventSynchronize(stop)); - float elapsed = 0.0f; - CUDA_CHECK(cudaEventElapsedTime(&elapsed, start, stop)); - CUDA_CHECK(cudaEventDestroy(start)); - CUDA_CHECK(cudaEventDestroy(stop)); - return elapsed / iterations; -} - -void* device_alloc_copy(const void* src, size_t bytes) { - void* dst = nullptr; - CUDA_CHECK(cudaMalloc(&dst, bytes)); - CUDA_CHECK(cudaMemcpy(dst, src, bytes, cudaMemcpyHostToDevice)); - return dst; -} - -int64_t arg_value(int argc, char** argv, const char* name, int64_t fallback) { - std::string key = std::string("--") + name + "="; - for (int i = 1; i < argc; ++i) { - std::string arg(argv[i]); - if (arg.rfind(key, 0) == 0) { - return std::stoll(arg.substr(key.size())); - } - } - return fallback; -} - -} // namespace - -int main(int argc, char** argv) { - const int64_t M = arg_value(argc, argv, "M", 4); - const int64_t K = arg_value(argc, argv, "K", 6656); - const int64_t N = arg_value(argc, argv, "N", 6656); - const int64_t gs = arg_value(argc, argv, "gs", 16); - const int64_t warmup = arg_value(argc, argv, "warmup", 30); - const int64_t iterations = arg_value(argc, argv, "iters", 200); - const int64_t seeds = arg_value(argc, argv, "seeds", 3); - - if (M < 1 || M > 4 || K % 256 != 0 || K % 32 != 0 || gs < 8 || (gs & (gs - 1))) { - std::fprintf(stderr, "unsupported shape M=%lld N=%lld K=%lld gs=%lld\n", M, N, K, gs); - return 2; - } - - CUDA_CHECK(cudaSetDevice(0)); - int32_t gs_shift = log2_pow2_host(static_cast(gs)); - int32_t n_groups = static_cast(K / gs); - int32_t n_super = static_cast(K / 256); - int32_t n_q8_blocks = static_cast(K / cuda_shims::Q8_BLOCK_SIZE_I6); - dim3 q8_grid((n_q8_blocks + cuda_shims::MV6_NWARPS - 1) / cuda_shims::MV6_NWARPS, M); - dim3 block(cuda_shims::MV6_WARP_SIZE, cuda_shims::MV6_NWARPS); - dim3 ref_grid((N + cuda_shims::MV6_NWARPS - 1) / cuda_shims::MV6_NWARPS, M); - dim3 cand_grid((N + cuda_shims::MV6_NWARPS - 1) / cuda_shims::MV6_NWARPS); - - std::printf( - "shape M=%lld N=%lld K=%lld gs=%lld q8_blocks=%d warmup=%lld iters=%lld seeds=%lld\n", - M, - N, - K, - gs, - n_q8_blocks, - warmup, - iterations, - seeds); - - double ref_total = 0.0; - double cand_total = 0.0; - for (int64_t seed_idx = 0; seed_idx < seeds; ++seed_idx) { - std::vector hA; - std::vector hql; - std::vector hqh; - std::vector hscale; - std::vector hsteps; - fill_case(M, N, K, gs, 1009 + seed_idx * 17, hA, hql, hqh, hscale, hsteps); - - auto* dA = static_cast<__nv_bfloat16*>(device_alloc_copy(hA.data(), hA.size() * sizeof(uint16_t))); - auto* dql = static_cast(device_alloc_copy(hql.data(), hql.size() * sizeof(uint8_t))); - auto* dqh = static_cast(device_alloc_copy(hqh.data(), hqh.size() * sizeof(uint8_t))); - auto* dscale = static_cast(device_alloc_copy(hscale.data(), hscale.size() * sizeof(int8_t))); - auto* dsteps = static_cast<__half*>(device_alloc_copy(hsteps.data(), hsteps.size() * sizeof(uint16_t))); - cuda_shims::Q8Block_i6* dq8 = nullptr; - __nv_bfloat16* dref = nullptr; - __nv_bfloat16* dcand = nullptr; - CUDA_CHECK(cudaMalloc(&dq8, M * n_q8_blocks * sizeof(cuda_shims::Q8Block_i6))); - CUDA_CHECK(cudaMalloc(&dref, M * N * sizeof(__nv_bfloat16))); - CUDA_CHECK(cudaMalloc(&dcand, M * N * sizeof(__nv_bfloat16))); - - cuda_shims::quantize_activations_q8_i6_kernel<<>>(dA, dq8, static_cast(K)); - CUDA_CHECK(cudaGetLastError()); - CUDA_CHECK(cudaDeviceSynchronize()); - - if (M == 3 && gs == 16) { - cuda_shims::int6_w6a8_matvec_m3_gs16_kernel<<>>( - dql, - dqh, - dscale, - dsteps, - dq8, - dref, - static_cast(N), - static_cast(K), - n_groups, - n_super); - } else { - int6_w6a8_matvec_reference_kernel<<>>( - dql, - dqh, - dscale, - dsteps, - dq8, - dref, - static_cast(N), - static_cast(K), - static_cast(M), - gs_shift, - n_groups, - n_super); - } - if (M == 4 && gs == 16) { - cuda_shims::int6_w6a8_matvec_m4_sum_gs16_kernel<<>>( - dql, - dqh, - dscale, - dsteps, - dq8, - dcand, - static_cast(N), - static_cast(K), - n_groups, - n_super); - } else if (M == 3 && gs == 16) { - cuda_shims::int6_w6a8_matvec_m3_sum_gs16_kernel<<>>( - dql, - dqh, - dscale, - dsteps, - dq8, - dcand, - static_cast(N), - static_cast(K), - n_groups, - n_super); - } else if (M == 2 && gs == 16) { - cuda_shims::int6_w6a8_matvec_m2_gs16_kernel<<>>( - dql, - dqh, - dscale, - dsteps, - dq8, - dcand, - static_cast(N), - static_cast(K), - n_groups, - n_super); - } else if (M == 1 && gs == 16) { - cuda_shims::int6_w6a8_matvec_m1_gs16_kernel<<>>( - dql, - dqh, - dscale, - dsteps, - dq8, - dcand, - static_cast(N), - static_cast(K), - n_groups, - n_super); - } else if (M == 4) { - cuda_shims::int6_w6a8_matvec_m4_kernel<<>>( - dql, - dqh, - dscale, - dsteps, - dq8, - dcand, - static_cast(N), - static_cast(K), - gs_shift, - n_groups, - n_super); - } else { - cuda_shims::int6_w6a8_matvec_kernel<<>>( - dql, - dqh, - dscale, - dsteps, - dq8, - dcand, - static_cast(N), - static_cast(K), - static_cast(M), - gs_shift, - n_groups, - n_super); - } - CUDA_CHECK(cudaGetLastError()); - CUDA_CHECK(cudaDeviceSynchronize()); - - std::vector href(M * N); - std::vector hcand(M * N); - CUDA_CHECK(cudaMemcpy(href.data(), dref, href.size() * sizeof(uint16_t), cudaMemcpyDeviceToHost)); - CUDA_CHECK(cudaMemcpy(hcand.data(), dcand, hcand.size() * sizeof(uint16_t), cudaMemcpyDeviceToHost)); - size_t mismatches = 0; - for (size_t i = 0; i < href.size(); ++i) { - if (href[i] != hcand[i]) { - if (mismatches < 8) { - std::printf("mismatch seed=%lld idx=%zu ref=0x%04x cand=0x%04x\n", seed_idx, i, href[i], hcand[i]); - } - ++mismatches; - } - } - if (mismatches != 0) { - std::printf("bitwise=FAIL mismatches=%zu/%zu\n", mismatches, href.size()); - return 1; - } - - float q8_ms = time_ms([&]() { - cuda_shims::quantize_activations_q8_i6_kernel<<>>(dA, dq8, static_cast(K)); - }, warmup, iterations); - float ref_ms = time_ms([&]() { - if (M == 3 && gs == 16) { - cuda_shims::int6_w6a8_matvec_m3_gs16_kernel<<>>( - dql, dqh, dscale, dsteps, dq8, dref, static_cast(N), static_cast(K), n_groups, n_super); - } else { - int6_w6a8_matvec_reference_kernel<<>>( - dql, dqh, dscale, dsteps, dq8, dref, static_cast(N), static_cast(K), static_cast(M), gs_shift, n_groups, n_super); - } - }, warmup, iterations); - float cand_ms = time_ms([&]() { - if (M == 4 && gs == 16) { - cuda_shims::int6_w6a8_matvec_m4_sum_gs16_kernel<<>>( - dql, dqh, dscale, dsteps, dq8, dcand, static_cast(N), static_cast(K), n_groups, n_super); - } else if (M == 3 && gs == 16) { - cuda_shims::int6_w6a8_matvec_m3_sum_gs16_kernel<<>>( - dql, dqh, dscale, dsteps, dq8, dcand, static_cast(N), static_cast(K), n_groups, n_super); - } else if (M == 2 && gs == 16) { - cuda_shims::int6_w6a8_matvec_m2_gs16_kernel<<>>( - dql, dqh, dscale, dsteps, dq8, dcand, static_cast(N), static_cast(K), n_groups, n_super); - } else if (M == 1 && gs == 16) { - cuda_shims::int6_w6a8_matvec_m1_gs16_kernel<<>>( - dql, dqh, dscale, dsteps, dq8, dcand, static_cast(N), static_cast(K), n_groups, n_super); - } else if (M == 4) { - cuda_shims::int6_w6a8_matvec_m4_kernel<<>>( - dql, dqh, dscale, dsteps, dq8, dcand, static_cast(N), static_cast(K), gs_shift, n_groups, n_super); - } else { - cuda_shims::int6_w6a8_matvec_kernel<<>>( - dql, dqh, dscale, dsteps, dq8, dcand, static_cast(N), static_cast(K), static_cast(M), gs_shift, n_groups, n_super); - } - }, warmup, iterations); - ref_total += ref_ms; - cand_total += cand_ms; - std::printf( - "seed=%lld bitwise=OK q8_ms=%.6f ref_matvec_ms=%.6f cand_matvec_ms=%.6f speedup=%.3fx\n", - seed_idx, - q8_ms, - ref_ms, - cand_ms, - ref_ms / cand_ms); - - CUDA_CHECK(cudaFree(dA)); - CUDA_CHECK(cudaFree(dql)); - CUDA_CHECK(cudaFree(dqh)); - CUDA_CHECK(cudaFree(dscale)); - CUDA_CHECK(cudaFree(dsteps)); - CUDA_CHECK(cudaFree(dq8)); - CUDA_CHECK(cudaFree(dref)); - CUDA_CHECK(cudaFree(dcand)); - } - - std::printf( - "avg ref_matvec_ms=%.6f cand_matvec_ms=%.6f speedup=%.3fx\n", - ref_total / seeds, - cand_total / seeds, - ref_total / cand_total); - return 0; -} diff --git a/backends/cuda/runtime/shims/tests/gen_plain_mm_test_vectors.py b/backends/cuda/runtime/shims/tests/gen_plain_mm_test_vectors.py index 4ca75320c9d..cadbabc29d4 100644 --- a/backends/cuda/runtime/shims/tests/gen_plain_mm_test_vectors.py +++ b/backends/cuda/runtime/shims/tests/gen_plain_mm_test_vectors.py @@ -5,12 +5,11 @@ # This source code is licensed under the BSD-style license found in the # LICENSE file in the root directory of this source tree. -"""Regenerate the hardcoded INT5/INT6 plain_mm dp4a test vectors. +"""Regenerate the hardcoded INT5 plain_mm dp4a test vectors. This script deterministically recreates every ``uint8_t``/``int8_t``/``uint16_t`` array embedded in the plain_mm gtest files: --dtype int5 -> test_aoti_torch_cuda_int5_plain_mm.cpp - --dtype int6 -> test_aoti_torch_cuda_int6_plain_mm.cpp Each dtype has a fixed seed, so the emitted vectors are reproducible by construction: the vectors in the .cpp are exactly this script's output. @@ -20,7 +19,7 @@ INT5 (W5A8, asymmetric Q5_K; pack path in dp4a_planar_int5_tensor.py): 1. torch.manual_seed(INT5_SEED) ONCE, then draw ALL cases in list order - (unlike int6, the int5 cases share one RNG stream, so build_int5 + (the int5 cases share one RNG stream, so build_int5 replays the earlier cases to stay reproducible per-case). 2. Per case, draw a scaled fp32 weight ``[N, K]`` (``randn * (0.5 + rand[N,1])``) then activation ``[M, K]`` (bf16). Weight-then-activation @@ -32,22 +31,11 @@ zero_point codes [N, K/gs] uint8, zero_point_step [N, K/256] fp16. 4. expected = F.linear(A, tensor.dequantize(bf16)). -INT6 (W6A8, symmetric Q6_K, NO zero tensor; pack path in -dp4a_planar_int6_tensor.py): - 1. torch.manual_seed(case.seed) on CPU. - 2. Draw symmetric Q6_K values q ``[N, K]`` in [-32, 31], a small positive - per-group scale ``[N, K/gs]`` (bf16), then an activation ``A`` ``[M, K]`` - (bf16). Draw order q -> scale -> A is part of the seed contract. - 3. pack_int6(q) -> planar ql [N, K/2] uint8, qh [N, K/4] uint8. - 4. _encode_int8_per_super(scale, gs) -> scale codes [N, K/gs] int8 + - per-256-super-block step [N, K/256] fp16 (scale = code * step[:, g//gps]). - 5. expected = F.linear(A, tensor.dequantize(bf16)). - -Both kernels quantize activations to int8, so the .cpp compares with a 0.5 atol. +The kernel quantizes activations to int8, so the .cpp compares with a 0.5 atol. Usage (from the executorch repo root, conda env with torch + torchao): python backends/cuda/runtime/shims/tests/gen_plain_mm_test_vectors.py \\ - --dtype {int5,int6} [--case NAME] [--check] + --dtype int5 [--case NAME] [--check] Without ``--check`` it prints the C++ array blocks for each case to stdout; paste them into the matching TEST_F body. With ``--check`` it re-derives the vectors @@ -86,13 +74,6 @@ class Case: Case("PackedShuffleMultiSuper", M=1, K=1024, N=8, gs=32, seed=INT5_SEED), ] -INT6_CASES: List[Case] = [ - Case("Q6KSingleSuperBlock", M=2, K=256, N=4, gs=16, seed=0), - Case("Q6KMultiSuperBlock", M=1, K=512, N=6, gs=16, seed=1), - Case("Q6KWideN", M=1, K=256, N=16, gs=16, seed=2), -] - - # --------------------------------------------------------------------------- # Shared bit/format helpers. # --------------------------------------------------------------------------- @@ -205,41 +186,6 @@ def build_int5(case: Case) -> Dict[str, tuple]: } -def build_int6(case: Case) -> Dict[str, tuple]: - """Return {array_name: (ctype, [ints])} for one INT6 case (CPU only).""" - from executorch.backends.cuda.dp4a_planar_int6_tensor import ( - _encode_int8_per_super, - CudaDp4aPlanarInt6Tensor, - pack_int6, - ) - - torch.manual_seed(case.seed) - # Symmetric q, then scale, then activation: fixed order is part of the seed - # contract. Matches test_int6_dispatch._make_int6_tensor's convention. - q = torch.randint(-32, 32, (case.N, case.K), dtype=torch.int8) - scale = (torch.rand(case.N, case.K // case.gs) * 0.1 + 0.01).to(torch.bfloat16) - A = torch.randn(case.M, case.K, dtype=torch.bfloat16) - - ql, qh = pack_int6(q) - scale_codes, steps = _encode_int8_per_super(scale.float(), case.gs) - tensor = CudaDp4aPlanarInt6Tensor( - ql, qh, scale_codes, steps, [1, case.gs], torch.Size([case.N, case.K]) - ) - - # bf16 dequant @ F.linear reference (kernel adds activation-quant noise). - w_deq = tensor.dequantize(torch.bfloat16) - expected = torch.nn.functional.linear(A, w_deq) - - return { - "ql_host": ("uint8_t", _u8(ql)), - "qh_host": ("uint8_t", _u8(qh)), - "scale_codes": ("int8_t", _i8(scale_codes)), - "scale_step": ("uint16_t", _fp16_bits(steps)), - "A_host": ("uint16_t", _bf16_bits(A)), - "expected": ("uint16_t", _bf16_bits(expected)), - } - - @dataclass(frozen=True) class DtypeSpec: name: str @@ -268,14 +214,6 @@ class DtypeSpec: test_class="AOTITorchInt5PlainMMTest", cpp_name="test_aoti_torch_cuda_int5_plain_mm.cpp", ), - "int6": DtypeSpec( - name="int6", - cases=INT6_CASES, - order=["ql_host", "qh_host", "scale_codes", "scale_step", "A_host", "expected"], - build=build_int6, - test_class="AOTITorchInt6PlainMMTest", - cpp_name="test_aoti_torch_cuda_int6_plain_mm.cpp", - ), } @@ -286,7 +224,7 @@ def _fmt_array(name: str, ctype: str, values: List[int]) -> str: if ctype == "uint8_t": per_line, cell = 12, lambda v: f"0x{v & 0xFF:02X}" elif ctype == "int8_t": - # Signed decimal, right-aligned like the .cpp (Q6_K scale codes). + # Signed decimal, right-aligned like the .cpp. per_line, cell = 12, lambda v: f"{v:4d}" elif ctype == "uint16_t": per_line, cell = 8, lambda v: f"0x{v & 0xFFFF:04X}" @@ -316,7 +254,7 @@ def _parse_cpp_array(text: str, test_class: str, case_name: str, arr: str) -> Li """Extract a single array's ints from the given TEST_F body in the .cpp. Handles both hex cells (e.g. ql/qh/scale_step/A/expected) and signed-decimal - cells (int6 scale_codes int8). + cells. """ m = re.search( rf"TEST_F\({re.escape(test_class)},\s*{re.escape(case_name)}\)", diff --git a/backends/cuda/runtime/shims/tests/test_aoti_torch_cuda_int6_plain_mm.cpp b/backends/cuda/runtime/shims/tests/test_aoti_torch_cuda_int6_plain_mm.cpp deleted file mode 100644 index 345c8b88d8f..00000000000 --- a/backends/cuda/runtime/shims/tests/test_aoti_torch_cuda_int6_plain_mm.cpp +++ /dev/null @@ -1,1032 +0,0 @@ -/* - * Copyright (c) Meta Platforms, Inc. and affiliates. - * All rights reserved. - * - * This source code is licensed under the BSD-style license found in the - * LICENSE file in the root directory of this source tree. - */ - -#include -#include -#include -#include -#include -#include -#include -#include -#include -#include -#include - -using executorch::backends::cuda::aoti_torch_cuda_int6_plain_mm; -using executorch::backends::cuda::aoti_torch_empty_strided; -using executorch::backends::cuda::AOTITorchError; -using executorch::runtime::Error; -namespace slim_c10 = executorch::backends::aoti::slim::c10; - -using Tensor = executorch::backends::aoti::slim::SlimTensor; - -// W6A8 dp4a matvec shim for packed-INT6 decode (CudaDp4aPlanarInt6Tensor -// layout, GGUF Q6_K). The 6-bit weight is split into two planes plus a -// per-group scale; there is NO zero tensor (Q6_K is symmetric, the -32 offset -// is applied in the kernel): -// ql : [N, K/2] uint8 — low-nibble plane, nibble-packed even/odd -// qh : [N, K/4] uint8 — high-2-bit plane, 4 values/byte (per 32-weight -// chunk: hi_even_packed[4] then hi_odd_packed[4]) -// scale : [N, K/gs] int8 — per-group signed scale *codes* (row-major) -// steps : [N, K/256] fp16 — per-256-super-block scale step; the real -// per-group scale is scale_code * steps[:, g / (256 / gs)]. The step -// is loaded once per super-block and __shfl-broadcast across each -// 8-lane subgroup (T3 super-block-cooperative decode). -// Test vectors are generated from the production pack path (pack_int6 + -// _encode_int8_per_super, dp4a_planar_int6_tensor.py) and the expected[] -// outputs from the tensor's own bf16 dequant @ F.linear (the same math the M>4 -// dispatch reference _unit_dq_mm_int6 computes): -// w[n, k] = q[n, k] * (scale_code[n, k//gs] * steps[n, (k//gs)/(256/gs)]); -// out = A @ w^T (q symmetric, in [-32, 31]). The kernel runs W6A8 (it also -// quantizes activations to int8), so a 0.5 atol absorbs the activation-quant -// noise. The cases cover single/multi super-block step indexing and the -// warp-shuffle broadcast (Q6KWideN spans multiple warps of N). -// -// ALL arrays below (ql_host/qh_host/scale_codes/scale_step/A_host/expected) are -// GENERATED — do not hand-edit. Regenerate them with the checked-in script -// (which is the source of truth for these vectors): -// python backends/cuda/runtime/shims/tests/gen_plain_mm_test_vectors.py \ -// --dtype int6 -// and verify they are in sync with: -// python backends/cuda/runtime/shims/tests/gen_plain_mm_test_vectors.py \ -// --dtype int6 --check -class AOTITorchInt6PlainMMTest : public ::testing::Test { - protected: - void SetUp() override { - et_pal_init(); - - int device_count = 0; - cudaError_t err = cudaGetDeviceCount(&device_count); - if (err != cudaSuccess || device_count == 0) { - GTEST_SKIP() << "CUDA not available"; - } - } - - Tensor* create_tensor( - const std::vector& sizes, - slim_c10::ScalarType dtype) { - Tensor* tensor; - AOTITorchError error = aoti_torch_empty_strided( - sizes.size(), - sizes.data(), - nullptr, - static_cast(dtype), - static_cast(slim_c10::DeviceType::CUDA), - 0, - &tensor); - return (error == Error::Ok) ? tensor : nullptr; - } - - Tensor* create_bf16(const std::vector& sizes) { - return create_tensor(sizes, slim_c10::ScalarType::BFloat16); - } - - // The per-256 scale step is fp16. - Tensor* create_fp16(const std::vector& sizes) { - return create_tensor(sizes, slim_c10::ScalarType::Half); - } - - // ql / qh are uint8 (ScalarType::Byte) packed planes. - Tensor* create_uint8(const std::vector& sizes) { - return create_tensor(sizes, slim_c10::ScalarType::Byte); - } - - // scale codes are signed int8 (ScalarType::Char). - Tensor* create_int8(const std::vector& sizes) { - return create_tensor(sizes, slim_c10::ScalarType::Char); - } - - // Upload raw bytes to a CUDA tensor. - void upload(Tensor* t, const void* host_data, size_t bytes) { - cudaMemcpy(t->data_ptr(), host_data, bytes, cudaMemcpyHostToDevice); - } - - // Download CUDA tensor to host buffer. - void download(const Tensor* t, void* host_data, size_t bytes) { - cudaMemcpy(host_data, t->data_ptr(), bytes, cudaMemcpyDeviceToHost); - } - - // Run the shim and return the output tensor (asserts success). - Tensor* run( - Tensor* A, - Tensor* ql, - Tensor* qh, - Tensor* scale, - Tensor* steps, - int64_t group_size) { - Tensor* output = nullptr; - AOTITorchError error = aoti_torch_cuda_int6_plain_mm( - A, ql, qh, scale, steps, group_size, &output); - EXPECT_EQ(error, Error::Ok); - EXPECT_NE(output, nullptr); - return output; - } - - // Check output bf16 values against expected, with absolute tolerance. - void check_bf16_output( - Tensor* output, - const uint16_t* expected_data, - int64_t count, - float atol = 0.5f) { - std::vector actual(count); - download(output, actual.data(), count * sizeof(uint16_t)); - cudaDeviceSynchronize(); - - for (int64_t i = 0; i < count; i++) { - // Convert bf16 raw bits to float: bf16 is the upper 16 bits of float32. - uint32_t actual_bits = static_cast(actual[i]) << 16; - uint32_t expected_bits = static_cast(expected_data[i]) << 16; - float actual_f, expected_f; - memcpy(&actual_f, &actual_bits, sizeof(float)); - memcpy(&expected_f, &expected_bits, sizeof(float)); - - EXPECT_NEAR(actual_f, expected_f, atol) - << "Mismatch at index " << i << ": actual=" << actual_f - << " expected=" << expected_f; - } - } - - // Upload data and run the shim. ql/qh are uint8; scale is int8 codes; steps - // is the per-256 fp16 scale step ([N, K/256]); A is bf16. - Tensor* setup_and_run( - int64_t M, - int64_t N, - int64_t K, - int64_t gs, - const uint8_t* ql_host, - const uint8_t* qh_host, - const int8_t* scale_codes, - const uint16_t* steps_host, - const uint16_t* A_host) { - int64_t ng = K / gs; - int64_t n_super = K / 256; - Tensor* A = create_bf16({M, K}); - Tensor* ql = create_uint8({N, K / 2}); - Tensor* qh = create_uint8({N, K / 4}); - Tensor* scale = create_int8({N, ng}); - Tensor* steps = create_fp16({N, n_super}); - EXPECT_NE(A, nullptr); - EXPECT_NE(ql, nullptr); - EXPECT_NE(qh, nullptr); - EXPECT_NE(scale, nullptr); - EXPECT_NE(steps, nullptr); - - upload(A, A_host, static_cast(M) * K * sizeof(uint16_t)); - upload(ql, ql_host, static_cast(N) * (K / 2) * sizeof(uint8_t)); - upload(qh, qh_host, static_cast(N) * (K / 4) * sizeof(uint8_t)); - upload(scale, scale_codes, static_cast(N) * ng * sizeof(int8_t)); - upload( - steps, steps_host, static_cast(N) * n_super * sizeof(uint16_t)); - - return run(A, ql, qh, scale, steps, gs); - } -}; - -// Q6KSingleSuperBlock: M=2, N=4, K=256, gs=16, ng=16, n_super=1 -TEST_F(AOTITorchInt6PlainMMTest, Q6KSingleSuperBlock) { - int64_t M = 2, N = 4, K = 256, gs = 16; - int64_t ng = K / gs; // 16 - int64_t n_super = K / 256; // 1 - // clang-format off - uint8_t ql_host[] = { - 0xFC, 0x05, 0xB3, 0x73, 0x39, 0x25, 0x74, 0x86, 0xC8, 0x1A, 0x76, 0xE7, - 0x18, 0x95, 0x8D, 0x49, 0x03, 0x53, 0xFE, 0x0F, 0x32, 0x18, 0xD3, 0x33, - 0x7E, 0x10, 0x99, 0x0F, 0xAF, 0x74, 0xE3, 0x2B, 0xC7, 0x02, 0x40, 0x55, - 0x86, 0x14, 0x4F, 0xA9, 0xFA, 0x18, 0x71, 0x99, 0x63, 0xB7, 0x2E, 0x0B, - 0x3E, 0xC5, 0xA9, 0xB4, 0x64, 0xF4, 0x4F, 0xC3, 0x44, 0xE8, 0x4F, 0xA3, - 0xF7, 0x5D, 0x05, 0x51, 0x39, 0xF0, 0xE5, 0x10, 0x42, 0x02, 0x3D, 0xA2, - 0x0D, 0x57, 0xF9, 0xA0, 0xB2, 0xFA, 0xB7, 0x92, 0xE2, 0xE3, 0x3B, 0xE2, - 0x43, 0x21, 0xEB, 0xA9, 0x41, 0x6A, 0x8B, 0x2B, 0x03, 0x60, 0x60, 0xE3, - 0x3A, 0xC8, 0xD8, 0x8E, 0x2D, 0xE3, 0xB2, 0x0D, 0x88, 0xF3, 0xA8, 0x82, - 0x34, 0xDC, 0x0E, 0x34, 0xBD, 0x6C, 0x9D, 0xBD, 0x08, 0xD8, 0xE5, 0x09, - 0x9C, 0x56, 0x3D, 0x81, 0x40, 0xBB, 0x69, 0xF5, 0x7D, 0x88, 0x9F, 0x82, - 0xF6, 0x6B, 0xDF, 0x19, 0x6F, 0x8C, 0x8D, 0x23, 0x3F, 0xCA, 0x36, 0xE6, - 0x5D, 0xB7, 0xB0, 0x48, 0x6A, 0x5B, 0xCD, 0x8F, 0xB2, 0x93, 0x57, 0xE3, - 0x54, 0x3D, 0x73, 0x99, 0x79, 0xE3, 0x2A, 0x3C, 0x9F, 0xBA, 0x7D, 0xD7, - 0x15, 0x2C, 0x82, 0x51, 0xF8, 0x04, 0x2B, 0xE5, 0xE5, 0xE0, 0xD8, 0x11, - 0x30, 0xB8, 0x48, 0x04, 0x39, 0xC7, 0x23, 0xDE, 0x11, 0x12, 0x4D, 0x2D, - 0x55, 0xCC, 0xD5, 0xF2, 0x5B, 0x77, 0x6B, 0x1E, 0x76, 0xC2, 0x3B, 0x91, - 0xED, 0x95, 0x29, 0xB0, 0xC9, 0x1B, 0xA9, 0x60, 0x0F, 0x4A, 0x8E, 0xF4, - 0x33, 0x88, 0xDB, 0xEB, 0x07, 0x83, 0x77, 0xAD, 0x81, 0x74, 0xB0, 0x4C, - 0x09, 0xEA, 0xC6, 0x24, 0x64, 0xA3, 0x73, 0xD8, 0x5D, 0x80, 0xFF, 0xB5, - 0x74, 0xDC, 0xAA, 0xE4, 0xB1, 0x33, 0x29, 0x25, 0xE3, 0xA5, 0x7B, 0xCC, - 0xD2, 0x17, 0xA6, 0x05, 0xBE, 0xA0, 0x3C, 0x1A, 0x9A, 0xDC, 0x69, 0x76, - 0xD8, 0x78, 0x80, 0x86, 0x9F, 0x38, 0x16, 0x47, 0x29, 0x0C, 0x28, 0x7D, - 0x48, 0xC4, 0x71, 0x6E, 0x49, 0x51, 0x79, 0xDB, 0x31, 0x75, 0x63, 0x76, - 0x19, 0x69, 0x30, 0x48, 0x41, 0x05, 0xE3, 0x1A, 0x44, 0x04, 0xF0, 0xDA, - 0xE8, 0xBB, 0xFF, 0x64, 0xC9, 0xFB, 0x33, 0x2C, 0xA1, 0x12, 0x3B, 0x4C, - 0x1A, 0x01, 0x87, 0x4A, 0x53, 0x36, 0x92, 0xC8, 0xC1, 0xA4, 0x80, 0x3D, - 0x59, 0x15, 0xD7, 0x8E, 0x46, 0x37, 0xB5, 0x3C, 0x46, 0x37, 0x0F, 0xB5, - 0xEB, 0x39, 0x57, 0xC5, 0xB8, 0x0E, 0x38, 0x96, 0x3A, 0xB2, 0x07, 0xA3, - 0x0E, 0x63, 0xC1, 0x9B, 0xD2, 0x49, 0xD9, 0x1B, 0x23, 0xD4, 0xF9, 0x47, - 0xE9, 0x14, 0xD2, 0xDE, 0x27, 0x93, 0x7A, 0x6A, 0xC6, 0x2F, 0x3A, 0xE6, - 0xD0, 0xFC, 0x8A, 0xA0, 0xDB, 0xE7, 0xD6, 0xE5, 0x9F, 0x56, 0x2B, 0xC7, - 0x91, 0x22, 0xEC, 0xB5, 0x46, 0xDF, 0xC2, 0xDE, 0x12, 0x0C, 0x09, 0x82, - 0xD3, 0x0A, 0xBA, 0x88, 0x01, 0x85, 0xF2, 0xC3, 0x5E, 0x3B, 0x68, 0x64, - 0x3C, 0x26, 0xCC, 0x56, 0x5B, 0x9D, 0x64, 0x5D, 0xFB, 0x31, 0x3D, 0x8A, - 0x9E, 0xD5, 0x65, 0x90, 0x7E, 0xF5, 0x51, 0xC6, 0xAA, 0x6B, 0xEE, 0x78, - 0x5D, 0xBF, 0x3A, 0x2A, 0xB9, 0x39, 0xED, 0x52, 0x4A, 0x51, 0xEF, 0x38, - 0x85, 0xA4, 0x71, 0x18, 0x12, 0x1E, 0x57, 0xEB, 0x40, 0x1F, 0xD1, 0x6C, - 0x6E, 0x20, 0x73, 0x9C, 0x2E, 0x4B, 0xE9, 0xC0, 0x96, 0x42, 0x37, 0xC0, - 0x45, 0x0D, 0x2D, 0x13, 0xA7, 0xE1, 0x3D, 0x4A, 0x1A, 0x47, 0xA0, 0xC2, - 0x7A, 0x04, 0x62, 0x29, 0xC4, 0x99, 0x5C, 0xA4, 0x94, 0xCA, 0x18, 0x75, - 0xA0, 0xA1, 0x93, 0x82, 0x42, 0xCE, 0x28, 0x89, 0xA7, 0xCB, 0x28, 0xC3, - 0xD3, 0x06, 0x3A, 0xB6, 0x93, 0x36, 0x22, 0x62, - }; - uint8_t qh_host[] = { - 0x0E, 0x24, 0x6D, 0x09, 0xB2, 0x5D, 0xA0, 0x45, 0xBF, 0x75, 0xC8, 0x2D, - 0x01, 0x5E, 0xB2, 0xF4, 0xCB, 0x8E, 0x8F, 0x6E, 0x21, 0x86, 0xE1, 0x12, - 0x7C, 0x30, 0x8A, 0xD1, 0x22, 0x38, 0xD3, 0x3C, 0x02, 0x9B, 0x08, 0x1F, - 0xD6, 0x8E, 0x37, 0xFE, 0x7A, 0x96, 0x20, 0x4B, 0x87, 0xCD, 0x26, 0xAD, - 0x6E, 0x23, 0x06, 0x0B, 0x4F, 0xE4, 0xD4, 0x2C, 0x8E, 0x57, 0xBF, 0x7E, - 0x34, 0x65, 0x9F, 0x4E, 0x4C, 0xE3, 0x39, 0x6A, 0xA4, 0x18, 0x57, 0xE0, - 0xA7, 0xB8, 0xD4, 0x25, 0x48, 0x9B, 0xB5, 0x23, 0x92, 0xA1, 0xA9, 0x0A, - 0x39, 0xC4, 0x1A, 0xE3, 0x07, 0xAF, 0x69, 0xD7, 0x7F, 0xF6, 0x7C, 0x16, - 0x30, 0x49, 0xE4, 0x95, 0xBA, 0x59, 0x94, 0x72, 0x28, 0xDB, 0xD0, 0x05, - 0x75, 0xD9, 0x9F, 0xD2, 0x4D, 0x59, 0x18, 0xC0, 0x8F, 0x52, 0xCD, 0x00, - 0xD4, 0xEF, 0xEF, 0xC5, 0x8A, 0x99, 0xFD, 0xE2, 0xF5, 0xBA, 0x5C, 0x00, - 0xE7, 0x80, 0xF0, 0x59, 0xBD, 0x79, 0x79, 0xFE, 0x74, 0xD8, 0x59, 0x1F, - 0xC2, 0x8D, 0x31, 0x76, 0x9F, 0x3D, 0xA6, 0x19, 0x24, 0x2B, 0xEF, 0xA2, - 0x5B, 0x88, 0x90, 0xF0, 0xA2, 0x40, 0x93, 0x12, 0xF5, 0x3B, 0x11, 0x45, - 0x82, 0xA8, 0x58, 0xEC, 0x55, 0x72, 0x6F, 0xE2, 0x7B, 0xF6, 0x72, 0x59, - 0x52, 0x88, 0x8F, 0xF2, 0x91, 0x3D, 0xC0, 0x28, 0xB8, 0x2B, 0x64, 0x5E, - 0x90, 0xB7, 0x23, 0xD3, 0x4E, 0x76, 0x1A, 0xDF, 0x20, 0x5A, 0x91, 0x2B, - 0xC4, 0x06, 0x93, 0x70, 0xE7, 0x67, 0x0A, 0x0D, 0x39, 0x59, 0x0F, 0x3C, - 0x2A, 0x77, 0x77, 0x78, 0xCC, 0x01, 0x10, 0xBE, 0x81, 0x1F, 0xFA, 0x2A, - 0xBE, 0xDB, 0x5B, 0x2F, 0xC7, 0xEA, 0x1C, 0xDE, 0x00, 0xDF, 0x49, 0x64, - 0x0D, 0xBB, 0xA0, 0xEA, 0x2D, 0x76, 0xA2, 0x57, 0x45, 0x1E, 0x0D, 0xE2, - 0x55, 0x96, 0x42, 0xB7, - }; - int8_t scale_codes[] = { - 72, 16, 40, 102, 55, 121, 127, 86, 72, 56, 73, 88, - 28, 61, 23, 116, 34, 127, 57, 38, 101, 81, 114, 98, - 14, 111, 61, 46, 111, 116, 19, 67, 119, 94, 120, 51, - 38, 121, 95, 127, 46, 33, 64, 15, 94, 114, 84, 65, - 99, 33, 127, 12, 14, 56, 118, 76, 92, 27, 91, 64, - 93, 111, 95, 65, - }; - uint16_t scale_step[] = { - 0x126D, 0x12CE, 0x12F6, 0x1306, - }; - uint16_t A_host[] = { - 0xBF90, 0x3F1F, 0xBF60, 0xBFE3, 0xBF4F, 0xBFBD, 0x3F2B, 0x3F2B, - 0xBD9A, 0x3E28, 0x3CE1, 0x3F15, 0x3E71, 0x3F89, 0xBF06, 0xBF39, - 0xBE53, 0xBF4F, 0x3EEE, 0x3EC8, 0xBFAC, 0x3F42, 0x3F3B, 0xBD89, - 0x3F0E, 0xBFB4, 0x3DE3, 0x3D10, 0xBF11, 0x3F6A, 0x3E99, 0x3E35, - 0xBEB0, 0x3E6B, 0x3EE0, 0x3E8A, 0xBE4D, 0x3F82, 0xBED5, 0xBCFC, - 0xBF7C, 0x3E4F, 0x3ED5, 0xBF9A, 0xBE08, 0xBE96, 0x3EEE, 0x3F14, - 0xBD6D, 0x3F96, 0xBFC2, 0x3D43, 0xBC29, 0x3FB8, 0x3E81, 0x3F3E, - 0xBF8D, 0x3E1F, 0xBFB3, 0x3D27, 0xBF8F, 0x3F34, 0xBF2B, 0x3E00, - 0x3ECF, 0xBF28, 0x3D56, 0x3EAE, 0xBE5A, 0x3FC8, 0xBF68, 0xBFC8, - 0x3D47, 0x3F7E, 0x4008, 0xBE3B, 0xBF47, 0x3FE3, 0xBF01, 0x3D68, - 0xBF0C, 0xBF31, 0xBF28, 0x3F32, 0xBEEF, 0xBE2C, 0x3F42, 0x3F98, - 0xBF21, 0xBD86, 0x3F1E, 0x4000, 0x3E97, 0x3F08, 0xBF06, 0x3DCC, - 0xBDEF, 0x3EBF, 0xBF83, 0x3F8A, 0x3F93, 0xBE31, 0x3D82, 0xBFA3, - 0xBF1F, 0xBE74, 0x3DB7, 0x3FE6, 0xC004, 0x3EA5, 0xBF92, 0x3F9F, - 0xBFC2, 0x3F9A, 0xBEE4, 0x3D4B, 0xBE32, 0x3F44, 0x3F1E, 0x4033, - 0x3EDD, 0x4008, 0x4023, 0x3E99, 0xBF15, 0x3EA7, 0x3E08, 0x3E00, - 0xBFBF, 0x3C95, 0x3EB1, 0x3D00, 0x3ECD, 0x3EDA, 0x3F57, 0x3F9F, - 0x3F06, 0xBFD4, 0x3F0C, 0xBFB3, 0xBFFD, 0x3F9E, 0xBF0C, 0xBD8E, - 0x3F3B, 0xBE1A, 0xBF35, 0x3EEF, 0x3E8F, 0x3FAB, 0x3F94, 0xBD67, - 0x3EDC, 0xBF6F, 0xBFC1, 0x3DD4, 0x3FC3, 0x3EBB, 0x3FB7, 0x3EBB, - 0x3FD8, 0x3F81, 0x3F1F, 0x4012, 0x3FCA, 0xBE9F, 0x3FBF, 0x3F80, - 0xBFB8, 0x3D31, 0xBFD3, 0xBFBB, 0x3EB3, 0xBDD9, 0xBF3E, 0x3F94, - 0x3FB9, 0x3E4B, 0xBF8D, 0xC002, 0x4008, 0xBE51, 0x3EF0, 0x3F9A, - 0x3EA1, 0x3FCB, 0x3F61, 0x4030, 0xBF05, 0xC013, 0xBF77, 0xBFF2, - 0xBEB1, 0x3F1F, 0x3F31, 0x3F3F, 0xBF6E, 0xBF2F, 0x3F21, 0x3FF9, - 0x3EC7, 0x3F62, 0xBF8E, 0xBF99, 0x3FB3, 0x3EFF, 0x3F2E, 0x3F8E, - 0x3F4E, 0xBE1D, 0xBE8D, 0x3E92, 0x3F81, 0xBF97, 0x3FD6, 0xBFE6, - 0xBFA6, 0xBF4F, 0x4008, 0x3EB7, 0x3EB1, 0xBE3E, 0xBE23, 0x3F9C, - 0xBF8F, 0xBEEB, 0xBFA6, 0x3F8A, 0x3DEE, 0x3E8A, 0xBFF1, 0xBF0C, - 0x3E64, 0xBE94, 0x3F2B, 0x3FBF, 0x3CC7, 0xBEFF, 0x3F3D, 0x3FA3, - 0xBE1E, 0xBE92, 0xBF21, 0xBE7F, 0x3F4E, 0xBF8D, 0xBF3E, 0x3F53, - 0xBF73, 0x3F29, 0xBF30, 0xBE95, 0x3FEF, 0xBF8D, 0xBE14, 0xBEB1, - 0x3F53, 0x3F71, 0xC00D, 0x3ED7, 0xBE7F, 0xBFAC, 0xBE9A, 0xC00C, - 0xBE9C, 0xBE56, 0x3F2F, 0xBE05, 0x3F9B, 0x401F, 0xBF3F, 0xBFDC, - 0x3FAF, 0x400D, 0x3E10, 0x3F39, 0x3FBC, 0xBFBA, 0xBFD1, 0xBF3E, - 0x3F87, 0xBFCC, 0x3F21, 0xBFC5, 0x3E61, 0x3F61, 0xBD5F, 0xBEE5, - 0x3F94, 0x3FA9, 0xBF8A, 0xBF26, 0x3F96, 0x3EC8, 0xBF8D, 0x3F58, - 0xBEF0, 0xBF24, 0x3DBE, 0x3F9E, 0x4041, 0xBEF9, 0x3E79, 0xBE20, - 0xBFD8, 0x3FAB, 0x4001, 0xBE70, 0xBF8A, 0xBF4D, 0x3F64, 0xBE90, - 0xBEF1, 0x3FF9, 0x3F32, 0xBF13, 0xBF76, 0xBE66, 0xBFB4, 0x3F88, - 0x3ED6, 0x3F37, 0x3EBB, 0xBF56, 0xBFF9, 0x3FEB, 0xBED6, 0xBE01, - 0xBFB9, 0x3AFD, 0x3FA8, 0xBF07, 0xBD0C, 0x3FBC, 0x4013, 0x3F9B, - 0x3D5B, 0xBEBE, 0x3E88, 0x3E27, 0xBF40, 0x3E92, 0x3ED7, 0x3F99, - 0xBDA4, 0x3F9A, 0x3FF3, 0x3FC6, 0x3E57, 0x3D3A, 0x3E27, 0xBD70, - 0xBF47, 0x3F83, 0x3F20, 0x3E95, 0x3F0F, 0xBF3F, 0x3D74, 0x3F8F, - 0x3C40, 0x3C6E, 0xBF39, 0xBE4E, 0xBF5E, 0x3FA0, 0xBF56, 0x3E66, - 0xBF58, 0x3F65, 0x3E6A, 0x3E2E, 0x3F48, 0xBF08, 0xBF6D, 0xBE50, - 0x3EF7, 0x3EEE, 0xBF5B, 0xBFC8, 0xBE21, 0x3F71, 0xBE19, 0xBE13, - 0x406E, 0x3FDE, 0xBF3E, 0x3E14, 0x3EA2, 0x3F21, 0xBEF1, 0x3E5C, - 0xBED3, 0xBFEA, 0xBD0B, 0x3EBB, 0xBCC0, 0xBEE7, 0xC00D, 0x3FBC, - 0xBFD7, 0x3F94, 0x3FE8, 0xBFC9, 0x3FB3, 0x3E8D, 0xBEAB, 0xBF2E, - 0x3FA5, 0x3EFF, 0xBF07, 0x3DDA, 0xBF79, 0x3E67, 0xBF20, 0x3FD5, - 0x3E4C, 0x3F78, 0xBF08, 0x3FD7, 0xBF5A, 0xC003, 0xBFAD, 0xBCED, - 0x3F4D, 0x3E87, 0x3F31, 0x3ECB, 0x3F0F, 0xBDD8, 0x3E8E, 0x3F86, - 0xBFEE, 0xBF3C, 0x3F5F, 0xBF42, 0x3FD1, 0xC016, 0xBE8E, 0x3E84, - 0x3F0F, 0x3FCB, 0xBEDD, 0x3EC2, 0xBF58, 0xBD15, 0x3DC9, 0x3F77, - 0xBF45, 0x3F96, 0xBFCC, 0x3F4F, 0x3ED5, 0xBE67, 0xBEA5, 0xBFE4, - 0x3EF0, 0xBF63, 0xBDA9, 0x3F50, 0x402D, 0xBF08, 0x3FC0, 0xBFD7, - 0xC01F, 0xBF9C, 0xBE83, 0x3FB3, 0x3DB9, 0x3DB0, 0xBEE4, 0x3EB7, - 0xBFEF, 0xBF5B, 0x3F89, 0x3DC5, 0xBF2F, 0x3ECC, 0x3E84, 0x3E4C, - 0xBE22, 0x3F3D, 0xBE80, 0xBE4A, 0xBFFE, 0xBF21, 0x3F1C, 0xBE11, - 0xBF2F, 0x3FB1, 0xBDA7, 0xBF2B, 0x3EEC, 0xBECE, 0x3F10, 0xBF6F, - 0xBEF1, 0xBF86, 0x3EDA, 0xBFC7, 0x3FC0, 0xBF81, 0xBC53, 0xBF4A, - 0x3FBA, 0x3E6E, 0xBBC8, 0xBCBD, 0x3F8F, 0x3E70, 0xBFD1, 0x3FA6, - }; - uint16_t expected[] = { - 0x3FC4, 0xC211, 0xC112, 0xC1F9, 0x4245, 0x4205, 0xBF33, 0x4201, - }; - // clang-format on - - Tensor* output = setup_and_run( - M, N, K, gs, ql_host, qh_host, scale_codes, scale_step, A_host); - ASSERT_NE(output, nullptr); - EXPECT_EQ(output->size(0), M); - EXPECT_EQ(output->size(1), N); - check_bf16_output(output, expected, M * N, 0.5f); -} - -// Q6KMultiSuperBlock: M=1, N=6, K=512, gs=16, ng=32, n_super=2 -TEST_F(AOTITorchInt6PlainMMTest, Q6KMultiSuperBlock) { - int64_t M = 1, N = 6, K = 512, gs = 16; - int64_t ng = K / gs; // 32 - int64_t n_super = K / 256; // 2 - // clang-format off - uint8_t ql_host[] = { - 0xB5, 0x8C, 0x9F, 0x5B, 0x0F, 0x10, 0x7C, 0xCD, 0x96, 0x42, 0x5E, 0x42, - 0xCB, 0xCA, 0xED, 0x42, 0x77, 0x19, 0xE7, 0x60, 0x9D, 0x79, 0xDF, 0x96, - 0x01, 0x1C, 0x88, 0x3D, 0xAF, 0x89, 0x7E, 0x63, 0x15, 0x39, 0x84, 0xB1, - 0xAC, 0x04, 0x3D, 0x29, 0x0F, 0x94, 0x2F, 0x77, 0x9A, 0x68, 0x39, 0xE7, - 0xD7, 0xB4, 0x95, 0x63, 0x08, 0xA2, 0xF7, 0x7F, 0x79, 0xC3, 0xEA, 0x0E, - 0x78, 0x17, 0xD1, 0x03, 0xD8, 0x6D, 0x54, 0x26, 0xBC, 0xB5, 0xD7, 0xB8, - 0xBC, 0xE4, 0x74, 0xD7, 0x94, 0xC0, 0x02, 0xAD, 0xC7, 0x71, 0x89, 0xCA, - 0x04, 0x1C, 0xCE, 0x89, 0x32, 0x21, 0xCC, 0x27, 0x6A, 0x90, 0x62, 0x26, - 0x77, 0x60, 0x5F, 0xA1, 0x64, 0x60, 0xD5, 0x21, 0x51, 0xA4, 0x70, 0x8D, - 0x59, 0x07, 0x39, 0x19, 0xDE, 0x44, 0x6B, 0x8D, 0xFB, 0xB8, 0x29, 0x57, - 0x45, 0x85, 0x85, 0xBD, 0x11, 0xDF, 0xE8, 0x7D, 0xC0, 0xDA, 0x43, 0xE2, - 0x0E, 0xFD, 0x3B, 0xD5, 0x21, 0x34, 0x60, 0xD0, 0x7C, 0xF2, 0xEC, 0x8A, - 0x3D, 0x80, 0xD4, 0xF2, 0x09, 0x3C, 0x18, 0x4A, 0xFC, 0xBE, 0x3E, 0xB3, - 0xBA, 0x76, 0x53, 0xEF, 0x3C, 0x42, 0x04, 0x33, 0x38, 0xE5, 0xF6, 0x57, - 0xFF, 0x71, 0xA0, 0x2A, 0x28, 0x1E, 0x04, 0x4C, 0x1E, 0x37, 0xC1, 0x66, - 0xA9, 0x6D, 0x69, 0x00, 0xE2, 0x9C, 0x06, 0x76, 0x30, 0x09, 0x3C, 0xAF, - 0x4B, 0xE7, 0x5F, 0xCE, 0x83, 0x08, 0x6E, 0x7F, 0xEC, 0xFD, 0x59, 0x94, - 0x5A, 0xE2, 0xCD, 0x65, 0x6B, 0xCC, 0xB8, 0xCA, 0x77, 0xD7, 0xCB, 0x62, - 0x0B, 0x5D, 0x12, 0xAB, 0x58, 0x49, 0xB9, 0x1B, 0xF2, 0x0E, 0x4A, 0x7E, - 0x60, 0x42, 0x63, 0x7C, 0x36, 0x0F, 0x6D, 0x4D, 0x67, 0xA2, 0x9C, 0xFA, - 0x95, 0x99, 0xD8, 0xD6, 0xC4, 0xA2, 0x9B, 0x04, 0xE0, 0xA3, 0x94, 0x93, - 0x21, 0xD5, 0xD4, 0x0E, 0x8B, 0x2C, 0xD3, 0x99, 0x44, 0x28, 0x61, 0x3F, - 0x8E, 0x79, 0xF0, 0x25, 0xC2, 0xBE, 0xAA, 0x8A, 0xA5, 0x0F, 0x5C, 0xD9, - 0xA8, 0x6B, 0x06, 0x4E, 0xB7, 0xF3, 0x0B, 0xBE, 0x61, 0x0A, 0x16, 0x46, - 0xB2, 0x5C, 0x64, 0xE2, 0x29, 0x7E, 0x05, 0xD7, 0xDB, 0xAB, 0x8E, 0x8E, - 0x08, 0x7E, 0xC2, 0x70, 0x1C, 0x1A, 0xCA, 0xB9, 0x15, 0x95, 0x6A, 0x94, - 0xE8, 0xA7, 0x15, 0x8A, 0x50, 0xC3, 0xD9, 0x40, 0x68, 0x2B, 0x34, 0xCA, - 0xC2, 0x00, 0xFD, 0xA4, 0xCA, 0xE2, 0x5B, 0xF0, 0xC0, 0xCB, 0xBF, 0xC3, - 0xE8, 0x5F, 0xD3, 0x41, 0xC7, 0x23, 0xA2, 0x2E, 0x66, 0x0A, 0x1A, 0x5E, - 0x6A, 0x5B, 0x88, 0x55, 0xE7, 0x95, 0x1E, 0x3B, 0xED, 0xEB, 0x39, 0x33, - 0x16, 0x3F, 0xC0, 0x05, 0xD5, 0xE2, 0x67, 0x4C, 0x0A, 0xFE, 0xDB, 0x42, - 0x8D, 0xC7, 0xD6, 0x77, 0xBA, 0xAC, 0x1C, 0xB7, 0xC7, 0x83, 0x03, 0xEC, - 0x36, 0xE0, 0x56, 0xFB, 0x69, 0x4C, 0x6D, 0x26, 0xDD, 0x42, 0xFA, 0x21, - 0x93, 0xDD, 0x3B, 0xF6, 0x07, 0x3C, 0xDB, 0x3B, 0x6F, 0xA8, 0xB6, 0x15, - 0xFD, 0x23, 0xD6, 0xF3, 0x6A, 0x7D, 0x82, 0x0C, 0xB1, 0x68, 0x00, 0x21, - 0x77, 0xE4, 0xA4, 0x10, 0x80, 0x65, 0xE2, 0x45, 0x03, 0x26, 0xB1, 0x9B, - 0xC4, 0xD4, 0x90, 0x78, 0x67, 0xA1, 0x27, 0x54, 0x6D, 0x27, 0xDC, 0x0E, - 0x25, 0xCD, 0xC7, 0x9F, 0x79, 0x11, 0x4D, 0x56, 0x6B, 0x14, 0xE1, 0x5F, - 0x21, 0xC6, 0x2B, 0xF3, 0x23, 0x03, 0x0A, 0x51, 0x50, 0x8E, 0xCB, 0x08, - 0x0D, 0x90, 0x58, 0xB9, 0x3E, 0x4E, 0xDE, 0x90, 0x8D, 0x36, 0x9A, 0x09, - 0x18, 0x66, 0x31, 0x7F, 0xDB, 0x3F, 0xFC, 0x2D, 0xDF, 0x03, 0x2A, 0xA8, - 0xE2, 0x09, 0xC1, 0x3E, 0x49, 0x88, 0xE8, 0x2E, 0x68, 0x94, 0x0C, 0x5A, - 0xAB, 0x65, 0x71, 0x56, 0x7B, 0xF1, 0x79, 0x65, 0x67, 0x7A, 0x5F, 0x86, - 0xC7, 0x24, 0x4A, 0x00, 0x3E, 0xEC, 0xBC, 0x95, 0x30, 0x56, 0xA1, 0xEB, - 0x1B, 0x48, 0x07, 0x1F, 0x17, 0xEE, 0xBD, 0x5D, 0xC4, 0x5A, 0xA6, 0x52, - 0xFA, 0xC3, 0x39, 0x59, 0xC1, 0xB9, 0x77, 0xF0, 0x8C, 0x6C, 0x21, 0xA0, - 0x44, 0x16, 0x90, 0x94, 0xE3, 0xF0, 0xFA, 0x15, 0x91, 0xA6, 0x0C, 0xD2, - 0x8D, 0xAC, 0x73, 0xF2, 0x5F, 0x06, 0x24, 0x3F, 0x1E, 0xE0, 0x75, 0x20, - 0xA2, 0xB0, 0x59, 0x11, 0x36, 0x2E, 0x11, 0xF2, 0xD7, 0x25, 0xCE, 0x9B, - 0xD4, 0xFF, 0xF7, 0xB3, 0xD5, 0x20, 0x1F, 0xE4, 0x49, 0x96, 0xB3, 0x8A, - 0x8C, 0xD7, 0xEC, 0xB1, 0xD8, 0x7A, 0x98, 0x22, 0xF5, 0xEA, 0x35, 0x95, - 0xDA, 0x72, 0xFD, 0x4D, 0x61, 0x9E, 0xDF, 0x18, 0x18, 0x26, 0x16, 0x38, - 0x07, 0xD7, 0x0E, 0xA0, 0x37, 0xB9, 0xB5, 0x2A, 0x15, 0xBB, 0xA2, 0x5B, - 0x3A, 0xBF, 0x63, 0xD1, 0xCB, 0x18, 0x68, 0x64, 0x9C, 0xF5, 0xB4, 0xAF, - 0xB7, 0x02, 0xBE, 0xC2, 0x50, 0x9F, 0x4A, 0x1F, 0x54, 0x92, 0x3C, 0x1B, - 0xD5, 0xAF, 0xDE, 0x11, 0x7E, 0x21, 0xD6, 0x70, 0x7D, 0x4C, 0xB3, 0x72, - 0x58, 0xC2, 0x4F, 0x29, 0x32, 0x5D, 0x69, 0xB4, 0xDB, 0x9C, 0x20, 0x38, - 0x7A, 0xD3, 0xB9, 0x2E, 0x83, 0xF8, 0x0C, 0x62, 0xAE, 0xF8, 0x63, 0xBA, - 0x4F, 0x79, 0xE6, 0x23, 0x9C, 0xBA, 0x51, 0xA5, 0x96, 0x64, 0xB7, 0x58, - 0x27, 0x5A, 0xD3, 0x4E, 0xCB, 0x95, 0x2A, 0xB4, 0xA0, 0xAD, 0xB3, 0xAC, - 0x25, 0xAC, 0xC6, 0x82, 0x0E, 0x2D, 0x46, 0x8B, 0x34, 0x4C, 0x9D, 0x38, - 0xF5, 0xEB, 0x2D, 0x0B, 0x97, 0xCC, 0x8D, 0x1B, 0x4D, 0xC6, 0x61, 0x39, - 0xFE, 0x34, 0x37, 0xDC, 0x8C, 0x3A, 0xF4, 0xFB, 0xC4, 0xA4, 0x36, 0xF0, - 0xA4, 0xA6, 0x94, 0x3E, 0x60, 0xF2, 0x29, 0x9F, 0xA0, 0x21, 0x64, 0x60, - 0xCD, 0x0D, 0x0B, 0x2A, 0x46, 0xAA, 0x51, 0x8D, 0xC2, 0x2D, 0xB6, 0x15, - 0x2E, 0x67, 0x47, 0xD5, 0x27, 0x73, 0x14, 0xD5, 0x8A, 0xA3, 0x01, 0x88, - 0xE4, 0x24, 0x86, 0xF5, 0xD7, 0xA7, 0x79, 0xAE, 0xED, 0xAB, 0x8A, 0xF8, - 0x49, 0xBB, 0xD9, 0x26, 0xD3, 0x4A, 0x3C, 0x25, 0xD6, 0xEE, 0x69, 0x0B, - 0x74, 0x3F, 0x67, 0x05, 0xA2, 0x99, 0x3C, 0x52, 0x1F, 0x97, 0x78, 0x13, - 0xE7, 0xD3, 0xD7, 0x9B, 0x3B, 0xA8, 0x47, 0xED, 0xF9, 0xE8, 0xA9, 0x65, - 0x30, 0x53, 0x1B, 0xE8, 0x5A, 0xFE, 0x32, 0x3A, 0xDF, 0x67, 0x78, 0xAA, - 0xF8, 0x2C, 0x77, 0xD8, 0xE7, 0xE2, 0x3B, 0x07, 0x4A, 0xF0, 0xE8, 0xAE, - 0xA9, 0x24, 0xFD, 0x87, 0x43, 0x70, 0xAF, 0x7E, 0x89, 0x17, 0x84, 0x3F, - 0xE9, 0x24, 0xCE, 0x09, 0x21, 0xBD, 0x8A, 0x63, 0xCB, 0xF9, 0x7D, 0xB8, - 0x6F, 0xF9, 0x99, 0xEF, 0x28, 0x61, 0xBC, 0xE4, 0xA9, 0x13, 0x3F, 0x6F, - 0xAE, 0x45, 0xB9, 0xD6, 0xA8, 0xFE, 0x04, 0x3A, 0x9F, 0x83, 0x82, 0x47, - 0x6E, 0xB4, 0x57, 0x45, 0xE7, 0x27, 0xAD, 0x8A, 0x1F, 0x17, 0xAD, 0xF5, - 0x81, 0xB1, 0x2B, 0x02, 0x3A, 0xAA, 0x43, 0x77, 0x3E, 0xFF, 0xCC, 0x7E, - 0x56, 0xC7, 0x72, 0xA5, 0x5F, 0x8E, 0x17, 0xB4, 0xFC, 0x2A, 0xDE, 0x41, - 0x7C, 0x46, 0x31, 0xBE, 0x0E, 0x1E, 0xC4, 0x00, 0xDD, 0xEB, 0x42, 0xDA, - 0xE8, 0x37, 0x1B, 0x0A, 0xA5, 0x7C, 0xA1, 0xE4, 0x48, 0xC2, 0x31, 0x1F, - 0x7D, 0x7C, 0x70, 0xA2, 0x43, 0xD0, 0xE4, 0x44, 0x24, 0x2B, 0xF4, 0x3C, - 0x20, 0x53, 0x01, 0xF9, 0x83, 0x6C, 0xCB, 0x7B, 0x0E, 0x08, 0x4D, 0x88, - 0x3B, 0xEA, 0x21, 0xAC, 0x7D, 0x76, 0xCF, 0x03, 0x33, 0x3F, 0xB1, 0xF2, - 0x0D, 0xC3, 0xD7, 0x85, 0xF0, 0x4E, 0xD1, 0xA4, 0xCC, 0xE6, 0xE8, 0xBC, - 0x08, 0x60, 0x86, 0xFD, 0x39, 0xF9, 0x7A, 0x1E, 0x0D, 0xDE, 0xF3, 0x91, - 0x93, 0x3D, 0x62, 0x71, 0xD3, 0x7E, 0x30, 0xF3, 0x98, 0xE0, 0x59, 0xFC, - 0x43, 0xA4, 0x80, 0xEE, 0xE8, 0xFE, 0x7B, 0xA8, 0xDE, 0x23, 0xAE, 0x13, - 0xFD, 0x99, 0x54, 0x49, 0x68, 0xBD, 0x05, 0xD8, 0x66, 0x8E, 0xA6, 0x04, - 0xC4, 0x45, 0x29, 0x20, 0x63, 0xF8, 0x27, 0x31, 0xF1, 0xFF, 0x2E, 0x32, - 0xD9, 0x62, 0xCA, 0x92, 0x11, 0x45, 0x48, 0x7E, 0xF2, 0x03, 0xE5, 0x28, - 0x49, 0x0E, 0xFF, 0x30, 0x4B, 0x69, 0x67, 0x08, 0xE8, 0x50, 0x02, 0x8C, - 0x16, 0x59, 0xC0, 0x80, 0x31, 0x7E, 0x15, 0xE2, 0x36, 0x12, 0x68, 0xA5, - 0x19, 0x5C, 0x1F, 0x6A, 0xCF, 0xC3, 0x50, 0xDC, 0x7E, 0x23, 0xCC, 0x49, - 0xEB, 0x4D, 0xF4, 0x99, 0xF9, 0xB4, 0x8E, 0xD1, 0x66, 0xED, 0x27, 0x54, - 0x37, 0xC7, 0xD1, 0x76, 0x49, 0xAE, 0x78, 0x0C, 0x14, 0x79, 0xC8, 0xD6, - 0x19, 0x8F, 0x43, 0xBB, 0x27, 0x62, 0x21, 0x12, 0xB5, 0x5D, 0x2B, 0xBB, - 0x33, 0x03, 0x5D, 0xB3, 0x01, 0x6E, 0x34, 0xA1, 0xB5, 0x95, 0xAE, 0x0D, - 0xA6, 0x7F, 0x8A, 0xFF, 0x2F, 0x58, 0xC0, 0x8F, 0x2A, 0x8A, 0x49, 0x21, - 0x61, 0xDE, 0x05, 0x16, 0x0E, 0x31, 0xEE, 0x5A, 0x74, 0xB6, 0x24, 0x39, - 0x82, 0xDE, 0xBA, 0xD4, 0x52, 0xA1, 0xEE, 0xD2, 0x5B, 0xC3, 0x33, 0x2C, - 0xAB, 0xB8, 0xB1, 0x69, 0xE4, 0xC3, 0x84, 0x3B, 0x78, 0x31, 0xCE, 0x04, - 0x32, 0xFB, 0x24, 0xB0, 0xCA, 0xFE, 0x5A, 0x78, 0x88, 0x64, 0xC3, 0xF4, - 0x15, 0xC4, 0x25, 0x04, 0xED, 0x8B, 0x3C, 0xE7, 0xDE, 0x30, 0xE5, 0x5B, - 0x28, 0xE1, 0x9C, 0x13, 0x82, 0x9F, 0xEB, 0xD9, 0xC9, 0xAD, 0x0D, 0xAF, - 0x1C, 0xE0, 0x13, 0x60, 0x3A, 0xEB, 0xA5, 0xAF, 0x15, 0x94, 0xE4, 0x70, - 0xB7, 0x5D, 0x4A, 0x73, 0x5D, 0x90, 0xBB, 0x7E, 0x0B, 0xE8, 0x77, 0xE4, - 0xE1, 0xC5, 0x0C, 0xEE, 0x2C, 0xF3, 0x77, 0x07, 0xD2, 0x65, 0xA2, 0xF3, - 0x8E, 0xCD, 0xB5, 0xFF, 0x1E, 0x62, 0x5A, 0x5A, 0xCB, 0x86, 0xEC, 0x2E, - 0xD7, 0x2C, 0xE9, 0x46, 0xF1, 0x87, 0xC9, 0x97, 0x4D, 0x7E, 0x57, 0xB6, - 0x6A, 0x31, 0xAD, 0x70, 0x8D, 0xD2, 0xF0, 0x43, 0x67, 0x81, 0x8A, 0x09, - 0x06, 0x47, 0x16, 0xDB, 0x62, 0xA3, 0xB2, 0x50, 0x91, 0xEB, 0x6A, 0xF8, - 0x29, 0x96, 0x16, 0xC0, 0x56, 0x4F, 0xE5, 0xE8, 0x64, 0x50, 0xB6, 0x57, - 0x09, 0x04, 0x6D, 0x49, 0x04, 0xF5, 0x15, 0x50, 0x55, 0xEC, 0x60, 0x10, - 0x73, 0x06, 0x96, 0x5C, 0x1A, 0xEB, 0xE7, 0xD4, 0xFC, 0x67, 0xF2, 0x0E, - 0xCA, 0x45, 0xFB, 0xD8, 0xE2, 0x6B, 0xA8, 0x81, 0x3B, 0xE9, 0x67, 0x65, - 0x09, 0xBB, 0x12, 0x46, 0x6C, 0x63, 0x0D, 0xF2, 0x9E, 0xEF, 0x72, 0x4A, - 0xA1, 0x62, 0xF2, 0x09, 0x7C, 0x85, 0x94, 0x1E, 0x2B, 0xB8, 0x7C, 0xDB, - 0x28, 0x26, 0x9D, 0xF2, 0xAD, 0x2F, 0x70, 0x1F, 0x1F, 0x1B, 0xC4, 0x01, - 0x5B, 0xBE, 0xCE, 0x40, 0xE1, 0x59, 0x03, 0x07, 0x8E, 0x50, 0x40, 0xC4, - 0x6A, 0xC7, 0x29, 0x2F, 0x4D, 0xD9, 0xB2, 0x03, 0x62, 0x50, 0x61, 0x2F, - 0x5E, 0x0E, 0x9B, 0x11, 0x3D, 0x9F, 0xCA, 0xC9, 0xCC, 0xA5, 0xEF, 0x9D, - 0xB6, 0xD7, 0xBB, 0x76, 0xEE, 0xD6, 0x77, 0xB2, 0x99, 0xD0, 0x4E, 0xB4, - 0xE9, 0x49, 0x1D, 0x95, 0xA9, 0x26, 0xDE, 0xBE, 0xA0, 0x32, 0x89, 0xA8, - 0xB4, 0x28, 0x46, 0x46, 0xCF, 0xE5, 0x07, 0x5E, 0x5B, 0x83, 0xE2, 0xC0, - 0x35, 0xEC, 0xC8, 0x8A, 0x5B, 0x8C, 0x7F, 0xCC, 0x94, 0x55, 0x70, 0x41, - 0xD4, 0xD8, 0xC2, 0xB2, 0xC6, 0x7F, 0xC1, 0x35, 0xF8, 0x04, 0xC4, 0x0C, - 0x5A, 0xE9, 0xA7, 0xDE, 0xEA, 0x59, 0x9C, 0xFF, 0xDA, 0xC5, 0xFC, 0x9C, - }; - uint8_t qh_host[] = { - 0x32, 0x84, 0x7C, 0xD8, 0x02, 0xC0, 0x65, 0x07, 0xB9, 0x70, 0x0C, 0x1E, - 0x5D, 0xF2, 0xD4, 0x02, 0xCD, 0xAC, 0x60, 0x89, 0x93, 0x75, 0x24, 0xC5, - 0x69, 0x98, 0x8D, 0x3D, 0xB7, 0x42, 0xB4, 0x01, 0x33, 0x04, 0x24, 0x1C, - 0x13, 0x01, 0xA6, 0x2F, 0x35, 0x91, 0xA9, 0x18, 0x53, 0x72, 0x93, 0x1B, - 0x4C, 0xED, 0xE6, 0x79, 0x8E, 0xC7, 0xF3, 0xD9, 0x24, 0x07, 0x18, 0xAB, - 0xCC, 0xCC, 0x6E, 0x43, 0x09, 0xC6, 0xB6, 0x15, 0x39, 0x6C, 0xF4, 0x17, - 0xDA, 0x9A, 0x46, 0xA5, 0x02, 0x74, 0xF4, 0xB2, 0x6B, 0xA9, 0x2C, 0xC5, - 0x2B, 0x25, 0x27, 0xC6, 0x58, 0xD2, 0x7B, 0x91, 0x99, 0xE0, 0xCB, 0x3D, - 0xB7, 0x24, 0x53, 0x9F, 0xD0, 0x37, 0xF0, 0x69, 0x71, 0xE8, 0xAD, 0x93, - 0x33, 0x8D, 0x63, 0x1C, 0x55, 0x08, 0x3A, 0x91, 0x9B, 0x8D, 0x8D, 0x7D, - 0xDB, 0x87, 0x9D, 0xA5, 0x6C, 0xE8, 0xB9, 0x64, 0xA5, 0xEA, 0xB6, 0x72, - 0xC0, 0xE6, 0x3A, 0xA1, 0x9D, 0x02, 0xB9, 0x1F, 0xA7, 0xA6, 0x6C, 0xA2, - 0xF5, 0x9B, 0x4C, 0x06, 0x69, 0x53, 0x2F, 0x97, 0x8C, 0xE0, 0x37, 0x57, - 0x19, 0xDD, 0x1C, 0x03, 0xC5, 0xB8, 0x74, 0x9A, 0x0F, 0xC8, 0x44, 0xAB, - 0x91, 0x22, 0xFE, 0x07, 0x07, 0x6B, 0x1E, 0x2E, 0x48, 0x08, 0x9F, 0x66, - 0xE8, 0x09, 0x1A, 0xC1, 0x78, 0xF9, 0x0D, 0x34, 0x9E, 0x97, 0x59, 0xEF, - 0x5A, 0xA2, 0xF5, 0x65, 0x47, 0x5D, 0x9D, 0x3B, 0x94, 0x94, 0xDD, 0x9A, - 0x77, 0x0C, 0xE8, 0x69, 0x85, 0xAF, 0x6F, 0x2C, 0x7A, 0x10, 0x15, 0xC7, - 0x83, 0x08, 0xCB, 0x8D, 0x18, 0x2F, 0xAF, 0x57, 0x8A, 0x45, 0x8B, 0x83, - 0x73, 0xFA, 0x77, 0x28, 0x68, 0xDB, 0xBE, 0x6A, 0x62, 0x5A, 0xF2, 0xBD, - 0x4E, 0x5F, 0xB6, 0xBE, 0x32, 0x00, 0xDB, 0x79, 0xBE, 0x73, 0xA7, 0x2A, - 0xF9, 0x43, 0xBC, 0x06, 0x99, 0x09, 0xED, 0xBF, 0xE9, 0x51, 0xF0, 0x83, - 0xA1, 0xCC, 0x7E, 0x6B, 0xA2, 0x62, 0x9C, 0xA3, 0x38, 0xBD, 0xD9, 0x61, - 0x3A, 0x72, 0x73, 0xAD, 0x48, 0x98, 0xE8, 0xEA, 0x3D, 0xE5, 0xEA, 0xC4, - 0x2B, 0x77, 0x64, 0xDC, 0x8B, 0x57, 0x27, 0xDC, 0x52, 0xDE, 0xCD, 0x50, - 0x07, 0xD4, 0x28, 0x57, 0xD8, 0xE6, 0xA2, 0x14, 0x3A, 0xD4, 0xFF, 0xC0, - 0x34, 0x3E, 0x43, 0x05, 0xB3, 0xA8, 0x0E, 0x30, 0x0F, 0x50, 0x60, 0x18, - 0x27, 0x70, 0x8A, 0x12, 0x93, 0x70, 0x16, 0xA5, 0x54, 0x52, 0x3C, 0xEF, - 0xBE, 0x12, 0xA0, 0xDE, 0xB4, 0xB9, 0xF5, 0xED, 0xD7, 0x37, 0x01, 0xC0, - 0x46, 0x9F, 0x5B, 0x94, 0x18, 0x09, 0x66, 0x61, 0xFC, 0x75, 0x83, 0x8B, - 0x79, 0x21, 0xA3, 0xDC, 0x1A, 0xF4, 0x91, 0x8D, 0x5C, 0x29, 0x74, 0x0D, - 0x8F, 0x3C, 0xBB, 0x4F, 0x6E, 0x62, 0x64, 0x3F, 0x9C, 0xB4, 0xF6, 0x1E, - 0x5B, 0x0F, 0xED, 0x1F, 0xCB, 0xA5, 0xF7, 0xCB, 0x65, 0x3E, 0x71, 0x5E, - 0x1C, 0x27, 0xF6, 0x51, 0x40, 0x3B, 0x98, 0x6A, 0x9B, 0x04, 0x61, 0x40, - 0x4D, 0xD3, 0xC2, 0xCA, 0x09, 0x49, 0xAE, 0xBF, 0x89, 0x74, 0x94, 0xD6, - 0x6F, 0xFA, 0xA5, 0x89, 0x39, 0x00, 0xF9, 0xDE, 0xB5, 0xD5, 0xC6, 0x90, - 0x6F, 0x26, 0x3F, 0x54, 0xA9, 0x57, 0x1D, 0x02, 0xD2, 0x8A, 0x5D, 0x5B, - 0xD2, 0xFB, 0xA2, 0x2B, 0x80, 0xF3, 0xE9, 0xD7, 0x06, 0x36, 0xD3, 0xF4, - 0x1C, 0x49, 0x67, 0x09, 0x34, 0x62, 0xB1, 0xAB, 0xF5, 0x97, 0x6C, 0x50, - 0x35, 0x42, 0x14, 0x0F, 0x34, 0xA9, 0x71, 0x1E, 0xE6, 0xF5, 0x28, 0xA6, - 0xB9, 0x70, 0x35, 0x0F, 0x28, 0xAD, 0x54, 0x78, 0xC7, 0xAE, 0x2E, 0x8E, - 0x50, 0x9E, 0x43, 0x29, 0x9F, 0xF2, 0xC3, 0x10, 0x7B, 0x17, 0x02, 0x33, - 0x57, 0x44, 0x97, 0x36, 0xF8, 0xA7, 0xA0, 0xD2, 0x3D, 0x99, 0xAE, 0x22, - 0xE6, 0x83, 0xDD, 0x43, 0x5B, 0x06, 0xDE, 0x8C, 0xDB, 0x2E, 0x8B, 0xC5, - 0x8D, 0x36, 0xB2, 0x48, 0x84, 0x5D, 0x2E, 0xF6, 0xE2, 0x92, 0x47, 0xA9, - 0x82, 0x40, 0x79, 0x77, 0xAD, 0x0C, 0x59, 0x7F, 0xD8, 0xF4, 0x0D, 0x5F, - 0x83, 0x75, 0x99, 0xC3, 0x84, 0xA5, 0xCC, 0x0A, 0xA0, 0xD2, 0x76, 0x6E, - 0x5D, 0x73, 0xD5, 0xFF, 0x91, 0xF4, 0xC6, 0xF1, 0x77, 0x00, 0xA7, 0xF9, - 0xA5, 0x06, 0x35, 0xD5, 0x7C, 0x67, 0x9C, 0x50, 0x5F, 0x60, 0xFE, 0x39, - 0xE6, 0x80, 0xE6, 0x91, 0xD2, 0xDB, 0x0D, 0x9C, 0x49, 0x51, 0x5F, 0xAC, - 0xFF, 0x94, 0x6D, 0x51, 0x83, 0xFD, 0xBD, 0x50, 0x92, 0xE4, 0x76, 0x5A, - 0x84, 0xBF, 0x75, 0xE3, 0x5F, 0x31, 0xAF, 0x3A, 0xEA, 0x54, 0xEC, 0xF3, - 0x18, 0xDA, 0x31, 0x00, 0x1A, 0x74, 0x13, 0x96, 0x5F, 0x38, 0x0B, 0xA4, - 0x14, 0xA2, 0xB2, 0x36, 0xFE, 0x3E, 0x2B, 0xF0, 0x30, 0xCE, 0x44, 0x3B, - 0x20, 0xEA, 0x82, 0x4C, 0x42, 0x43, 0x62, 0x65, 0x4E, 0x2E, 0xD3, 0x4A, - 0xC6, 0xA6, 0xF7, 0x04, 0x50, 0xE0, 0x73, 0x2B, 0xA3, 0x5E, 0xD2, 0x4B, - 0xBD, 0x6C, 0x91, 0xE9, 0xAE, 0x9C, 0x91, 0xE3, 0x03, 0x93, 0x3F, 0x4C, - 0xE2, 0xBF, 0x67, 0x71, 0xD0, 0x42, 0x9E, 0xEB, 0x33, 0xEE, 0x68, 0x94, - 0xD6, 0xAE, 0xAB, 0x8F, 0xA6, 0x78, 0x4F, 0x35, 0x71, 0x88, 0x96, 0xD5, - 0x82, 0xD6, 0xA4, 0xFB, 0xF3, 0xF0, 0x33, 0x16, 0x10, 0x8A, 0x24, 0x79, - 0xC1, 0x8F, 0xD6, 0xBC, 0x2E, 0x10, 0x68, 0xE2, 0x34, 0x1D, 0xF6, 0x26, - 0x3E, 0x9C, 0xF9, 0xCF, 0xD9, 0x6E, 0x10, 0xD6, 0x7F, 0x7C, 0x40, 0x96, - 0x14, 0xAC, 0xD9, 0x9E, 0x3C, 0x75, 0xA0, 0x87, 0x66, 0xF9, 0xFD, 0x1F, - 0xD9, 0x11, 0x54, 0x79, 0x78, 0x2E, 0x0E, 0x7B, 0x9C, 0xD6, 0xD2, 0x89, - }; - int8_t scale_codes[] = { - 123, 18, 34, 97, 81, 116, 107, 82, 89, 127, 17, 56, - 121, 27, 110, 123, 92, 33, 39, 86, 45, 96, 89, 83, - 127, 14, 14, 112, 52, 15, 52, 118, 73, 116, 123, 61, - 118, 103, 39, 69, 124, 110, 92, 74, 34, 58, 127, 79, - 37, 81, 23, 66, 71, 16, 32, 32, 31, 123, 98, 30, - 47, 119, 127, 80, 78, 28, 116, 59, 84, 64, 99, 92, - 63, 67, 77, 26, 74, 29, 127, 110, 64, 110, 75, 122, - 81, 97, 79, 37, 15, 46, 115, 99, 51, 34, 127, 76, - 20, 49, 74, 64, 47, 87, 30, 94, 30, 116, 116, 127, - 13, 48, 58, 16, 98, 47, 70, 99, 25, 79, 45, 113, - 122, 57, 33, 106, 29, 127, 118, 96, 44, 118, 118, 61, - 37, 74, 120, 125, 17, 127, 34, 84, 46, 38, 63, 75, - 23, 51, 115, 69, 41, 123, 16, 56, 121, 77, 50, 39, - 57, 38, 68, 127, 109, 55, 114, 17, 75, 127, 89, 24, - 86, 113, 70, 111, 22, 117, 14, 102, 24, 62, 76, 73, - 75, 30, 33, 127, 99, 91, 71, 101, 98, 68, 88, 88, - }; - uint16_t scale_step[] = { - 0x123C, 0x12CE, 0x11F4, 0x12AD, 0x125D, 0x130E, 0x130E, 0x12EE, - 0x12C6, 0x12DE, 0x1295, 0x11EC, - }; - uint16_t A_host[] = { - 0xC013, 0x3E95, 0xBF05, 0x3F14, 0x3F1C, 0x3EE9, 0xBEE7, 0x3D6D, - 0x3F22, 0x3D77, 0xBEA8, 0xBF0F, 0xBFBC, 0x3EC0, 0x3F21, 0xBF4F, - 0xBE9D, 0x3F6B, 0x3F9F, 0xBF22, 0xBE0A, 0xC026, 0x3F3F, 0xBF5E, - 0xBFBD, 0xBFA8, 0x3F76, 0xBF3F, 0x3FF1, 0xBED6, 0x3FA9, 0x3F12, - 0xBF96, 0x3E29, 0xBF0A, 0xBF45, 0xBF26, 0xBEF6, 0x3F42, 0xBFFE, - 0xBDFB, 0x3F21, 0x3D9E, 0x3FB7, 0xBED1, 0xBFF0, 0x3E53, 0xBF84, - 0xBF6E, 0xBF32, 0x3F98, 0x3F6C, 0xBEF6, 0xBF4A, 0xBEB3, 0x3EEE, - 0xBFB7, 0x3ECC, 0x3DD2, 0x3FB8, 0x3DA1, 0x3FEF, 0x3E5F, 0xBF4D, - 0x3F1B, 0x3C0C, 0xBD8C, 0xBF51, 0x3F29, 0xBF26, 0xBF94, 0xBF3F, - 0x3F0A, 0xBFA1, 0x3EBE, 0x3DB2, 0x4004, 0x3FB7, 0xBE16, 0x3EC5, - 0x3F6F, 0xBEB5, 0xBFC4, 0xBFAA, 0x3FB5, 0x3F9D, 0xBE6F, 0x3F49, - 0xBF1D, 0x3F74, 0xBE17, 0xBE68, 0x3FFF, 0xBE3E, 0x3E62, 0x3F52, - 0xBEBF, 0x3F2A, 0x3BC2, 0xBF35, 0x4003, 0xBE76, 0x3EBF, 0xBFB9, - 0x3ED5, 0xBD60, 0x3E42, 0x3FFA, 0x3FA5, 0x3EB6, 0xBF18, 0x3F4B, - 0x3E53, 0x3FDD, 0x3EB7, 0xBF9C, 0xBF88, 0x3F37, 0xC00B, 0x3F01, - 0x3E50, 0x3E0F, 0xBF1C, 0x3F48, 0xC032, 0x3FBA, 0xBF04, 0xBDFA, - 0x3F9A, 0x3FB2, 0x3DF1, 0xBF9F, 0xBF7E, 0x3E99, 0xBDE6, 0xBF75, - 0x3E20, 0xBFE4, 0x3F14, 0x3E57, 0x3D09, 0x3D14, 0x3F3E, 0xBF08, - 0xBF19, 0xBFB6, 0xBF34, 0xBF74, 0x3F91, 0x3FA5, 0x3E73, 0xBE2F, - 0x3F39, 0xBF89, 0xBFA0, 0xBF99, 0xBFAA, 0xBE9A, 0x3E86, 0x3FCA, - 0xBF98, 0xBF89, 0xBE92, 0xBD90, 0x3F48, 0x3F80, 0xBFD3, 0xBF07, - 0xBF85, 0xBEE1, 0x3FAD, 0xBEA7, 0xBE9C, 0xBF18, 0x3FB7, 0x3D9C, - 0x3F82, 0xBF81, 0x3F18, 0xBE1F, 0xBFE0, 0x3FDD, 0xBF48, 0xBFE1, - 0x3F05, 0x3FBB, 0xBED4, 0x3FB2, 0xBE52, 0xBE9A, 0xBEA1, 0x3D0A, - 0x3F58, 0x3CB4, 0x3DE4, 0x3FB2, 0x3F53, 0xBFA1, 0x3C42, 0xBF8A, - 0x3F8E, 0x3F0F, 0x3D2C, 0x3F5A, 0xBF36, 0x3F80, 0x3F91, 0x3FD6, - 0xBF3B, 0x3F3D, 0x3E64, 0xBF91, 0x3F88, 0x3E7D, 0x3FB0, 0xBF58, - 0x3F55, 0x3F0E, 0x3C81, 0xBFC6, 0x3EC0, 0xBEFB, 0x3ED5, 0x3EB8, - 0xBE3A, 0xBF0D, 0x3D07, 0x3F2D, 0x3F8C, 0x3FA1, 0x3EF0, 0xBFC5, - 0x3EE8, 0x3E23, 0xBED6, 0xBE35, 0xBF10, 0x3E75, 0x3F24, 0xBF6A, - 0x3F46, 0x3F8E, 0xBEA2, 0x3F0A, 0xBFC8, 0xBF2A, 0xBFD6, 0x3F2B, - 0xBF60, 0x3FDF, 0xBFD9, 0xBC7F, 0xC002, 0x3FA5, 0xBE32, 0xBF47, - 0x3F1E, 0x3F15, 0xBF99, 0x3F42, 0xBD86, 0x3EBD, 0xBFD9, 0x3E01, - 0x3E98, 0xBEF8, 0x3F0F, 0x3E9C, 0xBF73, 0x3F4C, 0x3E9A, 0xBEB8, - 0x3F2B, 0x3E38, 0x3F3D, 0x3F44, 0xBEBF, 0x3FFE, 0x3FAE, 0x3EFD, - 0x3E58, 0x3EEC, 0x3F89, 0xBD75, 0x3F69, 0xBE6B, 0xBD85, 0x3FAD, - 0xBFB8, 0xBDDE, 0x401E, 0xBFE9, 0xBF55, 0x3FB6, 0x3F83, 0x3F81, - 0x3F26, 0x3F57, 0x3FAD, 0x3F4D, 0x3F95, 0xBF6B, 0xBE92, 0x3EB7, - 0x3EEE, 0x3E20, 0x3E1C, 0x3E73, 0x3EBA, 0x3F39, 0xBF25, 0xBF55, - 0x3D40, 0x3F40, 0x3C70, 0x3F8C, 0x3FAF, 0xBF51, 0x3F9E, 0x3F1E, - 0x3DD0, 0xBF5D, 0xBF86, 0xBF5D, 0xBF9E, 0xBFBC, 0xBEA6, 0xBFFE, - 0x3EF5, 0x3E8A, 0xBF44, 0xBF89, 0x4009, 0xBF0B, 0xBF87, 0x3EC2, - 0x3F64, 0x3E1B, 0xBF29, 0x3BF5, 0x3F7F, 0x3ECE, 0x3E1C, 0xC006, - 0x3E35, 0xBE87, 0x3D5D, 0x3D76, 0xBE57, 0xBE6E, 0x3F4A, 0xBE87, - 0x3E1D, 0x3E7A, 0xBFD8, 0xBFBB, 0xBF76, 0x3E50, 0xBF28, 0x3EEB, - 0x3F52, 0xBDD3, 0xBDDF, 0xBF67, 0xC005, 0xBF22, 0xBF31, 0x3FD1, - 0xBF5E, 0x3F2C, 0x3F68, 0x3D20, 0x3F08, 0x3EEB, 0x3F8A, 0xBFA0, - 0x4012, 0x3EE6, 0xBF69, 0x3D77, 0xBF20, 0x3FF9, 0xBF30, 0xBE16, - 0x3E5D, 0xBE19, 0x4002, 0xBEAE, 0xBE85, 0x3F6E, 0xBEE8, 0xBFD7, - 0x3F18, 0xC046, 0xBF3E, 0xBF25, 0x3F22, 0xBE26, 0x3F89, 0xBF73, - 0xBD86, 0x3FAA, 0x3C3A, 0xBF87, 0xBD3D, 0x3E6A, 0xBC16, 0x3FFD, - 0xBF13, 0x3F54, 0xBE0F, 0x3F8B, 0x3F0B, 0x3D15, 0x3F5C, 0xBEC2, - 0xBED7, 0xBF74, 0xBF62, 0xBE42, 0xBF1E, 0x3F98, 0x3FBC, 0x3FA4, - 0xBEFF, 0x3FD9, 0x3DCA, 0xBE50, 0x3E89, 0xBF69, 0xBF22, 0xBEEE, - 0xC00B, 0xBDA4, 0xBF60, 0x3D5F, 0x3FEC, 0xBF70, 0x3F37, 0x3FA5, - 0xBF91, 0x3E1D, 0x3E8A, 0xBFF1, 0xBEA9, 0x3F57, 0xBFA3, 0xBE11, - 0xBF63, 0x3FED, 0xBFD6, 0xBE9D, 0xBF13, 0x3F30, 0x3EF2, 0x3F4D, - 0x3F8A, 0x3F98, 0x3E8C, 0xBE45, 0x3EA5, 0x3FB1, 0x3F35, 0xBEEA, - 0xBFD1, 0x3F10, 0xBE40, 0xBF7A, 0xBE90, 0x3DAE, 0xBEC3, 0x3F90, - 0x3FD0, 0xBF08, 0xBF0E, 0x3E66, 0x3F21, 0xBE04, 0x3EC1, 0x4000, - 0xBFA4, 0xBF48, 0xBF9E, 0xBFA1, 0xBCAF, 0xBE77, 0xBF29, 0x3F94, - 0xBF5A, 0x3E11, 0x3F89, 0x3F26, 0xBEBB, 0x3F3D, 0x3F92, 0xBEAB, - 0x4022, 0x3FFD, 0x3FA7, 0xBE44, 0x3F45, 0xBED7, 0x3EAB, 0x3F4F, - 0xBF09, 0x3E9D, 0xBFF4, 0x3F15, 0xBF10, 0xBEAA, 0x3F3F, 0xBFBE, - }; - uint16_t expected[] = { - 0xC1D2, 0xC21D, 0xC197, 0xC1DC, 0xC1B4, 0x4115, - }; - // clang-format on - - Tensor* output = setup_and_run( - M, N, K, gs, ql_host, qh_host, scale_codes, scale_step, A_host); - ASSERT_NE(output, nullptr); - EXPECT_EQ(output->size(0), M); - EXPECT_EQ(output->size(1), N); - check_bf16_output(output, expected, M * N, 0.5f); -} - -// Q6KWideN: M=1, N=16, K=256, gs=16, ng=16, n_super=1 -TEST_F(AOTITorchInt6PlainMMTest, Q6KWideN) { - int64_t M = 1, N = 16, K = 256, gs = 16; - int64_t ng = K / gs; // 16 - int64_t n_super = K / 256; // 1 - // clang-format off - uint8_t ql_host[] = { - 0xF8, 0x8D, 0xB6, 0xB2, 0x78, 0x12, 0xBF, 0xF5, 0xFF, 0x4A, 0x4C, 0x75, - 0x63, 0xA4, 0x3B, 0x67, 0x1A, 0xCA, 0x3F, 0x85, 0xE4, 0xF6, 0x93, 0x2F, - 0x40, 0xCE, 0x42, 0xFA, 0x1F, 0xF7, 0x8E, 0xD2, 0xFA, 0xD9, 0x78, 0x61, - 0xA8, 0x95, 0x99, 0xCE, 0xEA, 0x03, 0x0B, 0x2A, 0xC8, 0xAE, 0x28, 0x9C, - 0x56, 0x66, 0x36, 0x28, 0x41, 0x18, 0x6D, 0xD9, 0x15, 0x42, 0x7A, 0xB6, - 0x54, 0x8A, 0x3A, 0x0B, 0x50, 0xB7, 0xB5, 0x0D, 0x68, 0x5A, 0x1B, 0xEF, - 0xAD, 0xB7, 0xFF, 0x34, 0xB6, 0x41, 0xFA, 0x0B, 0x8E, 0x5A, 0x4C, 0xBF, - 0x2E, 0x79, 0x91, 0x12, 0xBF, 0x0E, 0x17, 0xC8, 0xBE, 0xB9, 0xEE, 0xFE, - 0x70, 0x50, 0xAB, 0x52, 0x1A, 0x3F, 0xCD, 0x13, 0xB8, 0x86, 0xBF, 0xB1, - 0xB5, 0x7F, 0x90, 0xFC, 0x1A, 0x95, 0x2F, 0xA0, 0x40, 0x36, 0x1F, 0x8D, - 0x95, 0x45, 0x72, 0x78, 0xB3, 0x4A, 0x68, 0x83, 0xC8, 0x5E, 0x31, 0xB3, - 0xF3, 0x62, 0x55, 0x52, 0xCD, 0x6C, 0xC1, 0x05, 0x45, 0xE8, 0x3A, 0x57, - 0xD9, 0xF4, 0x25, 0xD6, 0x0E, 0x75, 0x61, 0x07, 0x21, 0xFD, 0xCD, 0xA4, - 0x60, 0x86, 0x6D, 0x9B, 0x15, 0x95, 0x94, 0x8F, 0x43, 0xC6, 0xA2, 0xBC, - 0xFF, 0x4F, 0xE7, 0x15, 0xD3, 0x7C, 0x54, 0x21, 0xBC, 0xA3, 0xA6, 0xEA, - 0x1D, 0xB2, 0x69, 0xF0, 0xC5, 0x3E, 0x25, 0x8A, 0x9B, 0x7A, 0xAB, 0xC3, - 0x40, 0xCF, 0xD6, 0xE2, 0x97, 0x13, 0x1F, 0x99, 0x0B, 0x08, 0x8A, 0x35, - 0x85, 0x3E, 0xFE, 0x16, 0x6C, 0x21, 0x1F, 0x2D, 0x66, 0x1B, 0xD3, 0xCA, - 0x3F, 0xAD, 0xCC, 0xF2, 0x23, 0xAC, 0xBA, 0xAE, 0x3B, 0x3A, 0x0D, 0x5C, - 0xBA, 0xAE, 0x21, 0xF1, 0x26, 0xEE, 0xC1, 0xE3, 0x6F, 0x3D, 0x09, 0x4E, - 0xB2, 0x57, 0xB5, 0x3C, 0x19, 0xF8, 0xAF, 0x41, 0xD6, 0x2E, 0x76, 0xA5, - 0x10, 0x5F, 0x96, 0xDE, 0x8F, 0x28, 0xB6, 0x60, 0x0A, 0x34, 0x73, 0xFF, - 0x0C, 0x9F, 0xE2, 0x2E, 0x3A, 0x65, 0x60, 0x30, 0x37, 0xC8, 0x0F, 0x83, - 0x5C, 0x57, 0x47, 0xB1, 0x0F, 0xE6, 0xB7, 0x87, 0x94, 0xE2, 0x57, 0xA9, - 0x63, 0xFC, 0x04, 0x6F, 0xAA, 0x1E, 0x8F, 0x5E, 0x01, 0xBE, 0x05, 0x89, - 0xE2, 0xC6, 0x9B, 0xDF, 0x9F, 0x28, 0xC4, 0x11, 0xE2, 0x2E, 0x6C, 0x7E, - 0xB9, 0xF8, 0xB3, 0x49, 0x8D, 0x04, 0xB0, 0xC4, 0xD7, 0x95, 0xE2, 0x9A, - 0x06, 0xB3, 0x6B, 0x28, 0xAC, 0xFC, 0x01, 0xB7, 0xDD, 0x75, 0x6F, 0x39, - 0x40, 0x09, 0xB2, 0xB0, 0xFA, 0xEC, 0x49, 0x47, 0x56, 0x59, 0x5C, 0x29, - 0xF6, 0xAE, 0x92, 0x04, 0x89, 0x12, 0x6A, 0x91, 0x93, 0x22, 0x8A, 0x52, - 0xA2, 0x69, 0x56, 0x50, 0xAF, 0x3A, 0x9A, 0x70, 0xC6, 0xAF, 0x75, 0xC0, - 0x1F, 0xB9, 0xB1, 0x06, 0x12, 0x6F, 0xC5, 0xE5, 0x85, 0xA9, 0x1D, 0x14, - 0x5B, 0xAD, 0x9E, 0x59, 0xB3, 0x4C, 0x0D, 0x53, 0x79, 0x68, 0x21, 0xD0, - 0xBD, 0xE9, 0xCD, 0x01, 0x51, 0x24, 0xE8, 0x91, 0x68, 0x6D, 0xD4, 0x2D, - 0x2E, 0xE3, 0xFA, 0x47, 0x03, 0xEE, 0x2E, 0xCF, 0x1D, 0x25, 0x02, 0xC3, - 0xB9, 0x15, 0x28, 0x18, 0x90, 0x2D, 0x5C, 0x42, 0x56, 0x63, 0xDE, 0x39, - 0xBF, 0xC0, 0x8C, 0xBA, 0x04, 0xC2, 0x1D, 0x55, 0xCA, 0x84, 0x8F, 0x68, - 0x85, 0x7F, 0x14, 0x1B, 0x2C, 0x7E, 0x69, 0xA8, 0x69, 0x7F, 0x77, 0xB3, - 0xBD, 0xBD, 0xAE, 0xA8, 0x4B, 0x90, 0xA4, 0x4B, 0x02, 0x47, 0xA2, 0xD5, - 0xA5, 0xC1, 0xF0, 0x80, 0x11, 0xCF, 0xAA, 0x03, 0x82, 0x09, 0x25, 0x79, - 0x28, 0xC3, 0x22, 0x89, 0x3F, 0xE0, 0xAF, 0xE1, 0x7E, 0x58, 0x0E, 0xA7, - 0xB4, 0x2B, 0x38, 0xCD, 0x32, 0x71, 0xF3, 0x5A, 0xC7, 0xEA, 0xCF, 0x97, - 0x64, 0x43, 0x8B, 0x93, 0xF6, 0xF5, 0xD1, 0x8D, 0xD2, 0xA8, 0x6C, 0x6B, - 0xC6, 0xEA, 0xCC, 0x02, 0x59, 0x10, 0xB5, 0xCC, 0xFB, 0x1C, 0x23, 0x65, - 0x2F, 0xB1, 0x0C, 0xAA, 0x63, 0xA1, 0x76, 0xD4, 0x83, 0xE2, 0x3A, 0xC8, - 0xC0, 0xBC, 0xBE, 0x83, 0x8E, 0xF2, 0xDD, 0xED, 0xAC, 0x8C, 0xD7, 0x9A, - 0x9F, 0x5D, 0xC5, 0xCA, 0xD0, 0x6D, 0xA3, 0xA5, 0x39, 0xD7, 0x47, 0xC0, - 0x08, 0xDA, 0x88, 0x70, 0xD0, 0x88, 0x61, 0x75, 0x53, 0xB2, 0x39, 0xC1, - 0xE3, 0xC1, 0x3E, 0x3D, 0xDC, 0xDB, 0x45, 0x2B, 0xC0, 0xAE, 0x2A, 0xC9, - 0x5E, 0xAE, 0x95, 0x48, 0xB2, 0x0C, 0xB0, 0xFB, 0xA1, 0xB4, 0x68, 0x68, - 0x34, 0xC2, 0x47, 0x3A, 0x09, 0x32, 0x0F, 0x11, 0x69, 0x12, 0x35, 0xA4, - 0x7D, 0x1E, 0x1F, 0x91, 0x96, 0xB7, 0x2A, 0xD2, 0xB9, 0x72, 0xF1, 0xF9, - 0x88, 0xB2, 0x8E, 0x92, 0xDC, 0x92, 0x6E, 0xE9, 0xB6, 0xA7, 0x0B, 0x17, - 0x7B, 0x61, 0x9C, 0xC3, 0x1E, 0xF6, 0xE5, 0x22, 0xA5, 0x61, 0x8B, 0x0A, - 0x7A, 0x51, 0xF0, 0xCD, 0x4E, 0x0A, 0xED, 0x78, 0x09, 0xAB, 0xB2, 0xCB, - 0xF7, 0xB6, 0xD3, 0x91, 0xDB, 0x08, 0x37, 0x64, 0x0F, 0xBC, 0xF2, 0xF5, - 0x9C, 0x0A, 0x1D, 0xF6, 0x67, 0xEF, 0xD7, 0xA4, 0xEA, 0x49, 0xC2, 0x98, - 0xB7, 0x5B, 0x6B, 0xB0, 0x52, 0x37, 0x5D, 0x56, 0x8D, 0x1D, 0x52, 0x71, - 0x1E, 0x84, 0xDD, 0x30, 0x5A, 0x05, 0x03, 0xEA, 0xCA, 0x2C, 0x22, 0xE0, - 0xF5, 0x91, 0x7C, 0x8A, 0x16, 0xEC, 0x7E, 0x49, 0x20, 0xAD, 0x68, 0xAC, - 0x55, 0xF5, 0xDF, 0x55, 0x0C, 0xCA, 0x30, 0x32, 0xC9, 0xF3, 0x75, 0x9D, - 0x9A, 0xCD, 0xB0, 0x69, 0xAF, 0xB3, 0xE6, 0xE7, 0xF0, 0xBB, 0x04, 0x98, - 0x57, 0x0F, 0x16, 0xB7, 0xD1, 0x48, 0x7F, 0xF1, 0x47, 0xB0, 0x2B, 0x1E, - 0xBA, 0x34, 0x09, 0x84, 0xE8, 0x0D, 0xC1, 0x02, 0xD0, 0x17, 0xBF, 0x9E, - 0xC9, 0xAF, 0xC5, 0xE4, 0xBA, 0x1F, 0xBC, 0xCC, 0x71, 0x76, 0x55, 0x18, - 0x4D, 0x49, 0x7C, 0xA7, 0x59, 0xDC, 0x55, 0xB7, 0xAC, 0xD2, 0x72, 0x66, - 0x67, 0xC5, 0x5B, 0x04, 0x7B, 0x74, 0xB8, 0x9A, 0x78, 0x2E, 0x33, 0x92, - 0x38, 0x79, 0x3E, 0x8C, 0xF7, 0xEB, 0x19, 0x0D, 0xE5, 0x9F, 0x54, 0xA2, - 0x36, 0x84, 0x63, 0xFE, 0x3C, 0x10, 0xDA, 0x40, 0xB9, 0xB7, 0xB2, 0x9D, - 0x09, 0x18, 0x49, 0xFF, 0x13, 0x69, 0x9A, 0x70, 0x96, 0x9D, 0x41, 0xF9, - 0xB9, 0x4C, 0x7B, 0x79, 0xAB, 0xB9, 0x1F, 0xEC, 0xFF, 0xE0, 0x18, 0xD2, - 0x57, 0xC5, 0x40, 0xD6, 0x25, 0x1D, 0xAC, 0x33, 0xB9, 0x94, 0xAB, 0xF0, - 0x5F, 0x01, 0x69, 0x85, 0xB7, 0xD1, 0x11, 0xE7, 0x5B, 0xBA, 0x1E, 0x54, - 0x8E, 0xE8, 0x34, 0x0A, 0x4E, 0x92, 0xB0, 0x15, 0xBF, 0xCB, 0xDB, 0xD4, - 0x4C, 0x94, 0xD1, 0xA6, 0x6F, 0x2F, 0x6F, 0x3E, 0xE6, 0x13, 0xA4, 0x55, - 0x46, 0xB3, 0xDB, 0x8D, 0x2E, 0xE2, 0x21, 0x43, 0x80, 0x49, 0x58, 0xA5, - 0x06, 0x51, 0x6C, 0x55, 0xA9, 0xE1, 0x7E, 0x2C, 0x50, 0x44, 0x3B, 0x9F, - 0x2D, 0xD5, 0xC0, 0x70, 0x0D, 0x10, 0xA7, 0x17, 0xB4, 0x87, 0x1C, 0xD0, - 0x3B, 0xDB, 0x83, 0x9D, 0x39, 0xEC, 0x94, 0x5F, 0x73, 0x34, 0xDF, 0xBC, - 0x51, 0x66, 0x67, 0xBB, 0xF8, 0xE0, 0x1E, 0x11, 0x9F, 0x30, 0x0B, 0x20, - 0x29, 0x04, 0x42, 0x79, 0x57, 0x9F, 0x0A, 0xE8, 0x0F, 0xDD, 0x85, 0x63, - 0x3B, 0x72, 0x5C, 0x64, 0x2D, 0x3F, 0x24, 0xD0, 0xE5, 0xC1, 0xF7, 0x29, - 0x6B, 0xA1, 0x1E, 0xFE, 0x3A, 0x94, 0x22, 0x98, 0x3E, 0x73, 0x6B, 0xC5, - 0x66, 0x5B, 0x13, 0xC6, 0x81, 0x6B, 0xC8, 0x7C, 0xA2, 0xF3, 0x2F, 0x36, - 0x54, 0x14, 0x2C, 0x9A, 0x0D, 0x9D, 0xE5, 0xEB, 0x16, 0x66, 0x93, 0x72, - 0xDD, 0xEC, 0x9C, 0xBE, 0x3D, 0x99, 0x93, 0xAF, 0x34, 0x7B, 0x77, 0x7F, - 0x11, 0x56, 0xB1, 0xDA, 0xA5, 0xA2, 0x44, 0xF2, 0xEE, 0x92, 0xD5, 0x26, - 0x2B, 0xB2, 0x3E, 0xEC, 0xCB, 0xE9, 0x62, 0x95, 0x1D, 0x6D, 0x40, 0x49, - 0xA5, 0x41, 0xF5, 0x38, 0xF8, 0x18, 0xA6, 0x48, 0xF6, 0x65, 0xD0, 0xD1, - 0x6C, 0x17, 0x67, 0x60, 0x7B, 0x80, 0xD9, 0x3B, 0xA5, 0x49, 0x20, 0xFE, - 0x91, 0x06, 0x1E, 0xCB, 0x47, 0x4E, 0xF2, 0x10, 0xEF, 0x28, 0xEE, 0xDE, - 0x36, 0x83, 0xBE, 0xD6, 0x37, 0x23, 0x92, 0xE2, 0xFA, 0xB4, 0x8E, 0x69, - 0x72, 0xCF, 0x47, 0xAE, 0x6F, 0xA0, 0x76, 0x0C, 0x37, 0x04, 0xB6, 0xB6, - 0xE0, 0xEB, 0xDD, 0xD8, 0x1A, 0x08, 0x3F, 0xEC, 0xAE, 0xCD, 0xC4, 0x42, - 0xAD, 0xEE, 0xBF, 0x8C, 0x5C, 0xA0, 0xD9, 0x1F, 0xED, 0x31, 0xE9, 0xA0, - 0xE3, 0xCC, 0x09, 0x9A, 0x72, 0x45, 0xF3, 0x02, 0x17, 0x00, 0xEB, 0x21, - 0xCF, 0xD2, 0xD5, 0xA9, 0x35, 0xFB, 0x70, 0xC8, 0x06, 0x5D, 0x27, 0x88, - 0x5F, 0xD6, 0x26, 0x97, 0xE6, 0x54, 0xBD, 0xAA, 0x45, 0x9A, 0x29, 0xC9, - 0xE2, 0x7A, 0x15, 0x54, 0x8D, 0xB9, 0x76, 0xF1, 0x8B, 0x41, 0x3C, 0xD6, - 0x87, 0x3A, 0x00, 0x05, 0x92, 0xD4, 0x67, 0x72, 0xCB, 0x25, 0x33, 0xC5, - 0x67, 0xFF, 0x7E, 0x8D, 0x71, 0x51, 0xF7, 0x63, 0x4A, 0xDA, 0x91, 0x15, - 0x52, 0xC2, 0x5C, 0xEB, 0xB0, 0x63, 0x25, 0x60, 0xB2, 0xCC, 0xA3, 0x1A, - 0x36, 0x3F, 0xAA, 0x65, 0x46, 0x2A, 0x33, 0x75, 0x45, 0x98, 0x8C, 0x84, - 0x27, 0x05, 0x2C, 0xAD, 0x3F, 0xD2, 0x6D, 0x52, 0x73, 0x54, 0x44, 0xEB, - 0x39, 0xFC, 0x50, 0xCF, 0x81, 0x56, 0x2A, 0x12, 0x82, 0x2E, 0xDA, 0x9B, - 0x50, 0x00, 0xD7, 0x28, 0x22, 0x4B, 0x29, 0x2F, 0xFA, 0x56, 0x19, 0x4B, - 0x8F, 0xCB, 0xFA, 0xE8, 0x47, 0xAE, 0xC5, 0x41, 0xBF, 0x3D, 0xB0, 0x4E, - 0x13, 0x09, 0x5F, 0x63, 0x99, 0x61, 0x81, 0xD6, 0x09, 0xB3, 0xF1, 0xCE, - 0xC2, 0xB8, 0xA3, 0x44, 0x30, 0xC9, 0x98, 0xE1, 0x49, 0x3A, 0x91, 0x9C, - 0x9A, 0x5F, 0x52, 0x03, 0x55, 0x3D, 0xE7, 0x82, 0x1E, 0x3E, 0xC7, 0xBE, - 0x5B, 0xB2, 0x52, 0x67, 0xBA, 0xF5, 0xE6, 0x98, 0x8E, 0xDA, 0x4A, 0x13, - 0x9A, 0xCF, 0xA9, 0xA1, 0xB7, 0x7C, 0x93, 0xEE, 0x97, 0xEF, 0x2D, 0xF8, - 0x1D, 0x86, 0x00, 0xB9, 0xC2, 0x59, 0x6C, 0x9C, 0xB1, 0x03, 0xE3, 0x2C, - 0x83, 0x3A, 0xC5, 0x50, 0x11, 0x52, 0xFC, 0xCD, 0xF5, 0x80, 0x59, 0x1F, - 0x89, 0xB1, 0xAD, 0xCC, 0xAD, 0xC6, 0x27, 0xC9, 0xD3, 0x08, 0x5C, 0x52, - 0x4A, 0x4D, 0xE0, 0x37, 0x6B, 0xB8, 0xB1, 0xC3, 0xF0, 0xCC, 0x2B, 0xEC, - 0x4B, 0xD8, 0x6A, 0x4B, 0x24, 0x07, 0x3A, 0x59, 0xF9, 0xDF, 0xFD, 0xA4, - 0xCB, 0x39, 0x17, 0x21, 0xC4, 0xF1, 0x96, 0x66, 0xB7, 0x0B, 0x5A, 0xFC, - 0x78, 0x5A, 0x5C, 0x28, 0x45, 0x52, 0x3D, 0xC5, 0x33, 0x63, 0xFF, 0xF3, - 0x26, 0xD0, 0x5B, 0x74, 0xEE, 0xFC, 0x2C, 0xCA, 0xF8, 0xD4, 0x16, 0x20, - 0x1D, 0x14, 0x26, 0x6F, 0x77, 0x38, 0x69, 0x8E, 0x29, 0x02, 0xAA, 0xBF, - 0xE4, 0xB6, 0x17, 0x3D, 0xD8, 0x66, 0xDF, 0x14, 0x24, 0xEE, 0xC9, 0x78, - 0xCD, 0x6F, 0x87, 0x22, 0x52, 0xCD, 0x65, 0xD6, 0x7F, 0x67, 0xA9, 0xF3, - 0x2F, 0x63, 0x83, 0x0D, 0x0D, 0x65, 0x27, 0x04, 0xE6, 0xC3, 0x51, 0x94, - 0x2E, 0xAA, 0xB8, 0xBD, 0xDE, 0x0D, 0x31, 0x88, 0x70, 0xAA, 0x5F, 0x0A, - 0x72, 0x70, 0xD0, 0xC5, 0xE1, 0xFD, 0xB2, 0x94, 0x25, 0x87, 0x7B, 0xBC, - 0x4E, 0xA7, 0xED, 0x4F, 0xDE, 0x0E, 0xB1, 0xDF, 0xD2, 0x79, 0x95, 0xF7, - 0xCA, 0xB6, 0xB5, 0x9A, 0x9D, 0x86, 0x45, 0xB3, 0x80, 0x43, 0x60, 0x22, - 0xB2, 0x34, 0x5E, 0x60, 0x6B, 0x50, 0x7B, 0x7E, 0x56, 0x49, 0x51, 0x49, - 0x88, 0x4F, 0xD9, 0x76, 0x75, 0xA0, 0xA7, 0xDC, 0x8B, 0x46, 0xEB, 0x99, - 0x54, 0xD5, 0x49, 0x87, 0x8F, 0x57, 0x43, 0x74, 0xBC, 0x01, 0xD2, 0x53, - 0xB0, 0x50, 0x7F, 0x62, 0xE0, 0xD3, 0x9E, 0xAB, 0x64, 0x5D, 0x7C, 0xE3, - 0x4B, 0xFA, 0x13, 0x11, 0xAE, 0xF1, 0xC0, 0x6B, 0xA7, 0x46, 0x7C, 0xAC, - 0xE3, 0x42, 0x9B, 0xBA, 0x8E, 0x9C, 0x15, 0x8A, 0x01, 0xEF, 0xC9, 0xAE, - 0x8C, 0xEE, 0x4A, 0xE8, 0xDC, 0x60, 0x4E, 0x06, 0xD8, 0xEB, 0xAB, 0x35, - 0xF8, 0xD9, 0x7A, 0xE7, 0xE8, 0xC9, 0x52, 0xED, 0xB4, 0xA9, 0x68, 0x59, - 0x4D, 0xA8, 0x15, 0x70, 0x40, 0x0C, 0x45, 0x17, 0x01, 0x3F, 0xDA, 0x1E, - 0x0D, 0x3D, 0x61, 0x1E, 0x9C, 0x55, 0x46, 0x9D, 0xC2, 0xD2, 0xFD, 0x18, - 0x1B, 0xDE, 0x52, 0xDA, 0x51, 0x37, 0x0B, 0x16, 0x16, 0x7D, 0x71, 0x5F, - 0xA6, 0x18, 0x4E, 0xDA, 0xA4, 0x72, 0x83, 0x0B, 0x78, 0x28, 0x9D, 0x9C, - 0xDF, 0x6F, 0xBB, 0x67, 0xD2, 0x4B, 0x37, 0xFC, 0xBA, 0xEE, 0xE7, 0x82, - 0x1C, 0x9B, 0xC0, 0xC4, 0x4B, 0x77, 0x3E, 0x69, 0xF5, 0x7D, 0xF1, 0xD1, - 0xEC, 0x70, 0xCD, 0xB2, 0xCB, 0x75, 0x03, 0x9B, 0x18, 0x42, 0x32, 0x81, - 0x47, 0x90, 0x37, 0x9D, 0x2D, 0xD6, 0x62, 0x41, 0xC3, 0xBE, 0xBA, 0x43, - 0x60, 0x98, 0x7F, 0x80, 0xEB, 0x68, 0xFB, 0x9E, 0x1A, 0x59, 0xF3, 0x88, - 0xA3, 0xAA, 0x7D, 0xC9, 0xB5, 0x55, 0x28, 0x71, 0x1D, 0xF0, 0xEA, 0x90, - 0xAB, 0x55, 0xA1, 0x46, 0x02, 0x6C, 0x1E, 0x2C, 0x34, 0x18, 0xF2, 0xA9, - 0x85, 0xA2, 0x2A, 0x1C, 0x9C, 0x0C, 0xA3, 0x3F, 0x68, 0x4B, 0x86, 0xEC, - 0x52, 0xC7, 0x6E, 0xCD, 0xA1, 0x72, 0x3B, 0xC8, 0x74, 0xD5, 0xD7, 0x6B, - 0x85, 0xC5, 0x32, 0xCF, 0xDC, 0xD5, 0xA6, 0x23, 0x71, 0x7D, 0xDC, 0xA3, - 0xEE, 0x15, 0x3E, 0x64, 0x3E, 0xE8, 0xD2, 0x12, 0xEC, 0xFD, 0xE0, 0x65, - 0x7F, 0xAC, 0x03, 0x02, 0xAD, 0xDD, 0xA3, 0xA0, 0x41, 0xDE, 0xD0, 0x84, - 0x39, 0x4D, 0xDA, 0x00, 0xD7, 0x2E, 0x97, 0x7D, 0xBF, 0xEF, 0x78, 0xF9, - 0x60, 0x5F, 0x6D, 0xAC, 0x78, 0x7A, 0xDE, 0x9E, 0xA3, 0x20, 0xA0, 0xAD, - 0x5D, 0x13, 0xD8, 0x02, 0x1C, 0x35, 0xEA, 0x32, 0x4C, 0x21, 0xD4, 0x16, - 0x62, 0xB4, 0xD4, 0x63, 0x24, 0xD4, 0xAD, 0xAE, 0xCD, 0x8A, 0x6A, 0x0F, - 0xF6, 0x3E, 0x48, 0x28, 0x19, 0x45, 0x2B, 0x79, 0x5F, 0x66, 0x41, 0x5C, - 0x51, 0xBD, 0xC5, 0xEE, 0x1B, 0xDF, 0xAA, 0x02, 0x47, 0x46, 0xF9, 0xE5, - 0xF3, 0x10, 0x08, 0xD1, 0xD1, 0x33, 0x24, 0xB7, 0x73, 0xF3, 0xEF, 0x7F, - 0xBE, 0xB9, 0x63, 0x74, 0x37, 0x94, 0x37, 0x21, 0xE1, 0xF7, 0xB6, 0x6C, - 0x9C, 0x16, 0xAC, 0xA1, 0xAA, 0x01, 0x10, 0x26, 0x7F, 0xF3, 0xCE, 0x5B, - 0x71, 0xCD, 0xD1, 0x64, 0x43, 0xA1, 0x62, 0x40, 0x28, 0xB1, 0x62, 0xB6, - 0x9E, 0x61, 0x62, 0x24, 0x38, 0x91, 0xFF, 0x83, 0x14, 0x63, 0xF0, 0xD3, - 0x40, 0xB9, 0x4C, 0xFA, 0x6A, 0x71, 0xA1, 0x4D, 0x4C, 0xDD, 0xDF, 0x40, - 0x7F, 0x84, 0x38, 0xD8, 0xDC, 0xD7, 0x85, 0xAB, 0x61, 0x86, 0xD8, 0x0D, - 0x8C, 0x63, 0x6B, 0xC7, 0xD9, 0x21, 0xB9, 0x58, 0xD1, 0x6D, 0x01, 0x9B, - 0x6C, 0xB2, 0xD3, 0x25, 0x42, 0x99, 0x8E, 0xB7, 0x7F, 0x55, 0x35, 0xEF, - 0x4F, 0xC0, 0xBA, 0x2E, 0x16, 0x75, 0xF2, 0x18, 0xD6, 0xDD, 0x8A, 0x85, - 0x83, 0x0F, 0x15, 0x88, 0xB3, 0xE7, 0x99, 0x7D, 0x08, 0xFD, 0x39, 0xB8, - 0xCC, 0x93, 0x0A, 0x04, 0xB3, 0x49, 0xC1, 0x96, 0x83, 0xF0, 0xF4, 0x6A, - 0xED, 0xA1, 0x23, 0xB0, 0xC9, 0x23, 0xAD, 0x1B, - }; - uint8_t qh_host[] = { - 0x5A, 0x5A, 0xB7, 0xA0, 0x20, 0x8C, 0xB5, 0xBA, 0x3E, 0x50, 0x79, 0xE8, - 0x4E, 0x36, 0xC3, 0x8F, 0x4B, 0xBF, 0x65, 0xEC, 0x78, 0xE3, 0xEB, 0x20, - 0xC3, 0xF8, 0x56, 0x53, 0x56, 0xC4, 0x96, 0xA2, 0x5D, 0x32, 0x78, 0x7D, - 0x6F, 0x58, 0xE9, 0x3B, 0xDF, 0x72, 0xA8, 0x82, 0x2E, 0x33, 0xB9, 0x5B, - 0xA6, 0xFE, 0x24, 0x5C, 0x77, 0x97, 0x86, 0x23, 0x5B, 0x97, 0xB6, 0x65, - 0xAD, 0x16, 0xA8, 0xEB, 0xF0, 0x40, 0x35, 0x9C, 0x7F, 0x84, 0x60, 0x9F, - 0x59, 0x20, 0x67, 0xB9, 0x53, 0xA4, 0x51, 0xC4, 0xE5, 0x2D, 0x6B, 0x67, - 0xAB, 0xFC, 0x3C, 0x76, 0x6D, 0xB2, 0xFD, 0xDC, 0xC9, 0xEE, 0x14, 0xF6, - 0xB9, 0x55, 0x83, 0xCB, 0xF0, 0x1A, 0x8C, 0xCA, 0x5A, 0xEF, 0x2A, 0xCE, - 0xB8, 0x0B, 0x7D, 0x97, 0x9A, 0x14, 0xB1, 0x66, 0x02, 0xD8, 0xAA, 0x85, - 0xDE, 0xE8, 0x46, 0xCD, 0x54, 0x01, 0x54, 0x79, 0x1E, 0xB1, 0xF9, 0xDE, - 0x6B, 0xDC, 0xE1, 0xFF, 0xF8, 0xFF, 0x7A, 0x03, 0x4F, 0x3F, 0x79, 0x2D, - 0x41, 0x31, 0x0F, 0xDF, 0x87, 0x2C, 0x25, 0x60, 0xBA, 0x7C, 0x56, 0x2A, - 0x4C, 0xDA, 0x90, 0x18, 0x71, 0x5D, 0xB7, 0x12, 0x97, 0xEE, 0xA5, 0x43, - 0x71, 0x80, 0x91, 0xDB, 0x9D, 0xD9, 0xF4, 0xBA, 0x70, 0x2A, 0xAD, 0xCC, - 0x62, 0xC1, 0x47, 0x60, 0xA5, 0x28, 0x07, 0xE3, 0xCA, 0xD1, 0x41, 0xA9, - 0xBB, 0x13, 0x2F, 0xBC, 0x1E, 0x0E, 0x48, 0xC3, 0x0F, 0xAE, 0x43, 0xFD, - 0x9F, 0x01, 0xF4, 0x61, 0xC3, 0x03, 0xD5, 0x89, 0xE5, 0xF8, 0xBF, 0xF1, - 0xC7, 0xFE, 0x5D, 0xD1, 0xD8, 0xF7, 0x18, 0xAE, 0xAC, 0x28, 0x65, 0xBE, - 0x51, 0x2E, 0x86, 0x56, 0xDF, 0x4B, 0x3F, 0xB4, 0xD2, 0xAA, 0x62, 0x98, - 0xEE, 0x22, 0xA7, 0x33, 0x74, 0x9E, 0x5C, 0x69, 0x1D, 0x74, 0x1B, 0x04, - 0x7D, 0x48, 0xC0, 0xBA, 0x1E, 0x83, 0xD5, 0x6E, 0x41, 0x75, 0xAE, 0x4C, - 0x3C, 0xDF, 0x26, 0x47, 0x69, 0xD5, 0xBC, 0x25, 0xD2, 0x31, 0x96, 0x93, - 0xA5, 0x31, 0x83, 0x0E, 0x0A, 0x7B, 0x0E, 0x43, 0xB1, 0x52, 0x8E, 0x6C, - 0xC9, 0x02, 0xCF, 0x46, 0x4D, 0x13, 0x2F, 0xF0, 0x7E, 0x39, 0xEF, 0x59, - 0x80, 0x04, 0x7D, 0x0D, 0xDE, 0xFB, 0xEE, 0x80, 0xF9, 0x6A, 0x07, 0xE3, - 0x5A, 0x85, 0xF9, 0xAF, 0xF7, 0xD7, 0x1F, 0xD2, 0xEA, 0x43, 0xD6, 0xD5, - 0x8F, 0x73, 0x79, 0xE9, 0x56, 0x31, 0x8E, 0xF6, 0xA8, 0xD7, 0x6B, 0x0B, - 0xB1, 0x5C, 0x31, 0x0E, 0x06, 0x85, 0x42, 0x0F, 0x39, 0x15, 0x07, 0xAC, - 0xB4, 0x7D, 0xAB, 0x62, 0x92, 0x79, 0x7D, 0xC0, 0x82, 0x29, 0x3B, 0x11, - 0xB1, 0xCA, 0xF4, 0xCA, 0x47, 0x95, 0x95, 0x06, 0x7F, 0x31, 0x66, 0x08, - 0xC1, 0x55, 0xEE, 0xCD, 0xD9, 0xB8, 0xC4, 0xC7, 0xF2, 0x13, 0xDD, 0x0A, - 0x94, 0x49, 0x94, 0xD5, 0x0F, 0x18, 0x57, 0x28, 0x6B, 0xE2, 0xAC, 0xE1, - 0x5B, 0xD8, 0x64, 0x07, 0xD3, 0xF4, 0x9C, 0x53, 0x72, 0x82, 0x8F, 0x09, - 0x42, 0xF5, 0x59, 0xF3, 0x26, 0x0A, 0xF6, 0xF2, 0xBB, 0x02, 0x35, 0x6F, - 0xF1, 0xCA, 0x41, 0xC1, 0x29, 0x25, 0x21, 0xEB, 0x74, 0xF6, 0x22, 0x4E, - 0x6A, 0x3B, 0xEA, 0x33, 0xEC, 0x6C, 0x3F, 0x81, 0x86, 0x8B, 0xA5, 0x8A, - 0xC5, 0x1B, 0xF9, 0x4C, 0x7A, 0x27, 0x0A, 0x78, 0x77, 0x05, 0x4D, 0x35, - 0xEE, 0x68, 0x5A, 0x04, 0xF0, 0x74, 0xE9, 0xB2, 0x5E, 0xCD, 0x99, 0xB6, - 0x90, 0x64, 0x5D, 0x64, 0x32, 0x22, 0xA7, 0xB3, 0x4D, 0xC2, 0x46, 0xC3, - 0xDA, 0x0A, 0xA1, 0xBB, 0xD6, 0x8D, 0x6B, 0x15, 0x7D, 0xDB, 0xF6, 0xC6, - 0x8B, 0xAA, 0x31, 0x23, 0xD5, 0x08, 0xE1, 0x0C, 0x2C, 0x46, 0x39, 0x88, - 0x3F, 0x96, 0x04, 0xB5, 0x86, 0x2C, 0x96, 0x6A, 0xCD, 0x70, 0xAA, 0xC7, - 0xF3, 0x29, 0xEF, 0xCA, 0xB7, 0x58, 0x04, 0xCC, 0x54, 0x75, 0x3A, 0x61, - 0x43, 0xB6, 0x83, 0xD3, 0x53, 0x63, 0x5B, 0x7B, 0xB8, 0x9F, 0x7E, 0xDE, - 0x0A, 0x8F, 0xD4, 0x14, 0x3D, 0x88, 0x23, 0xB3, 0xB5, 0x2A, 0x46, 0xB1, - 0x98, 0x05, 0xBB, 0x07, 0xF7, 0x5C, 0xD2, 0xF0, 0xCF, 0x73, 0xF1, 0xD1, - 0x53, 0x65, 0x8C, 0x82, 0xA9, 0x9A, 0xE2, 0x54, 0xFD, 0xB3, 0xB7, 0x1B, - 0x72, 0xD7, 0xF9, 0x90, 0x11, 0x9D, 0xE8, 0x08, 0xB5, 0x9D, 0x24, 0xFC, - 0xBD, 0x80, 0xC4, 0xFB, 0x22, 0xE8, 0xBE, 0xC4, 0xEA, 0xAA, 0x6F, 0xF4, - 0xDB, 0x17, 0x06, 0x73, 0x55, 0x86, 0xFD, 0xEC, 0xDD, 0xCC, 0x91, 0xBD, - 0xC8, 0xB9, 0xFA, 0x92, 0xAD, 0x77, 0x42, 0x86, 0xF1, 0x4F, 0x15, 0xCF, - 0x11, 0x96, 0x5A, 0x0C, 0xCB, 0x12, 0xAF, 0xC7, 0xE2, 0xF8, 0xB4, 0x02, - 0xFC, 0x96, 0x09, 0x08, 0xD6, 0xF5, 0x5B, 0x1E, 0xE2, 0x93, 0x1E, 0xC9, - 0x69, 0x09, 0xB2, 0xA7, 0x38, 0x62, 0xAA, 0x81, 0xF9, 0x58, 0x28, 0xD1, - 0xAB, 0x1F, 0xBC, 0xEA, 0x25, 0x72, 0xB6, 0xB2, 0xFF, 0x5A, 0x55, 0x25, - 0x62, 0xAF, 0x96, 0x27, 0x39, 0x61, 0x70, 0x4F, 0x9E, 0xDE, 0xFA, 0x80, - 0x0E, 0xFF, 0xC5, 0x90, 0x5A, 0x2B, 0x9F, 0x00, 0x42, 0xD4, 0x72, 0xE9, - 0xC0, 0xBC, 0x78, 0xDB, 0x53, 0x15, 0x47, 0x14, 0x17, 0x38, 0x14, 0x50, - 0x6A, 0xF1, 0x4C, 0x3F, 0xF7, 0xAA, 0xD4, 0x92, 0xD5, 0x95, 0xD9, 0xEE, - 0x9C, 0x90, 0xBD, 0x01, 0xD5, 0xC6, 0x03, 0x90, 0xE6, 0x1E, 0xB2, 0xCC, - 0x5E, 0xDA, 0x69, 0xB4, 0x94, 0xD7, 0x1C, 0x9B, 0x33, 0xE6, 0xFF, 0x62, - 0xD8, 0x7C, 0xAF, 0x83, 0x1C, 0xBF, 0x76, 0x4B, 0xAF, 0xF9, 0x74, 0x62, - 0x0D, 0x34, 0xA2, 0xCF, 0x44, 0x1E, 0x68, 0x7D, 0x10, 0x44, 0xFF, 0xE4, - 0x36, 0x0A, 0x6E, 0x6F, 0x95, 0x01, 0xBB, 0xFA, 0x68, 0x78, 0xB6, 0xD9, - 0x9A, 0xC8, 0xEA, 0x1D, 0xBD, 0x53, 0x45, 0xE9, 0xEB, 0xDC, 0xD3, 0x44, - 0x65, 0x5A, 0x63, 0x6A, 0xAF, 0x58, 0x1C, 0x32, 0xA3, 0x0B, 0xBB, 0x1A, - 0xD6, 0x29, 0x69, 0xF8, 0x41, 0x91, 0x63, 0x47, 0x1B, 0x94, 0x26, 0x7B, - 0xA5, 0xF1, 0xA0, 0xC0, 0x86, 0xBE, 0xC3, 0xD1, 0x97, 0x53, 0xE7, 0xA4, - 0xC2, 0x8E, 0xA4, 0xE8, 0x4F, 0xC8, 0x4B, 0x4D, 0xC5, 0x71, 0x78, 0x75, - 0x7B, 0x73, 0x33, 0x9C, 0x33, 0xC9, 0x4B, 0xB6, 0xEF, 0x3C, 0x42, 0xC6, - 0x19, 0xEF, 0xC1, 0x3A, 0xA2, 0x58, 0x41, 0xAF, 0x37, 0x03, 0x1B, 0x1B, - 0x9D, 0xA8, 0x58, 0x7B, 0x89, 0x65, 0x9C, 0x9C, 0xFE, 0xF2, 0x94, 0x15, - 0x0C, 0x99, 0xB8, 0x45, 0x69, 0x7C, 0x5B, 0x02, 0xA2, 0x1A, 0x0D, 0xF0, - 0x5C, 0xCA, 0x6C, 0x2C, 0xAA, 0xAF, 0xC5, 0xD8, 0x29, 0x9F, 0x61, 0x01, - 0x55, 0x2E, 0x3B, 0x70, 0x19, 0xFF, 0x24, 0x85, 0xC6, 0x00, 0x14, 0x93, - 0x0D, 0x8C, 0x75, 0x34, 0xDC, 0xC4, 0x94, 0x4B, 0x26, 0x40, 0xE3, 0x84, - 0xCB, 0x09, 0xC0, 0x82, 0x81, 0xBD, 0xA0, 0x6A, 0x2F, 0xF9, 0xF0, 0x32, - 0xF9, 0xE0, 0x07, 0x54, 0x21, 0xEA, 0x2B, 0xEA, 0x6F, 0xC3, 0x8A, 0x16, - 0x9D, 0xA4, 0x9E, 0x79, 0x63, 0x4D, 0x8F, 0x8D, 0xDB, 0xDF, 0x94, 0x59, - 0x92, 0x9A, 0x7C, 0x52, 0x78, 0x17, 0x84, 0x8C, 0x11, 0xCA, 0xA5, 0x60, - 0xD0, 0x04, 0xC8, 0xD7, 0x74, 0x3A, 0x29, 0x2E, 0x80, 0x3D, 0x5B, 0xB2, - 0xEA, 0x71, 0xFF, 0x61, 0xB8, 0x86, 0x8D, 0xDC, 0xEC, 0x95, 0xC9, 0x93, - 0x87, 0x3F, 0x27, 0x5E, 0xC3, 0x27, 0xC0, 0xF7, 0xA3, 0x3C, 0xD6, 0xE7, - 0xC7, 0x0F, 0x38, 0xBF, 0x8F, 0x10, 0x27, 0x32, 0x00, 0x48, 0x96, 0xD4, - 0x77, 0x15, 0x88, 0xC3, - }; - int8_t scale_codes[] = { - 67, 127, 74, 14, 65, 45, 31, 99, 84, 118, 50, 99, - 76, 104, 20, 56, 50, 67, 103, 97, 64, 71, 12, 127, - 46, 60, 71, 16, 107, 47, 39, 101, 85, 40, 127, 90, - 107, 50, 89, 76, 114, 31, 79, 62, 25, 104, 122, 91, - 102, 63, 79, 86, 52, 21, 29, 109, 40, 109, 105, 96, - 57, 127, 65, 49, 81, 108, 100, 74, 96, 30, 100, 46, - 35, 32, 35, 63, 84, 42, 12, 127, 122, 39, 113, 24, - 32, 127, 71, 55, 50, 43, 47, 126, 25, 27, 53, 62, - 121, 76, 74, 105, 116, 35, 15, 95, 127, 53, 47, 101, - 15, 69, 15, 78, 17, 81, 118, 71, 52, 73, 91, 125, - 39, 104, 20, 127, 45, 61, 111, 103, 43, 67, 25, 73, - 17, 24, 104, 55, 47, 61, 14, 127, 70, 63, 46, 102, - 33, 127, 39, 65, 46, 46, 85, 38, 43, 18, 118, 35, - 87, 77, 31, 43, 65, 113, 54, 121, 92, 26, 88, 14, - 37, 126, 127, 37, 77, 13, 38, 80, 17, 51, 44, 27, - 104, 33, 125, 83, 92, 29, 75, 127, 41, 48, 49, 59, - 120, 115, 57, 127, 19, 60, 70, 78, 76, 103, 103, 23, - 108, 87, 124, 72, 18, 95, 50, 36, 65, 127, 24, 51, - 16, 66, 93, 108, 57, 123, 105, 113, 45, 123, 62, 127, - 125, 22, 85, 86, 67, 124, 103, 14, 124, 93, 99, 91, - 74, 40, 82, 72, 94, 57, 115, 69, 60, 49, 123, 69, - 48, 45, 127, 18, - }; - uint16_t scale_step[] = { - 0x12FE, 0x130E, 0x1306, 0x12CE, 0x12EE, 0x1316, 0x12F6, 0x130E, - 0x12D6, 0x111A, 0x1306, 0x12D6, 0x121C, 0x126D, 0x1245, 0x12F6, - }; - uint16_t A_host[] = { - 0xBF00, 0xBF5B, 0x3F88, 0x3F12, 0x3F54, 0x3E2F, 0x3FB3, 0x3E89, - 0xBECA, 0xBF5C, 0x3F71, 0xBF96, 0x3E9B, 0x3F3F, 0x3D17, 0xBDBE, - 0x3FC8, 0x3FA1, 0xBE37, 0x3FDF, 0xBF3B, 0xBDAF, 0xBD77, 0x3F26, - 0x3EC7, 0xBE69, 0xBE51, 0xBDB1, 0x3DAB, 0xBEC8, 0xBF36, 0xBF92, - 0xBEC1, 0xBEE8, 0xBFC6, 0xBD34, 0xBF3F, 0x3F4C, 0x3FB3, 0x3E35, - 0x3FC5, 0x3D8D, 0x3FC4, 0xBF0C, 0xBEBF, 0x3F27, 0x3E9F, 0xBFEB, - 0x3F5B, 0x3FCE, 0x3F89, 0xC015, 0x3FD8, 0x3E89, 0x401B, 0xBE14, - 0x3F47, 0xBFF2, 0xBF3F, 0xBF38, 0xBF6C, 0xBD45, 0xBE22, 0x3D77, - 0x3FC0, 0x3FE2, 0x3FE1, 0xBF57, 0xBF98, 0x3E27, 0x3E76, 0x3F10, - 0xBFFD, 0xBF01, 0xBF98, 0xBF2F, 0x3F85, 0x3F9D, 0xBFC9, 0x3FA5, - 0x3F2D, 0xBF14, 0x3FB0, 0xBF62, 0x3FA1, 0xBE3E, 0x3FA8, 0xBDD0, - 0x3F85, 0xBF79, 0xBF20, 0x3F68, 0x3F99, 0x3EF5, 0xBEAF, 0xBFBC, - 0x3E92, 0x3F24, 0xBF4B, 0xBE2C, 0x3E29, 0x3EDF, 0xBF6E, 0xBF82, - 0x3F6E, 0xBE9F, 0xBF78, 0xBEA2, 0xBFB4, 0xBC79, 0xBE5B, 0x3F59, - 0xBF91, 0xBF92, 0x3EF2, 0xBF82, 0xBE28, 0x3E3F, 0x3FC3, 0x3F76, - 0xBEF6, 0x3FD7, 0xBD12, 0xBDED, 0xBFCD, 0x400E, 0xBE99, 0xBF01, - 0xBE3B, 0xBF14, 0x4014, 0x3F83, 0x3EC8, 0x3F69, 0x3FFF, 0x3F85, - 0x3DC6, 0x3EFB, 0x3F55, 0x3F3A, 0x3A8B, 0xBF41, 0x3F44, 0x3D32, - 0xBEB5, 0x3F49, 0x3F83, 0x3EE7, 0xBF7F, 0x3F41, 0x3F09, 0xBE9B, - 0x3F9A, 0x3ED8, 0x3F5C, 0x3DB2, 0xBE49, 0x4012, 0xBF99, 0xBD2E, - 0xBFED, 0xBDD6, 0x3E8A, 0xBE8A, 0x3DDC, 0x3E84, 0x3EBA, 0x3FA2, - 0x3BFB, 0x3CB1, 0x3E2F, 0xBFCA, 0xBEE5, 0xBE21, 0xBE5D, 0xBE9C, - 0x3F9C, 0xBDAF, 0x3DCD, 0x3F8E, 0x3F76, 0xBE43, 0x3E1E, 0x3F43, - 0xBF56, 0x3FA9, 0x3F75, 0xBF26, 0xBFEC, 0x3F8D, 0x3F89, 0x3F00, - 0x3F37, 0xBF46, 0x3F58, 0x3F79, 0xBF2E, 0xBE0D, 0x3EC2, 0x3FE3, - 0x3F18, 0xBF0D, 0x3DA2, 0xBF23, 0xBEA4, 0x3FFC, 0x3E98, 0x3D69, - 0x3E66, 0xBF1B, 0xBF9A, 0xBEFC, 0xBFA4, 0xBF0B, 0x3F7C, 0xBD36, - 0xBEAC, 0x3FD5, 0xBF86, 0x3EAE, 0xBF41, 0x3F26, 0x3E5C, 0x3F06, - 0x3E6C, 0x3F15, 0x3F84, 0xBFC4, 0x3E33, 0x3DEE, 0x3F2D, 0xBF29, - 0xBF3C, 0x3EC0, 0xBF0B, 0x3FFF, 0xBF8B, 0xBF16, 0xBF2B, 0xBED5, - 0xBF33, 0xBFBA, 0xBF86, 0x3EE5, 0x3E19, 0xBFD8, 0x3F8E, 0x3F0D, - 0x3E24, 0xBF25, 0xBF33, 0xBFB9, 0x3F8D, 0x3D09, 0xBE71, 0xBD99, - }; - uint16_t expected[] = { - 0xC103, 0xBE0A, 0xC189, 0xC156, 0x4198, 0xC14A, 0xC170, 0x418F, - 0x40F6, 0xC118, 0x411B, 0x417E, 0x4197, 0xC13E, 0x417E, 0xC19E, - }; - // clang-format on - - Tensor* output = setup_and_run( - M, N, K, gs, ql_host, qh_host, scale_codes, scale_step, A_host); - ASSERT_NE(output, nullptr); - EXPECT_EQ(output->size(0), M); - EXPECT_EQ(output->size(1), N); - check_bf16_output(output, expected, M * N, 0.5f); -} - -TEST_F(AOTITorchInt6PlainMMTest, NullInputHandling) { - int64_t M = 2, K = 256, N = 64, gs = 16; - int64_t ng = K / gs; - int64_t n_super = K / 256; - - Tensor* A = create_bf16({M, K}); - Tensor* ql = create_uint8({N, K / 2}); - Tensor* qh = create_uint8({N, K / 4}); - Tensor* scale = create_int8({N, ng}); - Tensor* steps = create_fp16({N, n_super}); - Tensor* output = nullptr; - - EXPECT_EQ( - aoti_torch_cuda_int6_plain_mm(nullptr, ql, qh, scale, steps, gs, &output), - Error::InvalidArgument); - EXPECT_EQ( - aoti_torch_cuda_int6_plain_mm(A, nullptr, qh, scale, steps, gs, &output), - Error::InvalidArgument); - EXPECT_EQ( - aoti_torch_cuda_int6_plain_mm(A, ql, nullptr, scale, steps, gs, &output), - Error::InvalidArgument); - EXPECT_EQ( - aoti_torch_cuda_int6_plain_mm(A, ql, qh, nullptr, steps, gs, &output), - Error::InvalidArgument); - EXPECT_EQ( - aoti_torch_cuda_int6_plain_mm(A, ql, qh, scale, nullptr, gs, &output), - Error::InvalidArgument); - EXPECT_EQ( - aoti_torch_cuda_int6_plain_mm(A, ql, qh, scale, steps, gs, nullptr), - Error::InvalidArgument); -} diff --git a/backends/cuda/runtime/targets.bzl b/backends/cuda/runtime/targets.bzl index b0368d91abc..1131d91444e 100644 --- a/backends/cuda/runtime/targets.bzl +++ b/backends/cuda/runtime/targets.bzl @@ -37,7 +37,6 @@ def define_common_targets(is_fbcode = False): "shims/cuda_guard.cpp", "shims/int4mm.cu", "shims/int5_plain_mm.cu", - "shims/int6_plain_mm.cu", "shims/int8_plain_mm.cu", "shims/memory.cpp", "shims/rand.cu", @@ -50,8 +49,6 @@ def define_common_targets(is_fbcode = False): "shims/int4mm.h", "shims/int5_plain_mm.cuh", "shims/int5_plain_mm.h", - "shims/int6_plain_mm.cuh", - "shims/int6_plain_mm.h", "shims/int8_plain_mm.cuh", "shims/int8_plain_mm.h", "shims/memory.h", diff --git a/backends/cuda/tests/targets.bzl b/backends/cuda/tests/targets.bzl index de487acd542..8d0b1264bbb 100644 --- a/backends/cuda/tests/targets.bzl +++ b/backends/cuda/tests/targets.bzl @@ -138,6 +138,30 @@ def define_common_targets(is_fbcode = False): ), ) + python_unittest_remote_gpu( + name = "test_int6_quantized_gemm", + srcs = [ + "test_int6_dispatch.py", + "test_int6_quantized_gemm.py", + ], + visibility = [ + "//executorch/...", + ], + deps = [ + "//caffe2:torch", + "//executorch/backends/cuda:dp4a_planar_int6_tensor", + "//executorch/backends/cuda:quantize_op_dispatch", + "//executorch/backends/cuda:triton_kernels", + "//executorch/extension/llm/export:gguf", + "//pytorch/ao:torchao", + ], + keep_gpu_sections = True, + remote_execution = re_test_utils.remote_execution( + platform = "gpu-remote-execution", + subplatform = "A100-exclusive", + ), + ) + python_unittest_remote_gpu( name = "test_offgraph_kv", srcs = [ diff --git a/backends/cuda/tests/test_int6_dispatch.py b/backends/cuda/tests/test_int6_dispatch.py index 1b7f9181adc..12e60c0da45 100644 --- a/backends/cuda/tests/test_int6_dispatch.py +++ b/backends/cuda/tests/test_int6_dispatch.py @@ -8,14 +8,13 @@ """Tests for CudaDp4aPlanarInt6Tensor F.linear dispatch via int6_dispatch. These tests validate the eager / trace-time dispatch path — the same code that -torch.export traces through when building the AOTI graph. They do NOT test the -.pte runtime C shim (W6A8 dp4a kernel); that is covered by -test_aoti_torch_cuda_int6_plain_mm.cpp (C++ unit tests). +torch.export traces through when building the AOTI graph. The Triton kernels +themselves are covered by test_int6_quantized_gemm.py. The API contract: after importing int6_dispatch, F.linear / nn.Linear with a CudaDp4aPlanarInt6Tensor weight produce numerically correct results, routed by -batch size (decode M<=4 -> custom op, prefill M>4 -> inline dequant). Routing -tests run without a GPU by recording calls to the decode custom op. +batch size (decode M<=4 -> ``triton::int6_quantized_gemm_m{M}``, everything +else -> inline dequant, never an error). Usage: python -m pytest backends/cuda/tests/test_int6_dispatch.py -v @@ -77,28 +76,36 @@ def _ref_weight(q, scale, group_size, dtype=torch.bfloat16): @contextlib.contextmanager -def _record_int6_plain_mm(): - """Record calls to the decode custom op without needing a GPU. +def _record_int6_kernel_ops(): + """Record which INT6 Triton op the dispatch would launch, without a GPU. - Replaces ``torch.ops.executorch_cuda.int6_plain_mm`` (whose real impl is the - CUDA C shim) with a recorder that computes the result via the eager CPU - dequant, so the dispatch handler still returns a valid tensor. + Replaces ``INT6_QUANTIZED_GEMM.op`` with a recorder whose ops compute the + result via the eager dequant, so the dispatch handler still returns a valid + tensor. """ + from executorch.backends.cuda.triton.kernels.int6_quantized_gemm import ( + INT6_QUANTIZED_GEMM, + ) + calls = [] - def _fake(self, ql, qh, scale, steps, group_size): - calls.append((tuple(self.shape), group_size)) - return _unit_dq_mm_int6(self, ql, qh, scale, steps, group_size) + def _op(bucket): + def run(x, *weight_args): + calls.append((bucket, tuple(x.shape))) + return _unit_dq_mm_int6(x, *weight_args) - with mock.patch.object(torch.ops.executorch_cuda, "int6_plain_mm", _fake): + return run + + with mock.patch.object(INT6_QUANTIZED_GEMM, "op", side_effect=_op): yield calls class TestDispatchRouting(unittest.TestCase): - """Type-based routing: M<=4 -> int6_plain_mm op, M>4 -> inline dequant. + """Type-based routing on CPU: CudaDp4aPlanarInt6Tensor takes inline dequant. - Runs without a GPU by recording calls to the decode custom op and computing - the result with the eager CPU dequant. + These tests run without a GPU. Decode on CUDA traces the Triton kernels + (TestDecodeDispatch); CPU eager cannot launch Triton, so every M takes the + inline dequant and no INT6 op is reached. """ def setUp(self): @@ -109,32 +116,32 @@ def _rel_err(self, out, ref): (out.float() - ref.float()).abs().mean() / ref.float().abs().mean() ).item() - def test_decode_routes_to_int6_plain_mm(self): - """M<=4 routes to the decode custom op.""" + def test_cpu_decode_uses_dequant(self): + """M<=4 on CPU eager takes inline dequant, never a Triton op.""" t, _, _ = _make_int6_tensor(16, 256) x = torch.randn(1, 256, dtype=torch.bfloat16) # M=1 (decode regime) - with _record_int6_plain_mm() as calls: + with _record_int6_kernel_ops() as calls: out = F.linear(x, t) - self.assertEqual(len(calls), 1) + self.assertEqual(calls, []) self.assertEqual(out.shape, (1, 16)) def test_prefill_uses_dequant(self): """M>4 uses inline dequant (no custom op) and is numerically correct.""" t, q, scale = _make_int6_tensor(16, 256) x = torch.randn(8, 256, dtype=torch.bfloat16) # M=8 > 4 (prefill regime) - with _record_int6_plain_mm() as calls: + with _record_int6_kernel_ops() as calls: out = F.linear(x, t) self.assertEqual(calls, []) ref = F.linear(x, _ref_weight(q, scale, 16)) self.assertLess(self._rel_err(out, ref), 0.02) def test_decode_result_matches_reference(self): - """The decode op (eager -> dequant) is numerically correct.""" + """The CPU decode-sized result is numerically correct.""" t, q, scale = _make_int6_tensor(24, 512) x = torch.randn(2, 512, dtype=torch.bfloat16) - with _record_int6_plain_mm() as calls: + with _record_int6_kernel_ops() as calls: out = F.linear(x, t) - self.assertEqual(len(calls), 1) + self.assertEqual(calls, []) ref = F.linear(x, _ref_weight(q, scale, 16)) self.assertLess(self._rel_err(out, ref), 0.02) @@ -143,7 +150,7 @@ def test_with_bias(self): t, q, scale = _make_int6_tensor(16, 256) bias = torch.randn(16, dtype=torch.bfloat16) x = torch.randn(1, 256, dtype=torch.bfloat16) - with _record_int6_plain_mm(): + with _record_int6_kernel_ops(): out = F.linear(x, t, bias) ref = F.linear(x, _ref_weight(q, scale, 16), bias) self.assertLess(self._rel_err(out, ref), 0.02) @@ -153,13 +160,13 @@ def test_with_bias_kwarg(self): t, q, scale = _make_int6_tensor(16, 256) bias = torch.randn(16, dtype=torch.bfloat16) x = torch.randn(1, 256, dtype=torch.bfloat16) - with _record_int6_plain_mm(): + with _record_int6_kernel_ops(): out = F.linear(x, t, bias=bias) ref = F.linear(x, _ref_weight(q, scale, 16), bias) self.assertLess(self._rel_err(out, ref), 0.02) # Guard against a regression to dropping the keyword bias: the no-bias # result must differ from the bias result by exactly the bias. - with _record_int6_plain_mm(): + with _record_int6_kernel_ops(): out_no_bias = F.linear(x, t) self.assertTrue( torch.allclose(out, out_no_bias + bias, atol=1e-2), @@ -170,7 +177,7 @@ def test_3d_batched_input(self): """3D input is flattened and the output shape is restored.""" t, q, scale = _make_int6_tensor(16, 256) x = torch.randn(2, 8, 256, dtype=torch.bfloat16) # flattened M=16 > 4 - with _record_int6_plain_mm() as calls: + with _record_int6_kernel_ops() as calls: out = F.linear(x, t) self.assertEqual(calls, []) # prefill regime self.assertEqual(out.shape, (2, 8, 16)) @@ -195,9 +202,9 @@ def test_from_intx_int8_roundtrip(self): ) t = CudaDp4aPlanarInt6Tensor._from_intx_int8(intx) x = torch.randn(1, K, dtype=torch.bfloat16) - with _record_int6_plain_mm() as calls: + with _record_int6_kernel_ops() as calls: out = F.linear(x, t) - self.assertEqual(len(calls), 1) + self.assertEqual(calls, []) # CPU eager: inline dequant # The packer re-encodes scale as int8 code * per-256 fp16 step, so the # reference uses the effective decoded scale (the tensor's dequant), not # the raw input scale. @@ -317,5 +324,119 @@ def test_dequantize_matches_reference(self): self.assertTrue(torch.equal(t.dequantize(torch.bfloat16).cpu(), ref)) +class TestDecodeDispatch(unittest.TestCase): + """CUDA export: decode-sized M captures the INT6 bucket op.""" + + def setUp(self): + _require_cuda(self) + torch.manual_seed(0) + + def _module(self, n=256, k=512, group_size=16): + t, q, scale = _make_int6_tensor(n, k, group_size) + module = nn.Linear(k, n, bias=False, dtype=torch.bfloat16) + module.weight = nn.Parameter(t, requires_grad=False) + return module.cuda(), _ref_weight(q, scale, group_size).cuda() + + @staticmethod + def _targets(module, x, dynamic_m=None): + from torch.export import Dim + + dynamic = None + if dynamic_m is not None: + dynamic = ({0: Dim("m", min=dynamic_m[0], max=dynamic_m[1])},) + with torch.no_grad(): + program = torch.export.export(module, (x,), dynamic_shapes=dynamic) + return {str(node.target) for node in program.graph.nodes} + + @staticmethod + def _bucket_ops(targets): + return {t for t in targets if "int6_quantized_gemm" in t} + + def test_decode_sized_m_uses_its_bucket_op(self): + module, w_ref = self._module() + for m in (1, 2, 3, 4): + x = torch.randn(m, 512, dtype=torch.bfloat16, device="cuda") + ops = self._bucket_ops(self._targets(module, x)) + self.assertEqual(ops, {f"triton.int6_quantized_gemm_m{m}.default"}, m) + with torch.no_grad(): + out = module(x) + ref = F.linear(x, w_ref) + rel = (out.float() - ref.float()).abs().mean() / ref.float().abs().mean() + self.assertLess(rel.item(), 0.02, m) + + def test_dynamic_m_bounded_by_a_bucket_uses_that_bucket(self): + module, _ = self._module() + x = torch.randn(4, 512, dtype=torch.bfloat16, device="cuda") + self.assertEqual( + self._bucket_ops(self._targets(module, x, dynamic_m=(2, 4))), + {"triton.int6_quantized_gemm_m4.default"}, + ) + self.assertEqual( + self._bucket_ops(self._targets(module, x[:2], dynamic_m=(2, 3))), + {"triton.int6_quantized_gemm_m3.default"}, + ) + + def test_prefill_and_unbounded_dynamic_m_use_dequant(self): + module, _ = self._module() + x8 = torch.randn(8, 512, dtype=torch.bfloat16, device="cuda") + self.assertFalse(self._bucket_ops(self._targets(module, x8))) + self.assertFalse(self._bucket_ops(self._targets(module, x8, dynamic_m=(5, 64)))) + self.assertFalse(self._bucket_ops(self._targets(module, x8, dynamic_m=(1, 64)))) + + +class TestFallbacks(unittest.TestCase): + """Inputs the INT6 kernels do not serve take inline dequant: no error, no + Triton op in the graph, correct output.""" + + def setUp(self): + _require_cuda(self) + torch.manual_seed(5) + + def _check(self, t, q, scale, group_size, x): + with _record_int6_kernel_ops() as calls: + out = F.linear(x, t) + self.assertEqual(calls, []) + ref = F.linear(x.to(torch.bfloat16), _ref_weight(q, scale, group_size).cuda()).to(out.dtype) + rel = (out.float() - ref.float()).abs().mean() / ref.float().abs().mean() + self.assertLess(rel.item(), 0.05) + + def _cuda_tensor(self, n, k, group_size): + t, q, scale = _make_int6_tensor(n, k, group_size) + return t.cuda(), q, scale + + def test_fp16_activation(self): + t, q, scale = self._cuda_tensor(64, 512, 16) + self._check(t, q, scale, 16, torch.randn(1, 512, dtype=torch.float16, device="cuda")) + + def test_non_contiguous_activation(self): + t, q, scale = self._cuda_tensor(64, 512, 16) + x = torch.randn(512, 2, dtype=torch.bfloat16, device="cuda").t() + self._check(t, q, scale, 16, x) + + def test_group_size_other_than_16(self): + t, q, scale = self._cuda_tensor(64, 512, 32) + self._check(t, q, scale, 32, torch.randn(2, 512, dtype=torch.bfloat16, device="cuda")) + + def test_more_than_four_rows(self): + t, q, scale = self._cuda_tensor(64, 512, 16) + self._check(t, q, scale, 16, torch.randn(5, 512, dtype=torch.bfloat16, device="cuda")) + + def test_raw_uint8_scale_codes_are_signed(self): + """uint8 scale storage holds the same signed codes, on the dequant + fallback (gs = 32) as in the kernels.""" + t, _, _ = self._cuda_tensor(64, 512, 32) + codes = t.scale.clone() + codes[:, ::2] = -codes[:, ::2] + x = torch.randn(1, 512, dtype=torch.bfloat16, device="cuda") + ref = _unit_dq_mm_int6(x, t.ql, t.qh, codes, t.steps, 32) + raw = _unit_dq_mm_int6(x, t.ql, t.qh, codes.view(torch.uint8), t.steps, 32) + torch.testing.assert_close(raw, ref) + strided = codes.t().contiguous().t().view(torch.uint8) + self.assertFalse(strided.is_contiguous()) + torch.testing.assert_close( + _unit_dq_mm_int6(x, t.ql, t.qh, strided, t.steps, 32), ref + ) + + if __name__ == "__main__": unittest.main() diff --git a/backends/cuda/tests/test_int6_quantized_gemm.py b/backends/cuda/tests/test_int6_quantized_gemm.py new file mode 100644 index 00000000000..480030e94e7 --- /dev/null +++ b/backends/cuda/tests/test_int6_quantized_gemm.py @@ -0,0 +1,383 @@ +# Copyright (c) Meta Platforms, Inc. and affiliates. +# All rights reserved. +# +# This source code is licensed under the BSD-style license found in the +# LICENSE file in the root directory of this source tree. + +"""Tests for the decode-sized INT6 Triton GEMM (triton/kernels/int6_quantized_gemm.py). + +Every autotune candidate of every bucket must match the W6A16 dequant + +F.linear reference, statically and for a compiled dynamic M; every legality +rule makes ``supports`` False and ``validate`` raise. + + python -m pytest backends/cuda/tests/test_int6_quantized_gemm.py -v +""" + +import unittest +from unittest import mock + +import torch +import triton +from executorch.backends.cuda.quantize_op_dispatch.int6_dispatch import ( + _unit_dq_mm_int6, +) +from executorch.backends.cuda.tests.test_int6_dispatch import _make_int6_tensor +from executorch.backends.cuda.triton.kernels import int6_quantized_gemm as int6_kernel +from executorch.backends.cuda.triton.kernels.int6_quantized_gemm import ( + int6_autotune_configs, + INT6_QUANTIZED_GEMM, + SUPPORTED_BUCKETS, +) + +GROUP_SIZE = 16 + + +def _packed(n: int, k: int, seed: int = 0): + torch.manual_seed(seed) + weight, _, _ = _make_int6_tensor(n, k, GROUP_SIZE) + return tuple( + tensor.cuda() for tensor in (weight.ql, weight.qh, weight.scale, weight.steps) + ) + + +def _check_close(test: unittest.TestCase, out: torch.Tensor, ref: torch.Tensor) -> None: + test.assertTrue(torch.isfinite(out).all()) + torch.testing.assert_close(out.float(), ref.float(), rtol=0.02, atol=4.0) + mean_rel = (out.float() - ref.float()).abs().mean() / ref.float().abs().mean() + test.assertLess(mean_rel.item(), 0.01) + + +def _single_config(config: triton.Config): + return triton.autotune(configs=[config], key=["N", "K", "SPLIT_K"])( + int6_kernel._int6_w6a8_bucket_kernel + ) + + +class Int6QuantizedGemmRulesTest(unittest.TestCase): + def test_candidates_cover_the_generic_space(self) -> None: + configs = int6_autotune_configs() + seen = { + ( + config.kwargs["K_TILE"], + config.kwargs["BLOCK_N"], + config.num_warps, + config.kwargs["PIPELINE_STAGES"], + config.num_stages, + ) + for config in configs + } + self.assertEqual(len(configs), 24) + self.assertEqual( + seen, + { + (k_tile, warps, warps, stages, stages) + for k_tile in (32, 16) + for warps in (1, 2, 4, 8) + for stages in (1, 2, 3) + }, + ) + + def test_prune_drops_stages_above_main_loop_trips(self) -> None: + configs = int6_autotune_configs() + cases = ( + (256, 1, 1, 8), + (2048, 1, 3, 20), + (2048, 2, 2, 12), + (8192, 1, 3, 24), + (8192, 8, 2, 12), + ) + for k, split_k, max_stages, expected_count in cases: + with self.subTest(k=k, split_k=split_k): + kept = int6_kernel._prune(configs, {"K": k}, SPLIT_K=split_k) + self.assertEqual( + max(c.kwargs["PIPELINE_STAGES"] for c in kept), + max_stages, + ) + self.assertEqual(len(kept), expected_count) + + +class Int6QuantizedGemmLegalityTest(unittest.TestCase): + @classmethod + def setUpClass(cls) -> None: + if not torch.cuda.is_available(): + raise unittest.SkipTest("CUDA required") + + def _args(self, m: int = 1, n: int = 64, k: int = 512): + x = torch.randn(m, k, dtype=torch.bfloat16, device="cuda") + return [x, *_packed(n, k, seed=3), GROUP_SIZE] + + def _expect_unsupported(self, bucket, args, pattern) -> None: + self.assertFalse(INT6_QUANTIZED_GEMM.supports(bucket, *args)) + with self.assertRaisesRegex(RuntimeError, pattern): + INT6_QUANTIZED_GEMM.validate(bucket, *args) + + def test_valid_inputs_and_raw_byte_dtypes_are_supported(self) -> None: + for bucket in SUPPORTED_BUCKETS: + args = self._args(m=bucket) + self.assertTrue(INT6_QUANTIZED_GEMM.supports(bucket, *args)) + INT6_QUANTIZED_GEMM.validate(bucket, *args) + + args = self._args() + args[1] = args[1].view(torch.int8) + args[2] = args[2].view(torch.int8) + args[3] = args[3].view(torch.uint8) + self.assertTrue(INT6_QUANTIZED_GEMM.supports(1, *args)) + INT6_QUANTIZED_GEMM.validate(1, *args) + + def test_unsupported_bucket(self) -> None: + args = self._args() + self.assertFalse(INT6_QUANTIZED_GEMM.supports(5, *args)) + with self.assertRaisesRegex( + RuntimeError, "unsupported int6_quantized_gemm bucket 5" + ): + INT6_QUANTIZED_GEMM.validate(5, *args) + + def test_each_dtype_rule(self) -> None: + cases = { + "activation": (0, torch.float16, "activation must be bfloat16"), + "ql": (1, torch.float32, "ql must be"), + "qh": (2, torch.float32, "qh must be"), + "scale": (3, torch.int32, "scale codes must be"), + "steps": (4, torch.float32, "steps must be float16"), + } + for name, (index, dtype, pattern) in cases.items(): + with self.subTest(rule=name): + args = self._args() + args[index] = args[index].to(dtype) + self._expect_unsupported(1, args, pattern) + + def test_each_rank_rule(self) -> None: + for index, name in enumerate(("x", "ql", "qh", "scale", "steps")): + with self.subTest(tensor=name): + args = self._args() + args[index] = args[index].unsqueeze(0) + self._expect_unsupported(1, args, "rank-2") + + def test_group_size_and_static_m_rules(self) -> None: + args = self._args() + args[5] = 32 + self._expect_unsupported(1, args, "group_size must be 16") + + args = self._args(m=2) + self._expect_unsupported(1, args, "static M must equal") + + def test_k_must_be_positive_static_multiple_of_256(self) -> None: + n, k = 32, 384 + args = [ + torch.randn(1, k, dtype=torch.bfloat16, device="cuda"), + torch.zeros(n, k // 2, dtype=torch.uint8, device="cuda"), + torch.zeros(n, k // 4, dtype=torch.uint8, device="cuda"), + torch.zeros(n, k // GROUP_SIZE, dtype=torch.int8, device="cuda"), + torch.zeros(n, k // 256, dtype=torch.float16, device="cuda"), + GROUP_SIZE, + ] + self._expect_unsupported(1, args, "K must be a multiple of 256") + + args = [ + torch.empty(1, 0, dtype=torch.bfloat16, device="cuda"), + torch.empty(n, 0, dtype=torch.uint8, device="cuda"), + torch.empty(n, 0, dtype=torch.uint8, device="cuda"), + torch.empty(n, 0, dtype=torch.int8, device="cuda"), + torch.empty(n, 0, dtype=torch.float16, device="cuda"), + GROUP_SIZE, + ] + self._expect_unsupported(1, args, "K must be positive") + + def test_symbolic_k_and_unprovable_m_are_unsupported(self) -> None: + from torch._subclasses.fake_tensor import FakeTensorMode + from torch.fx.experimental.symbolic_shapes import ShapeEnv + + shape_env = ShapeEnv() + mode = FakeTensorMode(shape_env=shape_env) + with mode: + weights = [ + torch.empty(32, 256, dtype=torch.uint8, device="cuda"), + torch.empty(32, 128, dtype=torch.uint8, device="cuda"), + torch.empty(32, 32, dtype=torch.int8, device="cuda"), + torch.empty(32, 2, dtype=torch.float16, device="cuda"), + ] + symbolic_k = shape_env.create_unbacked_symint() + dynamic_k_args = [ + torch.empty(1, symbolic_k, dtype=torch.bfloat16, device="cuda"), + *weights, + GROUP_SIZE, + ] + self._expect_unsupported(1, dynamic_k_args, "K must be static") + + symbolic_m = shape_env.create_unbacked_symint() + dynamic_m_args = [ + torch.empty(symbolic_m, 512, dtype=torch.bfloat16, device="cuda"), + *weights, + GROUP_SIZE, + ] + self._expect_unsupported(4, dynamic_m_args, "dynamic M is not provably") + + def test_each_shape_rule(self) -> None: + cases = { + "ql": (1, "ql K/2 mismatch"), + "qh": (2, "qh shape"), + "scale": (3, "scale shape"), + "steps": (4, "steps shape"), + } + for name, (index, pattern) in cases.items(): + with self.subTest(tensor=name): + args = self._args() + args[index] = args[index][:, :-1].contiguous() + self._expect_unsupported(1, args, pattern) + + def test_contiguous_and_device_rules(self) -> None: + for index, name in enumerate(("x", "ql", "qh", "scale", "steps")): + with self.subTest(contiguous=name): + args = self._args() + args[index] = torch.stack((args[index], args[index]), dim=-1)[..., 0] + self.assertFalse(args[index].is_contiguous()) + self._expect_unsupported(1, args, "contiguous") + + args = self._args() + args[3] = args[3].cpu() + self._expect_unsupported(1, args, "same device") + + args = self._args() + for index in range(5): + args[index] = args[index].cpu() + self._expect_unsupported(1, args, "CUDA device") + + def test_fake_inputs_are_supported(self) -> None: + from torch._subclasses.fake_tensor import FakeTensorMode + + real = self._args() + mode = FakeTensorMode() + fake = [mode.from_tensor(t.cpu()) for t in real[:5]] + [GROUP_SIZE] + self.assertTrue(INT6_QUANTIZED_GEMM.supports(1, *fake)) + INT6_QUANTIZED_GEMM.validate(1, *fake) + + def test_op_validates_before_launching(self) -> None: + args = self._args(m=2) + with self.assertRaisesRegex(RuntimeError, "static M must equal"): + INT6_QUANTIZED_GEMM.op(1)(*args) + + +class Int6QuantizedGemmTest(unittest.TestCase): + @classmethod + def setUpClass(cls) -> None: + if not torch.cuda.is_available(): + raise unittest.SkipTest("CUDA required") + + def test_ops_are_registered_per_bucket(self) -> None: + self.assertEqual(SUPPORTED_BUCKETS, (1, 2, 3, 4)) + for bucket in SUPPORTED_BUCKETS: + self.assertTrue(hasattr(torch.ops.triton, f"int6_quantized_gemm_m{bucket}")) + + def test_every_static_candidate_matches_w6a16_reference(self) -> None: + shapes = ((37, 256), (53, 512), (37, 5376)) + for n, k in shapes: + weights = _packed(n, k, seed=n + k) + for bucket in SUPPORTED_BUCKETS: + x = torch.randn(bucket, k, dtype=torch.bfloat16, device="cuda") + ref = _unit_dq_mm_int6(x, *weights, GROUP_SIZE) + for config in int6_autotune_configs(): + with self.subTest( + n=n, + k=k, + bucket=bucket, + config=str(config), + ): + single = _single_config(config) + with mock.patch.dict( + int6_kernel._BUCKET_KERNELS, {bucket: single} + ): + out = int6_kernel._launch(bucket, x, *weights, GROUP_SIZE) + self.assertEqual(out.shape, (bucket, n)) + self.assertEqual(out.dtype, torch.bfloat16) + _check_close(self, out, ref) + + def test_every_candidate_serves_dynamic_m(self) -> None: + for n, k in ((37, 256), (37, 512), (37, 5376)): + weights = _packed(n, k, seed=71 + k) + for bucket in (3, 4): + for config in int6_autotune_configs(): + with self.subTest(n=n, k=k, bucket=bucket, config=str(config)): + torch._dynamo.reset() + single = _single_config(config) + with mock.patch.dict( + int6_kernel._BUCKET_KERNELS, {bucket: single} + ): + compiled = torch.compile( + lambda x: INT6_QUANTIZED_GEMM.op(bucket)( + x, *weights, GROUP_SIZE + ), + fullgraph=True, + ) + for m in range(2, bucket + 1): + x = torch.randn( + m, k, dtype=torch.bfloat16, device="cuda" + ) + torch._dynamo.mark_dynamic(x, 0, min=2, max=bucket) + out = compiled(x) + self.assertEqual(out.shape, (m, n)) + _check_close( + self, + out, + _unit_dq_mm_int6(x, *weights, GROUP_SIZE), + ) + + def test_dynamic_bucket4_exact_rows_matches_reference(self) -> None: + n, k = 37, 512 + weights = _packed(n, k, seed=79) + config = next( + config + for config in int6_autotune_configs() + if config.kwargs["K_TILE"] == 32 + and config.kwargs["BLOCK_N"] == 1 + and config.kwargs["PIPELINE_STAGES"] == 1 + ) + single = _single_config(config) + with mock.patch.dict(int6_kernel._BUCKET_KERNELS, {4: single}): + compiled = torch.compile( + lambda x: INT6_QUANTIZED_GEMM.op(4)(x, *weights, GROUP_SIZE), + fullgraph=True, + ) + for m in (2, 3, 4): + x = torch.randn(m, k, dtype=torch.bfloat16, device="cuda") + torch._dynamo.mark_dynamic(x, 0, min=2, max=4) + out = compiled(x) + self.assertEqual(out.shape, (m, n)) + _check_close( + self, + out, + _unit_dq_mm_int6(x, *weights, GROUP_SIZE), + ) + + def test_raw_byte_storage_dtypes_execute_correctly(self) -> None: + original = _packed(37, 512, seed=81) + weights = list(original) + weights[0] = weights[0].view(torch.int8) + weights[1] = weights[1].view(torch.int8) + weights[2] = weights[2].view(torch.uint8) + x = torch.randn(1, 512, dtype=torch.bfloat16, device="cuda") + out = INT6_QUANTIZED_GEMM.op(1)(x, *weights, GROUP_SIZE) + ref = _unit_dq_mm_int6(x, *original, GROUP_SIZE) + _check_close(self, out, ref) + + def test_split_k_extremes_match_reference(self) -> None: + weights = _packed(37, 768, seed=101) + for sm_count in (1, 132): + with mock.patch.object(int6_kernel, "_sm_count", return_value=sm_count): + for bucket in SUPPORTED_BUCKETS: + with self.subTest(sm_count=sm_count, bucket=bucket): + x = torch.randn( + bucket, + 768, + dtype=torch.bfloat16, + device="cuda", + ) + out = INT6_QUANTIZED_GEMM.op(bucket)(x, *weights, GROUP_SIZE) + _check_close( + self, + out, + _unit_dq_mm_int6(x, *weights, GROUP_SIZE), + ) + + +if __name__ == "__main__": + unittest.main() diff --git a/backends/cuda/tests/test_sort_shim.py b/backends/cuda/tests/test_sort_shim.py index 8bd0d98cbdc..50463259d76 100644 --- a/backends/cuda/tests/test_sort_shim.py +++ b/backends/cuda/tests/test_sort_shim.py @@ -43,8 +43,6 @@ "aoti_torch_cuda_randint_low_out", "executorch_cuda::int5_plain_mm", "aoti_torch_cuda_int5_plain_mm", - "executorch_cuda::int6_plain_mm", - "aoti_torch_cuda_int6_plain_mm", "executorch_cuda::int8_plain_mm", "aoti_torch_cuda_int8_plain_mm", } @@ -154,7 +152,7 @@ def test_cuda_shim_map_unchanged_by_rocm_gate(self): with patch.object(torch.version, "hip", None): options = CudaBackend.get_aoti_compile_options([]) - self.assertEqual(len(options["aot_inductor.custom_ops_to_c_shims"]), 3) + self.assertEqual(len(options["aot_inductor.custom_ops_to_c_shims"]), 2) self.assertNotIn("aot_inductor.precompile_headers", options) def test_rocm_advertises_no_unbuilt_shims(self): diff --git a/backends/cuda/triton/kernels/int6_quantized_gemm.py b/backends/cuda/triton/kernels/int6_quantized_gemm.py new file mode 100644 index 00000000000..705109a8bb7 --- /dev/null +++ b/backends/cuda/triton/kernels/int6_quantized_gemm.py @@ -0,0 +1,1205 @@ +# Copyright (c) Meta Platforms, Inc. and affiliates. +# All rights reserved. +# +# This source code is licensed under the BSD-style license found in the +# LICENSE file in the root directory of this source tree. +# Optimized by KernelAgent-Oink(https://github.com/meta-pytorch/KernelAgent) + +"""Decode-sized (M <= 4) INT6 GEMM on planar GGUF Q6_K weights. + +The kernels consume ``CudaDp4aPlanarInt6Tensor`` storage directly: ``ql`` holds +nibble-packed low bits, ``qh`` holds the two high bits in K32 even/odd planes, +``scale`` holds signed K16 scale codes, and ``steps`` holds FP16 K256 scale +steps. Activations are quantized once per K32 block to signed INT8. The GEMM +reconstructs unsigned six-bit weights, uses DP4A, applies the constant -32 and +both scales in FP32, accumulates in FP32, and returns BF16. + +One op is registered for each bucket in M={1,2,3,4}. Each bucket autotunes the +generic configuration space over explicit-row K16 and K32 implementations, +1/2/4/8 output warps, and 1/2/3 pipeline stages. Split-K is shape/device derived +and reduced deterministically. Dynamic buckets branch uniformly to the explicit +kernel with exactly the runtime row count. +""" + +from typing import Optional + +import torch +import triton +import triton.language as tl +from executorch.backends.cuda.triton.kernels.quantized_gemm_family import ( + launch_split_k_gemm, + QuantizedGemmFamily, +) +from executorch.backends.cuda.triton.kernels.quantized_gemm_utils import ( + _device_index, + _dp4a_u8_s8, + _sm_count, + _warp_sum_f32, + autotune_configs, + check_contiguous, + check_device, + check_dtypes, + check_k, + check_rows, + first_reason, + prune_by_main_loop_trips, + quantize_activations_q8, + split_k_for, +) + +_GROUP_SIZE = 16 +_SUPER_BLOCK = 256 +_TL_GROUP_SIZE = tl.constexpr(16) +_TL_Q8_BLOCK = tl.constexpr(32) + +SUPPORTED_BUCKETS = (1, 2, 3, 4) + + +@triton.jit +def _spread_high2(high_byte): + """Spread four packed two-bit fields into four uint8 lanes.""" + value = (high_byte | (high_byte << 12)) & 0x000F000F + return (value | (value << 6)) & 0x03030303 + + +@triton.jit +def _signed_byte_to_f32(value): + """Interpret an int8 or uint8 tensor element as a signed byte.""" + return value.to(tl.int8, bitcast=True).to(tl.float32) + + +@triton.jit +def _load_int6_k32( + ql, + qh, + scale, + steps, + offs_n, + block, + block_mask, + stride_qln: tl.constexpr, + stride_qlk: tl.constexpr, + stride_qhn: tl.constexpr, + stride_qhk: tl.constexpr, + stride_sn: tl.constexpr, + stride_sk: tl.constexpr, + stride_stn: tl.constexpr, + stride_stk: tl.constexpr, +): + """Load and reconstruct one K32 block for every lane.""" + word = tl.arange(0, 4) + ql_base = ql + offs_n * stride_qln + block * 16 * stride_qlk + packed_low = tl.load( + ql_base.to(tl.pointer_type(tl.uint32))[:, None] + word[None, :], + mask=block_mask[:, None], + other=0, + ) + low_even = packed_low & 0x0F0F0F0F + low_odd = (packed_low >> 4) & 0x0F0F0F0F + + qh_base = qh + offs_n * stride_qhn + block * 8 * stride_qhk + high_even_word = tl.load( + qh_base.to(tl.pointer_type(tl.uint32)), mask=block_mask, other=0 + ) + high_odd_word = tl.load( + qh_base.to(tl.pointer_type(tl.uint32)) + 1, mask=block_mask, other=0 + ) + shift = word * 8 + high_even = (high_even_word[:, None] >> shift[None, :]) & 0xFF + high_odd = (high_odd_word[:, None] >> shift[None, :]) & 0xFF + weight_even = low_even | (_spread_high2(high_even) << 4) + weight_odd = low_odd | (_spread_high2(high_odd) << 4) + + group = block[:, None] * 2 + tl.arange(0, 2)[None, :] + scale_code = tl.load( + scale + offs_n[:, None] * stride_sn + group * stride_sk, + mask=block_mask[:, None], + other=0, + ) + scale_code = _signed_byte_to_f32(scale_code) + step = tl.load( + steps + offs_n * stride_stn + (block // 8) * stride_stk, + mask=block_mask, + other=0.0, + ).to(tl.float32) + return weight_even, weight_odd, scale_code * step[:, None] + + +@triton.jit +def _int6_row_contribution( + qwords, + x_scale, + weight_even, + weight_odd, + weight_scale, + block, + block_mask, + row, + blocks32: tl.constexpr, + BLOCK_N: tl.constexpr, +): + """Return one activation row's K32 contribution for every lane.""" + word = tl.arange(0, 4) + qword_base = (row * blocks32 + block) * 8 + activation_even = tl.load( + qwords + qword_base[:, None] + word[None, :], + mask=block_mask[:, None], + other=0, + ) + activation_odd = tl.load( + qwords + qword_base[:, None] + 4 + word[None, :], + mask=block_mask[:, None], + other=0, + ) + dot_words = _dp4a_u8_s8( + weight_even, + activation_even, + tl.zeros((BLOCK_N * 32, 4), dtype=tl.int32), + ) + dot_words = _dp4a_u8_s8(weight_odd, activation_odd, dot_words) + ones = tl.full((BLOCK_N * 32, 4), 0x01010101, dtype=tl.uint32) + sum_words = _dp4a_u8_s8( + ones, + activation_even, + tl.zeros((BLOCK_N * 32, 4), dtype=tl.int32), + ) + sum_words = _dp4a_u8_s8(ones, activation_odd, sum_words) + corrected_words = dot_words - 32 * sum_words + corrected = tl.sum( + tl.reshape(corrected_words, (BLOCK_N * 32, 2, 2), can_reorder=False), + axis=2, + ).to(tl.float32) + activation_scale = tl.load( + x_scale + row * blocks32 + block, mask=block_mask, other=0.0 + ).to(tl.float32) + contribution = tl.sum(corrected * weight_scale, axis=1) + return tl.where(block_mask, contribution * activation_scale, 0.0) + + +@triton.jit +def _int6_row_contribution_if_active( + qwords, + x_scale, + weight_even, + weight_odd, + weight_scale, + block, + block_mask, + row, + M, + blocks32: tl.constexpr, + BLOCK_N: tl.constexpr, + DYNAMIC_M: tl.constexpr, +): + """Dynamic exact rows are active; mask a static inactive row.""" + if DYNAMIC_M: + return _int6_row_contribution( + qwords, + x_scale, + weight_even, + weight_odd, + weight_scale, + block, + block_mask, + row, + blocks32, + BLOCK_N, + ) + else: + return _int6_row_contribution( + qwords, + x_scale, + weight_even, + weight_odd, + weight_scale, + block, + block_mask & (row < M), + row, + blocks32, + BLOCK_N, + ) + + +@triton.jit +def _int6_w6a8_explicit_kernel( + qwords, + x_scale, + ql, + qh, + scale, + steps, + out, + M, + N: tl.constexpr, + K: tl.constexpr, + stride_qln: tl.constexpr, + stride_qlk: tl.constexpr, + stride_qhn: tl.constexpr, + stride_qhk: tl.constexpr, + stride_sn: tl.constexpr, + stride_sk: tl.constexpr, + stride_stn: tl.constexpr, + stride_stk: tl.constexpr, + stride_os: tl.constexpr, + stride_om: tl.constexpr, + stride_on: tl.constexpr, + BUCKET: tl.constexpr, + BLOCK_N: tl.constexpr, + SPLIT_K: tl.constexpr, + PIPELINE_STAGES: tl.constexpr, + DYNAMIC_M: tl.constexpr, +): + """Row-count-specialized W6A8 kernel with explicit accumulators.""" + pid_n = tl.program_id(0) + split_id = tl.program_id(2) + thread = tl.arange(0, BLOCK_N * 32) + warp = thread // 32 + lane = thread % 32 + offs_n = pid_n * BLOCK_N + warp + n_mask = offs_n < N + blocks32: tl.constexpr = K // _TL_Q8_BLOCK + blocks_per_split: tl.constexpr = tl.cdiv(blocks32, SPLIT_K) + first_block = split_id * blocks_per_split + last_block = tl.minimum(first_block + blocks_per_split, blocks32) + partial0 = tl.zeros((BLOCK_N * 32,), dtype=tl.float32) + partial1 = tl.zeros((BLOCK_N * 32,), dtype=tl.float32) + partial2 = tl.zeros((BLOCK_N * 32,), dtype=tl.float32) + partial3 = tl.zeros((BLOCK_N * 32,), dtype=tl.float32) + + for block_offset in tl.range(0, blocks_per_split, 32, num_stages=PIPELINE_STAGES): + block = first_block + block_offset + lane + block_mask = n_mask & (block < last_block) + weight_even, weight_odd, weight_scale = _load_int6_k32( + ql, + qh, + scale, + steps, + offs_n, + block, + block_mask, + stride_qln, + stride_qlk, + stride_qhn, + stride_qhk, + stride_sn, + stride_sk, + stride_stn, + stride_stk, + ) + partial0 += _int6_row_contribution( + qwords, + x_scale, + weight_even, + weight_odd, + weight_scale, + block, + block_mask, + 0, + blocks32, + BLOCK_N, + ) + if BUCKET >= 2: + partial1 += _int6_row_contribution_if_active( + qwords, + x_scale, + weight_even, + weight_odd, + weight_scale, + block, + block_mask, + 1, + M, + blocks32, + BLOCK_N, + DYNAMIC_M, + ) + if BUCKET >= 3: + partial2 += _int6_row_contribution_if_active( + qwords, + x_scale, + weight_even, + weight_odd, + weight_scale, + block, + block_mask, + 2, + M, + blocks32, + BLOCK_N, + DYNAMIC_M, + ) + if BUCKET >= 4: + partial3 += _int6_row_contribution_if_active( + qwords, + x_scale, + weight_even, + weight_odd, + weight_scale, + block, + block_mask, + 3, + M, + blocks32, + BLOCK_N, + DYNAMIC_M, + ) + + result0 = _warp_sum_f32(partial0) + result1 = _warp_sum_f32(partial1) + result2 = _warp_sum_f32(partial2) + result3 = _warp_sum_f32(partial3) + out_base = out + split_id * stride_os + offs_n * stride_on + store_mask = n_mask & (lane == 0) + tl.store( + out_base, + result0.to(tl.float32) if SPLIT_K > 1 else result0.to(tl.bfloat16), + mask=store_mask, + ) + if BUCKET >= 2: + tl.store( + out_base + stride_om, + result1.to(tl.float32) if SPLIT_K > 1 else result1.to(tl.bfloat16), + mask=store_mask if DYNAMIC_M else store_mask & (1 < M), + ) + if BUCKET >= 3: + tl.store( + out_base + 2 * stride_om, + result2.to(tl.float32) if SPLIT_K > 1 else result2.to(tl.bfloat16), + mask=store_mask if DYNAMIC_M else store_mask & (2 < M), + ) + if BUCKET >= 4: + tl.store( + out_base + 3 * stride_om, + result3.to(tl.float32) if SPLIT_K > 1 else result3.to(tl.bfloat16), + mask=store_mask if DYNAMIC_M else store_mask & (3 < M), + ) + + +@triton.jit +def _load_int6_k16( + ql, + qh, + scale, + steps, + offs_n, + group, + group_mask, + stride_qln: tl.constexpr, + stride_qlk: tl.constexpr, + stride_qhn: tl.constexpr, + stride_qhk: tl.constexpr, + stride_sn: tl.constexpr, + stride_sk: tl.constexpr, + stride_stn: tl.constexpr, + stride_stk: tl.constexpr, +): + """Load and reconstruct one K16 group for every lane.""" + word = tl.arange(0, 2) + ql_base = ql + offs_n * stride_qln + group * 8 * stride_qlk + packed_low = tl.load( + ql_base.to(tl.pointer_type(tl.uint32))[:, None] + word[None, :], + mask=group_mask[:, None], + other=0, + ) + low_even = packed_low & 0x0F0F0F0F + low_odd = (packed_low >> 4) & 0x0F0F0F0F + + block = group // 2 + word_in_block = (group % 2)[:, None] * 2 + word[None, :] + qh_base = qh + offs_n * stride_qhn + block * 8 * stride_qhk + high_even = ( + tl.load( + qh_base[:, None] + word_in_block * stride_qhk, + mask=group_mask[:, None], + other=0, + ) + .to(tl.uint8) + .to(tl.uint32) + ) + high_odd = ( + tl.load( + qh_base[:, None] + (4 + word_in_block) * stride_qhk, + mask=group_mask[:, None], + other=0, + ) + .to(tl.uint8) + .to(tl.uint32) + ) + weight_even = low_even | (_spread_high2(high_even) << 4) + weight_odd = low_odd | (_spread_high2(high_odd) << 4) + + scale_code = tl.load( + scale + offs_n * stride_sn + group * stride_sk, + mask=group_mask, + other=0, + ) + step = tl.load( + steps + offs_n * stride_stn + (block // 8) * stride_stk, + mask=group_mask, + other=0.0, + ).to(tl.float32) + return ( + weight_even, + weight_odd, + _signed_byte_to_f32(scale_code) * step, + block, + word_in_block, + ) + + +@triton.jit +def _int6_k16_row_contribution( + qwords, + x_scale, + weight_even, + weight_odd, + weight_scale, + block, + word_in_block, + group_mask, + row, + blocks32: tl.constexpr, + BLOCK_N: tl.constexpr, +): + qword_base = (row * blocks32 + block) * 8 + activation_even = tl.load( + qwords + qword_base[:, None] + word_in_block, + mask=group_mask[:, None], + other=0, + ) + activation_odd = tl.load( + qwords + qword_base[:, None] + 4 + word_in_block, + mask=group_mask[:, None], + other=0, + ) + dot_words = _dp4a_u8_s8( + weight_even, + activation_even, + tl.zeros((BLOCK_N * 32, 2), dtype=tl.int32), + ) + dot_words = _dp4a_u8_s8(weight_odd, activation_odd, dot_words) + ones = tl.full((BLOCK_N * 32, 2), 0x01010101, dtype=tl.uint32) + sum_words = _dp4a_u8_s8( + ones, + activation_even, + tl.zeros((BLOCK_N * 32, 2), dtype=tl.int32), + ) + sum_words = _dp4a_u8_s8(ones, activation_odd, sum_words) + corrected = tl.sum(dot_words - 32 * sum_words, axis=1).to(tl.float32) + activation_scale = tl.load( + x_scale + row * blocks32 + block, + mask=group_mask, + other=0.0, + ).to(tl.float32) + return tl.where( + group_mask, + corrected * weight_scale * activation_scale, + 0.0, + ) + + +@triton.jit +def _int6_k16_row_contribution_if_active( + qwords, + x_scale, + weight_even, + weight_odd, + weight_scale, + block, + word_in_block, + group_mask, + row, + M, + blocks32: tl.constexpr, + BLOCK_N: tl.constexpr, + DYNAMIC_M: tl.constexpr, +): + if DYNAMIC_M: + return _int6_k16_row_contribution( + qwords, + x_scale, + weight_even, + weight_odd, + weight_scale, + block, + word_in_block, + group_mask, + row, + blocks32, + BLOCK_N, + ) + else: + return _int6_k16_row_contribution( + qwords, + x_scale, + weight_even, + weight_odd, + weight_scale, + block, + word_in_block, + group_mask & (row < M), + row, + blocks32, + BLOCK_N, + ) + + +@triton.jit +def _int6_w6a8_k16_kernel( + qwords, + x_scale, + ql, + qh, + scale, + steps, + out, + M, + N: tl.constexpr, + K: tl.constexpr, + stride_qln: tl.constexpr, + stride_qlk: tl.constexpr, + stride_qhn: tl.constexpr, + stride_qhk: tl.constexpr, + stride_sn: tl.constexpr, + stride_sk: tl.constexpr, + stride_stn: tl.constexpr, + stride_stk: tl.constexpr, + stride_os: tl.constexpr, + stride_om: tl.constexpr, + stride_on: tl.constexpr, + BUCKET: tl.constexpr, + BLOCK_N: tl.constexpr, + SPLIT_K: tl.constexpr, + PIPELINE_STAGES: tl.constexpr, + DYNAMIC_M: tl.constexpr, +): + """Row-count-specialized K16 schedule with smaller per-lane vectors.""" + pid_n = tl.program_id(0) + split_id = tl.program_id(2) + thread = tl.arange(0, BLOCK_N * 32) + warp = thread // 32 + lane = thread % 32 + offs_n = pid_n * BLOCK_N + warp + n_mask = offs_n < N + groups16: tl.constexpr = K // _TL_GROUP_SIZE + blocks32: tl.constexpr = K // _TL_Q8_BLOCK + groups_per_split: tl.constexpr = tl.cdiv(groups16, SPLIT_K) + first_group = split_id * groups_per_split + last_group = tl.minimum(first_group + groups_per_split, groups16) + partial0 = tl.zeros((BLOCK_N * 32,), dtype=tl.float32) + partial1 = tl.zeros((BLOCK_N * 32,), dtype=tl.float32) + partial2 = tl.zeros((BLOCK_N * 32,), dtype=tl.float32) + partial3 = tl.zeros((BLOCK_N * 32,), dtype=tl.float32) + + for group_offset in tl.range(0, groups_per_split, 32, num_stages=PIPELINE_STAGES): + group = first_group + group_offset + lane + group_mask = n_mask & (group < last_group) + weight_even, weight_odd, weight_scale, block, word_in_block = _load_int6_k16( + ql, + qh, + scale, + steps, + offs_n, + group, + group_mask, + stride_qln, + stride_qlk, + stride_qhn, + stride_qhk, + stride_sn, + stride_sk, + stride_stn, + stride_stk, + ) + partial0 += _int6_k16_row_contribution( + qwords, + x_scale, + weight_even, + weight_odd, + weight_scale, + block, + word_in_block, + group_mask, + 0, + blocks32, + BLOCK_N, + ) + if BUCKET >= 2: + partial1 += _int6_k16_row_contribution_if_active( + qwords, + x_scale, + weight_even, + weight_odd, + weight_scale, + block, + word_in_block, + group_mask, + 1, + M, + blocks32, + BLOCK_N, + DYNAMIC_M, + ) + if BUCKET >= 3: + partial2 += _int6_k16_row_contribution_if_active( + qwords, + x_scale, + weight_even, + weight_odd, + weight_scale, + block, + word_in_block, + group_mask, + 2, + M, + blocks32, + BLOCK_N, + DYNAMIC_M, + ) + if BUCKET >= 4: + partial3 += _int6_k16_row_contribution_if_active( + qwords, + x_scale, + weight_even, + weight_odd, + weight_scale, + block, + word_in_block, + group_mask, + 3, + M, + blocks32, + BLOCK_N, + DYNAMIC_M, + ) + + result0 = _warp_sum_f32(partial0) + result1 = _warp_sum_f32(partial1) + result2 = _warp_sum_f32(partial2) + result3 = _warp_sum_f32(partial3) + out_base = out + split_id * stride_os + offs_n * stride_on + store_mask = n_mask & (lane == 0) + tl.store( + out_base, + result0.to(tl.float32) if SPLIT_K > 1 else result0.to(tl.bfloat16), + mask=store_mask, + ) + if BUCKET >= 2: + tl.store( + out_base + stride_om, + result1.to(tl.float32) if SPLIT_K > 1 else result1.to(tl.bfloat16), + mask=store_mask if DYNAMIC_M else store_mask & (1 < M), + ) + if BUCKET >= 3: + tl.store( + out_base + 2 * stride_om, + result2.to(tl.float32) if SPLIT_K > 1 else result2.to(tl.bfloat16), + mask=store_mask if DYNAMIC_M else store_mask & (2 < M), + ) + if BUCKET >= 4: + tl.store( + out_base + 3 * stride_om, + result3.to(tl.float32) if SPLIT_K > 1 else result3.to(tl.bfloat16), + mask=store_mask if DYNAMIC_M else store_mask & (3 < M), + ) + + +@triton.jit +def _int6_w6a8_rows( + qwords, + x_scale, + ql, + qh, + scale, + steps, + out, + M, + N: tl.constexpr, + K: tl.constexpr, + stride_qln: tl.constexpr, + stride_qlk: tl.constexpr, + stride_qhn: tl.constexpr, + stride_qhk: tl.constexpr, + stride_sn: tl.constexpr, + stride_sk: tl.constexpr, + stride_stn: tl.constexpr, + stride_stk: tl.constexpr, + stride_os: tl.constexpr, + stride_om: tl.constexpr, + stride_on: tl.constexpr, + ROWS: tl.constexpr, + BLOCK_N: tl.constexpr, + SPLIT_K: tl.constexpr, + PIPELINE_STAGES: tl.constexpr, + DYNAMIC_M: tl.constexpr, + K_TILE: tl.constexpr, +): + """The explicit ROWS-row kernel for the selected K tile.""" + if K_TILE == 32: + _int6_w6a8_explicit_kernel( + qwords, + x_scale, + ql, + qh, + scale, + steps, + out, + M, + N, + K, + stride_qln, + stride_qlk, + stride_qhn, + stride_qhk, + stride_sn, + stride_sk, + stride_stn, + stride_stk, + stride_os, + stride_om, + stride_on, + ROWS, + BLOCK_N, + SPLIT_K, + PIPELINE_STAGES, + DYNAMIC_M, + ) + else: + _int6_w6a8_k16_kernel( + qwords, + x_scale, + ql, + qh, + scale, + steps, + out, + M, + N, + K, + stride_qln, + stride_qlk, + stride_qhn, + stride_qhk, + stride_sn, + stride_sk, + stride_stn, + stride_stk, + stride_os, + stride_om, + stride_on, + ROWS, + BLOCK_N, + SPLIT_K, + PIPELINE_STAGES, + DYNAMIC_M, + ) + + +@triton.jit +def _int6_w6a8_exact_rows( + qwords, + x_scale, + ql, + qh, + scale, + steps, + out, + M, + N: tl.constexpr, + K: tl.constexpr, + stride_qln: tl.constexpr, + stride_qlk: tl.constexpr, + stride_qhn: tl.constexpr, + stride_qhk: tl.constexpr, + stride_sn: tl.constexpr, + stride_sk: tl.constexpr, + stride_stn: tl.constexpr, + stride_stk: tl.constexpr, + stride_os: tl.constexpr, + stride_om: tl.constexpr, + stride_on: tl.constexpr, + ROWS: tl.constexpr, + BLOCK_N: tl.constexpr, + SPLIT_K: tl.constexpr, + PIPELINE_STAGES: tl.constexpr, + DYNAMIC_M: tl.constexpr, + K_TILE: tl.constexpr, +): + """Dynamic M: branch uniformly to the explicit kernel with exactly M rows.""" + if M == 1: + _int6_w6a8_rows( + qwords, + x_scale, + ql, + qh, + scale, + steps, + out, + M, + N, + K, + stride_qln, + stride_qlk, + stride_qhn, + stride_qhk, + stride_sn, + stride_sk, + stride_stn, + stride_stk, + stride_os, + stride_om, + stride_on, + 1, + BLOCK_N, + SPLIT_K, + PIPELINE_STAGES, + DYNAMIC_M, + K_TILE, + ) + elif ROWS >= 2 and M == 2: + _int6_w6a8_rows( + qwords, + x_scale, + ql, + qh, + scale, + steps, + out, + M, + N, + K, + stride_qln, + stride_qlk, + stride_qhn, + stride_qhk, + stride_sn, + stride_sk, + stride_stn, + stride_stk, + stride_os, + stride_om, + stride_on, + 2, + BLOCK_N, + SPLIT_K, + PIPELINE_STAGES, + DYNAMIC_M, + K_TILE, + ) + elif ROWS >= 3 and M == 3: + _int6_w6a8_rows( + qwords, + x_scale, + ql, + qh, + scale, + steps, + out, + M, + N, + K, + stride_qln, + stride_qlk, + stride_qhn, + stride_qhk, + stride_sn, + stride_sk, + stride_stn, + stride_stk, + stride_os, + stride_om, + stride_on, + 3, + BLOCK_N, + SPLIT_K, + PIPELINE_STAGES, + DYNAMIC_M, + K_TILE, + ) + else: + _int6_w6a8_rows( + qwords, + x_scale, + ql, + qh, + scale, + steps, + out, + M, + N, + K, + stride_qln, + stride_qlk, + stride_qhn, + stride_qhk, + stride_sn, + stride_sk, + stride_stn, + stride_stk, + stride_os, + stride_om, + stride_on, + ROWS, + BLOCK_N, + SPLIT_K, + PIPELINE_STAGES, + DYNAMIC_M, + K_TILE, + ) + + +@triton.jit +def _int6_w6a8_bucket_kernel( + qwords, + x_scale, + ql, + qh, + scale, + steps, + out, + M, + N: tl.constexpr, + K: tl.constexpr, + stride_qln: tl.constexpr, + stride_qlk: tl.constexpr, + stride_qhn: tl.constexpr, + stride_qhk: tl.constexpr, + stride_sn: tl.constexpr, + stride_sk: tl.constexpr, + stride_stn: tl.constexpr, + stride_stk: tl.constexpr, + stride_os: tl.constexpr, + stride_om: tl.constexpr, + stride_on: tl.constexpr, + BUCKET: tl.constexpr, + BLOCK_N: tl.constexpr, + SPLIT_K: tl.constexpr, + PIPELINE_STAGES: tl.constexpr, + DYNAMIC_M: tl.constexpr, + K_TILE: tl.constexpr, +): + if DYNAMIC_M: + _int6_w6a8_exact_rows( + qwords, + x_scale, + ql, + qh, + scale, + steps, + out, + M, + N, + K, + stride_qln, + stride_qlk, + stride_qhn, + stride_qhk, + stride_sn, + stride_sk, + stride_stn, + stride_stk, + stride_os, + stride_om, + stride_on, + BUCKET, + BLOCK_N, + SPLIT_K, + PIPELINE_STAGES, + DYNAMIC_M, + K_TILE, + ) + else: + _int6_w6a8_rows( + qwords, + x_scale, + ql, + qh, + scale, + steps, + out, + M, + N, + K, + stride_qln, + stride_qlk, + stride_qhn, + stride_qhk, + stride_sn, + stride_sk, + stride_stn, + stride_stk, + stride_os, + stride_om, + stride_on, + BUCKET, + BLOCK_N, + SPLIT_K, + PIPELINE_STAGES, + DYNAMIC_M, + K_TILE, + ) + + +def _main_loop_trips(args) -> int: + """K_TILE units are consumed by 32 lanes per main-loop trip.""" + units_per_split = triton.cdiv( + int(args["K"]) // int(args["K_TILE"]), int(args["SPLIT_K"]) + ) + return triton.cdiv(units_per_split, 32) + + +def int6_autotune_configs() -> list[triton.Config]: + return autotune_configs([{"K_TILE": 32}, {"K_TILE": 16}]) + + +def _prune(configs, named_args, **kwargs): + kept = [] + for k_tile in (32, 16): + subset = [config for config in configs if config.kwargs["K_TILE"] == k_tile] + if subset: + kept.extend( + prune_by_main_loop_trips( + _main_loop_trips, + subset, + named_args, + K_TILE=k_tile, + **kwargs, + ) + ) + return kept + + +_BUCKET_KERNELS = { + bucket: triton.autotune( + configs=int6_autotune_configs(), + key=["N", "K", "SPLIT_K"], + prune_configs_by={"early_config_prune": _prune}, + )(_int6_w6a8_bucket_kernel) + for bucket in SUPPORTED_BUCKETS +} + + +def _unsupported_reason( + bucket: int, + x: torch.Tensor, + ql: torch.Tensor, + qh: torch.Tensor, + scale: torch.Tensor, + steps: torch.Tensor, + group_size: int, +) -> Optional[str]: + weights = (ql, qh, scale, steps) + if any(tensor.dim() != 2 for tensor in (x, *weights)): + return "expects rank-2 activation and weight tensors" + M, K = x.shape + N, K_half = ql.shape + if isinstance(K, int) and K <= 0: + return f"K must be positive, got {K}" + reason = first_reason( + check_dtypes( + ( + ("activation", x, (torch.bfloat16,)), + ("ql", ql, (torch.uint8, torch.int8)), + ("qh", qh, (torch.uint8, torch.int8)), + ("scale codes", scale, (torch.int8, torch.uint8)), + ("steps", steps, (torch.float16,)), + ) + ), + None + if group_size == _GROUP_SIZE + else f"group_size must be {_GROUP_SIZE}, got {group_size}", + check_k(K, _SUPER_BLOCK), + check_rows(M, bucket), + check_contiguous((x, *weights)), + check_device(x, weights), + ) + if reason is not None: + return reason + if K_half * 2 != K: + return f"ql K/2 mismatch: x K={K}, ql K/2={K_half}" + if tuple(qh.shape) != (N, K // 4): + return "qh shape does not match [N, K/4]" + if tuple(scale.shape) != (N, K // group_size): + return "scale shape does not match [N, K/group_size]" + if tuple(steps.shape) != (N, K // _SUPER_BLOCK): + return "steps shape does not match [N, K/256]" + return None + + +def _launch( + bucket: int, + x: torch.Tensor, + ql: torch.Tensor, + qh: torch.Tensor, + scale: torch.Tensor, + steps: torch.Tensor, + group_size: int, +) -> torch.Tensor: + M, K = x.shape + N = ql.shape[0] + split_k = split_k_for( + bucket, + int(N), + int(K), + _sm_count(_device_index(x.device)), + k_per_split_unit=_SUPER_BLOCK, + ) + qwords, x_scale, _ = quantize_activations_q8(x, store_sum=False) + block_m = triton.next_power_of_2(bucket) + return launch_split_k_gemm( + _BUCKET_KERNELS[bucket], + bucket=bucket, + m=M, + n=N, + device=x.device, + split_k=split_k, + block_m=block_m, + inputs=(qwords, x_scale, ql, qh, scale, steps), + shape_args=( + M, + N, + K, + ql.stride(0), + ql.stride(1), + qh.stride(0), + qh.stride(1), + scale.stride(0), + scale.stride(1), + steps.stride(0), + steps.stride(1), + ), + BUCKET=bucket, + DYNAMIC_M=not isinstance(M, int), + ) + + +def _prototype( + x: torch.Tensor, + ql: torch.Tensor, + qh: torch.Tensor, + scale: torch.Tensor, + steps: torch.Tensor, + group_size: int, +) -> torch.Tensor: + raise NotImplementedError + + +def _fake(bucket: int, x: torch.Tensor, ql: torch.Tensor, *args) -> torch.Tensor: + return torch.empty((x.shape[0], ql.shape[0]), dtype=torch.bfloat16, device=x.device) + + +INT6_QUANTIZED_GEMM = QuantizedGemmFamily( + "int6_quantized_gemm", + SUPPORTED_BUCKETS, + _prototype, + _launch, + _fake, + _unsupported_reason, +) + + +__all__ = [ + "INT6_QUANTIZED_GEMM", + "SUPPORTED_BUCKETS", + "int6_autotune_configs", +] diff --git a/examples/models/gemma4_31b/quant/tests/test_pack_cuda.py b/examples/models/gemma4_31b/quant/tests/test_pack_cuda.py index 25366717977..c288adb3248 100644 --- a/examples/models/gemma4_31b/quant/tests/test_pack_cuda.py +++ b/examples/models/gemma4_31b/quant/tests/test_pack_cuda.py @@ -250,14 +250,14 @@ def test_e2e_q6k_export_lower_decode(self): Builds a synthetic Q6_K ExportableGGUFTensor, packs it into a CudaDp4aPlanarInt6Tensor, exports a decode-shaped (M=1) nn.Linear, and asserts: - * the exported graph captured ``executorch_cuda.int6_plain_mm`` (the - decode custom op chosen for M<=4), + * the exported graph captured ``triton.int6_quantized_gemm_m1`` (the + decode op chosen for M=1), * lowering through the CUDA backend produces an ``executorch_call_delegate``, * running the exported graph matches the Q6_K dequant reference. - The lowered .pte is not executed here (that needs the built C-shim + The lowered .pte is not executed here (that needs the built CUDA runtime); the eager exported graph already exercises the int6 decode op - through its registered CUDA impl. + through its Triton kernel. """ _require_cuda(self) from executorch.backends.cuda.cuda_backend import CudaBackend @@ -282,11 +282,11 @@ def test_e2e_q6k_export_lower_decode(self): with torch.no_grad(): ep = export(module, (x,), strict=True) - # The decode (M<=4) path must capture the int6 decode custom op. + # The decode (M<=4) path must capture the int6 decode Triton op. targets = [str(n.target) for n in ep.graph.nodes if n.op == "call_function"] self.assertTrue( - any("int6_plain_mm" in t for t in targets), - f"int6_plain_mm not found in exported graph: {targets}", + any("int6_quantized_gemm_m1" in t for t in targets), + f"int6_quantized_gemm_m1 not found in exported graph: {targets}", ) # Run the exported graph and compare against the Q6_K dequant reference. diff --git a/examples/models/gemma4_31b/tests/test_cuda_packers.py b/examples/models/gemma4_31b/tests/test_cuda_packers.py index d2d7048f687..fb2ff85f659 100644 --- a/examples/models/gemma4_31b/tests/test_cuda_packers.py +++ b/examples/models/gemma4_31b/tests/test_cuda_packers.py @@ -253,14 +253,14 @@ def test_e2e_q6k_export_lower_decode(self): Builds a synthetic Q6_K ExportableGGUFTensor, packs it into a CudaDp4aPlanarInt6Tensor, exports a decode-shaped (M=1) nn.Linear, and asserts: - * the exported graph captured ``executorch_cuda.int6_plain_mm`` (the - decode custom op chosen for M<=4), + * the exported graph captured ``triton.int6_quantized_gemm_m1`` (the + decode op chosen for M=1), * lowering through the CUDA backend produces an ``executorch_call_delegate``, * running the exported graph matches the Q6_K dequant reference. - The lowered .pte is not executed here (that needs the built C-shim + The lowered .pte is not executed here (that needs the built CUDA runtime); the eager exported graph already exercises the int6 decode op - through its registered CUDA impl. + through its Triton kernel. """ _require_cuda(self) from executorch.backends.cuda.cuda_backend import CudaBackend @@ -285,11 +285,11 @@ def test_e2e_q6k_export_lower_decode(self): with torch.no_grad(): ep = export(module, (x,), strict=True) - # The decode (M<=4) path must capture the int6 decode custom op. + # The decode (M<=4) path must capture the int6 decode Triton op. targets = [str(n.target) for n in ep.graph.nodes if n.op == "call_function"] self.assertTrue( - any("int6_plain_mm" in t for t in targets), - f"int6_plain_mm not found in exported graph: {targets}", + any("int6_quantized_gemm_m1" in t for t in targets), + f"int6_quantized_gemm_m1 not found in exported graph: {targets}", ) # Run the exported graph and compare against the Q6_K dequant reference.