From 8c36d543637d4dd8957fdf890b2fd499285f3b9e Mon Sep 17 00:00:00 2001 From: Hari Seshadri Date: Thu, 4 Jun 2026 11:01:57 -0700 Subject: [PATCH 01/17] Stage --- cmake/onnxruntime_mlas.cmake | 6 + onnxruntime/core/mlas/lib/qnbitgemm.cpp | 143 ++- onnxruntime/core/mlas/lib/qnbitgemm.h | 33 + .../mlas/lib/sqnbitgemm_kernel_avx512.cpp | 3 + .../lib/sqnbitgemm_kernel_avx512_2bit.cpp | 259 ++++++ .../mlas/lib/sqnbitgemm_kernel_avx512_2bit.h | 197 +++++ .../mlas/lib/sqnbitgemm_kernel_avx512vnni.cpp | 10 + ...nbitgemm_kernel_avx512vnni_2bit_blklen64.h | 816 ++++++++++++++++++ onnxruntime/test/mlas/bench/bench_lutgemm.cpp | 119 ++- .../test/mlas/bench/bench_qnbitgemm.cpp | 54 ++ .../mlas/unittest/test_sqnbitgemm_2bit.cpp | 167 ++++ .../unittest/test_sqnbitgemm_2bit_gemm.cpp | 289 +++++++ 12 files changed, 2070 insertions(+), 26 deletions(-) create mode 100644 onnxruntime/core/mlas/lib/sqnbitgemm_kernel_avx512_2bit.cpp create mode 100644 onnxruntime/core/mlas/lib/sqnbitgemm_kernel_avx512_2bit.h create mode 100644 onnxruntime/core/mlas/lib/sqnbitgemm_kernel_avx512vnni_2bit_blklen64.h create mode 100644 onnxruntime/test/mlas/unittest/test_sqnbitgemm_2bit.cpp create mode 100644 onnxruntime/test/mlas/unittest/test_sqnbitgemm_2bit_gemm.cpp diff --git a/cmake/onnxruntime_mlas.cmake b/cmake/onnxruntime_mlas.cmake index b254b40f88e76..c7abd3c0c4345 100644 --- a/cmake/onnxruntime_mlas.cmake +++ b/cmake/onnxruntime_mlas.cmake @@ -241,6 +241,9 @@ function(setup_mlas_source_for_windows) ${MLAS_SRC_DIR}/sqnbitgemm_lut_kernel_avx2.h ${MLAS_SRC_DIR}/sqnbitgemm_lut_kernel_avx2.cpp ${MLAS_SRC_DIR}/sqnbitgemm_kernel_avx2.cpp + ${MLAS_SRC_DIR}/sqnbitgemm_kernel_avx512_2bit.h + ${MLAS_SRC_DIR}/sqnbitgemm_kernel_avx512_2bit.cpp + ${MLAS_SRC_DIR}/sqnbitgemm_kernel_avx512vnni_2bit_blklen64.h ${MLAS_SRC_DIR}/sqnbitgemm_kernel_avx512.cpp ${MLAS_SRC_DIR}/sqnbitgemm_kernel_avx512vnni.cpp ${MLAS_SRC_DIR}/qkv_quant_kernel_avx512vnni.cpp @@ -791,6 +794,9 @@ else() ${MLAS_SRC_DIR}/intrinsics/avx2/qdwconv_avx2.cpp ${MLAS_SRC_DIR}/intrinsics/avx2/saturation_check_avx2.cpp ${MLAS_SRC_DIR}/sqnbitgemm_kernel_avx2.cpp + ${MLAS_SRC_DIR}/sqnbitgemm_kernel_avx512_2bit.h + ${MLAS_SRC_DIR}/sqnbitgemm_kernel_avx512_2bit.cpp + ${MLAS_SRC_DIR}/sqnbitgemm_kernel_avx512vnni_2bit_blklen64.h ${MLAS_SRC_DIR}/sqnbitgemm_lut_kernel_avx2.h ${MLAS_SRC_DIR}/sqnbitgemm_lut_kernel_avx2.cpp ${MLAS_SRC_DIR}/rotary_embedding_kernel_avx2.h diff --git a/onnxruntime/core/mlas/lib/qnbitgemm.cpp b/onnxruntime/core/mlas/lib/qnbitgemm.cpp index f649d8ab38648..e9f9ce34d3977 100644 --- a/onnxruntime/core/mlas/lib/qnbitgemm.cpp +++ b/onnxruntime/core/mlas/lib/qnbitgemm.cpp @@ -33,6 +33,7 @@ enum QNBitGemmVariant { HQ4BitGemmVariant_CompFp16, SQ8BitGemmVariant_CompInt8, HQ8BitGemmVariant_CompFp16, + SQ2BitGemmVariant_CompInt8, // End of valid variants @@ -63,6 +64,10 @@ GetQNBitGemmVariant( } else if (ComputeType == HQNBIT_CompFp16) { return HQ8BitGemmVariant_CompFp16; } + } else if (BlkBitWidth == 2) { + if (ComputeType == SQNBIT_CompInt8) { + return SQ2BitGemmVariant_CompInt8; + } } } @@ -117,6 +122,12 @@ MlasIsQNBitGemmAvailable( Dispatch->HQ8BitBlkDequantBForHgemm_CompFp16 != nullptr && Dispatch->HQ4BitGemmKernel_CompFp16 != nullptr; } + case SQ2BitGemmVariant_CompInt8: { + return Dispatch->Q2BitGemmPackQuantBDataSize != nullptr && + Dispatch->SQ2BitGemmPackQuantBDataAndBlkSum != nullptr && + Dispatch->SQ2BitGemmKernel_BlkSum_CompInt8 != nullptr && + Dispatch->QuantizeARowComputeBlkSum_CompInt8 != nullptr; + } default: { return false; } @@ -143,7 +154,7 @@ QNBitGemmPerGemmWorkspaceSize( return 0; } - if (BlkBitWidth == 4 || BlkBitWidth == 8) { + if (BlkBitWidth == 4 || BlkBitWidth == 8 || BlkBitWidth == 2) { return Dispatch->QNBitGemmPerGemmWorkspaceSize(M, N, K, BlkLen, HasZeroPoint, ComputeType, BlkBitWidth, BackendKernelSelectorConfig); } @@ -162,7 +173,7 @@ QNBitGemmPerGemmWorkspaceAlignment( return 1; } - if (BlkBitWidth == 4 || BlkBitWidth == 8) { + if (BlkBitWidth == 4 || BlkBitWidth == 8 || BlkBitWidth == 2) { return Dispatch->QNBitGemmPerGemmWorkspaceAlignment(BlkLen, ComputeType); } @@ -240,6 +251,11 @@ MlasQNBitGemmPackQuantBDataSize( N, K, BlkLen, HasZeroPoint, ComputeType, BackendKernelSelectorConfig ); + } else if (BlkBitWidth == 2 && Dispatch->Q2BitGemmPackQuantBDataSize != nullptr) { + return Dispatch->Q2BitGemmPackQuantBDataSize( + N, K, BlkLen, HasZeroPoint, ComputeType, + BackendKernelSelectorConfig + ); } return 0; @@ -355,6 +371,24 @@ MlasQNBitGemmPackQuantBData( BackendKernelSelectorConfig ); } + } else if (BlkBitWidth == 2) { + if (ComputeType == SQNBIT_CompInt8 && Dispatch->SQ2BitGemmPackQuantBDataAndBlkSum != nullptr) { + const size_t BlockCountK = MlasDivRoundup(K, BlkLen); + PackedQuantBDataStruct packed_quant_b(PackedQuantBDataAndOrBlkSumWorkspace, N, BlockCountK, BlkLen, false); + Dispatch->SQ2BitGemmPackQuantBDataAndBlkSum( + N, + K, + BlkLen, + ComputeType, + static_cast(QuantBData), + static_cast(QuantBScale), + HasZeroPoint, + static_cast(QuantBZeroPoint), + packed_quant_b, + ThreadPool, + BackendKernelSelectorConfig + ); + } } } @@ -937,6 +971,88 @@ SQ8BitGemm_CompInt8( } } +// +// 2-bit weight CompInt8 wrapper. Mirrors SQ4BitGemm_CompInt8 but specialised +// for BlkBitWidth=2 (kBlkBytes = BlkLen/4) and uses the simple column-major +// QuantBBlkSum layout produced by SQ2BitGemmPackQuantBDataAndBlkSum_*. +// +void +SQ2BitGemm_CompInt8( + const size_t BlkLen, + const size_t K, + const MLAS_QNBIT_GEMM_DATA_PARAMS* const DataParams, + void* const PerGemmWorkspace, + const size_t RangeStartM, + const size_t RangeCountM, + const size_t RangeStartN, + const size_t RangeCountN, + const MLAS_BACKEND_KERNEL_SELECTOR_CONFIG* /*BackendKernelSelectorConfig*/ +) +{ + constexpr size_t BlkBitWidth = 2; + + const auto* Dispatch = GetMlasPlatform().QNBitGemmDispatch; + if (Dispatch == nullptr || Dispatch->SQ2BitGemmKernel_BlkSum_CompInt8 == nullptr) { + return; + } + + PerGemmQuantAWorkspace* const per_gemm_quant_a_workspace = + static_cast(PerGemmWorkspace); + + const size_t k_blks = MlasDivRoundup(K, BlkLen); + + const size_t lda = k_blks * (per_gemm_quant_a_workspace->QuantScale ? BlkLen : Q8BlkSize(BlkLen)); + const size_t ldc = DataParams->ldc; + const size_t ldb = k_blks * MlasQNBitBlkDataSizeInBytes(BlkBitWidth, BlkLen); // BlkLen / 4 bytes per block + + const std::byte* QuantA = per_gemm_quant_a_workspace->QuantData + RangeStartM * lda; + const float* QuantAScale = per_gemm_quant_a_workspace->QuantScale + RangeStartM * k_blks; + + const std::byte* QuantBData = static_cast(DataParams->PackedQuantBData) + RangeStartN * ldb; + const float* QuantBScale = DataParams->QuantBScale + RangeStartN * k_blks; + const float* ABlockSum = per_gemm_quant_a_workspace->BlockSum + RangeStartM * k_blks; + const float* QuantBBlkSum = DataParams->QuantBBlkSum + RangeStartN * k_blks; + float* C = DataParams->C + RangeStartM * ldc + RangeStartN; + + const float* Bias = (DataParams->Bias == nullptr) ? nullptr : DataParams->Bias + RangeStartN; + + size_t CountN; + for (size_t n = 0; n < RangeCountN; n += CountN) { + CountN = std::min(RangeCountN - n, size_t{128}); + + const std::byte* b_col = QuantBData + n * ldb; + const float* b_col_scale = QuantBScale + n * k_blks; + float* c_blk = C + n; + const float* bias = (Bias == nullptr) ? nullptr : Bias + n; + const float* b_blk_sum = QuantBBlkSum + n * k_blks; + + Dispatch->SQ2BitGemmKernel_BlkSum_CompInt8( + BlkLen, + QuantA, + QuantAScale, + b_col, + b_col_scale, + /*QuantBZeroPoint*/ nullptr, + c_blk, + RangeCountM, + CountN, + K, + k_blks, + bias, + ldc, + ABlockSum, + b_blk_sum + ); + + if (DataParams->PostProcessor != nullptr) { + DataParams->PostProcessor->Process( + DataParams->C, RangeStartM, RangeStartN + n, + RangeCountM, CountN, ldc + ); + } + } +} + template void InitializeWorkspace_CompInt8( @@ -1007,7 +1123,7 @@ InitializeWorkspace_CompInt8( }); } else { // TODO(hasesh): Clean-up the following logic so that it is clean AND it works as expected on all platforms - if (BlkBitWidth == 4) { + if (BlkBitWidth == 4 || BlkBitWidth == 2) { if (QuantizeARow) { MlasTrySimpleParallel(ThreadPool, BatchN, [&](ptrdiff_t gemm_idx) { const auto& data = DataParams[gemm_idx]; @@ -1090,6 +1206,7 @@ GetInitializeWorkspace(QNBitGemmVariant variant) switch (variant) { case SQ4BitGemmVariant_CompInt8: case SQ8BitGemmVariant_CompInt8: + case SQ2BitGemmVariant_CompInt8: return InitializeWorkspace_CompInt8; default: return nullptr; @@ -1132,6 +1249,11 @@ GetQNBitGemm(QNBitGemmVariant variant) return SQ4BitGemm_CompInt8; case SQ8BitGemmVariant_CompInt8: return SQ8BitGemm_CompInt8; + case SQ2BitGemmVariant_CompInt8: + // Phase 2b: scalar reference compute kernel registered by the + // AVX-512 / AVX-512-VNNI dispatch tables. Phase 3 swaps in the + // vectorized version at the dispatch slot only. + return SQ2BitGemm_CompInt8; default: return nullptr; } @@ -1255,6 +1377,13 @@ MlasQNBitGemmBatch( const_cast*>(Data)->QuantBScale = packed_quant_b.PackedQuantBScale; const_cast*>(Data)->BlkUnsignedQuantAZeroPointCorrection = packed_quant_b.BlkUnsignedQuantAZeroPointCorrection; + PerGemmQuantAWorkspace per_gemm_quant_a_workspace(PerGemmWorkspace, M, BlockCountK, BlkLen); + ComputeOperation(BlkLen, K, Data, &per_gemm_quant_a_workspace, 0, M, 0, N, BackendKernelSelectorConfig); + } else if (Variant == SQ2BitGemmVariant_CompInt8 && GetMlasPlatform().QNBitGemmDispatch->SQ2BitGemmKernel_BlkSum_CompInt8 != nullptr) { + PackedQuantBDataStruct packed_quant_b(const_cast(Data->QuantBDataWorkspace), N, BlockCountK, BlkLen, false); + const_cast*>(Data)->PackedQuantBData = packed_quant_b.PackedQuantBData; + const_cast*>(Data)->QuantBBlkSum = packed_quant_b.QuantBBlkSum; + const_cast*>(Data)->QuantBScale = packed_quant_b.PackedQuantBScale; PerGemmQuantAWorkspace per_gemm_quant_a_workspace(PerGemmWorkspace, M, BlockCountK, BlkLen); ComputeOperation(BlkLen, K, Data, &per_gemm_quant_a_workspace, 0, M, 0, N, BackendKernelSelectorConfig); } else { @@ -1336,6 +1465,14 @@ MlasQNBitGemmBatch( const_cast*>(Data)->QuantBScale = packed_quant_b.PackedQuantBScale; const_cast*>(Data)->BlkUnsignedQuantAZeroPointCorrection = packed_quant_b.BlkUnsignedQuantAZeroPointCorrection; + PerGemmQuantAWorkspace per_gemm_quant_a_workspace(PerGemmWorkspace, M, BlockCountK, BlkLen); + ComputeOperation(BlkLen, K, Data, &per_gemm_quant_a_workspace, RangeStartM, RangeCountM, RangeStartN, RangeCountN, BackendKernelSelectorConfig); + } else if (Variant == SQ2BitGemmVariant_CompInt8 && GetMlasPlatform().QNBitGemmDispatch->SQ2BitGemmKernel_BlkSum_CompInt8 != nullptr) { + PackedQuantBDataStruct packed_quant_b(const_cast(Data->QuantBDataWorkspace), N, BlockCountK, BlkLen, false); + const_cast*>(Data)->PackedQuantBData = packed_quant_b.PackedQuantBData; + const_cast*>(Data)->QuantBBlkSum = packed_quant_b.QuantBBlkSum; + const_cast*>(Data)->QuantBScale = packed_quant_b.PackedQuantBScale; + PerGemmQuantAWorkspace per_gemm_quant_a_workspace(PerGemmWorkspace, M, BlockCountK, BlkLen); ComputeOperation(BlkLen, K, Data, &per_gemm_quant_a_workspace, RangeStartM, RangeCountM, RangeStartN, RangeCountN, BackendKernelSelectorConfig); } else { diff --git a/onnxruntime/core/mlas/lib/qnbitgemm.h b/onnxruntime/core/mlas/lib/qnbitgemm.h index 6503f0108c823..421fe13009d26 100644 --- a/onnxruntime/core/mlas/lib/qnbitgemm.h +++ b/onnxruntime/core/mlas/lib/qnbitgemm.h @@ -457,6 +457,39 @@ struct MLAS_QNBIT_GEMM_DISPATCH { SQ8BitGemmKernel_BlkSum_CompInt8_Fn* SQ8BitGemmKernel_BlkSum_CompInt8 = nullptr; + // + // SQ2BIT_CompInt8 dispatch surface (mirrors the SQ4 set for 2-bit weights). + // + // These pointers are populated only on platforms that ship a native 2-bit + // VNNI kernel. When all three are nullptr, MlasIsQNBitGemmAvailable returns + // false for (BlkBitWidth=2, ComputeType=SQNBIT_CompInt8) and the LUT path + // continues to handle 2-bit weights. + // + + /** Gets size of packed quantized B data containing 2-bit integers. See MlasQNBitGemmPackQuantBDataSize(). */ + Q4BitGemmPackQuantBDataSize_Fn* Q2BitGemmPackQuantBDataSize = nullptr; + + /** Packs quantized B data + per-block sums for the 2-bit CompInt8 kernel. */ + typedef void(SQ2BitGemmPackQuantBDataAndSumBlk_Fn)( + size_t N, + size_t K, + size_t BlkLen, + MLAS_QNBIT_GEMM_COMPUTE_TYPE ComputeType, + const std::byte* QuantBDataBegin, + const float* QuantBScaleBegin, + bool HasZeroPoint, + const std::byte* QuantBZPBegin, + PackedQuantBDataStruct& PackedQuantB, + MLAS_THREADPOOL* ThreadPool, + const MLAS_BACKEND_KERNEL_SELECTOR_CONFIG* BackendKernelSelectorConfig + ); + + SQ2BitGemmPackQuantBDataAndSumBlk_Fn* SQ2BitGemmPackQuantBDataAndBlkSum = nullptr; + + /** Inner kernel for the 2-bit CompInt8 path. Same shape as the 4-bit version; + the packed B layout encodes 2-bit weights instead of 4-bit. */ + SQ4BitGemmKernel_BlkSum_CompInt8_Fn* SQ2BitGemmKernel_BlkSum_CompInt8 = nullptr; + /** * @brief Multiply quantized 8-bit integer matrix A with quantized 4-bit integer matrix B. * A and B are block quantized and B is column major. diff --git a/onnxruntime/core/mlas/lib/sqnbitgemm_kernel_avx512.cpp b/onnxruntime/core/mlas/lib/sqnbitgemm_kernel_avx512.cpp index b9757836b994b..add4784b55d01 100644 --- a/onnxruntime/core/mlas/lib/sqnbitgemm_kernel_avx512.cpp +++ b/onnxruntime/core/mlas/lib/sqnbitgemm_kernel_avx512.cpp @@ -494,5 +494,8 @@ const MLAS_QNBIT_GEMM_DISPATCH MlasSQNBitGemmDispatchAvx512 = []() { d.SQ8BitGemmKernel_BlkSum_CompInt8 = SQ8BitGemmKernel_BlkSum_CompInt8_avx512; d.QuantizeARowComputeBlkSum_CompInt8 = QuantizeARow_CompInt8_avx512; + // 2-bit native CompInt8 path is registered in the AVX-512-VNNI dispatch only. + // Hosts with AVX-512 but no VNNI fall through to the existing LUT path. + return d; }(); diff --git a/onnxruntime/core/mlas/lib/sqnbitgemm_kernel_avx512_2bit.cpp b/onnxruntime/core/mlas/lib/sqnbitgemm_kernel_avx512_2bit.cpp new file mode 100644 index 0000000000000..e6f2778dffe7c --- /dev/null +++ b/onnxruntime/core/mlas/lib/sqnbitgemm_kernel_avx512_2bit.cpp @@ -0,0 +1,259 @@ +/*++ + +Copyright (c) Microsoft Corporation. All rights reserved. + +Licensed under the MIT License. + +Module Name: + + sqnbitgemm_kernel_avx512_2bit.cpp + +Abstract: + + Phase 2b reference implementation of the 2-bit weight CompInt8 GEMM + (BlkBitWidth=2, BlkLen=64). This file contains only scalar C++ code; the + AVX-512-VNNI vectorized inner loop lands in Phase 3 as a separate header + (sqnbitgemm_kernel_avx512vnni_2bit_blklen64.h) that replaces the kernel + slot in the AVX-512-VNNI dispatch table. + + The scalar functions exposed here are linked into the AVX-512-VNNI + dispatch table only (the plain AVX-512 dispatch leaves the W2 slots + null so non-VNNI hosts fall through to the existing LUT kernel). They + are also reachable as plain C++ symbols from unit tests, which use them + as a correctness oracle for the vectorized path. + + Restrictions (Phase 2): + * BlkLen == 64 only. Other BlkLens are rejected by the pack helper + and the kernel returns 0 rows handled. + * Symmetric quantization only (no per-block zero-point tensor; + an implicit zero-point of 2 is used to recentre values in [0, 3]). + +--*/ + +#include "sqnbitgemm_kernel_avx512_2bit.h" + +#include +#include +#include +#include + +#include "mlasi.h" +#include "qnbitgemm.h" + +namespace onnxruntime { +namespace mlas { +namespace sq2bit_avx512 { + +// +// Workspace / pack-buffer size for the 2-bit CompInt8 path. +// +// Layout (in bytes): +// +// [PackedQuantBData] N * BlockCountK * kBlkBytes (BlkLen / 4 bytes per block) +// [BlkSum (float)] roundup_16(N) * BlockCountK +// [Scales (float)] N * BlockCountK +// +// Alignment slack is added so that the AVX-512 dequant in Phase 3 can use +// 64-byte aligned loads. +// +size_t MLASCALL +Q2BitGemmPackQuantBDataSize_Avx512( + size_t N, + size_t K, + size_t BlkLen, + bool /* HasZeroPoint */, + MLAS_QNBIT_GEMM_COMPUTE_TYPE ComputeType, + const MLAS_BACKEND_KERNEL_SELECTOR_CONFIG* /* BackendKernelSelectorConfig */ +) +{ + // Phase 2 supports only BlkLen=64 and SQNBIT_CompInt8. Anything else returns + // 0 so MlasQNBitGemmPackQuantBDataSize reports an unsupported configuration. + if (BlkLen != kBlkLen || ComputeType != SQNBIT_CompInt8) { + return 0; + } + + const size_t BlockCountK = MlasDivRoundup(K, BlkLen); + size_t PackedQuantBDataSize = N * BlockCountK * kBlkBytes; + const size_t ScaleSize = N * BlockCountK * sizeof(float); + size_t BlkSumSize = MlasDivRoundup(N, 16) * BlockCountK * 16 * sizeof(float); + + constexpr size_t kPackedQuantBDataAlignment = 64; // AVX-512 friendly + PackedQuantBDataSize += kPackedQuantBDataAlignment - 1; + + constexpr size_t kBlkSumAlignment = MlasQNBitQuantBBlkSumAlignment(); + BlkSumSize += kBlkSumAlignment - 1; + + return PackedQuantBDataSize + ScaleSize + BlkSumSize; +} + +// +// Pack quantized B data + scales + per-block sums for the 2-bit kernel. +// +// Layouts produced (all column-major in N): +// PackedQuantBData[n * BlockCountK * kBlkBytes + blk * kBlkBytes + i] +// Block (n, blk), byte i of the kPackedBlkBytes layout (see header). +// PackedQuantBScale[n * BlockCountK + blk] +// Copy of the input scale; column-major. +// QuantBBlkSum[n * BlockCountK + blk] +// = -scale * 2 (symmetric W2 uses an implicit zero point of 2). +// +// QuantBZPBegin / HasZeroPoint are accepted for ABI parity with the W4 path +// but ignored in Phase 2 because the customer model and the projection +// rely on the symmetric layout. +// +void MLASCALL +SQ2BitGemmPackQuantBDataAndBlkSum_Scalar( + size_t N, + size_t K, + size_t BlkLen, + MLAS_QNBIT_GEMM_COMPUTE_TYPE /* ComputeType */, + const std::byte* QuantBDataBegin, + const float* QuantBScaleBegin, + bool /* HasZeroPoint */, + const std::byte* /* QuantBZPBegin */, + PackedQuantBDataStruct& PackedQuantB, + MLAS_THREADPOOL* ThreadPool, + const MLAS_BACKEND_KERNEL_SELECTOR_CONFIG* /* BackendKernelSelectorConfig */ +) +{ + assert(BlkLen == kBlkLen); + if (BlkLen != kBlkLen) { + return; + } + + const size_t BlockCountK = MlasDivRoundup(K, BlkLen); + const size_t Iterations = N * BlockCountK; + + // Pack weight bytes (block-by-block, parallel over N * BlockCountK). + if (QuantBDataBegin != nullptr) { + std::byte* PackedQuantBData = PackedQuantB.PackedQuantBData; + MlasTrySimpleParallel( + ThreadPool, static_cast(Iterations), + [&](ptrdiff_t tid) { + const size_t n = static_cast(tid) / BlockCountK; + const size_t blk = static_cast(tid) % BlockCountK; + const size_t offset = (n * BlockCountK + blk) * kBlkBytes; + PackBlock_BlkLen64(QuantBDataBegin + offset, PackedQuantBData + offset); + } + ); + } + + // Copy scales as-is (column-major) and compute BlkSum. + // + // BlkSum uses the W4-style "width-16 row-major chunked" layout because the + // top-level kernel performs the zero-point correction via the float SGEMM + // micro-kernel (`GetMlasPlatform().GemmFloatKernel`), which expects this + // pre-packed B layout: + // + // BlkSum[(n / 16) * BlockCountK * 16 + blk * 16 + (n % 16)] + // = -scale_b * 2 (symmetric W2 uses an implicit ZP of 2) + // + // The allocated BlkSum buffer is sized at MlasDivRoundup(N, 16) * BlockCountK + // * 16 floats so the layout is well-defined even when N % 16 != 0 (the tail + // chunk's unused lanes hold whatever the buffer was initialised with, which + // for production callers must be zero so the SGEMM correction reads zeros). + if (QuantBScaleBegin != nullptr) { + float* PackedScales = PackedQuantB.PackedQuantBScale; + float* BlkSum = PackedQuantB.QuantBBlkSum; + MlasTrySimpleParallel( + ThreadPool, static_cast(Iterations), + [&](ptrdiff_t tid) { + const size_t n = static_cast(tid) / BlockCountK; + const size_t blk = static_cast(tid) % BlockCountK; + const float scale = QuantBScaleBegin[n * BlockCountK + blk]; + PackedScales[n * BlockCountK + blk] = scale; + const size_t blksum_offset = ((n / 16) * BlockCountK + blk) * 16 + (n % 16); + BlkSum[blksum_offset] = -scale * static_cast(kDefaultSymmetricZeroPoint2Bit); + } + ); + } +} + +// +// Scalar reference kernel for SQ2BitGemmVariant_CompInt8. +// +// Inputs match the SQ4BitGemmKernel_BlkSum_CompInt8_Fn typedef so the +// vectorized Phase 3 implementation can drop into the same dispatch slot. +// +// Math: +// C[m, n] = bias[n] +// + sum_blk( scale_a[m, blk] * scale_b[n, blk] +// * dot(int8 a[m, blk, :], uint8 b_unpacked[n, blk, :]) ) +// + sum_blk( ABlockSum[m, blk] * QuantBBlkSum[n, blk] ) +// +// The third term applies the symmetric W2 zero-point correction: +// QuantBBlkSum[n, blk] = -scale_b * 2, and ABlockSum[m, blk] = scale_a * sum(a). +// +size_t MLASCALL +SQ2BitGemmKernel_BlkSum_CompInt8_Scalar( + const size_t BlkLen, + const std::byte* QuantA, + const float* QuantAScale, + const std::byte* QuantBData, + const float* QuantBScale, + const std::byte* /* QuantBZeroPoint */, + float* C, + size_t CountM, + size_t CountN, + size_t /* CountK */, + size_t BlockCountK, + const float* Bias, + size_t ldc, + const float* ABlockSum, + const float* QuantBBlkSum +) +{ + if (BlkLen != kBlkLen) { + return 0; // Phase 2b only supports BlkLen=64. + } + + const size_t lda = BlockCountK * kBlkLen; // bytes per A row (int8) + const size_t lda_scale = BlockCountK; // floats per A scale row + const size_t ldb = BlockCountK * kBlkBytes; // bytes per B column + const size_t ldb_scale = BlockCountK; // floats per B column + + for (size_t m = 0; m < CountM; ++m) { + const int8_t* a_row = reinterpret_cast(QuantA + m * lda); + const float* a_scale_row = QuantAScale + m * lda_scale; + const float* a_blksum_row = ABlockSum + m * lda_scale; + float* c_row = C + m * ldc; + + for (size_t n = 0; n < CountN; ++n) { + const std::byte* b_col = QuantBData + n * ldb; + const float* b_scale_col = QuantBScale + n * ldb_scale; + const float* b_blksum_col = QuantBBlkSum + n * ldb_scale; + + float acc = (Bias != nullptr) ? Bias[n] : 0.0f; + + for (size_t blk = 0; blk < BlockCountK; ++blk) { + // Unpack 64 2-bit weights into 64 uint8 values (values in [0, 3]). + uint8_t b_unpacked[kBlkLen]; + UnpackBlock_BlkLen64_Reference(b_col + blk * kBlkBytes, b_unpacked); + + // int8 * uint8 dot product across the block. + const int8_t* a_blk = a_row + blk * kBlkLen; + int32_t dot = 0; + for (size_t i = 0; i < kBlkLen; ++i) { + dot += static_cast(a_blk[i]) * static_cast(b_unpacked[i]); + } + + // Integer term * scales. + acc += a_scale_row[blk] * b_scale_col[blk] * static_cast(dot); + + // Symmetric W2 zero-point correction: + // dot(a, b_signed) = dot(a, b_unsigned) - zp * sum(a) + // so we need C += scale_a * scale_b * (-zp) * sum(a) + // = ABlockSum * QuantBBlkSum (where QuantBBlkSum already encodes -scale_b * zp). + acc += a_blksum_row[blk] * b_blksum_col[blk]; + } + + c_row[n] = acc; + } + } + + return CountM; +} + +} // namespace sq2bit_avx512 +} // namespace mlas +} // namespace onnxruntime diff --git a/onnxruntime/core/mlas/lib/sqnbitgemm_kernel_avx512_2bit.h b/onnxruntime/core/mlas/lib/sqnbitgemm_kernel_avx512_2bit.h new file mode 100644 index 0000000000000..91c47fc13589a --- /dev/null +++ b/onnxruntime/core/mlas/lib/sqnbitgemm_kernel_avx512_2bit.h @@ -0,0 +1,197 @@ +/*++ + +Copyright (c) Microsoft Corporation. All rights reserved. + +Licensed under the MIT License. + +Module Name: + + sqnbitgemm_kernel_avx512_2bit.h + +Abstract: + + Pack-time helpers and reference (scalar) routines for the 2-bit, BlkLen=64 + AVX-512-VNNI weight GEMM path (SQNBIT_CompInt8, BlkBitWidth=2). + + This header is currently scalar / header-only and contains the pieces that + the Phase 2a round-trip unit tests exercise: + + * The packed-block layout used by the (future) AVX-512 dequant inner loop. + * A pack routine that converts standard ONNX MatMulNBits 2-bit input data + into the packed layout. + * A reference unpack routine (independent of the pack code) that + materialises one packed block back into 64 individual int8 values. + + The packed layout is designed so that the AVX-512BW dequant in Phase 3 is + one 128-bit broadcast plus a per-lane variable shift: + + __m128i p = _mm_loadu_si128(packed); // 16 bytes + __m512i p4 = _mm512_broadcast_i32x4(p); // 4 lanes + __m512i sh = _mm512_set_epi32(6,6,6,6, 4,4,4,4, + 2,2,2,2, 0,0,0,0); + __m512i v = _mm512_srlv_epi32(p4, sh); // per-lane shift + v = _mm512_and_si512(v, _mm512_set1_epi8(0x03)); + + With this layout the resulting ZMM holds weights 0..63 in their natural + order. + + NOTE: This file intentionally restricts itself to BlkLen == 64. Other + BlkLens fall through to the existing LUT path until later phases. + +--*/ + +#pragma once + +#include +#include +#include + +#include "mlas.h" +#include "mlas_qnbit.h" + +template +struct PackedQuantBDataStruct; // fwd decl, defined in qnbitgemm.h + +struct MLAS_BACKEND_KERNEL_SELECTOR_CONFIG; + +namespace onnxruntime { +namespace mlas { +namespace sq2bit_avx512 { + +// Each 2-bit weight occupies 2 bits; one byte holds 4 weights. +constexpr size_t kWeightsPerByte = 4; + +// Block constants for the BlkLen=64 variant. +constexpr size_t kBlkLen = 64; +constexpr size_t kBlkBytes = kBlkLen / kWeightsPerByte; // 16 packed src bytes per block +constexpr size_t kPackedBlkBytes = kBlkBytes; // packing is in-place: 16 -> 16 + +// Default zero point used when the input is symmetric (no zero-point tensor). +// For 2-bit unsigned values in [0, 3], the symmetric mid-point is 2. +constexpr uint8_t kDefaultSymmetricZeroPoint2Bit = 2; + +// +// Extract a single 2-bit weight from a standard ONNX MatMulNBits packed byte +// stream. `src` is the start of one block (kBlkBytes bytes). `i` is the +// in-block weight index in [0, kBlkLen). +// +inline uint8_t +ExtractSrcWeight(const std::byte* src, size_t i) +{ + const size_t byte_idx = i / kWeightsPerByte; + const size_t bit_off = (i % kWeightsPerByte) * 2; + return static_cast( + (static_cast(src[byte_idx]) >> bit_off) & 0x03u + ); +} + +// +// Pack one source block (16 bytes = 64 2-bit weights, standard ONNX layout) +// into the destination layout described at the top of this file. +// +// src_byte[i] holds val[4i .. 4i+3] at bit positions {0..1, 2..3, 4..5, 6..7}. +// dst_byte[i] holds val[i], val[i+16], val[i+32], val[i+48] at the same bit +// positions. +// +// This is a pure permutation of the 64 2-bit elements; the bit width and +// count are preserved (16 bytes in, 16 bytes out). +// +inline void +PackBlock_BlkLen64(const std::byte* src, std::byte* dst) +{ + for (size_t i = 0; i < kBlkBytes; ++i) { + const uint8_t v0 = ExtractSrcWeight(src, i + 0); + const uint8_t v1 = ExtractSrcWeight(src, i + 16); + const uint8_t v2 = ExtractSrcWeight(src, i + 32); + const uint8_t v3 = ExtractSrcWeight(src, i + 48); + dst[i] = static_cast( + static_cast(v0 | (v1 << 2) | (v2 << 4) | (v3 << 6)) + ); + } +} + +// +// Reference unpack of one packed block back into 64 int8 values in natural +// order ([val0, val1, ..., val63]). +// +// This routine is intentionally written without reference to PackBlock_BlkLen64 +// so that it can serve as an independent oracle for round-trip tests: +// it simply applies the documented dst_byte layout rule in reverse. +// +inline void +UnpackBlock_BlkLen64_Reference(const std::byte* packed, uint8_t out[kBlkLen]) +{ + for (size_t i = 0; i < kBlkBytes; ++i) { + const uint8_t b = static_cast(packed[i]); + out[i + 0] = static_cast((b >> 0) & 0x03u); + out[i + 16] = static_cast((b >> 2) & 0x03u); + out[i + 32] = static_cast((b >> 4) & 0x03u); + out[i + 48] = static_cast((b >> 6) & 0x03u); + } +} + +// +// Extract all 64 weights from a standard ONNX MatMulNBits source block into +// natural order. Used by tests as the "expected" sequence for round-tripping. +// +inline void +UnpackSourceBlock_BlkLen64_Reference(const std::byte* src, uint8_t out[kBlkLen]) +{ + for (size_t i = 0; i < kBlkLen; ++i) { + out[i] = ExtractSrcWeight(src, i); + } +} + +// +// Phase 2b reference / dispatch entry points. +// Defined in sqnbitgemm_kernel_avx512_2bit.cpp; registered into the +// AVX-512 / AVX-512-VNNI dispatch tables. +// + +size_t MLASCALL +Q2BitGemmPackQuantBDataSize_Avx512( + size_t N, + size_t K, + size_t BlkLen, + bool HasZeroPoint, + MLAS_QNBIT_GEMM_COMPUTE_TYPE ComputeType, + const MLAS_BACKEND_KERNEL_SELECTOR_CONFIG* BackendKernelSelectorConfig +); + +void MLASCALL +SQ2BitGemmPackQuantBDataAndBlkSum_Scalar( + size_t N, + size_t K, + size_t BlkLen, + MLAS_QNBIT_GEMM_COMPUTE_TYPE ComputeType, + const std::byte* QuantBDataBegin, + const float* QuantBScaleBegin, + bool HasZeroPoint, + const std::byte* QuantBZPBegin, + PackedQuantBDataStruct& PackedQuantB, + MLAS_THREADPOOL* ThreadPool, + const MLAS_BACKEND_KERNEL_SELECTOR_CONFIG* BackendKernelSelectorConfig +); + +size_t MLASCALL +SQ2BitGemmKernel_BlkSum_CompInt8_Scalar( + size_t BlkLen, + const std::byte* QuantA, + const float* QuantAScale, + const std::byte* QuantBData, + const float* QuantBScale, + const std::byte* QuantBZeroPoint, + float* C, + size_t CountM, + size_t CountN, + size_t CountK, + size_t BlockCountK, + const float* Bias, + size_t ldc, + const float* ABlockSum, + const float* QuantBBlkSum +); + +} // namespace sq2bit_avx512 +} // namespace mlas +} // namespace onnxruntime diff --git a/onnxruntime/core/mlas/lib/sqnbitgemm_kernel_avx512vnni.cpp b/onnxruntime/core/mlas/lib/sqnbitgemm_kernel_avx512vnni.cpp index 17b3b7bf0bfb3..c94638a8fd5c9 100644 --- a/onnxruntime/core/mlas/lib/sqnbitgemm_kernel_avx512vnni.cpp +++ b/onnxruntime/core/mlas/lib/sqnbitgemm_kernel_avx512vnni.cpp @@ -27,6 +27,8 @@ Module Name: #include "sqnbitgemm_kernel_avx512_int8_blklen32.h" #include "sqnbitgemm_kernel_avx512_int8_blklen64.h" #include "sqnbitgemm_kernel_avx512_int8_blklen128.h" +#include "sqnbitgemm_kernel_avx512_2bit.h" +#include "sqnbitgemm_kernel_avx512vnni_2bit_blklen64.h" MLAS_FORCEINLINE void SQ4BitGemmM1Kernel_CompFp32( @@ -479,5 +481,13 @@ const MLAS_QNBIT_GEMM_DISPATCH MlasSQNBitGemmDispatchAvx512vnni = []() { d.SQ8BitGemmKernel_BlkSum_CompInt8 = SQ8BitGemmKernel_BlkSum_CompInt8_avx512vnni; d.QuantizeARowComputeBlkSum_CompInt8 = QuantizeARow_CompInt8_avx512; + // 2-bit native CompInt8 path. Phase 2b: scalar reference. Phase 3: this slot + // gets swapped for the AVX-512-VNNI vectorized kernel (`_mm512_dpbusd_epi32`). + // Plain AVX-512 (non-VNNI) hosts intentionally do not register this path and + // fall through to the LUT kernel. + d.Q2BitGemmPackQuantBDataSize = onnxruntime::mlas::sq2bit_avx512::Q2BitGemmPackQuantBDataSize_Avx512; + d.SQ2BitGemmPackQuantBDataAndBlkSum = onnxruntime::mlas::sq2bit_avx512::SQ2BitGemmPackQuantBDataAndBlkSum_Scalar; + d.SQ2BitGemmKernel_BlkSum_CompInt8 = onnxruntime::mlas::sq2bit_avx512::SQ2BitGemmKernel_BlkSum_CompInt8_Avx512Vnni; + return d; }(); diff --git a/onnxruntime/core/mlas/lib/sqnbitgemm_kernel_avx512vnni_2bit_blklen64.h b/onnxruntime/core/mlas/lib/sqnbitgemm_kernel_avx512vnni_2bit_blklen64.h new file mode 100644 index 0000000000000..6178bbcbce423 --- /dev/null +++ b/onnxruntime/core/mlas/lib/sqnbitgemm_kernel_avx512vnni_2bit_blklen64.h @@ -0,0 +1,816 @@ +/*++ + +Copyright (c) Microsoft Corporation. All rights reserved. + +Licensed under the MIT License. + +Module Name: + + sqnbitgemm_kernel_avx512vnni_2bit_blklen64.h + +Abstract: + + Phase 4 AVX-512-VNNI tiled kernel for the 2-bit weight CompInt8 GEMM + path (BlkBitWidth=2, BlkLen=64). Header-only; included exactly once by + sqnbitgemm_kernel_avx512vnni.cpp, which carries the required compile + flags (-mavx512vnni -mavx512bw -mavx512dq -mavx512vl -mavx512f on + GCC/Clang; MSVC generates the right instructions from intrinsics). + + Architecture mirrors W4's MlasQ4Int8GemmKernelBlkLen64Avx512 + (in sqnbitgemm_kernel_avx512_int8_blklen64.h): + + * Per-(m,n) accumulator is a __m512 carrying lane-interleaved + scaled partial sums. Final _mm512_reduce_add_ps happens ONCE per + (m,n) tile, not per K-block. + * Outer tile shape is R2 x C4 (2 M-rows x 4 N-cols) so each A vector + load amortises across 4 N-cols and each B load across 2 M-rows. + 16 ZMM accumulators in flight. + * Inner unroll is PerAccuBlk2 = 2 K-blocks per iteration. Pairs of + dpbusd outputs are interleaved (unpacklo/hi + add) so a single + FMA applies both blocks' scales in one shot. + * BlkLen=64 only. + * Symmetric quantization only (QuantBZeroPoint must be null). + * Tail tiles R2xC1, R1xC4, R1xC1 cover M % 2 != 0 / N % 4 != 0. + + Zero-point correction is performed OUTSIDE the int8 kernel via the + platform float SGEMM kernel (GetMlasPlatform().GemmFloatKernel), exactly + as W4 does. This requires QuantBBlkSum to be in the W4 "width-16 row- + major chunked" layout, which is produced by SQ2BitGemmPackQuantBData + AndBlkSum_Scalar. + + Dequant prologue (per 64-element block): + + __m128i p = _mm_loadu_si128(packed); // 16 packed bytes + __m512i p4 = _mm512_broadcast_i32x4(p); // 4 lanes of those 16 bytes + __m512i sh = {0,0,0,0, 2,2,2,2, 4,4,4,4, 6,6,6,6}; // per-dword shifts + __m512i v = _mm512_srlv_epi32(p4, sh); + __m512i b = _mm512_and_si512(v, _mm512_set1_epi8(0x03)); + + DESIGN DEVIATION FROM W4 (worth re-revisiting in a follow-up if perf + still falls short): W4 packs weights in a 4-N-col-grouped layout so the + 4 cols of a tile's data lie consecutively in memory (one ColStride + advances within the same K-block-pair). W2 currently uses the simpler + column-major layout (each col is BlockCountK * 16 bytes apart in memory), + which the kernel addresses via a multi-stream stride pattern. Lower + spatial locality, larger working set; on streaming shapes the HW + prefetcher copes, but for compute-bound small shapes this likely costs + a few percent. + +--*/ + +#pragma once + +#include +#include +#include + +#include + +#include "mlasi.h" +#include "qnbitgemm.h" +#include "sqnbitgemm_kernel_avx512_2bit.h" + +namespace onnxruntime { +namespace mlas { +namespace sq2bit_avx512 { + +// Number of K-blocks the inner loop processes per iteration. Mirrors W4. +inline constexpr size_t kPerAccuBlk2 = 2; +// Outer tile shape. Mirrors W4 NCols4 / NRows2. +inline constexpr size_t kNCols4 = 4; +inline constexpr size_t kNRows2 = 2; + +// +// Dequant one 64-element 2-bit weight block from the packed (broadcast + +// shift) layout into a ZMM of 64 unsigned bytes in [0, 3]. +// +// Bytes 0..15 : weights[0..15] (shift 0, & 0x03) +// Bytes 16..31 : weights[16..31] (shift 2, & 0x03) +// Bytes 32..47 : weights[32..47] (shift 4, & 0x03) +// Bytes 48..63 : weights[48..63] (shift 6, & 0x03) +// +static MLAS_FORCEINLINE __m512i +unpack_w2_blk_to_zmm(__m128i p128) +{ + const __m512i p_dup = _mm512_broadcast_i32x4(p128); + // Per-dword right shifts, in memory order (lane 0 first): + // [0,0,0,0, 2,2,2,2, 4,4,4,4, 6,6,6,6] + // _mm512_set_epi32 takes args in reverse order (lane 15 first). + const __m512i shifts = _mm512_set_epi32( + 6, 6, 6, 6, + 4, 4, 4, 4, + 2, 2, 2, 2, + 0, 0, 0, 0); + const __m512i mask03 = _mm512_set1_epi8(0x03); + return _mm512_and_si512(_mm512_srlv_epi32(p_dup, shifts), mask03); +} + +// +// Load + dequant ONE block. Used by the single-block (tail) helpers. +// +static MLAS_FORCEINLINE __m512i +load_unpack_1blk_w2(const std::byte* packed) +{ + return unpack_w2_blk_to_zmm( + _mm_loadu_si128(reinterpret_cast(packed))); +} + +// +// Load + dequant TWO consecutive K-blocks via one 256-bit YMM load. +// 32 packed bytes -> 2 ZMMs of 64 weights each. +// +static MLAS_FORCEINLINE void +load_unpack_2blk_w2(const std::byte* packed, __m512i& bv0_64_epi8, __m512i& bv1_64_epi8) +{ + const __m256i p_ymm = _mm256_loadu_si256(reinterpret_cast(packed)); + bv0_64_epi8 = unpack_w2_blk_to_zmm(_mm256_castsi256_si128(p_ymm)); + bv1_64_epi8 = unpack_w2_blk_to_zmm(_mm256_extracti128_si256(p_ymm, 1)); +} + +// +// Lane-interleaved 2-K-block accumulator (single M-row, single N-col). +// Mirrors W4's dot_accumulate_2blkvnni; identical math for W2 because the +// dequanted block is already in the same uint8 [0,3] form W4 produces. +// +// acc += sum_2blks( cvt(dpbusd(bv, av)) * scale_a * scale_b ) +// +// `scale_a` and `scale_b` point to TWO consecutive floats (scales for blk0 +// and blk1). The double-broadcast trick gives the 16-lane pattern +// [s0,s1, s0,s1, s0,s1, s0,s1, s0,s1, s0,s1, s0,s1, s0,s1] +// which matches the post-unpacklo/hi+add lane layout (blk0 in even lanes, +// blk1 in odd lanes). +// +static MLAS_FORCEINLINE void +dot_accumulate_2blk_w2_vnni( + const __m512i& av0_64_epi8, + const __m512i& av1_64_epi8, + const float* scale_a, + const __m512i& bv0_64_epi8, + const __m512i& bv1_64_epi8, + const __m512& scale_b_16_ps, + __m512& acc) +{ + const __m512i dot0_16_epi32 = _mm512_dpbusd_epi32(_mm512_setzero_si512(), bv0_64_epi8, av0_64_epi8); + const __m512i dot1_16_epi32 = _mm512_dpbusd_epi32(_mm512_setzero_si512(), bv1_64_epi8, av1_64_epi8); + + const __m512i t1 = _mm512_unpacklo_epi32(dot0_16_epi32, dot1_16_epi32); + const __m512i t2 = _mm512_unpackhi_epi32(dot0_16_epi32, dot1_16_epi32); + const __m512i sum_16_epi32 = _mm512_add_epi32(t1, t2); + const __m512 sum_16_ps = _mm512_cvtepi32_ps(sum_16_epi32); + + const __m256 scale_a_8_ps = _mm256_castpd_ps(_mm256_broadcast_sd(reinterpret_cast(scale_a))); + const __m512 scale_a_16_ps = _mm512_broadcast_f32x8(scale_a_8_ps); + + acc = _mm512_fmadd_ps(sum_16_ps, _mm512_mul_ps(scale_a_16_ps, scale_b_16_ps), acc); +} + +// +// Single-K-block accumulator. Uses uniform 16-lane scale broadcast since +// there's only one block's scale in play. +// +static MLAS_FORCEINLINE void +dot_accumulate_1blk_w2_vnni( + const __m512i& av_64_epi8, + const float* scale_a, + const __m512i& bv_64_epi8, + const __m512& scale_b_16_ps, + __m512& acc) +{ + const __m512i dot_16_epi32 = _mm512_dpbusd_epi32(_mm512_setzero_si512(), bv_64_epi8, av_64_epi8); + const __m512 sum_16_ps = _mm512_cvtepi32_ps(dot_16_epi32); + + const __m128 scale_a_ps = _mm_broadcast_ss(scale_a); + const __m512 scale_a_16_ps = _mm512_broadcast_f32x2(scale_a_ps); + + acc = _mm512_fmadd_ps(sum_16_ps, _mm512_mul_ps(scale_a_16_ps, scale_b_16_ps), acc); +} + +// +// 2 M-rows x 1 N-col x 2 K-blocks accumulator. The 2-block B load is shared +// across the 2 M-rows. +// +static MLAS_FORCEINLINE void +accumulate_w2_blklen64_r2c1blk2_vnni( + const __m512i& av00_64_epi8, const __m512i& av01_64_epi8, + const __m512i& av10_64_epi8, const __m512i& av11_64_epi8, + const std::byte* QuantBDataPtr, + const float* scale_a0, + const float* scale_a1, + const float* scale_b, + __m512& acc0, + __m512& acc1) +{ + __m512i bv0, bv1; + load_unpack_2blk_w2(QuantBDataPtr, bv0, bv1); + + const __m256 scale_b_8_ps = _mm256_castpd_ps(_mm256_broadcast_sd(reinterpret_cast(scale_b))); + const __m512 scale_b_16_ps = _mm512_broadcast_f32x8(scale_b_8_ps); + + dot_accumulate_2blk_w2_vnni(av00_64_epi8, av01_64_epi8, scale_a0, bv0, bv1, scale_b_16_ps, acc0); + dot_accumulate_2blk_w2_vnni(av10_64_epi8, av11_64_epi8, scale_a1, bv0, bv1, scale_b_16_ps, acc1); +} + +// +// 2 M-rows x 1 N-col x 1 K-block accumulator (K-tail). +// +static MLAS_FORCEINLINE void +accumulate_w2_blklen64_r2c1blk1_vnni( + const __m512i& av0_64_epi8, + const __m512i& av1_64_epi8, + const std::byte* QuantBDataPtr, + const float* scale_a0, + const float* scale_a1, + const float* scale_b, + __m512& acc0, + __m512& acc1) +{ + const __m512i bv = load_unpack_1blk_w2(QuantBDataPtr); + + const __m128 scale_b_ps = _mm_broadcast_ss(scale_b); + const __m512 scale_b_16_ps = _mm512_broadcast_f32x2(scale_b_ps); + + dot_accumulate_1blk_w2_vnni(av0_64_epi8, scale_a0, bv, scale_b_16_ps, acc0); + dot_accumulate_1blk_w2_vnni(av1_64_epi8, scale_a1, bv, scale_b_16_ps, acc1); +} + +// +// 1 M-row x 1 N-col x 2 K-blocks accumulator. +// +static MLAS_FORCEINLINE void +accumulate_w2_blklen64_r1c1blk2_vnni( + const __m512i& av0_64_epi8, + const __m512i& av1_64_epi8, + const std::byte* QuantBDataPtr, + const float* scale_a, + const float* scale_b, + __m512& acc) +{ + __m512i bv0, bv1; + load_unpack_2blk_w2(QuantBDataPtr, bv0, bv1); + + const __m256 scale_b_8_ps = _mm256_castpd_ps(_mm256_broadcast_sd(reinterpret_cast(scale_b))); + const __m512 scale_b_16_ps = _mm512_broadcast_f32x8(scale_b_8_ps); + + dot_accumulate_2blk_w2_vnni(av0_64_epi8, av1_64_epi8, scale_a, bv0, bv1, scale_b_16_ps, acc); +} + +// +// 1 M-row x 1 N-col x 1 K-block accumulator. +// +static MLAS_FORCEINLINE void +accumulate_w2_blklen64_r1c1blk1_vnni( + const __m512i& av_64_epi8, + const std::byte* QuantBDataPtr, + const float* scale_a, + const float* scale_b, + __m512& acc) +{ + const __m512i bv = load_unpack_1blk_w2(QuantBDataPtr); + + const __m128 scale_b_ps = _mm_broadcast_ss(scale_b); + const __m512 scale_b_16_ps = _mm512_broadcast_f32x2(scale_b_ps); + + dot_accumulate_1blk_w2_vnni(av_64_epi8, scale_a, bv, scale_b_16_ps, acc); +} + +// +// R2 x C4 tile. Main hot path for the customer model. +// +// Layout assumptions: +// * QuantBData : column-major. Col n at offset n * BlockCountK * kBlkBytes. +// * QuantBScale: column-major. Col n at offset n * BlockCountK. +// * QuantA : row-major int8, BlockCountK * kBlkLen bytes per row. +// * QuantAScale: row-major float, BlockCountK floats per row. +// +MLAS_FORCEINLINE void +Q2Int8GemmR2xC4BlkLen64Avx512Vnni( + const std::byte* QuantA, + const float* QuantAScale, + const std::byte* QuantBData, + const float* QuantBScale, + float* C, + size_t CountM, + size_t CountN, + size_t BlockCountK, + const float* Bias, + size_t ldc) +{ + const size_t lda = BlockCountK * kBlkLen; + const size_t ColStrideBytes = BlockCountK * kBlkBytes; + const size_t ColStrideScale = BlockCountK; + + assert(CountM % kNRows2 == 0); + assert(CountN % kNCols4 == 0); + + for (size_t m = 0; m < CountM; m += kNRows2) { + const std::byte* QuantBDataColPtr = QuantBData; + const float* QuantBScaleColPtr = QuantBScale; + const float* BiasPtr = Bias; + float* SumPtr = C + m * ldc; + + for (size_t n = 0; n < CountN; n += kNCols4) { + const std::byte* QuantAPtr = QuantA + m * lda; + const float* QuantAScalePtr = QuantAScale + m * BlockCountK; + + const std::byte* QuantBDataPtr = QuantBDataColPtr; + const float* QuantBScalePtr = QuantBScaleColPtr; + + __m512 acc[kNCols4 * kNRows2] = { + _mm512_setzero_ps(), _mm512_setzero_ps(), _mm512_setzero_ps(), _mm512_setzero_ps(), + _mm512_setzero_ps(), _mm512_setzero_ps(), _mm512_setzero_ps(), _mm512_setzero_ps() + }; + + size_t k_blks_remaining = BlockCountK; + for (; k_blks_remaining > 1; k_blks_remaining -= kPerAccuBlk2) { + const __m512i av_00 = _mm512_loadu_si512(reinterpret_cast(QuantAPtr)); + const __m512i av_01 = _mm512_loadu_si512(reinterpret_cast(QuantAPtr + kBlkLen)); + const __m512i av_10 = _mm512_loadu_si512(reinterpret_cast(QuantAPtr + lda)); + const __m512i av_11 = _mm512_loadu_si512(reinterpret_cast(QuantAPtr + lda + kBlkLen)); + + accumulate_w2_blklen64_r2c1blk2_vnni( + av_00, av_01, av_10, av_11, + QuantBDataPtr + 0 * ColStrideBytes, + QuantAScalePtr, QuantAScalePtr + BlockCountK, + QuantBScalePtr + 0 * ColStrideScale, + acc[0], acc[kNCols4 + 0]); + accumulate_w2_blklen64_r2c1blk2_vnni( + av_00, av_01, av_10, av_11, + QuantBDataPtr + 1 * ColStrideBytes, + QuantAScalePtr, QuantAScalePtr + BlockCountK, + QuantBScalePtr + 1 * ColStrideScale, + acc[1], acc[kNCols4 + 1]); + accumulate_w2_blklen64_r2c1blk2_vnni( + av_00, av_01, av_10, av_11, + QuantBDataPtr + 2 * ColStrideBytes, + QuantAScalePtr, QuantAScalePtr + BlockCountK, + QuantBScalePtr + 2 * ColStrideScale, + acc[2], acc[kNCols4 + 2]); + accumulate_w2_blklen64_r2c1blk2_vnni( + av_00, av_01, av_10, av_11, + QuantBDataPtr + 3 * ColStrideBytes, + QuantAScalePtr, QuantAScalePtr + BlockCountK, + QuantBScalePtr + 3 * ColStrideScale, + acc[3], acc[kNCols4 + 3]); + + QuantAPtr += kBlkLen * kPerAccuBlk2; + QuantAScalePtr += kPerAccuBlk2; + QuantBDataPtr += kPerAccuBlk2 * kBlkBytes; + QuantBScalePtr += kPerAccuBlk2; + } + + while (k_blks_remaining-- > 0) { + const __m512i av_00 = _mm512_loadu_si512(reinterpret_cast(QuantAPtr)); + const __m512i av_10 = _mm512_loadu_si512(reinterpret_cast(QuantAPtr + lda)); + + accumulate_w2_blklen64_r2c1blk1_vnni( + av_00, av_10, + QuantBDataPtr + 0 * ColStrideBytes, + QuantAScalePtr, QuantAScalePtr + BlockCountK, + QuantBScalePtr + 0 * ColStrideScale, + acc[0], acc[kNCols4 + 0]); + accumulate_w2_blklen64_r2c1blk1_vnni( + av_00, av_10, + QuantBDataPtr + 1 * ColStrideBytes, + QuantAScalePtr, QuantAScalePtr + BlockCountK, + QuantBScalePtr + 1 * ColStrideScale, + acc[1], acc[kNCols4 + 1]); + accumulate_w2_blklen64_r2c1blk1_vnni( + av_00, av_10, + QuantBDataPtr + 2 * ColStrideBytes, + QuantAScalePtr, QuantAScalePtr + BlockCountK, + QuantBScalePtr + 2 * ColStrideScale, + acc[2], acc[kNCols4 + 2]); + accumulate_w2_blklen64_r2c1blk1_vnni( + av_00, av_10, + QuantBDataPtr + 3 * ColStrideBytes, + QuantAScalePtr, QuantAScalePtr + BlockCountK, + QuantBScalePtr + 3 * ColStrideScale, + acc[3], acc[kNCols4 + 3]); + + QuantAPtr += kBlkLen; + QuantAScalePtr++; + QuantBDataPtr += kBlkBytes; + QuantBScalePtr++; + } + + SumPtr[0] = _mm512_reduce_add_ps(acc[0]); + SumPtr[1] = _mm512_reduce_add_ps(acc[1]); + SumPtr[2] = _mm512_reduce_add_ps(acc[2]); + SumPtr[3] = _mm512_reduce_add_ps(acc[3]); + SumPtr[ldc + 0] = _mm512_reduce_add_ps(acc[kNCols4 + 0]); + SumPtr[ldc + 1] = _mm512_reduce_add_ps(acc[kNCols4 + 1]); + SumPtr[ldc + 2] = _mm512_reduce_add_ps(acc[kNCols4 + 2]); + SumPtr[ldc + 3] = _mm512_reduce_add_ps(acc[kNCols4 + 3]); + if (BiasPtr != nullptr) { + SumPtr[0] += BiasPtr[0]; + SumPtr[1] += BiasPtr[1]; + SumPtr[2] += BiasPtr[2]; + SumPtr[3] += BiasPtr[3]; + SumPtr[ldc + 0] += BiasPtr[0]; + SumPtr[ldc + 1] += BiasPtr[1]; + SumPtr[ldc + 2] += BiasPtr[2]; + SumPtr[ldc + 3] += BiasPtr[3]; + } + + QuantBDataColPtr += kNCols4 * ColStrideBytes; + QuantBScaleColPtr += kNCols4 * ColStrideScale; + BiasPtr += BiasPtr != nullptr ? kNCols4 : 0; + SumPtr += kNCols4; + } + } +} + +// +// R2 x C1 tile (N-tail). +// +MLAS_FORCEINLINE void +Q2Int8GemmR2xC1BlkLen64Avx512Vnni( + const std::byte* QuantA, + const float* QuantAScale, + const std::byte* QuantBData, + const float* QuantBScale, + float* C, + size_t CountM, + size_t CountN, + size_t BlockCountK, + const float* Bias, + size_t ldc) +{ + const size_t lda = BlockCountK * kBlkLen; + const size_t ColStrideBytes = BlockCountK * kBlkBytes; + const size_t ColStrideScale = BlockCountK; + + assert(CountM % kNRows2 == 0); + + for (size_t m = 0; m < CountM; m += kNRows2) { + const std::byte* QuantBDataColPtr = QuantBData; + const float* QuantBScaleColPtr = QuantBScale; + const float* BiasPtr = Bias; + float* SumPtr = C + m * ldc; + + for (size_t n = 0; n < CountN; ++n) { + const std::byte* QuantAPtr = QuantA + m * lda; + const float* QuantAScalePtr = QuantAScale + m * BlockCountK; + + const std::byte* QuantBDataPtr = QuantBDataColPtr; + const float* QuantBScalePtr = QuantBScaleColPtr; + + __m512 acc0 = _mm512_setzero_ps(); + __m512 acc1 = _mm512_setzero_ps(); + + size_t k_blks_remaining = BlockCountK; + for (; k_blks_remaining > 1; k_blks_remaining -= kPerAccuBlk2) { + const __m512i av_00 = _mm512_loadu_si512(reinterpret_cast(QuantAPtr)); + const __m512i av_01 = _mm512_loadu_si512(reinterpret_cast(QuantAPtr + kBlkLen)); + const __m512i av_10 = _mm512_loadu_si512(reinterpret_cast(QuantAPtr + lda)); + const __m512i av_11 = _mm512_loadu_si512(reinterpret_cast(QuantAPtr + lda + kBlkLen)); + + accumulate_w2_blklen64_r2c1blk2_vnni( + av_00, av_01, av_10, av_11, + QuantBDataPtr, + QuantAScalePtr, QuantAScalePtr + BlockCountK, + QuantBScalePtr, + acc0, acc1); + + QuantAPtr += kBlkLen * kPerAccuBlk2; + QuantAScalePtr += kPerAccuBlk2; + QuantBDataPtr += kPerAccuBlk2 * kBlkBytes; + QuantBScalePtr += kPerAccuBlk2; + } + + while (k_blks_remaining-- > 0) { + const __m512i av_00 = _mm512_loadu_si512(reinterpret_cast(QuantAPtr)); + const __m512i av_10 = _mm512_loadu_si512(reinterpret_cast(QuantAPtr + lda)); + + accumulate_w2_blklen64_r2c1blk1_vnni( + av_00, av_10, + QuantBDataPtr, + QuantAScalePtr, QuantAScalePtr + BlockCountK, + QuantBScalePtr, + acc0, acc1); + + QuantAPtr += kBlkLen; + QuantAScalePtr++; + QuantBDataPtr += kBlkBytes; + QuantBScalePtr++; + } + + SumPtr[0] = _mm512_reduce_add_ps(acc0); + SumPtr[ldc] = _mm512_reduce_add_ps(acc1); + if (BiasPtr != nullptr) { + SumPtr[0] += BiasPtr[0]; + SumPtr[ldc] += BiasPtr[0]; + } + + QuantBDataColPtr += ColStrideBytes; + QuantBScaleColPtr += ColStrideScale; + BiasPtr += BiasPtr != nullptr ? 1 : 0; + SumPtr += 1; + } + } +} + +// +// R1 x C4 tile (M-tail). +// +MLAS_FORCEINLINE void +Q2Int8GemmR1xC4BlkLen64Avx512Vnni( + const std::byte* QuantA, + const float* QuantAScale, + const std::byte* QuantBData, + const float* QuantBScale, + float* C, + size_t CountM, + size_t CountN, + size_t BlockCountK, + const float* Bias, + size_t ldc) +{ + const size_t lda = BlockCountK * kBlkLen; + const size_t ColStrideBytes = BlockCountK * kBlkBytes; + const size_t ColStrideScale = BlockCountK; + + assert(CountN % kNCols4 == 0); + + for (size_t m = 0; m < CountM; ++m) { + const std::byte* QuantBDataColPtr = QuantBData; + const float* QuantBScaleColPtr = QuantBScale; + const float* BiasPtr = Bias; + float* SumPtr = C + m * ldc; + + for (size_t n = 0; n < CountN; n += kNCols4) { + const std::byte* QuantAPtr = QuantA + m * lda; + const float* QuantAScalePtr = QuantAScale + m * BlockCountK; + + const std::byte* QuantBDataPtr = QuantBDataColPtr; + const float* QuantBScalePtr = QuantBScaleColPtr; + + __m512 acc[kNCols4] = { + _mm512_setzero_ps(), _mm512_setzero_ps(), + _mm512_setzero_ps(), _mm512_setzero_ps() + }; + + size_t k_blks_remaining = BlockCountK; + for (; k_blks_remaining > 1; k_blks_remaining -= kPerAccuBlk2) { + const __m512i av_0 = _mm512_loadu_si512(reinterpret_cast(QuantAPtr)); + const __m512i av_1 = _mm512_loadu_si512(reinterpret_cast(QuantAPtr + kBlkLen)); + + accumulate_w2_blklen64_r1c1blk2_vnni( + av_0, av_1, + QuantBDataPtr + 0 * ColStrideBytes, + QuantAScalePtr, QuantBScalePtr + 0 * ColStrideScale, acc[0]); + accumulate_w2_blklen64_r1c1blk2_vnni( + av_0, av_1, + QuantBDataPtr + 1 * ColStrideBytes, + QuantAScalePtr, QuantBScalePtr + 1 * ColStrideScale, acc[1]); + accumulate_w2_blklen64_r1c1blk2_vnni( + av_0, av_1, + QuantBDataPtr + 2 * ColStrideBytes, + QuantAScalePtr, QuantBScalePtr + 2 * ColStrideScale, acc[2]); + accumulate_w2_blklen64_r1c1blk2_vnni( + av_0, av_1, + QuantBDataPtr + 3 * ColStrideBytes, + QuantAScalePtr, QuantBScalePtr + 3 * ColStrideScale, acc[3]); + + QuantAPtr += kBlkLen * kPerAccuBlk2; + QuantAScalePtr += kPerAccuBlk2; + QuantBDataPtr += kPerAccuBlk2 * kBlkBytes; + QuantBScalePtr += kPerAccuBlk2; + } + + while (k_blks_remaining-- > 0) { + const __m512i av = _mm512_loadu_si512(reinterpret_cast(QuantAPtr)); + + accumulate_w2_blklen64_r1c1blk1_vnni( + av, QuantBDataPtr + 0 * ColStrideBytes, + QuantAScalePtr, QuantBScalePtr + 0 * ColStrideScale, acc[0]); + accumulate_w2_blklen64_r1c1blk1_vnni( + av, QuantBDataPtr + 1 * ColStrideBytes, + QuantAScalePtr, QuantBScalePtr + 1 * ColStrideScale, acc[1]); + accumulate_w2_blklen64_r1c1blk1_vnni( + av, QuantBDataPtr + 2 * ColStrideBytes, + QuantAScalePtr, QuantBScalePtr + 2 * ColStrideScale, acc[2]); + accumulate_w2_blklen64_r1c1blk1_vnni( + av, QuantBDataPtr + 3 * ColStrideBytes, + QuantAScalePtr, QuantBScalePtr + 3 * ColStrideScale, acc[3]); + + QuantAPtr += kBlkLen; + QuantAScalePtr++; + QuantBDataPtr += kBlkBytes; + QuantBScalePtr++; + } + + SumPtr[0] = _mm512_reduce_add_ps(acc[0]); + SumPtr[1] = _mm512_reduce_add_ps(acc[1]); + SumPtr[2] = _mm512_reduce_add_ps(acc[2]); + SumPtr[3] = _mm512_reduce_add_ps(acc[3]); + if (BiasPtr != nullptr) { + SumPtr[0] += BiasPtr[0]; + SumPtr[1] += BiasPtr[1]; + SumPtr[2] += BiasPtr[2]; + SumPtr[3] += BiasPtr[3]; + } + + QuantBDataColPtr += kNCols4 * ColStrideBytes; + QuantBScaleColPtr += kNCols4 * ColStrideScale; + BiasPtr += BiasPtr != nullptr ? kNCols4 : 0; + SumPtr += kNCols4; + } + } +} + +// +// R1 x C1 tile (corner). +// +MLAS_FORCEINLINE void +Q2Int8GemmR1xC1BlkLen64Avx512Vnni( + const std::byte* QuantA, + const float* QuantAScale, + const std::byte* QuantBData, + const float* QuantBScale, + float* C, + size_t CountM, + size_t CountN, + size_t BlockCountK, + const float* Bias, + size_t ldc) +{ + const size_t lda = BlockCountK * kBlkLen; + const size_t ColStrideBytes = BlockCountK * kBlkBytes; + const size_t ColStrideScale = BlockCountK; + + for (size_t m = 0; m < CountM; ++m) { + const std::byte* QuantBDataColPtr = QuantBData; + const float* QuantBScaleColPtr = QuantBScale; + const float* BiasPtr = Bias; + float* SumPtr = C + m * ldc; + + for (size_t n = 0; n < CountN; ++n) { + const std::byte* QuantAPtr = QuantA + m * lda; + const float* QuantAScalePtr = QuantAScale + m * BlockCountK; + + const std::byte* QuantBDataPtr = QuantBDataColPtr; + const float* QuantBScalePtr = QuantBScaleColPtr; + + __m512 acc = _mm512_setzero_ps(); + + size_t k_blks_remaining = BlockCountK; + for (; k_blks_remaining > 1; k_blks_remaining -= kPerAccuBlk2) { + const __m512i av_0 = _mm512_loadu_si512(reinterpret_cast(QuantAPtr)); + const __m512i av_1 = _mm512_loadu_si512(reinterpret_cast(QuantAPtr + kBlkLen)); + + accumulate_w2_blklen64_r1c1blk2_vnni( + av_0, av_1, QuantBDataPtr, QuantAScalePtr, QuantBScalePtr, acc); + + QuantAPtr += kBlkLen * kPerAccuBlk2; + QuantAScalePtr += kPerAccuBlk2; + QuantBDataPtr += kPerAccuBlk2 * kBlkBytes; + QuantBScalePtr += kPerAccuBlk2; + } + + while (k_blks_remaining-- > 0) { + const __m512i av = _mm512_loadu_si512(reinterpret_cast(QuantAPtr)); + + accumulate_w2_blklen64_r1c1blk1_vnni( + av, QuantBDataPtr, QuantAScalePtr, QuantBScalePtr, acc); + + QuantAPtr += kBlkLen; + QuantAScalePtr++; + QuantBDataPtr += kBlkBytes; + QuantBScalePtr++; + } + + SumPtr[0] = _mm512_reduce_add_ps(acc); + if (BiasPtr != nullptr) { + SumPtr[0] += BiasPtr[0]; + } + + QuantBDataColPtr += ColStrideBytes; + QuantBScaleColPtr += ColStrideScale; + BiasPtr += BiasPtr != nullptr ? 1 : 0; + SumPtr += 1; + } + } +} + +// +// Tile dispatcher. Mirrors W4's MlasQ4Int8GemmKernelBlkLen64Avx512. +// +MLAS_FORCEINLINE void +MlasQ2Int8GemmKernelBlkLen64Avx512Vnni( + const std::byte* QuantA, + const float* QuantAScale, + const std::byte* QuantBData, + const float* QuantBScale, + float* C, + size_t CountM, + size_t CountN, + size_t BlockCountK, + const float* Bias, + size_t ldc) +{ + const size_t lda = BlockCountK * kBlkLen; + const size_t lda_scale = BlockCountK; + const size_t ColStrideBytes = BlockCountK * kBlkBytes; + const size_t ColStrideScale = BlockCountK; + + const size_t remainingRows = CountM % kNRows2; + const size_t multipleRows = CountM - remainingRows; + const size_t remainingCols = CountN % kNCols4; + const size_t multipleCols = CountN - remainingCols; + + if (multipleRows > 0 && multipleCols > 0) { + Q2Int8GemmR2xC4BlkLen64Avx512Vnni( + QuantA, QuantAScale, QuantBData, QuantBScale, + C, multipleRows, multipleCols, BlockCountK, Bias, ldc); + } + if (remainingCols > 0 && multipleRows > 0) { + Q2Int8GemmR2xC1BlkLen64Avx512Vnni( + QuantA, QuantAScale, + QuantBData + multipleCols * ColStrideBytes, + QuantBScale + multipleCols * ColStrideScale, + C + multipleCols, + multipleRows, remainingCols, BlockCountK, + Bias ? Bias + multipleCols : nullptr, ldc); + } + if (remainingRows > 0 && multipleCols > 0) { + Q2Int8GemmR1xC4BlkLen64Avx512Vnni( + QuantA + multipleRows * lda, + QuantAScale + multipleRows * lda_scale, + QuantBData, QuantBScale, + C + multipleRows * ldc, + remainingRows, multipleCols, BlockCountK, Bias, ldc); + } + if (remainingRows > 0 && remainingCols > 0) { + Q2Int8GemmR1xC1BlkLen64Avx512Vnni( + QuantA + multipleRows * lda, + QuantAScale + multipleRows * lda_scale, + QuantBData + multipleCols * ColStrideBytes, + QuantBScale + multipleCols * ColStrideScale, + C + multipleRows * ldc + multipleCols, + remainingRows, remainingCols, BlockCountK, + Bias ? Bias + multipleCols : nullptr, ldc); + } +} + +// +// Top-level kernel registered into MlasSQNBitGemmDispatchAvx512vnni. +// +// 1) Calls the tile dispatcher above, which computes +// C[m,n] = bias[n] + sum_blk(scale_a * scale_b * dpbusd(b, a)) +// 2) Adds the symmetric zero-point correction +// C[m,n] += sum_blk(ABlockSum[m,blk] * QuantBBlkSum[n,blk]) +// via the platform float SGEMM micro-kernel. QuantBBlkSum is in the +// W4 width-16 row-major chunked layout (produced by the W2 pack +// function), which is what GemmFloatKernel expects for its packed-B +// operand. ZeroMode=false means SGEMM does `C += A @ B`, so the bias +// and int8 contribution already in C are preserved. +// +static MLAS_FORCEINLINE size_t MLASCALL +SQ2BitGemmKernel_BlkSum_CompInt8_Avx512Vnni( + const size_t BlkLen, + const std::byte* QuantA, + const float* QuantAScale, + const std::byte* QuantBData, + const float* QuantBScale, + const std::byte* /* QuantBZeroPoint */, + float* C, + size_t CountM, + size_t CountN, + size_t /* CountK */, + size_t BlockCountK, + const float* Bias, + size_t ldc, + const float* ABlockSum, + const float* QuantBBlkSum) +{ + if (BlkLen != kBlkLen) { + return 0; + } + + MlasQ2Int8GemmKernelBlkLen64Avx512Vnni( + QuantA, QuantAScale, QuantBData, QuantBScale, + C, CountM, CountN, BlockCountK, Bias, ldc); + + // BlkSum correction: C += ABlockSum [M x BlockCountK] @ QuantBBlkSum [BlockCountK x N]. + float* c_blk = C; + const float* b_blk_sum = QuantBBlkSum; + size_t RowsRemaining = CountM; + const float* a_blksum_row = ABlockSum; + while (RowsRemaining > 0) { + const auto RowsHandled = GetMlasPlatform().GemmFloatKernel( + a_blksum_row, b_blk_sum, c_blk, + BlockCountK, RowsRemaining, CountN, + BlockCountK, ldc, 1.0f, false); + + c_blk += ldc * RowsHandled; + a_blksum_row += BlockCountK * RowsHandled; + RowsRemaining -= RowsHandled; + } + + return CountM; +} + +} // namespace sq2bit_avx512 +} // namespace mlas +} // namespace onnxruntime diff --git a/onnxruntime/test/mlas/bench/bench_lutgemm.cpp b/onnxruntime/test/mlas/bench/bench_lutgemm.cpp index d94ccecc45eb0..235916fd20203 100644 --- a/onnxruntime/test/mlas/bench/bench_lutgemm.cpp +++ b/onnxruntime/test/mlas/bench/bench_lutgemm.cpp @@ -1,6 +1,7 @@ // Copyright (c) Microsoft Corporation. All rights reserved. // Licensed under the MIT License. +#include "mlas.h" #include "mlas_q4.h" #include "mlas_qnbit.h" #include "bench_util.h" @@ -97,10 +98,7 @@ void LUTGEMM_COMPUTE(benchmark::State& state) { const bool HasZeroPoint = static_cast(state.range(5)); const bool HasBias = static_cast(state.range(6)); - if (!MlasIsLutGemmAvailable(N, K, BlkBitWidth, BlkLen)) { - state.SkipWithMessage("LUT GEMM is not available with the given configuration."); - return; - } + const bool lut_available = MlasIsLutGemmAvailable(N, K, BlkBitWidth, BlkLen); OrtThreadPoolParams tpo; tpo.thread_pool_size = static_cast(Threads); @@ -130,21 +128,6 @@ void LUTGEMM_COMPUTE(benchmark::State& state) { static_cast(K), static_cast(N), static_cast(N), tp.get()); - MlasClearLutGemmKernelConfig(); - MlasInitLutGemmKernelConfig(N, K, BlkBitWidth, BlkLen, HasZeroPoint); - - size_t PackedBufSize = MlasLutGemmPackedSize(N, K, BlkBitWidth, BlkLen, HasZeroPoint); - std::vector PackedBuf(PackedBufSize); - - MlasLutGemmPack( - N, K, BlkBitWidth, BlkLen, HasZeroPoint, - reinterpret_cast(QuantBData.data()), - QuantBScale.data(), - HasZeroPoint ? QuantBZeroPoint.data() : nullptr, - false, // IsFloatZeroPoint - PackedBuf.data(), - tp.get()); - std::vector Bias; const float* BiasPtr = nullptr; if (HasBias) { @@ -152,14 +135,77 @@ void LUTGEMM_COMPUTE(benchmark::State& state) { BiasPtr = Bias.data(); } - MlasLutGemm(A.data(), BlkLen, PackedBuf.data(), C.data(), - static_cast(K), static_cast(M), static_cast(N), - HasZeroPoint, tp.get(), BiasPtr); + if (lut_available) { + state.SetLabel("path=LUT"); - for (auto _ : state) { + MlasClearLutGemmKernelConfig(); + MlasInitLutGemmKernelConfig(N, K, BlkBitWidth, BlkLen, HasZeroPoint); + + size_t PackedBufSize = MlasLutGemmPackedSize(N, K, BlkBitWidth, BlkLen, HasZeroPoint); + std::vector PackedBuf(PackedBufSize); + + MlasLutGemmPack( + N, K, BlkBitWidth, BlkLen, HasZeroPoint, + reinterpret_cast(QuantBData.data()), + QuantBScale.data(), + HasZeroPoint ? QuantBZeroPoint.data() : nullptr, + false, // IsFloatZeroPoint + PackedBuf.data(), + tp.get()); + + // warm-up MlasLutGemm(A.data(), BlkLen, PackedBuf.data(), C.data(), static_cast(K), static_cast(M), static_cast(N), HasZeroPoint, tp.get(), BiasPtr); + + for (auto _ : state) { + MlasLutGemm(A.data(), BlkLen, PackedBuf.data(), C.data(), + static_cast(K), static_cast(M), static_cast(N), + HasZeroPoint, tp.get(), BiasPtr); + } + } else { + // Fall back to ComputeBUnpacked-equivalent: dequant B to fp32 then SGEMM. + // This matches what the runtime does (matmul_nbits.cc ComputeBUnpacked) + // when LUT is gated out for the given shape (e.g. N % n_div != 0). + state.SetLabel("path=Dequant+SGEMM"); + + std::vector DequantB(K * N); // [K, N] row-major + + // Time dequant+SGEMM as a unit (this is what the runtime pays per call, + // since the dequantized buffer is not cached across MatMulNBits calls). + auto dequant_and_gemm = [&]() { + MlasDequantizeBlockwise( + DequantB.data(), QuantBData.data(), QuantBScale.data(), + HasZeroPoint ? QuantBZeroPoint.data() : nullptr, + static_cast(BlkLen), /*columnwise*/ true, + static_cast(K), static_cast(N), tp.get()); + + MlasGemm(CblasNoTrans, CblasNoTrans, + M, N, K, + 1.0f, + A.data(), K, + DequantB.data(), N, + 0.0f, + C.data(), N, + tp.get(), nullptr); + + if (BiasPtr != nullptr) { + // Broadcast-add bias [N] over rows of C [M, N]. + for (size_t m = 0; m < M; ++m) { + float* row = C.data() + m * N; + for (size_t n = 0; n < N; ++n) { + row[n] += BiasPtr[n]; + } + } + } + }; + + // warm-up + dequant_and_gemm(); + + for (auto _ : state) { + dequant_and_gemm(); + } } } @@ -187,11 +233,38 @@ static void LutGemmComputeArgs(benchmark::internal::Benchmark* b) { }); } +// Customer 2-bit MatMulNBits shapes (BlkLen=64). Five distinct (K, N) pairs: +// (K=384, N=1024): 20 nodes +// (K=1024, N=192): 40 nodes +// (K=1024, N=384): 20 nodes +// (K=1024, N=4096): 20 nodes +// (K=4096, N=1024): 20 nodes +// M=128 is the widest-gap prefill shape vs the W4 CompInt8 path. +// Pair with QNBITGEMM/QNBitGemmCustomerArgs for head-to-head. +static void LutGemmCustomerArgs(benchmark::internal::Benchmark* b) { + b->ArgNames(lutgemm_compute_arg_names); + // Five separate Args() entries (rather than ArgsProduct) so we only run the + // exact (K, N) pairs that appear in the customer model. + const int64_t M = 128; + const int64_t BlkLen = 64; + const int64_t Threads = 8; + const int64_t HasZP = 0; + const int64_t HasBias = 1; + for (auto kn : {std::pair{384, 1024}, + std::pair{1024, 192}, + std::pair{1024, 384}, + std::pair{1024, 4096}, + std::pair{4096, 1024}}) { + b->Args({BlkLen, M, kn.second, kn.first, Threads, HasZP, HasBias}); + } +} + [[maybe_unused]] static const bool benchmarks_registered = []() { const bool is_lutgemm_supported = MlasIsLutGemmAvailable(4096, 4096, 2, 128); if (is_lutgemm_supported) { BENCHMARK(LUTGEMM_PACK<2>)->Apply(LutGemmPackArgs)->UseRealTime(); BENCHMARK(LUTGEMM_COMPUTE<2>)->Apply(LutGemmComputeArgs)->UseRealTime(); + BENCHMARK(LUTGEMM_COMPUTE<2>)->Apply(LutGemmCustomerArgs)->UseRealTime(); return true; } return false; diff --git a/onnxruntime/test/mlas/bench/bench_qnbitgemm.cpp b/onnxruntime/test/mlas/bench/bench_qnbitgemm.cpp index 351895f352b4e..f068db7ba8227 100644 --- a/onnxruntime/test/mlas/bench/bench_qnbitgemm.cpp +++ b/onnxruntime/test/mlas/bench/bench_qnbitgemm.cpp @@ -138,6 +138,60 @@ BENCHMARK(QNBITGEMM)->Apply(QNBitGemmArgs)->UseRealTime(); BENCHMARK(QNBITGEMM)->Apply(QNBitGemmArgs)->UseRealTime(); BENCHMARK(QNBITGEMM)->Apply(QNBitGemmArgs)->UseRealTime(); +// Customer MatMulNBits shapes mirrored at 4-bit for a head-to-head comparison +// vs the W2 LUT path (LUTGEMM_COMPUTE/CUSTOMER). Customer model uses BlkLen=64. +// Five distinct (K, N) pairs: +// (K=384, N=1024): 20 nodes +// (K=1024, N=192): 40 nodes +// (K=1024, N=384): 20 nodes +// (K=1024, N=4096): 20 nodes +// (K=4096, N=1024): 20 nodes +// M=128 is the widest-gap prefill shape from e2e benchmarks. +static void QNBitGemmCustomerArgs(benchmark::internal::Benchmark* b) { + b->ArgNames({"BlkLen", "M", "N", "K", "Threads", "Symmetric", "HasBias", "ComputeType"}); + const int64_t M = 128; + const int64_t BlkLen = 64; + const int64_t Threads = 8; + const int64_t Symmetric = 1; + const int64_t HasBias = 1; + for (auto kn : {std::pair{384, 1024}, + std::pair{1024, 192}, + std::pair{1024, 384}, + std::pair{1024, 4096}, + std::pair{4096, 1024}}) { + for (int64_t ct : {int64_t{SQNBIT_CompFp32}, int64_t{SQNBIT_CompInt8}}) { + b->Args({BlkLen, M, kn.second, kn.first, Threads, Symmetric, HasBias, ct}); + } + } +} + +BENCHMARK(QNBITGEMM)->Apply(QNBitGemmCustomerArgs)->UseRealTime(); + +// 2-bit weight rows for the same customer shapes. Phase 2b is a scalar +// reference; Phase 3 swaps in the AVX-512 (+VNNI) vectorized kernel and +// these rows become the head-to-head LUT-vs-native comparison surface. +// W2 vectorized path supports only BlkLen=64 + SQNBIT_CompInt8; the +// SQNBIT_CompFp32 row exercises the LUT fallback. +static void QNBit2BitCustomerArgs(benchmark::internal::Benchmark* b) { + b->ArgNames({"BlkLen", "M", "N", "K", "Threads", "Symmetric", "HasBias", "ComputeType"}); + const int64_t M = 128; + const int64_t BlkLen = 64; + const int64_t Threads = 8; + const int64_t Symmetric = 1; // W2 vectorized path is symmetric-only. + const int64_t HasBias = 1; + for (auto kn : {std::pair{384, 1024}, + std::pair{1024, 192}, + std::pair{1024, 384}, + std::pair{1024, 4096}, + std::pair{4096, 1024}}) { + for (int64_t ct : {int64_t{SQNBIT_CompFp32}, int64_t{SQNBIT_CompInt8}}) { + b->Args({BlkLen, M, kn.second, kn.first, Threads, Symmetric, HasBias, ct}); + } + } +} + +BENCHMARK(QNBITGEMM)->Apply(QNBit2BitCustomerArgs)->UseRealTime(); + // This test gets benchmark arguments from environment variables. template void QNBITGEMM_ENV(benchmark::State& state) { diff --git a/onnxruntime/test/mlas/unittest/test_sqnbitgemm_2bit.cpp b/onnxruntime/test/mlas/unittest/test_sqnbitgemm_2bit.cpp new file mode 100644 index 0000000000000..c6434916e8bac --- /dev/null +++ b/onnxruntime/test/mlas/unittest/test_sqnbitgemm_2bit.cpp @@ -0,0 +1,167 @@ +/*++ + +Copyright (c) Microsoft Corporation. All rights reserved. + +Licensed under the MIT License. + +Module Name: + + test_sqnbitgemm_2bit.cpp + +Abstract: + + Unit tests for the 2-bit AVX-512-VNNI weight-GEMM helpers. + + Phase 2a coverage: pack / unpack round-trip of the BlkLen=64 packed + layout. The tests exercise sqnbitgemm_kernel_avx512_2bit.h directly; + they do not depend on platform dispatch being wired up, so they run on + every host (the helpers are pure scalar bit-twiddling). + +--*/ + +#include "gtest/gtest.h" + +#include +#include +#include +#include + +#include "core/mlas/lib/sqnbitgemm_kernel_avx512_2bit.h" + +namespace { + +namespace sq2 = onnxruntime::mlas::sq2bit_avx512; + +// Standard ONNX 2-bit packing: byte_i = w[4i] | w[4i+1]<<2 | w[4i+2]<<4 | w[4i+3]<<6. +void +PackSourceBlock_BlkLen64(const uint8_t weights[sq2::kBlkLen], std::byte* src_out) +{ + for (size_t i = 0; i < sq2::kBlkBytes; ++i) { + const uint8_t v0 = weights[4 * i + 0] & 0x03u; + const uint8_t v1 = weights[4 * i + 1] & 0x03u; + const uint8_t v2 = weights[4 * i + 2] & 0x03u; + const uint8_t v3 = weights[4 * i + 3] & 0x03u; + src_out[i] = static_cast( + static_cast(v0 | (v1 << 2) | (v2 << 4) | (v3 << 6)) + ); + } +} + +} // namespace + +// +// Pack then immediately unpack a single block. Each 2-bit position must +// survive the layout permutation exactly. +// +TEST(MlasSq2BitTest, PackUnpackRoundTrip_BlkLen64_DeterministicPattern) +{ + // Use a deterministic pattern that touches every (position, value) pair: + // weight i gets value (i % 4). This guarantees that if any bit-position + // accounting is off the failing index pinpoints the bug. + std::array weights{}; + for (size_t i = 0; i < weights.size(); ++i) { + weights[i] = static_cast(i % 4); + } + + std::array src{}; + PackSourceBlock_BlkLen64(weights.data(), src.data()); + + // Sanity check: the source unpack reproduces the original weights. + std::array via_src{}; + sq2::UnpackSourceBlock_BlkLen64_Reference(src.data(), via_src.data()); + for (size_t i = 0; i < weights.size(); ++i) { + ASSERT_EQ(via_src[i], weights[i]) << "Source-unpack disagrees at i=" << i; + } + + // Pack into the new layout, then unpack via the reference inverse. + std::array packed{}; + sq2::PackBlock_BlkLen64(src.data(), packed.data()); + + std::array recovered{}; + sq2::UnpackBlock_BlkLen64_Reference(packed.data(), recovered.data()); + + for (size_t i = 0; i < weights.size(); ++i) { + ASSERT_EQ(recovered[i], weights[i]) + << "Round-trip mismatch at i=" << i + << ": expected " << static_cast(weights[i]) + << ", got " << static_cast(recovered[i]); + } +} + +// +// Same round-trip but with pseudo-random weights, repeated across many +// blocks and seeds. Catches accidental position-dependent bugs that the +// deterministic pattern above might mask. +// +TEST(MlasSq2BitTest, PackUnpackRoundTrip_BlkLen64_Randomized) +{ + constexpr size_t kBlockCount = 17; // arbitrary, > 1, not a SIMD-friendly number + constexpr unsigned kSeeds = 8; + + for (unsigned seed = 0; seed < kSeeds; ++seed) { + std::mt19937 rng(seed * 7919u + 1u); + std::uniform_int_distribution dist(0u, 3u); + + std::vector weights(kBlockCount * sq2::kBlkLen); + for (auto& w : weights) { + w = static_cast(dist(rng)); + } + + std::vector src(kBlockCount * sq2::kBlkBytes); + for (size_t blk = 0; blk < kBlockCount; ++blk) { + PackSourceBlock_BlkLen64(weights.data() + blk * sq2::kBlkLen, + src.data() + blk * sq2::kBlkBytes); + } + + std::vector packed(kBlockCount * sq2::kPackedBlkBytes); + for (size_t blk = 0; blk < kBlockCount; ++blk) { + sq2::PackBlock_BlkLen64(src.data() + blk * sq2::kBlkBytes, + packed.data() + blk * sq2::kPackedBlkBytes); + } + + for (size_t blk = 0; blk < kBlockCount; ++blk) { + std::array recovered{}; + sq2::UnpackBlock_BlkLen64_Reference(packed.data() + blk * sq2::kPackedBlkBytes, + recovered.data()); + for (size_t i = 0; i < sq2::kBlkLen; ++i) { + ASSERT_EQ(recovered[i], weights[blk * sq2::kBlkLen + i]) + << "Random round-trip mismatch seed=" << seed + << " blk=" << blk + << " i=" << i; + } + } + } +} + +// +// Bit-position invariant: writing only value v into every weight slot must +// yield a packed buffer where every byte is 0x55 * v (== v repeated at +// positions 0,2,4,6). This catches confusion between low/high nibbles or +// reversed-bit packing. +// +TEST(MlasSq2BitTest, PackUnpackRoundTrip_BlkLen64_ConstantValues) +{ + for (uint8_t v = 0; v < 4; ++v) { + std::array weights{}; + weights.fill(v); + + std::array src{}; + PackSourceBlock_BlkLen64(weights.data(), src.data()); + + std::array packed{}; + sq2::PackBlock_BlkLen64(src.data(), packed.data()); + + const uint8_t expected_byte = static_cast(v * 0x55u); // v at bits {0..1,2..3,4..5,6..7} + for (size_t i = 0; i < sq2::kPackedBlkBytes; ++i) { + ASSERT_EQ(static_cast(packed[i]), expected_byte) + << "Constant-fill v=" << static_cast(v) + << " byte_i=" << i; + } + + std::array recovered{}; + sq2::UnpackBlock_BlkLen64_Reference(packed.data(), recovered.data()); + for (size_t i = 0; i < sq2::kBlkLen; ++i) { + ASSERT_EQ(recovered[i], v) << "Constant-fill v=" << static_cast(v) << " i=" << i; + } + } +} diff --git a/onnxruntime/test/mlas/unittest/test_sqnbitgemm_2bit_gemm.cpp b/onnxruntime/test/mlas/unittest/test_sqnbitgemm_2bit_gemm.cpp new file mode 100644 index 0000000000000..ea1c3c05c523c --- /dev/null +++ b/onnxruntime/test/mlas/unittest/test_sqnbitgemm_2bit_gemm.cpp @@ -0,0 +1,289 @@ +/*++ + +Copyright (c) Microsoft Corporation. All rights reserved. + +Licensed under the MIT License. + +Module Name: + + test_sqnbitgemm_2bit_gemm.cpp + +Abstract: + + Numerical correctness tests for the 2-bit weight CompInt8 GEMM path + (Phase 2b). Exercises the full MlasQNBitGemmBatch dispatch: + + MlasQNBitGemmPackQuantBDataSize -> Q2BitGemmPackQuantBDataSize_Avx512 + MlasQNBitGemmPackQuantBData -> SQ2BitGemmPackQuantBDataAndBlkSum_Scalar + MlasQNBitGemmBatch -> SQ2BitGemm_CompInt8 wrapper + -> InitializeWorkspace_CompInt8 + -> SQ2BitGemmKernel_BlkSum_CompInt8_Scalar + + Reference output is computed by reproducing the same per-block int8 + quantization that MLAS uses internally (amax/127 scale, symmetric), then + running an integer dot product against the raw 2-bit weights with the + implicit zero-point of 2. Because both paths use the identical A + quantization rule and dequant-free integer accumulation, results match + to a tight float tolerance. + +--*/ + +#include "gtest/gtest.h" + +#include +#include +#include +#include +#include +#include +#include + +#include "core/mlas/inc/mlas_qnbit.h" +#include "core/mlas/lib/sqnbitgemm_kernel_avx512_2bit.h" + +namespace { + +namespace sq2 = onnxruntime::mlas::sq2bit_avx512; + +constexpr size_t kBlkLen = sq2::kBlkLen; // 64 +constexpr size_t kBlkBytes = sq2::kBlkBytes; // 16 source bytes per block +constexpr size_t kBlkBitWidth = 2; +constexpr MLAS_QNBIT_GEMM_COMPUTE_TYPE kComputeType = SQNBIT_CompInt8; + +// Standard ONNX 2-bit packing: byte_i = w[4i] | w[4i+1]<<2 | w[4i+2]<<4 | w[4i+3]<<6. +inline void +PackSourceBlock_BlkLen64(const uint8_t weights[kBlkLen], std::byte* src_out) +{ + for (size_t i = 0; i < kBlkBytes; ++i) { + const uint8_t v0 = weights[4 * i + 0] & 0x03u; + const uint8_t v1 = weights[4 * i + 1] & 0x03u; + const uint8_t v2 = weights[4 * i + 2] & 0x03u; + const uint8_t v3 = weights[4 * i + 3] & 0x03u; + src_out[i] = static_cast( + static_cast(v0 | (v1 << 2) | (v2 << 4) | (v3 << 6)) + ); + } +} + +// +// Mirror the MLAS per-row int8 block quantizer used by the CompInt8 path: +// per-block symmetric scale = amax / 127, round-to-nearest, clamp to [-127, 127]. +// +void +QuantizeA_Reference(size_t M, + size_t K, + const float* A, + int8_t* QuantAData, + float* QuantAScale) +{ + const size_t BlockCountK = (K + kBlkLen - 1) / kBlkLen; + for (size_t m = 0; m < M; ++m) { + for (size_t k = 0, k_blk = 0; k < K; k += kBlkLen, ++k_blk) { + const size_t local_len = std::min(K - k, kBlkLen); + + float amax = 0.0f; + for (size_t kk = 0; kk < local_len; ++kk) { + amax = std::max(amax, std::fabs(A[m * K + k + kk])); + } + + constexpr float range_max = static_cast((1 << 7) - 1); + const float scale = amax / range_max; + const float scale_recip = scale != 0.0f ? 1.0f / scale : 0.0f; + + QuantAScale[m * BlockCountK + k_blk] = scale; + + for (size_t kk = 0; kk < kBlkLen; ++kk) { + const float a = (kk < local_len) ? A[m * K + k + kk] : 0.0f; + const float q = std::round(a * scale_recip); + QuantAData[m * BlockCountK * kBlkLen + k + kk] = + static_cast( + std::clamp(q, + static_cast(std::numeric_limits::min()), + static_cast(std::numeric_limits::max()))); + } + } + } +} + +// +// Reference GEMM that exactly mirrors the math performed by the MLAS W2 +// CompInt8 path: +// +// C[m,n] = bias[n] +// + sum_blk( scale_a[m,blk] * scale_b[n,blk] +// * dot(qa[m,blk,:], (qb[n,blk,:] - 2)) ) +// +// (Equivalent to the kernel's "dot with raw uint8 weights" + the BlkSum +// correction term, just written without the algebraic split.) +// +void +ReferenceGemm_W2_CompInt8(size_t M, + size_t N, + size_t K, + const float* A, + const std::vector& BWeights, // [N * K] in [0,3] + const float* QuantBScale, + const float* Bias, + float* C) +{ + const size_t BlockCountK = (K + kBlkLen - 1) / kBlkLen; + + std::vector QuantAData(M * BlockCountK * kBlkLen, int8_t{0}); + std::vector QuantAScale(M * BlockCountK, 0.0f); + QuantizeA_Reference(M, K, A, QuantAData.data(), QuantAScale.data()); + + for (size_t m = 0; m < M; ++m) { + for (size_t n = 0; n < N; ++n) { + float acc = (Bias != nullptr) ? Bias[n] : 0.0f; + for (size_t k = 0, blk = 0; k < K; k += kBlkLen, ++blk) { + const size_t local_len = std::min(K - k, kBlkLen); + const float a_scale = QuantAScale[m * BlockCountK + blk]; + const float b_scale = QuantBScale[n * BlockCountK + blk]; + + int32_t dot = 0; + for (size_t kk = 0; kk < local_len; ++kk) { + const int8_t qa = QuantAData[m * BlockCountK * kBlkLen + k + kk]; + const int32_t qb = static_cast(BWeights[n * K + k + kk]) + - static_cast(sq2::kDefaultSymmetricZeroPoint2Bit); + dot += static_cast(qa) * qb; + } + acc += static_cast(dot) * a_scale * b_scale; + } + C[m * N + n] = acc; + } + } +} + +class MlasSQ2BitGemmTest { + public: + static void Run(size_t M, size_t N, size_t K, bool WithBias, uint32_t seed) + { + const size_t BlockCountK = (K + kBlkLen - 1) / kBlkLen; + ASSERT_EQ(K % kBlkLen, 0u) << "Test K must be a multiple of BlkLen=64"; + + std::mt19937 rng(seed); + std::uniform_real_distribution a_dist(-1.0f, 1.0f); + std::uniform_int_distribution w_dist(0, 3); + std::uniform_real_distribution s_dist(0.05f, 0.5f); + + std::vector A(M * K); + for (auto& v : A) v = a_dist(rng); + + // Raw weights in [0,3], natural [n, k] order; the test owns this oracle + // copy and is the source of truth for the reference math. + std::vector BWeights(N * K); + for (auto& v : BWeights) v = static_cast(w_dist(rng)); + + // Source-packed B in the layout that MlasQNBitGemmPackQuantBData consumes: + // column-major in N, kBlkBytes per block, standard ONNX 4-weights-per-byte. + std::vector QuantBData(N * BlockCountK * kBlkBytes, std::byte{0}); + for (size_t n = 0; n < N; ++n) { + for (size_t blk = 0; blk < BlockCountK; ++blk) { + uint8_t blk_weights[kBlkLen]; + for (size_t kk = 0; kk < kBlkLen; ++kk) { + blk_weights[kk] = BWeights[n * K + blk * kBlkLen + kk]; + } + PackSourceBlock_BlkLen64( + blk_weights, + QuantBData.data() + (n * BlockCountK + blk) * kBlkBytes); + } + } + + std::vector QuantBScale(N * BlockCountK); + for (auto& v : QuantBScale) v = s_dist(rng); + + std::vector Bias; + const float* BiasPtr = nullptr; + if (WithBias) { + Bias.resize(N); + for (auto& v : Bias) v = a_dist(rng); + BiasPtr = Bias.data(); + } + + // Pack B through the public API. + const size_t PackedSize = MlasQNBitGemmPackQuantBDataSize( + N, K, kBlkBitWidth, kBlkLen, /*has_zero_point=*/false, kComputeType, nullptr); + ASSERT_GT(PackedSize, 0u); + std::vector PackedQuantB(PackedSize, std::byte{0}); + + MlasQNBitGemmPackQuantBData( + N, K, kBlkBitWidth, kBlkLen, kComputeType, + QuantBData.data(), PackedQuantB.data(), + QuantBScale.data(), /*has_zp_input=*/false, /*QuantBZeroPoint=*/nullptr, + nullptr, nullptr); + + const size_t WorkspaceSize = MlasQNBitGemmBatchWorkspaceSize( + M, N, K, 1, kBlkBitWidth, kBlkLen, /*has_zero_point=*/false, kComputeType, nullptr); + std::vector Workspace(std::max(WorkspaceSize, 1), std::byte{0}); + + std::vector C(M * N, 0.0f); + + MLAS_QNBIT_GEMM_DATA_PARAMS params{}; + params.A = A.data(); + params.lda = K; + params.QuantBDataWorkspace = PackedQuantB.data(); + params.PackedQuantBData = PackedQuantB.data(); + params.QuantBScale = QuantBScale.data(); + params.QuantBZeroPoint = nullptr; + params.Bias = BiasPtr; + params.C = C.data(); + params.ldc = N; + params.PostProcessor = nullptr; + + MlasQNBitGemmBatch(M, N, K, 1, kBlkBitWidth, kBlkLen, kComputeType, + ¶ms, Workspace.data(), nullptr, nullptr); + + std::vector CRef(M * N, 0.0f); + ReferenceGemm_W2_CompInt8(M, N, K, A.data(), BWeights, QuantBScale.data(), + BiasPtr, CRef.data()); + + // Both paths perform the identical integer-domain dot product followed + // by the same float multiply-add chain, so the result should agree to + // a small relative tolerance driven only by float accumulation order. + const float abs_tol = 1e-4f; + const float rel_tol = 1e-4f; + for (size_t i = 0; i < M * N; ++i) { + const float diff = std::fabs(C[i] - CRef[i]); + const float bound = abs_tol + rel_tol * std::fabs(CRef[i]); + ASSERT_LE(diff, bound) + << "Mismatch at i=" << i + << " (m=" << (i / N) << ", n=" << (i % N) << ")" + << " MLAS=" << C[i] << " Ref=" << CRef[i] + << " M=" << M << " N=" << N << " K=" << K + << " WithBias=" << WithBias; + } + } +}; + +} // namespace + +// +// Single gtest case that walks a small grid of shapes. Gated on platforms +// where the W2/BlkLen=64/CompInt8 dispatch is wired up. +// +TEST(MlasSq2BitTest, GemmCompInt8_BlkLen64) +{ + if (!MlasIsQNBitGemmAvailable(kBlkBitWidth, kBlkLen, kComputeType)) { + GTEST_SKIP() << "MlasQNBitGemm W2/BlkLen=64/CompInt8 not available on this host"; + } + + struct Shape { size_t M, N, K; }; + constexpr Shape shapes[] = { + {1, 16, 64}, + {1, 32, 128}, + {1, 64, 256}, + {4, 16, 64}, + {4, 33, 192}, + {7, 17, 128}, + {16, 64, 512}, + {32, 128, 256}, + }; + + for (uint32_t seed : {0xC0FFEEu, 0xBADC0DEu}) { + for (const Shape& s : shapes) { + for (bool bias : {false, true}) { + MlasSQ2BitGemmTest::Run(s.M, s.N, s.K, bias, seed + (bias ? 1u : 0u)); + } + } + } +} From c3062b6ec9f1afd8312e110d9fde359b86cc768b Mon Sep 17 00:00:00 2001 From: Hari Seshadri Date: Fri, 5 Jun 2026 11:07:17 -0700 Subject: [PATCH 02/17] AVX512 VNNI kernels --- .../lib/sqnbitgemm_kernel_avx512_2bit.cpp | 15 +- .../mlas/lib/sqnbitgemm_kernel_avx512_2bit.h | 77 ++++++++++ ...nbitgemm_kernel_avx512vnni_2bit_blklen64.h | 138 ++++++++++-------- 3 files changed, 167 insertions(+), 63 deletions(-) diff --git a/onnxruntime/core/mlas/lib/sqnbitgemm_kernel_avx512_2bit.cpp b/onnxruntime/core/mlas/lib/sqnbitgemm_kernel_avx512_2bit.cpp index e6f2778dffe7c..c939dd2c4f366 100644 --- a/onnxruntime/core/mlas/lib/sqnbitgemm_kernel_avx512_2bit.cpp +++ b/onnxruntime/core/mlas/lib/sqnbitgemm_kernel_avx512_2bit.cpp @@ -122,9 +122,13 @@ SQ2BitGemmPackQuantBDataAndBlkSum_Scalar( } const size_t BlockCountK = MlasDivRoundup(K, BlkLen); + const size_t NMain = (N / kNCols4) * kNCols4; const size_t Iterations = N * BlockCountK; - // Pack weight bytes (block-by-block, parallel over N * BlockCountK). + // Pack weight bytes in the 4-col-grouped + 2-K-block-paired layout for + // the main NMain cols; column-major for the tail N % 4 cols. See + // PackedQuantBOffsetBytes_W2 in sqnbitgemm_kernel_avx512_2bit.h for the + // exact mapping. if (QuantBDataBegin != nullptr) { std::byte* PackedQuantBData = PackedQuantB.PackedQuantBData; MlasTrySimpleParallel( @@ -132,13 +136,14 @@ SQ2BitGemmPackQuantBDataAndBlkSum_Scalar( [&](ptrdiff_t tid) { const size_t n = static_cast(tid) / BlockCountK; const size_t blk = static_cast(tid) % BlockCountK; - const size_t offset = (n * BlockCountK + blk) * kBlkBytes; - PackBlock_BlkLen64(QuantBDataBegin + offset, PackedQuantBData + offset); + const size_t src_offset = (n * BlockCountK + blk) * kBlkBytes; + const size_t dst_offset = PackedQuantBOffsetBytes_W2(n, blk, BlockCountK, NMain); + PackBlock_BlkLen64(QuantBDataBegin + src_offset, PackedQuantBData + dst_offset); } ); } - // Copy scales as-is (column-major) and compute BlkSum. + // Scales follow the same 4-col-grouped layout (see PackedQuantBScaleOffset_W2). // // BlkSum uses the W4-style "width-16 row-major chunked" layout because the // top-level kernel performs the zero-point correction via the float SGEMM @@ -161,7 +166,7 @@ SQ2BitGemmPackQuantBDataAndBlkSum_Scalar( const size_t n = static_cast(tid) / BlockCountK; const size_t blk = static_cast(tid) % BlockCountK; const float scale = QuantBScaleBegin[n * BlockCountK + blk]; - PackedScales[n * BlockCountK + blk] = scale; + PackedScales[PackedQuantBScaleOffset_W2(n, blk, BlockCountK, NMain)] = scale; const size_t blksum_offset = ((n / 16) * BlockCountK + blk) * 16 + (n % 16); BlkSum[blksum_offset] = -scale * static_cast(kDefaultSymmetricZeroPoint2Bit); } diff --git a/onnxruntime/core/mlas/lib/sqnbitgemm_kernel_avx512_2bit.h b/onnxruntime/core/mlas/lib/sqnbitgemm_kernel_avx512_2bit.h index 91c47fc13589a..e973df39f67c1 100644 --- a/onnxruntime/core/mlas/lib/sqnbitgemm_kernel_avx512_2bit.h +++ b/onnxruntime/core/mlas/lib/sqnbitgemm_kernel_avx512_2bit.h @@ -70,6 +70,83 @@ constexpr size_t kPackedBlkBytes = kBlkBytes; // packing is in-place // For 2-bit unsigned values in [0, 3], the symmetric mid-point is 2. constexpr uint8_t kDefaultSymmetricZeroPoint2Bit = 2; +// ----------------------------------------------------------------------------- +// Tile shape (must match the AVX-512-VNNI kernel in +// sqnbitgemm_kernel_avx512vnni_2bit_blklen64.h). The pack layout below is +// keyed off these values: the main NMain = floor(N / kNCols4) * kNCols4 cols +// are stored in a 4-col-grouped + 2-K-block-paired arrangement that lets the +// kernel's R2xC4 hot loop read each tile of B as a single contiguous stream, +// matching W4's pack layout in PackQuantB (sqnbitgemm_kernel_avx_common.h). +// ----------------------------------------------------------------------------- + +constexpr size_t kNCols4 = 4; +constexpr size_t kPerAccuBlk2 = 2; + +// +// Byte offset into the packed B-data buffer for a logical (n, blk) cell. +// +// Main region (n < NMain = floor(N/kNCols4)*kNCols4): +// 4-N-col group g = n / 4, col within group c = n % 4. +// K-block-pair p = blk / 2, block within pair = blk % 2. +// Pair slots run first, then a single-block trailing slot when +// BlockCountK is odd. +// +// Tail region (n >= NMain): plain column-major. The tail base lies +// exactly at NMain * BlockCountK * kBlkBytes, so the dispatcher's +// `multipleCols * ColStrideBytes` offset (where ColStrideBytes is +// BlockCountK * kBlkBytes) is unchanged from the column-major layout. +// +inline size_t +PackedQuantBOffsetBytes_W2(size_t n, size_t blk, size_t BlockCountK, size_t NMain) +{ + if (n < NMain) { + const size_t g = n / kNCols4; + const size_t c = n % kNCols4; + const size_t per_group_bytes = BlockCountK * kNCols4 * kBlkBytes; + const size_t pair_idx = blk / kPerAccuBlk2; + const size_t blk_in_pair = blk % kPerAccuBlk2; + const size_t full_pairs = BlockCountK / kPerAccuBlk2; + if (pair_idx < full_pairs) { + return g * per_group_bytes + + pair_idx * (kNCols4 * kPerAccuBlk2 * kBlkBytes) + + c * (kPerAccuBlk2 * kBlkBytes) + + blk_in_pair * kBlkBytes; + } + return g * per_group_bytes + + full_pairs * (kNCols4 * kPerAccuBlk2 * kBlkBytes) + + c * kBlkBytes; + } + return (n * BlockCountK + blk) * kBlkBytes; +} + +// +// Float offset into the packed B-scale buffer for a logical (n, blk) cell. +// Same grouping rule as the B-data, two scales per pair, one scale per +// single-block trailing slot. +// +inline size_t +PackedQuantBScaleOffset_W2(size_t n, size_t blk, size_t BlockCountK, size_t NMain) +{ + if (n < NMain) { + const size_t g = n / kNCols4; + const size_t c = n % kNCols4; + const size_t per_group_scales = BlockCountK * kNCols4; + const size_t pair_idx = blk / kPerAccuBlk2; + const size_t blk_in_pair = blk % kPerAccuBlk2; + const size_t full_pairs = BlockCountK / kPerAccuBlk2; + if (pair_idx < full_pairs) { + return g * per_group_scales + + pair_idx * (kNCols4 * kPerAccuBlk2) + + c * kPerAccuBlk2 + + blk_in_pair; + } + return g * per_group_scales + + full_pairs * (kNCols4 * kPerAccuBlk2) + + c; + } + return n * BlockCountK + blk; +} + // // Extract a single 2-bit weight from a standard ONNX MatMulNBits packed byte // stream. `src` is the start of one block (kBlkBytes bytes). `i` is the diff --git a/onnxruntime/core/mlas/lib/sqnbitgemm_kernel_avx512vnni_2bit_blklen64.h b/onnxruntime/core/mlas/lib/sqnbitgemm_kernel_avx512vnni_2bit_blklen64.h index 6178bbcbce423..8c76b4f8518f0 100644 --- a/onnxruntime/core/mlas/lib/sqnbitgemm_kernel_avx512vnni_2bit_blklen64.h +++ b/onnxruntime/core/mlas/lib/sqnbitgemm_kernel_avx512vnni_2bit_blklen64.h @@ -74,10 +74,9 @@ namespace onnxruntime { namespace mlas { namespace sq2bit_avx512 { -// Number of K-blocks the inner loop processes per iteration. Mirrors W4. -inline constexpr size_t kPerAccuBlk2 = 2; -// Outer tile shape. Mirrors W4 NCols4 / NRows2. -inline constexpr size_t kNCols4 = 4; +// kNCols4 (= 4) and kPerAccuBlk2 (= 2) come from the shared header so the +// pack layout and the kernel use the same constants. kNRows2 is the M-tile +// shape, kernel-only. inline constexpr size_t kNRows2 = 2; // @@ -276,9 +275,13 @@ accumulate_w2_blklen64_r1c1blk1_vnni( // // R2 x C4 tile. Main hot path for the customer model. // -// Layout assumptions: -// * QuantBData : column-major. Col n at offset n * BlockCountK * kBlkBytes. -// * QuantBScale: column-major. Col n at offset n * BlockCountK. +// Layout assumptions (matches the W4 layout produced by W2's pack function): +// * QuantBData : 4-N-col grouped, 2-K-block paired (W4-style). Within a +// group the inner col stride is kPerAccuBlk2 * kBlkBytes (32 B) for pair +// blocks and kBlkBytes (16 B) for the single-block trailing slot when +// BlockCountK is odd. Per-group total = BlockCountK * kNCols4 * kBlkBytes. +// * QuantBScale: same grouping; col stride kPerAccuBlk2 (2 floats) / +// 1 float in the single-block slot. Per-group total = BlockCountK * kNCols4. // * QuantA : row-major int8, BlockCountK * kBlkLen bytes per row. // * QuantAScale: row-major float, BlockCountK floats per row. // @@ -296,8 +299,16 @@ Q2Int8GemmR2xC4BlkLen64Avx512Vnni( size_t ldc) { const size_t lda = BlockCountK * kBlkLen; - const size_t ColStrideBytes = BlockCountK * kBlkBytes; - const size_t ColStrideScale = BlockCountK; + // Per-group strides for the 4-N-col grouped layout. + constexpr size_t PerColPairBytes = kPerAccuBlk2 * kBlkBytes; // 32 B per col in a K-pair slot + constexpr size_t PerColSingleBytes = kBlkBytes; // 16 B per col in a K-single slot + constexpr size_t PerColPairScale = kPerAccuBlk2; // 2 floats per col in a K-pair slot + constexpr size_t PerKPairAdvanceBytes = kNCols4 * PerColPairBytes; // 128 B per K-pair iteration + constexpr size_t PerKSingleAdvanceBytes = kNCols4 * PerColSingleBytes; + constexpr size_t PerKPairAdvanceScale = kNCols4 * PerColPairScale; // 8 floats per K-pair iteration + constexpr size_t PerKSingleAdvanceScale = kNCols4; + const size_t GroupStrideBytes = BlockCountK * kNCols4 * kBlkBytes; + const size_t GroupStrideScale = BlockCountK * kNCols4; assert(CountM % kNRows2 == 0); assert(CountN % kNCols4 == 0); @@ -329,33 +340,33 @@ Q2Int8GemmR2xC4BlkLen64Avx512Vnni( accumulate_w2_blklen64_r2c1blk2_vnni( av_00, av_01, av_10, av_11, - QuantBDataPtr + 0 * ColStrideBytes, + QuantBDataPtr + 0 * PerColPairBytes, QuantAScalePtr, QuantAScalePtr + BlockCountK, - QuantBScalePtr + 0 * ColStrideScale, + QuantBScalePtr + 0 * PerColPairScale, acc[0], acc[kNCols4 + 0]); accumulate_w2_blklen64_r2c1blk2_vnni( av_00, av_01, av_10, av_11, - QuantBDataPtr + 1 * ColStrideBytes, + QuantBDataPtr + 1 * PerColPairBytes, QuantAScalePtr, QuantAScalePtr + BlockCountK, - QuantBScalePtr + 1 * ColStrideScale, + QuantBScalePtr + 1 * PerColPairScale, acc[1], acc[kNCols4 + 1]); accumulate_w2_blklen64_r2c1blk2_vnni( av_00, av_01, av_10, av_11, - QuantBDataPtr + 2 * ColStrideBytes, + QuantBDataPtr + 2 * PerColPairBytes, QuantAScalePtr, QuantAScalePtr + BlockCountK, - QuantBScalePtr + 2 * ColStrideScale, + QuantBScalePtr + 2 * PerColPairScale, acc[2], acc[kNCols4 + 2]); accumulate_w2_blklen64_r2c1blk2_vnni( av_00, av_01, av_10, av_11, - QuantBDataPtr + 3 * ColStrideBytes, + QuantBDataPtr + 3 * PerColPairBytes, QuantAScalePtr, QuantAScalePtr + BlockCountK, - QuantBScalePtr + 3 * ColStrideScale, + QuantBScalePtr + 3 * PerColPairScale, acc[3], acc[kNCols4 + 3]); QuantAPtr += kBlkLen * kPerAccuBlk2; QuantAScalePtr += kPerAccuBlk2; - QuantBDataPtr += kPerAccuBlk2 * kBlkBytes; - QuantBScalePtr += kPerAccuBlk2; + QuantBDataPtr += PerKPairAdvanceBytes; + QuantBScalePtr += PerKPairAdvanceScale; } while (k_blks_remaining-- > 0) { @@ -364,33 +375,33 @@ Q2Int8GemmR2xC4BlkLen64Avx512Vnni( accumulate_w2_blklen64_r2c1blk1_vnni( av_00, av_10, - QuantBDataPtr + 0 * ColStrideBytes, + QuantBDataPtr + 0 * PerColSingleBytes, QuantAScalePtr, QuantAScalePtr + BlockCountK, - QuantBScalePtr + 0 * ColStrideScale, + QuantBScalePtr + 0, acc[0], acc[kNCols4 + 0]); accumulate_w2_blklen64_r2c1blk1_vnni( av_00, av_10, - QuantBDataPtr + 1 * ColStrideBytes, + QuantBDataPtr + 1 * PerColSingleBytes, QuantAScalePtr, QuantAScalePtr + BlockCountK, - QuantBScalePtr + 1 * ColStrideScale, + QuantBScalePtr + 1, acc[1], acc[kNCols4 + 1]); accumulate_w2_blklen64_r2c1blk1_vnni( av_00, av_10, - QuantBDataPtr + 2 * ColStrideBytes, + QuantBDataPtr + 2 * PerColSingleBytes, QuantAScalePtr, QuantAScalePtr + BlockCountK, - QuantBScalePtr + 2 * ColStrideScale, + QuantBScalePtr + 2, acc[2], acc[kNCols4 + 2]); accumulate_w2_blklen64_r2c1blk1_vnni( av_00, av_10, - QuantBDataPtr + 3 * ColStrideBytes, + QuantBDataPtr + 3 * PerColSingleBytes, QuantAScalePtr, QuantAScalePtr + BlockCountK, - QuantBScalePtr + 3 * ColStrideScale, + QuantBScalePtr + 3, acc[3], acc[kNCols4 + 3]); QuantAPtr += kBlkLen; QuantAScalePtr++; - QuantBDataPtr += kBlkBytes; - QuantBScalePtr++; + QuantBDataPtr += PerKSingleAdvanceBytes; + QuantBScalePtr += PerKSingleAdvanceScale; } SumPtr[0] = _mm512_reduce_add_ps(acc[0]); @@ -412,8 +423,8 @@ Q2Int8GemmR2xC4BlkLen64Avx512Vnni( SumPtr[ldc + 3] += BiasPtr[3]; } - QuantBDataColPtr += kNCols4 * ColStrideBytes; - QuantBScaleColPtr += kNCols4 * ColStrideScale; + QuantBDataColPtr += GroupStrideBytes; + QuantBScaleColPtr += GroupStrideScale; BiasPtr += BiasPtr != nullptr ? kNCols4 : 0; SumPtr += kNCols4; } @@ -421,7 +432,11 @@ Q2Int8GemmR2xC4BlkLen64Avx512Vnni( } // -// R2 x C1 tile (N-tail). +// R2 x C1 tile (N-tail). Operates on the column-major tail region of the +// packed B buffer (cols NMain..N-1), which the dispatcher addresses with +// the same `multipleCols * (BlockCountK * kBlkBytes)` offset as the old +// pure column-major layout did (the grouped main region was sized to fit +// in exactly that many bytes by construction). // MLAS_FORCEINLINE void Q2Int8GemmR2xC1BlkLen64Avx512Vnni( @@ -511,7 +526,7 @@ Q2Int8GemmR2xC1BlkLen64Avx512Vnni( } // -// R1 x C4 tile (M-tail). +// R1 x C4 tile (M-tail). Uses the same 4-N-col grouped layout as R2 x C4. // MLAS_FORCEINLINE void Q2Int8GemmR1xC4BlkLen64Avx512Vnni( @@ -527,8 +542,15 @@ Q2Int8GemmR1xC4BlkLen64Avx512Vnni( size_t ldc) { const size_t lda = BlockCountK * kBlkLen; - const size_t ColStrideBytes = BlockCountK * kBlkBytes; - const size_t ColStrideScale = BlockCountK; + constexpr size_t PerColPairBytes = kPerAccuBlk2 * kBlkBytes; + constexpr size_t PerColSingleBytes = kBlkBytes; + constexpr size_t PerColPairScale = kPerAccuBlk2; + constexpr size_t PerKPairAdvanceBytes = kNCols4 * PerColPairBytes; + constexpr size_t PerKSingleAdvanceBytes = kNCols4 * PerColSingleBytes; + constexpr size_t PerKPairAdvanceScale = kNCols4 * PerColPairScale; + constexpr size_t PerKSingleAdvanceScale = kNCols4; + const size_t GroupStrideBytes = BlockCountK * kNCols4 * kBlkBytes; + const size_t GroupStrideScale = BlockCountK * kNCols4; assert(CountN % kNCols4 == 0); @@ -557,47 +579,47 @@ Q2Int8GemmR1xC4BlkLen64Avx512Vnni( accumulate_w2_blklen64_r1c1blk2_vnni( av_0, av_1, - QuantBDataPtr + 0 * ColStrideBytes, - QuantAScalePtr, QuantBScalePtr + 0 * ColStrideScale, acc[0]); + QuantBDataPtr + 0 * PerColPairBytes, + QuantAScalePtr, QuantBScalePtr + 0 * PerColPairScale, acc[0]); accumulate_w2_blklen64_r1c1blk2_vnni( av_0, av_1, - QuantBDataPtr + 1 * ColStrideBytes, - QuantAScalePtr, QuantBScalePtr + 1 * ColStrideScale, acc[1]); + QuantBDataPtr + 1 * PerColPairBytes, + QuantAScalePtr, QuantBScalePtr + 1 * PerColPairScale, acc[1]); accumulate_w2_blklen64_r1c1blk2_vnni( av_0, av_1, - QuantBDataPtr + 2 * ColStrideBytes, - QuantAScalePtr, QuantBScalePtr + 2 * ColStrideScale, acc[2]); + QuantBDataPtr + 2 * PerColPairBytes, + QuantAScalePtr, QuantBScalePtr + 2 * PerColPairScale, acc[2]); accumulate_w2_blklen64_r1c1blk2_vnni( av_0, av_1, - QuantBDataPtr + 3 * ColStrideBytes, - QuantAScalePtr, QuantBScalePtr + 3 * ColStrideScale, acc[3]); + QuantBDataPtr + 3 * PerColPairBytes, + QuantAScalePtr, QuantBScalePtr + 3 * PerColPairScale, acc[3]); QuantAPtr += kBlkLen * kPerAccuBlk2; QuantAScalePtr += kPerAccuBlk2; - QuantBDataPtr += kPerAccuBlk2 * kBlkBytes; - QuantBScalePtr += kPerAccuBlk2; + QuantBDataPtr += PerKPairAdvanceBytes; + QuantBScalePtr += PerKPairAdvanceScale; } while (k_blks_remaining-- > 0) { const __m512i av = _mm512_loadu_si512(reinterpret_cast(QuantAPtr)); accumulate_w2_blklen64_r1c1blk1_vnni( - av, QuantBDataPtr + 0 * ColStrideBytes, - QuantAScalePtr, QuantBScalePtr + 0 * ColStrideScale, acc[0]); + av, QuantBDataPtr + 0 * PerColSingleBytes, + QuantAScalePtr, QuantBScalePtr + 0, acc[0]); accumulate_w2_blklen64_r1c1blk1_vnni( - av, QuantBDataPtr + 1 * ColStrideBytes, - QuantAScalePtr, QuantBScalePtr + 1 * ColStrideScale, acc[1]); + av, QuantBDataPtr + 1 * PerColSingleBytes, + QuantAScalePtr, QuantBScalePtr + 1, acc[1]); accumulate_w2_blklen64_r1c1blk1_vnni( - av, QuantBDataPtr + 2 * ColStrideBytes, - QuantAScalePtr, QuantBScalePtr + 2 * ColStrideScale, acc[2]); + av, QuantBDataPtr + 2 * PerColSingleBytes, + QuantAScalePtr, QuantBScalePtr + 2, acc[2]); accumulate_w2_blklen64_r1c1blk1_vnni( - av, QuantBDataPtr + 3 * ColStrideBytes, - QuantAScalePtr, QuantBScalePtr + 3 * ColStrideScale, acc[3]); + av, QuantBDataPtr + 3 * PerColSingleBytes, + QuantAScalePtr, QuantBScalePtr + 3, acc[3]); QuantAPtr += kBlkLen; QuantAScalePtr++; - QuantBDataPtr += kBlkBytes; - QuantBScalePtr++; + QuantBDataPtr += PerKSingleAdvanceBytes; + QuantBScalePtr += PerKSingleAdvanceScale; } SumPtr[0] = _mm512_reduce_add_ps(acc[0]); @@ -611,8 +633,8 @@ Q2Int8GemmR1xC4BlkLen64Avx512Vnni( SumPtr[3] += BiasPtr[3]; } - QuantBDataColPtr += kNCols4 * ColStrideBytes; - QuantBScaleColPtr += kNCols4 * ColStrideScale; + QuantBDataColPtr += GroupStrideBytes; + QuantBScaleColPtr += GroupStrideScale; BiasPtr += BiasPtr != nullptr ? kNCols4 : 0; SumPtr += kNCols4; } @@ -620,7 +642,7 @@ Q2Int8GemmR1xC4BlkLen64Avx512Vnni( } // -// R1 x C1 tile (corner). +// R1 x C1 tile (corner). Same column-major tail-region addressing as R2 x C1. // MLAS_FORCEINLINE void Q2Int8GemmR1xC1BlkLen64Avx512Vnni( From 140696004d457146678348acd7c17427a8d3810e Mon Sep 17 00:00:00 2001 From: Hari Seshadri Date: Fri, 5 Jun 2026 17:17:11 -0700 Subject: [PATCH 03/17] Stage 2 --- onnxruntime/core/mlas/lib/qnbitgemm.cpp | 5 +- .../mlas/lib/sqnbitgemm_kernel_avx512.cpp | 42 ++- .../lib/sqnbitgemm_kernel_avx512_2bit.cpp | 34 ++- .../mlas/lib/sqnbitgemm_kernel_avx512_2bit.h | 67 ++++- .../mlas/lib/sqnbitgemm_kernel_avx512vnni.cpp | 39 ++- ...nbitgemm_kernel_avx512vnni_2bit_blklen64.h | 266 +++++++++++++----- .../test/mlas/bench/bench_qnbitgemm.cpp | 15 +- .../unittest/test_sqnbitgemm_2bit_gemm.cpp | 262 +++++++++++++++-- 8 files changed, 599 insertions(+), 131 deletions(-) diff --git a/onnxruntime/core/mlas/lib/qnbitgemm.cpp b/onnxruntime/core/mlas/lib/qnbitgemm.cpp index e9f9ce34d3977..ce0824bc8b198 100644 --- a/onnxruntime/core/mlas/lib/qnbitgemm.cpp +++ b/onnxruntime/core/mlas/lib/qnbitgemm.cpp @@ -1250,9 +1250,8 @@ GetQNBitGemm(QNBitGemmVariant variant) case SQ8BitGemmVariant_CompInt8: return SQ8BitGemm_CompInt8; case SQ2BitGemmVariant_CompInt8: - // Phase 2b: scalar reference compute kernel registered by the - // AVX-512 / AVX-512-VNNI dispatch tables. Phase 3 swaps in the - // vectorized version at the dispatch slot only. + // W2 CompInt8 compute kernel registered by the AVX-512 / + // AVX-512-VNNI dispatch tables. return SQ2BitGemm_CompInt8; default: return nullptr; diff --git a/onnxruntime/core/mlas/lib/sqnbitgemm_kernel_avx512.cpp b/onnxruntime/core/mlas/lib/sqnbitgemm_kernel_avx512.cpp index add4784b55d01..ad5a257dd1306 100644 --- a/onnxruntime/core/mlas/lib/sqnbitgemm_kernel_avx512.cpp +++ b/onnxruntime/core/mlas/lib/sqnbitgemm_kernel_avx512.cpp @@ -26,6 +26,8 @@ Module Name: #include "sqnbitgemm_kernel_avx512_int8_blklen32.h" #include "sqnbitgemm_kernel_avx512_int8_blklen64.h" #include "sqnbitgemm_kernel_avx512_int8_blklen128.h" +#include "sqnbitgemm_kernel_avx512_2bit.h" +#include "sqnbitgemm_kernel_avx512vnni_2bit_blklen64.h" // // SQNBIT_CompFp32 kernel implementation. @@ -475,6 +477,37 @@ SQ8BitGemmPackQuantBDataAndBlkSum512( HasZeroPoint, QuantBZPBegin, PackedQuantB, ThreadPool); } +// +// Unit-test entry point for the AVX-512BW (non-VNNI) W2 kernel. Linked into +// the test binary so the test TU (compiled without AVX-512 flags) can invoke +// this kernel directly without going through the platform dispatcher. The +// caller must verify AVX-512BW is available on the host before calling. +// +namespace onnxruntime::mlas::sq2bit_avx512 { +size_t MLASCALL +SQ2BitGemmKernel_BlkSum_CompInt8_Avx512_TestEntry( + size_t BlkLen, + const std::byte* QuantA, + const float* QuantAScale, + const std::byte* QuantBData, + const float* QuantBScale, + const std::byte* QuantBZeroPoint, + float* C, + size_t CountM, + size_t CountN, + size_t CountK, + size_t BlockCountK, + const float* Bias, + size_t ldc, + const float* ABlockSum, + const float* QuantBBlkSum) +{ + return SQ2BitGemmKernel_BlkSum_CompInt8_Avx512( + BlkLen, QuantA, QuantAScale, QuantBData, QuantBScale, QuantBZeroPoint, + C, CountM, CountN, CountK, BlockCountK, Bias, ldc, ABlockSum, QuantBBlkSum); +} +} // namespace onnxruntime::mlas::sq2bit_avx512 + const MLAS_QNBIT_GEMM_DISPATCH MlasSQNBitGemmDispatchAvx512 = []() { MLAS_QNBIT_GEMM_DISPATCH d; @@ -494,8 +527,13 @@ const MLAS_QNBIT_GEMM_DISPATCH MlasSQNBitGemmDispatchAvx512 = []() { d.SQ8BitGemmKernel_BlkSum_CompInt8 = SQ8BitGemmKernel_BlkSum_CompInt8_avx512; d.QuantizeARowComputeBlkSum_CompInt8 = QuantizeARow_CompInt8_avx512; - // 2-bit native CompInt8 path is registered in the AVX-512-VNNI dispatch only. - // Hosts with AVX-512 but no VNNI fall through to the existing LUT path. + // 2-bit native CompInt8 path: AVX-512BW variant (no VNNI). Uses the same + // tile + pack layout as the VNNI variant; the per-block MAC is + // `vpmaddubsw + vpmaddwd + vpaddd` instead of `_mm512_dpbusd_epi32`. + // Pack-size and pack functions are identical between AVX-512 and AVX-512-VNNI. + d.Q2BitGemmPackQuantBDataSize = onnxruntime::mlas::sq2bit_avx512::Q2BitGemmPackQuantBDataSize_Avx512; + d.SQ2BitGemmPackQuantBDataAndBlkSum = onnxruntime::mlas::sq2bit_avx512::SQ2BitGemmPackQuantBDataAndBlkSum_Scalar; + d.SQ2BitGemmKernel_BlkSum_CompInt8 = onnxruntime::mlas::sq2bit_avx512::SQ2BitGemmKernel_BlkSum_CompInt8_Avx512; return d; }(); diff --git a/onnxruntime/core/mlas/lib/sqnbitgemm_kernel_avx512_2bit.cpp b/onnxruntime/core/mlas/lib/sqnbitgemm_kernel_avx512_2bit.cpp index c939dd2c4f366..c45365f400a5b 100644 --- a/onnxruntime/core/mlas/lib/sqnbitgemm_kernel_avx512_2bit.cpp +++ b/onnxruntime/core/mlas/lib/sqnbitgemm_kernel_avx512_2bit.cpp @@ -10,19 +10,17 @@ Module Name: Abstract: - Phase 2b reference implementation of the 2-bit weight CompInt8 GEMM - (BlkBitWidth=2, BlkLen=64). This file contains only scalar C++ code; the - AVX-512-VNNI vectorized inner loop lands in Phase 3 as a separate header - (sqnbitgemm_kernel_avx512vnni_2bit_blklen64.h) that replaces the kernel - slot in the AVX-512-VNNI dispatch table. - - The scalar functions exposed here are linked into the AVX-512-VNNI - dispatch table only (the plain AVX-512 dispatch leaves the W2 slots - null so non-VNNI hosts fall through to the existing LUT kernel). They - are also reachable as plain C++ symbols from unit tests, which use them - as a correctness oracle for the vectorized path. - - Restrictions (Phase 2): + Reference implementation and pack-time helpers for the 2-bit weight + CompInt8 GEMM (BlkBitWidth=2, BlkLen=64). Scalar-only; the vectorized + inner loop lives in sqnbitgemm_kernel_avx512vnni_2bit_blklen64.h and is + registered into the AVX-512 and AVX-512-VNNI dispatch tables (with VNNI + / non-VNNI MAC variants templated from the same source). + + The scalar functions exposed here back the pack-time helpers used by + both dispatch tables and are reachable as plain C++ symbols from unit + tests, which use them as a correctness oracle for the vectorized paths. + + Restrictions: * BlkLen == 64 only. Other BlkLens are rejected by the pack helper and the kernel returns 0 rows handled. * Symmetric quantization only (no per-block zero-point tensor; @@ -53,7 +51,7 @@ namespace sq2bit_avx512 { // [BlkSum (float)] roundup_16(N) * BlockCountK // [Scales (float)] N * BlockCountK // -// Alignment slack is added so that the AVX-512 dequant in Phase 3 can use +// Alignment slack is added so that the AVX-512 dequant can use // 64-byte aligned loads. // size_t MLASCALL @@ -66,7 +64,7 @@ Q2BitGemmPackQuantBDataSize_Avx512( const MLAS_BACKEND_KERNEL_SELECTOR_CONFIG* /* BackendKernelSelectorConfig */ ) { - // Phase 2 supports only BlkLen=64 and SQNBIT_CompInt8. Anything else returns + // Only BlkLen=64 and SQNBIT_CompInt8 are supported. Anything else returns // 0 so MlasQNBitGemmPackQuantBDataSize reports an unsupported configuration. if (BlkLen != kBlkLen || ComputeType != SQNBIT_CompInt8) { return 0; @@ -98,7 +96,7 @@ Q2BitGemmPackQuantBDataSize_Avx512( // = -scale * 2 (symmetric W2 uses an implicit zero point of 2). // // QuantBZPBegin / HasZeroPoint are accepted for ABI parity with the W4 path -// but ignored in Phase 2 because the customer model and the projection +// but ignored here because the customer model and the production callers // rely on the symmetric layout. // void MLASCALL @@ -178,7 +176,7 @@ SQ2BitGemmPackQuantBDataAndBlkSum_Scalar( // Scalar reference kernel for SQ2BitGemmVariant_CompInt8. // // Inputs match the SQ4BitGemmKernel_BlkSum_CompInt8_Fn typedef so the -// vectorized Phase 3 implementation can drop into the same dispatch slot. +// vectorized implementation can drop into the same dispatch slot. // // Math: // C[m, n] = bias[n] @@ -209,7 +207,7 @@ SQ2BitGemmKernel_BlkSum_CompInt8_Scalar( ) { if (BlkLen != kBlkLen) { - return 0; // Phase 2b only supports BlkLen=64. + return 0; // Only BlkLen=64 is supported. } const size_t lda = BlockCountK * kBlkLen; // bytes per A row (int8) diff --git a/onnxruntime/core/mlas/lib/sqnbitgemm_kernel_avx512_2bit.h b/onnxruntime/core/mlas/lib/sqnbitgemm_kernel_avx512_2bit.h index e973df39f67c1..c9a3a56b2bb57 100644 --- a/onnxruntime/core/mlas/lib/sqnbitgemm_kernel_avx512_2bit.h +++ b/onnxruntime/core/mlas/lib/sqnbitgemm_kernel_avx512_2bit.h @@ -11,19 +11,20 @@ Module Name: Abstract: Pack-time helpers and reference (scalar) routines for the 2-bit, BlkLen=64 - AVX-512-VNNI weight GEMM path (SQNBIT_CompInt8, BlkBitWidth=2). + AVX-512 weight GEMM path (SQNBIT_CompInt8, BlkBitWidth=2). The vectorized + kernels (AVX-512 and AVX-512-VNNI variants) live in + sqnbitgemm_kernel_avx512vnni_2bit_blklen64.h. - This header is currently scalar / header-only and contains the pieces that - the Phase 2a round-trip unit tests exercise: + This header is scalar / header-only and contains: - * The packed-block layout used by the (future) AVX-512 dequant inner loop. + * The packed-block layout used by the AVX-512 dequant inner loop. * A pack routine that converts standard ONNX MatMulNBits 2-bit input data into the packed layout. * A reference unpack routine (independent of the pack code) that materialises one packed block back into 64 individual int8 values. - The packed layout is designed so that the AVX-512BW dequant in Phase 3 is - one 128-bit broadcast plus a per-lane variable shift: + The packed layout is designed so that the AVX-512BW dequant is one + 128-bit broadcast plus a per-lane variable shift: __m128i p = _mm_loadu_si128(packed); // 16 bytes __m512i p4 = _mm512_broadcast_i32x4(p); // 4 lanes @@ -220,9 +221,7 @@ UnpackSourceBlock_BlkLen64_Reference(const std::byte* src, uint8_t out[kBlkLen]) } // -// Phase 2b reference / dispatch entry points. -// Defined in sqnbitgemm_kernel_avx512_2bit.cpp; registered into the -// AVX-512 / AVX-512-VNNI dispatch tables. +// Reference / dispatch entry points defined in sqnbitgemm_kernel_avx512_2bit.cpp. // size_t MLASCALL @@ -269,6 +268,56 @@ SQ2BitGemmKernel_BlkSum_CompInt8_Scalar( const float* QuantBBlkSum ); +// +// Unit-test forwarders for the two AVX-512 vectorized kernel variants. These +// are non-inline symbols defined in `sqnbitgemm_kernel_avx512vnni.cpp` (VNNI) +// and `sqnbitgemm_kernel_avx512.cpp` (non-VNNI), each of which is compiled +// with the appropriate ISA flags. The test TU (which is NOT compiled with +// AVX-512 flags) calls these by `extern` linkage to exercise both kernels +// independently of the platform dispatcher. +// +// Callers MUST gate on `GetMlasPlatform().Avx512Supported_` before invoking +// these symbols, since they execute AVX-512BW (and, for the VNNI variant, +// AVX-512-VNNI) instructions. +// +size_t MLASCALL +SQ2BitGemmKernel_BlkSum_CompInt8_Avx512_TestEntry( + size_t BlkLen, + const std::byte* QuantA, + const float* QuantAScale, + const std::byte* QuantBData, + const float* QuantBScale, + const std::byte* QuantBZeroPoint, + float* C, + size_t CountM, + size_t CountN, + size_t CountK, + size_t BlockCountK, + const float* Bias, + size_t ldc, + const float* ABlockSum, + const float* QuantBBlkSum +); + +size_t MLASCALL +SQ2BitGemmKernel_BlkSum_CompInt8_Avx512Vnni_TestEntry( + size_t BlkLen, + const std::byte* QuantA, + const float* QuantAScale, + const std::byte* QuantBData, + const float* QuantBScale, + const std::byte* QuantBZeroPoint, + float* C, + size_t CountM, + size_t CountN, + size_t CountK, + size_t BlockCountK, + const float* Bias, + size_t ldc, + const float* ABlockSum, + const float* QuantBBlkSum +); + } // namespace sq2bit_avx512 } // namespace mlas } // namespace onnxruntime diff --git a/onnxruntime/core/mlas/lib/sqnbitgemm_kernel_avx512vnni.cpp b/onnxruntime/core/mlas/lib/sqnbitgemm_kernel_avx512vnni.cpp index c94638a8fd5c9..4d7acd9db9a47 100644 --- a/onnxruntime/core/mlas/lib/sqnbitgemm_kernel_avx512vnni.cpp +++ b/onnxruntime/core/mlas/lib/sqnbitgemm_kernel_avx512vnni.cpp @@ -459,6 +459,37 @@ SQ8BitGemmPackQuantBDataAndBlkSum512vnni( HasZeroPoint, QuantBZPBegin, PackedQuantB, ThreadPool); } +// +// Unit-test entry point for the AVX-512-VNNI W2 kernel. Mirrors the +// AVX-512BW (non-VNNI) entry point in sqnbitgemm_kernel_avx512.cpp; see the +// comment there for the rationale. The caller must verify AVX-512-VNNI is +// available on the host before calling. +// +namespace onnxruntime::mlas::sq2bit_avx512 { +size_t MLASCALL +SQ2BitGemmKernel_BlkSum_CompInt8_Avx512Vnni_TestEntry( + size_t BlkLen, + const std::byte* QuantA, + const float* QuantAScale, + const std::byte* QuantBData, + const float* QuantBScale, + const std::byte* QuantBZeroPoint, + float* C, + size_t CountM, + size_t CountN, + size_t CountK, + size_t BlockCountK, + const float* Bias, + size_t ldc, + const float* ABlockSum, + const float* QuantBBlkSum) +{ + return SQ2BitGemmKernel_BlkSum_CompInt8_Avx512Vnni( + BlkLen, QuantA, QuantAScale, QuantBData, QuantBScale, QuantBZeroPoint, + C, CountM, CountN, CountK, BlockCountK, Bias, ldc, ABlockSum, QuantBBlkSum); +} +} // namespace onnxruntime::mlas::sq2bit_avx512 + // // Kernel dispatch structure definition. // @@ -481,10 +512,10 @@ const MLAS_QNBIT_GEMM_DISPATCH MlasSQNBitGemmDispatchAvx512vnni = []() { d.SQ8BitGemmKernel_BlkSum_CompInt8 = SQ8BitGemmKernel_BlkSum_CompInt8_avx512vnni; d.QuantizeARowComputeBlkSum_CompInt8 = QuantizeARow_CompInt8_avx512; - // 2-bit native CompInt8 path. Phase 2b: scalar reference. Phase 3: this slot - // gets swapped for the AVX-512-VNNI vectorized kernel (`_mm512_dpbusd_epi32`). - // Plain AVX-512 (non-VNNI) hosts intentionally do not register this path and - // fall through to the LUT kernel. + // 2-bit native CompInt8 path: AVX-512-VNNI variant. Uses the same tile + + // pack layout as the non-VNNI variant (registered in the AVX-512 dispatch); + // the per-block MAC is `_mm512_dpbusd_epi32` instead of the + // `vpmaddubsw + vpmaddwd + vpaddd` chain. d.Q2BitGemmPackQuantBDataSize = onnxruntime::mlas::sq2bit_avx512::Q2BitGemmPackQuantBDataSize_Avx512; d.SQ2BitGemmPackQuantBDataAndBlkSum = onnxruntime::mlas::sq2bit_avx512::SQ2BitGemmPackQuantBDataAndBlkSum_Scalar; d.SQ2BitGemmKernel_BlkSum_CompInt8 = onnxruntime::mlas::sq2bit_avx512::SQ2BitGemmKernel_BlkSum_CompInt8_Avx512Vnni; diff --git a/onnxruntime/core/mlas/lib/sqnbitgemm_kernel_avx512vnni_2bit_blklen64.h b/onnxruntime/core/mlas/lib/sqnbitgemm_kernel_avx512vnni_2bit_blklen64.h index 8c76b4f8518f0..102f05bfacdae 100644 --- a/onnxruntime/core/mlas/lib/sqnbitgemm_kernel_avx512vnni_2bit_blklen64.h +++ b/onnxruntime/core/mlas/lib/sqnbitgemm_kernel_avx512vnni_2bit_blklen64.h @@ -10,11 +10,20 @@ Module Name: Abstract: - Phase 4 AVX-512-VNNI tiled kernel for the 2-bit weight CompInt8 GEMM - path (BlkBitWidth=2, BlkLen=64). Header-only; included exactly once by - sqnbitgemm_kernel_avx512vnni.cpp, which carries the required compile - flags (-mavx512vnni -mavx512bw -mavx512dq -mavx512vl -mavx512f on - GCC/Clang; MSVC generates the right instructions from intrinsics). + AVX-512 tiled kernel for the 2-bit weight CompInt8 GEMM path + (BlkBitWidth=2, BlkLen=64). Header-only; the kernel is templated on a + `bool kVnni` parameter so the same source supports both: + + * AVX-512-VNNI host: `kVnni == true` -> single `_mm512_dpbusd_epi32` + per block MAC. Included by sqnbitgemm_kernel_avx512vnni.cpp. + * AVX-512 (no VNNI): `kVnni == false` -> three-instruction MAC chain + `vpmaddubsw + vpmaddwd + vpaddd` (AVX-512BW only). Included by + sqnbitgemm_kernel_avx512.cpp. + + The file name still carries `vnni` for git-history continuity; the + kernel handles both ISA targets via the template parameter, mirroring + W4's `dot_accumulate_2blk` / `dot_accumulate_2blkvnni` split in + sqnbitgemm_kernel_avx512_int8_blklen64.h. Architecture mirrors W4's MlasQ4Int8GemmKernelBlkLen64Avx512 (in sqnbitgemm_kernel_avx512_int8_blklen64.h): @@ -28,6 +37,8 @@ Module Name: * Inner unroll is PerAccuBlk2 = 2 K-blocks per iteration. Pairs of dpbusd outputs are interleaved (unpacklo/hi + add) so a single FMA applies both blocks' scales in one shot. + * Packed-B layout is 4-N-col-grouped + 2-K-block-paired (matches W4), + so the 4 cols of a tile read as a contiguous 128-byte stream. * BlkLen=64 only. * Symmetric quantization only (QuantBZeroPoint must be null). * Tail tiles R2xC1, R1xC4, R1xC1 cover M % 2 != 0 / N % 4 != 0. @@ -35,8 +46,8 @@ Module Name: Zero-point correction is performed OUTSIDE the int8 kernel via the platform float SGEMM kernel (GetMlasPlatform().GemmFloatKernel), exactly as W4 does. This requires QuantBBlkSum to be in the W4 "width-16 row- - major chunked" layout, which is produced by SQ2BitGemmPackQuantBData - AndBlkSum_Scalar. + major chunked" layout, which is produced by + SQ2BitGemmPackQuantBDataAndBlkSum_Scalar. Dequant prologue (per 64-element block): @@ -46,16 +57,6 @@ Module Name: __m512i v = _mm512_srlv_epi32(p4, sh); __m512i b = _mm512_and_si512(v, _mm512_set1_epi8(0x03)); - DESIGN DEVIATION FROM W4 (worth re-revisiting in a follow-up if perf - still falls short): W4 packs weights in a 4-N-col-grouped layout so the - 4 cols of a tile's data lie consecutively in memory (one ColStride - advances within the same K-block-pair). W2 currently uses the simpler - column-major layout (each col is BlockCountK * 16 bytes apart in memory), - which the kernel addresses via a multi-stream stride pattern. Lower - spatial locality, larger working set; on streaming shapes the HW - prefetcher copes, but for compute-bound small shapes this likely costs - a few percent. - --*/ #pragma once @@ -128,17 +129,60 @@ load_unpack_2blk_w2(const std::byte* packed, __m512i& bv0_64_epi8, __m512i& bv1_ // // Lane-interleaved 2-K-block accumulator (single M-row, single N-col). -// Mirrors W4's dot_accumulate_2blkvnni; identical math for W2 because the -// dequanted block is already in the same uint8 [0,3] form W4 produces. +// Mirrors W4's dot_accumulate_2blk / dot_accumulate_2blkvnni split (in +// sqnbitgemm_kernel_avx512_int8_blklen64.h). Identical math for W2 because +// the dequanted block is already in the same uint8 [0,3] form W4 produces. // -// acc += sum_2blks( cvt(dpbusd(bv, av)) * scale_a * scale_b ) +// acc += sum_2blks( cvt(dot(bv, av)) * scale_a * scale_b ) +// +// Two variants: +// * dot_accumulate_2blk_w2: AVX-512BW only (vpmaddubsw + vpmaddwd + add). +// Reduces at epi16 granularity, then ONE vpmaddwd +// folds adjacent pairs to epi32 (one madd_epi16 +// saved per K-block-pair vs an epi32-interleave +// approach). +// * dot_accumulate_2blk_w2_vnni: VNNI variant using _mm512_dpbusd_epi32 +// for the inner MAC. // // `scale_a` and `scale_b` point to TWO consecutive floats (scales for blk0 // and blk1). The double-broadcast trick gives the 16-lane pattern // [s0,s1, s0,s1, s0,s1, s0,s1, s0,s1, s0,s1, s0,s1, s0,s1] -// which matches the post-unpacklo/hi+add lane layout (blk0 in even lanes, +// which matches the post-interleave lane layout (blk0 in even lanes, // blk1 in odd lanes). // +static MLAS_FORCEINLINE __m512i +ones_32_epi16_w2() +{ + const __m512i zeros = _mm512_setzero_si512(); + return _mm512_srli_epi16(_mm512_ternarylogic_epi64(zeros, zeros, zeros, 1), 15); +} + +static MLAS_FORCEINLINE void +dot_accumulate_2blk_w2( + const __m512i& av0_64_epi8, + const __m512i& av1_64_epi8, + const float* scale_a, + const __m512i& bv0_64_epi8, + const __m512i& bv1_64_epi8, + const __m512& scale_b_16_ps, + __m512& acc) +{ + const __m512i dot0_32_epi16 = _mm512_maddubs_epi16(bv0_64_epi8, av0_64_epi8); + const __m512i dot1_32_epi16 = _mm512_maddubs_epi16(bv1_64_epi8, av1_64_epi8); + + const __m512i t1 = _mm512_unpacklo_epi32(dot0_32_epi16, dot1_32_epi16); + const __m512i t2 = _mm512_unpackhi_epi32(dot0_32_epi16, dot1_32_epi16); + const __m512i sum_32_epi16 = _mm512_add_epi16(t1, t2); // [b0 b0 b1 b1 ...] in epi16 + const __m512i ones = ones_32_epi16_w2(); + const __m512i sum_16_epi32 = _mm512_madd_epi16(ones, sum_32_epi16); // [b0 b1 b0 b1 ...] in epi32 + const __m512 sum_16_ps = _mm512_cvtepi32_ps(sum_16_epi32); + + const __m256 scale_a_8_ps = _mm256_castpd_ps(_mm256_broadcast_sd(reinterpret_cast(scale_a))); + const __m512 scale_a_16_ps = _mm512_broadcast_f32x8(scale_a_8_ps); + + acc = _mm512_fmadd_ps(sum_16_ps, _mm512_mul_ps(scale_a_16_ps, scale_b_16_ps), acc); +} + static MLAS_FORCEINLINE void dot_accumulate_2blk_w2_vnni( const __m512i& av0_64_epi8, @@ -167,15 +211,23 @@ dot_accumulate_2blk_w2_vnni( // Single-K-block accumulator. Uses uniform 16-lane scale broadcast since // there's only one block's scale in play. // +template static MLAS_FORCEINLINE void -dot_accumulate_1blk_w2_vnni( +dot_accumulate_1blk_w2( const __m512i& av_64_epi8, const float* scale_a, const __m512i& bv_64_epi8, const __m512& scale_b_16_ps, __m512& acc) { - const __m512i dot_16_epi32 = _mm512_dpbusd_epi32(_mm512_setzero_si512(), bv_64_epi8, av_64_epi8); + __m512i dot_16_epi32; + if constexpr (kVnni) { + dot_16_epi32 = _mm512_dpbusd_epi32(_mm512_setzero_si512(), bv_64_epi8, av_64_epi8); + } else { + const __m512i ones = ones_32_epi16_w2(); + const __m512i dot_32_epi16 = _mm512_maddubs_epi16(bv_64_epi8, av_64_epi8); + dot_16_epi32 = _mm512_madd_epi16(dot_32_epi16, ones); + } const __m512 sum_16_ps = _mm512_cvtepi32_ps(dot_16_epi32); const __m128 scale_a_ps = _mm_broadcast_ss(scale_a); @@ -188,8 +240,9 @@ dot_accumulate_1blk_w2_vnni( // 2 M-rows x 1 N-col x 2 K-blocks accumulator. The 2-block B load is shared // across the 2 M-rows. // +template static MLAS_FORCEINLINE void -accumulate_w2_blklen64_r2c1blk2_vnni( +accumulate_w2_blklen64_r2c1blk2( const __m512i& av00_64_epi8, const __m512i& av01_64_epi8, const __m512i& av10_64_epi8, const __m512i& av11_64_epi8, const std::byte* QuantBDataPtr, @@ -205,15 +258,21 @@ accumulate_w2_blklen64_r2c1blk2_vnni( const __m256 scale_b_8_ps = _mm256_castpd_ps(_mm256_broadcast_sd(reinterpret_cast(scale_b))); const __m512 scale_b_16_ps = _mm512_broadcast_f32x8(scale_b_8_ps); - dot_accumulate_2blk_w2_vnni(av00_64_epi8, av01_64_epi8, scale_a0, bv0, bv1, scale_b_16_ps, acc0); - dot_accumulate_2blk_w2_vnni(av10_64_epi8, av11_64_epi8, scale_a1, bv0, bv1, scale_b_16_ps, acc1); + if constexpr (kVnni) { + dot_accumulate_2blk_w2_vnni(av00_64_epi8, av01_64_epi8, scale_a0, bv0, bv1, scale_b_16_ps, acc0); + dot_accumulate_2blk_w2_vnni(av10_64_epi8, av11_64_epi8, scale_a1, bv0, bv1, scale_b_16_ps, acc1); + } else { + dot_accumulate_2blk_w2(av00_64_epi8, av01_64_epi8, scale_a0, bv0, bv1, scale_b_16_ps, acc0); + dot_accumulate_2blk_w2(av10_64_epi8, av11_64_epi8, scale_a1, bv0, bv1, scale_b_16_ps, acc1); + } } // // 2 M-rows x 1 N-col x 1 K-block accumulator (K-tail). // +template static MLAS_FORCEINLINE void -accumulate_w2_blklen64_r2c1blk1_vnni( +accumulate_w2_blklen64_r2c1blk1( const __m512i& av0_64_epi8, const __m512i& av1_64_epi8, const std::byte* QuantBDataPtr, @@ -228,15 +287,16 @@ accumulate_w2_blklen64_r2c1blk1_vnni( const __m128 scale_b_ps = _mm_broadcast_ss(scale_b); const __m512 scale_b_16_ps = _mm512_broadcast_f32x2(scale_b_ps); - dot_accumulate_1blk_w2_vnni(av0_64_epi8, scale_a0, bv, scale_b_16_ps, acc0); - dot_accumulate_1blk_w2_vnni(av1_64_epi8, scale_a1, bv, scale_b_16_ps, acc1); + dot_accumulate_1blk_w2(av0_64_epi8, scale_a0, bv, scale_b_16_ps, acc0); + dot_accumulate_1blk_w2(av1_64_epi8, scale_a1, bv, scale_b_16_ps, acc1); } // // 1 M-row x 1 N-col x 2 K-blocks accumulator. // +template static MLAS_FORCEINLINE void -accumulate_w2_blklen64_r1c1blk2_vnni( +accumulate_w2_blklen64_r1c1blk2( const __m512i& av0_64_epi8, const __m512i& av1_64_epi8, const std::byte* QuantBDataPtr, @@ -250,14 +310,19 @@ accumulate_w2_blklen64_r1c1blk2_vnni( const __m256 scale_b_8_ps = _mm256_castpd_ps(_mm256_broadcast_sd(reinterpret_cast(scale_b))); const __m512 scale_b_16_ps = _mm512_broadcast_f32x8(scale_b_8_ps); - dot_accumulate_2blk_w2_vnni(av0_64_epi8, av1_64_epi8, scale_a, bv0, bv1, scale_b_16_ps, acc); + if constexpr (kVnni) { + dot_accumulate_2blk_w2_vnni(av0_64_epi8, av1_64_epi8, scale_a, bv0, bv1, scale_b_16_ps, acc); + } else { + dot_accumulate_2blk_w2(av0_64_epi8, av1_64_epi8, scale_a, bv0, bv1, scale_b_16_ps, acc); + } } // // 1 M-row x 1 N-col x 1 K-block accumulator. // +template static MLAS_FORCEINLINE void -accumulate_w2_blklen64_r1c1blk1_vnni( +accumulate_w2_blklen64_r1c1blk1( const __m512i& av_64_epi8, const std::byte* QuantBDataPtr, const float* scale_a, @@ -269,7 +334,7 @@ accumulate_w2_blklen64_r1c1blk1_vnni( const __m128 scale_b_ps = _mm_broadcast_ss(scale_b); const __m512 scale_b_16_ps = _mm512_broadcast_f32x2(scale_b_ps); - dot_accumulate_1blk_w2_vnni(av_64_epi8, scale_a, bv, scale_b_16_ps, acc); + dot_accumulate_1blk_w2(av_64_epi8, scale_a, bv, scale_b_16_ps, acc); } // @@ -285,8 +350,9 @@ accumulate_w2_blklen64_r1c1blk1_vnni( // * QuantA : row-major int8, BlockCountK * kBlkLen bytes per row. // * QuantAScale: row-major float, BlockCountK floats per row. // +template MLAS_FORCEINLINE void -Q2Int8GemmR2xC4BlkLen64Avx512Vnni( +Q2Int8GemmR2xC4BlkLen64Avx512( const std::byte* QuantA, const float* QuantAScale, const std::byte* QuantBData, @@ -338,25 +404,25 @@ Q2Int8GemmR2xC4BlkLen64Avx512Vnni( const __m512i av_10 = _mm512_loadu_si512(reinterpret_cast(QuantAPtr + lda)); const __m512i av_11 = _mm512_loadu_si512(reinterpret_cast(QuantAPtr + lda + kBlkLen)); - accumulate_w2_blklen64_r2c1blk2_vnni( + accumulate_w2_blklen64_r2c1blk2( av_00, av_01, av_10, av_11, QuantBDataPtr + 0 * PerColPairBytes, QuantAScalePtr, QuantAScalePtr + BlockCountK, QuantBScalePtr + 0 * PerColPairScale, acc[0], acc[kNCols4 + 0]); - accumulate_w2_blklen64_r2c1blk2_vnni( + accumulate_w2_blklen64_r2c1blk2( av_00, av_01, av_10, av_11, QuantBDataPtr + 1 * PerColPairBytes, QuantAScalePtr, QuantAScalePtr + BlockCountK, QuantBScalePtr + 1 * PerColPairScale, acc[1], acc[kNCols4 + 1]); - accumulate_w2_blklen64_r2c1blk2_vnni( + accumulate_w2_blklen64_r2c1blk2( av_00, av_01, av_10, av_11, QuantBDataPtr + 2 * PerColPairBytes, QuantAScalePtr, QuantAScalePtr + BlockCountK, QuantBScalePtr + 2 * PerColPairScale, acc[2], acc[kNCols4 + 2]); - accumulate_w2_blklen64_r2c1blk2_vnni( + accumulate_w2_blklen64_r2c1blk2( av_00, av_01, av_10, av_11, QuantBDataPtr + 3 * PerColPairBytes, QuantAScalePtr, QuantAScalePtr + BlockCountK, @@ -373,25 +439,25 @@ Q2Int8GemmR2xC4BlkLen64Avx512Vnni( const __m512i av_00 = _mm512_loadu_si512(reinterpret_cast(QuantAPtr)); const __m512i av_10 = _mm512_loadu_si512(reinterpret_cast(QuantAPtr + lda)); - accumulate_w2_blklen64_r2c1blk1_vnni( + accumulate_w2_blklen64_r2c1blk1( av_00, av_10, QuantBDataPtr + 0 * PerColSingleBytes, QuantAScalePtr, QuantAScalePtr + BlockCountK, QuantBScalePtr + 0, acc[0], acc[kNCols4 + 0]); - accumulate_w2_blklen64_r2c1blk1_vnni( + accumulate_w2_blklen64_r2c1blk1( av_00, av_10, QuantBDataPtr + 1 * PerColSingleBytes, QuantAScalePtr, QuantAScalePtr + BlockCountK, QuantBScalePtr + 1, acc[1], acc[kNCols4 + 1]); - accumulate_w2_blklen64_r2c1blk1_vnni( + accumulate_w2_blklen64_r2c1blk1( av_00, av_10, QuantBDataPtr + 2 * PerColSingleBytes, QuantAScalePtr, QuantAScalePtr + BlockCountK, QuantBScalePtr + 2, acc[2], acc[kNCols4 + 2]); - accumulate_w2_blklen64_r2c1blk1_vnni( + accumulate_w2_blklen64_r2c1blk1( av_00, av_10, QuantBDataPtr + 3 * PerColSingleBytes, QuantAScalePtr, QuantAScalePtr + BlockCountK, @@ -438,8 +504,9 @@ Q2Int8GemmR2xC4BlkLen64Avx512Vnni( // pure column-major layout did (the grouped main region was sized to fit // in exactly that many bytes by construction). // +template MLAS_FORCEINLINE void -Q2Int8GemmR2xC1BlkLen64Avx512Vnni( +Q2Int8GemmR2xC1BlkLen64Avx512( const std::byte* QuantA, const float* QuantAScale, const std::byte* QuantBData, @@ -480,7 +547,7 @@ Q2Int8GemmR2xC1BlkLen64Avx512Vnni( const __m512i av_10 = _mm512_loadu_si512(reinterpret_cast(QuantAPtr + lda)); const __m512i av_11 = _mm512_loadu_si512(reinterpret_cast(QuantAPtr + lda + kBlkLen)); - accumulate_w2_blklen64_r2c1blk2_vnni( + accumulate_w2_blklen64_r2c1blk2( av_00, av_01, av_10, av_11, QuantBDataPtr, QuantAScalePtr, QuantAScalePtr + BlockCountK, @@ -497,7 +564,7 @@ Q2Int8GemmR2xC1BlkLen64Avx512Vnni( const __m512i av_00 = _mm512_loadu_si512(reinterpret_cast(QuantAPtr)); const __m512i av_10 = _mm512_loadu_si512(reinterpret_cast(QuantAPtr + lda)); - accumulate_w2_blklen64_r2c1blk1_vnni( + accumulate_w2_blklen64_r2c1blk1( av_00, av_10, QuantBDataPtr, QuantAScalePtr, QuantAScalePtr + BlockCountK, @@ -528,8 +595,9 @@ Q2Int8GemmR2xC1BlkLen64Avx512Vnni( // // R1 x C4 tile (M-tail). Uses the same 4-N-col grouped layout as R2 x C4. // +template MLAS_FORCEINLINE void -Q2Int8GemmR1xC4BlkLen64Avx512Vnni( +Q2Int8GemmR1xC4BlkLen64Avx512( const std::byte* QuantA, const float* QuantAScale, const std::byte* QuantBData, @@ -577,19 +645,19 @@ Q2Int8GemmR1xC4BlkLen64Avx512Vnni( const __m512i av_0 = _mm512_loadu_si512(reinterpret_cast(QuantAPtr)); const __m512i av_1 = _mm512_loadu_si512(reinterpret_cast(QuantAPtr + kBlkLen)); - accumulate_w2_blklen64_r1c1blk2_vnni( + accumulate_w2_blklen64_r1c1blk2( av_0, av_1, QuantBDataPtr + 0 * PerColPairBytes, QuantAScalePtr, QuantBScalePtr + 0 * PerColPairScale, acc[0]); - accumulate_w2_blklen64_r1c1blk2_vnni( + accumulate_w2_blklen64_r1c1blk2( av_0, av_1, QuantBDataPtr + 1 * PerColPairBytes, QuantAScalePtr, QuantBScalePtr + 1 * PerColPairScale, acc[1]); - accumulate_w2_blklen64_r1c1blk2_vnni( + accumulate_w2_blklen64_r1c1blk2( av_0, av_1, QuantBDataPtr + 2 * PerColPairBytes, QuantAScalePtr, QuantBScalePtr + 2 * PerColPairScale, acc[2]); - accumulate_w2_blklen64_r1c1blk2_vnni( + accumulate_w2_blklen64_r1c1blk2( av_0, av_1, QuantBDataPtr + 3 * PerColPairBytes, QuantAScalePtr, QuantBScalePtr + 3 * PerColPairScale, acc[3]); @@ -603,16 +671,16 @@ Q2Int8GemmR1xC4BlkLen64Avx512Vnni( while (k_blks_remaining-- > 0) { const __m512i av = _mm512_loadu_si512(reinterpret_cast(QuantAPtr)); - accumulate_w2_blklen64_r1c1blk1_vnni( + accumulate_w2_blklen64_r1c1blk1( av, QuantBDataPtr + 0 * PerColSingleBytes, QuantAScalePtr, QuantBScalePtr + 0, acc[0]); - accumulate_w2_blklen64_r1c1blk1_vnni( + accumulate_w2_blklen64_r1c1blk1( av, QuantBDataPtr + 1 * PerColSingleBytes, QuantAScalePtr, QuantBScalePtr + 1, acc[1]); - accumulate_w2_blklen64_r1c1blk1_vnni( + accumulate_w2_blklen64_r1c1blk1( av, QuantBDataPtr + 2 * PerColSingleBytes, QuantAScalePtr, QuantBScalePtr + 2, acc[2]); - accumulate_w2_blklen64_r1c1blk1_vnni( + accumulate_w2_blklen64_r1c1blk1( av, QuantBDataPtr + 3 * PerColSingleBytes, QuantAScalePtr, QuantBScalePtr + 3, acc[3]); @@ -644,8 +712,9 @@ Q2Int8GemmR1xC4BlkLen64Avx512Vnni( // // R1 x C1 tile (corner). Same column-major tail-region addressing as R2 x C1. // +template MLAS_FORCEINLINE void -Q2Int8GemmR1xC1BlkLen64Avx512Vnni( +Q2Int8GemmR1xC1BlkLen64Avx512( const std::byte* QuantA, const float* QuantAScale, const std::byte* QuantBData, @@ -681,7 +750,7 @@ Q2Int8GemmR1xC1BlkLen64Avx512Vnni( const __m512i av_0 = _mm512_loadu_si512(reinterpret_cast(QuantAPtr)); const __m512i av_1 = _mm512_loadu_si512(reinterpret_cast(QuantAPtr + kBlkLen)); - accumulate_w2_blklen64_r1c1blk2_vnni( + accumulate_w2_blklen64_r1c1blk2( av_0, av_1, QuantBDataPtr, QuantAScalePtr, QuantBScalePtr, acc); QuantAPtr += kBlkLen * kPerAccuBlk2; @@ -693,7 +762,7 @@ Q2Int8GemmR1xC1BlkLen64Avx512Vnni( while (k_blks_remaining-- > 0) { const __m512i av = _mm512_loadu_si512(reinterpret_cast(QuantAPtr)); - accumulate_w2_blklen64_r1c1blk1_vnni( + accumulate_w2_blklen64_r1c1blk1( av, QuantBDataPtr, QuantAScalePtr, QuantBScalePtr, acc); QuantAPtr += kBlkLen; @@ -718,8 +787,9 @@ Q2Int8GemmR1xC1BlkLen64Avx512Vnni( // // Tile dispatcher. Mirrors W4's MlasQ4Int8GemmKernelBlkLen64Avx512. // +template MLAS_FORCEINLINE void -MlasQ2Int8GemmKernelBlkLen64Avx512Vnni( +MlasQ2Int8GemmKernelBlkLen64Avx512( const std::byte* QuantA, const float* QuantAScale, const std::byte* QuantBData, @@ -742,12 +812,12 @@ MlasQ2Int8GemmKernelBlkLen64Avx512Vnni( const size_t multipleCols = CountN - remainingCols; if (multipleRows > 0 && multipleCols > 0) { - Q2Int8GemmR2xC4BlkLen64Avx512Vnni( + Q2Int8GemmR2xC4BlkLen64Avx512( QuantA, QuantAScale, QuantBData, QuantBScale, C, multipleRows, multipleCols, BlockCountK, Bias, ldc); } if (remainingCols > 0 && multipleRows > 0) { - Q2Int8GemmR2xC1BlkLen64Avx512Vnni( + Q2Int8GemmR2xC1BlkLen64Avx512( QuantA, QuantAScale, QuantBData + multipleCols * ColStrideBytes, QuantBScale + multipleCols * ColStrideScale, @@ -756,7 +826,7 @@ MlasQ2Int8GemmKernelBlkLen64Avx512Vnni( Bias ? Bias + multipleCols : nullptr, ldc); } if (remainingRows > 0 && multipleCols > 0) { - Q2Int8GemmR1xC4BlkLen64Avx512Vnni( + Q2Int8GemmR1xC4BlkLen64Avx512( QuantA + multipleRows * lda, QuantAScale + multipleRows * lda_scale, QuantBData, QuantBScale, @@ -764,7 +834,7 @@ MlasQ2Int8GemmKernelBlkLen64Avx512Vnni( remainingRows, multipleCols, BlockCountK, Bias, ldc); } if (remainingRows > 0 && remainingCols > 0) { - Q2Int8GemmR1xC1BlkLen64Avx512Vnni( + Q2Int8GemmR1xC1BlkLen64Avx512( QuantA + multipleRows * lda, QuantAScale + multipleRows * lda_scale, QuantBData + multipleCols * ColStrideBytes, @@ -776,10 +846,14 @@ MlasQ2Int8GemmKernelBlkLen64Avx512Vnni( } // -// Top-level kernel registered into MlasSQNBitGemmDispatchAvx512vnni. +// Common dispatched-kernel body. Templated on ; the two top-level +// wrappers below instantiate it. // -// 1) Calls the tile dispatcher above, which computes -// C[m,n] = bias[n] + sum_blk(scale_a * scale_b * dpbusd(b, a)) +// Steps: +// 1) Calls the tile dispatcher, which computes +// C[m,n] = bias[n] + sum_blk(scale_a * scale_b * dot(b, a)) +// using either VNNI (`_mm512_dpbusd_epi32`) or AVX-512BW +// (`vpmaddubsw + vpmaddwd + vpaddd`) depending on the template parameter. // 2) Adds the symmetric zero-point correction // C[m,n] += sum_blk(ABlockSum[m,blk] * QuantBBlkSum[n,blk]) // via the platform float SGEMM micro-kernel. QuantBBlkSum is in the @@ -788,8 +862,9 @@ MlasQ2Int8GemmKernelBlkLen64Avx512Vnni( // operand. ZeroMode=false means SGEMM does `C += A @ B`, so the bias // and int8 contribution already in C are preserved. // -static MLAS_FORCEINLINE size_t MLASCALL -SQ2BitGemmKernel_BlkSum_CompInt8_Avx512Vnni( +template +static MLAS_FORCEINLINE size_t +SQ2BitGemmKernel_BlkSum_CompInt8_Impl( const size_t BlkLen, const std::byte* QuantA, const float* QuantAScale, @@ -810,7 +885,7 @@ SQ2BitGemmKernel_BlkSum_CompInt8_Avx512Vnni( return 0; } - MlasQ2Int8GemmKernelBlkLen64Avx512Vnni( + MlasQ2Int8GemmKernelBlkLen64Avx512( QuantA, QuantAScale, QuantBData, QuantBScale, C, CountM, CountN, BlockCountK, Bias, ldc); @@ -833,6 +908,61 @@ SQ2BitGemmKernel_BlkSum_CompInt8_Avx512Vnni( return CountM; } +// +// Top-level VNNI variant registered into MlasSQNBitGemmDispatchAvx512vnni. +// +static MLAS_FORCEINLINE size_t MLASCALL +SQ2BitGemmKernel_BlkSum_CompInt8_Avx512Vnni( + const size_t BlkLen, + const std::byte* QuantA, + const float* QuantAScale, + const std::byte* QuantBData, + const float* QuantBScale, + const std::byte* QuantBZeroPoint, + float* C, + size_t CountM, + size_t CountN, + size_t CountK, + size_t BlockCountK, + const float* Bias, + size_t ldc, + const float* ABlockSum, + const float* QuantBBlkSum) +{ + return SQ2BitGemmKernel_BlkSum_CompInt8_Impl( + BlkLen, QuantA, QuantAScale, QuantBData, QuantBScale, QuantBZeroPoint, + C, CountM, CountN, CountK, BlockCountK, Bias, ldc, ABlockSum, QuantBBlkSum); +} + +// +// Top-level non-VNNI variant registered into MlasSQNBitGemmDispatchAvx512. +// Uses the AVX-512BW MAC chain (`vpmaddubsw + vpmaddwd + vpaddd`) instead of +// `_mm512_dpbusd_epi32`. Same tile shapes, same pack layout, same numerical +// result. +// +static MLAS_FORCEINLINE size_t MLASCALL +SQ2BitGemmKernel_BlkSum_CompInt8_Avx512( + const size_t BlkLen, + const std::byte* QuantA, + const float* QuantAScale, + const std::byte* QuantBData, + const float* QuantBScale, + const std::byte* QuantBZeroPoint, + float* C, + size_t CountM, + size_t CountN, + size_t CountK, + size_t BlockCountK, + const float* Bias, + size_t ldc, + const float* ABlockSum, + const float* QuantBBlkSum) +{ + return SQ2BitGemmKernel_BlkSum_CompInt8_Impl( + BlkLen, QuantA, QuantAScale, QuantBData, QuantBScale, QuantBZeroPoint, + C, CountM, CountN, CountK, BlockCountK, Bias, ldc, ABlockSum, QuantBBlkSum); +} + } // namespace sq2bit_avx512 } // namespace mlas } // namespace onnxruntime diff --git a/onnxruntime/test/mlas/bench/bench_qnbitgemm.cpp b/onnxruntime/test/mlas/bench/bench_qnbitgemm.cpp index f068db7ba8227..67c8af5015872 100644 --- a/onnxruntime/test/mlas/bench/bench_qnbitgemm.cpp +++ b/onnxruntime/test/mlas/bench/bench_qnbitgemm.cpp @@ -167,26 +167,23 @@ static void QNBitGemmCustomerArgs(benchmark::internal::Benchmark* b) { BENCHMARK(QNBITGEMM)->Apply(QNBitGemmCustomerArgs)->UseRealTime(); -// 2-bit weight rows for the same customer shapes. Phase 2b is a scalar -// reference; Phase 3 swaps in the AVX-512 (+VNNI) vectorized kernel and -// these rows become the head-to-head LUT-vs-native comparison surface. -// W2 vectorized path supports only BlkLen=64 + SQNBIT_CompInt8; the -// SQNBIT_CompFp32 row exercises the LUT fallback. +// 2-bit weight rows for the customer shapes. Exercises the AVX-512 W2 native +// path (VNNI variant on AVX-512-VNNI hosts; non-VNNI variant on AVX-512BW +// hosts). W2 is registered only for SQNBIT_CompInt8 and BlkLen=64, so we +// emit just that one ComputeType. static void QNBit2BitCustomerArgs(benchmark::internal::Benchmark* b) { b->ArgNames({"BlkLen", "M", "N", "K", "Threads", "Symmetric", "HasBias", "ComputeType"}); const int64_t M = 128; const int64_t BlkLen = 64; const int64_t Threads = 8; - const int64_t Symmetric = 1; // W2 vectorized path is symmetric-only. + const int64_t Symmetric = 1; // W2 native path is symmetric-only. const int64_t HasBias = 1; for (auto kn : {std::pair{384, 1024}, std::pair{1024, 192}, std::pair{1024, 384}, std::pair{1024, 4096}, std::pair{4096, 1024}}) { - for (int64_t ct : {int64_t{SQNBIT_CompFp32}, int64_t{SQNBIT_CompInt8}}) { - b->Args({BlkLen, M, kn.second, kn.first, Threads, Symmetric, HasBias, ct}); - } + b->Args({BlkLen, M, kn.second, kn.first, Threads, Symmetric, HasBias, int64_t{SQNBIT_CompInt8}}); } } diff --git a/onnxruntime/test/mlas/unittest/test_sqnbitgemm_2bit_gemm.cpp b/onnxruntime/test/mlas/unittest/test_sqnbitgemm_2bit_gemm.cpp index ea1c3c05c523c..01e677d8142ba 100644 --- a/onnxruntime/test/mlas/unittest/test_sqnbitgemm_2bit_gemm.cpp +++ b/onnxruntime/test/mlas/unittest/test_sqnbitgemm_2bit_gemm.cpp @@ -10,21 +10,27 @@ Module Name: Abstract: - Numerical correctness tests for the 2-bit weight CompInt8 GEMM path - (Phase 2b). Exercises the full MlasQNBitGemmBatch dispatch: - - MlasQNBitGemmPackQuantBDataSize -> Q2BitGemmPackQuantBDataSize_Avx512 - MlasQNBitGemmPackQuantBData -> SQ2BitGemmPackQuantBDataAndBlkSum_Scalar - MlasQNBitGemmBatch -> SQ2BitGemm_CompInt8 wrapper - -> InitializeWorkspace_CompInt8 - -> SQ2BitGemmKernel_BlkSum_CompInt8_Scalar - - Reference output is computed by reproducing the same per-block int8 - quantization that MLAS uses internally (amax/127 scale, symmetric), then - running an integer dot product against the raw 2-bit weights with the - implicit zero-point of 2. Because both paths use the identical A - quantization rule and dequant-free integer accumulation, results match - to a tight float tolerance. + Numerical correctness tests for the 2-bit weight CompInt8 GEMM path on + AVX-512 hosts. Covers three execution routes: + + 1) MlasQNBitGemmBatch (public API) -> platform-selected dispatch -> + AVX-512-VNNI W2 kernel. This is what production callers hit on a + VNNI host. + + 2) AVX-512-VNNI W2 kernel via direct test-entry forwarder, bypassing + the platform dispatcher. Same kernel as (1); validates the + forwarder mechanism used by (3). + + 3) AVX-512BW (non-VNNI) W2 kernel via direct test-entry forwarder. + Validates the kernel a non-VNNI AVX-512 host would normally run. + On a VNNI host the platform dispatcher never picks this path, so + the direct-call route is the only way to exercise it. + + All three are compared against a single bit-exact integer-domain + reference (`ReferenceGemm_W2_CompInt8`) that reproduces the same per- + block int8 A quantization MLAS uses internally (amax/127, symmetric) + and the same dequant-free integer dot product against raw 2-bit + weights with an implicit zero-point of 2. --*/ @@ -40,6 +46,8 @@ Module Name: #include "core/mlas/inc/mlas_qnbit.h" #include "core/mlas/lib/sqnbitgemm_kernel_avx512_2bit.h" +#include "core/mlas/lib/qnbitgemm.h" // for PackedQuantBDataStruct (test direct-call path) +#include "core/mlas/lib/mlasi.h" // for GetMlasPlatform().Avx512Supported_ namespace { @@ -258,10 +266,10 @@ class MlasSQ2BitGemmTest { } // namespace // -// Single gtest case that walks a small grid of shapes. Gated on platforms -// where the W2/BlkLen=64/CompInt8 dispatch is wired up. +// Public-API correctness test. On a VNNI host the platform dispatcher +// resolves to the AVX-512-VNNI W2 kernel. Skips if no W2 path is available. // -TEST(MlasSq2BitTest, GemmCompInt8_BlkLen64) +TEST(MlasSq2BitTest, GemmCompInt8_BlkLen64_PublicApi) { if (!MlasIsQNBitGemmAvailable(kBlkBitWidth, kBlkLen, kComputeType)) { GTEST_SKIP() << "MlasQNBitGemm W2/BlkLen=64/CompInt8 not available on this host"; @@ -287,3 +295,221 @@ TEST(MlasSq2BitTest, GemmCompInt8_BlkLen64) } } } + +// +// Direct-call test harness for the AVX-512 W2 kernel variants. Calls the +// kernel through its non-inline test-entry forwarder (declared in +// sqnbitgemm_kernel_avx512_2bit.h), bypassing the platform dispatcher. +// This is the only way to exercise the non-VNNI kernel on a VNNI host. +// +// Setup re-uses the public pack API (MlasQNBitGemmPackQuantBData) and the +// reference int8 A-quantizer in this file (QuantizeA_Reference, which is +// bit-identical to QuantizeARow_CompInt8_avx512). Output is compared +// against ReferenceGemm_W2_CompInt8 within a tight tolerance. +// +class MlasSQ2BitGemmDirectCallTest { + public: + static void Run(size_t M, size_t N, size_t K, bool WithBias, uint32_t seed, + bool TestVnni) + { + const size_t BlockCountK = (K + kBlkLen - 1) / kBlkLen; + ASSERT_EQ(K % kBlkLen, 0u) << "Test K must be a multiple of BlkLen=64"; + + std::mt19937 rng(seed); + std::uniform_real_distribution a_dist(-1.0f, 1.0f); + std::uniform_int_distribution w_dist(0, 3); + std::uniform_real_distribution s_dist(0.05f, 0.5f); + + std::vector A(M * K); + for (auto& v : A) v = a_dist(rng); + + std::vector BWeights(N * K); + for (auto& v : BWeights) v = static_cast(w_dist(rng)); + + std::vector QuantBData(N * BlockCountK * kBlkBytes, std::byte{0}); + for (size_t n = 0; n < N; ++n) { + for (size_t blk = 0; blk < BlockCountK; ++blk) { + uint8_t blk_weights[kBlkLen]; + for (size_t kk = 0; kk < kBlkLen; ++kk) { + blk_weights[kk] = BWeights[n * K + blk * kBlkLen + kk]; + } + PackSourceBlock_BlkLen64( + blk_weights, + QuantBData.data() + (n * BlockCountK + blk) * kBlkBytes); + } + } + + std::vector QuantBScale(N * BlockCountK); + for (auto& v : QuantBScale) v = s_dist(rng); + + std::vector Bias; + const float* BiasPtr = nullptr; + if (WithBias) { + Bias.resize(N); + for (auto& v : Bias) v = a_dist(rng); + BiasPtr = Bias.data(); + } + + // Pack B through the public API (produces the same buffer for both + // VNNI and non-VNNI consumers; the kernel doesn't care which one). + const size_t PackedSize = MlasQNBitGemmPackQuantBDataSize( + N, K, kBlkBitWidth, kBlkLen, /*has_zero_point=*/false, kComputeType, nullptr); + ASSERT_GT(PackedSize, 0u); + std::vector PackedQuantBBuf(PackedSize, std::byte{0}); + + MlasQNBitGemmPackQuantBData( + N, K, kBlkBitWidth, kBlkLen, kComputeType, + QuantBData.data(), PackedQuantBBuf.data(), + QuantBScale.data(), /*has_zp_input=*/false, /*QuantBZeroPoint=*/nullptr, + nullptr, nullptr); + + // Reconstruct the packed-B view so we can pass the right sub-pointers + // to the kernel forwarder (PackedQuantBData, PackedQuantBScale, + // QuantBBlkSum). The struct is a layout overlay over the buffer. + PackedQuantBDataStruct packed_b( + PackedQuantBBuf.data(), N, BlockCountK, kBlkLen, /*QuantAUnsigned=*/false); + + // Quantize A (bit-identical to MLAS's QuantizeARow_CompInt8_avx512 for + // BlkLen=64: per-block symmetric scale = amax / 127). Also compute the + // scaled block sums the kernel's BlkSum correction needs. + std::vector QuantAData(M * BlockCountK * kBlkLen, int8_t{0}); + std::vector QuantAScale(M * BlockCountK, 0.0f); + QuantizeA_Reference(M, K, A.data(), QuantAData.data(), QuantAScale.data()); + + std::vector ABlockSum(M * BlockCountK, 0.0f); + for (size_t m = 0; m < M; ++m) { + for (size_t blk = 0; blk < BlockCountK; ++blk) { + int32_t sum = 0; + for (size_t kk = 0; kk < kBlkLen; ++kk) { + sum += static_cast( + QuantAData[m * BlockCountK * kBlkLen + blk * kBlkLen + kk]); + } + ABlockSum[m * BlockCountK + blk] = + QuantAScale[m * BlockCountK + blk] * static_cast(sum); + } + } + + // Run the kernel under test directly via the test-entry forwarder. + std::vector C(M * N, 0.0f); + if (TestVnni) { + onnxruntime::mlas::sq2bit_avx512::SQ2BitGemmKernel_BlkSum_CompInt8_Avx512Vnni_TestEntry( + kBlkLen, + reinterpret_cast(QuantAData.data()), + QuantAScale.data(), + packed_b.PackedQuantBData, + packed_b.PackedQuantBScale, + /*QuantBZeroPoint=*/nullptr, + C.data(), + M, N, /*CountK=*/K, BlockCountK, + BiasPtr, + /*ldc=*/N, + ABlockSum.data(), + packed_b.QuantBBlkSum); + } else { + onnxruntime::mlas::sq2bit_avx512::SQ2BitGemmKernel_BlkSum_CompInt8_Avx512_TestEntry( + kBlkLen, + reinterpret_cast(QuantAData.data()), + QuantAScale.data(), + packed_b.PackedQuantBData, + packed_b.PackedQuantBScale, + /*QuantBZeroPoint=*/nullptr, + C.data(), + M, N, /*CountK=*/K, BlockCountK, + BiasPtr, + /*ldc=*/N, + ABlockSum.data(), + packed_b.QuantBBlkSum); + } + + // Reference: bit-exact integer-domain math. + std::vector CRef(M * N, 0.0f); + ReferenceGemm_W2_CompInt8(M, N, K, A.data(), BWeights, QuantBScale.data(), + BiasPtr, CRef.data()); + + const float abs_tol = 1e-4f; + const float rel_tol = 1e-4f; + for (size_t i = 0; i < M * N; ++i) { + const float diff = std::fabs(C[i] - CRef[i]); + const float bound = abs_tol + rel_tol * std::fabs(CRef[i]); + ASSERT_LE(diff, bound) + << (TestVnni ? "VNNI" : "non-VNNI") << " direct-call mismatch at i=" << i + << " (m=" << (i / N) << ", n=" << (i % N) << ")" + << " MLAS=" << C[i] << " Ref=" << CRef[i] + << " M=" << M << " N=" << N << " K=" << K + << " WithBias=" << WithBias; + } + } +}; + +// +// Exercises the non-VNNI W2 kernel (vpmaddubsw + vpmaddwd + vpaddd MAC chain). +// Gated on AVX-512BW availability. Runs even on VNNI hosts where the platform +// dispatcher would never select this kernel naturally. +// +TEST(MlasSq2BitTest, GemmCompInt8_BlkLen64_Avx512) +{ + if (!GetMlasPlatform().Avx512Supported_) { + GTEST_SKIP() << "AVX-512BW (DQ/VL) not available on this host"; + } + + struct Shape { size_t M, N, K; }; + constexpr Shape shapes[] = { + {1, 16, 64}, + {1, 32, 128}, + {1, 64, 256}, + {4, 16, 64}, + {4, 33, 192}, + {7, 17, 128}, + {16, 64, 512}, + {32, 128, 256}, + }; + + for (uint32_t seed : {0xC0FFEEu, 0xBADC0DEu}) { + for (const Shape& s : shapes) { + for (bool bias : {false, true}) { + MlasSQ2BitGemmDirectCallTest::Run( + s.M, s.N, s.K, bias, seed + (bias ? 1u : 0u), + /*TestVnni=*/false); + } + } + } +} + +// +// Exercises the AVX-512-VNNI W2 kernel (_mm512_dpbusd_epi32 MAC) via the +// direct-call forwarder. On a VNNI host this is the same kernel that the +// public-API test above hits through the dispatcher; the explicit invocation +// here keeps the harness symmetric and validates the forwarder mechanism. +// +// Gating: requires the platform to have actually selected the VNNI dispatch +// table. MlasIsQNBitGemmAvailable alone is insufficient because the W2 path +// is now registered into both AVX-512 and AVX-512-VNNI dispatch tables. +// +TEST(MlasSq2BitTest, GemmCompInt8_BlkLen64_Avx512Vnni) +{ + if (GetMlasPlatform().QNBitGemmDispatch != &MlasSQNBitGemmDispatchAvx512vnni) { + GTEST_SKIP() << "AVX-512-VNNI not selected as the active dispatch on this host"; + } + + struct Shape { size_t M, N, K; }; + constexpr Shape shapes[] = { + {1, 16, 64}, + {1, 32, 128}, + {1, 64, 256}, + {4, 16, 64}, + {4, 33, 192}, + {7, 17, 128}, + {16, 64, 512}, + {32, 128, 256}, + }; + + for (uint32_t seed : {0xC0FFEEu, 0xBADC0DEu}) { + for (const Shape& s : shapes) { + for (bool bias : {false, true}) { + MlasSQ2BitGemmDirectCallTest::Run( + s.M, s.N, s.K, bias, seed + (bias ? 1u : 0u), + /*TestVnni=*/true); + } + } + } +} From 5b26abf47e06a73aefb88fe7dfdf6bd62f02e685 Mon Sep 17 00:00:00 2001 From: Hari Seshadri Date: Sat, 6 Jun 2026 19:07:13 -0700 Subject: [PATCH 04/17] Stage 3 --- .../cpu/quantization/matmul_nbits.cc | 15 +- .../lib/sqnbitgemm_kernel_avx512_2bit.cpp | 74 +++++- ...nbitgemm_kernel_avx512vnni_2bit_blklen64.h | 48 +++- .../test/mlas/bench/bench_qnbitgemm.cpp | 44 ++-- .../unittest/test_sqnbitgemm_2bit_gemm.cpp | 241 ++++++++++++++++-- 5 files changed, 358 insertions(+), 64 deletions(-) diff --git a/onnxruntime/contrib_ops/cpu/quantization/matmul_nbits.cc b/onnxruntime/contrib_ops/cpu/quantization/matmul_nbits.cc index 2c83b5aff3594..457d2c7c3af18 100644 --- a/onnxruntime/contrib_ops/cpu/quantization/matmul_nbits.cc +++ b/onnxruntime/contrib_ops/cpu/quantization/matmul_nbits.cc @@ -1074,7 +1074,18 @@ Status MatMulNBits::ComputeBUnpacked(const Tensor* a, auto tmp_b_data_ptr = IAllocator::MakeUniquePtr(allocator, SafeInt(K_) * N_, true); if ((reorder_idx_data == nullptr) && (!zero_points || !zero_points->IsDataType())) { - if (nbits_ == 4) { + if (nbits_ == 2) { + MlasDequantizeBlockwise( + tmp_b_data_ptr.get(), // dequantized output + b_data, // quantized input + scales_ptr, // quantization scales + static_cast(zero_points_data), // quantization zero points + static_cast(block_size_), // quantization block size + column_wise_quant_, // columnwise quantization or row-wise + static_cast(K_), // number of rows in quantized input + static_cast(N_), // number of columns in quantized input + thread_pool); + } else if (nbits_ == 4) { MlasDequantizeBlockwise( tmp_b_data_ptr.get(), // dequantized output b_data, // quantized input @@ -1085,7 +1096,7 @@ Status MatMulNBits::ComputeBUnpacked(const Tensor* a, static_cast(K_), // number of rows in quantized input static_cast(N_), // number of columns in quantized input thread_pool); - } else { // If it isn't 4bit, it has to be 8-bit quantization + } else { // If it isn't 2bit or 4bit, it has to be 8-bit quantization ORT_ENFORCE(nbits_ == 8); MlasDequantizeBlockwise( tmp_b_data_ptr.get(), // dequantized output diff --git a/onnxruntime/core/mlas/lib/sqnbitgemm_kernel_avx512_2bit.cpp b/onnxruntime/core/mlas/lib/sqnbitgemm_kernel_avx512_2bit.cpp index c45365f400a5b..09cea17e43685 100644 --- a/onnxruntime/core/mlas/lib/sqnbitgemm_kernel_avx512_2bit.cpp +++ b/onnxruntime/core/mlas/lib/sqnbitgemm_kernel_avx512_2bit.cpp @@ -23,8 +23,9 @@ Module Name: Restrictions: * BlkLen == 64 only. Other BlkLens are rejected by the pack helper and the kernel returns 0 rows handled. - * Symmetric quantization only (no per-block zero-point tensor; - an implicit zero-point of 2 is used to recentre values in [0, 3]). + * Per-block zero-point input is supported (standard ONNX W2 layout, + 4 ZPs per packed byte along K). When no zero-point tensor is + supplied the symmetric default of 2 is used. --*/ @@ -95,9 +96,16 @@ Q2BitGemmPackQuantBDataSize_Avx512( // QuantBBlkSum[n * BlockCountK + blk] // = -scale * 2 (symmetric W2 uses an implicit zero point of 2). // -// QuantBZPBegin / HasZeroPoint are accepted for ABI parity with the W4 path -// but ignored here because the customer model and the production callers -// rely on the symmetric layout. +// When QuantBZPBegin is non-null, the per-block zero-point byte stream is in +// the standard ONNX MatMulNBits W2 layout: 4 zero-points per byte, packed +// along the K-block axis, row-major in N. Row stride is +// ZPCountK = ceil(BlockCountK / 4) bytes; the ZP for (n, blk) lives at byte +// index (n * ZPCountK + blk / 4), at bit offset (blk % 4) * 2. +// +// When QuantBZPBegin is null, we fall back to the symmetric default +// (kDefaultSymmetricZeroPoint2Bit = 2), preserving the prior behavior. The +// HasZeroPoint flag is informational only: if QuantBZPBegin is non-null we +// always consume it. // void MLASCALL SQ2BitGemmPackQuantBDataAndBlkSum_Scalar( @@ -108,7 +116,7 @@ SQ2BitGemmPackQuantBDataAndBlkSum_Scalar( const std::byte* QuantBDataBegin, const float* QuantBScaleBegin, bool /* HasZeroPoint */, - const std::byte* /* QuantBZPBegin */, + const std::byte* QuantBZPBegin, PackedQuantBDataStruct& PackedQuantB, MLAS_THREADPOOL* ThreadPool, const MLAS_BACKEND_KERNEL_SELECTOR_CONFIG* /* BackendKernelSelectorConfig */ @@ -149,15 +157,25 @@ SQ2BitGemmPackQuantBDataAndBlkSum_Scalar( // pre-packed B layout: // // BlkSum[(n / 16) * BlockCountK * 16 + blk * 16 + (n % 16)] - // = -scale_b * 2 (symmetric W2 uses an implicit ZP of 2) + // = -scale_b * zero_point // // The allocated BlkSum buffer is sized at MlasDivRoundup(N, 16) * BlockCountK // * 16 floats so the layout is well-defined even when N % 16 != 0 (the tail // chunk's unused lanes hold whatever the buffer was initialised with, which // for production callers must be zero so the SGEMM correction reads zeros). + // + // Important: ORT's matmul_nbits.cc prepack flow invokes this function up to + // three times per node — once each for B, scales, and zero_points — so on + // any given invocation only one of (QuantBScaleBegin, QuantBZPBegin) may be + // non-null. To get the correct BlkSum we therefore (a) copy scales into the + // packed buffer when QuantBScaleBegin is provided, and (b) recompute BlkSum + // whenever EITHER scales or zero-points are provided, reading the scales + // from the already-packed PackedQuantBScale buffer (which is populated by a + // previous call when only ZPs arrive in the current call). This mirrors + // the W4 ComputePackBlkSum helper. + if (QuantBScaleBegin != nullptr) { float* PackedScales = PackedQuantB.PackedQuantBScale; - float* BlkSum = PackedQuantB.QuantBBlkSum; MlasTrySimpleParallel( ThreadPool, static_cast(Iterations), [&](ptrdiff_t tid) { @@ -165,8 +183,35 @@ SQ2BitGemmPackQuantBDataAndBlkSum_Scalar( const size_t blk = static_cast(tid) % BlockCountK; const float scale = QuantBScaleBegin[n * BlockCountK + blk]; PackedScales[PackedQuantBScaleOffset_W2(n, blk, BlockCountK, NMain)] = scale; + } + ); + } + + // BlkSum needs to be (re)computed whenever scales or zero-points change. + // Source of scales is the already-packed buffer, so this works even when + // only zero_points are provided in the current invocation. + if (QuantBScaleBegin != nullptr || QuantBZPBegin != nullptr) { + const float* PackedScales = PackedQuantB.PackedQuantBScale; + float* BlkSum = PackedQuantB.QuantBBlkSum; + const size_t ZPCountK = MlasDivRoundup(BlockCountK, 4); + MlasTrySimpleParallel( + ThreadPool, static_cast(Iterations), + [&](ptrdiff_t tid) { + const size_t n = static_cast(tid) / BlockCountK; + const size_t blk = static_cast(tid) % BlockCountK; + const float scale = + PackedScales[PackedQuantBScaleOffset_W2(n, blk, BlockCountK, NMain)]; + + uint8_t zp = kDefaultSymmetricZeroPoint2Bit; + if (QuantBZPBegin != nullptr) { + const size_t zp_byte_idx = n * ZPCountK + (blk / 4); + const size_t zp_bit_off = (blk % 4) * 2; + zp = static_cast( + (static_cast(QuantBZPBegin[zp_byte_idx]) >> zp_bit_off) & 0x03u); + } + const size_t blksum_offset = ((n / 16) * BlockCountK + blk) * 16 + (n % 16); - BlkSum[blksum_offset] = -scale * static_cast(kDefaultSymmetricZeroPoint2Bit); + BlkSum[blksum_offset] = -scale * static_cast(zp); } ); } @@ -184,8 +229,10 @@ SQ2BitGemmPackQuantBDataAndBlkSum_Scalar( // * dot(int8 a[m, blk, :], uint8 b_unpacked[n, blk, :]) ) // + sum_blk( ABlockSum[m, blk] * QuantBBlkSum[n, blk] ) // -// The third term applies the symmetric W2 zero-point correction: -// QuantBBlkSum[n, blk] = -scale_b * 2, and ABlockSum[m, blk] = scale_a * sum(a). +// The third term applies the W2 zero-point correction: +// QuantBBlkSum[n, blk] = -scale_b * zp, and ABlockSum[m, blk] = scale_a * sum(a). +// `zp` is the per-block zero point baked in at pack time (defaults to +// kDefaultSymmetricZeroPoint2Bit = 2 when no zero-point tensor is supplied). // size_t MLASCALL SQ2BitGemmKernel_BlkSum_CompInt8_Scalar( @@ -243,10 +290,11 @@ SQ2BitGemmKernel_BlkSum_CompInt8_Scalar( // Integer term * scales. acc += a_scale_row[blk] * b_scale_col[blk] * static_cast(dot); - // Symmetric W2 zero-point correction: + // W2 zero-point correction: // dot(a, b_signed) = dot(a, b_unsigned) - zp * sum(a) // so we need C += scale_a * scale_b * (-zp) * sum(a) - // = ABlockSum * QuantBBlkSum (where QuantBBlkSum already encodes -scale_b * zp). + // = ABlockSum * QuantBBlkSum (QuantBBlkSum encodes -scale_b * zp, + // with zp = per-block ZP or 2 when no ZP tensor was supplied). acc += a_blksum_row[blk] * b_blksum_col[blk]; } diff --git a/onnxruntime/core/mlas/lib/sqnbitgemm_kernel_avx512vnni_2bit_blklen64.h b/onnxruntime/core/mlas/lib/sqnbitgemm_kernel_avx512vnni_2bit_blklen64.h index 102f05bfacdae..126b89a2893b4 100644 --- a/onnxruntime/core/mlas/lib/sqnbitgemm_kernel_avx512vnni_2bit_blklen64.h +++ b/onnxruntime/core/mlas/lib/sqnbitgemm_kernel_avx512vnni_2bit_blklen64.h @@ -890,20 +890,42 @@ SQ2BitGemmKernel_BlkSum_CompInt8_Impl( C, CountM, CountN, BlockCountK, Bias, ldc); // BlkSum correction: C += ABlockSum [M x BlockCountK] @ QuantBBlkSum [BlockCountK x N]. - float* c_blk = C; - const float* b_blk_sum = QuantBBlkSum; - size_t RowsRemaining = CountM; - const float* a_blksum_row = ABlockSum; - while (RowsRemaining > 0) { - const auto RowsHandled = GetMlasPlatform().GemmFloatKernel( - a_blksum_row, b_blk_sum, c_blk, - BlockCountK, RowsRemaining, CountN, - BlockCountK, ldc, 1.0f, false); - - c_blk += ldc * RowsHandled; - a_blksum_row += BlockCountK * RowsHandled; - RowsRemaining -= RowsHandled; + // + // TEMP DEBUG: scalar reference instead of GetMlasPlatform().GemmFloatKernel + // to test whether the SGEMM-call shape/layout assumptions are wrong. + // QuantBBlkSum is in the "width-16 chunked" layout: + // BlkSum[(n/16) * BlockCountK * 16 + blk * 16 + (n%16)] + { + for (size_t m = 0; m < CountM; ++m) { + const float* a_row = ABlockSum + m * BlockCountK; + float* c_row = C + m * ldc; + for (size_t n = 0; n < CountN; ++n) { + const size_t chunk = n / 16; + const size_t lane = n % 16; + float acc = 0.0f; + for (size_t blk = 0; blk < BlockCountK; ++blk) { + const float b = QuantBBlkSum[(chunk * BlockCountK + blk) * 16 + lane]; + acc += a_row[blk] * b; + } + c_row[n] += acc; + } + } } + // Original fast path (disabled for debug): + // float* c_blk = C; + // const float* b_blk_sum = QuantBBlkSum; + // size_t RowsRemaining = CountM; + // const float* a_blksum_row = ABlockSum; + // while (RowsRemaining > 0) { + // const auto RowsHandled = GetMlasPlatform().GemmFloatKernel( + // a_blksum_row, b_blk_sum, c_blk, + // BlockCountK, RowsRemaining, CountN, + // BlockCountK, ldc, 1.0f, false); + // + // c_blk += ldc * RowsHandled; + // a_blksum_row += BlockCountK * RowsHandled; + // RowsRemaining -= RowsHandled; + // } return CountM; } diff --git a/onnxruntime/test/mlas/bench/bench_qnbitgemm.cpp b/onnxruntime/test/mlas/bench/bench_qnbitgemm.cpp index 67c8af5015872..6e2f50f33d81a 100644 --- a/onnxruntime/test/mlas/bench/bench_qnbitgemm.cpp +++ b/onnxruntime/test/mlas/bench/bench_qnbitgemm.cpp @@ -146,21 +146,23 @@ BENCHMARK(QNBITGEMM)->Apply(QNBitGemmArgs)->UseRealTime // (K=1024, N=384): 20 nodes // (K=1024, N=4096): 20 nodes // (K=4096, N=1024): 20 nodes -// M=128 is the widest-gap prefill shape from e2e benchmarks. +// Both M=1 (decode) and M=128 (prefill) are exercised — paired with the W2 +// rows below so we get a 3-way (W2 / W4 / W8) comparison at each M. static void QNBitGemmCustomerArgs(benchmark::internal::Benchmark* b) { b->ArgNames({"BlkLen", "M", "N", "K", "Threads", "Symmetric", "HasBias", "ComputeType"}); - const int64_t M = 128; const int64_t BlkLen = 64; const int64_t Threads = 8; const int64_t Symmetric = 1; const int64_t HasBias = 1; - for (auto kn : {std::pair{384, 1024}, - std::pair{1024, 192}, - std::pair{1024, 384}, - std::pair{1024, 4096}, - std::pair{4096, 1024}}) { - for (int64_t ct : {int64_t{SQNBIT_CompFp32}, int64_t{SQNBIT_CompInt8}}) { - b->Args({BlkLen, M, kn.second, kn.first, Threads, Symmetric, HasBias, ct}); + for (int64_t M : {int64_t{1}, int64_t{128}}) { + for (auto kn : {std::pair{384, 1024}, + std::pair{1024, 192}, + std::pair{1024, 384}, + std::pair{1024, 4096}, + std::pair{4096, 1024}}) { + for (int64_t ct : {int64_t{SQNBIT_CompFp32}, int64_t{SQNBIT_CompInt8}}) { + b->Args({BlkLen, M, kn.second, kn.first, Threads, Symmetric, HasBias, ct}); + } } } } @@ -170,25 +172,33 @@ BENCHMARK(QNBITGEMM)->Apply(QNBitGemmCustomerArgs)->UseRealTime(); // 2-bit weight rows for the customer shapes. Exercises the AVX-512 W2 native // path (VNNI variant on AVX-512-VNNI hosts; non-VNNI variant on AVX-512BW // hosts). W2 is registered only for SQNBIT_CompInt8 and BlkLen=64, so we -// emit just that one ComputeType. +// emit just that one ComputeType. Covers both M=1 (decode) and M=128 (prefill). static void QNBit2BitCustomerArgs(benchmark::internal::Benchmark* b) { b->ArgNames({"BlkLen", "M", "N", "K", "Threads", "Symmetric", "HasBias", "ComputeType"}); - const int64_t M = 128; const int64_t BlkLen = 64; const int64_t Threads = 8; const int64_t Symmetric = 1; // W2 native path is symmetric-only. const int64_t HasBias = 1; - for (auto kn : {std::pair{384, 1024}, - std::pair{1024, 192}, - std::pair{1024, 384}, - std::pair{1024, 4096}, - std::pair{4096, 1024}}) { - b->Args({BlkLen, M, kn.second, kn.first, Threads, Symmetric, HasBias, int64_t{SQNBIT_CompInt8}}); + for (int64_t M : {int64_t{1}, int64_t{128}}) { + for (auto kn : {std::pair{384, 1024}, + std::pair{1024, 192}, + std::pair{1024, 384}, + std::pair{1024, 4096}, + std::pair{4096, 1024}}) { + b->Args({BlkLen, M, kn.second, kn.first, Threads, Symmetric, HasBias, int64_t{SQNBIT_CompInt8}}); + } } } BENCHMARK(QNBITGEMM)->Apply(QNBit2BitCustomerArgs)->UseRealTime(); +// 8-bit weight rows on the customer shapes. Used to confirm whether the W2-vs-W4 +// gap is driven by B-weight unpacking cost: W8 has zero unpacking (one byte per +// weight, direct vmovdqu8 + dpbusd), W4 has cheap nibble extraction, W2 has the +// most expensive unpack path. If unpack-density is the bottleneck, expect +// W8 < W4 < W2 in per-MAC cycles at the larger N shapes. +BENCHMARK(QNBITGEMM)->Apply(QNBit2BitCustomerArgs)->UseRealTime(); + // This test gets benchmark arguments from environment variables. template void QNBITGEMM_ENV(benchmark::State& state) { diff --git a/onnxruntime/test/mlas/unittest/test_sqnbitgemm_2bit_gemm.cpp b/onnxruntime/test/mlas/unittest/test_sqnbitgemm_2bit_gemm.cpp index 01e677d8142ba..c1905ec56a8a9 100644 --- a/onnxruntime/test/mlas/unittest/test_sqnbitgemm_2bit_gemm.cpp +++ b/onnxruntime/test/mlas/unittest/test_sqnbitgemm_2bit_gemm.cpp @@ -96,13 +96,25 @@ QuantizeA_Reference(size_t M, constexpr float range_max = static_cast((1 << 7) - 1); const float scale = amax / range_max; - const float scale_recip = scale != 0.0f ? 1.0f / scale : 0.0f; + // Match MLAS's QuantizeARow_CompInt8_avx512 bit-for-bit: it computes + // `inverse_scale = 127 / amax` directly from amax (single op), not + // `1 / (amax/127)` (two ops). The two are equal in real math but + // differ by 1 ULP in float, which is enough to flip the rounded + // int8 by ±1 for A values that land right on a half-integer. + const float scale_recip = amax != 0.0f ? range_max / amax : 0.0f; QuantAScale[m * BlockCountK + k_blk] = scale; for (size_t kk = 0; kk < kBlkLen; ++kk) { const float a = (kk < local_len) ? A[m * K + k + kk] : 0.0f; - const float q = std::round(a * scale_recip); + // `std::nearbyint` uses the current floating-point rounding mode + // (default round-half-to-even) and matches what MLAS's + // `_mm512_roundscale_ps(v, _MM_ROUND_NEAREST)` does at .5 + // boundaries. `std::round` would diverge (rounds half away from + // zero), introducing systematic mismatches in the public-API path + // where MLAS quantises A and the reference computes ABlockSum + // from its own quantization. + const float q = std::nearbyint(a * scale_recip); QuantAData[m * BlockCountK * kBlkLen + k + kk] = static_cast( std::clamp(q, @@ -119,10 +131,12 @@ QuantizeA_Reference(size_t M, // // C[m,n] = bias[n] // + sum_blk( scale_a[m,blk] * scale_b[n,blk] -// * dot(qa[m,blk,:], (qb[n,blk,:] - 2)) ) +// * dot(qa[m,blk,:], (qb[n,blk,:] - zp[n,blk])) ) // -// (Equivalent to the kernel's "dot with raw uint8 weights" + the BlkSum -// correction term, just written without the algebraic split.) +// When BZeroPoints is null the symmetric default ZP = 2 is used for every +// block (matches the kernel's behavior when no zero-point tensor is supplied). +// When non-null, BZeroPoints[n * BlockCountK + blk] gives the per-block ZP in +// [0, 3]. // void ReferenceGemm_W2_CompInt8(size_t M, @@ -131,6 +145,7 @@ ReferenceGemm_W2_CompInt8(size_t M, const float* A, const std::vector& BWeights, // [N * K] in [0,3] const float* QuantBScale, + const uint8_t* BZeroPoints, // [N * BlockCountK] in [0,3] or nullptr const float* Bias, float* C) { @@ -147,12 +162,14 @@ ReferenceGemm_W2_CompInt8(size_t M, const size_t local_len = std::min(K - k, kBlkLen); const float a_scale = QuantAScale[m * BlockCountK + blk]; const float b_scale = QuantBScale[n * BlockCountK + blk]; + const int32_t zp = BZeroPoints != nullptr + ? static_cast(BZeroPoints[n * BlockCountK + blk]) + : static_cast(sq2::kDefaultSymmetricZeroPoint2Bit); int32_t dot = 0; for (size_t kk = 0; kk < local_len; ++kk) { const int8_t qa = QuantAData[m * BlockCountK * kBlkLen + k + kk]; - const int32_t qb = static_cast(BWeights[n * K + k + kk]) - - static_cast(sq2::kDefaultSymmetricZeroPoint2Bit); + const int32_t qb = static_cast(BWeights[n * K + k + kk]) - zp; dot += static_cast(qa) * qb; } acc += static_cast(dot) * a_scale * b_scale; @@ -162,9 +179,32 @@ ReferenceGemm_W2_CompInt8(size_t M, } } +// +// Pack per-block 2-bit zero points into the standard ONNX MatMulNBits W2 +// layout: 4 ZPs per byte along K, row-major in N. Row stride is +// ceil(BlockCountK / 4) bytes. +// +inline std::vector +PackW2ZeroPoints(size_t N, size_t BlockCountK, const std::vector& BZeroPoints) +{ + const size_t ZPCountK = (BlockCountK + 3) / 4; + std::vector packed(N * ZPCountK, std::byte{0}); + for (size_t n = 0; n < N; ++n) { + for (size_t blk = 0; blk < BlockCountK; ++blk) { + const uint8_t zp = BZeroPoints[n * BlockCountK + blk] & 0x03u; + const size_t byte_idx = n * ZPCountK + (blk / 4); + const size_t bit_off = (blk % 4) * 2; + packed[byte_idx] = static_cast( + static_cast(packed[byte_idx]) | (zp << bit_off)); + } + } + return packed; +} + class MlasSQ2BitGemmTest { public: - static void Run(size_t M, size_t N, size_t K, bool WithBias, uint32_t seed) + static void Run(size_t M, size_t N, size_t K, bool WithBias, uint32_t seed, + bool WithZeroPoints = false) { const size_t BlockCountK = (K + kBlkLen - 1) / kBlkLen; ASSERT_EQ(K % kBlkLen, 0u) << "Test K must be a multiple of BlkLen=64"; @@ -200,6 +240,21 @@ class MlasSQ2BitGemmTest { std::vector QuantBScale(N * BlockCountK); for (auto& v : QuantBScale) v = s_dist(rng); + // Per-block zero points (W2: each ZP in [0,3]). The reference path + // uses the raw [N * BlockCountK] uint8 vector; MLAS gets the standard + // ONNX-packed byte stream produced by PackW2ZeroPoints. + std::vector BZeroPoints; + std::vector BZeroPointsPacked; + const uint8_t* BZeroPointsRef = nullptr; + const std::byte* BZeroPointsMlas = nullptr; + if (WithZeroPoints) { + BZeroPoints.resize(N * BlockCountK); + for (auto& v : BZeroPoints) v = static_cast(w_dist(rng)); + BZeroPointsRef = BZeroPoints.data(); + BZeroPointsPacked = PackW2ZeroPoints(N, BlockCountK, BZeroPoints); + BZeroPointsMlas = BZeroPointsPacked.data(); + } + std::vector Bias; const float* BiasPtr = nullptr; if (WithBias) { @@ -210,18 +265,35 @@ class MlasSQ2BitGemmTest { // Pack B through the public API. const size_t PackedSize = MlasQNBitGemmPackQuantBDataSize( - N, K, kBlkBitWidth, kBlkLen, /*has_zero_point=*/false, kComputeType, nullptr); + N, K, kBlkBitWidth, kBlkLen, WithZeroPoints, kComputeType, nullptr); ASSERT_GT(PackedSize, 0u); std::vector PackedQuantB(PackedSize, std::byte{0}); + // Mirror the matmul_nbits.cc prepack flow on AMD64: three separate + // calls, one per input (B data, scales, zero_points). Each call passes + // only its own input and nullptr for the others. The pack function + // must update the BlkSum buffer on the zero_points call by reading + // scales from the already-packed buffer. MlasQNBitGemmPackQuantBData( N, K, kBlkBitWidth, kBlkLen, kComputeType, QuantBData.data(), PackedQuantB.data(), - QuantBScale.data(), /*has_zp_input=*/false, /*QuantBZeroPoint=*/nullptr, + /*QuantBScale=*/nullptr, WithZeroPoints, /*QuantBZeroPoint=*/nullptr, + nullptr, nullptr); + MlasQNBitGemmPackQuantBData( + N, K, kBlkBitWidth, kBlkLen, kComputeType, + /*QuantBData=*/nullptr, PackedQuantB.data(), + QuantBScale.data(), WithZeroPoints, /*QuantBZeroPoint=*/nullptr, nullptr, nullptr); + if (WithZeroPoints) { + MlasQNBitGemmPackQuantBData( + N, K, kBlkBitWidth, kBlkLen, kComputeType, + /*QuantBData=*/nullptr, PackedQuantB.data(), + /*QuantBScale=*/nullptr, WithZeroPoints, BZeroPointsMlas, + nullptr, nullptr); + } const size_t WorkspaceSize = MlasQNBitGemmBatchWorkspaceSize( - M, N, K, 1, kBlkBitWidth, kBlkLen, /*has_zero_point=*/false, kComputeType, nullptr); + M, N, K, 1, kBlkBitWidth, kBlkLen, WithZeroPoints, kComputeType, nullptr); std::vector Workspace(std::max(WorkspaceSize, 1), std::byte{0}); std::vector C(M * N, 0.0f); @@ -232,7 +304,7 @@ class MlasSQ2BitGemmTest { params.QuantBDataWorkspace = PackedQuantB.data(); params.PackedQuantBData = PackedQuantB.data(); params.QuantBScale = QuantBScale.data(); - params.QuantBZeroPoint = nullptr; + params.QuantBZeroPoint = BZeroPointsMlas; params.Bias = BiasPtr; params.C = C.data(); params.ldc = N; @@ -243,7 +315,7 @@ class MlasSQ2BitGemmTest { std::vector CRef(M * N, 0.0f); ReferenceGemm_W2_CompInt8(M, N, K, A.data(), BWeights, QuantBScale.data(), - BiasPtr, CRef.data()); + BZeroPointsRef, BiasPtr, CRef.data()); // Both paths perform the identical integer-domain dot product followed // by the same float multiply-add chain, so the result should agree to @@ -258,7 +330,8 @@ class MlasSQ2BitGemmTest { << " (m=" << (i / N) << ", n=" << (i % N) << ")" << " MLAS=" << C[i] << " Ref=" << CRef[i] << " M=" << M << " N=" << N << " K=" << K - << " WithBias=" << WithBias; + << " WithBias=" << WithBias + << " WithZeroPoints=" << WithZeroPoints; } } }; @@ -285,6 +358,12 @@ TEST(MlasSq2BitTest, GemmCompInt8_BlkLen64_PublicApi) {7, 17, 128}, {16, 64, 512}, {32, 128, 256}, + // Customer model shapes routed through the full MlasQNBitGemmBatch + // dispatcher (threading, PerGemmQuantAWorkspace, packed-B). Both decode + // (M=1) and prefill (M=128) sizes; M=128 forces the multi-threaded + // path that the direct-call test cannot reach. + { 1, 1024, 384}, { 1, 192, 1024}, { 1, 384, 1024}, { 1, 4096, 1024}, { 1, 1024, 4096}, + {128, 1024, 384}, {128, 192, 1024}, {128, 384, 1024}, {128, 4096, 1024}, {128, 1024, 4096}, }; for (uint32_t seed : {0xC0FFEEu, 0xBADC0DEu}) { @@ -296,6 +375,43 @@ TEST(MlasSq2BitTest, GemmCompInt8_BlkLen64_PublicApi) } } +// +// Same coverage as GemmCompInt8_BlkLen64_PublicApi but with per-block +// non-default zero points (random ZP in [0, 3] per block). This is the +// configuration the customer model uses: 4-input MatMulNBits nodes with an +// explicit zero_points initializer. +// +TEST(MlasSq2BitTest, GemmCompInt8_BlkLen64_PublicApi_WithZeroPoints) +{ + if (!MlasIsQNBitGemmAvailable(kBlkBitWidth, kBlkLen, kComputeType)) { + GTEST_SKIP() << "MlasQNBitGemm W2/BlkLen=64/CompInt8 not available on this host"; + } + + struct Shape { size_t M, N, K; }; + constexpr Shape shapes[] = { + {1, 16, 64}, + {1, 32, 128}, + {1, 64, 256}, + {4, 16, 64}, + {4, 33, 192}, + {7, 17, 128}, + {16, 64, 512}, + {32, 128, 256}, + // Customer model shapes (BlkLen=64, asymmetric ZP). + { 1, 1024, 384}, { 1, 192, 1024}, { 1, 384, 1024}, { 1, 4096, 1024}, { 1, 1024, 4096}, + {128, 1024, 384}, {128, 192, 1024}, {128, 384, 1024}, {128, 4096, 1024}, {128, 1024, 4096}, + }; + + for (uint32_t seed : {0xC0FFEEu, 0xBADC0DEu}) { + for (const Shape& s : shapes) { + for (bool bias : {false, true}) { + MlasSQ2BitGemmTest::Run(s.M, s.N, s.K, bias, seed + (bias ? 1u : 0u), + /*WithZeroPoints=*/true); + } + } + } +} + // // Direct-call test harness for the AVX-512 W2 kernel variants. Calls the // kernel through its non-inline test-entry forwarder (declared in @@ -310,7 +426,7 @@ TEST(MlasSq2BitTest, GemmCompInt8_BlkLen64_PublicApi) class MlasSQ2BitGemmDirectCallTest { public: static void Run(size_t M, size_t N, size_t K, bool WithBias, uint32_t seed, - bool TestVnni) + bool TestVnni, bool WithZeroPoints = false) { const size_t BlockCountK = (K + kBlkLen - 1) / kBlkLen; ASSERT_EQ(K % kBlkLen, 0u) << "Test K must be a multiple of BlkLen=64"; @@ -342,6 +458,19 @@ class MlasSQ2BitGemmDirectCallTest { std::vector QuantBScale(N * BlockCountK); for (auto& v : QuantBScale) v = s_dist(rng); + // Per-block zero points (same conventions as the public-API harness). + std::vector BZeroPoints; + std::vector BZeroPointsPacked; + const uint8_t* BZeroPointsRef = nullptr; + const std::byte* BZeroPointsMlas = nullptr; + if (WithZeroPoints) { + BZeroPoints.resize(N * BlockCountK); + for (auto& v : BZeroPoints) v = static_cast(w_dist(rng)); + BZeroPointsRef = BZeroPoints.data(); + BZeroPointsPacked = PackW2ZeroPoints(N, BlockCountK, BZeroPoints); + BZeroPointsMlas = BZeroPointsPacked.data(); + } + std::vector Bias; const float* BiasPtr = nullptr; if (WithBias) { @@ -353,14 +482,14 @@ class MlasSQ2BitGemmDirectCallTest { // Pack B through the public API (produces the same buffer for both // VNNI and non-VNNI consumers; the kernel doesn't care which one). const size_t PackedSize = MlasQNBitGemmPackQuantBDataSize( - N, K, kBlkBitWidth, kBlkLen, /*has_zero_point=*/false, kComputeType, nullptr); + N, K, kBlkBitWidth, kBlkLen, WithZeroPoints, kComputeType, nullptr); ASSERT_GT(PackedSize, 0u); std::vector PackedQuantBBuf(PackedSize, std::byte{0}); MlasQNBitGemmPackQuantBData( N, K, kBlkBitWidth, kBlkLen, kComputeType, QuantBData.data(), PackedQuantBBuf.data(), - QuantBScale.data(), /*has_zp_input=*/false, /*QuantBZeroPoint=*/nullptr, + QuantBScale.data(), WithZeroPoints, BZeroPointsMlas, nullptr, nullptr); // Reconstruct the packed-B view so we can pass the right sub-pointers @@ -424,7 +553,7 @@ class MlasSQ2BitGemmDirectCallTest { // Reference: bit-exact integer-domain math. std::vector CRef(M * N, 0.0f); ReferenceGemm_W2_CompInt8(M, N, K, A.data(), BWeights, QuantBScale.data(), - BiasPtr, CRef.data()); + BZeroPointsRef, BiasPtr, CRef.data()); const float abs_tol = 1e-4f; const float rel_tol = 1e-4f; @@ -436,7 +565,8 @@ class MlasSQ2BitGemmDirectCallTest { << " (m=" << (i / N) << ", n=" << (i % N) << ")" << " MLAS=" << C[i] << " Ref=" << CRef[i] << " M=" << M << " N=" << N << " K=" << K - << " WithBias=" << WithBias; + << " WithBias=" << WithBias + << " WithZeroPoints=" << WithZeroPoints; } } }; @@ -475,6 +605,38 @@ TEST(MlasSq2BitTest, GemmCompInt8_BlkLen64_Avx512) } } +// +// Same as GemmCompInt8_BlkLen64_Avx512 but with random per-block zero points. +// +TEST(MlasSq2BitTest, GemmCompInt8_BlkLen64_Avx512_WithZeroPoints) +{ + if (!GetMlasPlatform().Avx512Supported_) { + GTEST_SKIP() << "AVX-512BW (DQ/VL) not available on this host"; + } + + struct Shape { size_t M, N, K; }; + constexpr Shape shapes[] = { + {1, 16, 64}, + {1, 32, 128}, + {1, 64, 256}, + {4, 16, 64}, + {4, 33, 192}, + {7, 17, 128}, + {16, 64, 512}, + {32, 128, 256}, + }; + + for (uint32_t seed : {0xC0FFEEu, 0xBADC0DEu}) { + for (const Shape& s : shapes) { + for (bool bias : {false, true}) { + MlasSQ2BitGemmDirectCallTest::Run( + s.M, s.N, s.K, bias, seed + (bias ? 1u : 0u), + /*TestVnni=*/false, /*WithZeroPoints=*/true); + } + } + } +} + // // Exercises the AVX-512-VNNI W2 kernel (_mm512_dpbusd_epi32 MAC) via the // direct-call forwarder. On a VNNI host this is the same kernel that the @@ -501,6 +663,11 @@ TEST(MlasSq2BitTest, GemmCompInt8_BlkLen64_Avx512Vnni) {7, 17, 128}, {16, 64, 512}, {32, 128, 256}, + // Customer model shapes (BlkLen=64). Both decode (M=1) and prefill + // (M=128) rows; these are far larger than the small-N shapes above + // and exercise the R1xC4 / R2xC4 tile paths at production sizes. + { 1, 1024, 384}, { 1, 192, 1024}, { 1, 384, 1024}, { 1, 4096, 1024}, { 1, 1024, 4096}, + {128, 1024, 384}, {128, 192, 1024}, {128, 384, 1024}, {128, 4096, 1024}, {128, 1024, 4096}, }; for (uint32_t seed : {0xC0FFEEu, 0xBADC0DEu}) { @@ -513,3 +680,39 @@ TEST(MlasSq2BitTest, GemmCompInt8_BlkLen64_Avx512Vnni) } } } + +// +// Same as GemmCompInt8_BlkLen64_Avx512Vnni but with random per-block zero points. +// +TEST(MlasSq2BitTest, GemmCompInt8_BlkLen64_Avx512Vnni_WithZeroPoints) +{ + if (GetMlasPlatform().QNBitGemmDispatch != &MlasSQNBitGemmDispatchAvx512vnni) { + GTEST_SKIP() << "AVX-512-VNNI not selected as the active dispatch on this host"; + } + + struct Shape { size_t M, N, K; }; + constexpr Shape shapes[] = { + {1, 16, 64}, + {1, 32, 128}, + {1, 64, 256}, + {4, 16, 64}, + {4, 33, 192}, + {7, 17, 128}, + {16, 64, 512}, + {32, 128, 256}, + // Customer model shapes with explicit per-block zero points (matches the + // production accuracy concern). + { 1, 1024, 384}, { 1, 192, 1024}, { 1, 384, 1024}, { 1, 4096, 1024}, { 1, 1024, 4096}, + {128, 1024, 384}, {128, 192, 1024}, {128, 384, 1024}, {128, 4096, 1024}, {128, 1024, 4096}, + }; + + for (uint32_t seed : {0xC0FFEEu, 0xBADC0DEu}) { + for (const Shape& s : shapes) { + for (bool bias : {false, true}) { + MlasSQ2BitGemmDirectCallTest::Run( + s.M, s.N, s.K, bias, seed + (bias ? 1u : 0u), + /*TestVnni=*/true, /*WithZeroPoints=*/true); + } + } + } +} From 1d5b4eb631d9a9c60c5f3d30c8cdaf80cc42522d Mon Sep 17 00:00:00 2001 From: Hari Seshadri Date: Sat, 6 Jun 2026 20:52:18 -0700 Subject: [PATCH 05/17] Add sqnbitgemm avx512vnni 2bit blklen64 kernel --- .../core/mlas/lib/sqnbitgemm_kernel_avx512vnni_2bit_blklen64.h | 3 +-- 1 file changed, 1 insertion(+), 2 deletions(-) diff --git a/onnxruntime/core/mlas/lib/sqnbitgemm_kernel_avx512vnni_2bit_blklen64.h b/onnxruntime/core/mlas/lib/sqnbitgemm_kernel_avx512vnni_2bit_blklen64.h index 126b89a2893b4..6d679ad3a5ce0 100644 --- a/onnxruntime/core/mlas/lib/sqnbitgemm_kernel_avx512vnni_2bit_blklen64.h +++ b/onnxruntime/core/mlas/lib/sqnbitgemm_kernel_avx512vnni_2bit_blklen64.h @@ -84,8 +84,7 @@ inline constexpr size_t kNRows2 = 2; // Dequant one 64-element 2-bit weight block from the packed (broadcast + // shift) layout into a ZMM of 64 unsigned bytes in [0, 3]. // -// Bytes 0..15 : weights[0..15] (shift 0, & 0x03) -// Bytes 16..31 : weights[16..31] (shift 2, & 0x03) +// Bytes 0..15 : weights[0..15] (shift 0, & 0x03) / Bytes 16..31 : weights[16..31] (shift 2, & 0x03) // Bytes 32..47 : weights[32..47] (shift 4, & 0x03) // Bytes 48..63 : weights[48..63] (shift 6, & 0x03) // From b9f085a6f58ad4e2e11e3aadbce8b2d53b4311fc Mon Sep 17 00:00:00 2001 From: Hari Seshadri Date: Wed, 10 Jun 2026 14:40:48 -0700 Subject: [PATCH 06/17] W2 unification --- cmake/onnxruntime_mlas.cmake | 4 + onnxruntime/core/mlas/lib/qnbitgemm.cpp | 17 +- onnxruntime/core/mlas/lib/qnbitgemm.h | 38 +- .../mlas/lib/sqnbitgemm_kernel_avx512.cpp | 57 +- .../lib/sqnbitgemm_kernel_avx512_2bit.cpp | 12 +- .../mlas/lib/sqnbitgemm_kernel_avx512_2bit.h | 86 ++ ...nbitgemm_kernel_avx512_2bit_superblock.cpp | 359 ++++++++ ...sqnbitgemm_kernel_avx512_2bit_superblock.h | 249 +++++ .../mlas/lib/sqnbitgemm_kernel_avx512vnni.cpp | 58 +- ...rnel_avx512vnni_2bit_blklen64_superblock.h | 864 ++++++++++++++++++ onnxruntime/test/mlas/bench/bench_lutgemm.cpp | 25 +- .../test/mlas/bench/bench_qnbitgemm.cpp | 202 +++- .../mlas/unittest/test_sqnbitgemm_2bit.cpp | 161 ++++ .../unittest/test_sqnbitgemm_2bit_gemm.cpp | 320 +------ .../test_sqnbitgemm_2bit_superblock.cpp | 529 +++++++++++ 15 files changed, 2642 insertions(+), 339 deletions(-) create mode 100644 onnxruntime/core/mlas/lib/sqnbitgemm_kernel_avx512_2bit_superblock.cpp create mode 100644 onnxruntime/core/mlas/lib/sqnbitgemm_kernel_avx512_2bit_superblock.h create mode 100644 onnxruntime/core/mlas/lib/sqnbitgemm_kernel_avx512vnni_2bit_blklen64_superblock.h create mode 100644 onnxruntime/test/mlas/unittest/test_sqnbitgemm_2bit_superblock.cpp diff --git a/cmake/onnxruntime_mlas.cmake b/cmake/onnxruntime_mlas.cmake index c7abd3c0c4345..5fe2b844b0111 100644 --- a/cmake/onnxruntime_mlas.cmake +++ b/cmake/onnxruntime_mlas.cmake @@ -243,6 +243,8 @@ function(setup_mlas_source_for_windows) ${MLAS_SRC_DIR}/sqnbitgemm_kernel_avx2.cpp ${MLAS_SRC_DIR}/sqnbitgemm_kernel_avx512_2bit.h ${MLAS_SRC_DIR}/sqnbitgemm_kernel_avx512_2bit.cpp + ${MLAS_SRC_DIR}/sqnbitgemm_kernel_avx512_2bit_superblock.h + ${MLAS_SRC_DIR}/sqnbitgemm_kernel_avx512_2bit_superblock.cpp ${MLAS_SRC_DIR}/sqnbitgemm_kernel_avx512vnni_2bit_blklen64.h ${MLAS_SRC_DIR}/sqnbitgemm_kernel_avx512.cpp ${MLAS_SRC_DIR}/sqnbitgemm_kernel_avx512vnni.cpp @@ -796,6 +798,8 @@ else() ${MLAS_SRC_DIR}/sqnbitgemm_kernel_avx2.cpp ${MLAS_SRC_DIR}/sqnbitgemm_kernel_avx512_2bit.h ${MLAS_SRC_DIR}/sqnbitgemm_kernel_avx512_2bit.cpp + ${MLAS_SRC_DIR}/sqnbitgemm_kernel_avx512_2bit_superblock.h + ${MLAS_SRC_DIR}/sqnbitgemm_kernel_avx512_2bit_superblock.cpp ${MLAS_SRC_DIR}/sqnbitgemm_kernel_avx512vnni_2bit_blklen64.h ${MLAS_SRC_DIR}/sqnbitgemm_lut_kernel_avx2.h ${MLAS_SRC_DIR}/sqnbitgemm_lut_kernel_avx2.cpp diff --git a/onnxruntime/core/mlas/lib/qnbitgemm.cpp b/onnxruntime/core/mlas/lib/qnbitgemm.cpp index ce0824bc8b198..3200d3216f650 100644 --- a/onnxruntime/core/mlas/lib/qnbitgemm.cpp +++ b/onnxruntime/core/mlas/lib/qnbitgemm.cpp @@ -1001,15 +1001,26 @@ SQ2BitGemm_CompInt8( const size_t k_blks = MlasDivRoundup(K, BlkLen); + // The packed-B-data and per-block-scale buffers may use a layout that + // addresses each N-col at a stride larger than `k_blks` (e.g. the AVX-512 + // super-block / W2-v2 kernel rounds up to a multiple of 4 to amortise + // unpack across 4 K-blocks). Ask the dispatch for the effective per-N-col + // block count; default to logical k_blks if the dispatch doesn't override. + // BlkSum keeps the logical stride because the SGEMM correction step + // expects width-16 chunked PackB layout independent of the W2 variant. + const size_t k_blks_eff = (Dispatch->Q2BitGemmEffectiveBlockCountK != nullptr) + ? Dispatch->Q2BitGemmEffectiveBlockCountK(k_blks) + : k_blks; + const size_t lda = k_blks * (per_gemm_quant_a_workspace->QuantScale ? BlkLen : Q8BlkSize(BlkLen)); const size_t ldc = DataParams->ldc; - const size_t ldb = k_blks * MlasQNBitBlkDataSizeInBytes(BlkBitWidth, BlkLen); // BlkLen / 4 bytes per block + const size_t ldb = k_blks_eff * MlasQNBitBlkDataSizeInBytes(BlkBitWidth, BlkLen); // BlkLen / 4 bytes per block const std::byte* QuantA = per_gemm_quant_a_workspace->QuantData + RangeStartM * lda; const float* QuantAScale = per_gemm_quant_a_workspace->QuantScale + RangeStartM * k_blks; const std::byte* QuantBData = static_cast(DataParams->PackedQuantBData) + RangeStartN * ldb; - const float* QuantBScale = DataParams->QuantBScale + RangeStartN * k_blks; + const float* QuantBScale = DataParams->QuantBScale + RangeStartN * k_blks_eff; const float* ABlockSum = per_gemm_quant_a_workspace->BlockSum + RangeStartM * k_blks; const float* QuantBBlkSum = DataParams->QuantBBlkSum + RangeStartN * k_blks; float* C = DataParams->C + RangeStartM * ldc + RangeStartN; @@ -1021,7 +1032,7 @@ SQ2BitGemm_CompInt8( CountN = std::min(RangeCountN - n, size_t{128}); const std::byte* b_col = QuantBData + n * ldb; - const float* b_col_scale = QuantBScale + n * k_blks; + const float* b_col_scale = QuantBScale + n * k_blks_eff; float* c_blk = C + n; const float* bias = (Bias == nullptr) ? nullptr : Bias + n; const float* b_blk_sum = QuantBBlkSum + n * k_blks; diff --git a/onnxruntime/core/mlas/lib/qnbitgemm.h b/onnxruntime/core/mlas/lib/qnbitgemm.h index 421fe13009d26..1d5916a5d5ed5 100644 --- a/onnxruntime/core/mlas/lib/qnbitgemm.h +++ b/onnxruntime/core/mlas/lib/qnbitgemm.h @@ -51,8 +51,19 @@ struct PackedQuantBDataStruct { PackedQuantBDataStruct(void* PackedQuantBWorkspace, size_t N, size_t BlockCountK, size_t BlkLen, bool QuantAUnsigned) : QuantBWorkspace_(PackedQuantBWorkspace), N_(N), BlockCountK_(BlockCountK), BlkLen_(BlkLen) { - const size_t PackedQuantBDataSize = N * BlockCountK * MlasQNBitBlkDataSizeInBytes(BlkBitWidth, BlkLen); - size_t BlkSumSize = MlasDivRoundup(N, 16) * BlockCountK * 16 * sizeof(T); + // For 2-bit weights, the AVX-512 super-block (W2-v2) layout requires + // BlockCountK to be a multiple of 4 (one super-block packs 4 consecutive + // K-blocks together). The legacy W2-v1 layout doesn't require this but + // accepts the small storage padding (<= 48 bytes per N-col), so we pad + // unconditionally for BlkBitWidth=2 to keep a single buffer ABI for both + // kernel variants. The matching pack-size dispatch functions + // (Q2BitGemmPackQuantBDataSize_Avx512 / _SuperBlock) round up by the + // same amount, so the allocated buffer always matches the slab layout + // computed below. + const size_t EffectiveBlockCountK = + (BlkBitWidth == 2) ? ((BlockCountK + 3) / 4) * 4 : BlockCountK; + const size_t PackedQuantBDataSize = N * EffectiveBlockCountK * MlasQNBitBlkDataSizeInBytes(BlkBitWidth, BlkLen); + size_t BlkSumSize = MlasDivRoundup(N, 16) * EffectiveBlockCountK * 16 * sizeof(T); #if defined(MLAS_TARGET_AMD64_IX86) // avx512 requires alignment on a 64-byte boundary PackedQuantBData = (std::byte*)MlasAlignAddress(PackedQuantBWorkspace, 64); @@ -490,6 +501,29 @@ struct MLAS_QNBIT_GEMM_DISPATCH { the packed B layout encodes 2-bit weights instead of 4-bit. */ SQ4BitGemmKernel_BlkSum_CompInt8_Fn* SQ2BitGemmKernel_BlkSum_CompInt8 = nullptr; + /** + * @brief Returns the effective per-N-col block count used by the 2-bit packed + * B-data and B-scale layouts. The layout addresses each N-col at a + * stride of `effective_block_count * `, + * and so does the dispatcher's per-N-tile pointer arithmetic. Some + * 2-bit kernels (notably the AVX-512 super-block / W2-v2 layout) round + * BlockCountK up to a multiple of 4 internally to amortise unpack + * cost across 4 consecutive K-blocks; the buffer is sized accordingly + * (see PackedQuantBDataStruct, which always pads for BlkBitWidth==2) + * so the dispatcher must use the matching stride or it will step + * past the data when n != 0. + * + * BlkSum is NOT affected: it uses the SGEMM-style width-16 chunked + * layout with stride `BlockCountK * 16 * sizeof(float)` per chunk + * regardless of which kernel variant is active. + * + * Returns 0 if not set; the dispatcher falls back to the logical + * `BlockCountK` (the W2-v1 / column-major-ish stride convention). + */ + typedef size_t(Q2BitGemmEffectiveBlockCountK_Fn)(size_t BlockCountK); + + Q2BitGemmEffectiveBlockCountK_Fn* Q2BitGemmEffectiveBlockCountK = nullptr; + /** * @brief Multiply quantized 8-bit integer matrix A with quantized 4-bit integer matrix B. * A and B are block quantized and B is column major. diff --git a/onnxruntime/core/mlas/lib/sqnbitgemm_kernel_avx512.cpp b/onnxruntime/core/mlas/lib/sqnbitgemm_kernel_avx512.cpp index ad5a257dd1306..303e4ee6a2bdc 100644 --- a/onnxruntime/core/mlas/lib/sqnbitgemm_kernel_avx512.cpp +++ b/onnxruntime/core/mlas/lib/sqnbitgemm_kernel_avx512.cpp @@ -28,6 +28,8 @@ Module Name: #include "sqnbitgemm_kernel_avx512_int8_blklen128.h" #include "sqnbitgemm_kernel_avx512_2bit.h" #include "sqnbitgemm_kernel_avx512vnni_2bit_blklen64.h" +#include "sqnbitgemm_kernel_avx512_2bit_superblock.h" +#include "sqnbitgemm_kernel_avx512vnni_2bit_blklen64_superblock.h" // // SQNBIT_CompFp32 kernel implementation. @@ -508,6 +510,35 @@ SQ2BitGemmKernel_BlkSum_CompInt8_Avx512_TestEntry( } } // namespace onnxruntime::mlas::sq2bit_avx512 +// +// Unit-test entry point for the AVX-512BW (non-VNNI) W2 SUPER-BLOCK kernel. +// Sibling of the VNNI variant in sqnbitgemm_kernel_avx512vnni.cpp. +// +namespace onnxruntime::mlas::sq2bit_avx512_super { +size_t MLASCALL +SQ2BitGemmKernel_BlkSum_CompInt8_Super_Avx512_TestEntry( + size_t BlkLen, + const std::byte* QuantA, + const float* QuantAScale, + const std::byte* QuantBData, + const float* QuantBScale, + const std::byte* QuantBZeroPoint, + float* C, + size_t CountM, + size_t CountN, + size_t CountK, + size_t BlockCountK, + const float* Bias, + size_t ldc, + const float* ABlockSum, + const float* QuantBBlkSum) +{ + return SQ2BitGemmKernel_BlkSum_CompInt8_Super_Avx512( + BlkLen, QuantA, QuantAScale, QuantBData, QuantBScale, QuantBZeroPoint, + C, CountM, CountN, CountK, BlockCountK, Bias, ldc, ABlockSum, QuantBBlkSum); +} +} // namespace onnxruntime::mlas::sq2bit_avx512_super + const MLAS_QNBIT_GEMM_DISPATCH MlasSQNBitGemmDispatchAvx512 = []() { MLAS_QNBIT_GEMM_DISPATCH d; @@ -527,13 +558,25 @@ const MLAS_QNBIT_GEMM_DISPATCH MlasSQNBitGemmDispatchAvx512 = []() { d.SQ8BitGemmKernel_BlkSum_CompInt8 = SQ8BitGemmKernel_BlkSum_CompInt8_avx512; d.QuantizeARowComputeBlkSum_CompInt8 = QuantizeARow_CompInt8_avx512; - // 2-bit native CompInt8 path: AVX-512BW variant (no VNNI). Uses the same - // tile + pack layout as the VNNI variant; the per-block MAC is - // `vpmaddubsw + vpmaddwd + vpaddd` instead of `_mm512_dpbusd_epi32`. - // Pack-size and pack functions are identical between AVX-512 and AVX-512-VNNI. - d.Q2BitGemmPackQuantBDataSize = onnxruntime::mlas::sq2bit_avx512::Q2BitGemmPackQuantBDataSize_Avx512; - d.SQ2BitGemmPackQuantBDataAndBlkSum = onnxruntime::mlas::sq2bit_avx512::SQ2BitGemmPackQuantBDataAndBlkSum_Scalar; - d.SQ2BitGemmKernel_BlkSum_CompInt8 = onnxruntime::mlas::sq2bit_avx512::SQ2BitGemmKernel_BlkSum_CompInt8_Avx512; + // 2-bit native CompInt8 path: AVX-512BW (no VNNI) variant of the + // super-block (W2-v2) kernel -- single 64-byte load + four fixed + // shift/mask pairs to unpack 4 K-blocks at once. Pack-size and pack + // functions are shared with the AVX-512-VNNI variant; only the inner + // integer MAC differs (`vpmaddubsw + vpmaddwd + vpaddd` here vs + // `_mm512_dpbusd_epi32` in the VNNI variant). + // + // The legacy `sq2bit_avx512::*` (W2-v1) symbols remain in the build but + // are no longer reached at runtime. A follow-up will remove them once + // W2-v2 has soaked in production. + d.Q2BitGemmPackQuantBDataSize = onnxruntime::mlas::sq2bit_avx512_super::Q2BitGemmPackQuantBDataSize_SuperBlock; + d.SQ2BitGemmPackQuantBDataAndBlkSum = onnxruntime::mlas::sq2bit_avx512_super::SQ2BitGemmPackQuantBDataAndBlkSum_SuperBlockScalar; + d.SQ2BitGemmKernel_BlkSum_CompInt8 = onnxruntime::mlas::sq2bit_avx512_super::SQ2BitGemmKernel_BlkSum_CompInt8_Super_Avx512_TestEntry; + // W2-v2 packs and addresses each N-col at a stride of + // SuperBlockCountKPadded * kSuperBlockBlks blocks (BlockCountK rounded + // UP to a multiple of 4) to keep the inner K-loop's super-block stride + // constant. The dispatcher needs this stride for its per-N-tile pointer + // arithmetic in SQ2BitGemm_CompInt8. + d.Q2BitGemmEffectiveBlockCountK = [](size_t BlockCountK) { return ((BlockCountK + 3) / 4) * 4; }; return d; }(); diff --git a/onnxruntime/core/mlas/lib/sqnbitgemm_kernel_avx512_2bit.cpp b/onnxruntime/core/mlas/lib/sqnbitgemm_kernel_avx512_2bit.cpp index 09cea17e43685..7cb69e6288587 100644 --- a/onnxruntime/core/mlas/lib/sqnbitgemm_kernel_avx512_2bit.cpp +++ b/onnxruntime/core/mlas/lib/sqnbitgemm_kernel_avx512_2bit.cpp @@ -72,9 +72,15 @@ Q2BitGemmPackQuantBDataSize_Avx512( } const size_t BlockCountK = MlasDivRoundup(K, BlkLen); - size_t PackedQuantBDataSize = N * BlockCountK * kBlkBytes; - const size_t ScaleSize = N * BlockCountK * sizeof(float); - size_t BlkSumSize = MlasDivRoundup(N, 16) * BlockCountK * 16 * sizeof(float); + // Pad BlockCountK to a multiple of 4 to match the W2 buffer ABI used by + // PackedQuantBDataStruct (BlkBitWidth=2 always pads). The padding is a + // few extra K-block slots per N-col (<= 48 bytes data + a few floats for + // scales / BlkSum) -- negligible -- but lets the W2-v1 (this path) and + // W2-v2 (super-block) pack helpers share a single buffer layout. + const size_t BlockCountKPadded = ((BlockCountK + 3) / 4) * 4; + size_t PackedQuantBDataSize = N * BlockCountKPadded * kBlkBytes; + const size_t ScaleSize = N * BlockCountKPadded * sizeof(float); + size_t BlkSumSize = MlasDivRoundup(N, 16) * BlockCountKPadded * 16 * sizeof(float); constexpr size_t kPackedQuantBDataAlignment = 64; // AVX-512 friendly PackedQuantBDataSize += kPackedQuantBDataAlignment - 1; diff --git a/onnxruntime/core/mlas/lib/sqnbitgemm_kernel_avx512_2bit.h b/onnxruntime/core/mlas/lib/sqnbitgemm_kernel_avx512_2bit.h index c9a3a56b2bb57..7b71c79b71ff8 100644 --- a/onnxruntime/core/mlas/lib/sqnbitgemm_kernel_avx512_2bit.h +++ b/onnxruntime/core/mlas/lib/sqnbitgemm_kernel_avx512_2bit.h @@ -220,6 +220,92 @@ UnpackSourceBlock_BlkLen64_Reference(const std::byte* src, uint8_t out[kBlkLen]) } } +// ----------------------------------------------------------------------------- +// EXPERIMENTAL: super-block packing for fast unpack +// ----------------------------------------------------------------------------- +// +// The existing per-block pack (PackBlock_BlkLen64) stores 16 bytes per K-block +// such that each byte holds 4 weights from 4 different "rows" of the block. +// Unpacking that layout into a ZMM of 64 natural-order weights requires a +// broadcast + per-dword variable shift (`vpsrlvd`, ~3c latency, 1c throughput +// on Zen5) -- a cost we measured at ~30-35% of total W2 kernel time. +// +// The super-block layout below groups FOUR consecutive K-blocks together into +// a single 64-byte buffer (= 4 * 16 bytes, same total storage) such that +// byte[i] (i = 0..63) holds: +// bits[0..1] = block_0.weight[i] +// bits[2..3] = block_1.weight[i] +// bits[4..5] = block_2.weight[i] +// bits[6..7] = block_3.weight[i] +// +// Unpack with a single ZMM load and four fixed-shift+mask pairs: +// +// __m512i super = _mm512_loadu_si512(packed); +// __m512i mask = _mm512_set1_epi8(0x03); +// __m512i bv0 = _mm512_and_si512(super, mask); +// __m512i bv1 = _mm512_and_si512(_mm512_srli_epi16(super, 2), mask); +// __m512i bv2 = _mm512_and_si512(_mm512_srli_epi16(super, 4), mask); +// __m512i bv3 = _mm512_and_si512(_mm512_srli_epi16(super, 6), mask); +// +// Each shift+mask is ~2c (1c srli_epi16 + 1c andd) and the four chains are +// fully independent, so the critical path is ~4c for ALL four blocks combined +// vs the current ~20c (4 broadcasts * ~5c each, partially overlapped). +// +// Note on `_mm512_srli_epi16`: it shifts each 16-bit lane by N. For a byte +// that is the LOW byte of a 16-bit lane, the shift pulls in bits from the +// adjacent HIGH byte. The subsequent AND with 0x03 discards those leaked +// bits, leaving the correct per-byte result. For the HIGH byte, zeros are +// shifted in from the top, which is what we want. + +constexpr size_t kSuperBlockBlks = 4; // 4 K-blocks per super +constexpr size_t kSuperBlockBytes = kSuperBlockBlks * kBlkBytes; // 64 bytes +constexpr size_t kSuperBlockWeights = kSuperBlockBlks * kBlkLen; // 256 weights + +// +// Pack 4 consecutive K-blocks (4 * 16 = 64 source bytes in standard ONNX +// layout) into a 64-byte super-block. Pure permutation of the 256 2-bit +// elements; bit-identical round-trip with UnpackSuperBlock4_BlkLen64_Reference. +// +inline void +PackSuperBlock4_BlkLen64(const std::byte* src_block_0, + const std::byte* src_block_1, + const std::byte* src_block_2, + const std::byte* src_block_3, + std::byte* dst) +{ + for (size_t i = 0; i < kBlkLen; ++i) { + const uint8_t v0 = ExtractSrcWeight(src_block_0, i); + const uint8_t v1 = ExtractSrcWeight(src_block_1, i); + const uint8_t v2 = ExtractSrcWeight(src_block_2, i); + const uint8_t v3 = ExtractSrcWeight(src_block_3, i); + dst[i] = static_cast( + static_cast(v0 | (v1 << 2) | (v2 << 4) | (v3 << 6)) + ); + } +} + +// +// Reference unpack of one 64-byte super-block back into 4 K-blocks worth of +// natural-order uint8 weights ([0, 3]). Written from the documented layout +// rule -- intentionally independent of PackSuperBlock4_BlkLen64 so it can +// serve as a round-trip oracle. +// +inline void +UnpackSuperBlock4_BlkLen64_Reference(const std::byte* packed, + uint8_t out_block_0[kBlkLen], + uint8_t out_block_1[kBlkLen], + uint8_t out_block_2[kBlkLen], + uint8_t out_block_3[kBlkLen]) +{ + for (size_t i = 0; i < kBlkLen; ++i) { + const uint8_t b = static_cast(packed[i]); + out_block_0[i] = static_cast((b >> 0) & 0x03u); + out_block_1[i] = static_cast((b >> 2) & 0x03u); + out_block_2[i] = static_cast((b >> 4) & 0x03u); + out_block_3[i] = static_cast((b >> 6) & 0x03u); + } +} + // // Reference / dispatch entry points defined in sqnbitgemm_kernel_avx512_2bit.cpp. // diff --git a/onnxruntime/core/mlas/lib/sqnbitgemm_kernel_avx512_2bit_superblock.cpp b/onnxruntime/core/mlas/lib/sqnbitgemm_kernel_avx512_2bit_superblock.cpp new file mode 100644 index 0000000000000..01be1d153feae --- /dev/null +++ b/onnxruntime/core/mlas/lib/sqnbitgemm_kernel_avx512_2bit_superblock.cpp @@ -0,0 +1,359 @@ +/*++ + +Copyright (c) Microsoft Corporation. All rights reserved. + +Licensed under the MIT License. + +Module Name: + + sqnbitgemm_kernel_avx512_2bit_superblock.cpp + +Abstract: + + Pack helpers and a scalar reference kernel for the super-block W2 layout. + + See sqnbitgemm_kernel_avx512_2bit_superblock.h for the layout description + and the rationale (closing the W2-vs-W4 prefill gap by replacing the + per-block broadcast + variable shift unpack with a single 64-byte load + + four fixed-shift+mask pairs). + + This translation unit is scalar / portable. The vectorized inner loop + that consumes the super-block layout lives in a separate header to be + added in phase 3, and is wired into the dispatch tables in phase 4. + +--*/ + +#include "sqnbitgemm_kernel_avx512_2bit_superblock.h" + +#include +#include +#include +#include + +#include "mlasi.h" +#include "qnbitgemm.h" + +namespace onnxruntime { +namespace mlas { +namespace sq2bit_avx512_super { + +namespace sq2 = ::onnxruntime::mlas::sq2bit_avx512; + +// +// Workspace / pack-buffer size for the super-block W2 path. Returns 0 if any +// of the configuration constraints is violated; the caller (MlasQNBitGemmPackQuantBDataSize) +// treats that as "unsupported" and falls back to the original W2 path. +// +// Constraints: +// * BlkLen == 64 +// * ComputeType == SQNBIT_CompInt8 +// +// K-tail handling: BlockCountK is rounded UP to a multiple of kSuperBlockBlks +// for the storage that the inner K-loop walks (PackedQuantBData, +// PackedQuantBScale). Padding slots hold zeroed weights and scales, so they +// contribute exactly 0 to the dot product. The BlkSum buffer is kept at the +// LOGICAL BlockCountK because it is consumed by the SGEMM correction step, +// not by the inner K-loop. +// +// Storage matches the original W2 layout total bytes when BlockCountK is a +// multiple of 4. When not a multiple of 4, storage grows by at most 3 K-blocks +// per N-col (= up to 48 bytes per col -- negligible at production N values). +// +// [PackedQuantBData] N * BlockCountKPadded * kBlkBytes +// [PackedQuantBScale] N * BlockCountKPadded * sizeof(float) +// [QuantBBlkSum] roundup_16(N) * BlockCountK (logical) * 16 floats +// +size_t MLASCALL +Q2BitGemmPackQuantBDataSize_SuperBlock( + size_t N, + size_t K, + size_t BlkLen, + bool /* HasZeroPoint */, + MLAS_QNBIT_GEMM_COMPUTE_TYPE ComputeType, + const MLAS_BACKEND_KERNEL_SELECTOR_CONFIG* /* BackendKernelSelectorConfig */ +) +{ + if (BlkLen != sq2::kBlkLen || ComputeType != SQNBIT_CompInt8) { + return 0; + } + const size_t BlockCountK = MlasDivRoundup(K, BlkLen); + if (BlockCountK == 0) { + return 0; + } + const size_t BlockCountKPadded = + MlasDivRoundup(BlockCountK, kSuperBlockBlks) * kSuperBlockBlks; + + // Use BlockCountKPadded for BlkSum sizing too. The actual SGEMM-correction + // step only reads LOGICAL BlockCountK entries, but PackedQuantBDataStruct + // is constructed by the caller with a single BlockCountK value that + // controls BOTH the packed-B size and the BlkSum offset. If we sized + // BlkSum at the logical BlockCountK while sizing packed-B at padded, + // the struct's BlkSum pointer would land inside the packed-B region + // (because the caller's struct uses one BlockCountK consistently). The + // extra storage from padding the BlkSum is ~16 floats per N -- trivial. + size_t PackedQuantBDataSize = N * BlockCountKPadded * kBlkBytes; + const size_t ScaleSize = N * BlockCountKPadded * sizeof(float); + size_t BlkSumSize = MlasDivRoundup(N, 16) * BlockCountKPadded * 16 * sizeof(float); + + constexpr size_t kPackedQuantBDataAlignment = 64; + PackedQuantBDataSize += kPackedQuantBDataAlignment - 1; + + constexpr size_t kBlkSumAlignment = MlasQNBitQuantBBlkSumAlignment(); + BlkSumSize += kBlkSumAlignment - 1; + + return PackedQuantBDataSize + ScaleSize + BlkSumSize; +} + +// +// Pack quantized B data + scales + per-block sums for the super-block W2 path. +// +// PackedQuantBData layout (super-blocks of 4 K-blocks, 64 bytes each): +// The super-block at logical (n, blk_super=blk/4) lives at byte offset +// PackedQuantBOffsetBytes_W2_SuperBlock(n, blk_super, SuperBlockCountK, NMain). +// Byte b within the super-block holds 2-bit weight b from each of the 4 +// constituent K-blocks at bit positions {0..1, 2..3, 4..5, 6..7}. +// +// PackedQuantBScale layout: one float per K-block, four floats per super-block, +// addressed by PackedQuantBScaleOffset_W2_SuperBlock. +// +// QuantBBlkSum layout: the same width-16 row-major chunked layout used by the +// existing W2 path, so the SGEMM correction step (MlasGemmFloatKernel) can be +// shared verbatim with the existing kernel. +// +// Mirrors the SQ2BitGemmPackQuantBDataAndBlkSum_Scalar prepack 3-call pattern: +// ORT's matmul_nbits.cc invokes this function up to three times (B, scales, ZP). +// We write scales when scales arrive, then re-derive BlkSum whenever either +// scales or zero-points arrive, reading scales from the already-packed buffer. +// +void MLASCALL +SQ2BitGemmPackQuantBDataAndBlkSum_SuperBlockScalar( + size_t N, + size_t K, + size_t BlkLen, + MLAS_QNBIT_GEMM_COMPUTE_TYPE /* ComputeType */, + const std::byte* QuantBDataBegin, + const float* QuantBScaleBegin, + bool /* HasZeroPoint */, + const std::byte* QuantBZPBegin, + PackedQuantBDataStruct& PackedQuantB, + MLAS_THREADPOOL* ThreadPool, + const MLAS_BACKEND_KERNEL_SELECTOR_CONFIG* /* BackendKernelSelectorConfig */ +) +{ + assert(BlkLen == sq2::kBlkLen); + if (BlkLen != sq2::kBlkLen) { + return; + } + + const size_t BlockCountK = MlasDivRoundup(K, BlkLen); + if (BlockCountK == 0) { + return; + } + + // Pad BlockCountK up to a multiple of kSuperBlockBlks so the inner K-loop + // can iterate whole super-blocks uniformly. Padding slots store zeroed + // weights / scales and contribute exactly 0 to the dot product. + const size_t SuperBlockCountKPadded = + MlasDivRoundup(BlockCountK, kSuperBlockBlks); + const size_t BlockCountKPadded = SuperBlockCountKPadded * kSuperBlockBlks; + const size_t NMain = (N / kNCols4) * kNCols4; + + // Zero source block used when packing a super whose K-range crosses the + // logical BlockCountK boundary. We point to this static zero buffer in + // the missing slots so the existing 4-block pack helper does the right + // thing without any branching inside it. + static const std::byte kZeroBlock[kBlkBytes] = {}; + + // ----- B-data pack ----- + if (QuantBDataBegin != nullptr) { + std::byte* PackedQuantBData = PackedQuantB.PackedQuantBData; + const size_t Iterations = N * SuperBlockCountKPadded; + MlasTrySimpleParallel( + ThreadPool, static_cast(Iterations), + [&](ptrdiff_t tid) { + const size_t n = static_cast(tid) / SuperBlockCountKPadded; + const size_t blk_super = static_cast(tid) % SuperBlockCountKPadded; + const size_t blk0 = blk_super * kSuperBlockBlks; + + // Pick real source block pointers for slots that exist; the + // static zero buffer for slots past the logical BlockCountK. + auto src_for = [&](size_t blk) -> const std::byte* { + if (blk < BlockCountK) { + return QuantBDataBegin + (n * BlockCountK + blk) * kBlkBytes; + } + return kZeroBlock; + }; + const std::byte* src_blk_0 = src_for(blk0 + 0); + const std::byte* src_blk_1 = src_for(blk0 + 1); + const std::byte* src_blk_2 = src_for(blk0 + 2); + const std::byte* src_blk_3 = src_for(blk0 + 3); + + const size_t dst_offset = + PackedQuantBOffsetBytes_W2_SuperBlock(n, blk_super, SuperBlockCountKPadded, NMain); + sq2::PackSuperBlock4_BlkLen64(src_blk_0, src_blk_1, src_blk_2, src_blk_3, + PackedQuantBData + dst_offset); + } + ); + } + + // ----- Scales ----- + // Iterate over the PADDED block count so trailing padding slots get + // explicit zero scales (otherwise they could hold uninitialised noise + // and the kernel's K-loop would read those into the FMA). + if (QuantBScaleBegin != nullptr) { + float* PackedScales = PackedQuantB.PackedQuantBScale; + const size_t Iterations = N * BlockCountKPadded; + MlasTrySimpleParallel( + ThreadPool, static_cast(Iterations), + [&](ptrdiff_t tid) { + const size_t n = static_cast(tid) / BlockCountKPadded; + const size_t blk = static_cast(tid) % BlockCountKPadded; + const float scale = (blk < BlockCountK) + ? QuantBScaleBegin[n * BlockCountK + blk] + : 0.0f; + PackedScales[PackedQuantBScaleOffset_W2_SuperBlock(n, blk, BlockCountKPadded, NMain)] = scale; + } + ); + } + + // ----- BlkSum (recomputed whenever scales or ZPs arrive) ----- + // BlkSum is consumed by the SGEMM correction step (MlasGemmFloatKernel), + // which is called outside the inner K-loop with the LOGICAL BlockCountK + // and the per-row ABlockSum the dispatcher produced for that logical K. + // We therefore only need to fill the first BlockCountK entries; the buffer + // is sized at MlasDivRoundup(N, 16) * BlockCountK * 16 floats (logical). + if (QuantBScaleBegin != nullptr || QuantBZPBegin != nullptr) { + const float* PackedScales = PackedQuantB.PackedQuantBScale; + float* BlkSum = PackedQuantB.QuantBBlkSum; + const size_t ZPCountK = MlasDivRoundup(BlockCountK, 4); + const size_t Iterations = N * BlockCountK; + MlasTrySimpleParallel( + ThreadPool, static_cast(Iterations), + [&](ptrdiff_t tid) { + const size_t n = static_cast(tid) / BlockCountK; + const size_t blk = static_cast(tid) % BlockCountK; + const float scale = + PackedScales[PackedQuantBScaleOffset_W2_SuperBlock(n, blk, BlockCountKPadded, NMain)]; + + uint8_t zp = kDefaultSymmetricZeroPoint2Bit; + if (QuantBZPBegin != nullptr) { + const size_t zp_byte_idx = n * ZPCountK + (blk / 4); + const size_t zp_bit_off = (blk % 4) * 2; + zp = static_cast( + (static_cast(QuantBZPBegin[zp_byte_idx]) >> zp_bit_off) & 0x03u); + } + + const size_t blksum_offset = ((n / 16) * BlockCountK + blk) * 16 + (n % 16); + BlkSum[blksum_offset] = -scale * static_cast(zp); + } + ); + } +} + +// +// Scalar reference kernel that consumes the super-block packed layout. +// Same math as the existing reference kernel; differs only in how it walks +// PackedQuantBData (super-block-major) and PackedQuantBScale (super-block-major). +// +// This is the correctness oracle for the SIMD super-block kernel coming in +// Phase 3. It also lets us validate the pack layout end-to-end via the +// existing MlasQNBitGemmBatch dispatch path once we wire it up. +// +size_t MLASCALL +SQ2BitGemmKernel_BlkSum_CompInt8_SuperBlockScalar( + const size_t BlkLen, + const std::byte* QuantA, + const float* QuantAScale, + const std::byte* QuantBData, + const float* QuantBScale, + const std::byte* /* QuantBZeroPoint */, + float* C, + size_t CountM, + size_t CountN, + size_t /* CountK */, + size_t BlockCountK, + const float* Bias, + size_t ldc, + const float* ABlockSum, + const float* QuantBBlkSum +) +{ + if (BlkLen != sq2::kBlkLen) { + return 0; + } + if (BlockCountK == 0) { + return 0; + } + + // PackedQuantBData and PackedQuantBScale are addressed via padded counts + // (K-tail handling -- see Q2BitGemmPackQuantBDataSize_SuperBlock). The K + // dot-product loop itself iterates only LOGICAL BlockCountK steps because + // A is unpadded; the kernel never reads past the real A rows. + const size_t SuperBlockCountKPadded = + MlasDivRoundup(BlockCountK, kSuperBlockBlks); + const size_t BlockCountKPadded = SuperBlockCountKPadded * kSuperBlockBlks; + + // The kernel is called by SQ2BitGemm_CompInt8 with the full CountN range; that + // function selects an N-tile boundary (kNCols4) up the stack. For a scalar + // reference path we don't depend on the 4-N-col grouping, but we DO need to + // index PackedQuantBData/PackedQuantBScale via the super-block offset helpers + // so we read the right bytes regardless of caller tile choice. + // + // CountN may not be a multiple of kNCols4 in the tail case. Detect that and + // fall back to plain column-major for the tail cols (the layout helpers + // already encode this rule). + const size_t NMainLocal = (CountN / kNCols4) * kNCols4; + + const size_t lda = BlockCountK * sq2::kBlkLen; // bytes per A row (int8) + const size_t lda_scale = BlockCountK; // floats per A scale row + + for (size_t m = 0; m < CountM; ++m) { + const int8_t* a_row = reinterpret_cast(QuantA + m * lda); + const float* a_scale_row = QuantAScale + m * lda_scale; + const float* a_blksum_row = ABlockSum + m * lda_scale; + float* c_row = C + m * ldc; + + for (size_t n = 0; n < CountN; ++n) { + float acc = (Bias != nullptr) ? Bias[n] : 0.0f; + + for (size_t blk = 0; blk < BlockCountK; ++blk) { + // Pull the super-block this K-block belongs to and unpack only the + // slot we need (block_in_super = blk % 4 selects the 2-bit field). + const size_t blk_super = blk / kSuperBlockBlks; + const size_t blk_in_super = blk % kSuperBlockBlks; + const size_t super_offset = + PackedQuantBOffsetBytes_W2_SuperBlock(n, blk_super, SuperBlockCountKPadded, NMainLocal); + const std::byte* super = QuantBData + super_offset; + + uint8_t b_unpacked[sq2::kBlkLen]; + for (size_t i = 0; i < sq2::kBlkLen; ++i) { + const uint8_t byte = static_cast(super[i]); + b_unpacked[i] = static_cast((byte >> (2 * blk_in_super)) & 0x03u); + } + + const int8_t* a_blk = a_row + blk * sq2::kBlkLen; + int32_t dot = 0; + for (size_t i = 0; i < sq2::kBlkLen; ++i) { + dot += static_cast(a_blk[i]) * static_cast(b_unpacked[i]); + } + + const float b_scale = + QuantBScale[PackedQuantBScaleOffset_W2_SuperBlock(n, blk, BlockCountKPadded, NMainLocal)]; + acc += a_scale_row[blk] * b_scale * static_cast(dot); + + // The width-16 row-major BlkSum layout is column-major in n + // (one float per (n, blk)); same as the existing W2 path. + const size_t blksum_offset = ((n / 16) * BlockCountK + blk) * 16 + (n % 16); + acc += a_blksum_row[blk] * QuantBBlkSum[blksum_offset]; + } + + c_row[n] = acc; + } + } + + return CountM; +} + +} // namespace sq2bit_avx512_super +} // namespace mlas +} // namespace onnxruntime diff --git a/onnxruntime/core/mlas/lib/sqnbitgemm_kernel_avx512_2bit_superblock.h b/onnxruntime/core/mlas/lib/sqnbitgemm_kernel_avx512_2bit_superblock.h new file mode 100644 index 0000000000000..a2a6b93d4c7ee --- /dev/null +++ b/onnxruntime/core/mlas/lib/sqnbitgemm_kernel_avx512_2bit_superblock.h @@ -0,0 +1,249 @@ +/*++ + +Copyright (c) Microsoft Corporation. All rights reserved. + +Licensed under the MIT License. + +Module Name: + + sqnbitgemm_kernel_avx512_2bit_superblock.h + +Abstract: + + Pack-time helpers and scalar reference routines for the EXPERIMENTAL + super-block W2 layout: groups of 4 K-blocks share a single 64-byte + packed buffer that allows the AVX-512 unpack to be one ZMM load plus + four fixed `vpsrlw+vpand` pairs (instead of the current per-block + broadcast + variable shift). + + Layout summary (BlkLen=64 only): + + * Each "super-block" packs FOUR consecutive K-blocks (256 weights total). + * Total storage per super-block = kBlkBytes * 4 = 64 bytes (identical to + 4 separately-packed blocks under the current scheme). + * Byte b of the super-block holds: + bits[0..1] = block_0.weight[b] + bits[2..3] = block_1.weight[b] + bits[4..5] = block_2.weight[b] + bits[6..7] = block_3.weight[b] + * The N-dimension uses the same 4-col-grouped layout as the existing + W2 kernel (kNCols4 = 4), so the main NMain region groups 4 N-cols + per "row" of super-blocks. + + Restrictions: + + * BlkLen == 64 only. + * BlockCountK must be a multiple of kSuperBlockBlks = 4. The prototype + path returns 0 from the pack-size helper for non-multiples; the + caller falls back to the existing W2 path. + * The customer model's K dimensions (384, 1024, 4096) are all multiples + of 256 (= 4 * 64), so all customer shapes satisfy this constraint. + +--*/ + +#pragma once + +#include +#include +#include + +#include "mlas.h" +#include "mlas_qnbit.h" + +#include "sqnbitgemm_kernel_avx512_2bit.h" // Re-uses kBlkLen, kBlkBytes, kNCols4, etc. + +template +struct PackedQuantBDataStruct; // fwd decl, defined in qnbitgemm.h + +struct MLAS_BACKEND_KERNEL_SELECTOR_CONFIG; + +namespace onnxruntime { +namespace mlas { +namespace sq2bit_avx512_super { + +using ::onnxruntime::mlas::sq2bit_avx512::kBlkBytes; +using ::onnxruntime::mlas::sq2bit_avx512::kBlkLen; +using ::onnxruntime::mlas::sq2bit_avx512::kDefaultSymmetricZeroPoint2Bit; +using ::onnxruntime::mlas::sq2bit_avx512::kNCols4; +using ::onnxruntime::mlas::sq2bit_avx512::kSuperBlockBytes; +using ::onnxruntime::mlas::sq2bit_avx512::kSuperBlockBlks; +using ::onnxruntime::mlas::sq2bit_avx512::kWeightsPerByte; + +// ----------------------------------------------------------------------------- +// Super-block packed-data layout. +// +// Main region (n < NMain = floor(N / kNCols4) * kNCols4): +// 4-N-col groups of g = n / 4, col within group c = n % 4. +// K-super-block index s = blk / 4 (s in [0, BlockCountK / 4)). +// Within a group, super-blocks run consecutively across the 4 cols, +// so each (s, group) slot is a contiguous (kNCols4 * kSuperBlockBytes) +// = 256 byte chunk. +// +// Tail region (n >= NMain): plain column-major super-blocks, identical +// shape to the main region but flat in N. +// +// K-tail handling (BlockCountK not a multiple of kSuperBlockBlks): +// The pack helpers round BlockCountK up to a multiple of kSuperBlockBlks +// (= 4) for storage purposes -- the padding 1-3 blocks at the trailing +// super-block contain zeroed weight bytes and zeroed scales, so they +// contribute 0 to the integer dot product and 0 to the BlkSum correction. +// This lets the SIMD kernel iterate the super-block K-loop without a +// special tail handler for B, and avoids dual packing layouts. Storage +// waste is at most (kSuperBlockBlks - 1) blocks per N-col, i.e. <= 48 +// bytes per col -- negligible at any realistic N. +// +// Conventions used by the offset helpers below: +// * `SuperBlockCountKPadded = ceil(BlockCountK / kSuperBlockBlks)` is +// the number of super-blocks the kernel actually iterates. +// * `BlockCountKPadded = SuperBlockCountKPadded * kSuperBlockBlks` is +// the K-block count used to address the scale buffer. +// * Callers must pass `SuperBlockCountKPadded` and `BlockCountKPadded` +// to these helpers; the original logical BlockCountK is only used +// for sizing the BlkSum buffer (which is consumed by the SGEMM +// correction step, not the inner K-loop). +// +// Caller-side constraints: BlkLen == 64; BlockCountK >= 1 (any K, padded +// internally to a multiple of kSuperBlockBlks). +// ----------------------------------------------------------------------------- + +inline size_t +PackedQuantBOffsetBytes_W2_SuperBlock(size_t n, size_t blk_super, + size_t SuperBlockCountKPadded, size_t NMain) +{ + if (n < NMain) { + const size_t g = n / kNCols4; + const size_t c = n % kNCols4; + const size_t per_group_bytes = SuperBlockCountKPadded * kNCols4 * kSuperBlockBytes; + return g * per_group_bytes + + blk_super * (kNCols4 * kSuperBlockBytes) + + c * kSuperBlockBytes; + } + return (n * SuperBlockCountKPadded + blk_super) * kSuperBlockBytes; +} + +// +// Float offset into the packed B-scale buffer for a logical (n, blk) cell. +// Scales remain per-block (one float per K-block), 4 per super. Caller +// passes BlockCountKPadded (= SuperBlockCountKPadded * kSuperBlockBlks); +// scale slots in [BlockCountK, BlockCountKPadded) contain zeros so the +// kernel can index uniformly. +// +inline size_t +PackedQuantBScaleOffset_W2_SuperBlock(size_t n, size_t blk, + size_t BlockCountKPadded, size_t NMain) +{ + const size_t SuperBlockCountKPadded = BlockCountKPadded / kSuperBlockBlks; + const size_t blk_super = blk / kSuperBlockBlks; + const size_t blk_in_super = blk % kSuperBlockBlks; + if (n < NMain) { + const size_t g = n / kNCols4; + const size_t c = n % kNCols4; + const size_t per_group_scales = SuperBlockCountKPadded * kNCols4 * kSuperBlockBlks; + return g * per_group_scales + + blk_super * (kNCols4 * kSuperBlockBlks) + + c * kSuperBlockBlks + + blk_in_super; + } + return n * BlockCountKPadded + blk; +} + +// +// Reference (scalar) entry points -- defined in sqnbitgemm_kernel_avx512_2bit_superblock.cpp. +// +// These cover Phase 2 of the super-block prototype: pack + scalar GEMM oracle +// against which the SIMD inner loop (Phase 3) will be validated. They are +// reachable from unit tests via direct linkage; production dispatch wiring +// happens in Phase 4 after the SIMD path is correctness-clean. +// + +size_t MLASCALL +Q2BitGemmPackQuantBDataSize_SuperBlock( + size_t N, + size_t K, + size_t BlkLen, + bool HasZeroPoint, + MLAS_QNBIT_GEMM_COMPUTE_TYPE ComputeType, + const MLAS_BACKEND_KERNEL_SELECTOR_CONFIG* BackendKernelSelectorConfig +); + +void MLASCALL +SQ2BitGemmPackQuantBDataAndBlkSum_SuperBlockScalar( + size_t N, + size_t K, + size_t BlkLen, + MLAS_QNBIT_GEMM_COMPUTE_TYPE ComputeType, + const std::byte* QuantBDataBegin, + const float* QuantBScaleBegin, + bool HasZeroPoint, + const std::byte* QuantBZPBegin, + PackedQuantBDataStruct& PackedQuantB, + MLAS_THREADPOOL* ThreadPool, + const MLAS_BACKEND_KERNEL_SELECTOR_CONFIG* BackendKernelSelectorConfig +); + +size_t MLASCALL +SQ2BitGemmKernel_BlkSum_CompInt8_SuperBlockScalar( + size_t BlkLen, + const std::byte* QuantA, + const float* QuantAScale, + const std::byte* QuantBData, + const float* QuantBScale, + const std::byte* QuantBZeroPoint, + float* C, + size_t CountM, + size_t CountN, + size_t CountK, + size_t BlockCountK, + const float* Bias, + size_t ldc, + const float* ABlockSum, + const float* QuantBBlkSum +); + +// +// Unit-test forwarders for the AVX-512 SIMD super-block kernels. Same gating +// rules as the existing W2 test entries: the caller MUST verify +// GetMlasPlatform().Avx512Supported_ (and, for the VNNI variant, that the +// active dispatch is the AVX-512-VNNI one) before invoking these symbols. +// +size_t MLASCALL +SQ2BitGemmKernel_BlkSum_CompInt8_Super_Avx512_TestEntry( + size_t BlkLen, + const std::byte* QuantA, + const float* QuantAScale, + const std::byte* QuantBData, + const float* QuantBScale, + const std::byte* QuantBZeroPoint, + float* C, + size_t CountM, + size_t CountN, + size_t CountK, + size_t BlockCountK, + const float* Bias, + size_t ldc, + const float* ABlockSum, + const float* QuantBBlkSum +); + +size_t MLASCALL +SQ2BitGemmKernel_BlkSum_CompInt8_Super_Avx512Vnni_TestEntry( + size_t BlkLen, + const std::byte* QuantA, + const float* QuantAScale, + const std::byte* QuantBData, + const float* QuantBScale, + const std::byte* QuantBZeroPoint, + float* C, + size_t CountM, + size_t CountN, + size_t CountK, + size_t BlockCountK, + const float* Bias, + size_t ldc, + const float* ABlockSum, + const float* QuantBBlkSum +); + +} // namespace sq2bit_avx512_super +} // namespace mlas +} // namespace onnxruntime diff --git a/onnxruntime/core/mlas/lib/sqnbitgemm_kernel_avx512vnni.cpp b/onnxruntime/core/mlas/lib/sqnbitgemm_kernel_avx512vnni.cpp index 4d7acd9db9a47..24c3fdbdcf159 100644 --- a/onnxruntime/core/mlas/lib/sqnbitgemm_kernel_avx512vnni.cpp +++ b/onnxruntime/core/mlas/lib/sqnbitgemm_kernel_avx512vnni.cpp @@ -29,6 +29,8 @@ Module Name: #include "sqnbitgemm_kernel_avx512_int8_blklen128.h" #include "sqnbitgemm_kernel_avx512_2bit.h" #include "sqnbitgemm_kernel_avx512vnni_2bit_blklen64.h" +#include "sqnbitgemm_kernel_avx512_2bit_superblock.h" +#include "sqnbitgemm_kernel_avx512vnni_2bit_blklen64_superblock.h" MLAS_FORCEINLINE void SQ4BitGemmM1Kernel_CompFp32( @@ -490,6 +492,37 @@ SQ2BitGemmKernel_BlkSum_CompInt8_Avx512Vnni_TestEntry( } } // namespace onnxruntime::mlas::sq2bit_avx512 +// +// Unit-test entry point for the AVX-512-VNNI W2 SUPER-BLOCK kernel +// (sqnbitgemm_kernel_avx512vnni_2bit_blklen64_superblock.h). Exposed for +// direct invocation from tests; production wiring (Phase 4) will add the +// runtime dispatcher integration. +// +namespace onnxruntime::mlas::sq2bit_avx512_super { +size_t MLASCALL +SQ2BitGemmKernel_BlkSum_CompInt8_Super_Avx512Vnni_TestEntry( + size_t BlkLen, + const std::byte* QuantA, + const float* QuantAScale, + const std::byte* QuantBData, + const float* QuantBScale, + const std::byte* QuantBZeroPoint, + float* C, + size_t CountM, + size_t CountN, + size_t CountK, + size_t BlockCountK, + const float* Bias, + size_t ldc, + const float* ABlockSum, + const float* QuantBBlkSum) +{ + return SQ2BitGemmKernel_BlkSum_CompInt8_Super_Avx512Vnni( + BlkLen, QuantA, QuantAScale, QuantBData, QuantBScale, QuantBZeroPoint, + C, CountM, CountN, CountK, BlockCountK, Bias, ldc, ABlockSum, QuantBBlkSum); +} +} // namespace onnxruntime::mlas::sq2bit_avx512_super + // // Kernel dispatch structure definition. // @@ -512,13 +545,24 @@ const MLAS_QNBIT_GEMM_DISPATCH MlasSQNBitGemmDispatchAvx512vnni = []() { d.SQ8BitGemmKernel_BlkSum_CompInt8 = SQ8BitGemmKernel_BlkSum_CompInt8_avx512vnni; d.QuantizeARowComputeBlkSum_CompInt8 = QuantizeARow_CompInt8_avx512; - // 2-bit native CompInt8 path: AVX-512-VNNI variant. Uses the same tile + - // pack layout as the non-VNNI variant (registered in the AVX-512 dispatch); - // the per-block MAC is `_mm512_dpbusd_epi32` instead of the - // `vpmaddubsw + vpmaddwd + vpaddd` chain. - d.Q2BitGemmPackQuantBDataSize = onnxruntime::mlas::sq2bit_avx512::Q2BitGemmPackQuantBDataSize_Avx512; - d.SQ2BitGemmPackQuantBDataAndBlkSum = onnxruntime::mlas::sq2bit_avx512::SQ2BitGemmPackQuantBDataAndBlkSum_Scalar; - d.SQ2BitGemmKernel_BlkSum_CompInt8 = onnxruntime::mlas::sq2bit_avx512::SQ2BitGemmKernel_BlkSum_CompInt8_Avx512Vnni; + // 2-bit native CompInt8 path: AVX-512-VNNI variant of the super-block + // (W2-v2) kernel -- single 64-byte load + four fixed shift/mask pairs + // to unpack 4 K-blocks at once. Pack-size and pack functions are shared + // with the AVX-512BW variant; the per-block MAC is `_mm512_dpbusd_epi32` + // instead of `vpmaddubsw + vpmaddwd + vpaddd`. + // + // The legacy `sq2bit_avx512::*` (W2-v1) symbols remain in the build but + // are no longer reached at runtime. A follow-up will remove them once + // W2-v2 has soaked in production. + d.Q2BitGemmPackQuantBDataSize = onnxruntime::mlas::sq2bit_avx512_super::Q2BitGemmPackQuantBDataSize_SuperBlock; + d.SQ2BitGemmPackQuantBDataAndBlkSum = onnxruntime::mlas::sq2bit_avx512_super::SQ2BitGemmPackQuantBDataAndBlkSum_SuperBlockScalar; + d.SQ2BitGemmKernel_BlkSum_CompInt8 = onnxruntime::mlas::sq2bit_avx512_super::SQ2BitGemmKernel_BlkSum_CompInt8_Super_Avx512Vnni_TestEntry; + // W2-v2 packs and addresses each N-col at a stride of + // SuperBlockCountKPadded * kSuperBlockBlks blocks (BlockCountK rounded + // UP to a multiple of 4) to keep the inner K-loop's super-block stride + // constant. The dispatcher needs this stride for its per-N-tile pointer + // arithmetic in SQ2BitGemm_CompInt8. + d.Q2BitGemmEffectiveBlockCountK = [](size_t BlockCountK) { return ((BlockCountK + 3) / 4) * 4; }; return d; }(); diff --git a/onnxruntime/core/mlas/lib/sqnbitgemm_kernel_avx512vnni_2bit_blklen64_superblock.h b/onnxruntime/core/mlas/lib/sqnbitgemm_kernel_avx512vnni_2bit_blklen64_superblock.h new file mode 100644 index 0000000000000..41ed3d8a3b6ef --- /dev/null +++ b/onnxruntime/core/mlas/lib/sqnbitgemm_kernel_avx512vnni_2bit_blklen64_superblock.h @@ -0,0 +1,864 @@ +/*++ + +Copyright (c) Microsoft Corporation. All rights reserved. + +Licensed under the MIT License. + +Module Name: + + sqnbitgemm_kernel_avx512vnni_2bit_blklen64_superblock.h + +Abstract: + + EXPERIMENTAL AVX-512 (-VNNI) W2 kernel that consumes the super-block packed + layout (sqnbitgemm_kernel_avx512_2bit_superblock.h). Replaces the existing + per-K-block broadcast + variable-shift unpack with one ZMM load and four + fixed-shift+mask pairs, halving the inner-loop unpack cost. + + Templated on `` like its sibling header + (sqnbitgemm_kernel_avx512vnni_2bit_blklen64.h): VNNI variant uses + `_mm512_dpbusd_epi32` for the integer MAC; non-VNNI uses the + `vpmaddubsw + vpmaddwd` chain. Both produce bit-identical results. + + Constraints (Phase 3): + * BlkLen == 64 only. + * BlockCountK must be a multiple of kSuperBlockBlks (= 4). + * CountM must be a multiple of 2 (R2 tile); CountN must be a multiple + of 4 (C4 tile). Tail handling is provided by the same R2xC1/R1xC4/ + R1xC1 helpers as the existing W2 kernel via a downgrade path the + dispatcher will pick when these constraints don't hold. + + Layout reference: + * 64-byte super-block: byte b holds 2-bit weight b from each of 4 + consecutive K-blocks at bit positions {0..1, 2..3, 4..5, 6..7}. + * In a tile slot, 4 N-cols of super-block live consecutively: each + N-col's super starts at offset c * kSuperBlockBytes within the slot. + +--*/ + +#pragma once + +#include +#include +#include + +#include + +#include "mlasi.h" +#include "qnbitgemm.h" +#include "sqnbitgemm_kernel_avx512_2bit.h" +#include "sqnbitgemm_kernel_avx512_2bit_superblock.h" + +namespace onnxruntime { +namespace mlas { +namespace sq2bit_avx512_super { + +inline constexpr size_t kNRows2 = 2; // matches the existing W2 R2 tile shape + +// +// Cheap super-block unpack: 1x ZMM load + 4x (fixed-shift + AND). +// Critical path ~4c (load + and / load + srli + and parallel chains). +// +// Bit layout of each byte b of `packed`: +// bits[0..1] = block_0.weight[b] +// bits[2..3] = block_1.weight[b] +// bits[4..5] = block_2.weight[b] +// bits[6..7] = block_3.weight[b] +// +// `_mm512_srli_epi16` shifts each 16-bit lane by N. For the LOW byte of a +// 16-bit lane, the shift pulls in bits from the adjacent HIGH byte; the +// subsequent AND with 0x03 discards those leaked bits. For the HIGH byte, +// zeros are shifted in from the top -- exactly what we want. +// +static MLAS_FORCEINLINE void +load_unpack_super_w2(const std::byte* packed, + __m512i& bv0_64_epi8, + __m512i& bv1_64_epi8, + __m512i& bv2_64_epi8, + __m512i& bv3_64_epi8) +{ + const __m512i super = _mm512_loadu_si512(reinterpret_cast(packed)); + const __m512i mask03 = _mm512_set1_epi8(0x03); + bv0_64_epi8 = _mm512_and_si512(super, mask03); + bv1_64_epi8 = _mm512_and_si512(_mm512_srli_epi16(super, 2), mask03); + bv2_64_epi8 = _mm512_and_si512(_mm512_srli_epi16(super, 4), mask03); + bv3_64_epi8 = _mm512_and_si512(_mm512_srli_epi16(super, 6), mask03); +} + +// +// Per single M-row, 4-K-block dot-and-accumulate. Each block produces its own +// uniformly-scaled FMA into a sub-accumulator; we keep two sub-accumulators +// (alternating per K-block) so the per-cell FMA dependency chain is two FMAs +// deep instead of four. The two sub-accumulators are summed into `acc` at the +// end with one extra vector add per super-block per cell. +// +// Math per K-block: acc += scale_a[blk] * scale_b[blk] * dot(av[blk], bv[blk]) +// +// scale_a and scale_b each point to 4 consecutive floats in their packed +// buffers (one per K-block of the super). +// +// Critical path analysis (Zen5 / SKX): +// * Single-chain (prior version): +// FMA latency * 4 = ~16 cycles per super-block per cell. +// * Two sub-accumulators (current): +// FMA latency * 2 = ~8 cycles per super-block per cell + one vaddps. +// This roughly halves the FP critical path; the integer dpbusd chain +// (one per block) is independent and runs in parallel with the FMAs. +// +template +static MLAS_FORCEINLINE void +dot_accumulate_4blk_w2_super(const __m512i& av0_64_epi8, const __m512i& av1_64_epi8, + const __m512i& av2_64_epi8, const __m512i& av3_64_epi8, + const __m512i& bv0_64_epi8, const __m512i& bv1_64_epi8, + const __m512i& bv2_64_epi8, const __m512i& bv3_64_epi8, + const float* scale_a, + const float* scale_b, + __m512& acc) +{ + __m512i d0, d1, d2, d3; + if constexpr (kVnni) { + d0 = _mm512_dpbusd_epi32(_mm512_setzero_si512(), bv0_64_epi8, av0_64_epi8); + d1 = _mm512_dpbusd_epi32(_mm512_setzero_si512(), bv1_64_epi8, av1_64_epi8); + d2 = _mm512_dpbusd_epi32(_mm512_setzero_si512(), bv2_64_epi8, av2_64_epi8); + d3 = _mm512_dpbusd_epi32(_mm512_setzero_si512(), bv3_64_epi8, av3_64_epi8); + } else { + // Non-VNNI: vpmaddubsw producing 32 lanes of int16, then vpmaddwd against + // a ones-vector reduces pairs to 16 lanes of int32. Same final layout as + // dpbusd, so the downstream cvt + FMA is identical. + const __m512i ones = _mm512_set1_epi16(1); + const __m512i t0 = _mm512_maddubs_epi16(bv0_64_epi8, av0_64_epi8); + const __m512i t1 = _mm512_maddubs_epi16(bv1_64_epi8, av1_64_epi8); + const __m512i t2 = _mm512_maddubs_epi16(bv2_64_epi8, av2_64_epi8); + const __m512i t3 = _mm512_maddubs_epi16(bv3_64_epi8, av3_64_epi8); + d0 = _mm512_madd_epi16(t0, ones); + d1 = _mm512_madd_epi16(t1, ones); + d2 = _mm512_madd_epi16(t2, ones); + d3 = _mm512_madd_epi16(t3, ones); + } + + // Pre-multiplied per-block scales (scalar broadcast, uniform across 16 lanes). + const __m512 s0 = _mm512_set1_ps(scale_a[0] * scale_b[0]); + const __m512 s1 = _mm512_set1_ps(scale_a[1] * scale_b[1]); + const __m512 s2 = _mm512_set1_ps(scale_a[2] * scale_b[2]); + const __m512 s3 = _mm512_set1_ps(scale_a[3] * scale_b[3]); + + // Two interleaved sub-accumulators: lo gets blocks {0, 2}, hi gets {1, 3}. + // Each sub-accumulator chain is 2 FMAs deep (~8c) vs the 4-FMA single chain. + __m512 acc_lo = _mm512_fmadd_ps(_mm512_cvtepi32_ps(d0), s0, _mm512_setzero_ps()); + __m512 acc_hi = _mm512_fmadd_ps(_mm512_cvtepi32_ps(d1), s1, _mm512_setzero_ps()); + acc_lo = _mm512_fmadd_ps(_mm512_cvtepi32_ps(d2), s2, acc_lo); + acc_hi = _mm512_fmadd_ps(_mm512_cvtepi32_ps(d3), s3, acc_hi); + acc = _mm512_add_ps(acc, _mm512_add_ps(acc_lo, acc_hi)); +} + +// +// 2 M-rows x 1 N-col x 4 K-blocks (one super-block) accumulator. The +// super-block B load + unpack is shared across the 2 M-rows. +// +template +static MLAS_FORCEINLINE void +accumulate_w2_blklen64_r2c1blk4_super( + const __m512i& av00, const __m512i& av01, const __m512i& av02, const __m512i& av03, + const __m512i& av10, const __m512i& av11, const __m512i& av12, const __m512i& av13, + const std::byte* QuantBDataPtr, + const float* scale_a0, + const float* scale_a1, + const float* scale_b, + __m512& acc0, + __m512& acc1) +{ + __m512i bv0, bv1, bv2, bv3; + load_unpack_super_w2(QuantBDataPtr, bv0, bv1, bv2, bv3); + + dot_accumulate_4blk_w2_super( + av00, av01, av02, av03, bv0, bv1, bv2, bv3, scale_a0, scale_b, acc0); + dot_accumulate_4blk_w2_super( + av10, av11, av12, av13, bv0, bv1, bv2, bv3, scale_a1, scale_b, acc1); +} + +// +// 1 M-row x 1 N-col x 4 K-blocks (one super-block) accumulator. Used by the +// R1xC4 tile for M=1 decode and as the trailing odd-row handler of R2xC4 +// when CountM is odd. +// +template +static MLAS_FORCEINLINE void +accumulate_w2_blklen64_r1c1blk4_super( + const __m512i& av00, const __m512i& av01, const __m512i& av02, const __m512i& av03, + const std::byte* QuantBDataPtr, + const float* scale_a0, + const float* scale_b, + __m512& acc0) +{ + __m512i bv0, bv1, bv2, bv3; + load_unpack_super_w2(QuantBDataPtr, bv0, bv1, bv2, bv3); + + dot_accumulate_4blk_w2_super( + av00, av01, av02, av03, bv0, bv1, bv2, bv3, scale_a0, scale_b, acc0); +} + +// +// R1 x C4 tile -- the M=1 decode path and the trailing odd-row handler for +// the R2xC4 tile when CountM is odd. Identical N-tile structure as R2xC4 +// (4 N-cols, super-block K stride) but processes a single M-row, so it uses +// half the registers (4 accumulators instead of 8) and half the MAC count +// per super-block iteration. +// +template +MLAS_FORCEINLINE void +Q2Int8GemmR1xC4BlkLen64Avx512_Super( + const std::byte* QuantA, + const float* QuantAScale, + const std::byte* QuantBData, + const float* QuantBScale, + float* C, + size_t CountM, // expected to be 1 (caller-enforced) + size_t CountN, + size_t BlockCountK, + const float* Bias, + size_t ldc) +{ + const size_t lda = BlockCountK * kBlkLen; + constexpr size_t PerColSuperBytes = kSuperBlockBytes; + constexpr size_t PerColSuperScale = kSuperBlockBlks; + constexpr size_t PerKSuperAdvanceBytes = kNCols4 * PerColSuperBytes; + constexpr size_t PerKSuperAdvanceScale = kNCols4 * PerColSuperScale; + // GroupStride uses the PADDED BlockCountK because the packed B layout + // walks N-groups at intervals of `SuperBlockCountKPadded * kNCols4 * + // kSuperBlockBytes` (see PackedQuantBOffsetBytes_W2_SuperBlock). When + // BlockCountK is a multiple of kSuperBlockBlks (== 4) the padded and + // logical strides are identical; when it isn't, the kernel must step + // past the padded slots to land on the next N-group correctly. + const size_t SuperBlockCountKPadded = + MlasDivRoundup(BlockCountK, kSuperBlockBlks); + const size_t BlockCountKPadded = SuperBlockCountKPadded * kSuperBlockBlks; + const size_t GroupStrideBytes = BlockCountKPadded * kNCols4 * kBlkBytes; + const size_t GroupStrideScale = BlockCountKPadded * kNCols4; + + assert(CountN % kNCols4 == 0); + // BlockCountK no longer required to be a multiple of kSuperBlockBlks: + // the main K-loop iterates full supers; an optional tail handler picks up + // the trailing 1-3 K-blocks (padded weights and scales contribute 0). + const size_t FullSupers = BlockCountK / kSuperBlockBlks; + const size_t TailBlocks = BlockCountK % kSuperBlockBlks; // 0, 1, 2, or 3 + + for (size_t m = 0; m < CountM; ++m) { + const std::byte* QuantBDataColPtr = QuantBData; + const float* QuantBScaleColPtr = QuantBScale; + const float* BiasPtr = Bias; + float* SumPtr = C + m * ldc; + + for (size_t n = 0; n < CountN; n += kNCols4) { + const std::byte* QuantAPtr = QuantA + m * lda; + const float* QuantAScalePtr = QuantAScale + m * BlockCountK; + + const std::byte* QuantBDataPtr = QuantBDataColPtr; + const float* QuantBScalePtr = QuantBScaleColPtr; + + __m512 acc[kNCols4] = { + _mm512_setzero_ps(), _mm512_setzero_ps(), + _mm512_setzero_ps(), _mm512_setzero_ps() + }; + + for (size_t sb = 0; sb < FullSupers; ++sb) { + const __m512i av00 = _mm512_loadu_si512( + reinterpret_cast(QuantAPtr)); + const __m512i av01 = _mm512_loadu_si512( + reinterpret_cast(QuantAPtr + kBlkLen)); + const __m512i av02 = _mm512_loadu_si512( + reinterpret_cast(QuantAPtr + 2 * kBlkLen)); + const __m512i av03 = _mm512_loadu_si512( + reinterpret_cast(QuantAPtr + 3 * kBlkLen)); + + accumulate_w2_blklen64_r1c1blk4_super( + av00, av01, av02, av03, + QuantBDataPtr + 0 * PerColSuperBytes, + QuantAScalePtr, + QuantBScalePtr + 0 * PerColSuperScale, + acc[0]); + accumulate_w2_blklen64_r1c1blk4_super( + av00, av01, av02, av03, + QuantBDataPtr + 1 * PerColSuperBytes, + QuantAScalePtr, + QuantBScalePtr + 1 * PerColSuperScale, + acc[1]); + accumulate_w2_blklen64_r1c1blk4_super( + av00, av01, av02, av03, + QuantBDataPtr + 2 * PerColSuperBytes, + QuantAScalePtr, + QuantBScalePtr + 2 * PerColSuperScale, + acc[2]); + accumulate_w2_blklen64_r1c1blk4_super( + av00, av01, av02, av03, + QuantBDataPtr + 3 * PerColSuperBytes, + QuantAScalePtr, + QuantBScalePtr + 3 * PerColSuperScale, + acc[3]); + + QuantAPtr += kBlkLen * kSuperBlockBlks; + QuantAScalePtr += kSuperBlockBlks; + QuantBDataPtr += PerKSuperAdvanceBytes; + QuantBScalePtr += PerKSuperAdvanceScale; + } + + // K-tail: 1-3 trailing real K-blocks. Pack helper zero-padded the + // missing K-block slots in B and the corresponding scale slots, + // so the 4-block accumulator can run safely on the packed buffer. + // + // Two safety concerns: + // * A bytes: don't load past row end -- use zero ZMM for missing slots. + // * A scales: don't read past the row's logical BlockCountK scales -- + // uninitialised memory there can contain NaN, which propagates + // through `0 * NaN = NaN` in the scale fmadd. We materialise a + // local 4-float scale_a buffer with real scales in [0, TailBlocks) + // and 0.0 in the trailing slots. + if (TailBlocks > 0) { + const __m512i zero = _mm512_setzero_si512(); + const __m512i av00 = _mm512_loadu_si512( + reinterpret_cast(QuantAPtr + 0 * kBlkLen)); + const __m512i av01 = (TailBlocks > 1) + ? _mm512_loadu_si512(reinterpret_cast(QuantAPtr + 1 * kBlkLen)) + : zero; + const __m512i av02 = (TailBlocks > 2) + ? _mm512_loadu_si512(reinterpret_cast(QuantAPtr + 2 * kBlkLen)) + : zero; + const __m512i av03 = zero; // TailBlocks at most 3 + + // Bounded scale_a copy. + float scale_a0_safe[kSuperBlockBlks] = {0.0f, 0.0f, 0.0f, 0.0f}; + for (size_t i = 0; i < TailBlocks; ++i) { + scale_a0_safe[i] = QuantAScalePtr[i]; + } + + accumulate_w2_blklen64_r1c1blk4_super( + av00, av01, av02, av03, + QuantBDataPtr + 0 * PerColSuperBytes, + scale_a0_safe, + QuantBScalePtr + 0 * PerColSuperScale, + acc[0]); + accumulate_w2_blklen64_r1c1blk4_super( + av00, av01, av02, av03, + QuantBDataPtr + 1 * PerColSuperBytes, + scale_a0_safe, + QuantBScalePtr + 1 * PerColSuperScale, + acc[1]); + accumulate_w2_blklen64_r1c1blk4_super( + av00, av01, av02, av03, + QuantBDataPtr + 2 * PerColSuperBytes, + scale_a0_safe, + QuantBScalePtr + 2 * PerColSuperScale, + acc[2]); + accumulate_w2_blklen64_r1c1blk4_super( + av00, av01, av02, av03, + QuantBDataPtr + 3 * PerColSuperBytes, + scale_a0_safe, + QuantBScalePtr + 3 * PerColSuperScale, + acc[3]); + } + + SumPtr[0] = _mm512_reduce_add_ps(acc[0]); + SumPtr[1] = _mm512_reduce_add_ps(acc[1]); + SumPtr[2] = _mm512_reduce_add_ps(acc[2]); + SumPtr[3] = _mm512_reduce_add_ps(acc[3]); + if (BiasPtr != nullptr) { + SumPtr[0] += BiasPtr[0]; + SumPtr[1] += BiasPtr[1]; + SumPtr[2] += BiasPtr[2]; + SumPtr[3] += BiasPtr[3]; + } + + QuantBDataColPtr += GroupStrideBytes; + QuantBScaleColPtr += GroupStrideScale; + BiasPtr += BiasPtr != nullptr ? kNCols4 : 0; + SumPtr += kNCols4; + } + } +} + +// +// R2 x C4 tile -- the main hot path for prefill (M >= 2). Iterates the K +// dimension in super-block strides of kSuperBlockBlks (= 4) K-blocks at a +// time. Assumes BlockCountK is a multiple of kSuperBlockBlks; the dispatcher +// must verify this before selecting the super-block kernel. +// +template +MLAS_FORCEINLINE void +Q2Int8GemmR2xC4BlkLen64Avx512_Super( + const std::byte* QuantA, + const float* QuantAScale, + const std::byte* QuantBData, + const float* QuantBScale, + float* C, + size_t CountM, + size_t CountN, + size_t BlockCountK, + const float* Bias, + size_t ldc) +{ + const size_t lda = BlockCountK * kBlkLen; + constexpr size_t PerColSuperBytes = kSuperBlockBytes; // 64 B per col per super + constexpr size_t PerColSuperScale = kSuperBlockBlks; // 4 scales per col per super + constexpr size_t PerKSuperAdvanceBytes = kNCols4 * PerColSuperBytes; // 256 B per K-super iter + constexpr size_t PerKSuperAdvanceScale = kNCols4 * PerColSuperScale; // 16 scales per K-super iter + // GroupStride uses the PADDED BlockCountK because the packed B layout + // walks N-groups at intervals of `SuperBlockCountKPadded * kNCols4 * + // kSuperBlockBytes` (see PackedQuantBOffsetBytes_W2_SuperBlock). When + // BlockCountK is a multiple of kSuperBlockBlks (== 4) the padded and + // logical strides are identical; when it isn't, the kernel must step + // past the padded slots to land on the next N-group correctly. + const size_t SuperBlockCountKPadded = + MlasDivRoundup(BlockCountK, kSuperBlockBlks); + const size_t BlockCountKPadded = SuperBlockCountKPadded * kSuperBlockBlks; + const size_t GroupStrideBytes = BlockCountKPadded * kNCols4 * kBlkBytes; + const size_t GroupStrideScale = BlockCountKPadded * kNCols4; + + assert(CountM % kNRows2 == 0); + assert(CountN % kNCols4 == 0); + // BlockCountK no longer required to be a multiple of kSuperBlockBlks: + // the main K-loop iterates full supers; an optional tail handler picks up + // the trailing 1-3 K-blocks (padded weights and scales contribute 0). + const size_t FullSupers = BlockCountK / kSuperBlockBlks; + const size_t TailBlocks = BlockCountK % kSuperBlockBlks; // 0, 1, 2, or 3 + + for (size_t m = 0; m < CountM; m += kNRows2) { + const std::byte* QuantBDataColPtr = QuantBData; + const float* QuantBScaleColPtr = QuantBScale; + const float* BiasPtr = Bias; + float* SumPtr = C + m * ldc; + + for (size_t n = 0; n < CountN; n += kNCols4) { + const std::byte* QuantAPtr = QuantA + m * lda; + const float* QuantAScalePtr = QuantAScale + m * BlockCountK; + + const std::byte* QuantBDataPtr = QuantBDataColPtr; + const float* QuantBScalePtr = QuantBScaleColPtr; + + __m512 acc[kNCols4 * kNRows2] = { + _mm512_setzero_ps(), _mm512_setzero_ps(), _mm512_setzero_ps(), _mm512_setzero_ps(), + _mm512_setzero_ps(), _mm512_setzero_ps(), _mm512_setzero_ps(), _mm512_setzero_ps() + }; + + for (size_t sb = 0; sb < FullSupers; ++sb) { + const __m512i av00 = _mm512_loadu_si512( + reinterpret_cast(QuantAPtr)); + const __m512i av01 = _mm512_loadu_si512( + reinterpret_cast(QuantAPtr + kBlkLen)); + const __m512i av02 = _mm512_loadu_si512( + reinterpret_cast(QuantAPtr + 2 * kBlkLen)); + const __m512i av03 = _mm512_loadu_si512( + reinterpret_cast(QuantAPtr + 3 * kBlkLen)); + const __m512i av10 = _mm512_loadu_si512( + reinterpret_cast(QuantAPtr + lda)); + const __m512i av11 = _mm512_loadu_si512( + reinterpret_cast(QuantAPtr + lda + kBlkLen)); + const __m512i av12 = _mm512_loadu_si512( + reinterpret_cast(QuantAPtr + lda + 2 * kBlkLen)); + const __m512i av13 = _mm512_loadu_si512( + reinterpret_cast(QuantAPtr + lda + 3 * kBlkLen)); + + accumulate_w2_blklen64_r2c1blk4_super( + av00, av01, av02, av03, av10, av11, av12, av13, + QuantBDataPtr + 0 * PerColSuperBytes, + QuantAScalePtr, QuantAScalePtr + BlockCountK, + QuantBScalePtr + 0 * PerColSuperScale, + acc[0], acc[kNCols4 + 0]); + accumulate_w2_blklen64_r2c1blk4_super( + av00, av01, av02, av03, av10, av11, av12, av13, + QuantBDataPtr + 1 * PerColSuperBytes, + QuantAScalePtr, QuantAScalePtr + BlockCountK, + QuantBScalePtr + 1 * PerColSuperScale, + acc[1], acc[kNCols4 + 1]); + accumulate_w2_blklen64_r2c1blk4_super( + av00, av01, av02, av03, av10, av11, av12, av13, + QuantBDataPtr + 2 * PerColSuperBytes, + QuantAScalePtr, QuantAScalePtr + BlockCountK, + QuantBScalePtr + 2 * PerColSuperScale, + acc[2], acc[kNCols4 + 2]); + accumulate_w2_blklen64_r2c1blk4_super( + av00, av01, av02, av03, av10, av11, av12, av13, + QuantBDataPtr + 3 * PerColSuperBytes, + QuantAScalePtr, QuantAScalePtr + BlockCountK, + QuantBScalePtr + 3 * PerColSuperScale, + acc[3], acc[kNCols4 + 3]); + + QuantAPtr += kBlkLen * kSuperBlockBlks; + QuantAScalePtr += kSuperBlockBlks; + QuantBDataPtr += PerKSuperAdvanceBytes; + QuantBScalePtr += PerKSuperAdvanceScale; + } + + // K-tail: 1-3 trailing real K-blocks. See R1 tile comment above. + if (TailBlocks > 0) { + const __m512i zero = _mm512_setzero_si512(); + const __m512i av00 = _mm512_loadu_si512( + reinterpret_cast(QuantAPtr + 0 * kBlkLen)); + const __m512i av01 = (TailBlocks > 1) + ? _mm512_loadu_si512(reinterpret_cast(QuantAPtr + 1 * kBlkLen)) + : zero; + const __m512i av02 = (TailBlocks > 2) + ? _mm512_loadu_si512(reinterpret_cast(QuantAPtr + 2 * kBlkLen)) + : zero; + const __m512i av03 = zero; + const __m512i av10 = _mm512_loadu_si512( + reinterpret_cast(QuantAPtr + lda + 0 * kBlkLen)); + const __m512i av11 = (TailBlocks > 1) + ? _mm512_loadu_si512(reinterpret_cast(QuantAPtr + lda + 1 * kBlkLen)) + : zero; + const __m512i av12 = (TailBlocks > 2) + ? _mm512_loadu_si512(reinterpret_cast(QuantAPtr + lda + 2 * kBlkLen)) + : zero; + const __m512i av13 = zero; + + // Bounded scale_a copies for both M-rows (see R1 K-tail comment). + float scale_a0_safe[kSuperBlockBlks] = {0.0f, 0.0f, 0.0f, 0.0f}; + float scale_a1_safe[kSuperBlockBlks] = {0.0f, 0.0f, 0.0f, 0.0f}; + for (size_t i = 0; i < TailBlocks; ++i) { + scale_a0_safe[i] = QuantAScalePtr[i]; + scale_a1_safe[i] = QuantAScalePtr[BlockCountK + i]; + } + + accumulate_w2_blklen64_r2c1blk4_super( + av00, av01, av02, av03, av10, av11, av12, av13, + QuantBDataPtr + 0 * PerColSuperBytes, + scale_a0_safe, scale_a1_safe, + QuantBScalePtr + 0 * PerColSuperScale, + acc[0], acc[kNCols4 + 0]); + accumulate_w2_blklen64_r2c1blk4_super( + av00, av01, av02, av03, av10, av11, av12, av13, + QuantBDataPtr + 1 * PerColSuperBytes, + scale_a0_safe, scale_a1_safe, + QuantBScalePtr + 1 * PerColSuperScale, + acc[1], acc[kNCols4 + 1]); + accumulate_w2_blklen64_r2c1blk4_super( + av00, av01, av02, av03, av10, av11, av12, av13, + QuantBDataPtr + 2 * PerColSuperBytes, + scale_a0_safe, scale_a1_safe, + QuantBScalePtr + 2 * PerColSuperScale, + acc[2], acc[kNCols4 + 2]); + accumulate_w2_blklen64_r2c1blk4_super( + av00, av01, av02, av03, av10, av11, av12, av13, + QuantBDataPtr + 3 * PerColSuperBytes, + scale_a0_safe, scale_a1_safe, + QuantBScalePtr + 3 * PerColSuperScale, + acc[3], acc[kNCols4 + 3]); + } + + SumPtr[0] = _mm512_reduce_add_ps(acc[0]); + SumPtr[1] = _mm512_reduce_add_ps(acc[1]); + SumPtr[2] = _mm512_reduce_add_ps(acc[2]); + SumPtr[3] = _mm512_reduce_add_ps(acc[3]); + SumPtr[ldc + 0] = _mm512_reduce_add_ps(acc[kNCols4 + 0]); + SumPtr[ldc + 1] = _mm512_reduce_add_ps(acc[kNCols4 + 1]); + SumPtr[ldc + 2] = _mm512_reduce_add_ps(acc[kNCols4 + 2]); + SumPtr[ldc + 3] = _mm512_reduce_add_ps(acc[kNCols4 + 3]); + if (BiasPtr != nullptr) { + SumPtr[0] += BiasPtr[0]; + SumPtr[1] += BiasPtr[1]; + SumPtr[2] += BiasPtr[2]; + SumPtr[3] += BiasPtr[3]; + SumPtr[ldc + 0] += BiasPtr[0]; + SumPtr[ldc + 1] += BiasPtr[1]; + SumPtr[ldc + 2] += BiasPtr[2]; + SumPtr[ldc + 3] += BiasPtr[3]; + } + + QuantBDataColPtr += GroupStrideBytes; + QuantBScaleColPtr += GroupStrideScale; + BiasPtr += BiasPtr != nullptr ? kNCols4 : 0; + SumPtr += kNCols4; + } + } +} + +// +// 1 M-row x 1 N-col N-tail tile. Handles the 1-3 trailing N-cols when +// CountN is not a multiple of kNCols4. The tail region of the packed B +// buffer is column-major (one super-block per K-super per N-col, see +// PackedQuantBOffsetBytes_W2_SuperBlock for n >= NMain), so this tile +// walks one column at a time and reuses the same accumulator helper used +// by R1xC4 (`accumulate_w2_blklen64_r1c1blk4_super`). Slower than the +// R2xC4 main tile, but it processes at most 3 N-cols per call -- a +// trivial fraction of total work even on the worst-case shape. +// +// Pointer convention (caller-supplied bases): +// QuantBDataTail : start of the tail region in packed B -- exactly +// NMain * SuperBlockCountKPadded * kSuperBlockBytes +// bytes past PackedQuantBData (see callsite below). +// QuantBScaleTail : same convention for the scale buffer (NMain * +// SuperBlockCountKPadded * kSuperBlockBlks floats). +// +// K-tail handling: identical to R1xC4 -- conditional A loads for the 1-3 +// trailing real K-blocks and a bounded scale_a copy to avoid NaN from +// uninitialised QuantAScale slots. +// +template +MLAS_FORCEINLINE void +Q2Int8GemmRMxC_Tail_BlkLen64Avx512_Super( + const std::byte* QuantA, + const float* QuantAScale, + const std::byte* QuantBDataTail, + const float* QuantBScaleTail, + float* C, + size_t CountM, + size_t TailN, // 1, 2, or 3 + size_t BlockCountK, + const float* BiasTail, // null OR points at the first tail N-col bias + size_t ldc) +{ + assert(TailN >= 1 && TailN <= 3); + constexpr size_t PerColSuperBytes = kSuperBlockBytes; // 64 B per col per super + constexpr size_t PerColSuperScale = kSuperBlockBlks; // 4 scales per col per super + + const size_t lda = BlockCountK * kBlkLen; + const size_t SuperBlockCountKPadded = + MlasDivRoundup(BlockCountK, kSuperBlockBlks); + const size_t FullSupers = BlockCountK / kSuperBlockBlks; + const size_t TailBlocks = BlockCountK % kSuperBlockBlks; // 0, 1, 2, or 3 + // In the tail region each N-col occupies SuperBlockCountKPadded + // super-blocks back-to-back (column-major). + const size_t ColStrideBytes = SuperBlockCountKPadded * PerColSuperBytes; + const size_t ColStrideScale = SuperBlockCountKPadded * PerColSuperScale; + + for (size_t m = 0; m < CountM; ++m) { + for (size_t c = 0; c < TailN; ++c) { + const std::byte* QuantAPtr = QuantA + m * lda; + const float* QuantAScalePtr = QuantAScale + m * BlockCountK; + + const std::byte* QuantBDataPtr = QuantBDataTail + c * ColStrideBytes; + const float* QuantBScalePtr = QuantBScaleTail + c * ColStrideScale; + + __m512 acc = _mm512_setzero_ps(); + + for (size_t sb = 0; sb < FullSupers; ++sb) { + const __m512i av00 = _mm512_loadu_si512( + reinterpret_cast(QuantAPtr)); + const __m512i av01 = _mm512_loadu_si512( + reinterpret_cast(QuantAPtr + kBlkLen)); + const __m512i av02 = _mm512_loadu_si512( + reinterpret_cast(QuantAPtr + 2 * kBlkLen)); + const __m512i av03 = _mm512_loadu_si512( + reinterpret_cast(QuantAPtr + 3 * kBlkLen)); + + accumulate_w2_blklen64_r1c1blk4_super( + av00, av01, av02, av03, + QuantBDataPtr, + QuantAScalePtr, + QuantBScalePtr, + acc); + + QuantAPtr += kBlkLen * kSuperBlockBlks; + QuantAScalePtr += kSuperBlockBlks; + QuantBDataPtr += PerColSuperBytes; + QuantBScalePtr += PerColSuperScale; + } + + if (TailBlocks > 0) { + const __m512i zero = _mm512_setzero_si512(); + const __m512i av00 = _mm512_loadu_si512( + reinterpret_cast(QuantAPtr + 0 * kBlkLen)); + const __m512i av01 = (TailBlocks > 1) + ? _mm512_loadu_si512(reinterpret_cast(QuantAPtr + 1 * kBlkLen)) + : zero; + const __m512i av02 = (TailBlocks > 2) + ? _mm512_loadu_si512(reinterpret_cast(QuantAPtr + 2 * kBlkLen)) + : zero; + const __m512i av03 = zero; + + float scale_a0_safe[kSuperBlockBlks] = {0.0f, 0.0f, 0.0f, 0.0f}; + for (size_t i = 0; i < TailBlocks; ++i) { + scale_a0_safe[i] = QuantAScalePtr[i]; + } + + accumulate_w2_blklen64_r1c1blk4_super( + av00, av01, av02, av03, + QuantBDataPtr, + scale_a0_safe, + QuantBScalePtr, + acc); + } + + float* SumPtr = C + m * ldc + c; + float v = _mm512_reduce_add_ps(acc); + if (BiasTail != nullptr) v += BiasTail[c]; + *SumPtr = v; + } + } +} + +// +// Top-level dispatched-kernel body. Mirrors the production W2 kernel's +// `SQ2BitGemmKernel_BlkSum_CompInt8_Impl` and the helper-mediated SGEMM +// correction step. Templated on ``; the AVX-512 and AVX-512-VNNI +// .cpp files instantiate it via test-entry forwarders. +// +// Restrictions enforced here: +// * BlkLen must equal kBlkLen (64). Otherwise returns 0 to signal "did not +// handle these rows" (the dispatcher will fall back). +// * CountM has no alignment requirement: the R2xC4 tile handles the +// M-aligned head and a single R1xC4 invocation picks up any trailing +// odd row. CountM == 1 (decode) lands directly on the R1 path. +// * CountN has no alignment requirement: the R2/R1 tiles process +// NMain = floor(CountN/4)*4 cols against the 4-N-col-grouped packed +// layout; a per-1-col tail tile picks up the trailing 1-3 cols against +// the column-major tail region of the same packed buffer. +// * BlockCountK has no alignment requirement: the R2/R1 tiles and the +// N-tail tile each run a partial-super K-tail handler that loads only +// the real trailing K-blocks (zero ZMM for missing slots, bounded +// scale_a copy) and lets the pre-zeroed packed-B / scale slots +// contribute 0 to the dot product. +// +template +static MLAS_FORCEINLINE size_t +SQ2BitGemmKernel_BlkSum_CompInt8_Super_Impl( + const size_t BlkLen, + const std::byte* QuantA, + const float* QuantAScale, + const std::byte* QuantBData, + const float* QuantBScale, + const std::byte* /* QuantBZeroPoint */, + float* C, + size_t CountM, + size_t CountN, + size_t /* CountK */, + size_t BlockCountK, + const float* Bias, + size_t ldc, + const float* ABlockSum, + const float* QuantBBlkSum) +{ + if (BlkLen != kBlkLen) { + return 0; + } + if (BlockCountK == 0 || CountM == 0 || CountN == 0) { + return 0; + } + + // Split CountN into a 4-aligned head (NMain) handled by the R2xC4/R1xC4 + // tiles and a 1-3-col tail handled by Q2Int8GemmRMxC_Tail. + const size_t NMain = (CountN / kNCols4) * kNCols4; + const size_t NTail = CountN - NMain; + + // Split CountM into an R2 head and an optional R1 tail row. + const size_t M_pairs = CountM / kNRows2; + const size_t M_main = M_pairs * kNRows2; + const size_t M_tail = CountM - M_main; + const size_t lda = BlockCountK * kBlkLen; + + if (NMain > 0) { + if (M_main > 0) { + Q2Int8GemmR2xC4BlkLen64Avx512_Super( + QuantA, QuantAScale, QuantBData, QuantBScale, + C, M_main, NMain, BlockCountK, Bias, ldc); + } + if (M_tail > 0) { + // R1 picks up the single trailing row. Pointers advance past the M_main + // rows the R2 tile already consumed: A advances M_main*lda bytes, A-scale + // advances M_main*BlockCountK floats, C advances M_main*ldc floats. The + // packed-B buffer is reused (column-major over N). + Q2Int8GemmR1xC4BlkLen64Avx512_Super( + QuantA + M_main * lda, + QuantAScale + M_main * BlockCountK, + QuantBData, QuantBScale, + C + M_main * ldc, + /*CountM=*/M_tail, + NMain, BlockCountK, Bias, ldc); + } + } + + if (NTail > 0) { + // The tail region of the packed B buffer is column-major and starts + // immediately after the NMain-cols grouped region: + // tail_base_bytes = NMain * SuperBlockCountKPadded * kSuperBlockBytes + // tail_base_scales = NMain * SuperBlockCountKPadded * kSuperBlockBlks + const size_t SuperBlockCountKPadded = + MlasDivRoundup(BlockCountK, kSuperBlockBlks); + const std::byte* QuantBDataTail = + QuantBData + NMain * SuperBlockCountKPadded * kSuperBlockBytes; + const float* QuantBScaleTail = + QuantBScale + NMain * SuperBlockCountKPadded * kSuperBlockBlks; + const float* BiasTail = (Bias != nullptr) ? Bias + NMain : nullptr; + + Q2Int8GemmRMxC_Tail_BlkLen64Avx512_Super( + QuantA, QuantAScale, + QuantBDataTail, QuantBScaleTail, + C + NMain, + CountM, NTail, BlockCountK, BiasTail, ldc); + } + + // BlkSum correction: same width-16 chunked layout as the existing W2 + // kernel, so we reuse the production SGEMM micro-kernel. The BlkSum + // buffer covers ALL N (including the tail), so this runs once over + // the full CountN. + float* c_blk = C; + const float* b_blk_sum = QuantBBlkSum; + size_t RowsRemaining = CountM; + const float* a_blksum_row = ABlockSum; + while (RowsRemaining > 0) { + const auto RowsHandled = GetMlasPlatform().GemmFloatKernel( + a_blksum_row, b_blk_sum, c_blk, + BlockCountK, RowsRemaining, CountN, BlockCountK, ldc, + 1.0f, /*ZeroMode=*/false); + + c_blk += ldc * RowsHandled; + a_blksum_row += BlockCountK * RowsHandled; + RowsRemaining -= RowsHandled; + } + return CountM; +} + +// +// Top-level VNNI variant. Compiled into AVX-512-VNNI sources only. +// +static MLAS_FORCEINLINE size_t MLASCALL +SQ2BitGemmKernel_BlkSum_CompInt8_Super_Avx512Vnni( + const size_t BlkLen, + const std::byte* QuantA, + const float* QuantAScale, + const std::byte* QuantBData, + const float* QuantBScale, + const std::byte* QuantBZeroPoint, + float* C, + size_t CountM, + size_t CountN, + size_t CountK, + size_t BlockCountK, + const float* Bias, + size_t ldc, + const float* ABlockSum, + const float* QuantBBlkSum) +{ + return SQ2BitGemmKernel_BlkSum_CompInt8_Super_Impl( + BlkLen, QuantA, QuantAScale, QuantBData, QuantBScale, QuantBZeroPoint, + C, CountM, CountN, CountK, BlockCountK, Bias, ldc, ABlockSum, QuantBBlkSum); +} + +// +// Top-level non-VNNI variant. Same tile + layout; integer MAC uses the +// vpmaddubsw + vpmaddwd chain instead of dpbusd. +// +static MLAS_FORCEINLINE size_t MLASCALL +SQ2BitGemmKernel_BlkSum_CompInt8_Super_Avx512( + const size_t BlkLen, + const std::byte* QuantA, + const float* QuantAScale, + const std::byte* QuantBData, + const float* QuantBScale, + const std::byte* QuantBZeroPoint, + float* C, + size_t CountM, + size_t CountN, + size_t CountK, + size_t BlockCountK, + const float* Bias, + size_t ldc, + const float* ABlockSum, + const float* QuantBBlkSum) +{ + return SQ2BitGemmKernel_BlkSum_CompInt8_Super_Impl( + BlkLen, QuantA, QuantAScale, QuantBData, QuantBScale, QuantBZeroPoint, + C, CountM, CountN, CountK, BlockCountK, Bias, ldc, ABlockSum, QuantBBlkSum); +} + +} // namespace sq2bit_avx512_super +} // namespace mlas +} // namespace onnxruntime diff --git a/onnxruntime/test/mlas/bench/bench_lutgemm.cpp b/onnxruntime/test/mlas/bench/bench_lutgemm.cpp index 235916fd20203..f021b36d8d23f 100644 --- a/onnxruntime/test/mlas/bench/bench_lutgemm.cpp +++ b/onnxruntime/test/mlas/bench/bench_lutgemm.cpp @@ -239,23 +239,26 @@ static void LutGemmComputeArgs(benchmark::internal::Benchmark* b) { // (K=1024, N=384): 20 nodes // (K=1024, N=4096): 20 nodes // (K=4096, N=1024): 20 nodes -// M=128 is the widest-gap prefill shape vs the W4 CompInt8 path. -// Pair with QNBITGEMM/QNBitGemmCustomerArgs for head-to-head. +// Covers both M=1 (decode) and M=128 (prefill) so the LUT path can be +// compared apples-to-apples against the W4 CompInt8 and W2-prod/W2-super +// kernels (QNBITGEMM/QNBitGemmCustomerArgs and +// QNBITGEMM/QNBit2BitCustomerArgs). static void LutGemmCustomerArgs(benchmark::internal::Benchmark* b) { b->ArgNames(lutgemm_compute_arg_names); - // Five separate Args() entries (rather than ArgsProduct) so we only run the - // exact (K, N) pairs that appear in the customer model. - const int64_t M = 128; + // Separate Args() entries so we only run the exact (M, K, N) tuples that + // appear in the customer model. const int64_t BlkLen = 64; const int64_t Threads = 8; const int64_t HasZP = 0; const int64_t HasBias = 1; - for (auto kn : {std::pair{384, 1024}, - std::pair{1024, 192}, - std::pair{1024, 384}, - std::pair{1024, 4096}, - std::pair{4096, 1024}}) { - b->Args({BlkLen, M, kn.second, kn.first, Threads, HasZP, HasBias}); + for (int64_t M : {int64_t{1}, int64_t{128}}) { + for (auto kn : {std::pair{384, 1024}, + std::pair{1024, 192}, + std::pair{1024, 384}, + std::pair{1024, 4096}, + std::pair{4096, 1024}}) { + b->Args({BlkLen, M, kn.second, kn.first, Threads, HasZP, HasBias}); + } } } diff --git a/onnxruntime/test/mlas/bench/bench_qnbitgemm.cpp b/onnxruntime/test/mlas/bench/bench_qnbitgemm.cpp index 6e2f50f33d81a..00fe98d604622 100644 --- a/onnxruntime/test/mlas/bench/bench_qnbitgemm.cpp +++ b/onnxruntime/test/mlas/bench/bench_qnbitgemm.cpp @@ -17,6 +17,14 @@ #include "core/util/thread_utils.h" #include "core/platform/env_var_utils.h" +// Prototype W2 super-block kernel + scalar pack helper (Phase 3 of the +// W2-vs-W4 parity work). Not yet wired into the platform dispatch, so the +// bench drives it directly via the test-entry forwarder. +#include "core/mlas/lib/mlasi.h" +#include "core/mlas/lib/qnbitgemm.h" +#include "core/mlas/lib/sqnbitgemm_kernel_avx512_2bit.h" +#include "core/mlas/lib/sqnbitgemm_kernel_avx512_2bit_superblock.h" + template void RunQNBitGemmBenchmark(size_t BlkLen, size_t M, size_t N, size_t K, @@ -192,12 +200,194 @@ static void QNBit2BitCustomerArgs(benchmark::internal::Benchmark* b) { BENCHMARK(QNBITGEMM)->Apply(QNBit2BitCustomerArgs)->UseRealTime(); -// 8-bit weight rows on the customer shapes. Used to confirm whether the W2-vs-W4 -// gap is driven by B-weight unpacking cost: W8 has zero unpacking (one byte per -// weight, direct vmovdqu8 + dpbusd), W4 has cheap nibble extraction, W2 has the -// most expensive unpack path. If unpack-density is the bottleneck, expect -// W8 < W4 < W2 in per-MAC cycles at the larger N shapes. -BENCHMARK(QNBITGEMM)->Apply(QNBit2BitCustomerArgs)->UseRealTime(); +// --------------------------------------------------------------------------- +// W2 SUPER-BLOCK PROTOTYPE BENCHMARK +// --------------------------------------------------------------------------- +// Drives the Phase-3 super-block W2 kernel directly via its test-entry +// forwarder. Mirrors what MlasQNBitGemmBatch would do for our path: pre-pack +// B via the super-block helpers, quantize each A row using the dispatch's +// AVX-512 A-quantizer (the same one the production W2 path uses), and call +// the SIMD kernel per-thread on the N-tile chunks the dispatcher splits over. +// +// Constraints (must be satisfied for the kernel to handle the rows; bench +// skips otherwise): +// * BlkLen == 64 +// * K a multiple of (BlkLen * kSuperBlockBlks) = 256 +// * M a multiple of kNRows2 (=2) +// * N a multiple of kNCols4 (=4) +// +namespace bench_super { + +namespace sq2sb = onnxruntime::mlas::sq2bit_avx512_super; +namespace sq2 = onnxruntime::mlas::sq2bit_avx512; + +void RunQ2SuperBlockBenchmark(size_t M, size_t N, size_t K, size_t Threads, + bool HasBias, benchmark::State& state) { + using onnxruntime::narrow; + constexpr size_t BlkBitWidth = 2; + constexpr size_t BlkLen = sq2::kBlkLen; + constexpr size_t kSuperBlockBlks = sq2::kSuperBlockBlks; + // R2xC4 SIMD tile (kNRows2 lives in the SIMD-only header which the bench + // TU is not built against -- mirror its value here, used only for the + // N alignment gate below). The kernel handles any M >= 1 and any K >= 1 + // (K-tail handler covers K not a multiple of 256). + constexpr size_t kNCols4 = sq2::kNCols4; + (void)kSuperBlockBlks; // retained for documentation reference only + + // Gate on the kernel's hard constraints. + if (K % BlkLen != 0 || K == 0) { + state.SkipWithMessage("Super-block requires K > 0 and K % BlkLen == 0."); + return; + } + if (M == 0 || (N % kNCols4) != 0) { + state.SkipWithMessage("Super-block requires M>=1 and N%4==0."); + return; + } + + // Gate on host having AVX-512(-VNNI) -- the kernel is AVX-512 BW + (optional) VNNI. + // QuantizeARowComputeBlkSum_CompInt8 is AVX-512 and required for the A-quant step. + const auto& platform = GetMlasPlatform(); + if (platform.QNBitGemmDispatch == nullptr || + platform.QNBitGemmDispatch->QuantizeARowComputeBlkSum_CompInt8 == nullptr) { + state.SkipWithMessage("AVX-512 dispatch table not available on this host."); + return; + } + + const size_t BlockCountK = K / BlkLen; + + OrtThreadPoolParams tpo; + tpo.thread_pool_size = static_cast(Threads); + tpo.auto_set_affinity = true; + std::unique_ptr tp( + onnxruntime::concurrency::CreateThreadPool(&onnxruntime::Env::Default(), + tpo, + onnxruntime::concurrency::ThreadPoolType::INTRA_OP)); + + // ----- Source data ----- + const auto A = RandomVectorUniform(M * K, float{-1.0f}, float{1.0f}); + const auto B = RandomVectorUniform(K * N, float{-1.0f}, float{1.0f}); + const auto Bias = HasBias ? RandomVectorUniform(N, float{-1.0f}, float{1.0f}) : std::vector(); + + size_t QuantBDataSizeInBytes, QuantBScaleSize, QuantBZeroPointSizeInBytes; + MlasBlockwiseQuantizedBufferSizes( + static_cast(BlkLen), /*columnwise=*/true, + static_cast(K), static_cast(N), + QuantBDataSizeInBytes, QuantBScaleSize, &QuantBZeroPointSizeInBytes); + + std::vector QuantBDataSrc(QuantBDataSizeInBytes); + std::vector QuantBScale(QuantBScaleSize); + MlasQuantizeBlockwise( + QuantBDataSrc.data(), QuantBScale.data(), /*zp=*/nullptr, + B.data(), static_cast(BlkLen), /*columnwise=*/true, + static_cast(K), static_cast(N), static_cast(N), tp.get()); + + // ----- Pack into the super-block layout ----- + const size_t PackedSize = sq2sb::Q2BitGemmPackQuantBDataSize_SuperBlock( + N, K, BlkLen, /*HasZeroPoint=*/false, SQNBIT_CompInt8, nullptr); + if (PackedSize == 0) { + state.SkipWithMessage("Super-block pack size returned 0 for this shape."); + return; + } + std::vector PackedBuf(PackedSize, std::byte{0}); + PackedQuantBDataStruct packed_b( + PackedBuf.data(), N, BlockCountK, BlkLen, /*QuantAUnsigned=*/false); + + // Same 3-call prepack pattern matmul_nbits.cc uses. + sq2sb::SQ2BitGemmPackQuantBDataAndBlkSum_SuperBlockScalar( + N, K, BlkLen, SQNBIT_CompInt8, + reinterpret_cast(QuantBDataSrc.data()), + /*scales=*/nullptr, + /*has_zp=*/false, /*zp=*/nullptr, + packed_b, tp.get(), nullptr); + sq2sb::SQ2BitGemmPackQuantBDataAndBlkSum_SuperBlockScalar( + N, K, BlkLen, SQNBIT_CompInt8, + /*B=*/nullptr, QuantBScale.data(), + /*has_zp=*/false, /*zp=*/nullptr, + packed_b, tp.get(), nullptr); + + // ----- Quantize A once via the dispatch's AVX-512 A-quantizer ----- + std::vector QuantAData(M * BlockCountK * BlkLen, std::byte{0}); + std::vector QuantAScale(M * BlockCountK, 0.0f); + std::vector ABlockSum(M * BlockCountK, 0.0f); + auto QuantizeARow = platform.QNBitGemmDispatch->QuantizeARowComputeBlkSum_CompInt8; + for (size_t m = 0; m < M; ++m) { + QuantizeARow(BlkLen, A.data() + m * K, K, + QuantAData.data() + m * BlockCountK * BlkLen, + QuantAScale.data() + m * BlockCountK, + ABlockSum.data() + m * BlockCountK); + } + + std::vector C(M * N, 0.0f); + + // Pick the best SIMD variant for the host. Prefer the VNNI variant when + // the platform dispatch table is the VNNI one (same logic the production + // dispatch wiring would use). + const bool use_vnni = (platform.QNBitGemmDispatch == &MlasSQNBitGemmDispatchAvx512vnni); + auto kernel = use_vnni + ? sq2sb::SQ2BitGemmKernel_BlkSum_CompInt8_Super_Avx512Vnni_TestEntry + : sq2sb::SQ2BitGemmKernel_BlkSum_CompInt8_Super_Avx512_TestEntry; + + // Mirror the production dispatcher's N-tile parallel split. SQ2BitGemm_CompInt8 + // tiles N in chunks of 128 and parallelizes across them; we do the same. + constexpr size_t kNTile = 128; + const size_t num_n_tiles = (N + kNTile - 1) / kNTile; + + auto run_one = [&]() { + onnxruntime::concurrency::ThreadPool::TryBatchParallelFor( + tp.get(), static_cast(num_n_tiles), + [&](ptrdiff_t t) { + const size_t n_start = static_cast(t) * kNTile; + const size_t n_count = std::min(kNTile, N - n_start); + if (n_count == 0) return; + + // Pointer arithmetic mirroring SQ2BitGemm_CompInt8. + const size_t ldb_bytes = BlockCountK * sq2::kBlkBytes; + const size_t ldb_scale = BlockCountK; + const std::byte* b_tile = packed_b.PackedQuantBData + n_start * ldb_bytes; + const float* bscale_tile = packed_b.PackedQuantBScale + n_start * ldb_scale; + const float* bblksum_tile = packed_b.QuantBBlkSum + n_start * ldb_scale; + float* c_tile = C.data() + n_start; + const float* bias_tile = HasBias ? (Bias.data() + n_start) : nullptr; + + kernel(BlkLen, + QuantAData.data(), QuantAScale.data(), + b_tile, bscale_tile, + /*QuantBZeroPoint=*/nullptr, + c_tile, + M, n_count, K, BlockCountK, + bias_tile, + /*ldc=*/N, + ABlockSum.data(), bblksum_tile); + }, + /*cost=*/0); + }; + + run_one(); // warm up + for (auto _ : state) { + run_one(); + } +} + +} // namespace bench_super + +void QNBITGEMM_SUPER(benchmark::State& state) { + using onnxruntime::narrow; + const auto BlkLen = narrow(state.range(0)); + (void)BlkLen; // The super-block kernel only supports BlkLen=64; gated inside. + const auto M = narrow(state.range(1)); + const auto N = narrow(state.range(2)); + const auto K = narrow(state.range(3)); + const auto Threads = narrow(state.range(4)); + // state.range(5) (Symmetric) and (7) (ComputeType) are unused; the + // super-block kernel is symmetric CompInt8 only. + const bool HasBias = narrow(state.range(6)); + bench_super::RunQ2SuperBlockBenchmark(M, N, K, Threads, HasBias, state); +} + +// Customer-shape rows for the super-block prototype. Uses the same argument +// schema as QNBITGEMM so the bench rows line up one-to-one in the +// output (modulo skipped shapes that violate the super-block K constraint). +BENCHMARK(QNBITGEMM_SUPER)->Apply(QNBit2BitCustomerArgs)->UseRealTime(); // This test gets benchmark arguments from environment variables. template diff --git a/onnxruntime/test/mlas/unittest/test_sqnbitgemm_2bit.cpp b/onnxruntime/test/mlas/unittest/test_sqnbitgemm_2bit.cpp index c6434916e8bac..002426335759a 100644 --- a/onnxruntime/test/mlas/unittest/test_sqnbitgemm_2bit.cpp +++ b/onnxruntime/test/mlas/unittest/test_sqnbitgemm_2bit.cpp @@ -165,3 +165,164 @@ TEST(MlasSq2BitTest, PackUnpackRoundTrip_BlkLen64_ConstantValues) } } } + +// ----------------------------------------------------------------------------- +// EXPERIMENTAL: super-block (4-K-block) round-trip tests +// ----------------------------------------------------------------------------- +// +// These exercise PackSuperBlock4_BlkLen64 / UnpackSuperBlock4_BlkLen64_Reference, +// which underpin the fast-unpack prototype (single 64-byte load + 4 fixed +// shift-and-mask producing 4 block ZMMs vs the current broadcast + variable +// shift per block). The tests live alongside the per-block tests above so a +// regression in either layout is caught by the same test target. + +// +// Deterministic-pattern super-block round-trip. Block k assigns weight i the +// value ((i + k) % 4), giving every (block_index, position, value) a unique +// fingerprint that pinpoints a layout swap if any. +// +TEST(MlasSq2BitTest, PackUnpackRoundTrip_SuperBlock4_BlkLen64_DeterministicPattern) +{ + std::array, sq2::kSuperBlockBlks> weights{}; + for (size_t k = 0; k < sq2::kSuperBlockBlks; ++k) { + for (size_t i = 0; i < sq2::kBlkLen; ++i) { + weights[k][i] = static_cast((i + k) % 4); + } + } + + std::array, sq2::kSuperBlockBlks> src{}; + for (size_t k = 0; k < sq2::kSuperBlockBlks; ++k) { + PackSourceBlock_BlkLen64(weights[k].data(), src[k].data()); + } + + std::array packed{}; + sq2::PackSuperBlock4_BlkLen64(src[0].data(), src[1].data(), src[2].data(), src[3].data(), + packed.data()); + + std::array, sq2::kSuperBlockBlks> recovered{}; + sq2::UnpackSuperBlock4_BlkLen64_Reference(packed.data(), + recovered[0].data(), recovered[1].data(), + recovered[2].data(), recovered[3].data()); + + for (size_t k = 0; k < sq2::kSuperBlockBlks; ++k) { + for (size_t i = 0; i < sq2::kBlkLen; ++i) { + ASSERT_EQ(recovered[k][i], weights[k][i]) + << "Super-block round-trip mismatch k=" << k << " i=" << i; + } + } +} + +// +// Randomized super-block round-trip across several seeds. The four input +// blocks are independent random fills; the test fails fast if any (block, +// weight) entry is mis-routed by the packed-byte layout. +// +TEST(MlasSq2BitTest, PackUnpackRoundTrip_SuperBlock4_BlkLen64_Randomized) +{ + constexpr unsigned kSeeds = 8; + + for (unsigned seed = 0; seed < kSeeds; ++seed) { + std::mt19937 rng(seed * 5051u + 13u); + std::uniform_int_distribution dist(0u, 3u); + + std::array, sq2::kSuperBlockBlks> weights{}; + for (auto& blk : weights) { + for (auto& w : blk) { + w = static_cast(dist(rng)); + } + } + + std::array, sq2::kSuperBlockBlks> src{}; + for (size_t k = 0; k < sq2::kSuperBlockBlks; ++k) { + PackSourceBlock_BlkLen64(weights[k].data(), src[k].data()); + } + + std::array packed{}; + sq2::PackSuperBlock4_BlkLen64(src[0].data(), src[1].data(), src[2].data(), src[3].data(), + packed.data()); + + std::array, sq2::kSuperBlockBlks> recovered{}; + sq2::UnpackSuperBlock4_BlkLen64_Reference(packed.data(), + recovered[0].data(), recovered[1].data(), + recovered[2].data(), recovered[3].data()); + + for (size_t k = 0; k < sq2::kSuperBlockBlks; ++k) { + for (size_t i = 0; i < sq2::kBlkLen; ++i) { + ASSERT_EQ(recovered[k][i], weights[k][i]) + << "Random super-block mismatch seed=" << seed << " k=" << k << " i=" << i; + } + } + } +} + +// +// Constant-value invariants for the super-block layout: +// - All four blocks set to the same value v produces packed bytes equal to +// 0x55 * v (v repeated at bit positions {0..1,2..3,4..5,6..7}). +// - Block_k set to value v with all other blocks zero produces packed bytes +// equal to (v << (2*k)) -- exclusively occupying the k-th bit slot. +// +TEST(MlasSq2BitTest, PackUnpackRoundTrip_SuperBlock4_BlkLen64_ConstantValues) +{ + // Case 1: every block filled with v. + for (uint8_t v = 0; v < 4; ++v) { + std::array, sq2::kSuperBlockBlks> weights{}; + for (auto& blk : weights) { + blk.fill(v); + } + + std::array, sq2::kSuperBlockBlks> src{}; + for (size_t k = 0; k < sq2::kSuperBlockBlks; ++k) { + PackSourceBlock_BlkLen64(weights[k].data(), src[k].data()); + } + + std::array packed{}; + sq2::PackSuperBlock4_BlkLen64(src[0].data(), src[1].data(), src[2].data(), src[3].data(), + packed.data()); + + const uint8_t expected_byte = static_cast(v * 0x55u); + for (size_t i = 0; i < sq2::kSuperBlockBytes; ++i) { + ASSERT_EQ(static_cast(packed[i]), expected_byte) + << "Uniform-fill v=" << static_cast(v) << " byte_i=" << i; + } + } + + // Case 2: only one block at a time carries a non-zero value. + for (size_t target_k = 0; target_k < sq2::kSuperBlockBlks; ++target_k) { + for (uint8_t v = 1; v < 4; ++v) { + std::array, sq2::kSuperBlockBlks> weights{}; + weights[target_k].fill(v); + + std::array, sq2::kSuperBlockBlks> src{}; + for (size_t k = 0; k < sq2::kSuperBlockBlks; ++k) { + PackSourceBlock_BlkLen64(weights[k].data(), src[k].data()); + } + + std::array packed{}; + sq2::PackSuperBlock4_BlkLen64(src[0].data(), src[1].data(), src[2].data(), src[3].data(), + packed.data()); + + const uint8_t expected_byte = static_cast(v << (2 * target_k)); + for (size_t i = 0; i < sq2::kSuperBlockBytes; ++i) { + ASSERT_EQ(static_cast(packed[i]), expected_byte) + << "Isolated block target_k=" << target_k + << " v=" << static_cast(v) + << " byte_i=" << i; + } + + std::array, sq2::kSuperBlockBlks> recovered{}; + sq2::UnpackSuperBlock4_BlkLen64_Reference( + packed.data(), + recovered[0].data(), recovered[1].data(), + recovered[2].data(), recovered[3].data()); + for (size_t k = 0; k < sq2::kSuperBlockBlks; ++k) { + const uint8_t expect_val = (k == target_k) ? v : uint8_t{0}; + for (size_t i = 0; i < sq2::kBlkLen; ++i) { + ASSERT_EQ(recovered[k][i], expect_val) + << "Isolated round-trip target_k=" << target_k + << " k=" << k << " v=" << static_cast(v) << " i=" << i; + } + } + } + } +} diff --git a/onnxruntime/test/mlas/unittest/test_sqnbitgemm_2bit_gemm.cpp b/onnxruntime/test/mlas/unittest/test_sqnbitgemm_2bit_gemm.cpp index c1905ec56a8a9..f7c55f8eb86f9 100644 --- a/onnxruntime/test/mlas/unittest/test_sqnbitgemm_2bit_gemm.cpp +++ b/onnxruntime/test/mlas/unittest/test_sqnbitgemm_2bit_gemm.cpp @@ -413,306 +413,26 @@ TEST(MlasSq2BitTest, GemmCompInt8_BlkLen64_PublicApi_WithZeroPoints) } // -// Direct-call test harness for the AVX-512 W2 kernel variants. Calls the -// kernel through its non-inline test-entry forwarder (declared in -// sqnbitgemm_kernel_avx512_2bit.h), bypassing the platform dispatcher. -// This is the only way to exercise the non-VNNI kernel on a VNNI host. +// ----------------------------------------------------------------------------- +// W2-v1 direct-call kernel tests (removed). // -// Setup re-uses the public pack API (MlasQNBitGemmPackQuantBData) and the -// reference int8 A-quantizer in this file (QuantizeA_Reference, which is -// bit-identical to QuantizeARow_CompInt8_avx512). Output is compared -// against ReferenceGemm_W2_CompInt8 within a tight tolerance. +// Earlier revisions of this file exercised the AVX-512 / AVX-512-VNNI W2-v1 +// kernels (`sq2bit_avx512::SQ2BitGemmKernel_BlkSum_CompInt8_Avx512[Vnni]_TestEntry`) +// directly via test-entry forwarders, bypassing the platform dispatcher. That +// arrangement made sense when the production dispatch used the W2-v1 packed-B +// layout: the public pack API and the direct kernel call shared a buffer ABI. // -class MlasSQ2BitGemmDirectCallTest { - public: - static void Run(size_t M, size_t N, size_t K, bool WithBias, uint32_t seed, - bool TestVnni, bool WithZeroPoints = false) - { - const size_t BlockCountK = (K + kBlkLen - 1) / kBlkLen; - ASSERT_EQ(K % kBlkLen, 0u) << "Test K must be a multiple of BlkLen=64"; - - std::mt19937 rng(seed); - std::uniform_real_distribution a_dist(-1.0f, 1.0f); - std::uniform_int_distribution w_dist(0, 3); - std::uniform_real_distribution s_dist(0.05f, 0.5f); - - std::vector A(M * K); - for (auto& v : A) v = a_dist(rng); - - std::vector BWeights(N * K); - for (auto& v : BWeights) v = static_cast(w_dist(rng)); - - std::vector QuantBData(N * BlockCountK * kBlkBytes, std::byte{0}); - for (size_t n = 0; n < N; ++n) { - for (size_t blk = 0; blk < BlockCountK; ++blk) { - uint8_t blk_weights[kBlkLen]; - for (size_t kk = 0; kk < kBlkLen; ++kk) { - blk_weights[kk] = BWeights[n * K + blk * kBlkLen + kk]; - } - PackSourceBlock_BlkLen64( - blk_weights, - QuantBData.data() + (n * BlockCountK + blk) * kBlkBytes); - } - } - - std::vector QuantBScale(N * BlockCountK); - for (auto& v : QuantBScale) v = s_dist(rng); - - // Per-block zero points (same conventions as the public-API harness). - std::vector BZeroPoints; - std::vector BZeroPointsPacked; - const uint8_t* BZeroPointsRef = nullptr; - const std::byte* BZeroPointsMlas = nullptr; - if (WithZeroPoints) { - BZeroPoints.resize(N * BlockCountK); - for (auto& v : BZeroPoints) v = static_cast(w_dist(rng)); - BZeroPointsRef = BZeroPoints.data(); - BZeroPointsPacked = PackW2ZeroPoints(N, BlockCountK, BZeroPoints); - BZeroPointsMlas = BZeroPointsPacked.data(); - } - - std::vector Bias; - const float* BiasPtr = nullptr; - if (WithBias) { - Bias.resize(N); - for (auto& v : Bias) v = a_dist(rng); - BiasPtr = Bias.data(); - } - - // Pack B through the public API (produces the same buffer for both - // VNNI and non-VNNI consumers; the kernel doesn't care which one). - const size_t PackedSize = MlasQNBitGemmPackQuantBDataSize( - N, K, kBlkBitWidth, kBlkLen, WithZeroPoints, kComputeType, nullptr); - ASSERT_GT(PackedSize, 0u); - std::vector PackedQuantBBuf(PackedSize, std::byte{0}); - - MlasQNBitGemmPackQuantBData( - N, K, kBlkBitWidth, kBlkLen, kComputeType, - QuantBData.data(), PackedQuantBBuf.data(), - QuantBScale.data(), WithZeroPoints, BZeroPointsMlas, - nullptr, nullptr); - - // Reconstruct the packed-B view so we can pass the right sub-pointers - // to the kernel forwarder (PackedQuantBData, PackedQuantBScale, - // QuantBBlkSum). The struct is a layout overlay over the buffer. - PackedQuantBDataStruct packed_b( - PackedQuantBBuf.data(), N, BlockCountK, kBlkLen, /*QuantAUnsigned=*/false); - - // Quantize A (bit-identical to MLAS's QuantizeARow_CompInt8_avx512 for - // BlkLen=64: per-block symmetric scale = amax / 127). Also compute the - // scaled block sums the kernel's BlkSum correction needs. - std::vector QuantAData(M * BlockCountK * kBlkLen, int8_t{0}); - std::vector QuantAScale(M * BlockCountK, 0.0f); - QuantizeA_Reference(M, K, A.data(), QuantAData.data(), QuantAScale.data()); - - std::vector ABlockSum(M * BlockCountK, 0.0f); - for (size_t m = 0; m < M; ++m) { - for (size_t blk = 0; blk < BlockCountK; ++blk) { - int32_t sum = 0; - for (size_t kk = 0; kk < kBlkLen; ++kk) { - sum += static_cast( - QuantAData[m * BlockCountK * kBlkLen + blk * kBlkLen + kk]); - } - ABlockSum[m * BlockCountK + blk] = - QuantAScale[m * BlockCountK + blk] * static_cast(sum); - } - } - - // Run the kernel under test directly via the test-entry forwarder. - std::vector C(M * N, 0.0f); - if (TestVnni) { - onnxruntime::mlas::sq2bit_avx512::SQ2BitGemmKernel_BlkSum_CompInt8_Avx512Vnni_TestEntry( - kBlkLen, - reinterpret_cast(QuantAData.data()), - QuantAScale.data(), - packed_b.PackedQuantBData, - packed_b.PackedQuantBScale, - /*QuantBZeroPoint=*/nullptr, - C.data(), - M, N, /*CountK=*/K, BlockCountK, - BiasPtr, - /*ldc=*/N, - ABlockSum.data(), - packed_b.QuantBBlkSum); - } else { - onnxruntime::mlas::sq2bit_avx512::SQ2BitGemmKernel_BlkSum_CompInt8_Avx512_TestEntry( - kBlkLen, - reinterpret_cast(QuantAData.data()), - QuantAScale.data(), - packed_b.PackedQuantBData, - packed_b.PackedQuantBScale, - /*QuantBZeroPoint=*/nullptr, - C.data(), - M, N, /*CountK=*/K, BlockCountK, - BiasPtr, - /*ldc=*/N, - ABlockSum.data(), - packed_b.QuantBBlkSum); - } - - // Reference: bit-exact integer-domain math. - std::vector CRef(M * N, 0.0f); - ReferenceGemm_W2_CompInt8(M, N, K, A.data(), BWeights, QuantBScale.data(), - BZeroPointsRef, BiasPtr, CRef.data()); - - const float abs_tol = 1e-4f; - const float rel_tol = 1e-4f; - for (size_t i = 0; i < M * N; ++i) { - const float diff = std::fabs(C[i] - CRef[i]); - const float bound = abs_tol + rel_tol * std::fabs(CRef[i]); - ASSERT_LE(diff, bound) - << (TestVnni ? "VNNI" : "non-VNNI") << " direct-call mismatch at i=" << i - << " (m=" << (i / N) << ", n=" << (i % N) << ")" - << " MLAS=" << C[i] << " Ref=" << CRef[i] - << " M=" << M << " N=" << N << " K=" << K - << " WithBias=" << WithBias - << " WithZeroPoints=" << WithZeroPoints; - } - } -}; - -// -// Exercises the non-VNNI W2 kernel (vpmaddubsw + vpmaddwd + vpaddd MAC chain). -// Gated on AVX-512BW availability. Runs even on VNNI hosts where the platform -// dispatcher would never select this kernel naturally. -// -TEST(MlasSq2BitTest, GemmCompInt8_BlkLen64_Avx512) -{ - if (!GetMlasPlatform().Avx512Supported_) { - GTEST_SKIP() << "AVX-512BW (DQ/VL) not available on this host"; - } - - struct Shape { size_t M, N, K; }; - constexpr Shape shapes[] = { - {1, 16, 64}, - {1, 32, 128}, - {1, 64, 256}, - {4, 16, 64}, - {4, 33, 192}, - {7, 17, 128}, - {16, 64, 512}, - {32, 128, 256}, - }; - - for (uint32_t seed : {0xC0FFEEu, 0xBADC0DEu}) { - for (const Shape& s : shapes) { - for (bool bias : {false, true}) { - MlasSQ2BitGemmDirectCallTest::Run( - s.M, s.N, s.K, bias, seed + (bias ? 1u : 0u), - /*TestVnni=*/false); - } - } - } -} - -// -// Same as GemmCompInt8_BlkLen64_Avx512 but with random per-block zero points. -// -TEST(MlasSq2BitTest, GemmCompInt8_BlkLen64_Avx512_WithZeroPoints) -{ - if (!GetMlasPlatform().Avx512Supported_) { - GTEST_SKIP() << "AVX-512BW (DQ/VL) not available on this host"; - } - - struct Shape { size_t M, N, K; }; - constexpr Shape shapes[] = { - {1, 16, 64}, - {1, 32, 128}, - {1, 64, 256}, - {4, 16, 64}, - {4, 33, 192}, - {7, 17, 128}, - {16, 64, 512}, - {32, 128, 256}, - }; - - for (uint32_t seed : {0xC0FFEEu, 0xBADC0DEu}) { - for (const Shape& s : shapes) { - for (bool bias : {false, true}) { - MlasSQ2BitGemmDirectCallTest::Run( - s.M, s.N, s.K, bias, seed + (bias ? 1u : 0u), - /*TestVnni=*/false, /*WithZeroPoints=*/true); - } - } - } -} - -// -// Exercises the AVX-512-VNNI W2 kernel (_mm512_dpbusd_epi32 MAC) via the -// direct-call forwarder. On a VNNI host this is the same kernel that the -// public-API test above hits through the dispatcher; the explicit invocation -// here keeps the harness symmetric and validates the forwarder mechanism. -// -// Gating: requires the platform to have actually selected the VNNI dispatch -// table. MlasIsQNBitGemmAvailable alone is insufficient because the W2 path -// is now registered into both AVX-512 and AVX-512-VNNI dispatch tables. -// -TEST(MlasSq2BitTest, GemmCompInt8_BlkLen64_Avx512Vnni) -{ - if (GetMlasPlatform().QNBitGemmDispatch != &MlasSQNBitGemmDispatchAvx512vnni) { - GTEST_SKIP() << "AVX-512-VNNI not selected as the active dispatch on this host"; - } - - struct Shape { size_t M, N, K; }; - constexpr Shape shapes[] = { - {1, 16, 64}, - {1, 32, 128}, - {1, 64, 256}, - {4, 16, 64}, - {4, 33, 192}, - {7, 17, 128}, - {16, 64, 512}, - {32, 128, 256}, - // Customer model shapes (BlkLen=64). Both decode (M=1) and prefill - // (M=128) rows; these are far larger than the small-N shapes above - // and exercise the R1xC4 / R2xC4 tile paths at production sizes. - { 1, 1024, 384}, { 1, 192, 1024}, { 1, 384, 1024}, { 1, 4096, 1024}, { 1, 1024, 4096}, - {128, 1024, 384}, {128, 192, 1024}, {128, 384, 1024}, {128, 4096, 1024}, {128, 1024, 4096}, - }; - - for (uint32_t seed : {0xC0FFEEu, 0xBADC0DEu}) { - for (const Shape& s : shapes) { - for (bool bias : {false, true}) { - MlasSQ2BitGemmDirectCallTest::Run( - s.M, s.N, s.K, bias, seed + (bias ? 1u : 0u), - /*TestVnni=*/true); - } - } - } -} - -// -// Same as GemmCompInt8_BlkLen64_Avx512Vnni but with random per-block zero points. +// The production dispatch now routes through the W2-v2 super-block kernel +// (`sq2bit_avx512_super::*`) which uses a different packed-B layout (4-N-col +// grouped + super-block K stride; see PackedQuantBOffsetBytes_W2_SuperBlock). +// Driving the W2-v1 kernel with a W2-v2 pack would compare apples to oranges +// and produce spurious failures, so the W2-v1 direct-call harness has been +// retired. W2-v1 sources are kept in the build for now as a fallback while +// customers validate W2-v2 in production. // -TEST(MlasSq2BitTest, GemmCompInt8_BlkLen64_Avx512Vnni_WithZeroPoints) -{ - if (GetMlasPlatform().QNBitGemmDispatch != &MlasSQNBitGemmDispatchAvx512vnni) { - GTEST_SKIP() << "AVX-512-VNNI not selected as the active dispatch on this host"; - } - - struct Shape { size_t M, N, K; }; - constexpr Shape shapes[] = { - {1, 16, 64}, - {1, 32, 128}, - {1, 64, 256}, - {4, 16, 64}, - {4, 33, 192}, - {7, 17, 128}, - {16, 64, 512}, - {32, 128, 256}, - // Customer model shapes with explicit per-block zero points (matches the - // production accuracy concern). - { 1, 1024, 384}, { 1, 192, 1024}, { 1, 384, 1024}, { 1, 4096, 1024}, { 1, 1024, 4096}, - {128, 1024, 384}, {128, 192, 1024}, {128, 384, 1024}, {128, 4096, 1024}, {128, 1024, 4096}, - }; - - for (uint32_t seed : {0xC0FFEEu, 0xBADC0DEu}) { - for (const Shape& s : shapes) { - for (bool bias : {false, true}) { - MlasSQ2BitGemmDirectCallTest::Run( - s.M, s.N, s.K, bias, seed + (bias ? 1u : 0u), - /*TestVnni=*/true, /*WithZeroPoints=*/true); - } - } - } -} +// Coverage is preserved by: +// * `GemmCompInt8_BlkLen64_PublicApi[_WithZeroPoints]` -- end-to-end through +// `MlasQNBitGemmBatch`; this is what production code paths invoke. +// * `SuperBlock*` (test_sqnbitgemm_2bit_superblock.cpp) -- direct-call +// coverage of the new default W2-v2 kernel using the matching pack helper. +// ----------------------------------------------------------------------------- diff --git a/onnxruntime/test/mlas/unittest/test_sqnbitgemm_2bit_superblock.cpp b/onnxruntime/test/mlas/unittest/test_sqnbitgemm_2bit_superblock.cpp new file mode 100644 index 0000000000000..5117149f00548 --- /dev/null +++ b/onnxruntime/test/mlas/unittest/test_sqnbitgemm_2bit_superblock.cpp @@ -0,0 +1,529 @@ +/*++ + +Copyright (c) Microsoft Corporation. All rights reserved. + +Licensed under the MIT License. + +Module Name: + + test_sqnbitgemm_2bit_superblock.cpp + +Abstract: + + Unit tests for the EXPERIMENTAL super-block W2 path + (sqnbitgemm_kernel_avx512_2bit_superblock.{h,cpp}). + + Phase 2 coverage: end-to-end pack + scalar GEMM correctness against the + same integer-domain reference used by the production W2 tests. These + tests do NOT exercise any SIMD path -- they validate the layout, the + pack 3-call sequence, and the scalar oracle kernel that will back the + Phase 3 SIMD work. + + Tests deliberately use the same shapes as the production W2 tests so a + side-by-side comparison is straightforward. + +--*/ + +#include "gtest/gtest.h" + +#include +#include +#include +#include +#include + +#include "core/mlas/inc/mlas_qnbit.h" +#include "core/mlas/lib/qnbitgemm.h" +#include "core/mlas/lib/mlasi.h" +#include "core/mlas/lib/sqnbitgemm_kernel_avx512_2bit.h" +#include "core/mlas/lib/sqnbitgemm_kernel_avx512_2bit_superblock.h" + +namespace { + +namespace sq2 = onnxruntime::mlas::sq2bit_avx512; +namespace sq2sb = onnxruntime::mlas::sq2bit_avx512_super; + +constexpr size_t kBlkLen = sq2::kBlkLen; // 64 +constexpr size_t kBlkBytes = sq2::kBlkBytes; // 16 +constexpr size_t kSuperBlockBlks = sq2::kSuperBlockBlks; // 4 + +// Standard ONNX 2-bit source packing (1 byte = 4 weights). +void +PackSourceBlock_BlkLen64(const uint8_t weights[kBlkLen], std::byte* src_out) +{ + for (size_t i = 0; i < kBlkBytes; ++i) { + const uint8_t v0 = weights[4 * i + 0] & 0x03u; + const uint8_t v1 = weights[4 * i + 1] & 0x03u; + const uint8_t v2 = weights[4 * i + 2] & 0x03u; + const uint8_t v3 = weights[4 * i + 3] & 0x03u; + src_out[i] = static_cast( + static_cast(v0 | (v1 << 2) | (v2 << 4) | (v3 << 6)) + ); + } +} + +// Bit-exact mirror of MlasQNBitGemm's per-block int8 quantizer (amax/127, +// round-half-to-even via std::nearbyint, scale_recip = 127/amax). +void +QuantizeA_Reference(size_t M, size_t K, const float* A, + int8_t* QuantAData, float* QuantAScale) +{ + const size_t BlockCountK = (K + kBlkLen - 1) / kBlkLen; + for (size_t m = 0; m < M; ++m) { + for (size_t k = 0, k_blk = 0; k < K; k += kBlkLen, ++k_blk) { + const size_t local_len = std::min(K - k, kBlkLen); + float amax = 0.0f; + for (size_t kk = 0; kk < local_len; ++kk) { + amax = std::max(amax, std::fabs(A[m * K + k + kk])); + } + constexpr float range_max = 127.0f; + const float scale = amax / range_max; + const float scale_recip = amax != 0.0f ? range_max / amax : 0.0f; + QuantAScale[m * BlockCountK + k_blk] = scale; + for (size_t kk = 0; kk < kBlkLen; ++kk) { + const float a = (kk < local_len) ? A[m * K + k + kk] : 0.0f; + const float q = std::nearbyint(a * scale_recip); + QuantAData[m * BlockCountK * kBlkLen + k + kk] = + static_cast(std::clamp(q, -127.0f, 127.0f)); + } + } + } +} + +// +// Integer-domain GEMM oracle: bit-exact match to the math the MLAS W2 path +// performs (kernel int8 GEMM + SGEMM zero-point correction collapsed into a +// single direct dot of (qa * (qb - zp))). +// +void +ReferenceGemm_W2_CompInt8(size_t M, size_t N, size_t K, + const float* A, + const std::vector& BWeights, // [N * K] in [0, 3] + const float* QuantBScale, // [N * BlockCountK] + const uint8_t* BZeroPoints, // [N * BlockCountK] or nullptr + const float* Bias, + float* C) +{ + const size_t BlockCountK = (K + kBlkLen - 1) / kBlkLen; + std::vector QuantAData(M * BlockCountK * kBlkLen, int8_t{0}); + std::vector QuantAScale(M * BlockCountK, 0.0f); + QuantizeA_Reference(M, K, A, QuantAData.data(), QuantAScale.data()); + + for (size_t m = 0; m < M; ++m) { + for (size_t n = 0; n < N; ++n) { + float acc = (Bias != nullptr) ? Bias[n] : 0.0f; + for (size_t k = 0, blk = 0; k < K; k += kBlkLen, ++blk) { + const size_t local_len = std::min(K - k, kBlkLen); + const float a_scale = QuantAScale[m * BlockCountK + blk]; + const float b_scale = QuantBScale[n * BlockCountK + blk]; + const int32_t zp = BZeroPoints != nullptr + ? static_cast(BZeroPoints[n * BlockCountK + blk]) + : static_cast(sq2::kDefaultSymmetricZeroPoint2Bit); + int32_t dot = 0; + for (size_t kk = 0; kk < local_len; ++kk) { + const int8_t qa = QuantAData[m * BlockCountK * kBlkLen + k + kk]; + const int32_t qb = + static_cast(BWeights[n * K + k + kk]) - zp; + dot += static_cast(qa) * qb; + } + acc += static_cast(dot) * a_scale * b_scale; + } + C[m * N + n] = acc; + } + } +} + +// +// Pack per-block W2 zero points into the standard ONNX byte stream +// (4 zp per byte along K, row-major in N). +// +std::vector +PackW2ZeroPoints(size_t N, size_t BlockCountK, const std::vector& BZeroPoints) +{ + const size_t ZPCountK = (BlockCountK + 3) / 4; + std::vector packed(N * ZPCountK, std::byte{0}); + for (size_t n = 0; n < N; ++n) { + for (size_t blk = 0; blk < BlockCountK; ++blk) { + const uint8_t zp = BZeroPoints[n * BlockCountK + blk] & 0x03u; + const size_t byte_idx = n * ZPCountK + (blk / 4); + const size_t bit_off = (blk % 4) * 2; + packed[byte_idx] = static_cast( + static_cast(packed[byte_idx]) | (zp << bit_off)); + } + } + return packed; +} + +// +// Test harness that drives the super-block path directly (without going through +// MlasQNBitGemmBatch). Builds the same packed buffer the dispatcher would +// construct, runs the chosen super-block kernel (scalar / AVX-512BW / VNNI), +// and compares to ReferenceGemm_W2_CompInt8. +// +// `KernelFn` matches the SQ4BitGemmKernel_BlkSum_CompInt8_Fn signature, which +// every super-block kernel variant honors via direct-call forwarders declared +// in sqnbitgemm_kernel_avx512_2bit_superblock.h. +// +using SuperBlockKernelFn = size_t (MLASCALL*)( + size_t, const std::byte*, const float*, const std::byte*, const float*, + const std::byte*, float*, size_t, size_t, size_t, size_t, + const float*, size_t, const float*, const float*); + +void +RunSuperBlockCase(size_t M, size_t N, size_t K, bool WithBias, uint32_t seed, + bool WithZeroPoints, SuperBlockKernelFn kernel, + const char* kernel_name) +{ + const size_t BlockCountK = (K + kBlkLen - 1) / kBlkLen; + ASSERT_EQ(K % kBlkLen, 0u) << "Test K must be a multiple of BlkLen=64"; + // BlockCountK no longer required to be a multiple of kSuperBlockBlks -- + // the K-tail handler picks up the trailing 1-3 blocks. + + std::mt19937 rng(seed); + std::uniform_real_distribution a_dist(-1.0f, 1.0f); + std::uniform_int_distribution w_dist(0, 3); + std::uniform_real_distribution s_dist(0.05f, 0.5f); + + std::vector A(M * K); + for (auto& v : A) v = a_dist(rng); + + std::vector BWeights(N * K); + for (auto& v : BWeights) v = static_cast(w_dist(rng)); + + // Source-packed B (standard ONNX layout) -- the input to the pack helper. + std::vector QuantBData(N * BlockCountK * kBlkBytes, std::byte{0}); + for (size_t n = 0; n < N; ++n) { + for (size_t blk = 0; blk < BlockCountK; ++blk) { + uint8_t blk_weights[kBlkLen]; + for (size_t kk = 0; kk < kBlkLen; ++kk) { + blk_weights[kk] = BWeights[n * K + blk * kBlkLen + kk]; + } + PackSourceBlock_BlkLen64(blk_weights, + QuantBData.data() + (n * BlockCountK + blk) * kBlkBytes); + } + } + + std::vector QuantBScale(N * BlockCountK); + for (auto& v : QuantBScale) v = s_dist(rng); + + std::vector BZeroPoints; + std::vector BZeroPointsPacked; + const uint8_t* BZeroPointsRef = nullptr; + const std::byte* BZeroPointsMlas = nullptr; + if (WithZeroPoints) { + BZeroPoints.resize(N * BlockCountK); + for (auto& v : BZeroPoints) v = static_cast(w_dist(rng)); + BZeroPointsRef = BZeroPoints.data(); + BZeroPointsPacked = PackW2ZeroPoints(N, BlockCountK, BZeroPoints); + BZeroPointsMlas = BZeroPointsPacked.data(); + } + + std::vector Bias; + const float* BiasPtr = nullptr; + if (WithBias) { + Bias.resize(N); + for (auto& v : Bias) v = a_dist(rng); + BiasPtr = Bias.data(); + } + + // Allocate the packed-B buffer (same total size as the production path). + const size_t PackedSize = sq2sb::Q2BitGemmPackQuantBDataSize_SuperBlock( + N, K, kBlkLen, WithZeroPoints, SQNBIT_CompInt8, nullptr); + ASSERT_GT(PackedSize, 0u) << "Super-block pack size unsupported for the chosen shape"; + + std::vector PackedQuantBBuf(PackedSize, std::byte{0}); + // The W2 PackedQuantBDataStruct constructor pads BlockCountK to a multiple + // of 4 internally (see qnbitgemm.h) so the slab layout matches what the + // super-block pack helper writes regardless of whether the caller passes + // the logical or padded BlockCountK. We pass the logical value to mirror + // exactly what matmul_nbits.cc does in production. + PackedQuantBDataStruct packed_b( + PackedQuantBBuf.data(), N, BlockCountK, kBlkLen, /*QuantAUnsigned=*/false); + + // Mirror the matmul_nbits.cc prepack 3-call pattern (B, scales, ZP) so the + // pack code path is exercised exactly as the production dispatcher would. + sq2sb::SQ2BitGemmPackQuantBDataAndBlkSum_SuperBlockScalar( + N, K, kBlkLen, SQNBIT_CompInt8, + QuantBData.data(), /*scales=*/nullptr, + WithZeroPoints, /*zp=*/nullptr, + packed_b, /*tp=*/nullptr, /*cfg=*/nullptr); + sq2sb::SQ2BitGemmPackQuantBDataAndBlkSum_SuperBlockScalar( + N, K, kBlkLen, SQNBIT_CompInt8, + /*B=*/nullptr, QuantBScale.data(), + WithZeroPoints, /*zp=*/nullptr, + packed_b, /*tp=*/nullptr, /*cfg=*/nullptr); + if (WithZeroPoints) { + sq2sb::SQ2BitGemmPackQuantBDataAndBlkSum_SuperBlockScalar( + N, K, kBlkLen, SQNBIT_CompInt8, + /*B=*/nullptr, /*scales=*/nullptr, + WithZeroPoints, BZeroPointsMlas, + packed_b, /*tp=*/nullptr, /*cfg=*/nullptr); + } + + // Quantize A the same way MLAS would (per-block amax/127, banker rounding). + std::vector QuantAData(M * BlockCountK * kBlkLen, int8_t{0}); + std::vector QuantAScale(M * BlockCountK, 0.0f); + QuantizeA_Reference(M, K, A.data(), QuantAData.data(), QuantAScale.data()); + + std::vector ABlockSum(M * BlockCountK, 0.0f); + for (size_t m = 0; m < M; ++m) { + for (size_t blk = 0; blk < BlockCountK; ++blk) { + int32_t sum = 0; + for (size_t kk = 0; kk < kBlkLen; ++kk) { + sum += static_cast( + QuantAData[m * BlockCountK * kBlkLen + blk * kBlkLen + kk]); + } + ABlockSum[m * BlockCountK + blk] = + QuantAScale[m * BlockCountK + blk] * static_cast(sum); + } + } + + std::vector C(M * N, 0.0f); + kernel( + kBlkLen, + reinterpret_cast(QuantAData.data()), + QuantAScale.data(), + packed_b.PackedQuantBData, + packed_b.PackedQuantBScale, + /*QuantBZeroPoint=*/nullptr, + C.data(), + M, N, /*CountK=*/K, BlockCountK, + BiasPtr, + /*ldc=*/N, + ABlockSum.data(), + packed_b.QuantBBlkSum); + + std::vector CRef(M * N, 0.0f); + ReferenceGemm_W2_CompInt8(M, N, K, A.data(), BWeights, QuantBScale.data(), + BZeroPointsRef, BiasPtr, CRef.data()); + + const float abs_tol = 1e-4f; + const float rel_tol = 1e-4f; + for (size_t i = 0; i < M * N; ++i) { + const float diff = std::fabs(C[i] - CRef[i]); + const float bound = abs_tol + rel_tol * std::fabs(CRef[i]); + ASSERT_LE(diff, bound) + << "Super-block " << kernel_name << " mismatch at i=" << i + << " (m=" << (i / N) << ", n=" << (i % N) << ")" + << " out=" << C[i] << " ref=" << CRef[i] + << " M=" << M << " N=" << N << " K=" << K + << " WithBias=" << WithBias + << " WithZeroPoints=" << WithZeroPoints; + } +} + +} // namespace + +// +// Scalar super-block test, no zero-points. Covers the same small synthetic +// shapes + customer prefill sizes used by the production W2 tests. All shapes +// have K as a multiple of (kBlkLen * kSuperBlockBlks) = 256. Customer K=384 +// is NOT a multiple of 256 so it's excluded; that shape will need a tail +// handler in a follow-up. +// +TEST(MlasSq2BitTest, SuperBlockScalar_BlkLen64) +{ + struct Shape { size_t M, N, K; }; + constexpr Shape shapes[] = { + {1, 16, 256}, + {1, 32, 256}, + {1, 64, 512}, + {4, 16, 256}, + {4, 33, 256}, + {7, 17, 256}, + {16, 64, 512}, + {32, 128, 256}, + // Customer prefill (only the K values that are multiples of 256). + { 1, 1024, 1024}, { 1, 192, 1024}, { 1, 384, 1024}, { 1, 4096, 1024}, { 1, 1024, 4096}, + {128, 1024, 1024}, {128, 192, 1024}, {128, 384, 1024}, {128, 4096, 1024}, {128, 1024, 4096}, + }; + + for (uint32_t seed : {0xC0FFEEu, 0xBADC0DEu}) { + for (const Shape& s : shapes) { + for (bool bias : {false, true}) { + RunSuperBlockCase(s.M, s.N, s.K, bias, seed + (bias ? 1u : 0u), + /*WithZeroPoints=*/false, + sq2sb::SQ2BitGemmKernel_BlkSum_CompInt8_SuperBlockScalar, + "scalar"); + } + } + } +} + +// +// Same coverage with per-block non-default zero points. +// +TEST(MlasSq2BitTest, SuperBlockScalar_BlkLen64_WithZeroPoints) +{ + struct Shape { size_t M, N, K; }; + constexpr Shape shapes[] = { + {1, 16, 256}, + {1, 32, 256}, + {1, 64, 512}, + {4, 16, 256}, + {4, 33, 256}, + {7, 17, 256}, + {16, 64, 512}, + {32, 128, 256}, + { 1, 1024, 1024}, { 1, 192, 1024}, { 1, 384, 1024}, { 1, 4096, 1024}, { 1, 1024, 4096}, + {128, 1024, 1024}, {128, 192, 1024}, {128, 384, 1024}, {128, 4096, 1024}, {128, 1024, 4096}, + }; + + for (uint32_t seed : {0xC0FFEEu, 0xBADC0DEu}) { + for (const Shape& s : shapes) { + for (bool bias : {false, true}) { + RunSuperBlockCase(s.M, s.N, s.K, bias, seed + (bias ? 1u : 0u), + /*WithZeroPoints=*/true, + sq2sb::SQ2BitGemmKernel_BlkSum_CompInt8_SuperBlockScalar, + "scalar"); + } + } + } +} + +// +// SIMD super-block shape coverage. Phase-3+K-tail kernel requires: +// * BlkLen == 64 +// * CountN a multiple of kNCols4 (=4) +// CountM and BlockCountK have NO alignment requirements: +// - R2xC4 handles the M-aligned head; a single R1xC4 picks up the optional +// trailing odd row. CountM == 1 dispatches directly to R1xC4. +// - The K-loop iterates `BlockCountK / 4` full super-blocks plus a partial +// "tail super" of 1-3 trailing K-blocks. The pack helpers zero-pad the +// trailing slots so they contribute 0 to the dot product; the tile loads +// only valid A blocks (zero ZMM for missing ones) to avoid OOB. +// +// Customer prefill shapes (M in {1, 128}, N in {192, 384, 1024, 4096}, K in +// {1024, 4096}) are all covered. M=3 and M=5 exercise the M-tail path. +// +// K-tail handler (BlockCountK not a multiple of kSuperBlockBlks=4): the +// pack helper zero-pads the trailing 1-3 K-block slots; the SIMD K-loop +// processes them via the 4-block accumulator with zero ZMM for the missing +// A blocks. GroupStride uses BlockCountKPadded so N-group advances land on +// the right packed-B address regardless of K % 4. Customer K=384 and the +// synthetic K=320, K=448 shapes exercise this path. +// +constexpr struct { size_t M, N, K; } kSimdShapes[] = { + {1, 16, 256}, // R1 only + {1, 192, 1024}, // R1 only, customer N + {1, 1024, 4096}, // R1 only, customer N + {2, 16, 256}, + {2, 32, 256}, + {2, 64, 512}, + {3, 16, 256}, // R2 head (1 pair) + R1 tail + {3, 384, 1024}, + {4, 16, 256}, + {4, 32, 256}, + {5, 64, 512}, // R2 head (2 pairs) + R1 tail + {16, 64, 512}, + {32, 128, 256}, + // Customer prefill (M=128) at all (K, N) pairs, including K=384. + {128, 1024, 384}, // K-tail: BlockCountK=6, 1 full super + tail of 2 blocks + {128, 1024, 1024}, {128, 192, 1024}, {128, 384, 1024}, + {128, 4096, 1024}, {128, 1024, 4096}, + // Customer decode (M=1) at K=384 (the case the K%4 gate previously blocked). + { 1, 1024, 384}, + // Synthetic K-tail stress shapes covering all (TailBlocks in {1, 2, 3}). + { 2, 16, 320}, // tail=1 + { 4, 16, 320}, + {128, 1024, 320}, + { 2, 16, 448}, // tail=3 + { 4, 16, 448}, + {128, 1024, 448}, + // N-tail stress (CountN % 4 != 0). The R2/R1 main tiles handle the + // NMain = floor(CountN/4)*4 cols; the per-1-col tail tile picks up + // the trailing 1-3 cols against the column-major tail region of the + // packed buffer. NMain = 0 cases (N in {1,2,3}) exercise the tail + // tile in isolation. + { 1, 1, 256}, // NMain=0, NTail=1, single-column decode + { 1, 3, 256}, // NMain=0, NTail=3 + { 4, 3, 256}, // NMain=0, NTail=3, R2+R1 head still empty + { 1, 17, 256}, // NMain=16, NTail=1, decode + { 4, 17, 256}, + {128, 17, 256}, + { 1, 33, 256}, // NMain=32, NTail=1 + { 4, 33, 256}, // exact shape that failed the dispatch swap + {128, 33, 256}, + { 1, 18, 256}, // NMain=16, NTail=2 + { 4, 18, 256}, + {128, 19, 256}, // NMain=16, NTail=3 + // N-tail combined with K-tail (the most generic case). + { 1, 17, 384}, + { 4, 33, 384}, + {128, 19, 448}, +}; + +// +// AVX-512BW (non-VNNI) SIMD super-block kernel. +// +TEST(MlasSq2BitTest, SuperBlock_BlkLen64_Avx512) +{ + if (!GetMlasPlatform().Avx512Supported_) { + GTEST_SKIP() << "AVX-512BW (DQ/VL) not available on this host"; + } + for (uint32_t seed : {0xC0FFEEu, 0xBADC0DEu}) { + for (const auto& s : kSimdShapes) { + for (bool bias : {false, true}) { + RunSuperBlockCase(s.M, s.N, s.K, bias, seed + (bias ? 1u : 0u), + /*WithZeroPoints=*/false, + sq2sb::SQ2BitGemmKernel_BlkSum_CompInt8_Super_Avx512_TestEntry, + "AVX-512BW"); + } + } + } +} + +TEST(MlasSq2BitTest, SuperBlock_BlkLen64_Avx512_WithZeroPoints) +{ + if (!GetMlasPlatform().Avx512Supported_) { + GTEST_SKIP() << "AVX-512BW (DQ/VL) not available on this host"; + } + for (uint32_t seed : {0xC0FFEEu, 0xBADC0DEu}) { + for (const auto& s : kSimdShapes) { + for (bool bias : {false, true}) { + RunSuperBlockCase(s.M, s.N, s.K, bias, seed + (bias ? 1u : 0u), + /*WithZeroPoints=*/true, + sq2sb::SQ2BitGemmKernel_BlkSum_CompInt8_Super_Avx512_TestEntry, + "AVX-512BW"); + } + } + } +} + +// +// AVX-512-VNNI SIMD super-block kernel. Gated on the platform having selected +// the VNNI dispatch table (the SIMD path uses `_mm512_dpbusd_epi32`). +// +TEST(MlasSq2BitTest, SuperBlock_BlkLen64_Avx512Vnni) +{ + if (GetMlasPlatform().QNBitGemmDispatch != &MlasSQNBitGemmDispatchAvx512vnni) { + GTEST_SKIP() << "AVX-512-VNNI not selected as the active dispatch on this host"; + } + for (uint32_t seed : {0xC0FFEEu, 0xBADC0DEu}) { + for (const auto& s : kSimdShapes) { + for (bool bias : {false, true}) { + RunSuperBlockCase(s.M, s.N, s.K, bias, seed + (bias ? 1u : 0u), + /*WithZeroPoints=*/false, + sq2sb::SQ2BitGemmKernel_BlkSum_CompInt8_Super_Avx512Vnni_TestEntry, + "AVX-512-VNNI"); + } + } + } +} + +TEST(MlasSq2BitTest, SuperBlock_BlkLen64_Avx512Vnni_WithZeroPoints) +{ + if (GetMlasPlatform().QNBitGemmDispatch != &MlasSQNBitGemmDispatchAvx512vnni) { + GTEST_SKIP() << "AVX-512-VNNI not selected as the active dispatch on this host"; + } + for (uint32_t seed : {0xC0FFEEu, 0xBADC0DEu}) { + for (const auto& s : kSimdShapes) { + for (bool bias : {false, true}) { + RunSuperBlockCase(s.M, s.N, s.K, bias, seed + (bias ? 1u : 0u), + /*WithZeroPoints=*/true, + sq2sb::SQ2BitGemmKernel_BlkSum_CompInt8_Super_Avx512Vnni_TestEntry, + "AVX-512-VNNI"); + } + } + } +} From a8c932fa8e406eed074dd2d381f49f9609d67ab4 Mon Sep 17 00:00:00 2001 From: Hari Seshadri Date: Wed, 10 Jun 2026 18:09:27 -0700 Subject: [PATCH 07/17] LUT fix --- onnxruntime/contrib_ops/cpu/quantization/matmul_nbits.cc | 9 ++++++++- 1 file changed, 8 insertions(+), 1 deletion(-) diff --git a/onnxruntime/contrib_ops/cpu/quantization/matmul_nbits.cc b/onnxruntime/contrib_ops/cpu/quantization/matmul_nbits.cc index 457d2c7c3af18..9ad418d8eaf2d 100644 --- a/onnxruntime/contrib_ops/cpu/quantization/matmul_nbits.cc +++ b/onnxruntime/contrib_ops/cpu/quantization/matmul_nbits.cc @@ -359,8 +359,15 @@ Status MatMulNBits::PrePack(const Tensor& tensor, int input_idx, /*out*/ All #endif // MLAS_TARGET_ARM64 } is_packed = true; - } else if (compute_type_ == SQNBIT_CompInt8) { + } else if (compute_type_ == SQNBIT_CompInt8 && !prefer_lut_gemm_) { // Packing scales and zero points + // Guard: for LUT-eligible nodes, scales/ZP are already packed inside + // packed_b_ by the LUT branch above (or by the LUT scale-pack path at + // the bottom of this function). Re-running the non-LUT pack here would + // overwrite the LUT-packed buffer with W2-super layout bytes, causing + // heap corruption when the LUT compute path later reads back the + // packed-B contents. prefer_lut_gemm_ is gated to T1==float (see ctor), + // so checking it here is sufficient. bool should_pack_scale_and_zp_inputs = [&]() { #if defined(MLAS_TARGET_AMD64_IX86) return true; From 37114895377bbb6c8e9b6c113f2225d5c544b2f9 Mon Sep 17 00:00:00 2001 From: Hari Seshadri Date: Wed, 10 Jun 2026 22:16:24 -0700 Subject: [PATCH 08/17] W2-v1+ Cleanup --- cmake/onnxruntime_mlas.cmake | 8 +- .../cpu/quantization/matmul_nbits.cc | 8 +- onnxruntime/core/mlas/lib/qnbitgemm.cpp | 9 +- onnxruntime/core/mlas/lib/qnbitgemm.h | 32 +- .../mlas/lib/sqnbitgemm_kernel_avx512.cpp | 67 +- .../lib/sqnbitgemm_kernel_avx512_2bit.cpp | 296 +++--- .../mlas/lib/sqnbitgemm_kernel_avx512_2bit.h | 352 +++---- ... sqnbitgemm_kernel_avx512_2bit_blklen64.h} | 335 +++--- ...nbitgemm_kernel_avx512_2bit_superblock.cpp | 359 ------- ...sqnbitgemm_kernel_avx512_2bit_superblock.h | 249 ----- .../mlas/lib/sqnbitgemm_kernel_avx512vnni.cpp | 69 +- ...nbitgemm_kernel_avx512vnni_2bit_blklen64.h | 989 ------------------ .../test/contrib_ops/matmul_2bits_test.cc | 68 ++ onnxruntime/test/mlas/bench/bench_lutgemm.cpp | 4 +- .../test/mlas/bench/bench_qnbitgemm.cpp | 216 +--- .../mlas/unittest/test_sqnbitgemm_2bit.cpp | 218 +--- .../unittest/test_sqnbitgemm_2bit_gemm.cpp | 636 ++++++----- .../test_sqnbitgemm_2bit_superblock.cpp | 529 ---------- 18 files changed, 1025 insertions(+), 3419 deletions(-) rename onnxruntime/core/mlas/lib/{sqnbitgemm_kernel_avx512vnni_2bit_blklen64_superblock.h => sqnbitgemm_kernel_avx512_2bit_blklen64.h} (73%) delete mode 100644 onnxruntime/core/mlas/lib/sqnbitgemm_kernel_avx512_2bit_superblock.cpp delete mode 100644 onnxruntime/core/mlas/lib/sqnbitgemm_kernel_avx512_2bit_superblock.h delete mode 100644 onnxruntime/core/mlas/lib/sqnbitgemm_kernel_avx512vnni_2bit_blklen64.h delete mode 100644 onnxruntime/test/mlas/unittest/test_sqnbitgemm_2bit_superblock.cpp diff --git a/cmake/onnxruntime_mlas.cmake b/cmake/onnxruntime_mlas.cmake index 5fe2b844b0111..f700a662becba 100644 --- a/cmake/onnxruntime_mlas.cmake +++ b/cmake/onnxruntime_mlas.cmake @@ -243,9 +243,7 @@ function(setup_mlas_source_for_windows) ${MLAS_SRC_DIR}/sqnbitgemm_kernel_avx2.cpp ${MLAS_SRC_DIR}/sqnbitgemm_kernel_avx512_2bit.h ${MLAS_SRC_DIR}/sqnbitgemm_kernel_avx512_2bit.cpp - ${MLAS_SRC_DIR}/sqnbitgemm_kernel_avx512_2bit_superblock.h - ${MLAS_SRC_DIR}/sqnbitgemm_kernel_avx512_2bit_superblock.cpp - ${MLAS_SRC_DIR}/sqnbitgemm_kernel_avx512vnni_2bit_blklen64.h + ${MLAS_SRC_DIR}/sqnbitgemm_kernel_avx512_2bit_blklen64.h ${MLAS_SRC_DIR}/sqnbitgemm_kernel_avx512.cpp ${MLAS_SRC_DIR}/sqnbitgemm_kernel_avx512vnni.cpp ${MLAS_SRC_DIR}/qkv_quant_kernel_avx512vnni.cpp @@ -798,9 +796,7 @@ else() ${MLAS_SRC_DIR}/sqnbitgemm_kernel_avx2.cpp ${MLAS_SRC_DIR}/sqnbitgemm_kernel_avx512_2bit.h ${MLAS_SRC_DIR}/sqnbitgemm_kernel_avx512_2bit.cpp - ${MLAS_SRC_DIR}/sqnbitgemm_kernel_avx512_2bit_superblock.h - ${MLAS_SRC_DIR}/sqnbitgemm_kernel_avx512_2bit_superblock.cpp - ${MLAS_SRC_DIR}/sqnbitgemm_kernel_avx512vnni_2bit_blklen64.h + ${MLAS_SRC_DIR}/sqnbitgemm_kernel_avx512_2bit_blklen64.h ${MLAS_SRC_DIR}/sqnbitgemm_lut_kernel_avx2.h ${MLAS_SRC_DIR}/sqnbitgemm_lut_kernel_avx2.cpp ${MLAS_SRC_DIR}/rotary_embedding_kernel_avx2.h diff --git a/onnxruntime/contrib_ops/cpu/quantization/matmul_nbits.cc b/onnxruntime/contrib_ops/cpu/quantization/matmul_nbits.cc index 9ad418d8eaf2d..6f33b34a049fb 100644 --- a/onnxruntime/contrib_ops/cpu/quantization/matmul_nbits.cc +++ b/onnxruntime/contrib_ops/cpu/quantization/matmul_nbits.cc @@ -364,10 +364,10 @@ Status MatMulNBits::PrePack(const Tensor& tensor, int input_idx, /*out*/ All // Guard: for LUT-eligible nodes, scales/ZP are already packed inside // packed_b_ by the LUT branch above (or by the LUT scale-pack path at // the bottom of this function). Re-running the non-LUT pack here would - // overwrite the LUT-packed buffer with W2-super layout bytes, causing - // heap corruption when the LUT compute path later reads back the - // packed-B contents. prefer_lut_gemm_ is gated to T1==float (see ctor), - // so checking it here is sufficient. + // overwrite the LUT-packed buffer with W2 layout bytes, causing heap + // corruption when the LUT compute path later reads back the packed-B + // contents. prefer_lut_gemm_ is gated to T1==float (see ctor), so + // checking it here is sufficient. bool should_pack_scale_and_zp_inputs = [&]() { #if defined(MLAS_TARGET_AMD64_IX86) return true; diff --git a/onnxruntime/core/mlas/lib/qnbitgemm.cpp b/onnxruntime/core/mlas/lib/qnbitgemm.cpp index 3200d3216f650..5f1357020accd 100644 --- a/onnxruntime/core/mlas/lib/qnbitgemm.cpp +++ b/onnxruntime/core/mlas/lib/qnbitgemm.cpp @@ -1003,11 +1003,12 @@ SQ2BitGemm_CompInt8( // The packed-B-data and per-block-scale buffers may use a layout that // addresses each N-col at a stride larger than `k_blks` (e.g. the AVX-512 - // super-block / W2-v2 kernel rounds up to a multiple of 4 to amortise - // unpack across 4 K-blocks). Ask the dispatch for the effective per-N-col - // block count; default to logical k_blks if the dispatch doesn't override. + // W2 kernel rounds up to a multiple of 4 to amortise unpack across 4 + // K-blocks). Ask the dispatch for the effective per-N-col block count; + // default to logical k_blks if the dispatch doesn't override. // BlkSum keeps the logical stride because the SGEMM correction step - // expects width-16 chunked PackB layout independent of the W2 variant. + // expects width-16 chunked PackB layout independent of the W2 packed-B + // stride convention. const size_t k_blks_eff = (Dispatch->Q2BitGemmEffectiveBlockCountK != nullptr) ? Dispatch->Q2BitGemmEffectiveBlockCountK(k_blks) : k_blks; diff --git a/onnxruntime/core/mlas/lib/qnbitgemm.h b/onnxruntime/core/mlas/lib/qnbitgemm.h index 1d5916a5d5ed5..0e56e20579d8b 100644 --- a/onnxruntime/core/mlas/lib/qnbitgemm.h +++ b/onnxruntime/core/mlas/lib/qnbitgemm.h @@ -51,15 +51,12 @@ struct PackedQuantBDataStruct { PackedQuantBDataStruct(void* PackedQuantBWorkspace, size_t N, size_t BlockCountK, size_t BlkLen, bool QuantAUnsigned) : QuantBWorkspace_(PackedQuantBWorkspace), N_(N), BlockCountK_(BlockCountK), BlkLen_(BlkLen) { - // For 2-bit weights, the AVX-512 super-block (W2-v2) layout requires - // BlockCountK to be a multiple of 4 (one super-block packs 4 consecutive - // K-blocks together). The legacy W2-v1 layout doesn't require this but - // accepts the small storage padding (<= 48 bytes per N-col), so we pad - // unconditionally for BlkBitWidth=2 to keep a single buffer ABI for both - // kernel variants. The matching pack-size dispatch functions - // (Q2BitGemmPackQuantBDataSize_Avx512 / _SuperBlock) round up by the - // same amount, so the allocated buffer always matches the slab layout - // computed below. + // For 2-bit weights, the AVX-512 W2 packed layout groups 4 consecutive + // K-blocks into a single 64-byte slot so the SIMD unpack is one ZMM + // load + four fixed shift/mask pairs. The pack-size dispatch + // (Q2BitGemmPackQuantBDataSize_Avx512) rounds BlockCountK up to a + // multiple of 4 internally; we mirror that rounding here so the + // allocated buffer always matches the slab layout computed below. const size_t EffectiveBlockCountK = (BlkBitWidth == 2) ? ((BlockCountK + 3) / 4) * 4 : BlockCountK; const size_t PackedQuantBDataSize = N * EffectiveBlockCountK * MlasQNBitBlkDataSizeInBytes(BlkBitWidth, BlkLen); @@ -505,20 +502,19 @@ struct MLAS_QNBIT_GEMM_DISPATCH { * @brief Returns the effective per-N-col block count used by the 2-bit packed * B-data and B-scale layouts. The layout addresses each N-col at a * stride of `effective_block_count * `, - * and so does the dispatcher's per-N-tile pointer arithmetic. Some - * 2-bit kernels (notably the AVX-512 super-block / W2-v2 layout) round - * BlockCountK up to a multiple of 4 internally to amortise unpack - * cost across 4 consecutive K-blocks; the buffer is sized accordingly - * (see PackedQuantBDataStruct, which always pads for BlkBitWidth==2) - * so the dispatcher must use the matching stride or it will step - * past the data when n != 0. + * and so does the dispatcher's per-N-tile pointer arithmetic. The + * AVX-512 W2 kernel rounds BlockCountK up to a multiple of 4 + * internally to amortise unpack cost across 4 consecutive K-blocks; + * the buffer is sized accordingly (see PackedQuantBDataStruct, which + * always pads for BlkBitWidth==2) so the dispatcher must use the + * matching stride or it will step past the data when n != 0. * * BlkSum is NOT affected: it uses the SGEMM-style width-16 chunked * layout with stride `BlockCountK * 16 * sizeof(float)` per chunk - * regardless of which kernel variant is active. + * regardless of which W2 dispatch implementation is active. * * Returns 0 if not set; the dispatcher falls back to the logical - * `BlockCountK` (the W2-v1 / column-major-ish stride convention). + * `BlockCountK` (the plain column-major-ish stride convention). */ typedef size_t(Q2BitGemmEffectiveBlockCountK_Fn)(size_t BlockCountK); diff --git a/onnxruntime/core/mlas/lib/sqnbitgemm_kernel_avx512.cpp b/onnxruntime/core/mlas/lib/sqnbitgemm_kernel_avx512.cpp index 303e4ee6a2bdc..c979841e20d24 100644 --- a/onnxruntime/core/mlas/lib/sqnbitgemm_kernel_avx512.cpp +++ b/onnxruntime/core/mlas/lib/sqnbitgemm_kernel_avx512.cpp @@ -27,9 +27,7 @@ Module Name: #include "sqnbitgemm_kernel_avx512_int8_blklen64.h" #include "sqnbitgemm_kernel_avx512_int8_blklen128.h" #include "sqnbitgemm_kernel_avx512_2bit.h" -#include "sqnbitgemm_kernel_avx512vnni_2bit_blklen64.h" -#include "sqnbitgemm_kernel_avx512_2bit_superblock.h" -#include "sqnbitgemm_kernel_avx512vnni_2bit_blklen64_superblock.h" +#include "sqnbitgemm_kernel_avx512_2bit_blklen64.h" // // SQNBIT_CompFp32 kernel implementation. @@ -480,10 +478,8 @@ SQ8BitGemmPackQuantBDataAndBlkSum512( } // -// Unit-test entry point for the AVX-512BW (non-VNNI) W2 kernel. Linked into -// the test binary so the test TU (compiled without AVX-512 flags) can invoke -// this kernel directly without going through the platform dispatcher. The -// caller must verify AVX-512BW is available on the host before calling. +// Unit-test entry point for the AVX-512BW (non-VNNI) W2 kernel. +// Sibling of the VNNI variant in sqnbitgemm_kernel_avx512vnni.cpp. // namespace onnxruntime::mlas::sq2bit_avx512 { size_t MLASCALL @@ -510,35 +506,6 @@ SQ2BitGemmKernel_BlkSum_CompInt8_Avx512_TestEntry( } } // namespace onnxruntime::mlas::sq2bit_avx512 -// -// Unit-test entry point for the AVX-512BW (non-VNNI) W2 SUPER-BLOCK kernel. -// Sibling of the VNNI variant in sqnbitgemm_kernel_avx512vnni.cpp. -// -namespace onnxruntime::mlas::sq2bit_avx512_super { -size_t MLASCALL -SQ2BitGemmKernel_BlkSum_CompInt8_Super_Avx512_TestEntry( - size_t BlkLen, - const std::byte* QuantA, - const float* QuantAScale, - const std::byte* QuantBData, - const float* QuantBScale, - const std::byte* QuantBZeroPoint, - float* C, - size_t CountM, - size_t CountN, - size_t CountK, - size_t BlockCountK, - const float* Bias, - size_t ldc, - const float* ABlockSum, - const float* QuantBBlkSum) -{ - return SQ2BitGemmKernel_BlkSum_CompInt8_Super_Avx512( - BlkLen, QuantA, QuantAScale, QuantBData, QuantBScale, QuantBZeroPoint, - C, CountM, CountN, CountK, BlockCountK, Bias, ldc, ABlockSum, QuantBBlkSum); -} -} // namespace onnxruntime::mlas::sq2bit_avx512_super - const MLAS_QNBIT_GEMM_DISPATCH MlasSQNBitGemmDispatchAvx512 = []() { MLAS_QNBIT_GEMM_DISPATCH d; @@ -558,25 +525,15 @@ const MLAS_QNBIT_GEMM_DISPATCH MlasSQNBitGemmDispatchAvx512 = []() { d.SQ8BitGemmKernel_BlkSum_CompInt8 = SQ8BitGemmKernel_BlkSum_CompInt8_avx512; d.QuantizeARowComputeBlkSum_CompInt8 = QuantizeARow_CompInt8_avx512; - // 2-bit native CompInt8 path: AVX-512BW (no VNNI) variant of the - // super-block (W2-v2) kernel -- single 64-byte load + four fixed - // shift/mask pairs to unpack 4 K-blocks at once. Pack-size and pack - // functions are shared with the AVX-512-VNNI variant; only the inner - // integer MAC differs (`vpmaddubsw + vpmaddwd + vpaddd` here vs - // `_mm512_dpbusd_epi32` in the VNNI variant). - // - // The legacy `sq2bit_avx512::*` (W2-v1) symbols remain in the build but - // are no longer reached at runtime. A follow-up will remove them once - // W2-v2 has soaked in production. - d.Q2BitGemmPackQuantBDataSize = onnxruntime::mlas::sq2bit_avx512_super::Q2BitGemmPackQuantBDataSize_SuperBlock; - d.SQ2BitGemmPackQuantBDataAndBlkSum = onnxruntime::mlas::sq2bit_avx512_super::SQ2BitGemmPackQuantBDataAndBlkSum_SuperBlockScalar; - d.SQ2BitGemmKernel_BlkSum_CompInt8 = onnxruntime::mlas::sq2bit_avx512_super::SQ2BitGemmKernel_BlkSum_CompInt8_Super_Avx512_TestEntry; - // W2-v2 packs and addresses each N-col at a stride of - // SuperBlockCountKPadded * kSuperBlockBlks blocks (BlockCountK rounded - // UP to a multiple of 4) to keep the inner K-loop's super-block stride - // constant. The dispatcher needs this stride for its per-N-tile pointer - // arithmetic in SQ2BitGemm_CompInt8. - d.Q2BitGemmEffectiveBlockCountK = [](size_t BlockCountK) { return ((BlockCountK + 3) / 4) * 4; }; + // 2-bit native CompInt8 path. Single dispatch entry (W2): + // 64-byte ZMM load + four fixed shift/mask pairs to unpack 4 K-blocks at + // once. Packs B with the K dimension rounded up to a multiple of + // kBlockGroupBlks (= 4); the kernel iterates the padded count, so the + // dispatcher reports the rounded-up stride via Q2BitGemmEffectiveBlockCountK. + d.Q2BitGemmPackQuantBDataSize = onnxruntime::mlas::sq2bit_avx512::Q2BitGemmPackQuantBDataSize_Avx512; + d.SQ2BitGemmPackQuantBDataAndBlkSum = onnxruntime::mlas::sq2bit_avx512::SQ2BitGemmPackQuantBDataAndBlkSum_Scalar; + d.SQ2BitGemmKernel_BlkSum_CompInt8 = onnxruntime::mlas::sq2bit_avx512::SQ2BitGemmKernel_BlkSum_CompInt8_Avx512_TestEntry; + d.Q2BitGemmEffectiveBlockCountK = [](size_t BlockCountK) { return ((BlockCountK + 3) / 4) * 4; }; return d; }(); diff --git a/onnxruntime/core/mlas/lib/sqnbitgemm_kernel_avx512_2bit.cpp b/onnxruntime/core/mlas/lib/sqnbitgemm_kernel_avx512_2bit.cpp index 7cb69e6288587..19fc022c62ab9 100644 --- a/onnxruntime/core/mlas/lib/sqnbitgemm_kernel_avx512_2bit.cpp +++ b/onnxruntime/core/mlas/lib/sqnbitgemm_kernel_avx512_2bit.cpp @@ -10,22 +10,17 @@ Module Name: Abstract: - Reference implementation and pack-time helpers for the 2-bit weight - CompInt8 GEMM (BlkBitWidth=2, BlkLen=64). Scalar-only; the vectorized - inner loop lives in sqnbitgemm_kernel_avx512vnni_2bit_blklen64.h and is - registered into the AVX-512 and AVX-512-VNNI dispatch tables (with VNNI - / non-VNNI MAC variants templated from the same source). - - The scalar functions exposed here back the pack-time helpers used by - both dispatch tables and are reachable as plain C++ symbols from unit - tests, which use them as a correctness oracle for the vectorized paths. - - Restrictions: - * BlkLen == 64 only. Other BlkLens are rejected by the pack helper - and the kernel returns 0 rows handled. - * Per-block zero-point input is supported (standard ONNX W2 layout, - 4 ZPs per packed byte along K). When no zero-point tensor is - supplied the symmetric default of 2 is used. + Pack helpers and a scalar reference kernel for the block-group W2 layout. + + See sqnbitgemm_kernel_avx512_2bit.h for the layout description + and the rationale (closing the W2-vs-W4 prefill gap by replacing the + per-block broadcast + variable shift unpack with a single 64-byte load + + four fixed-shift+mask pairs). + + This translation unit is scalar / portable. The vectorized inner loop + that consumes the block-group layout lives in a separate header + (sqnbitgemm_kernel_avx512_2bit_blklen64.h) and is wired into the AVX-512 + and AVX-512-VNNI dispatch tables. --*/ @@ -44,16 +39,28 @@ namespace mlas { namespace sq2bit_avx512 { // -// Workspace / pack-buffer size for the 2-bit CompInt8 path. +// Workspace / pack-buffer size for the block-group W2 path. Returns 0 if any +// of the configuration constraints is violated; the caller (MlasQNBitGemmPackQuantBDataSize) +// treats that as "unsupported" and falls back to the original W2 path. // -// Layout (in bytes): +// Constraints: +// * BlkLen == 64 +// * ComputeType == SQNBIT_CompInt8 // -// [PackedQuantBData] N * BlockCountK * kBlkBytes (BlkLen / 4 bytes per block) -// [BlkSum (float)] roundup_16(N) * BlockCountK -// [Scales (float)] N * BlockCountK +// K-tail handling: BlockCountK is rounded UP to a multiple of kBlockGroupBlks +// for the storage that the inner K-loop walks (PackedQuantBData, +// PackedQuantBScale). Padding slots hold zeroed weights and scales, so they +// contribute exactly 0 to the dot product. The BlkSum buffer is kept at the +// LOGICAL BlockCountK because it is consumed by the SGEMM correction step, +// not by the inner K-loop. // -// Alignment slack is added so that the AVX-512 dequant can use -// 64-byte aligned loads. +// Storage matches the original W2 layout total bytes when BlockCountK is a +// multiple of 4. When not a multiple of 4, storage grows by at most 3 K-blocks +// per N-col (= up to 48 bytes per col -- negligible at production N values). +// +// [PackedQuantBData] N * BlockCountKPadded * kBlkBytes +// [PackedQuantBScale] N * BlockCountKPadded * sizeof(float) +// [QuantBBlkSum] roundup_16(N) * BlockCountK (logical) * 16 floats // size_t MLASCALL Q2BitGemmPackQuantBDataSize_Avx512( @@ -65,24 +72,29 @@ Q2BitGemmPackQuantBDataSize_Avx512( const MLAS_BACKEND_KERNEL_SELECTOR_CONFIG* /* BackendKernelSelectorConfig */ ) { - // Only BlkLen=64 and SQNBIT_CompInt8 are supported. Anything else returns - // 0 so MlasQNBitGemmPackQuantBDataSize reports an unsupported configuration. if (BlkLen != kBlkLen || ComputeType != SQNBIT_CompInt8) { return 0; } - const size_t BlockCountK = MlasDivRoundup(K, BlkLen); - // Pad BlockCountK to a multiple of 4 to match the W2 buffer ABI used by - // PackedQuantBDataStruct (BlkBitWidth=2 always pads). The padding is a - // few extra K-block slots per N-col (<= 48 bytes data + a few floats for - // scales / BlkSum) -- negligible -- but lets the W2-v1 (this path) and - // W2-v2 (super-block) pack helpers share a single buffer layout. - const size_t BlockCountKPadded = ((BlockCountK + 3) / 4) * 4; + if (BlockCountK == 0) { + return 0; + } + const size_t BlockCountKPadded = + MlasDivRoundup(BlockCountK, kBlockGroupBlks) * kBlockGroupBlks; + + // Use BlockCountKPadded for BlkSum sizing too. The actual SGEMM-correction + // step only reads LOGICAL BlockCountK entries, but PackedQuantBDataStruct + // is constructed by the caller with a single BlockCountK value that + // controls BOTH the packed-B size and the BlkSum offset. If we sized + // BlkSum at the logical BlockCountK while sizing packed-B at padded, + // the struct's BlkSum pointer would land inside the packed-B region + // (because the caller's struct uses one BlockCountK consistently). The + // extra storage from padding the BlkSum is ~16 floats per N -- trivial. size_t PackedQuantBDataSize = N * BlockCountKPadded * kBlkBytes; const size_t ScaleSize = N * BlockCountKPadded * sizeof(float); size_t BlkSumSize = MlasDivRoundup(N, 16) * BlockCountKPadded * 16 * sizeof(float); - constexpr size_t kPackedQuantBDataAlignment = 64; // AVX-512 friendly + constexpr size_t kPackedQuantBDataAlignment = 64; PackedQuantBDataSize += kPackedQuantBDataAlignment - 1; constexpr size_t kBlkSumAlignment = MlasQNBitQuantBBlkSumAlignment(); @@ -92,26 +104,25 @@ Q2BitGemmPackQuantBDataSize_Avx512( } // -// Pack quantized B data + scales + per-block sums for the 2-bit kernel. +// Pack quantized B data + scales + per-block sums for the block-group W2 path. // -// Layouts produced (all column-major in N): -// PackedQuantBData[n * BlockCountK * kBlkBytes + blk * kBlkBytes + i] -// Block (n, blk), byte i of the kPackedBlkBytes layout (see header). -// PackedQuantBScale[n * BlockCountK + blk] -// Copy of the input scale; column-major. -// QuantBBlkSum[n * BlockCountK + blk] -// = -scale * 2 (symmetric W2 uses an implicit zero point of 2). +// PackedQuantBData layout (block-groups of 4 K-blocks, 64 bytes each): +// The block-group at logical (n, blk_group=blk/4) lives at byte offset +// PackedQuantBOffsetBytes_W2(n, blk_group, BlockGroupCountK, NMain). +// Byte b within the block-group holds 2-bit weight b from each of the 4 +// constituent K-blocks at bit positions {0..1, 2..3, 4..5, 6..7}. // -// When QuantBZPBegin is non-null, the per-block zero-point byte stream is in -// the standard ONNX MatMulNBits W2 layout: 4 zero-points per byte, packed -// along the K-block axis, row-major in N. Row stride is -// ZPCountK = ceil(BlockCountK / 4) bytes; the ZP for (n, blk) lives at byte -// index (n * ZPCountK + blk / 4), at bit offset (blk % 4) * 2. +// PackedQuantBScale layout: one float per K-block, four floats per block-group, +// addressed by PackedQuantBScaleOffset_W2. // -// When QuantBZPBegin is null, we fall back to the symmetric default -// (kDefaultSymmetricZeroPoint2Bit = 2), preserving the prior behavior. The -// HasZeroPoint flag is informational only: if QuantBZPBegin is non-null we -// always consume it. +// QuantBBlkSum layout: the same width-16 row-major chunked layout used by the +// existing W2 path, so the SGEMM correction step (MlasGemmFloatKernel) can be +// shared verbatim with the existing kernel. +// +// Mirrors the SQ2BitGemmPackQuantBDataAndBlkSum_Scalar prepack 3-call pattern: +// ORT's matmul_nbits.cc invokes this function up to three times (B, scales, ZP). +// We write scales when scales arrive, then re-derive BlkSum whenever either +// scales or zero-points arrive, reading scales from the already-packed buffer. // void MLASCALL SQ2BitGemmPackQuantBDataAndBlkSum_Scalar( @@ -134,79 +145,94 @@ SQ2BitGemmPackQuantBDataAndBlkSum_Scalar( } const size_t BlockCountK = MlasDivRoundup(K, BlkLen); + if (BlockCountK == 0) { + return; + } + + // Pad BlockCountK up to a multiple of kBlockGroupBlks so the inner K-loop + // can iterate whole block-groups uniformly. Padding slots store zeroed + // weights / scales and contribute exactly 0 to the dot product. + const size_t BlockGroupCountKPadded = + MlasDivRoundup(BlockCountK, kBlockGroupBlks); + const size_t BlockCountKPadded = BlockGroupCountKPadded * kBlockGroupBlks; const size_t NMain = (N / kNCols4) * kNCols4; - const size_t Iterations = N * BlockCountK; - // Pack weight bytes in the 4-col-grouped + 2-K-block-paired layout for - // the main NMain cols; column-major for the tail N % 4 cols. See - // PackedQuantBOffsetBytes_W2 in sqnbitgemm_kernel_avx512_2bit.h for the - // exact mapping. + // Zero source block used when packing a group whose K-range crosses the + // logical BlockCountK boundary. We point to this static zero buffer in + // the missing slots so the existing 4-block pack helper does the right + // thing without any branching inside it. + static const std::byte kZeroBlock[kBlkBytes] = {}; + + // ----- B-data pack ----- if (QuantBDataBegin != nullptr) { std::byte* PackedQuantBData = PackedQuantB.PackedQuantBData; + const size_t Iterations = N * BlockGroupCountKPadded; MlasTrySimpleParallel( ThreadPool, static_cast(Iterations), [&](ptrdiff_t tid) { - const size_t n = static_cast(tid) / BlockCountK; - const size_t blk = static_cast(tid) % BlockCountK; - const size_t src_offset = (n * BlockCountK + blk) * kBlkBytes; - const size_t dst_offset = PackedQuantBOffsetBytes_W2(n, blk, BlockCountK, NMain); - PackBlock_BlkLen64(QuantBDataBegin + src_offset, PackedQuantBData + dst_offset); + const size_t n = static_cast(tid) / BlockGroupCountKPadded; + const size_t blk_group = static_cast(tid) % BlockGroupCountKPadded; + const size_t blk0 = blk_group * kBlockGroupBlks; + + // Pick real source block pointers for slots that exist; the + // static zero buffer for slots past the logical BlockCountK. + auto src_for = [&](size_t blk) -> const std::byte* { + if (blk < BlockCountK) { + return QuantBDataBegin + (n * BlockCountK + blk) * kBlkBytes; + } + return kZeroBlock; + }; + const std::byte* src_blk_0 = src_for(blk0 + 0); + const std::byte* src_blk_1 = src_for(blk0 + 1); + const std::byte* src_blk_2 = src_for(blk0 + 2); + const std::byte* src_blk_3 = src_for(blk0 + 3); + + const size_t dst_offset = + PackedQuantBOffsetBytes_W2(n, blk_group, BlockGroupCountKPadded, NMain); + PackBlockGroup_BlkLen64(src_blk_0, src_blk_1, src_blk_2, src_blk_3, + PackedQuantBData + dst_offset); } ); } - // Scales follow the same 4-col-grouped layout (see PackedQuantBScaleOffset_W2). - // - // BlkSum uses the W4-style "width-16 row-major chunked" layout because the - // top-level kernel performs the zero-point correction via the float SGEMM - // micro-kernel (`GetMlasPlatform().GemmFloatKernel`), which expects this - // pre-packed B layout: - // - // BlkSum[(n / 16) * BlockCountK * 16 + blk * 16 + (n % 16)] - // = -scale_b * zero_point - // - // The allocated BlkSum buffer is sized at MlasDivRoundup(N, 16) * BlockCountK - // * 16 floats so the layout is well-defined even when N % 16 != 0 (the tail - // chunk's unused lanes hold whatever the buffer was initialised with, which - // for production callers must be zero so the SGEMM correction reads zeros). - // - // Important: ORT's matmul_nbits.cc prepack flow invokes this function up to - // three times per node — once each for B, scales, and zero_points — so on - // any given invocation only one of (QuantBScaleBegin, QuantBZPBegin) may be - // non-null. To get the correct BlkSum we therefore (a) copy scales into the - // packed buffer when QuantBScaleBegin is provided, and (b) recompute BlkSum - // whenever EITHER scales or zero-points are provided, reading the scales - // from the already-packed PackedQuantBScale buffer (which is populated by a - // previous call when only ZPs arrive in the current call). This mirrors - // the W4 ComputePackBlkSum helper. - + // ----- Scales ----- + // Iterate over the PADDED block count so trailing padding slots get + // explicit zero scales (otherwise they could hold uninitialised noise + // and the kernel's K-loop would read those into the FMA). if (QuantBScaleBegin != nullptr) { float* PackedScales = PackedQuantB.PackedQuantBScale; + const size_t Iterations = N * BlockCountKPadded; MlasTrySimpleParallel( ThreadPool, static_cast(Iterations), [&](ptrdiff_t tid) { - const size_t n = static_cast(tid) / BlockCountK; - const size_t blk = static_cast(tid) % BlockCountK; - const float scale = QuantBScaleBegin[n * BlockCountK + blk]; - PackedScales[PackedQuantBScaleOffset_W2(n, blk, BlockCountK, NMain)] = scale; + const size_t n = static_cast(tid) / BlockCountKPadded; + const size_t blk = static_cast(tid) % BlockCountKPadded; + const float scale = (blk < BlockCountK) + ? QuantBScaleBegin[n * BlockCountK + blk] + : 0.0f; + PackedScales[PackedQuantBScaleOffset_W2(n, blk, BlockCountKPadded, NMain)] = scale; } ); } - // BlkSum needs to be (re)computed whenever scales or zero-points change. - // Source of scales is the already-packed buffer, so this works even when - // only zero_points are provided in the current invocation. + // ----- BlkSum (recomputed whenever scales or ZPs arrive) ----- + // BlkSum is consumed by the SGEMM correction step (MlasGemmFloatKernel), + // which is called outside the inner K-loop with the LOGICAL BlockCountK + // and the per-row ABlockSum the dispatcher produced for that logical K. + // We therefore only need to fill the first BlockCountK entries; the buffer + // is sized at MlasDivRoundup(N, 16) * BlockCountK * 16 floats (logical). if (QuantBScaleBegin != nullptr || QuantBZPBegin != nullptr) { const float* PackedScales = PackedQuantB.PackedQuantBScale; float* BlkSum = PackedQuantB.QuantBBlkSum; const size_t ZPCountK = MlasDivRoundup(BlockCountK, 4); + const size_t Iterations = N * BlockCountK; MlasTrySimpleParallel( ThreadPool, static_cast(Iterations), [&](ptrdiff_t tid) { const size_t n = static_cast(tid) / BlockCountK; const size_t blk = static_cast(tid) % BlockCountK; const float scale = - PackedScales[PackedQuantBScaleOffset_W2(n, blk, BlockCountK, NMain)]; + PackedScales[PackedQuantBScaleOffset_W2(n, blk, BlockCountKPadded, NMain)]; uint8_t zp = kDefaultSymmetricZeroPoint2Bit; if (QuantBZPBegin != nullptr) { @@ -224,21 +250,13 @@ SQ2BitGemmPackQuantBDataAndBlkSum_Scalar( } // -// Scalar reference kernel for SQ2BitGemmVariant_CompInt8. +// Scalar reference kernel that consumes the block-group packed layout. +// Same math as the existing reference kernel; differs only in how it walks +// PackedQuantBData (block-group-major) and PackedQuantBScale (block-group-major). // -// Inputs match the SQ4BitGemmKernel_BlkSum_CompInt8_Fn typedef so the -// vectorized implementation can drop into the same dispatch slot. -// -// Math: -// C[m, n] = bias[n] -// + sum_blk( scale_a[m, blk] * scale_b[n, blk] -// * dot(int8 a[m, blk, :], uint8 b_unpacked[n, blk, :]) ) -// + sum_blk( ABlockSum[m, blk] * QuantBBlkSum[n, blk] ) -// -// The third term applies the W2 zero-point correction: -// QuantBBlkSum[n, blk] = -scale_b * zp, and ABlockSum[m, blk] = scale_a * sum(a). -// `zp` is the per-block zero point baked in at pack time (defaults to -// kDefaultSymmetricZeroPoint2Bit = 2 when no zero-point tensor is supplied). +// This is the correctness oracle for the SIMD block-group kernel coming in +// SIMD path. It also lets us validate the pack layout end-to-end via the +// existing MlasQNBitGemmBatch dispatch path once we wire it up. // size_t MLASCALL SQ2BitGemmKernel_BlkSum_CompInt8_Scalar( @@ -260,13 +278,33 @@ SQ2BitGemmKernel_BlkSum_CompInt8_Scalar( ) { if (BlkLen != kBlkLen) { - return 0; // Only BlkLen=64 is supported. + return 0; + } + if (BlockCountK == 0) { + return 0; } - const size_t lda = BlockCountK * kBlkLen; // bytes per A row (int8) - const size_t lda_scale = BlockCountK; // floats per A scale row - const size_t ldb = BlockCountK * kBlkBytes; // bytes per B column - const size_t ldb_scale = BlockCountK; // floats per B column + // PackedQuantBData and PackedQuantBScale are addressed via padded counts + // (K-tail handling -- see Q2BitGemmPackQuantBDataSize_Avx512). The K + // dot-product loop itself iterates only LOGICAL BlockCountK steps because + // A is unpadded; the kernel never reads past the real A rows. + const size_t BlockGroupCountKPadded = + MlasDivRoundup(BlockCountK, kBlockGroupBlks); + const size_t BlockCountKPadded = BlockGroupCountKPadded * kBlockGroupBlks; + + // The kernel is called by SQ2BitGemm_CompInt8 with the full CountN range; that + // function selects an N-tile boundary (kNCols4) up the stack. For a scalar + // reference path we don't depend on the 4-N-col grouping, but we DO need to + // index PackedQuantBData/PackedQuantBScale via the block-group offset helpers + // so we read the right bytes regardless of caller tile choice. + // + // CountN may not be a multiple of kNCols4 in the tail case. Detect that and + // fall back to plain column-major for the tail cols (the layout helpers + // already encode this rule). + const size_t NMainLocal = (CountN / kNCols4) * kNCols4; + + const size_t lda = BlockCountK * kBlkLen; // bytes per A row (int8) + const size_t lda_scale = BlockCountK; // floats per A scale row for (size_t m = 0; m < CountM; ++m) { const int8_t* a_row = reinterpret_cast(QuantA + m * lda); @@ -275,33 +313,37 @@ SQ2BitGemmKernel_BlkSum_CompInt8_Scalar( float* c_row = C + m * ldc; for (size_t n = 0; n < CountN; ++n) { - const std::byte* b_col = QuantBData + n * ldb; - const float* b_scale_col = QuantBScale + n * ldb_scale; - const float* b_blksum_col = QuantBBlkSum + n * ldb_scale; - float acc = (Bias != nullptr) ? Bias[n] : 0.0f; for (size_t blk = 0; blk < BlockCountK; ++blk) { - // Unpack 64 2-bit weights into 64 uint8 values (values in [0, 3]). + // Pull the block-group this K-block belongs to and unpack only the + // slot we need (block_in_group = blk % 4 selects the 2-bit field). + const size_t blk_group = blk / kBlockGroupBlks; + const size_t blk_in_group = blk % kBlockGroupBlks; + const size_t block_group_offset = + PackedQuantBOffsetBytes_W2(n, blk_group, BlockGroupCountKPadded, NMainLocal); + const std::byte* block_group = QuantBData + block_group_offset; + uint8_t b_unpacked[kBlkLen]; - UnpackBlock_BlkLen64_Reference(b_col + blk * kBlkBytes, b_unpacked); + for (size_t i = 0; i < kBlkLen; ++i) { + const uint8_t byte = static_cast(block_group[i]); + b_unpacked[i] = static_cast((byte >> (2 * blk_in_group)) & 0x03u); + } - // int8 * uint8 dot product across the block. const int8_t* a_blk = a_row + blk * kBlkLen; int32_t dot = 0; for (size_t i = 0; i < kBlkLen; ++i) { dot += static_cast(a_blk[i]) * static_cast(b_unpacked[i]); } - // Integer term * scales. - acc += a_scale_row[blk] * b_scale_col[blk] * static_cast(dot); + const float b_scale = + QuantBScale[PackedQuantBScaleOffset_W2(n, blk, BlockCountKPadded, NMainLocal)]; + acc += a_scale_row[blk] * b_scale * static_cast(dot); - // W2 zero-point correction: - // dot(a, b_signed) = dot(a, b_unsigned) - zp * sum(a) - // so we need C += scale_a * scale_b * (-zp) * sum(a) - // = ABlockSum * QuantBBlkSum (QuantBBlkSum encodes -scale_b * zp, - // with zp = per-block ZP or 2 when no ZP tensor was supplied). - acc += a_blksum_row[blk] * b_blksum_col[blk]; + // The width-16 row-major BlkSum layout is column-major in n + // (one float per (n, blk)); same as the existing W2 path. + const size_t blksum_offset = ((n / 16) * BlockCountK + blk) * 16 + (n % 16); + acc += a_blksum_row[blk] * QuantBBlkSum[blksum_offset]; } c_row[n] = acc; diff --git a/onnxruntime/core/mlas/lib/sqnbitgemm_kernel_avx512_2bit.h b/onnxruntime/core/mlas/lib/sqnbitgemm_kernel_avx512_2bit.h index 7b71c79b71ff8..34b08771802ca 100644 --- a/onnxruntime/core/mlas/lib/sqnbitgemm_kernel_avx512_2bit.h +++ b/onnxruntime/core/mlas/lib/sqnbitgemm_kernel_avx512_2bit.h @@ -10,34 +10,33 @@ Module Name: Abstract: - Pack-time helpers and reference (scalar) routines for the 2-bit, BlkLen=64 - AVX-512 weight GEMM path (SQNBIT_CompInt8, BlkBitWidth=2). The vectorized - kernels (AVX-512 and AVX-512-VNNI variants) live in - sqnbitgemm_kernel_avx512vnni_2bit_blklen64.h. - - This header is scalar / header-only and contains: - - * The packed-block layout used by the AVX-512 dequant inner loop. - * A pack routine that converts standard ONNX MatMulNBits 2-bit input data - into the packed layout. - * A reference unpack routine (independent of the pack code) that - materialises one packed block back into 64 individual int8 values. - - The packed layout is designed so that the AVX-512BW dequant is one - 128-bit broadcast plus a per-lane variable shift: - - __m128i p = _mm_loadu_si128(packed); // 16 bytes - __m512i p4 = _mm512_broadcast_i32x4(p); // 4 lanes - __m512i sh = _mm512_set_epi32(6,6,6,6, 4,4,4,4, - 2,2,2,2, 0,0,0,0); - __m512i v = _mm512_srlv_epi32(p4, sh); // per-lane shift - v = _mm512_and_si512(v, _mm512_set1_epi8(0x03)); - - With this layout the resulting ZMM holds weights 0..63 in their natural - order. - - NOTE: This file intentionally restricts itself to BlkLen == 64. Other - BlkLens fall through to the existing LUT path until later phases. + Pack-time helpers and scalar reference routines for the block-group W2 layout: groups of 4 K-blocks share a single 64-byte + packed buffer that allows the AVX-512 unpack to be one ZMM load plus + four fixed `vpsrlw+vpand` pairs (instead of the current per-block + broadcast + variable shift). + + Layout summary (BlkLen=64 only): + + * Each "block-group" packs FOUR consecutive K-blocks (256 weights total). + * Total storage per block-group = kBlkBytes * 4 = 64 bytes (identical to + 4 separately-packed blocks under the current scheme). + * Byte b of the block-group holds: + bits[0..1] = block_0.weight[b] + bits[2..3] = block_1.weight[b] + bits[4..5] = block_2.weight[b] + bits[6..7] = block_3.weight[b] + * The N-dimension uses the same 4-col-grouped layout as the existing + W2 kernel (kNCols4 = 4), so the main NMain region groups 4 N-cols + per "row" of block-groups. + + Restrictions: + + * BlkLen == 64 only. + * BlockCountK must be a multiple of kBlockGroupBlks = 4. The + path returns 0 from the pack-size helper for non-multiples; the + caller falls back to the existing W2 path. + * The customer model's K dimensions (384, 1024, 4096) are all multiples + of 256 (= 4 * 64), so all customer shapes satisfy this constraint. --*/ @@ -59,94 +58,29 @@ namespace onnxruntime { namespace mlas { namespace sq2bit_avx512 { +// ----------------------------------------------------------------------------- +// Block / block-group constants (BlkLen=64 W2 native path). +// ----------------------------------------------------------------------------- + // Each 2-bit weight occupies 2 bits; one byte holds 4 weights. constexpr size_t kWeightsPerByte = 4; // Block constants for the BlkLen=64 variant. constexpr size_t kBlkLen = 64; constexpr size_t kBlkBytes = kBlkLen / kWeightsPerByte; // 16 packed src bytes per block -constexpr size_t kPackedBlkBytes = kBlkBytes; // packing is in-place: 16 -> 16 // Default zero point used when the input is symmetric (no zero-point tensor). // For 2-bit unsigned values in [0, 3], the symmetric mid-point is 2. constexpr uint8_t kDefaultSymmetricZeroPoint2Bit = 2; -// ----------------------------------------------------------------------------- -// Tile shape (must match the AVX-512-VNNI kernel in -// sqnbitgemm_kernel_avx512vnni_2bit_blklen64.h). The pack layout below is -// keyed off these values: the main NMain = floor(N / kNCols4) * kNCols4 cols -// are stored in a 4-col-grouped + 2-K-block-paired arrangement that lets the -// kernel's R2xC4 hot loop read each tile of B as a single contiguous stream, -// matching W4's pack layout in PackQuantB (sqnbitgemm_kernel_avx_common.h). -// ----------------------------------------------------------------------------- - +// Tile shape used by the SIMD kernel; the pack layout is keyed off these. constexpr size_t kNCols4 = 4; -constexpr size_t kPerAccuBlk2 = 2; -// -// Byte offset into the packed B-data buffer for a logical (n, blk) cell. -// -// Main region (n < NMain = floor(N/kNCols4)*kNCols4): -// 4-N-col group g = n / 4, col within group c = n % 4. -// K-block-pair p = blk / 2, block within pair = blk % 2. -// Pair slots run first, then a single-block trailing slot when -// BlockCountK is odd. -// -// Tail region (n >= NMain): plain column-major. The tail base lies -// exactly at NMain * BlockCountK * kBlkBytes, so the dispatcher's -// `multipleCols * ColStrideBytes` offset (where ColStrideBytes is -// BlockCountK * kBlkBytes) is unchanged from the column-major layout. -// -inline size_t -PackedQuantBOffsetBytes_W2(size_t n, size_t blk, size_t BlockCountK, size_t NMain) -{ - if (n < NMain) { - const size_t g = n / kNCols4; - const size_t c = n % kNCols4; - const size_t per_group_bytes = BlockCountK * kNCols4 * kBlkBytes; - const size_t pair_idx = blk / kPerAccuBlk2; - const size_t blk_in_pair = blk % kPerAccuBlk2; - const size_t full_pairs = BlockCountK / kPerAccuBlk2; - if (pair_idx < full_pairs) { - return g * per_group_bytes - + pair_idx * (kNCols4 * kPerAccuBlk2 * kBlkBytes) - + c * (kPerAccuBlk2 * kBlkBytes) - + blk_in_pair * kBlkBytes; - } - return g * per_group_bytes - + full_pairs * (kNCols4 * kPerAccuBlk2 * kBlkBytes) - + c * kBlkBytes; - } - return (n * BlockCountK + blk) * kBlkBytes; -} - -// -// Float offset into the packed B-scale buffer for a logical (n, blk) cell. -// Same grouping rule as the B-data, two scales per pair, one scale per -// single-block trailing slot. -// -inline size_t -PackedQuantBScaleOffset_W2(size_t n, size_t blk, size_t BlockCountK, size_t NMain) -{ - if (n < NMain) { - const size_t g = n / kNCols4; - const size_t c = n % kNCols4; - const size_t per_group_scales = BlockCountK * kNCols4; - const size_t pair_idx = blk / kPerAccuBlk2; - const size_t blk_in_pair = blk % kPerAccuBlk2; - const size_t full_pairs = BlockCountK / kPerAccuBlk2; - if (pair_idx < full_pairs) { - return g * per_group_scales - + pair_idx * (kNCols4 * kPerAccuBlk2) - + c * kPerAccuBlk2 - + blk_in_pair; - } - return g * per_group_scales - + full_pairs * (kNCols4 * kPerAccuBlk2) - + c; - } - return n * BlockCountK + blk; -} +// block-group constants. 4 consecutive K-blocks share a single 64-byte buffer +// so the SIMD unpack is one ZMM load + four fixed shift/mask pairs. +constexpr size_t kBlockGroupBlks = 4; // 4 K-blocks per group +constexpr size_t kBlockGroupBytes = kBlockGroupBlks * kBlkBytes; // 64 bytes +constexpr size_t kBlockGroupWeights = kBlockGroupBlks * kBlkLen; // 256 weights // // Extract a single 2-bit weight from a standard ONNX MatMulNBits packed byte @@ -164,24 +98,26 @@ ExtractSrcWeight(const std::byte* src, size_t i) } // -// Pack one source block (16 bytes = 64 2-bit weights, standard ONNX layout) -// into the destination layout described at the top of this file. +// Pack 4 consecutive K-blocks (4 * 16 = 64 source bytes in standard ONNX +// layout) into a 64-byte block-group. Pure permutation of the 256 2-bit +// elements; bit-identical round-trip with UnPackBlockGroup_BlkLen64_Reference. // // src_byte[i] holds val[4i .. 4i+3] at bit positions {0..1, 2..3, 4..5, 6..7}. -// dst_byte[i] holds val[i], val[i+16], val[i+32], val[i+48] at the same bit -// positions. -// -// This is a pure permutation of the 64 2-bit elements; the bit width and -// count are preserved (16 bytes in, 16 bytes out). +// dst_byte[i] holds block0.weight[i], block1.weight[i], block2.weight[i], +// block3.weight[i] at the same bit positions. // inline void -PackBlock_BlkLen64(const std::byte* src, std::byte* dst) +PackBlockGroup_BlkLen64(const std::byte* src_block_0, + const std::byte* src_block_1, + const std::byte* src_block_2, + const std::byte* src_block_3, + std::byte* dst) { - for (size_t i = 0; i < kBlkBytes; ++i) { - const uint8_t v0 = ExtractSrcWeight(src, i + 0); - const uint8_t v1 = ExtractSrcWeight(src, i + 16); - const uint8_t v2 = ExtractSrcWeight(src, i + 32); - const uint8_t v3 = ExtractSrcWeight(src, i + 48); + for (size_t i = 0; i < kBlkLen; ++i) { + const uint8_t v0 = ExtractSrcWeight(src_block_0, i); + const uint8_t v1 = ExtractSrcWeight(src_block_1, i); + const uint8_t v2 = ExtractSrcWeight(src_block_2, i); + const uint8_t v3 = ExtractSrcWeight(src_block_3, i); dst[i] = static_cast( static_cast(v0 | (v1 << 2) | (v2 << 4) | (v3 << 6)) ); @@ -189,22 +125,24 @@ PackBlock_BlkLen64(const std::byte* src, std::byte* dst) } // -// Reference unpack of one packed block back into 64 int8 values in natural -// order ([val0, val1, ..., val63]). -// -// This routine is intentionally written without reference to PackBlock_BlkLen64 -// so that it can serve as an independent oracle for round-trip tests: -// it simply applies the documented dst_byte layout rule in reverse. +// Reference unpack of one 64-byte block-group back into 4 K-blocks worth of +// natural-order uint8 weights ([0, 3]). Written from the documented layout +// rule -- intentionally independent of PackBlockGroup_BlkLen64 so it can +// serve as a round-trip oracle. // inline void -UnpackBlock_BlkLen64_Reference(const std::byte* packed, uint8_t out[kBlkLen]) +UnPackBlockGroup_BlkLen64_Reference(const std::byte* packed, + uint8_t out_block_0[kBlkLen], + uint8_t out_block_1[kBlkLen], + uint8_t out_block_2[kBlkLen], + uint8_t out_block_3[kBlkLen]) { - for (size_t i = 0; i < kBlkBytes; ++i) { + for (size_t i = 0; i < kBlkLen; ++i) { const uint8_t b = static_cast(packed[i]); - out[i + 0] = static_cast((b >> 0) & 0x03u); - out[i + 16] = static_cast((b >> 2) & 0x03u); - out[i + 32] = static_cast((b >> 4) & 0x03u); - out[i + 48] = static_cast((b >> 6) & 0x03u); + out_block_0[i] = static_cast((b >> 0) & 0x03u); + out_block_1[i] = static_cast((b >> 2) & 0x03u); + out_block_2[i] = static_cast((b >> 4) & 0x03u); + out_block_3[i] = static_cast((b >> 6) & 0x03u); } } @@ -221,93 +159,89 @@ UnpackSourceBlock_BlkLen64_Reference(const std::byte* src, uint8_t out[kBlkLen]) } // ----------------------------------------------------------------------------- -// EXPERIMENTAL: super-block packing for fast unpack +// block-group packed-data layout. +// +// Main region (n < NMain = floor(N / kNCols4) * kNCols4): +// 4-N-col groups of g = n / 4, col within group c = n % 4. +// K-block-group index s = blk / 4 (s in [0, BlockCountK / 4)). +// Within a group, block-groups run consecutively across the 4 cols, +// so each (s, group) slot is a contiguous (kNCols4 * kBlockGroupBytes) +// = 256 byte chunk. +// +// Tail region (n >= NMain): plain column-major block-groups, identical +// shape to the main region but flat in N. +// +// K-tail handling (BlockCountK not a multiple of kBlockGroupBlks): +// The pack helpers round BlockCountK up to a multiple of kBlockGroupBlks +// (= 4) for storage purposes -- the padding 1-3 blocks at the trailing +// block-group contain zeroed weight bytes and zeroed scales, so they +// contribute 0 to the integer dot product and 0 to the BlkSum correction. +// This lets the SIMD kernel iterate the block-group K-loop without a +// special tail handler for B, and avoids dual packing layouts. Storage +// waste is at most (kBlockGroupBlks - 1) blocks per N-col, i.e. <= 48 +// bytes per col -- negligible at any realistic N. +// +// Conventions used by the offset helpers below: +// * `BlockGroupCountKPadded = ceil(BlockCountK / kBlockGroupBlks)` is +// the number of block-groups the kernel actually iterates. +// * `BlockCountKPadded = BlockGroupCountKPadded * kBlockGroupBlks` is +// the K-block count used to address the scale buffer. +// * Callers must pass `BlockGroupCountKPadded` and `BlockCountKPadded` +// to these helpers; the original logical BlockCountK is only used +// for sizing the BlkSum buffer (which is consumed by the SGEMM +// correction step, not the inner K-loop). +// +// Caller-side constraints: BlkLen == 64; BlockCountK >= 1 (any K, padded +// internally to a multiple of kBlockGroupBlks). // ----------------------------------------------------------------------------- -// -// The existing per-block pack (PackBlock_BlkLen64) stores 16 bytes per K-block -// such that each byte holds 4 weights from 4 different "rows" of the block. -// Unpacking that layout into a ZMM of 64 natural-order weights requires a -// broadcast + per-dword variable shift (`vpsrlvd`, ~3c latency, 1c throughput -// on Zen5) -- a cost we measured at ~30-35% of total W2 kernel time. -// -// The super-block layout below groups FOUR consecutive K-blocks together into -// a single 64-byte buffer (= 4 * 16 bytes, same total storage) such that -// byte[i] (i = 0..63) holds: -// bits[0..1] = block_0.weight[i] -// bits[2..3] = block_1.weight[i] -// bits[4..5] = block_2.weight[i] -// bits[6..7] = block_3.weight[i] -// -// Unpack with a single ZMM load and four fixed-shift+mask pairs: -// -// __m512i super = _mm512_loadu_si512(packed); -// __m512i mask = _mm512_set1_epi8(0x03); -// __m512i bv0 = _mm512_and_si512(super, mask); -// __m512i bv1 = _mm512_and_si512(_mm512_srli_epi16(super, 2), mask); -// __m512i bv2 = _mm512_and_si512(_mm512_srli_epi16(super, 4), mask); -// __m512i bv3 = _mm512_and_si512(_mm512_srli_epi16(super, 6), mask); -// -// Each shift+mask is ~2c (1c srli_epi16 + 1c andd) and the four chains are -// fully independent, so the critical path is ~4c for ALL four blocks combined -// vs the current ~20c (4 broadcasts * ~5c each, partially overlapped). -// -// Note on `_mm512_srli_epi16`: it shifts each 16-bit lane by N. For a byte -// that is the LOW byte of a 16-bit lane, the shift pulls in bits from the -// adjacent HIGH byte. The subsequent AND with 0x03 discards those leaked -// bits, leaving the correct per-byte result. For the HIGH byte, zeros are -// shifted in from the top, which is what we want. -constexpr size_t kSuperBlockBlks = 4; // 4 K-blocks per super -constexpr size_t kSuperBlockBytes = kSuperBlockBlks * kBlkBytes; // 64 bytes -constexpr size_t kSuperBlockWeights = kSuperBlockBlks * kBlkLen; // 256 weights - -// -// Pack 4 consecutive K-blocks (4 * 16 = 64 source bytes in standard ONNX -// layout) into a 64-byte super-block. Pure permutation of the 256 2-bit -// elements; bit-identical round-trip with UnpackSuperBlock4_BlkLen64_Reference. -// -inline void -PackSuperBlock4_BlkLen64(const std::byte* src_block_0, - const std::byte* src_block_1, - const std::byte* src_block_2, - const std::byte* src_block_3, - std::byte* dst) +inline size_t +PackedQuantBOffsetBytes_W2(size_t n, size_t blk_group, + size_t BlockGroupCountKPadded, size_t NMain) { - for (size_t i = 0; i < kBlkLen; ++i) { - const uint8_t v0 = ExtractSrcWeight(src_block_0, i); - const uint8_t v1 = ExtractSrcWeight(src_block_1, i); - const uint8_t v2 = ExtractSrcWeight(src_block_2, i); - const uint8_t v3 = ExtractSrcWeight(src_block_3, i); - dst[i] = static_cast( - static_cast(v0 | (v1 << 2) | (v2 << 4) | (v3 << 6)) - ); + if (n < NMain) { + const size_t g = n / kNCols4; + const size_t c = n % kNCols4; + const size_t per_group_bytes = BlockGroupCountKPadded * kNCols4 * kBlockGroupBytes; + return g * per_group_bytes + + blk_group * (kNCols4 * kBlockGroupBytes) + + c * kBlockGroupBytes; } + return (n * BlockGroupCountKPadded + blk_group) * kBlockGroupBytes; } // -// Reference unpack of one 64-byte super-block back into 4 K-blocks worth of -// natural-order uint8 weights ([0, 3]). Written from the documented layout -// rule -- intentionally independent of PackSuperBlock4_BlkLen64 so it can -// serve as a round-trip oracle. +// Float offset into the packed B-scale buffer for a logical (n, blk) cell. +// Scales remain per-block (one float per K-block), 4 per group. Caller +// passes BlockCountKPadded (= BlockGroupCountKPadded * kBlockGroupBlks); +// scale slots in [BlockCountK, BlockCountKPadded) contain zeros so the +// kernel can index uniformly. // -inline void -UnpackSuperBlock4_BlkLen64_Reference(const std::byte* packed, - uint8_t out_block_0[kBlkLen], - uint8_t out_block_1[kBlkLen], - uint8_t out_block_2[kBlkLen], - uint8_t out_block_3[kBlkLen]) +inline size_t +PackedQuantBScaleOffset_W2(size_t n, size_t blk, + size_t BlockCountKPadded, size_t NMain) { - for (size_t i = 0; i < kBlkLen; ++i) { - const uint8_t b = static_cast(packed[i]); - out_block_0[i] = static_cast((b >> 0) & 0x03u); - out_block_1[i] = static_cast((b >> 2) & 0x03u); - out_block_2[i] = static_cast((b >> 4) & 0x03u); - out_block_3[i] = static_cast((b >> 6) & 0x03u); + const size_t BlockGroupCountKPadded = BlockCountKPadded / kBlockGroupBlks; + const size_t blk_group = blk / kBlockGroupBlks; + const size_t blk_in_group = blk % kBlockGroupBlks; + if (n < NMain) { + const size_t g = n / kNCols4; + const size_t c = n % kNCols4; + const size_t per_group_scales = BlockGroupCountKPadded * kNCols4 * kBlockGroupBlks; + return g * per_group_scales + + blk_group * (kNCols4 * kBlockGroupBlks) + + c * kBlockGroupBlks + + blk_in_group; } + return n * BlockCountKPadded + blk; } // -// Reference / dispatch entry points defined in sqnbitgemm_kernel_avx512_2bit.cpp. +// Reference (scalar) entry points -- defined in sqnbitgemm_kernel_avx512_2bit.cpp. +// +// These cover Pack helpers and scalar oracle. They are +// reachable from unit tests via direct linkage; production dispatch wiring +// is performed by the platform dispatcher. // size_t MLASCALL @@ -355,16 +289,10 @@ SQ2BitGemmKernel_BlkSum_CompInt8_Scalar( ); // -// Unit-test forwarders for the two AVX-512 vectorized kernel variants. These -// are non-inline symbols defined in `sqnbitgemm_kernel_avx512vnni.cpp` (VNNI) -// and `sqnbitgemm_kernel_avx512.cpp` (non-VNNI), each of which is compiled -// with the appropriate ISA flags. The test TU (which is NOT compiled with -// AVX-512 flags) calls these by `extern` linkage to exercise both kernels -// independently of the platform dispatcher. -// -// Callers MUST gate on `GetMlasPlatform().Avx512Supported_` before invoking -// these symbols, since they execute AVX-512BW (and, for the VNNI variant, -// AVX-512-VNNI) instructions. +// Unit-test forwarders for the AVX-512 SIMD block-group kernels. Same gating +// rules as the existing W2 test entries: the caller MUST verify +// GetMlasPlatform().Avx512Supported_ (and, for the VNNI variant, that the +// active dispatch is the AVX-512-VNNI one) before invoking these symbols. // size_t MLASCALL SQ2BitGemmKernel_BlkSum_CompInt8_Avx512_TestEntry( diff --git a/onnxruntime/core/mlas/lib/sqnbitgemm_kernel_avx512vnni_2bit_blklen64_superblock.h b/onnxruntime/core/mlas/lib/sqnbitgemm_kernel_avx512_2bit_blklen64.h similarity index 73% rename from onnxruntime/core/mlas/lib/sqnbitgemm_kernel_avx512vnni_2bit_blklen64_superblock.h rename to onnxruntime/core/mlas/lib/sqnbitgemm_kernel_avx512_2bit_blklen64.h index 41ed3d8a3b6ef..53ac8b3d18f5b 100644 --- a/onnxruntime/core/mlas/lib/sqnbitgemm_kernel_avx512vnni_2bit_blklen64_superblock.h +++ b/onnxruntime/core/mlas/lib/sqnbitgemm_kernel_avx512_2bit_blklen64.h @@ -6,12 +6,12 @@ Licensed under the MIT License. Module Name: - sqnbitgemm_kernel_avx512vnni_2bit_blklen64_superblock.h + sqnbitgemm_kernel_avx512_2bit_blklen64.h Abstract: - EXPERIMENTAL AVX-512 (-VNNI) W2 kernel that consumes the super-block packed - layout (sqnbitgemm_kernel_avx512_2bit_superblock.h). Replaces the existing + AVX-512 (-VNNI) W2 kernel that consumes the block-group packed + layout (sqnbitgemm_kernel_avx512_2bit.h). Replaces the existing per-K-block broadcast + variable-shift unpack with one ZMM load and four fixed-shift+mask pairs, halving the inner-loop unpack cost. @@ -20,19 +20,19 @@ Module Name: `_mm512_dpbusd_epi32` for the integer MAC; non-VNNI uses the `vpmaddubsw + vpmaddwd` chain. Both produce bit-identical results. - Constraints (Phase 3): + Constraints (SIMD): * BlkLen == 64 only. - * BlockCountK must be a multiple of kSuperBlockBlks (= 4). + * BlockCountK must be a multiple of kBlockGroupBlks (= 4). * CountM must be a multiple of 2 (R2 tile); CountN must be a multiple of 4 (C4 tile). Tail handling is provided by the same R2xC1/R1xC4/ R1xC1 helpers as the existing W2 kernel via a downgrade path the dispatcher will pick when these constraints don't hold. Layout reference: - * 64-byte super-block: byte b holds 2-bit weight b from each of 4 + * 64-byte block-group: byte b holds 2-bit weight b from each of 4 consecutive K-blocks at bit positions {0..1, 2..3, 4..5, 6..7}. - * In a tile slot, 4 N-cols of super-block live consecutively: each - N-col's super starts at offset c * kSuperBlockBytes within the slot. + * In a tile slot, 4 N-cols of block-group live consecutively: each + N-col's group starts at offset c * kBlockGroupBytes within the slot. --*/ @@ -47,16 +47,15 @@ Module Name: #include "mlasi.h" #include "qnbitgemm.h" #include "sqnbitgemm_kernel_avx512_2bit.h" -#include "sqnbitgemm_kernel_avx512_2bit_superblock.h" namespace onnxruntime { namespace mlas { -namespace sq2bit_avx512_super { +namespace sq2bit_avx512 { inline constexpr size_t kNRows2 = 2; // matches the existing W2 R2 tile shape // -// Cheap super-block unpack: 1x ZMM load + 4x (fixed-shift + AND). +// Cheap block-group unpack: 1x ZMM load + 4x (fixed-shift + AND). // Critical path ~4c (load + and / load + srli + and parallel chains). // // Bit layout of each byte b of `packed`: @@ -71,18 +70,18 @@ inline constexpr size_t kNRows2 = 2; // matches the existing W2 R2 tile shape // zeros are shifted in from the top -- exactly what we want. // static MLAS_FORCEINLINE void -load_unpack_super_w2(const std::byte* packed, +load_unpack_w2_block_group(const std::byte* packed, __m512i& bv0_64_epi8, __m512i& bv1_64_epi8, __m512i& bv2_64_epi8, __m512i& bv3_64_epi8) { - const __m512i super = _mm512_loadu_si512(reinterpret_cast(packed)); + const __m512i block_group = _mm512_loadu_si512(reinterpret_cast(packed)); const __m512i mask03 = _mm512_set1_epi8(0x03); - bv0_64_epi8 = _mm512_and_si512(super, mask03); - bv1_64_epi8 = _mm512_and_si512(_mm512_srli_epi16(super, 2), mask03); - bv2_64_epi8 = _mm512_and_si512(_mm512_srli_epi16(super, 4), mask03); - bv3_64_epi8 = _mm512_and_si512(_mm512_srli_epi16(super, 6), mask03); + bv0_64_epi8 = _mm512_and_si512(block_group, mask03); + bv1_64_epi8 = _mm512_and_si512(_mm512_srli_epi16(block_group, 2), mask03); + bv2_64_epi8 = _mm512_and_si512(_mm512_srli_epi16(block_group, 4), mask03); + bv3_64_epi8 = _mm512_and_si512(_mm512_srli_epi16(block_group, 6), mask03); } // @@ -90,24 +89,24 @@ load_unpack_super_w2(const std::byte* packed, // uniformly-scaled FMA into a sub-accumulator; we keep two sub-accumulators // (alternating per K-block) so the per-cell FMA dependency chain is two FMAs // deep instead of four. The two sub-accumulators are summed into `acc` at the -// end with one extra vector add per super-block per cell. +// end with one extra vector add per block-group per cell. // // Math per K-block: acc += scale_a[blk] * scale_b[blk] * dot(av[blk], bv[blk]) // // scale_a and scale_b each point to 4 consecutive floats in their packed -// buffers (one per K-block of the super). +// buffers (one per K-block of the group). // // Critical path analysis (Zen5 / SKX): // * Single-chain (prior version): -// FMA latency * 4 = ~16 cycles per super-block per cell. +// FMA latency * 4 = ~16 cycles per block-group per cell. // * Two sub-accumulators (current): -// FMA latency * 2 = ~8 cycles per super-block per cell + one vaddps. +// FMA latency * 2 = ~8 cycles per block-group per cell + one vaddps. // This roughly halves the FP critical path; the integer dpbusd chain // (one per block) is independent and runs in parallel with the FMAs. // template static MLAS_FORCEINLINE void -dot_accumulate_4blk_w2_super(const __m512i& av0_64_epi8, const __m512i& av1_64_epi8, +dot_accumulate_4blk_w2(const __m512i& av0_64_epi8, const __m512i& av1_64_epi8, const __m512i& av2_64_epi8, const __m512i& av3_64_epi8, const __m512i& bv0_64_epi8, const __m512i& bv1_64_epi8, const __m512i& bv2_64_epi8, const __m512i& bv3_64_epi8, @@ -152,12 +151,12 @@ dot_accumulate_4blk_w2_super(const __m512i& av0_64_epi8, const __m512i& av1_64_e } // -// 2 M-rows x 1 N-col x 4 K-blocks (one super-block) accumulator. The -// super-block B load + unpack is shared across the 2 M-rows. +// 2 M-rows x 1 N-col x 4 K-blocks (one block-group) accumulator. The +// block-group B load + unpack is shared across the 2 M-rows. // template static MLAS_FORCEINLINE void -accumulate_w2_blklen64_r2c1blk4_super( +accumulate_w2_blklen64_r2c1blk4( const __m512i& av00, const __m512i& av01, const __m512i& av02, const __m512i& av03, const __m512i& av10, const __m512i& av11, const __m512i& av12, const __m512i& av13, const std::byte* QuantBDataPtr, @@ -168,22 +167,22 @@ accumulate_w2_blklen64_r2c1blk4_super( __m512& acc1) { __m512i bv0, bv1, bv2, bv3; - load_unpack_super_w2(QuantBDataPtr, bv0, bv1, bv2, bv3); + load_unpack_w2_block_group(QuantBDataPtr, bv0, bv1, bv2, bv3); - dot_accumulate_4blk_w2_super( + dot_accumulate_4blk_w2( av00, av01, av02, av03, bv0, bv1, bv2, bv3, scale_a0, scale_b, acc0); - dot_accumulate_4blk_w2_super( + dot_accumulate_4blk_w2( av10, av11, av12, av13, bv0, bv1, bv2, bv3, scale_a1, scale_b, acc1); } // -// 1 M-row x 1 N-col x 4 K-blocks (one super-block) accumulator. Used by the +// 1 M-row x 1 N-col x 4 K-blocks (one block-group) accumulator. Used by the // R1xC4 tile for M=1 decode and as the trailing odd-row handler of R2xC4 // when CountM is odd. // template static MLAS_FORCEINLINE void -accumulate_w2_blklen64_r1c1blk4_super( +accumulate_w2_blklen64_r1c1blk4( const __m512i& av00, const __m512i& av01, const __m512i& av02, const __m512i& av03, const std::byte* QuantBDataPtr, const float* scale_a0, @@ -191,22 +190,22 @@ accumulate_w2_blklen64_r1c1blk4_super( __m512& acc0) { __m512i bv0, bv1, bv2, bv3; - load_unpack_super_w2(QuantBDataPtr, bv0, bv1, bv2, bv3); + load_unpack_w2_block_group(QuantBDataPtr, bv0, bv1, bv2, bv3); - dot_accumulate_4blk_w2_super( + dot_accumulate_4blk_w2( av00, av01, av02, av03, bv0, bv1, bv2, bv3, scale_a0, scale_b, acc0); } // // R1 x C4 tile -- the M=1 decode path and the trailing odd-row handler for // the R2xC4 tile when CountM is odd. Identical N-tile structure as R2xC4 -// (4 N-cols, super-block K stride) but processes a single M-row, so it uses +// (4 N-cols, block-group K stride) but processes a single M-row, so it uses // half the registers (4 accumulators instead of 8) and half the MAC count -// per super-block iteration. +// per block-group iteration. // template MLAS_FORCEINLINE void -Q2Int8GemmR1xC4BlkLen64Avx512_Super( +Q2Int8GemmR1xC4BlkLen64Avx512( const std::byte* QuantA, const float* QuantAScale, const std::byte* QuantBData, @@ -219,28 +218,28 @@ Q2Int8GemmR1xC4BlkLen64Avx512_Super( size_t ldc) { const size_t lda = BlockCountK * kBlkLen; - constexpr size_t PerColSuperBytes = kSuperBlockBytes; - constexpr size_t PerColSuperScale = kSuperBlockBlks; - constexpr size_t PerKSuperAdvanceBytes = kNCols4 * PerColSuperBytes; - constexpr size_t PerKSuperAdvanceScale = kNCols4 * PerColSuperScale; + constexpr size_t PerColGroupBytes = kBlockGroupBytes; + constexpr size_t PerColGroupScale = kBlockGroupBlks; + constexpr size_t PerKGroupAdvanceBytes = kNCols4 * PerColGroupBytes; + constexpr size_t PerKGroupAdvanceScale = kNCols4 * PerColGroupScale; // GroupStride uses the PADDED BlockCountK because the packed B layout - // walks N-groups at intervals of `SuperBlockCountKPadded * kNCols4 * - // kSuperBlockBytes` (see PackedQuantBOffsetBytes_W2_SuperBlock). When - // BlockCountK is a multiple of kSuperBlockBlks (== 4) the padded and + // walks N-groups at intervals of `BlockGroupCountKPadded * kNCols4 * + // kBlockGroupBytes` (see PackedQuantBOffsetBytes_W2). When + // BlockCountK is a multiple of kBlockGroupBlks (== 4) the padded and // logical strides are identical; when it isn't, the kernel must step // past the padded slots to land on the next N-group correctly. - const size_t SuperBlockCountKPadded = - MlasDivRoundup(BlockCountK, kSuperBlockBlks); - const size_t BlockCountKPadded = SuperBlockCountKPadded * kSuperBlockBlks; + const size_t BlockGroupCountKPadded = + MlasDivRoundup(BlockCountK, kBlockGroupBlks); + const size_t BlockCountKPadded = BlockGroupCountKPadded * kBlockGroupBlks; const size_t GroupStrideBytes = BlockCountKPadded * kNCols4 * kBlkBytes; const size_t GroupStrideScale = BlockCountKPadded * kNCols4; assert(CountN % kNCols4 == 0); - // BlockCountK no longer required to be a multiple of kSuperBlockBlks: - // the main K-loop iterates full supers; an optional tail handler picks up + // BlockCountK no longer required to be a multiple of kBlockGroupBlks: + // the main K-loop iterates full groups; an optional tail handler picks up // the trailing 1-3 K-blocks (padded weights and scales contribute 0). - const size_t FullSupers = BlockCountK / kSuperBlockBlks; - const size_t TailBlocks = BlockCountK % kSuperBlockBlks; // 0, 1, 2, or 3 + const size_t FullGroups = BlockCountK / kBlockGroupBlks; + const size_t TailBlocks = BlockCountK % kBlockGroupBlks; // 0, 1, 2, or 3 for (size_t m = 0; m < CountM; ++m) { const std::byte* QuantBDataColPtr = QuantBData; @@ -260,7 +259,7 @@ Q2Int8GemmR1xC4BlkLen64Avx512_Super( _mm512_setzero_ps(), _mm512_setzero_ps() }; - for (size_t sb = 0; sb < FullSupers; ++sb) { + for (size_t sb = 0; sb < FullGroups; ++sb) { const __m512i av00 = _mm512_loadu_si512( reinterpret_cast(QuantAPtr)); const __m512i av01 = _mm512_loadu_si512( @@ -270,35 +269,35 @@ Q2Int8GemmR1xC4BlkLen64Avx512_Super( const __m512i av03 = _mm512_loadu_si512( reinterpret_cast(QuantAPtr + 3 * kBlkLen)); - accumulate_w2_blklen64_r1c1blk4_super( + accumulate_w2_blklen64_r1c1blk4( av00, av01, av02, av03, - QuantBDataPtr + 0 * PerColSuperBytes, + QuantBDataPtr + 0 * PerColGroupBytes, QuantAScalePtr, - QuantBScalePtr + 0 * PerColSuperScale, + QuantBScalePtr + 0 * PerColGroupScale, acc[0]); - accumulate_w2_blklen64_r1c1blk4_super( + accumulate_w2_blklen64_r1c1blk4( av00, av01, av02, av03, - QuantBDataPtr + 1 * PerColSuperBytes, + QuantBDataPtr + 1 * PerColGroupBytes, QuantAScalePtr, - QuantBScalePtr + 1 * PerColSuperScale, + QuantBScalePtr + 1 * PerColGroupScale, acc[1]); - accumulate_w2_blklen64_r1c1blk4_super( + accumulate_w2_blklen64_r1c1blk4( av00, av01, av02, av03, - QuantBDataPtr + 2 * PerColSuperBytes, + QuantBDataPtr + 2 * PerColGroupBytes, QuantAScalePtr, - QuantBScalePtr + 2 * PerColSuperScale, + QuantBScalePtr + 2 * PerColGroupScale, acc[2]); - accumulate_w2_blklen64_r1c1blk4_super( + accumulate_w2_blklen64_r1c1blk4( av00, av01, av02, av03, - QuantBDataPtr + 3 * PerColSuperBytes, + QuantBDataPtr + 3 * PerColGroupBytes, QuantAScalePtr, - QuantBScalePtr + 3 * PerColSuperScale, + QuantBScalePtr + 3 * PerColGroupScale, acc[3]); - QuantAPtr += kBlkLen * kSuperBlockBlks; - QuantAScalePtr += kSuperBlockBlks; - QuantBDataPtr += PerKSuperAdvanceBytes; - QuantBScalePtr += PerKSuperAdvanceScale; + QuantAPtr += kBlkLen * kBlockGroupBlks; + QuantAScalePtr += kBlockGroupBlks; + QuantBDataPtr += PerKGroupAdvanceBytes; + QuantBScalePtr += PerKGroupAdvanceScale; } // K-tail: 1-3 trailing real K-blocks. Pack helper zero-padded the @@ -325,34 +324,34 @@ Q2Int8GemmR1xC4BlkLen64Avx512_Super( const __m512i av03 = zero; // TailBlocks at most 3 // Bounded scale_a copy. - float scale_a0_safe[kSuperBlockBlks] = {0.0f, 0.0f, 0.0f, 0.0f}; + float scale_a0_safe[kBlockGroupBlks] = {0.0f, 0.0f, 0.0f, 0.0f}; for (size_t i = 0; i < TailBlocks; ++i) { scale_a0_safe[i] = QuantAScalePtr[i]; } - accumulate_w2_blklen64_r1c1blk4_super( + accumulate_w2_blklen64_r1c1blk4( av00, av01, av02, av03, - QuantBDataPtr + 0 * PerColSuperBytes, + QuantBDataPtr + 0 * PerColGroupBytes, scale_a0_safe, - QuantBScalePtr + 0 * PerColSuperScale, + QuantBScalePtr + 0 * PerColGroupScale, acc[0]); - accumulate_w2_blklen64_r1c1blk4_super( + accumulate_w2_blklen64_r1c1blk4( av00, av01, av02, av03, - QuantBDataPtr + 1 * PerColSuperBytes, + QuantBDataPtr + 1 * PerColGroupBytes, scale_a0_safe, - QuantBScalePtr + 1 * PerColSuperScale, + QuantBScalePtr + 1 * PerColGroupScale, acc[1]); - accumulate_w2_blklen64_r1c1blk4_super( + accumulate_w2_blklen64_r1c1blk4( av00, av01, av02, av03, - QuantBDataPtr + 2 * PerColSuperBytes, + QuantBDataPtr + 2 * PerColGroupBytes, scale_a0_safe, - QuantBScalePtr + 2 * PerColSuperScale, + QuantBScalePtr + 2 * PerColGroupScale, acc[2]); - accumulate_w2_blklen64_r1c1blk4_super( + accumulate_w2_blklen64_r1c1blk4( av00, av01, av02, av03, - QuantBDataPtr + 3 * PerColSuperBytes, + QuantBDataPtr + 3 * PerColGroupBytes, scale_a0_safe, - QuantBScalePtr + 3 * PerColSuperScale, + QuantBScalePtr + 3 * PerColGroupScale, acc[3]); } @@ -377,13 +376,13 @@ Q2Int8GemmR1xC4BlkLen64Avx512_Super( // // R2 x C4 tile -- the main hot path for prefill (M >= 2). Iterates the K -// dimension in super-block strides of kSuperBlockBlks (= 4) K-blocks at a -// time. Assumes BlockCountK is a multiple of kSuperBlockBlks; the dispatcher -// must verify this before selecting the super-block kernel. +// dimension in block-group strides of kBlockGroupBlks (= 4) K-blocks at a +// time. Assumes BlockCountK is a multiple of kBlockGroupBlks; the dispatcher +// must verify this before selecting the block-group kernel. // template MLAS_FORCEINLINE void -Q2Int8GemmR2xC4BlkLen64Avx512_Super( +Q2Int8GemmR2xC4BlkLen64Avx512( const std::byte* QuantA, const float* QuantAScale, const std::byte* QuantBData, @@ -396,29 +395,29 @@ Q2Int8GemmR2xC4BlkLen64Avx512_Super( size_t ldc) { const size_t lda = BlockCountK * kBlkLen; - constexpr size_t PerColSuperBytes = kSuperBlockBytes; // 64 B per col per super - constexpr size_t PerColSuperScale = kSuperBlockBlks; // 4 scales per col per super - constexpr size_t PerKSuperAdvanceBytes = kNCols4 * PerColSuperBytes; // 256 B per K-super iter - constexpr size_t PerKSuperAdvanceScale = kNCols4 * PerColSuperScale; // 16 scales per K-super iter + constexpr size_t PerColGroupBytes = kBlockGroupBytes; // 64 B per col per group + constexpr size_t PerColGroupScale = kBlockGroupBlks; // 4 scales per col per group + constexpr size_t PerKGroupAdvanceBytes = kNCols4 * PerColGroupBytes; // 256 B per K-group iter + constexpr size_t PerKGroupAdvanceScale = kNCols4 * PerColGroupScale; // 16 scales per K-group iter // GroupStride uses the PADDED BlockCountK because the packed B layout - // walks N-groups at intervals of `SuperBlockCountKPadded * kNCols4 * - // kSuperBlockBytes` (see PackedQuantBOffsetBytes_W2_SuperBlock). When - // BlockCountK is a multiple of kSuperBlockBlks (== 4) the padded and + // walks N-groups at intervals of `BlockGroupCountKPadded * kNCols4 * + // kBlockGroupBytes` (see PackedQuantBOffsetBytes_W2). When + // BlockCountK is a multiple of kBlockGroupBlks (== 4) the padded and // logical strides are identical; when it isn't, the kernel must step // past the padded slots to land on the next N-group correctly. - const size_t SuperBlockCountKPadded = - MlasDivRoundup(BlockCountK, kSuperBlockBlks); - const size_t BlockCountKPadded = SuperBlockCountKPadded * kSuperBlockBlks; + const size_t BlockGroupCountKPadded = + MlasDivRoundup(BlockCountK, kBlockGroupBlks); + const size_t BlockCountKPadded = BlockGroupCountKPadded * kBlockGroupBlks; const size_t GroupStrideBytes = BlockCountKPadded * kNCols4 * kBlkBytes; const size_t GroupStrideScale = BlockCountKPadded * kNCols4; assert(CountM % kNRows2 == 0); assert(CountN % kNCols4 == 0); - // BlockCountK no longer required to be a multiple of kSuperBlockBlks: - // the main K-loop iterates full supers; an optional tail handler picks up + // BlockCountK no longer required to be a multiple of kBlockGroupBlks: + // the main K-loop iterates full groups; an optional tail handler picks up // the trailing 1-3 K-blocks (padded weights and scales contribute 0). - const size_t FullSupers = BlockCountK / kSuperBlockBlks; - const size_t TailBlocks = BlockCountK % kSuperBlockBlks; // 0, 1, 2, or 3 + const size_t FullGroups = BlockCountK / kBlockGroupBlks; + const size_t TailBlocks = BlockCountK % kBlockGroupBlks; // 0, 1, 2, or 3 for (size_t m = 0; m < CountM; m += kNRows2) { const std::byte* QuantBDataColPtr = QuantBData; @@ -438,7 +437,7 @@ Q2Int8GemmR2xC4BlkLen64Avx512_Super( _mm512_setzero_ps(), _mm512_setzero_ps(), _mm512_setzero_ps(), _mm512_setzero_ps() }; - for (size_t sb = 0; sb < FullSupers; ++sb) { + for (size_t sb = 0; sb < FullGroups; ++sb) { const __m512i av00 = _mm512_loadu_si512( reinterpret_cast(QuantAPtr)); const __m512i av01 = _mm512_loadu_si512( @@ -456,35 +455,35 @@ Q2Int8GemmR2xC4BlkLen64Avx512_Super( const __m512i av13 = _mm512_loadu_si512( reinterpret_cast(QuantAPtr + lda + 3 * kBlkLen)); - accumulate_w2_blklen64_r2c1blk4_super( + accumulate_w2_blklen64_r2c1blk4( av00, av01, av02, av03, av10, av11, av12, av13, - QuantBDataPtr + 0 * PerColSuperBytes, + QuantBDataPtr + 0 * PerColGroupBytes, QuantAScalePtr, QuantAScalePtr + BlockCountK, - QuantBScalePtr + 0 * PerColSuperScale, + QuantBScalePtr + 0 * PerColGroupScale, acc[0], acc[kNCols4 + 0]); - accumulate_w2_blklen64_r2c1blk4_super( + accumulate_w2_blklen64_r2c1blk4( av00, av01, av02, av03, av10, av11, av12, av13, - QuantBDataPtr + 1 * PerColSuperBytes, + QuantBDataPtr + 1 * PerColGroupBytes, QuantAScalePtr, QuantAScalePtr + BlockCountK, - QuantBScalePtr + 1 * PerColSuperScale, + QuantBScalePtr + 1 * PerColGroupScale, acc[1], acc[kNCols4 + 1]); - accumulate_w2_blklen64_r2c1blk4_super( + accumulate_w2_blklen64_r2c1blk4( av00, av01, av02, av03, av10, av11, av12, av13, - QuantBDataPtr + 2 * PerColSuperBytes, + QuantBDataPtr + 2 * PerColGroupBytes, QuantAScalePtr, QuantAScalePtr + BlockCountK, - QuantBScalePtr + 2 * PerColSuperScale, + QuantBScalePtr + 2 * PerColGroupScale, acc[2], acc[kNCols4 + 2]); - accumulate_w2_blklen64_r2c1blk4_super( + accumulate_w2_blklen64_r2c1blk4( av00, av01, av02, av03, av10, av11, av12, av13, - QuantBDataPtr + 3 * PerColSuperBytes, + QuantBDataPtr + 3 * PerColGroupBytes, QuantAScalePtr, QuantAScalePtr + BlockCountK, - QuantBScalePtr + 3 * PerColSuperScale, + QuantBScalePtr + 3 * PerColGroupScale, acc[3], acc[kNCols4 + 3]); - QuantAPtr += kBlkLen * kSuperBlockBlks; - QuantAScalePtr += kSuperBlockBlks; - QuantBDataPtr += PerKSuperAdvanceBytes; - QuantBScalePtr += PerKSuperAdvanceScale; + QuantAPtr += kBlkLen * kBlockGroupBlks; + QuantAScalePtr += kBlockGroupBlks; + QuantBDataPtr += PerKGroupAdvanceBytes; + QuantBScalePtr += PerKGroupAdvanceScale; } // K-tail: 1-3 trailing real K-blocks. See R1 tile comment above. @@ -510,36 +509,36 @@ Q2Int8GemmR2xC4BlkLen64Avx512_Super( const __m512i av13 = zero; // Bounded scale_a copies for both M-rows (see R1 K-tail comment). - float scale_a0_safe[kSuperBlockBlks] = {0.0f, 0.0f, 0.0f, 0.0f}; - float scale_a1_safe[kSuperBlockBlks] = {0.0f, 0.0f, 0.0f, 0.0f}; + float scale_a0_safe[kBlockGroupBlks] = {0.0f, 0.0f, 0.0f, 0.0f}; + float scale_a1_safe[kBlockGroupBlks] = {0.0f, 0.0f, 0.0f, 0.0f}; for (size_t i = 0; i < TailBlocks; ++i) { scale_a0_safe[i] = QuantAScalePtr[i]; scale_a1_safe[i] = QuantAScalePtr[BlockCountK + i]; } - accumulate_w2_blklen64_r2c1blk4_super( + accumulate_w2_blklen64_r2c1blk4( av00, av01, av02, av03, av10, av11, av12, av13, - QuantBDataPtr + 0 * PerColSuperBytes, + QuantBDataPtr + 0 * PerColGroupBytes, scale_a0_safe, scale_a1_safe, - QuantBScalePtr + 0 * PerColSuperScale, + QuantBScalePtr + 0 * PerColGroupScale, acc[0], acc[kNCols4 + 0]); - accumulate_w2_blklen64_r2c1blk4_super( + accumulate_w2_blklen64_r2c1blk4( av00, av01, av02, av03, av10, av11, av12, av13, - QuantBDataPtr + 1 * PerColSuperBytes, + QuantBDataPtr + 1 * PerColGroupBytes, scale_a0_safe, scale_a1_safe, - QuantBScalePtr + 1 * PerColSuperScale, + QuantBScalePtr + 1 * PerColGroupScale, acc[1], acc[kNCols4 + 1]); - accumulate_w2_blklen64_r2c1blk4_super( + accumulate_w2_blklen64_r2c1blk4( av00, av01, av02, av03, av10, av11, av12, av13, - QuantBDataPtr + 2 * PerColSuperBytes, + QuantBDataPtr + 2 * PerColGroupBytes, scale_a0_safe, scale_a1_safe, - QuantBScalePtr + 2 * PerColSuperScale, + QuantBScalePtr + 2 * PerColGroupScale, acc[2], acc[kNCols4 + 2]); - accumulate_w2_blklen64_r2c1blk4_super( + accumulate_w2_blklen64_r2c1blk4( av00, av01, av02, av03, av10, av11, av12, av13, - QuantBDataPtr + 3 * PerColSuperBytes, + QuantBDataPtr + 3 * PerColGroupBytes, scale_a0_safe, scale_a1_safe, - QuantBScalePtr + 3 * PerColSuperScale, + QuantBScalePtr + 3 * PerColGroupScale, acc[3], acc[kNCols4 + 3]); } @@ -573,19 +572,19 @@ Q2Int8GemmR2xC4BlkLen64Avx512_Super( // // 1 M-row x 1 N-col N-tail tile. Handles the 1-3 trailing N-cols when // CountN is not a multiple of kNCols4. The tail region of the packed B -// buffer is column-major (one super-block per K-super per N-col, see -// PackedQuantBOffsetBytes_W2_SuperBlock for n >= NMain), so this tile +// buffer is column-major (one block-group per K-group per N-col, see +// PackedQuantBOffsetBytes_W2 for n >= NMain), so this tile // walks one column at a time and reuses the same accumulator helper used -// by R1xC4 (`accumulate_w2_blklen64_r1c1blk4_super`). Slower than the +// by R1xC4 (`accumulate_w2_blklen64_r1c1blk4`). Slower than the // R2xC4 main tile, but it processes at most 3 N-cols per call -- a // trivial fraction of total work even on the worst-case shape. // // Pointer convention (caller-supplied bases): // QuantBDataTail : start of the tail region in packed B -- exactly -// NMain * SuperBlockCountKPadded * kSuperBlockBytes +// NMain * BlockGroupCountKPadded * kBlockGroupBytes // bytes past PackedQuantBData (see callsite below). // QuantBScaleTail : same convention for the scale buffer (NMain * -// SuperBlockCountKPadded * kSuperBlockBlks floats). +// BlockGroupCountKPadded * kBlockGroupBlks floats). // // K-tail handling: identical to R1xC4 -- conditional A loads for the 1-3 // trailing real K-blocks and a bounded scale_a copy to avoid NaN from @@ -593,7 +592,7 @@ Q2Int8GemmR2xC4BlkLen64Avx512_Super( // template MLAS_FORCEINLINE void -Q2Int8GemmRMxC_Tail_BlkLen64Avx512_Super( +Q2Int8GemmRMxC_Tail_BlkLen64Avx512( const std::byte* QuantA, const float* QuantAScale, const std::byte* QuantBDataTail, @@ -606,18 +605,18 @@ Q2Int8GemmRMxC_Tail_BlkLen64Avx512_Super( size_t ldc) { assert(TailN >= 1 && TailN <= 3); - constexpr size_t PerColSuperBytes = kSuperBlockBytes; // 64 B per col per super - constexpr size_t PerColSuperScale = kSuperBlockBlks; // 4 scales per col per super + constexpr size_t PerColGroupBytes = kBlockGroupBytes; // 64 B per col per group + constexpr size_t PerColGroupScale = kBlockGroupBlks; // 4 scales per col per group const size_t lda = BlockCountK * kBlkLen; - const size_t SuperBlockCountKPadded = - MlasDivRoundup(BlockCountK, kSuperBlockBlks); - const size_t FullSupers = BlockCountK / kSuperBlockBlks; - const size_t TailBlocks = BlockCountK % kSuperBlockBlks; // 0, 1, 2, or 3 - // In the tail region each N-col occupies SuperBlockCountKPadded - // super-blocks back-to-back (column-major). - const size_t ColStrideBytes = SuperBlockCountKPadded * PerColSuperBytes; - const size_t ColStrideScale = SuperBlockCountKPadded * PerColSuperScale; + const size_t BlockGroupCountKPadded = + MlasDivRoundup(BlockCountK, kBlockGroupBlks); + const size_t FullGroups = BlockCountK / kBlockGroupBlks; + const size_t TailBlocks = BlockCountK % kBlockGroupBlks; // 0, 1, 2, or 3 + // In the tail region each N-col occupies BlockGroupCountKPadded + // block-groups back-to-back (column-major). + const size_t ColStrideBytes = BlockGroupCountKPadded * PerColGroupBytes; + const size_t ColStrideScale = BlockGroupCountKPadded * PerColGroupScale; for (size_t m = 0; m < CountM; ++m) { for (size_t c = 0; c < TailN; ++c) { @@ -629,7 +628,7 @@ Q2Int8GemmRMxC_Tail_BlkLen64Avx512_Super( __m512 acc = _mm512_setzero_ps(); - for (size_t sb = 0; sb < FullSupers; ++sb) { + for (size_t sb = 0; sb < FullGroups; ++sb) { const __m512i av00 = _mm512_loadu_si512( reinterpret_cast(QuantAPtr)); const __m512i av01 = _mm512_loadu_si512( @@ -639,17 +638,17 @@ Q2Int8GemmRMxC_Tail_BlkLen64Avx512_Super( const __m512i av03 = _mm512_loadu_si512( reinterpret_cast(QuantAPtr + 3 * kBlkLen)); - accumulate_w2_blklen64_r1c1blk4_super( + accumulate_w2_blklen64_r1c1blk4( av00, av01, av02, av03, QuantBDataPtr, QuantAScalePtr, QuantBScalePtr, acc); - QuantAPtr += kBlkLen * kSuperBlockBlks; - QuantAScalePtr += kSuperBlockBlks; - QuantBDataPtr += PerColSuperBytes; - QuantBScalePtr += PerColSuperScale; + QuantAPtr += kBlkLen * kBlockGroupBlks; + QuantAScalePtr += kBlockGroupBlks; + QuantBDataPtr += PerColGroupBytes; + QuantBScalePtr += PerColGroupScale; } if (TailBlocks > 0) { @@ -664,12 +663,12 @@ Q2Int8GemmRMxC_Tail_BlkLen64Avx512_Super( : zero; const __m512i av03 = zero; - float scale_a0_safe[kSuperBlockBlks] = {0.0f, 0.0f, 0.0f, 0.0f}; + float scale_a0_safe[kBlockGroupBlks] = {0.0f, 0.0f, 0.0f, 0.0f}; for (size_t i = 0; i < TailBlocks; ++i) { scale_a0_safe[i] = QuantAScalePtr[i]; } - accumulate_w2_blklen64_r1c1blk4_super( + accumulate_w2_blklen64_r1c1blk4( av00, av01, av02, av03, QuantBDataPtr, scale_a0_safe, @@ -702,14 +701,14 @@ Q2Int8GemmRMxC_Tail_BlkLen64Avx512_Super( // layout; a per-1-col tail tile picks up the trailing 1-3 cols against // the column-major tail region of the same packed buffer. // * BlockCountK has no alignment requirement: the R2/R1 tiles and the -// N-tail tile each run a partial-super K-tail handler that loads only +// N-tail tile each run a partial-group K-tail handler that loads only // the real trailing K-blocks (zero ZMM for missing slots, bounded // scale_a copy) and lets the pre-zeroed packed-B / scale slots // contribute 0 to the dot product. // template static MLAS_FORCEINLINE size_t -SQ2BitGemmKernel_BlkSum_CompInt8_Super_Impl( +SQ2BitGemmKernel_BlkSum_CompInt8_Impl( const size_t BlkLen, const std::byte* QuantA, const float* QuantAScale, @@ -746,7 +745,7 @@ SQ2BitGemmKernel_BlkSum_CompInt8_Super_Impl( if (NMain > 0) { if (M_main > 0) { - Q2Int8GemmR2xC4BlkLen64Avx512_Super( + Q2Int8GemmR2xC4BlkLen64Avx512( QuantA, QuantAScale, QuantBData, QuantBScale, C, M_main, NMain, BlockCountK, Bias, ldc); } @@ -755,7 +754,7 @@ SQ2BitGemmKernel_BlkSum_CompInt8_Super_Impl( // rows the R2 tile already consumed: A advances M_main*lda bytes, A-scale // advances M_main*BlockCountK floats, C advances M_main*ldc floats. The // packed-B buffer is reused (column-major over N). - Q2Int8GemmR1xC4BlkLen64Avx512_Super( + Q2Int8GemmR1xC4BlkLen64Avx512( QuantA + M_main * lda, QuantAScale + M_main * BlockCountK, QuantBData, QuantBScale, @@ -768,17 +767,17 @@ SQ2BitGemmKernel_BlkSum_CompInt8_Super_Impl( if (NTail > 0) { // The tail region of the packed B buffer is column-major and starts // immediately after the NMain-cols grouped region: - // tail_base_bytes = NMain * SuperBlockCountKPadded * kSuperBlockBytes - // tail_base_scales = NMain * SuperBlockCountKPadded * kSuperBlockBlks - const size_t SuperBlockCountKPadded = - MlasDivRoundup(BlockCountK, kSuperBlockBlks); + // tail_base_bytes = NMain * BlockGroupCountKPadded * kBlockGroupBytes + // tail_base_scales = NMain * BlockGroupCountKPadded * kBlockGroupBlks + const size_t BlockGroupCountKPadded = + MlasDivRoundup(BlockCountK, kBlockGroupBlks); const std::byte* QuantBDataTail = - QuantBData + NMain * SuperBlockCountKPadded * kSuperBlockBytes; + QuantBData + NMain * BlockGroupCountKPadded * kBlockGroupBytes; const float* QuantBScaleTail = - QuantBScale + NMain * SuperBlockCountKPadded * kSuperBlockBlks; + QuantBScale + NMain * BlockGroupCountKPadded * kBlockGroupBlks; const float* BiasTail = (Bias != nullptr) ? Bias + NMain : nullptr; - Q2Int8GemmRMxC_Tail_BlkLen64Avx512_Super( + Q2Int8GemmRMxC_Tail_BlkLen64Avx512( QuantA, QuantAScale, QuantBDataTail, QuantBScaleTail, C + NMain, @@ -810,7 +809,7 @@ SQ2BitGemmKernel_BlkSum_CompInt8_Super_Impl( // Top-level VNNI variant. Compiled into AVX-512-VNNI sources only. // static MLAS_FORCEINLINE size_t MLASCALL -SQ2BitGemmKernel_BlkSum_CompInt8_Super_Avx512Vnni( +SQ2BitGemmKernel_BlkSum_CompInt8_Avx512Vnni( const size_t BlkLen, const std::byte* QuantA, const float* QuantAScale, @@ -827,7 +826,7 @@ SQ2BitGemmKernel_BlkSum_CompInt8_Super_Avx512Vnni( const float* ABlockSum, const float* QuantBBlkSum) { - return SQ2BitGemmKernel_BlkSum_CompInt8_Super_Impl( + return SQ2BitGemmKernel_BlkSum_CompInt8_Impl( BlkLen, QuantA, QuantAScale, QuantBData, QuantBScale, QuantBZeroPoint, C, CountM, CountN, CountK, BlockCountK, Bias, ldc, ABlockSum, QuantBBlkSum); } @@ -837,7 +836,7 @@ SQ2BitGemmKernel_BlkSum_CompInt8_Super_Avx512Vnni( // vpmaddubsw + vpmaddwd chain instead of dpbusd. // static MLAS_FORCEINLINE size_t MLASCALL -SQ2BitGemmKernel_BlkSum_CompInt8_Super_Avx512( +SQ2BitGemmKernel_BlkSum_CompInt8_Avx512( const size_t BlkLen, const std::byte* QuantA, const float* QuantAScale, @@ -854,11 +853,11 @@ SQ2BitGemmKernel_BlkSum_CompInt8_Super_Avx512( const float* ABlockSum, const float* QuantBBlkSum) { - return SQ2BitGemmKernel_BlkSum_CompInt8_Super_Impl( + return SQ2BitGemmKernel_BlkSum_CompInt8_Impl( BlkLen, QuantA, QuantAScale, QuantBData, QuantBScale, QuantBZeroPoint, C, CountM, CountN, CountK, BlockCountK, Bias, ldc, ABlockSum, QuantBBlkSum); } -} // namespace sq2bit_avx512_super +} // namespace sq2bit_avx512 } // namespace mlas } // namespace onnxruntime diff --git a/onnxruntime/core/mlas/lib/sqnbitgemm_kernel_avx512_2bit_superblock.cpp b/onnxruntime/core/mlas/lib/sqnbitgemm_kernel_avx512_2bit_superblock.cpp deleted file mode 100644 index 01be1d153feae..0000000000000 --- a/onnxruntime/core/mlas/lib/sqnbitgemm_kernel_avx512_2bit_superblock.cpp +++ /dev/null @@ -1,359 +0,0 @@ -/*++ - -Copyright (c) Microsoft Corporation. All rights reserved. - -Licensed under the MIT License. - -Module Name: - - sqnbitgemm_kernel_avx512_2bit_superblock.cpp - -Abstract: - - Pack helpers and a scalar reference kernel for the super-block W2 layout. - - See sqnbitgemm_kernel_avx512_2bit_superblock.h for the layout description - and the rationale (closing the W2-vs-W4 prefill gap by replacing the - per-block broadcast + variable shift unpack with a single 64-byte load + - four fixed-shift+mask pairs). - - This translation unit is scalar / portable. The vectorized inner loop - that consumes the super-block layout lives in a separate header to be - added in phase 3, and is wired into the dispatch tables in phase 4. - ---*/ - -#include "sqnbitgemm_kernel_avx512_2bit_superblock.h" - -#include -#include -#include -#include - -#include "mlasi.h" -#include "qnbitgemm.h" - -namespace onnxruntime { -namespace mlas { -namespace sq2bit_avx512_super { - -namespace sq2 = ::onnxruntime::mlas::sq2bit_avx512; - -// -// Workspace / pack-buffer size for the super-block W2 path. Returns 0 if any -// of the configuration constraints is violated; the caller (MlasQNBitGemmPackQuantBDataSize) -// treats that as "unsupported" and falls back to the original W2 path. -// -// Constraints: -// * BlkLen == 64 -// * ComputeType == SQNBIT_CompInt8 -// -// K-tail handling: BlockCountK is rounded UP to a multiple of kSuperBlockBlks -// for the storage that the inner K-loop walks (PackedQuantBData, -// PackedQuantBScale). Padding slots hold zeroed weights and scales, so they -// contribute exactly 0 to the dot product. The BlkSum buffer is kept at the -// LOGICAL BlockCountK because it is consumed by the SGEMM correction step, -// not by the inner K-loop. -// -// Storage matches the original W2 layout total bytes when BlockCountK is a -// multiple of 4. When not a multiple of 4, storage grows by at most 3 K-blocks -// per N-col (= up to 48 bytes per col -- negligible at production N values). -// -// [PackedQuantBData] N * BlockCountKPadded * kBlkBytes -// [PackedQuantBScale] N * BlockCountKPadded * sizeof(float) -// [QuantBBlkSum] roundup_16(N) * BlockCountK (logical) * 16 floats -// -size_t MLASCALL -Q2BitGemmPackQuantBDataSize_SuperBlock( - size_t N, - size_t K, - size_t BlkLen, - bool /* HasZeroPoint */, - MLAS_QNBIT_GEMM_COMPUTE_TYPE ComputeType, - const MLAS_BACKEND_KERNEL_SELECTOR_CONFIG* /* BackendKernelSelectorConfig */ -) -{ - if (BlkLen != sq2::kBlkLen || ComputeType != SQNBIT_CompInt8) { - return 0; - } - const size_t BlockCountK = MlasDivRoundup(K, BlkLen); - if (BlockCountK == 0) { - return 0; - } - const size_t BlockCountKPadded = - MlasDivRoundup(BlockCountK, kSuperBlockBlks) * kSuperBlockBlks; - - // Use BlockCountKPadded for BlkSum sizing too. The actual SGEMM-correction - // step only reads LOGICAL BlockCountK entries, but PackedQuantBDataStruct - // is constructed by the caller with a single BlockCountK value that - // controls BOTH the packed-B size and the BlkSum offset. If we sized - // BlkSum at the logical BlockCountK while sizing packed-B at padded, - // the struct's BlkSum pointer would land inside the packed-B region - // (because the caller's struct uses one BlockCountK consistently). The - // extra storage from padding the BlkSum is ~16 floats per N -- trivial. - size_t PackedQuantBDataSize = N * BlockCountKPadded * kBlkBytes; - const size_t ScaleSize = N * BlockCountKPadded * sizeof(float); - size_t BlkSumSize = MlasDivRoundup(N, 16) * BlockCountKPadded * 16 * sizeof(float); - - constexpr size_t kPackedQuantBDataAlignment = 64; - PackedQuantBDataSize += kPackedQuantBDataAlignment - 1; - - constexpr size_t kBlkSumAlignment = MlasQNBitQuantBBlkSumAlignment(); - BlkSumSize += kBlkSumAlignment - 1; - - return PackedQuantBDataSize + ScaleSize + BlkSumSize; -} - -// -// Pack quantized B data + scales + per-block sums for the super-block W2 path. -// -// PackedQuantBData layout (super-blocks of 4 K-blocks, 64 bytes each): -// The super-block at logical (n, blk_super=blk/4) lives at byte offset -// PackedQuantBOffsetBytes_W2_SuperBlock(n, blk_super, SuperBlockCountK, NMain). -// Byte b within the super-block holds 2-bit weight b from each of the 4 -// constituent K-blocks at bit positions {0..1, 2..3, 4..5, 6..7}. -// -// PackedQuantBScale layout: one float per K-block, four floats per super-block, -// addressed by PackedQuantBScaleOffset_W2_SuperBlock. -// -// QuantBBlkSum layout: the same width-16 row-major chunked layout used by the -// existing W2 path, so the SGEMM correction step (MlasGemmFloatKernel) can be -// shared verbatim with the existing kernel. -// -// Mirrors the SQ2BitGemmPackQuantBDataAndBlkSum_Scalar prepack 3-call pattern: -// ORT's matmul_nbits.cc invokes this function up to three times (B, scales, ZP). -// We write scales when scales arrive, then re-derive BlkSum whenever either -// scales or zero-points arrive, reading scales from the already-packed buffer. -// -void MLASCALL -SQ2BitGemmPackQuantBDataAndBlkSum_SuperBlockScalar( - size_t N, - size_t K, - size_t BlkLen, - MLAS_QNBIT_GEMM_COMPUTE_TYPE /* ComputeType */, - const std::byte* QuantBDataBegin, - const float* QuantBScaleBegin, - bool /* HasZeroPoint */, - const std::byte* QuantBZPBegin, - PackedQuantBDataStruct& PackedQuantB, - MLAS_THREADPOOL* ThreadPool, - const MLAS_BACKEND_KERNEL_SELECTOR_CONFIG* /* BackendKernelSelectorConfig */ -) -{ - assert(BlkLen == sq2::kBlkLen); - if (BlkLen != sq2::kBlkLen) { - return; - } - - const size_t BlockCountK = MlasDivRoundup(K, BlkLen); - if (BlockCountK == 0) { - return; - } - - // Pad BlockCountK up to a multiple of kSuperBlockBlks so the inner K-loop - // can iterate whole super-blocks uniformly. Padding slots store zeroed - // weights / scales and contribute exactly 0 to the dot product. - const size_t SuperBlockCountKPadded = - MlasDivRoundup(BlockCountK, kSuperBlockBlks); - const size_t BlockCountKPadded = SuperBlockCountKPadded * kSuperBlockBlks; - const size_t NMain = (N / kNCols4) * kNCols4; - - // Zero source block used when packing a super whose K-range crosses the - // logical BlockCountK boundary. We point to this static zero buffer in - // the missing slots so the existing 4-block pack helper does the right - // thing without any branching inside it. - static const std::byte kZeroBlock[kBlkBytes] = {}; - - // ----- B-data pack ----- - if (QuantBDataBegin != nullptr) { - std::byte* PackedQuantBData = PackedQuantB.PackedQuantBData; - const size_t Iterations = N * SuperBlockCountKPadded; - MlasTrySimpleParallel( - ThreadPool, static_cast(Iterations), - [&](ptrdiff_t tid) { - const size_t n = static_cast(tid) / SuperBlockCountKPadded; - const size_t blk_super = static_cast(tid) % SuperBlockCountKPadded; - const size_t blk0 = blk_super * kSuperBlockBlks; - - // Pick real source block pointers for slots that exist; the - // static zero buffer for slots past the logical BlockCountK. - auto src_for = [&](size_t blk) -> const std::byte* { - if (blk < BlockCountK) { - return QuantBDataBegin + (n * BlockCountK + blk) * kBlkBytes; - } - return kZeroBlock; - }; - const std::byte* src_blk_0 = src_for(blk0 + 0); - const std::byte* src_blk_1 = src_for(blk0 + 1); - const std::byte* src_blk_2 = src_for(blk0 + 2); - const std::byte* src_blk_3 = src_for(blk0 + 3); - - const size_t dst_offset = - PackedQuantBOffsetBytes_W2_SuperBlock(n, blk_super, SuperBlockCountKPadded, NMain); - sq2::PackSuperBlock4_BlkLen64(src_blk_0, src_blk_1, src_blk_2, src_blk_3, - PackedQuantBData + dst_offset); - } - ); - } - - // ----- Scales ----- - // Iterate over the PADDED block count so trailing padding slots get - // explicit zero scales (otherwise they could hold uninitialised noise - // and the kernel's K-loop would read those into the FMA). - if (QuantBScaleBegin != nullptr) { - float* PackedScales = PackedQuantB.PackedQuantBScale; - const size_t Iterations = N * BlockCountKPadded; - MlasTrySimpleParallel( - ThreadPool, static_cast(Iterations), - [&](ptrdiff_t tid) { - const size_t n = static_cast(tid) / BlockCountKPadded; - const size_t blk = static_cast(tid) % BlockCountKPadded; - const float scale = (blk < BlockCountK) - ? QuantBScaleBegin[n * BlockCountK + blk] - : 0.0f; - PackedScales[PackedQuantBScaleOffset_W2_SuperBlock(n, blk, BlockCountKPadded, NMain)] = scale; - } - ); - } - - // ----- BlkSum (recomputed whenever scales or ZPs arrive) ----- - // BlkSum is consumed by the SGEMM correction step (MlasGemmFloatKernel), - // which is called outside the inner K-loop with the LOGICAL BlockCountK - // and the per-row ABlockSum the dispatcher produced for that logical K. - // We therefore only need to fill the first BlockCountK entries; the buffer - // is sized at MlasDivRoundup(N, 16) * BlockCountK * 16 floats (logical). - if (QuantBScaleBegin != nullptr || QuantBZPBegin != nullptr) { - const float* PackedScales = PackedQuantB.PackedQuantBScale; - float* BlkSum = PackedQuantB.QuantBBlkSum; - const size_t ZPCountK = MlasDivRoundup(BlockCountK, 4); - const size_t Iterations = N * BlockCountK; - MlasTrySimpleParallel( - ThreadPool, static_cast(Iterations), - [&](ptrdiff_t tid) { - const size_t n = static_cast(tid) / BlockCountK; - const size_t blk = static_cast(tid) % BlockCountK; - const float scale = - PackedScales[PackedQuantBScaleOffset_W2_SuperBlock(n, blk, BlockCountKPadded, NMain)]; - - uint8_t zp = kDefaultSymmetricZeroPoint2Bit; - if (QuantBZPBegin != nullptr) { - const size_t zp_byte_idx = n * ZPCountK + (blk / 4); - const size_t zp_bit_off = (blk % 4) * 2; - zp = static_cast( - (static_cast(QuantBZPBegin[zp_byte_idx]) >> zp_bit_off) & 0x03u); - } - - const size_t blksum_offset = ((n / 16) * BlockCountK + blk) * 16 + (n % 16); - BlkSum[blksum_offset] = -scale * static_cast(zp); - } - ); - } -} - -// -// Scalar reference kernel that consumes the super-block packed layout. -// Same math as the existing reference kernel; differs only in how it walks -// PackedQuantBData (super-block-major) and PackedQuantBScale (super-block-major). -// -// This is the correctness oracle for the SIMD super-block kernel coming in -// Phase 3. It also lets us validate the pack layout end-to-end via the -// existing MlasQNBitGemmBatch dispatch path once we wire it up. -// -size_t MLASCALL -SQ2BitGemmKernel_BlkSum_CompInt8_SuperBlockScalar( - const size_t BlkLen, - const std::byte* QuantA, - const float* QuantAScale, - const std::byte* QuantBData, - const float* QuantBScale, - const std::byte* /* QuantBZeroPoint */, - float* C, - size_t CountM, - size_t CountN, - size_t /* CountK */, - size_t BlockCountK, - const float* Bias, - size_t ldc, - const float* ABlockSum, - const float* QuantBBlkSum -) -{ - if (BlkLen != sq2::kBlkLen) { - return 0; - } - if (BlockCountK == 0) { - return 0; - } - - // PackedQuantBData and PackedQuantBScale are addressed via padded counts - // (K-tail handling -- see Q2BitGemmPackQuantBDataSize_SuperBlock). The K - // dot-product loop itself iterates only LOGICAL BlockCountK steps because - // A is unpadded; the kernel never reads past the real A rows. - const size_t SuperBlockCountKPadded = - MlasDivRoundup(BlockCountK, kSuperBlockBlks); - const size_t BlockCountKPadded = SuperBlockCountKPadded * kSuperBlockBlks; - - // The kernel is called by SQ2BitGemm_CompInt8 with the full CountN range; that - // function selects an N-tile boundary (kNCols4) up the stack. For a scalar - // reference path we don't depend on the 4-N-col grouping, but we DO need to - // index PackedQuantBData/PackedQuantBScale via the super-block offset helpers - // so we read the right bytes regardless of caller tile choice. - // - // CountN may not be a multiple of kNCols4 in the tail case. Detect that and - // fall back to plain column-major for the tail cols (the layout helpers - // already encode this rule). - const size_t NMainLocal = (CountN / kNCols4) * kNCols4; - - const size_t lda = BlockCountK * sq2::kBlkLen; // bytes per A row (int8) - const size_t lda_scale = BlockCountK; // floats per A scale row - - for (size_t m = 0; m < CountM; ++m) { - const int8_t* a_row = reinterpret_cast(QuantA + m * lda); - const float* a_scale_row = QuantAScale + m * lda_scale; - const float* a_blksum_row = ABlockSum + m * lda_scale; - float* c_row = C + m * ldc; - - for (size_t n = 0; n < CountN; ++n) { - float acc = (Bias != nullptr) ? Bias[n] : 0.0f; - - for (size_t blk = 0; blk < BlockCountK; ++blk) { - // Pull the super-block this K-block belongs to and unpack only the - // slot we need (block_in_super = blk % 4 selects the 2-bit field). - const size_t blk_super = blk / kSuperBlockBlks; - const size_t blk_in_super = blk % kSuperBlockBlks; - const size_t super_offset = - PackedQuantBOffsetBytes_W2_SuperBlock(n, blk_super, SuperBlockCountKPadded, NMainLocal); - const std::byte* super = QuantBData + super_offset; - - uint8_t b_unpacked[sq2::kBlkLen]; - for (size_t i = 0; i < sq2::kBlkLen; ++i) { - const uint8_t byte = static_cast(super[i]); - b_unpacked[i] = static_cast((byte >> (2 * blk_in_super)) & 0x03u); - } - - const int8_t* a_blk = a_row + blk * sq2::kBlkLen; - int32_t dot = 0; - for (size_t i = 0; i < sq2::kBlkLen; ++i) { - dot += static_cast(a_blk[i]) * static_cast(b_unpacked[i]); - } - - const float b_scale = - QuantBScale[PackedQuantBScaleOffset_W2_SuperBlock(n, blk, BlockCountKPadded, NMainLocal)]; - acc += a_scale_row[blk] * b_scale * static_cast(dot); - - // The width-16 row-major BlkSum layout is column-major in n - // (one float per (n, blk)); same as the existing W2 path. - const size_t blksum_offset = ((n / 16) * BlockCountK + blk) * 16 + (n % 16); - acc += a_blksum_row[blk] * QuantBBlkSum[blksum_offset]; - } - - c_row[n] = acc; - } - } - - return CountM; -} - -} // namespace sq2bit_avx512_super -} // namespace mlas -} // namespace onnxruntime diff --git a/onnxruntime/core/mlas/lib/sqnbitgemm_kernel_avx512_2bit_superblock.h b/onnxruntime/core/mlas/lib/sqnbitgemm_kernel_avx512_2bit_superblock.h deleted file mode 100644 index a2a6b93d4c7ee..0000000000000 --- a/onnxruntime/core/mlas/lib/sqnbitgemm_kernel_avx512_2bit_superblock.h +++ /dev/null @@ -1,249 +0,0 @@ -/*++ - -Copyright (c) Microsoft Corporation. All rights reserved. - -Licensed under the MIT License. - -Module Name: - - sqnbitgemm_kernel_avx512_2bit_superblock.h - -Abstract: - - Pack-time helpers and scalar reference routines for the EXPERIMENTAL - super-block W2 layout: groups of 4 K-blocks share a single 64-byte - packed buffer that allows the AVX-512 unpack to be one ZMM load plus - four fixed `vpsrlw+vpand` pairs (instead of the current per-block - broadcast + variable shift). - - Layout summary (BlkLen=64 only): - - * Each "super-block" packs FOUR consecutive K-blocks (256 weights total). - * Total storage per super-block = kBlkBytes * 4 = 64 bytes (identical to - 4 separately-packed blocks under the current scheme). - * Byte b of the super-block holds: - bits[0..1] = block_0.weight[b] - bits[2..3] = block_1.weight[b] - bits[4..5] = block_2.weight[b] - bits[6..7] = block_3.weight[b] - * The N-dimension uses the same 4-col-grouped layout as the existing - W2 kernel (kNCols4 = 4), so the main NMain region groups 4 N-cols - per "row" of super-blocks. - - Restrictions: - - * BlkLen == 64 only. - * BlockCountK must be a multiple of kSuperBlockBlks = 4. The prototype - path returns 0 from the pack-size helper for non-multiples; the - caller falls back to the existing W2 path. - * The customer model's K dimensions (384, 1024, 4096) are all multiples - of 256 (= 4 * 64), so all customer shapes satisfy this constraint. - ---*/ - -#pragma once - -#include -#include -#include - -#include "mlas.h" -#include "mlas_qnbit.h" - -#include "sqnbitgemm_kernel_avx512_2bit.h" // Re-uses kBlkLen, kBlkBytes, kNCols4, etc. - -template -struct PackedQuantBDataStruct; // fwd decl, defined in qnbitgemm.h - -struct MLAS_BACKEND_KERNEL_SELECTOR_CONFIG; - -namespace onnxruntime { -namespace mlas { -namespace sq2bit_avx512_super { - -using ::onnxruntime::mlas::sq2bit_avx512::kBlkBytes; -using ::onnxruntime::mlas::sq2bit_avx512::kBlkLen; -using ::onnxruntime::mlas::sq2bit_avx512::kDefaultSymmetricZeroPoint2Bit; -using ::onnxruntime::mlas::sq2bit_avx512::kNCols4; -using ::onnxruntime::mlas::sq2bit_avx512::kSuperBlockBytes; -using ::onnxruntime::mlas::sq2bit_avx512::kSuperBlockBlks; -using ::onnxruntime::mlas::sq2bit_avx512::kWeightsPerByte; - -// ----------------------------------------------------------------------------- -// Super-block packed-data layout. -// -// Main region (n < NMain = floor(N / kNCols4) * kNCols4): -// 4-N-col groups of g = n / 4, col within group c = n % 4. -// K-super-block index s = blk / 4 (s in [0, BlockCountK / 4)). -// Within a group, super-blocks run consecutively across the 4 cols, -// so each (s, group) slot is a contiguous (kNCols4 * kSuperBlockBytes) -// = 256 byte chunk. -// -// Tail region (n >= NMain): plain column-major super-blocks, identical -// shape to the main region but flat in N. -// -// K-tail handling (BlockCountK not a multiple of kSuperBlockBlks): -// The pack helpers round BlockCountK up to a multiple of kSuperBlockBlks -// (= 4) for storage purposes -- the padding 1-3 blocks at the trailing -// super-block contain zeroed weight bytes and zeroed scales, so they -// contribute 0 to the integer dot product and 0 to the BlkSum correction. -// This lets the SIMD kernel iterate the super-block K-loop without a -// special tail handler for B, and avoids dual packing layouts. Storage -// waste is at most (kSuperBlockBlks - 1) blocks per N-col, i.e. <= 48 -// bytes per col -- negligible at any realistic N. -// -// Conventions used by the offset helpers below: -// * `SuperBlockCountKPadded = ceil(BlockCountK / kSuperBlockBlks)` is -// the number of super-blocks the kernel actually iterates. -// * `BlockCountKPadded = SuperBlockCountKPadded * kSuperBlockBlks` is -// the K-block count used to address the scale buffer. -// * Callers must pass `SuperBlockCountKPadded` and `BlockCountKPadded` -// to these helpers; the original logical BlockCountK is only used -// for sizing the BlkSum buffer (which is consumed by the SGEMM -// correction step, not the inner K-loop). -// -// Caller-side constraints: BlkLen == 64; BlockCountK >= 1 (any K, padded -// internally to a multiple of kSuperBlockBlks). -// ----------------------------------------------------------------------------- - -inline size_t -PackedQuantBOffsetBytes_W2_SuperBlock(size_t n, size_t blk_super, - size_t SuperBlockCountKPadded, size_t NMain) -{ - if (n < NMain) { - const size_t g = n / kNCols4; - const size_t c = n % kNCols4; - const size_t per_group_bytes = SuperBlockCountKPadded * kNCols4 * kSuperBlockBytes; - return g * per_group_bytes - + blk_super * (kNCols4 * kSuperBlockBytes) - + c * kSuperBlockBytes; - } - return (n * SuperBlockCountKPadded + blk_super) * kSuperBlockBytes; -} - -// -// Float offset into the packed B-scale buffer for a logical (n, blk) cell. -// Scales remain per-block (one float per K-block), 4 per super. Caller -// passes BlockCountKPadded (= SuperBlockCountKPadded * kSuperBlockBlks); -// scale slots in [BlockCountK, BlockCountKPadded) contain zeros so the -// kernel can index uniformly. -// -inline size_t -PackedQuantBScaleOffset_W2_SuperBlock(size_t n, size_t blk, - size_t BlockCountKPadded, size_t NMain) -{ - const size_t SuperBlockCountKPadded = BlockCountKPadded / kSuperBlockBlks; - const size_t blk_super = blk / kSuperBlockBlks; - const size_t blk_in_super = blk % kSuperBlockBlks; - if (n < NMain) { - const size_t g = n / kNCols4; - const size_t c = n % kNCols4; - const size_t per_group_scales = SuperBlockCountKPadded * kNCols4 * kSuperBlockBlks; - return g * per_group_scales - + blk_super * (kNCols4 * kSuperBlockBlks) - + c * kSuperBlockBlks - + blk_in_super; - } - return n * BlockCountKPadded + blk; -} - -// -// Reference (scalar) entry points -- defined in sqnbitgemm_kernel_avx512_2bit_superblock.cpp. -// -// These cover Phase 2 of the super-block prototype: pack + scalar GEMM oracle -// against which the SIMD inner loop (Phase 3) will be validated. They are -// reachable from unit tests via direct linkage; production dispatch wiring -// happens in Phase 4 after the SIMD path is correctness-clean. -// - -size_t MLASCALL -Q2BitGemmPackQuantBDataSize_SuperBlock( - size_t N, - size_t K, - size_t BlkLen, - bool HasZeroPoint, - MLAS_QNBIT_GEMM_COMPUTE_TYPE ComputeType, - const MLAS_BACKEND_KERNEL_SELECTOR_CONFIG* BackendKernelSelectorConfig -); - -void MLASCALL -SQ2BitGemmPackQuantBDataAndBlkSum_SuperBlockScalar( - size_t N, - size_t K, - size_t BlkLen, - MLAS_QNBIT_GEMM_COMPUTE_TYPE ComputeType, - const std::byte* QuantBDataBegin, - const float* QuantBScaleBegin, - bool HasZeroPoint, - const std::byte* QuantBZPBegin, - PackedQuantBDataStruct& PackedQuantB, - MLAS_THREADPOOL* ThreadPool, - const MLAS_BACKEND_KERNEL_SELECTOR_CONFIG* BackendKernelSelectorConfig -); - -size_t MLASCALL -SQ2BitGemmKernel_BlkSum_CompInt8_SuperBlockScalar( - size_t BlkLen, - const std::byte* QuantA, - const float* QuantAScale, - const std::byte* QuantBData, - const float* QuantBScale, - const std::byte* QuantBZeroPoint, - float* C, - size_t CountM, - size_t CountN, - size_t CountK, - size_t BlockCountK, - const float* Bias, - size_t ldc, - const float* ABlockSum, - const float* QuantBBlkSum -); - -// -// Unit-test forwarders for the AVX-512 SIMD super-block kernels. Same gating -// rules as the existing W2 test entries: the caller MUST verify -// GetMlasPlatform().Avx512Supported_ (and, for the VNNI variant, that the -// active dispatch is the AVX-512-VNNI one) before invoking these symbols. -// -size_t MLASCALL -SQ2BitGemmKernel_BlkSum_CompInt8_Super_Avx512_TestEntry( - size_t BlkLen, - const std::byte* QuantA, - const float* QuantAScale, - const std::byte* QuantBData, - const float* QuantBScale, - const std::byte* QuantBZeroPoint, - float* C, - size_t CountM, - size_t CountN, - size_t CountK, - size_t BlockCountK, - const float* Bias, - size_t ldc, - const float* ABlockSum, - const float* QuantBBlkSum -); - -size_t MLASCALL -SQ2BitGemmKernel_BlkSum_CompInt8_Super_Avx512Vnni_TestEntry( - size_t BlkLen, - const std::byte* QuantA, - const float* QuantAScale, - const std::byte* QuantBData, - const float* QuantBScale, - const std::byte* QuantBZeroPoint, - float* C, - size_t CountM, - size_t CountN, - size_t CountK, - size_t BlockCountK, - const float* Bias, - size_t ldc, - const float* ABlockSum, - const float* QuantBBlkSum -); - -} // namespace sq2bit_avx512_super -} // namespace mlas -} // namespace onnxruntime diff --git a/onnxruntime/core/mlas/lib/sqnbitgemm_kernel_avx512vnni.cpp b/onnxruntime/core/mlas/lib/sqnbitgemm_kernel_avx512vnni.cpp index 24c3fdbdcf159..8e4d5a18d517c 100644 --- a/onnxruntime/core/mlas/lib/sqnbitgemm_kernel_avx512vnni.cpp +++ b/onnxruntime/core/mlas/lib/sqnbitgemm_kernel_avx512vnni.cpp @@ -28,9 +28,7 @@ Module Name: #include "sqnbitgemm_kernel_avx512_int8_blklen64.h" #include "sqnbitgemm_kernel_avx512_int8_blklen128.h" #include "sqnbitgemm_kernel_avx512_2bit.h" -#include "sqnbitgemm_kernel_avx512vnni_2bit_blklen64.h" -#include "sqnbitgemm_kernel_avx512_2bit_superblock.h" -#include "sqnbitgemm_kernel_avx512vnni_2bit_blklen64_superblock.h" +#include "sqnbitgemm_kernel_avx512_2bit_blklen64.h" MLAS_FORCEINLINE void SQ4BitGemmM1Kernel_CompFp32( @@ -462,10 +460,10 @@ SQ8BitGemmPackQuantBDataAndBlkSum512vnni( } // -// Unit-test entry point for the AVX-512-VNNI W2 kernel. Mirrors the -// AVX-512BW (non-VNNI) entry point in sqnbitgemm_kernel_avx512.cpp; see the -// comment there for the rationale. The caller must verify AVX-512-VNNI is -// available on the host before calling. +// Unit-test entry point for the AVX-512-VNNI W2 kernel +// (sqnbitgemm_kernel_avx512_2bit_blklen64.h). Exposed for +// direct invocation from tests; production wiring (Phase 4) will add the +// runtime dispatcher integration. // namespace onnxruntime::mlas::sq2bit_avx512 { size_t MLASCALL @@ -492,37 +490,6 @@ SQ2BitGemmKernel_BlkSum_CompInt8_Avx512Vnni_TestEntry( } } // namespace onnxruntime::mlas::sq2bit_avx512 -// -// Unit-test entry point for the AVX-512-VNNI W2 SUPER-BLOCK kernel -// (sqnbitgemm_kernel_avx512vnni_2bit_blklen64_superblock.h). Exposed for -// direct invocation from tests; production wiring (Phase 4) will add the -// runtime dispatcher integration. -// -namespace onnxruntime::mlas::sq2bit_avx512_super { -size_t MLASCALL -SQ2BitGemmKernel_BlkSum_CompInt8_Super_Avx512Vnni_TestEntry( - size_t BlkLen, - const std::byte* QuantA, - const float* QuantAScale, - const std::byte* QuantBData, - const float* QuantBScale, - const std::byte* QuantBZeroPoint, - float* C, - size_t CountM, - size_t CountN, - size_t CountK, - size_t BlockCountK, - const float* Bias, - size_t ldc, - const float* ABlockSum, - const float* QuantBBlkSum) -{ - return SQ2BitGemmKernel_BlkSum_CompInt8_Super_Avx512Vnni( - BlkLen, QuantA, QuantAScale, QuantBData, QuantBScale, QuantBZeroPoint, - C, CountM, CountN, CountK, BlockCountK, Bias, ldc, ABlockSum, QuantBBlkSum); -} -} // namespace onnxruntime::mlas::sq2bit_avx512_super - // // Kernel dispatch structure definition. // @@ -545,24 +512,14 @@ const MLAS_QNBIT_GEMM_DISPATCH MlasSQNBitGemmDispatchAvx512vnni = []() { d.SQ8BitGemmKernel_BlkSum_CompInt8 = SQ8BitGemmKernel_BlkSum_CompInt8_avx512vnni; d.QuantizeARowComputeBlkSum_CompInt8 = QuantizeARow_CompInt8_avx512; - // 2-bit native CompInt8 path: AVX-512-VNNI variant of the super-block - // (W2-v2) kernel -- single 64-byte load + four fixed shift/mask pairs - // to unpack 4 K-blocks at once. Pack-size and pack functions are shared - // with the AVX-512BW variant; the per-block MAC is `_mm512_dpbusd_epi32` - // instead of `vpmaddubsw + vpmaddwd + vpaddd`. - // - // The legacy `sq2bit_avx512::*` (W2-v1) symbols remain in the build but - // are no longer reached at runtime. A follow-up will remove them once - // W2-v2 has soaked in production. - d.Q2BitGemmPackQuantBDataSize = onnxruntime::mlas::sq2bit_avx512_super::Q2BitGemmPackQuantBDataSize_SuperBlock; - d.SQ2BitGemmPackQuantBDataAndBlkSum = onnxruntime::mlas::sq2bit_avx512_super::SQ2BitGemmPackQuantBDataAndBlkSum_SuperBlockScalar; - d.SQ2BitGemmKernel_BlkSum_CompInt8 = onnxruntime::mlas::sq2bit_avx512_super::SQ2BitGemmKernel_BlkSum_CompInt8_Super_Avx512Vnni_TestEntry; - // W2-v2 packs and addresses each N-col at a stride of - // SuperBlockCountKPadded * kSuperBlockBlks blocks (BlockCountK rounded - // UP to a multiple of 4) to keep the inner K-loop's super-block stride - // constant. The dispatcher needs this stride for its per-N-tile pointer - // arithmetic in SQ2BitGemm_CompInt8. - d.Q2BitGemmEffectiveBlockCountK = [](size_t BlockCountK) { return ((BlockCountK + 3) / 4) * 4; }; + // 2-bit native CompInt8 path. Single dispatch entry (W2): + // 64-byte ZMM load + four fixed shift/mask pairs to unpack 4 K-blocks at + // once, with VNNI's `vpdpbusd` for the integer MAC. See the matching + // comment in sqnbitgemm_kernel_avx512.cpp for the K-padding contract. + d.Q2BitGemmPackQuantBDataSize = onnxruntime::mlas::sq2bit_avx512::Q2BitGemmPackQuantBDataSize_Avx512; + d.SQ2BitGemmPackQuantBDataAndBlkSum = onnxruntime::mlas::sq2bit_avx512::SQ2BitGemmPackQuantBDataAndBlkSum_Scalar; + d.SQ2BitGemmKernel_BlkSum_CompInt8 = onnxruntime::mlas::sq2bit_avx512::SQ2BitGemmKernel_BlkSum_CompInt8_Avx512Vnni_TestEntry; + d.Q2BitGemmEffectiveBlockCountK = [](size_t BlockCountK) { return ((BlockCountK + 3) / 4) * 4; }; return d; }(); diff --git a/onnxruntime/core/mlas/lib/sqnbitgemm_kernel_avx512vnni_2bit_blklen64.h b/onnxruntime/core/mlas/lib/sqnbitgemm_kernel_avx512vnni_2bit_blklen64.h deleted file mode 100644 index 6d679ad3a5ce0..0000000000000 --- a/onnxruntime/core/mlas/lib/sqnbitgemm_kernel_avx512vnni_2bit_blklen64.h +++ /dev/null @@ -1,989 +0,0 @@ -/*++ - -Copyright (c) Microsoft Corporation. All rights reserved. - -Licensed under the MIT License. - -Module Name: - - sqnbitgemm_kernel_avx512vnni_2bit_blklen64.h - -Abstract: - - AVX-512 tiled kernel for the 2-bit weight CompInt8 GEMM path - (BlkBitWidth=2, BlkLen=64). Header-only; the kernel is templated on a - `bool kVnni` parameter so the same source supports both: - - * AVX-512-VNNI host: `kVnni == true` -> single `_mm512_dpbusd_epi32` - per block MAC. Included by sqnbitgemm_kernel_avx512vnni.cpp. - * AVX-512 (no VNNI): `kVnni == false` -> three-instruction MAC chain - `vpmaddubsw + vpmaddwd + vpaddd` (AVX-512BW only). Included by - sqnbitgemm_kernel_avx512.cpp. - - The file name still carries `vnni` for git-history continuity; the - kernel handles both ISA targets via the template parameter, mirroring - W4's `dot_accumulate_2blk` / `dot_accumulate_2blkvnni` split in - sqnbitgemm_kernel_avx512_int8_blklen64.h. - - Architecture mirrors W4's MlasQ4Int8GemmKernelBlkLen64Avx512 - (in sqnbitgemm_kernel_avx512_int8_blklen64.h): - - * Per-(m,n) accumulator is a __m512 carrying lane-interleaved - scaled partial sums. Final _mm512_reduce_add_ps happens ONCE per - (m,n) tile, not per K-block. - * Outer tile shape is R2 x C4 (2 M-rows x 4 N-cols) so each A vector - load amortises across 4 N-cols and each B load across 2 M-rows. - 16 ZMM accumulators in flight. - * Inner unroll is PerAccuBlk2 = 2 K-blocks per iteration. Pairs of - dpbusd outputs are interleaved (unpacklo/hi + add) so a single - FMA applies both blocks' scales in one shot. - * Packed-B layout is 4-N-col-grouped + 2-K-block-paired (matches W4), - so the 4 cols of a tile read as a contiguous 128-byte stream. - * BlkLen=64 only. - * Symmetric quantization only (QuantBZeroPoint must be null). - * Tail tiles R2xC1, R1xC4, R1xC1 cover M % 2 != 0 / N % 4 != 0. - - Zero-point correction is performed OUTSIDE the int8 kernel via the - platform float SGEMM kernel (GetMlasPlatform().GemmFloatKernel), exactly - as W4 does. This requires QuantBBlkSum to be in the W4 "width-16 row- - major chunked" layout, which is produced by - SQ2BitGemmPackQuantBDataAndBlkSum_Scalar. - - Dequant prologue (per 64-element block): - - __m128i p = _mm_loadu_si128(packed); // 16 packed bytes - __m512i p4 = _mm512_broadcast_i32x4(p); // 4 lanes of those 16 bytes - __m512i sh = {0,0,0,0, 2,2,2,2, 4,4,4,4, 6,6,6,6}; // per-dword shifts - __m512i v = _mm512_srlv_epi32(p4, sh); - __m512i b = _mm512_and_si512(v, _mm512_set1_epi8(0x03)); - ---*/ - -#pragma once - -#include -#include -#include - -#include - -#include "mlasi.h" -#include "qnbitgemm.h" -#include "sqnbitgemm_kernel_avx512_2bit.h" - -namespace onnxruntime { -namespace mlas { -namespace sq2bit_avx512 { - -// kNCols4 (= 4) and kPerAccuBlk2 (= 2) come from the shared header so the -// pack layout and the kernel use the same constants. kNRows2 is the M-tile -// shape, kernel-only. -inline constexpr size_t kNRows2 = 2; - -// -// Dequant one 64-element 2-bit weight block from the packed (broadcast + -// shift) layout into a ZMM of 64 unsigned bytes in [0, 3]. -// -// Bytes 0..15 : weights[0..15] (shift 0, & 0x03) / Bytes 16..31 : weights[16..31] (shift 2, & 0x03) -// Bytes 32..47 : weights[32..47] (shift 4, & 0x03) -// Bytes 48..63 : weights[48..63] (shift 6, & 0x03) -// -static MLAS_FORCEINLINE __m512i -unpack_w2_blk_to_zmm(__m128i p128) -{ - const __m512i p_dup = _mm512_broadcast_i32x4(p128); - // Per-dword right shifts, in memory order (lane 0 first): - // [0,0,0,0, 2,2,2,2, 4,4,4,4, 6,6,6,6] - // _mm512_set_epi32 takes args in reverse order (lane 15 first). - const __m512i shifts = _mm512_set_epi32( - 6, 6, 6, 6, - 4, 4, 4, 4, - 2, 2, 2, 2, - 0, 0, 0, 0); - const __m512i mask03 = _mm512_set1_epi8(0x03); - return _mm512_and_si512(_mm512_srlv_epi32(p_dup, shifts), mask03); -} - -// -// Load + dequant ONE block. Used by the single-block (tail) helpers. -// -static MLAS_FORCEINLINE __m512i -load_unpack_1blk_w2(const std::byte* packed) -{ - return unpack_w2_blk_to_zmm( - _mm_loadu_si128(reinterpret_cast(packed))); -} - -// -// Load + dequant TWO consecutive K-blocks via one 256-bit YMM load. -// 32 packed bytes -> 2 ZMMs of 64 weights each. -// -static MLAS_FORCEINLINE void -load_unpack_2blk_w2(const std::byte* packed, __m512i& bv0_64_epi8, __m512i& bv1_64_epi8) -{ - const __m256i p_ymm = _mm256_loadu_si256(reinterpret_cast(packed)); - bv0_64_epi8 = unpack_w2_blk_to_zmm(_mm256_castsi256_si128(p_ymm)); - bv1_64_epi8 = unpack_w2_blk_to_zmm(_mm256_extracti128_si256(p_ymm, 1)); -} - -// -// Lane-interleaved 2-K-block accumulator (single M-row, single N-col). -// Mirrors W4's dot_accumulate_2blk / dot_accumulate_2blkvnni split (in -// sqnbitgemm_kernel_avx512_int8_blklen64.h). Identical math for W2 because -// the dequanted block is already in the same uint8 [0,3] form W4 produces. -// -// acc += sum_2blks( cvt(dot(bv, av)) * scale_a * scale_b ) -// -// Two variants: -// * dot_accumulate_2blk_w2: AVX-512BW only (vpmaddubsw + vpmaddwd + add). -// Reduces at epi16 granularity, then ONE vpmaddwd -// folds adjacent pairs to epi32 (one madd_epi16 -// saved per K-block-pair vs an epi32-interleave -// approach). -// * dot_accumulate_2blk_w2_vnni: VNNI variant using _mm512_dpbusd_epi32 -// for the inner MAC. -// -// `scale_a` and `scale_b` point to TWO consecutive floats (scales for blk0 -// and blk1). The double-broadcast trick gives the 16-lane pattern -// [s0,s1, s0,s1, s0,s1, s0,s1, s0,s1, s0,s1, s0,s1, s0,s1] -// which matches the post-interleave lane layout (blk0 in even lanes, -// blk1 in odd lanes). -// -static MLAS_FORCEINLINE __m512i -ones_32_epi16_w2() -{ - const __m512i zeros = _mm512_setzero_si512(); - return _mm512_srli_epi16(_mm512_ternarylogic_epi64(zeros, zeros, zeros, 1), 15); -} - -static MLAS_FORCEINLINE void -dot_accumulate_2blk_w2( - const __m512i& av0_64_epi8, - const __m512i& av1_64_epi8, - const float* scale_a, - const __m512i& bv0_64_epi8, - const __m512i& bv1_64_epi8, - const __m512& scale_b_16_ps, - __m512& acc) -{ - const __m512i dot0_32_epi16 = _mm512_maddubs_epi16(bv0_64_epi8, av0_64_epi8); - const __m512i dot1_32_epi16 = _mm512_maddubs_epi16(bv1_64_epi8, av1_64_epi8); - - const __m512i t1 = _mm512_unpacklo_epi32(dot0_32_epi16, dot1_32_epi16); - const __m512i t2 = _mm512_unpackhi_epi32(dot0_32_epi16, dot1_32_epi16); - const __m512i sum_32_epi16 = _mm512_add_epi16(t1, t2); // [b0 b0 b1 b1 ...] in epi16 - const __m512i ones = ones_32_epi16_w2(); - const __m512i sum_16_epi32 = _mm512_madd_epi16(ones, sum_32_epi16); // [b0 b1 b0 b1 ...] in epi32 - const __m512 sum_16_ps = _mm512_cvtepi32_ps(sum_16_epi32); - - const __m256 scale_a_8_ps = _mm256_castpd_ps(_mm256_broadcast_sd(reinterpret_cast(scale_a))); - const __m512 scale_a_16_ps = _mm512_broadcast_f32x8(scale_a_8_ps); - - acc = _mm512_fmadd_ps(sum_16_ps, _mm512_mul_ps(scale_a_16_ps, scale_b_16_ps), acc); -} - -static MLAS_FORCEINLINE void -dot_accumulate_2blk_w2_vnni( - const __m512i& av0_64_epi8, - const __m512i& av1_64_epi8, - const float* scale_a, - const __m512i& bv0_64_epi8, - const __m512i& bv1_64_epi8, - const __m512& scale_b_16_ps, - __m512& acc) -{ - const __m512i dot0_16_epi32 = _mm512_dpbusd_epi32(_mm512_setzero_si512(), bv0_64_epi8, av0_64_epi8); - const __m512i dot1_16_epi32 = _mm512_dpbusd_epi32(_mm512_setzero_si512(), bv1_64_epi8, av1_64_epi8); - - const __m512i t1 = _mm512_unpacklo_epi32(dot0_16_epi32, dot1_16_epi32); - const __m512i t2 = _mm512_unpackhi_epi32(dot0_16_epi32, dot1_16_epi32); - const __m512i sum_16_epi32 = _mm512_add_epi32(t1, t2); - const __m512 sum_16_ps = _mm512_cvtepi32_ps(sum_16_epi32); - - const __m256 scale_a_8_ps = _mm256_castpd_ps(_mm256_broadcast_sd(reinterpret_cast(scale_a))); - const __m512 scale_a_16_ps = _mm512_broadcast_f32x8(scale_a_8_ps); - - acc = _mm512_fmadd_ps(sum_16_ps, _mm512_mul_ps(scale_a_16_ps, scale_b_16_ps), acc); -} - -// -// Single-K-block accumulator. Uses uniform 16-lane scale broadcast since -// there's only one block's scale in play. -// -template -static MLAS_FORCEINLINE void -dot_accumulate_1blk_w2( - const __m512i& av_64_epi8, - const float* scale_a, - const __m512i& bv_64_epi8, - const __m512& scale_b_16_ps, - __m512& acc) -{ - __m512i dot_16_epi32; - if constexpr (kVnni) { - dot_16_epi32 = _mm512_dpbusd_epi32(_mm512_setzero_si512(), bv_64_epi8, av_64_epi8); - } else { - const __m512i ones = ones_32_epi16_w2(); - const __m512i dot_32_epi16 = _mm512_maddubs_epi16(bv_64_epi8, av_64_epi8); - dot_16_epi32 = _mm512_madd_epi16(dot_32_epi16, ones); - } - const __m512 sum_16_ps = _mm512_cvtepi32_ps(dot_16_epi32); - - const __m128 scale_a_ps = _mm_broadcast_ss(scale_a); - const __m512 scale_a_16_ps = _mm512_broadcast_f32x2(scale_a_ps); - - acc = _mm512_fmadd_ps(sum_16_ps, _mm512_mul_ps(scale_a_16_ps, scale_b_16_ps), acc); -} - -// -// 2 M-rows x 1 N-col x 2 K-blocks accumulator. The 2-block B load is shared -// across the 2 M-rows. -// -template -static MLAS_FORCEINLINE void -accumulate_w2_blklen64_r2c1blk2( - const __m512i& av00_64_epi8, const __m512i& av01_64_epi8, - const __m512i& av10_64_epi8, const __m512i& av11_64_epi8, - const std::byte* QuantBDataPtr, - const float* scale_a0, - const float* scale_a1, - const float* scale_b, - __m512& acc0, - __m512& acc1) -{ - __m512i bv0, bv1; - load_unpack_2blk_w2(QuantBDataPtr, bv0, bv1); - - const __m256 scale_b_8_ps = _mm256_castpd_ps(_mm256_broadcast_sd(reinterpret_cast(scale_b))); - const __m512 scale_b_16_ps = _mm512_broadcast_f32x8(scale_b_8_ps); - - if constexpr (kVnni) { - dot_accumulate_2blk_w2_vnni(av00_64_epi8, av01_64_epi8, scale_a0, bv0, bv1, scale_b_16_ps, acc0); - dot_accumulate_2blk_w2_vnni(av10_64_epi8, av11_64_epi8, scale_a1, bv0, bv1, scale_b_16_ps, acc1); - } else { - dot_accumulate_2blk_w2(av00_64_epi8, av01_64_epi8, scale_a0, bv0, bv1, scale_b_16_ps, acc0); - dot_accumulate_2blk_w2(av10_64_epi8, av11_64_epi8, scale_a1, bv0, bv1, scale_b_16_ps, acc1); - } -} - -// -// 2 M-rows x 1 N-col x 1 K-block accumulator (K-tail). -// -template -static MLAS_FORCEINLINE void -accumulate_w2_blklen64_r2c1blk1( - const __m512i& av0_64_epi8, - const __m512i& av1_64_epi8, - const std::byte* QuantBDataPtr, - const float* scale_a0, - const float* scale_a1, - const float* scale_b, - __m512& acc0, - __m512& acc1) -{ - const __m512i bv = load_unpack_1blk_w2(QuantBDataPtr); - - const __m128 scale_b_ps = _mm_broadcast_ss(scale_b); - const __m512 scale_b_16_ps = _mm512_broadcast_f32x2(scale_b_ps); - - dot_accumulate_1blk_w2(av0_64_epi8, scale_a0, bv, scale_b_16_ps, acc0); - dot_accumulate_1blk_w2(av1_64_epi8, scale_a1, bv, scale_b_16_ps, acc1); -} - -// -// 1 M-row x 1 N-col x 2 K-blocks accumulator. -// -template -static MLAS_FORCEINLINE void -accumulate_w2_blklen64_r1c1blk2( - const __m512i& av0_64_epi8, - const __m512i& av1_64_epi8, - const std::byte* QuantBDataPtr, - const float* scale_a, - const float* scale_b, - __m512& acc) -{ - __m512i bv0, bv1; - load_unpack_2blk_w2(QuantBDataPtr, bv0, bv1); - - const __m256 scale_b_8_ps = _mm256_castpd_ps(_mm256_broadcast_sd(reinterpret_cast(scale_b))); - const __m512 scale_b_16_ps = _mm512_broadcast_f32x8(scale_b_8_ps); - - if constexpr (kVnni) { - dot_accumulate_2blk_w2_vnni(av0_64_epi8, av1_64_epi8, scale_a, bv0, bv1, scale_b_16_ps, acc); - } else { - dot_accumulate_2blk_w2(av0_64_epi8, av1_64_epi8, scale_a, bv0, bv1, scale_b_16_ps, acc); - } -} - -// -// 1 M-row x 1 N-col x 1 K-block accumulator. -// -template -static MLAS_FORCEINLINE void -accumulate_w2_blklen64_r1c1blk1( - const __m512i& av_64_epi8, - const std::byte* QuantBDataPtr, - const float* scale_a, - const float* scale_b, - __m512& acc) -{ - const __m512i bv = load_unpack_1blk_w2(QuantBDataPtr); - - const __m128 scale_b_ps = _mm_broadcast_ss(scale_b); - const __m512 scale_b_16_ps = _mm512_broadcast_f32x2(scale_b_ps); - - dot_accumulate_1blk_w2(av_64_epi8, scale_a, bv, scale_b_16_ps, acc); -} - -// -// R2 x C4 tile. Main hot path for the customer model. -// -// Layout assumptions (matches the W4 layout produced by W2's pack function): -// * QuantBData : 4-N-col grouped, 2-K-block paired (W4-style). Within a -// group the inner col stride is kPerAccuBlk2 * kBlkBytes (32 B) for pair -// blocks and kBlkBytes (16 B) for the single-block trailing slot when -// BlockCountK is odd. Per-group total = BlockCountK * kNCols4 * kBlkBytes. -// * QuantBScale: same grouping; col stride kPerAccuBlk2 (2 floats) / -// 1 float in the single-block slot. Per-group total = BlockCountK * kNCols4. -// * QuantA : row-major int8, BlockCountK * kBlkLen bytes per row. -// * QuantAScale: row-major float, BlockCountK floats per row. -// -template -MLAS_FORCEINLINE void -Q2Int8GemmR2xC4BlkLen64Avx512( - const std::byte* QuantA, - const float* QuantAScale, - const std::byte* QuantBData, - const float* QuantBScale, - float* C, - size_t CountM, - size_t CountN, - size_t BlockCountK, - const float* Bias, - size_t ldc) -{ - const size_t lda = BlockCountK * kBlkLen; - // Per-group strides for the 4-N-col grouped layout. - constexpr size_t PerColPairBytes = kPerAccuBlk2 * kBlkBytes; // 32 B per col in a K-pair slot - constexpr size_t PerColSingleBytes = kBlkBytes; // 16 B per col in a K-single slot - constexpr size_t PerColPairScale = kPerAccuBlk2; // 2 floats per col in a K-pair slot - constexpr size_t PerKPairAdvanceBytes = kNCols4 * PerColPairBytes; // 128 B per K-pair iteration - constexpr size_t PerKSingleAdvanceBytes = kNCols4 * PerColSingleBytes; - constexpr size_t PerKPairAdvanceScale = kNCols4 * PerColPairScale; // 8 floats per K-pair iteration - constexpr size_t PerKSingleAdvanceScale = kNCols4; - const size_t GroupStrideBytes = BlockCountK * kNCols4 * kBlkBytes; - const size_t GroupStrideScale = BlockCountK * kNCols4; - - assert(CountM % kNRows2 == 0); - assert(CountN % kNCols4 == 0); - - for (size_t m = 0; m < CountM; m += kNRows2) { - const std::byte* QuantBDataColPtr = QuantBData; - const float* QuantBScaleColPtr = QuantBScale; - const float* BiasPtr = Bias; - float* SumPtr = C + m * ldc; - - for (size_t n = 0; n < CountN; n += kNCols4) { - const std::byte* QuantAPtr = QuantA + m * lda; - const float* QuantAScalePtr = QuantAScale + m * BlockCountK; - - const std::byte* QuantBDataPtr = QuantBDataColPtr; - const float* QuantBScalePtr = QuantBScaleColPtr; - - __m512 acc[kNCols4 * kNRows2] = { - _mm512_setzero_ps(), _mm512_setzero_ps(), _mm512_setzero_ps(), _mm512_setzero_ps(), - _mm512_setzero_ps(), _mm512_setzero_ps(), _mm512_setzero_ps(), _mm512_setzero_ps() - }; - - size_t k_blks_remaining = BlockCountK; - for (; k_blks_remaining > 1; k_blks_remaining -= kPerAccuBlk2) { - const __m512i av_00 = _mm512_loadu_si512(reinterpret_cast(QuantAPtr)); - const __m512i av_01 = _mm512_loadu_si512(reinterpret_cast(QuantAPtr + kBlkLen)); - const __m512i av_10 = _mm512_loadu_si512(reinterpret_cast(QuantAPtr + lda)); - const __m512i av_11 = _mm512_loadu_si512(reinterpret_cast(QuantAPtr + lda + kBlkLen)); - - accumulate_w2_blklen64_r2c1blk2( - av_00, av_01, av_10, av_11, - QuantBDataPtr + 0 * PerColPairBytes, - QuantAScalePtr, QuantAScalePtr + BlockCountK, - QuantBScalePtr + 0 * PerColPairScale, - acc[0], acc[kNCols4 + 0]); - accumulate_w2_blklen64_r2c1blk2( - av_00, av_01, av_10, av_11, - QuantBDataPtr + 1 * PerColPairBytes, - QuantAScalePtr, QuantAScalePtr + BlockCountK, - QuantBScalePtr + 1 * PerColPairScale, - acc[1], acc[kNCols4 + 1]); - accumulate_w2_blklen64_r2c1blk2( - av_00, av_01, av_10, av_11, - QuantBDataPtr + 2 * PerColPairBytes, - QuantAScalePtr, QuantAScalePtr + BlockCountK, - QuantBScalePtr + 2 * PerColPairScale, - acc[2], acc[kNCols4 + 2]); - accumulate_w2_blklen64_r2c1blk2( - av_00, av_01, av_10, av_11, - QuantBDataPtr + 3 * PerColPairBytes, - QuantAScalePtr, QuantAScalePtr + BlockCountK, - QuantBScalePtr + 3 * PerColPairScale, - acc[3], acc[kNCols4 + 3]); - - QuantAPtr += kBlkLen * kPerAccuBlk2; - QuantAScalePtr += kPerAccuBlk2; - QuantBDataPtr += PerKPairAdvanceBytes; - QuantBScalePtr += PerKPairAdvanceScale; - } - - while (k_blks_remaining-- > 0) { - const __m512i av_00 = _mm512_loadu_si512(reinterpret_cast(QuantAPtr)); - const __m512i av_10 = _mm512_loadu_si512(reinterpret_cast(QuantAPtr + lda)); - - accumulate_w2_blklen64_r2c1blk1( - av_00, av_10, - QuantBDataPtr + 0 * PerColSingleBytes, - QuantAScalePtr, QuantAScalePtr + BlockCountK, - QuantBScalePtr + 0, - acc[0], acc[kNCols4 + 0]); - accumulate_w2_blklen64_r2c1blk1( - av_00, av_10, - QuantBDataPtr + 1 * PerColSingleBytes, - QuantAScalePtr, QuantAScalePtr + BlockCountK, - QuantBScalePtr + 1, - acc[1], acc[kNCols4 + 1]); - accumulate_w2_blklen64_r2c1blk1( - av_00, av_10, - QuantBDataPtr + 2 * PerColSingleBytes, - QuantAScalePtr, QuantAScalePtr + BlockCountK, - QuantBScalePtr + 2, - acc[2], acc[kNCols4 + 2]); - accumulate_w2_blklen64_r2c1blk1( - av_00, av_10, - QuantBDataPtr + 3 * PerColSingleBytes, - QuantAScalePtr, QuantAScalePtr + BlockCountK, - QuantBScalePtr + 3, - acc[3], acc[kNCols4 + 3]); - - QuantAPtr += kBlkLen; - QuantAScalePtr++; - QuantBDataPtr += PerKSingleAdvanceBytes; - QuantBScalePtr += PerKSingleAdvanceScale; - } - - SumPtr[0] = _mm512_reduce_add_ps(acc[0]); - SumPtr[1] = _mm512_reduce_add_ps(acc[1]); - SumPtr[2] = _mm512_reduce_add_ps(acc[2]); - SumPtr[3] = _mm512_reduce_add_ps(acc[3]); - SumPtr[ldc + 0] = _mm512_reduce_add_ps(acc[kNCols4 + 0]); - SumPtr[ldc + 1] = _mm512_reduce_add_ps(acc[kNCols4 + 1]); - SumPtr[ldc + 2] = _mm512_reduce_add_ps(acc[kNCols4 + 2]); - SumPtr[ldc + 3] = _mm512_reduce_add_ps(acc[kNCols4 + 3]); - if (BiasPtr != nullptr) { - SumPtr[0] += BiasPtr[0]; - SumPtr[1] += BiasPtr[1]; - SumPtr[2] += BiasPtr[2]; - SumPtr[3] += BiasPtr[3]; - SumPtr[ldc + 0] += BiasPtr[0]; - SumPtr[ldc + 1] += BiasPtr[1]; - SumPtr[ldc + 2] += BiasPtr[2]; - SumPtr[ldc + 3] += BiasPtr[3]; - } - - QuantBDataColPtr += GroupStrideBytes; - QuantBScaleColPtr += GroupStrideScale; - BiasPtr += BiasPtr != nullptr ? kNCols4 : 0; - SumPtr += kNCols4; - } - } -} - -// -// R2 x C1 tile (N-tail). Operates on the column-major tail region of the -// packed B buffer (cols NMain..N-1), which the dispatcher addresses with -// the same `multipleCols * (BlockCountK * kBlkBytes)` offset as the old -// pure column-major layout did (the grouped main region was sized to fit -// in exactly that many bytes by construction). -// -template -MLAS_FORCEINLINE void -Q2Int8GemmR2xC1BlkLen64Avx512( - const std::byte* QuantA, - const float* QuantAScale, - const std::byte* QuantBData, - const float* QuantBScale, - float* C, - size_t CountM, - size_t CountN, - size_t BlockCountK, - const float* Bias, - size_t ldc) -{ - const size_t lda = BlockCountK * kBlkLen; - const size_t ColStrideBytes = BlockCountK * kBlkBytes; - const size_t ColStrideScale = BlockCountK; - - assert(CountM % kNRows2 == 0); - - for (size_t m = 0; m < CountM; m += kNRows2) { - const std::byte* QuantBDataColPtr = QuantBData; - const float* QuantBScaleColPtr = QuantBScale; - const float* BiasPtr = Bias; - float* SumPtr = C + m * ldc; - - for (size_t n = 0; n < CountN; ++n) { - const std::byte* QuantAPtr = QuantA + m * lda; - const float* QuantAScalePtr = QuantAScale + m * BlockCountK; - - const std::byte* QuantBDataPtr = QuantBDataColPtr; - const float* QuantBScalePtr = QuantBScaleColPtr; - - __m512 acc0 = _mm512_setzero_ps(); - __m512 acc1 = _mm512_setzero_ps(); - - size_t k_blks_remaining = BlockCountK; - for (; k_blks_remaining > 1; k_blks_remaining -= kPerAccuBlk2) { - const __m512i av_00 = _mm512_loadu_si512(reinterpret_cast(QuantAPtr)); - const __m512i av_01 = _mm512_loadu_si512(reinterpret_cast(QuantAPtr + kBlkLen)); - const __m512i av_10 = _mm512_loadu_si512(reinterpret_cast(QuantAPtr + lda)); - const __m512i av_11 = _mm512_loadu_si512(reinterpret_cast(QuantAPtr + lda + kBlkLen)); - - accumulate_w2_blklen64_r2c1blk2( - av_00, av_01, av_10, av_11, - QuantBDataPtr, - QuantAScalePtr, QuantAScalePtr + BlockCountK, - QuantBScalePtr, - acc0, acc1); - - QuantAPtr += kBlkLen * kPerAccuBlk2; - QuantAScalePtr += kPerAccuBlk2; - QuantBDataPtr += kPerAccuBlk2 * kBlkBytes; - QuantBScalePtr += kPerAccuBlk2; - } - - while (k_blks_remaining-- > 0) { - const __m512i av_00 = _mm512_loadu_si512(reinterpret_cast(QuantAPtr)); - const __m512i av_10 = _mm512_loadu_si512(reinterpret_cast(QuantAPtr + lda)); - - accumulate_w2_blklen64_r2c1blk1( - av_00, av_10, - QuantBDataPtr, - QuantAScalePtr, QuantAScalePtr + BlockCountK, - QuantBScalePtr, - acc0, acc1); - - QuantAPtr += kBlkLen; - QuantAScalePtr++; - QuantBDataPtr += kBlkBytes; - QuantBScalePtr++; - } - - SumPtr[0] = _mm512_reduce_add_ps(acc0); - SumPtr[ldc] = _mm512_reduce_add_ps(acc1); - if (BiasPtr != nullptr) { - SumPtr[0] += BiasPtr[0]; - SumPtr[ldc] += BiasPtr[0]; - } - - QuantBDataColPtr += ColStrideBytes; - QuantBScaleColPtr += ColStrideScale; - BiasPtr += BiasPtr != nullptr ? 1 : 0; - SumPtr += 1; - } - } -} - -// -// R1 x C4 tile (M-tail). Uses the same 4-N-col grouped layout as R2 x C4. -// -template -MLAS_FORCEINLINE void -Q2Int8GemmR1xC4BlkLen64Avx512( - const std::byte* QuantA, - const float* QuantAScale, - const std::byte* QuantBData, - const float* QuantBScale, - float* C, - size_t CountM, - size_t CountN, - size_t BlockCountK, - const float* Bias, - size_t ldc) -{ - const size_t lda = BlockCountK * kBlkLen; - constexpr size_t PerColPairBytes = kPerAccuBlk2 * kBlkBytes; - constexpr size_t PerColSingleBytes = kBlkBytes; - constexpr size_t PerColPairScale = kPerAccuBlk2; - constexpr size_t PerKPairAdvanceBytes = kNCols4 * PerColPairBytes; - constexpr size_t PerKSingleAdvanceBytes = kNCols4 * PerColSingleBytes; - constexpr size_t PerKPairAdvanceScale = kNCols4 * PerColPairScale; - constexpr size_t PerKSingleAdvanceScale = kNCols4; - const size_t GroupStrideBytes = BlockCountK * kNCols4 * kBlkBytes; - const size_t GroupStrideScale = BlockCountK * kNCols4; - - assert(CountN % kNCols4 == 0); - - for (size_t m = 0; m < CountM; ++m) { - const std::byte* QuantBDataColPtr = QuantBData; - const float* QuantBScaleColPtr = QuantBScale; - const float* BiasPtr = Bias; - float* SumPtr = C + m * ldc; - - for (size_t n = 0; n < CountN; n += kNCols4) { - const std::byte* QuantAPtr = QuantA + m * lda; - const float* QuantAScalePtr = QuantAScale + m * BlockCountK; - - const std::byte* QuantBDataPtr = QuantBDataColPtr; - const float* QuantBScalePtr = QuantBScaleColPtr; - - __m512 acc[kNCols4] = { - _mm512_setzero_ps(), _mm512_setzero_ps(), - _mm512_setzero_ps(), _mm512_setzero_ps() - }; - - size_t k_blks_remaining = BlockCountK; - for (; k_blks_remaining > 1; k_blks_remaining -= kPerAccuBlk2) { - const __m512i av_0 = _mm512_loadu_si512(reinterpret_cast(QuantAPtr)); - const __m512i av_1 = _mm512_loadu_si512(reinterpret_cast(QuantAPtr + kBlkLen)); - - accumulate_w2_blklen64_r1c1blk2( - av_0, av_1, - QuantBDataPtr + 0 * PerColPairBytes, - QuantAScalePtr, QuantBScalePtr + 0 * PerColPairScale, acc[0]); - accumulate_w2_blklen64_r1c1blk2( - av_0, av_1, - QuantBDataPtr + 1 * PerColPairBytes, - QuantAScalePtr, QuantBScalePtr + 1 * PerColPairScale, acc[1]); - accumulate_w2_blklen64_r1c1blk2( - av_0, av_1, - QuantBDataPtr + 2 * PerColPairBytes, - QuantAScalePtr, QuantBScalePtr + 2 * PerColPairScale, acc[2]); - accumulate_w2_blklen64_r1c1blk2( - av_0, av_1, - QuantBDataPtr + 3 * PerColPairBytes, - QuantAScalePtr, QuantBScalePtr + 3 * PerColPairScale, acc[3]); - - QuantAPtr += kBlkLen * kPerAccuBlk2; - QuantAScalePtr += kPerAccuBlk2; - QuantBDataPtr += PerKPairAdvanceBytes; - QuantBScalePtr += PerKPairAdvanceScale; - } - - while (k_blks_remaining-- > 0) { - const __m512i av = _mm512_loadu_si512(reinterpret_cast(QuantAPtr)); - - accumulate_w2_blklen64_r1c1blk1( - av, QuantBDataPtr + 0 * PerColSingleBytes, - QuantAScalePtr, QuantBScalePtr + 0, acc[0]); - accumulate_w2_blklen64_r1c1blk1( - av, QuantBDataPtr + 1 * PerColSingleBytes, - QuantAScalePtr, QuantBScalePtr + 1, acc[1]); - accumulate_w2_blklen64_r1c1blk1( - av, QuantBDataPtr + 2 * PerColSingleBytes, - QuantAScalePtr, QuantBScalePtr + 2, acc[2]); - accumulate_w2_blklen64_r1c1blk1( - av, QuantBDataPtr + 3 * PerColSingleBytes, - QuantAScalePtr, QuantBScalePtr + 3, acc[3]); - - QuantAPtr += kBlkLen; - QuantAScalePtr++; - QuantBDataPtr += PerKSingleAdvanceBytes; - QuantBScalePtr += PerKSingleAdvanceScale; - } - - SumPtr[0] = _mm512_reduce_add_ps(acc[0]); - SumPtr[1] = _mm512_reduce_add_ps(acc[1]); - SumPtr[2] = _mm512_reduce_add_ps(acc[2]); - SumPtr[3] = _mm512_reduce_add_ps(acc[3]); - if (BiasPtr != nullptr) { - SumPtr[0] += BiasPtr[0]; - SumPtr[1] += BiasPtr[1]; - SumPtr[2] += BiasPtr[2]; - SumPtr[3] += BiasPtr[3]; - } - - QuantBDataColPtr += GroupStrideBytes; - QuantBScaleColPtr += GroupStrideScale; - BiasPtr += BiasPtr != nullptr ? kNCols4 : 0; - SumPtr += kNCols4; - } - } -} - -// -// R1 x C1 tile (corner). Same column-major tail-region addressing as R2 x C1. -// -template -MLAS_FORCEINLINE void -Q2Int8GemmR1xC1BlkLen64Avx512( - const std::byte* QuantA, - const float* QuantAScale, - const std::byte* QuantBData, - const float* QuantBScale, - float* C, - size_t CountM, - size_t CountN, - size_t BlockCountK, - const float* Bias, - size_t ldc) -{ - const size_t lda = BlockCountK * kBlkLen; - const size_t ColStrideBytes = BlockCountK * kBlkBytes; - const size_t ColStrideScale = BlockCountK; - - for (size_t m = 0; m < CountM; ++m) { - const std::byte* QuantBDataColPtr = QuantBData; - const float* QuantBScaleColPtr = QuantBScale; - const float* BiasPtr = Bias; - float* SumPtr = C + m * ldc; - - for (size_t n = 0; n < CountN; ++n) { - const std::byte* QuantAPtr = QuantA + m * lda; - const float* QuantAScalePtr = QuantAScale + m * BlockCountK; - - const std::byte* QuantBDataPtr = QuantBDataColPtr; - const float* QuantBScalePtr = QuantBScaleColPtr; - - __m512 acc = _mm512_setzero_ps(); - - size_t k_blks_remaining = BlockCountK; - for (; k_blks_remaining > 1; k_blks_remaining -= kPerAccuBlk2) { - const __m512i av_0 = _mm512_loadu_si512(reinterpret_cast(QuantAPtr)); - const __m512i av_1 = _mm512_loadu_si512(reinterpret_cast(QuantAPtr + kBlkLen)); - - accumulate_w2_blklen64_r1c1blk2( - av_0, av_1, QuantBDataPtr, QuantAScalePtr, QuantBScalePtr, acc); - - QuantAPtr += kBlkLen * kPerAccuBlk2; - QuantAScalePtr += kPerAccuBlk2; - QuantBDataPtr += kPerAccuBlk2 * kBlkBytes; - QuantBScalePtr += kPerAccuBlk2; - } - - while (k_blks_remaining-- > 0) { - const __m512i av = _mm512_loadu_si512(reinterpret_cast(QuantAPtr)); - - accumulate_w2_blklen64_r1c1blk1( - av, QuantBDataPtr, QuantAScalePtr, QuantBScalePtr, acc); - - QuantAPtr += kBlkLen; - QuantAScalePtr++; - QuantBDataPtr += kBlkBytes; - QuantBScalePtr++; - } - - SumPtr[0] = _mm512_reduce_add_ps(acc); - if (BiasPtr != nullptr) { - SumPtr[0] += BiasPtr[0]; - } - - QuantBDataColPtr += ColStrideBytes; - QuantBScaleColPtr += ColStrideScale; - BiasPtr += BiasPtr != nullptr ? 1 : 0; - SumPtr += 1; - } - } -} - -// -// Tile dispatcher. Mirrors W4's MlasQ4Int8GemmKernelBlkLen64Avx512. -// -template -MLAS_FORCEINLINE void -MlasQ2Int8GemmKernelBlkLen64Avx512( - const std::byte* QuantA, - const float* QuantAScale, - const std::byte* QuantBData, - const float* QuantBScale, - float* C, - size_t CountM, - size_t CountN, - size_t BlockCountK, - const float* Bias, - size_t ldc) -{ - const size_t lda = BlockCountK * kBlkLen; - const size_t lda_scale = BlockCountK; - const size_t ColStrideBytes = BlockCountK * kBlkBytes; - const size_t ColStrideScale = BlockCountK; - - const size_t remainingRows = CountM % kNRows2; - const size_t multipleRows = CountM - remainingRows; - const size_t remainingCols = CountN % kNCols4; - const size_t multipleCols = CountN - remainingCols; - - if (multipleRows > 0 && multipleCols > 0) { - Q2Int8GemmR2xC4BlkLen64Avx512( - QuantA, QuantAScale, QuantBData, QuantBScale, - C, multipleRows, multipleCols, BlockCountK, Bias, ldc); - } - if (remainingCols > 0 && multipleRows > 0) { - Q2Int8GemmR2xC1BlkLen64Avx512( - QuantA, QuantAScale, - QuantBData + multipleCols * ColStrideBytes, - QuantBScale + multipleCols * ColStrideScale, - C + multipleCols, - multipleRows, remainingCols, BlockCountK, - Bias ? Bias + multipleCols : nullptr, ldc); - } - if (remainingRows > 0 && multipleCols > 0) { - Q2Int8GemmR1xC4BlkLen64Avx512( - QuantA + multipleRows * lda, - QuantAScale + multipleRows * lda_scale, - QuantBData, QuantBScale, - C + multipleRows * ldc, - remainingRows, multipleCols, BlockCountK, Bias, ldc); - } - if (remainingRows > 0 && remainingCols > 0) { - Q2Int8GemmR1xC1BlkLen64Avx512( - QuantA + multipleRows * lda, - QuantAScale + multipleRows * lda_scale, - QuantBData + multipleCols * ColStrideBytes, - QuantBScale + multipleCols * ColStrideScale, - C + multipleRows * ldc + multipleCols, - remainingRows, remainingCols, BlockCountK, - Bias ? Bias + multipleCols : nullptr, ldc); - } -} - -// -// Common dispatched-kernel body. Templated on ; the two top-level -// wrappers below instantiate it. -// -// Steps: -// 1) Calls the tile dispatcher, which computes -// C[m,n] = bias[n] + sum_blk(scale_a * scale_b * dot(b, a)) -// using either VNNI (`_mm512_dpbusd_epi32`) or AVX-512BW -// (`vpmaddubsw + vpmaddwd + vpaddd`) depending on the template parameter. -// 2) Adds the symmetric zero-point correction -// C[m,n] += sum_blk(ABlockSum[m,blk] * QuantBBlkSum[n,blk]) -// via the platform float SGEMM micro-kernel. QuantBBlkSum is in the -// W4 width-16 row-major chunked layout (produced by the W2 pack -// function), which is what GemmFloatKernel expects for its packed-B -// operand. ZeroMode=false means SGEMM does `C += A @ B`, so the bias -// and int8 contribution already in C are preserved. -// -template -static MLAS_FORCEINLINE size_t -SQ2BitGemmKernel_BlkSum_CompInt8_Impl( - const size_t BlkLen, - const std::byte* QuantA, - const float* QuantAScale, - const std::byte* QuantBData, - const float* QuantBScale, - const std::byte* /* QuantBZeroPoint */, - float* C, - size_t CountM, - size_t CountN, - size_t /* CountK */, - size_t BlockCountK, - const float* Bias, - size_t ldc, - const float* ABlockSum, - const float* QuantBBlkSum) -{ - if (BlkLen != kBlkLen) { - return 0; - } - - MlasQ2Int8GemmKernelBlkLen64Avx512( - QuantA, QuantAScale, QuantBData, QuantBScale, - C, CountM, CountN, BlockCountK, Bias, ldc); - - // BlkSum correction: C += ABlockSum [M x BlockCountK] @ QuantBBlkSum [BlockCountK x N]. - // - // TEMP DEBUG: scalar reference instead of GetMlasPlatform().GemmFloatKernel - // to test whether the SGEMM-call shape/layout assumptions are wrong. - // QuantBBlkSum is in the "width-16 chunked" layout: - // BlkSum[(n/16) * BlockCountK * 16 + blk * 16 + (n%16)] - { - for (size_t m = 0; m < CountM; ++m) { - const float* a_row = ABlockSum + m * BlockCountK; - float* c_row = C + m * ldc; - for (size_t n = 0; n < CountN; ++n) { - const size_t chunk = n / 16; - const size_t lane = n % 16; - float acc = 0.0f; - for (size_t blk = 0; blk < BlockCountK; ++blk) { - const float b = QuantBBlkSum[(chunk * BlockCountK + blk) * 16 + lane]; - acc += a_row[blk] * b; - } - c_row[n] += acc; - } - } - } - // Original fast path (disabled for debug): - // float* c_blk = C; - // const float* b_blk_sum = QuantBBlkSum; - // size_t RowsRemaining = CountM; - // const float* a_blksum_row = ABlockSum; - // while (RowsRemaining > 0) { - // const auto RowsHandled = GetMlasPlatform().GemmFloatKernel( - // a_blksum_row, b_blk_sum, c_blk, - // BlockCountK, RowsRemaining, CountN, - // BlockCountK, ldc, 1.0f, false); - // - // c_blk += ldc * RowsHandled; - // a_blksum_row += BlockCountK * RowsHandled; - // RowsRemaining -= RowsHandled; - // } - - return CountM; -} - -// -// Top-level VNNI variant registered into MlasSQNBitGemmDispatchAvx512vnni. -// -static MLAS_FORCEINLINE size_t MLASCALL -SQ2BitGemmKernel_BlkSum_CompInt8_Avx512Vnni( - const size_t BlkLen, - const std::byte* QuantA, - const float* QuantAScale, - const std::byte* QuantBData, - const float* QuantBScale, - const std::byte* QuantBZeroPoint, - float* C, - size_t CountM, - size_t CountN, - size_t CountK, - size_t BlockCountK, - const float* Bias, - size_t ldc, - const float* ABlockSum, - const float* QuantBBlkSum) -{ - return SQ2BitGemmKernel_BlkSum_CompInt8_Impl( - BlkLen, QuantA, QuantAScale, QuantBData, QuantBScale, QuantBZeroPoint, - C, CountM, CountN, CountK, BlockCountK, Bias, ldc, ABlockSum, QuantBBlkSum); -} - -// -// Top-level non-VNNI variant registered into MlasSQNBitGemmDispatchAvx512. -// Uses the AVX-512BW MAC chain (`vpmaddubsw + vpmaddwd + vpaddd`) instead of -// `_mm512_dpbusd_epi32`. Same tile shapes, same pack layout, same numerical -// result. -// -static MLAS_FORCEINLINE size_t MLASCALL -SQ2BitGemmKernel_BlkSum_CompInt8_Avx512( - const size_t BlkLen, - const std::byte* QuantA, - const float* QuantAScale, - const std::byte* QuantBData, - const float* QuantBScale, - const std::byte* QuantBZeroPoint, - float* C, - size_t CountM, - size_t CountN, - size_t CountK, - size_t BlockCountK, - const float* Bias, - size_t ldc, - const float* ABlockSum, - const float* QuantBBlkSum) -{ - return SQ2BitGemmKernel_BlkSum_CompInt8_Impl( - BlkLen, QuantA, QuantAScale, QuantBData, QuantBScale, QuantBZeroPoint, - C, CountM, CountN, CountK, BlockCountK, Bias, ldc, ABlockSum, QuantBBlkSum); -} - -} // namespace sq2bit_avx512 -} // namespace mlas -} // namespace onnxruntime diff --git a/onnxruntime/test/contrib_ops/matmul_2bits_test.cc b/onnxruntime/test/contrib_ops/matmul_2bits_test.cc index 5ea9726bf2c20..d45ba429c0490 100644 --- a/onnxruntime/test/contrib_ops/matmul_2bits_test.cc +++ b/onnxruntime/test/contrib_ops/matmul_2bits_test.cc @@ -1027,6 +1027,74 @@ TEST(MatMul2Bits, Float32_2b_Accuracy4) { TestMatMul2BitsTyped(); } +// MatMulNBits operator-level coverage for the native AVX-512 W2 CPU kernel. +// The kernel is gated on (BlkBitWidth=2, BlkLen=64, ComputeType=SQNBIT_CompInt8) +// and is reachable through the public MatMulNBits op when accuracy_level=4 and +// block_size=64 (so the platform dispatcher picks the CompInt8 path on AVX-512 +// hosts). On non-AVX-512 hosts the path is unavailable and the op falls back +// to LUT or scalar -- the OpTester correctness check still passes because the +// expected output is computed via dequantize-and-matmul. +// +// The K-values are chosen to exercise: +// * Single block (BlockCountK=1) -> K=64 +// * BlockCountK=2,3 not multiple of 4 -> K=128, K=192 (K-tail handler) +// * Exact block-group multiples -> K=256, K=512, K=1024 +// * BlockCountK=6 (1 full + 2 tail) -> K=384 (matches customer model) +// and combinations of M ∈ {1 (decode), 2, 4, 100 (prefill)} and +// N ∈ {16 (N-tail R1xC1/R2xC1), 32, 288, 1024 (kNCols4 multiples)}. +// +// TestMatMul2BitsTyped fans each shape out to 4 sub-tests: ±has_zero_point +// crossed with ±has_bias. +TEST(MatMul2Bits, Float32_2b_BlkLen64_Accuracy4) { + // Single-block K (K = BlkLen = 64). + TestMatMul2BitsTyped(); + TestMatMul2BitsTyped(); + TestMatMul2BitsTyped(); + TestMatMul2BitsTyped(); + + // BlockCountK=2,3 -- exercises K-tail handler (BlockCountK not a multiple + // of kBlockGroupBlks=4). + TestMatMul2BitsTyped(); + TestMatMul2BitsTyped(); + TestMatMul2BitsTyped(); + + // Exact block-group multiples (BlockCountK = 4, 8, 16; no K-tail). + TestMatMul2BitsTyped(); + TestMatMul2BitsTyped(); + TestMatMul2BitsTyped(); + + // Customer-model proportion: BlockCountK=6 = 1 full block-group + 2-block tail. + TestMatMul2BitsTyped(); + + // Larger M (multiple R2 tile iterations) and customer-shape K=1024. + TestMatMul2BitsTyped(); + TestMatMul2BitsTyped(); +} + +// Same BlkLen=64 grid at accuracy_level=0. accuracy_level=0 lets the runtime +// pick the best path; for BlkBitWidth=2 + BlkLen=64 it still routes to the +// native W2 CompInt8 kernel on AVX-512 hosts. This catches any dispatch-table +// wiring bug that only manifests when accuracy_level isn't explicitly 4. +TEST(MatMul2Bits, Float32_2b_BlkLen64_Accuracy0) { + TestMatMul2BitsTyped(); + TestMatMul2BitsTyped(); + TestMatMul2BitsTyped(); + TestMatMul2BitsTyped(); + + TestMatMul2BitsTyped(); + TestMatMul2BitsTyped(); + TestMatMul2BitsTyped(); + + TestMatMul2BitsTyped(); + TestMatMul2BitsTyped(); + TestMatMul2BitsTyped(); + + TestMatMul2BitsTyped(); + + TestMatMul2BitsTyped(); + TestMatMul2BitsTyped(); +} + #if defined(USE_WEBGPU) && !defined(ORT_USE_EP_API_ADAPTERS) namespace { diff --git a/onnxruntime/test/mlas/bench/bench_lutgemm.cpp b/onnxruntime/test/mlas/bench/bench_lutgemm.cpp index f021b36d8d23f..422df80e9d247 100644 --- a/onnxruntime/test/mlas/bench/bench_lutgemm.cpp +++ b/onnxruntime/test/mlas/bench/bench_lutgemm.cpp @@ -240,8 +240,8 @@ static void LutGemmComputeArgs(benchmark::internal::Benchmark* b) { // (K=1024, N=4096): 20 nodes // (K=4096, N=1024): 20 nodes // Covers both M=1 (decode) and M=128 (prefill) so the LUT path can be -// compared apples-to-apples against the W4 CompInt8 and W2-prod/W2-super -// kernels (QNBITGEMM/QNBitGemmCustomerArgs and +// compared apples-to-apples against the W4 CompInt8 and W2 kernels +// (QNBITGEMM/QNBitGemmCustomerArgs and // QNBITGEMM/QNBit2BitCustomerArgs). static void LutGemmCustomerArgs(benchmark::internal::Benchmark* b) { b->ArgNames(lutgemm_compute_arg_names); diff --git a/onnxruntime/test/mlas/bench/bench_qnbitgemm.cpp b/onnxruntime/test/mlas/bench/bench_qnbitgemm.cpp index 00fe98d604622..4dba8b916bb75 100644 --- a/onnxruntime/test/mlas/bench/bench_qnbitgemm.cpp +++ b/onnxruntime/test/mlas/bench/bench_qnbitgemm.cpp @@ -17,14 +17,6 @@ #include "core/util/thread_utils.h" #include "core/platform/env_var_utils.h" -// Prototype W2 super-block kernel + scalar pack helper (Phase 3 of the -// W2-vs-W4 parity work). Not yet wired into the platform dispatch, so the -// bench drives it directly via the test-entry forwarder. -#include "core/mlas/lib/mlasi.h" -#include "core/mlas/lib/qnbitgemm.h" -#include "core/mlas/lib/sqnbitgemm_kernel_avx512_2bit.h" -#include "core/mlas/lib/sqnbitgemm_kernel_avx512_2bit_superblock.h" - template void RunQNBitGemmBenchmark(size_t BlkLen, size_t M, size_t N, size_t K, @@ -142,9 +134,28 @@ static void QNBitGemmArgs(benchmark::internal::Benchmark* b) { }); } +// Standard sweep for the native W2 kernel. W2 has fewer free dimensions than +// W4 (symmetric-only, BlkLen=64 only, SQNBIT_CompInt8 only), so the grid +// uses fixed values for those axes and sweeps the rest like QNBitGemmArgs. +static void QNBit2BitArgs(benchmark::internal::Benchmark* b) { + b->ArgNames({"BlkLen", "M", "N", "K", "Threads", "Symmetric", "HasBias", "ComputeType"}); + + b->ArgsProduct({ + {64}, // BlkLen (W2 native kernel constraint) + {1, 4096}, // M (decode + prefill) + {4096, 11008}, // N + {4096, 11008}, // K + {1, 8}, // Threads + {int64_t{true}}, // Symmetric (W2 native kernel constraint) + {int64_t{false}, int64_t{true}}, // HasBias + {int64_t{SQNBIT_CompInt8}}, // ComputeType (W2 native kernel constraint) + }); +} + BENCHMARK(QNBITGEMM)->Apply(QNBitGemmArgs)->UseRealTime(); BENCHMARK(QNBITGEMM)->Apply(QNBitGemmArgs)->UseRealTime(); BENCHMARK(QNBITGEMM)->Apply(QNBitGemmArgs)->UseRealTime(); +BENCHMARK(QNBITGEMM)->Apply(QNBit2BitArgs)->UseRealTime(); // Customer MatMulNBits shapes mirrored at 4-bit for a head-to-head comparison // vs the W2 LUT path (LUTGEMM_COMPUTE/CUSTOMER). Customer model uses BlkLen=64. @@ -200,195 +211,6 @@ static void QNBit2BitCustomerArgs(benchmark::internal::Benchmark* b) { BENCHMARK(QNBITGEMM)->Apply(QNBit2BitCustomerArgs)->UseRealTime(); -// --------------------------------------------------------------------------- -// W2 SUPER-BLOCK PROTOTYPE BENCHMARK -// --------------------------------------------------------------------------- -// Drives the Phase-3 super-block W2 kernel directly via its test-entry -// forwarder. Mirrors what MlasQNBitGemmBatch would do for our path: pre-pack -// B via the super-block helpers, quantize each A row using the dispatch's -// AVX-512 A-quantizer (the same one the production W2 path uses), and call -// the SIMD kernel per-thread on the N-tile chunks the dispatcher splits over. -// -// Constraints (must be satisfied for the kernel to handle the rows; bench -// skips otherwise): -// * BlkLen == 64 -// * K a multiple of (BlkLen * kSuperBlockBlks) = 256 -// * M a multiple of kNRows2 (=2) -// * N a multiple of kNCols4 (=4) -// -namespace bench_super { - -namespace sq2sb = onnxruntime::mlas::sq2bit_avx512_super; -namespace sq2 = onnxruntime::mlas::sq2bit_avx512; - -void RunQ2SuperBlockBenchmark(size_t M, size_t N, size_t K, size_t Threads, - bool HasBias, benchmark::State& state) { - using onnxruntime::narrow; - constexpr size_t BlkBitWidth = 2; - constexpr size_t BlkLen = sq2::kBlkLen; - constexpr size_t kSuperBlockBlks = sq2::kSuperBlockBlks; - // R2xC4 SIMD tile (kNRows2 lives in the SIMD-only header which the bench - // TU is not built against -- mirror its value here, used only for the - // N alignment gate below). The kernel handles any M >= 1 and any K >= 1 - // (K-tail handler covers K not a multiple of 256). - constexpr size_t kNCols4 = sq2::kNCols4; - (void)kSuperBlockBlks; // retained for documentation reference only - - // Gate on the kernel's hard constraints. - if (K % BlkLen != 0 || K == 0) { - state.SkipWithMessage("Super-block requires K > 0 and K % BlkLen == 0."); - return; - } - if (M == 0 || (N % kNCols4) != 0) { - state.SkipWithMessage("Super-block requires M>=1 and N%4==0."); - return; - } - - // Gate on host having AVX-512(-VNNI) -- the kernel is AVX-512 BW + (optional) VNNI. - // QuantizeARowComputeBlkSum_CompInt8 is AVX-512 and required for the A-quant step. - const auto& platform = GetMlasPlatform(); - if (platform.QNBitGemmDispatch == nullptr || - platform.QNBitGemmDispatch->QuantizeARowComputeBlkSum_CompInt8 == nullptr) { - state.SkipWithMessage("AVX-512 dispatch table not available on this host."); - return; - } - - const size_t BlockCountK = K / BlkLen; - - OrtThreadPoolParams tpo; - tpo.thread_pool_size = static_cast(Threads); - tpo.auto_set_affinity = true; - std::unique_ptr tp( - onnxruntime::concurrency::CreateThreadPool(&onnxruntime::Env::Default(), - tpo, - onnxruntime::concurrency::ThreadPoolType::INTRA_OP)); - - // ----- Source data ----- - const auto A = RandomVectorUniform(M * K, float{-1.0f}, float{1.0f}); - const auto B = RandomVectorUniform(K * N, float{-1.0f}, float{1.0f}); - const auto Bias = HasBias ? RandomVectorUniform(N, float{-1.0f}, float{1.0f}) : std::vector(); - - size_t QuantBDataSizeInBytes, QuantBScaleSize, QuantBZeroPointSizeInBytes; - MlasBlockwiseQuantizedBufferSizes( - static_cast(BlkLen), /*columnwise=*/true, - static_cast(K), static_cast(N), - QuantBDataSizeInBytes, QuantBScaleSize, &QuantBZeroPointSizeInBytes); - - std::vector QuantBDataSrc(QuantBDataSizeInBytes); - std::vector QuantBScale(QuantBScaleSize); - MlasQuantizeBlockwise( - QuantBDataSrc.data(), QuantBScale.data(), /*zp=*/nullptr, - B.data(), static_cast(BlkLen), /*columnwise=*/true, - static_cast(K), static_cast(N), static_cast(N), tp.get()); - - // ----- Pack into the super-block layout ----- - const size_t PackedSize = sq2sb::Q2BitGemmPackQuantBDataSize_SuperBlock( - N, K, BlkLen, /*HasZeroPoint=*/false, SQNBIT_CompInt8, nullptr); - if (PackedSize == 0) { - state.SkipWithMessage("Super-block pack size returned 0 for this shape."); - return; - } - std::vector PackedBuf(PackedSize, std::byte{0}); - PackedQuantBDataStruct packed_b( - PackedBuf.data(), N, BlockCountK, BlkLen, /*QuantAUnsigned=*/false); - - // Same 3-call prepack pattern matmul_nbits.cc uses. - sq2sb::SQ2BitGemmPackQuantBDataAndBlkSum_SuperBlockScalar( - N, K, BlkLen, SQNBIT_CompInt8, - reinterpret_cast(QuantBDataSrc.data()), - /*scales=*/nullptr, - /*has_zp=*/false, /*zp=*/nullptr, - packed_b, tp.get(), nullptr); - sq2sb::SQ2BitGemmPackQuantBDataAndBlkSum_SuperBlockScalar( - N, K, BlkLen, SQNBIT_CompInt8, - /*B=*/nullptr, QuantBScale.data(), - /*has_zp=*/false, /*zp=*/nullptr, - packed_b, tp.get(), nullptr); - - // ----- Quantize A once via the dispatch's AVX-512 A-quantizer ----- - std::vector QuantAData(M * BlockCountK * BlkLen, std::byte{0}); - std::vector QuantAScale(M * BlockCountK, 0.0f); - std::vector ABlockSum(M * BlockCountK, 0.0f); - auto QuantizeARow = platform.QNBitGemmDispatch->QuantizeARowComputeBlkSum_CompInt8; - for (size_t m = 0; m < M; ++m) { - QuantizeARow(BlkLen, A.data() + m * K, K, - QuantAData.data() + m * BlockCountK * BlkLen, - QuantAScale.data() + m * BlockCountK, - ABlockSum.data() + m * BlockCountK); - } - - std::vector C(M * N, 0.0f); - - // Pick the best SIMD variant for the host. Prefer the VNNI variant when - // the platform dispatch table is the VNNI one (same logic the production - // dispatch wiring would use). - const bool use_vnni = (platform.QNBitGemmDispatch == &MlasSQNBitGemmDispatchAvx512vnni); - auto kernel = use_vnni - ? sq2sb::SQ2BitGemmKernel_BlkSum_CompInt8_Super_Avx512Vnni_TestEntry - : sq2sb::SQ2BitGemmKernel_BlkSum_CompInt8_Super_Avx512_TestEntry; - - // Mirror the production dispatcher's N-tile parallel split. SQ2BitGemm_CompInt8 - // tiles N in chunks of 128 and parallelizes across them; we do the same. - constexpr size_t kNTile = 128; - const size_t num_n_tiles = (N + kNTile - 1) / kNTile; - - auto run_one = [&]() { - onnxruntime::concurrency::ThreadPool::TryBatchParallelFor( - tp.get(), static_cast(num_n_tiles), - [&](ptrdiff_t t) { - const size_t n_start = static_cast(t) * kNTile; - const size_t n_count = std::min(kNTile, N - n_start); - if (n_count == 0) return; - - // Pointer arithmetic mirroring SQ2BitGemm_CompInt8. - const size_t ldb_bytes = BlockCountK * sq2::kBlkBytes; - const size_t ldb_scale = BlockCountK; - const std::byte* b_tile = packed_b.PackedQuantBData + n_start * ldb_bytes; - const float* bscale_tile = packed_b.PackedQuantBScale + n_start * ldb_scale; - const float* bblksum_tile = packed_b.QuantBBlkSum + n_start * ldb_scale; - float* c_tile = C.data() + n_start; - const float* bias_tile = HasBias ? (Bias.data() + n_start) : nullptr; - - kernel(BlkLen, - QuantAData.data(), QuantAScale.data(), - b_tile, bscale_tile, - /*QuantBZeroPoint=*/nullptr, - c_tile, - M, n_count, K, BlockCountK, - bias_tile, - /*ldc=*/N, - ABlockSum.data(), bblksum_tile); - }, - /*cost=*/0); - }; - - run_one(); // warm up - for (auto _ : state) { - run_one(); - } -} - -} // namespace bench_super - -void QNBITGEMM_SUPER(benchmark::State& state) { - using onnxruntime::narrow; - const auto BlkLen = narrow(state.range(0)); - (void)BlkLen; // The super-block kernel only supports BlkLen=64; gated inside. - const auto M = narrow(state.range(1)); - const auto N = narrow(state.range(2)); - const auto K = narrow(state.range(3)); - const auto Threads = narrow(state.range(4)); - // state.range(5) (Symmetric) and (7) (ComputeType) are unused; the - // super-block kernel is symmetric CompInt8 only. - const bool HasBias = narrow(state.range(6)); - bench_super::RunQ2SuperBlockBenchmark(M, N, K, Threads, HasBias, state); -} - -// Customer-shape rows for the super-block prototype. Uses the same argument -// schema as QNBITGEMM so the bench rows line up one-to-one in the -// output (modulo skipped shapes that violate the super-block K constraint). -BENCHMARK(QNBITGEMM_SUPER)->Apply(QNBit2BitCustomerArgs)->UseRealTime(); - // This test gets benchmark arguments from environment variables. template void QNBITGEMM_ENV(benchmark::State& state) { diff --git a/onnxruntime/test/mlas/unittest/test_sqnbitgemm_2bit.cpp b/onnxruntime/test/mlas/unittest/test_sqnbitgemm_2bit.cpp index 002426335759a..159bb8c4fa8f9 100644 --- a/onnxruntime/test/mlas/unittest/test_sqnbitgemm_2bit.cpp +++ b/onnxruntime/test/mlas/unittest/test_sqnbitgemm_2bit.cpp @@ -10,12 +10,9 @@ Module Name: Abstract: - Unit tests for the 2-bit AVX-512-VNNI weight-GEMM helpers. - - Phase 2a coverage: pack / unpack round-trip of the BlkLen=64 packed - layout. The tests exercise sqnbitgemm_kernel_avx512_2bit.h directly; - they do not depend on platform dispatch being wired up, so they run on - every host (the helpers are pure scalar bit-twiddling). + Unit tests for the 2-bit AVX-512 weight-GEMM pack/unpack helpers + (block-group layout, BlkLen=64). Pure scalar bit-twiddling; runs on + every host because the exercised helpers do not contain SIMD code. --*/ @@ -49,175 +46,56 @@ PackSourceBlock_BlkLen64(const uint8_t weights[sq2::kBlkLen], std::byte* src_out } // namespace -// -// Pack then immediately unpack a single block. Each 2-bit position must -// survive the layout permutation exactly. -// -TEST(MlasSq2BitTest, PackUnpackRoundTrip_BlkLen64_DeterministicPattern) -{ - // Use a deterministic pattern that touches every (position, value) pair: - // weight i gets value (i % 4). This guarantees that if any bit-position - // accounting is off the failing index pinpoints the bug. - std::array weights{}; - for (size_t i = 0; i < weights.size(); ++i) { - weights[i] = static_cast(i % 4); - } - - std::array src{}; - PackSourceBlock_BlkLen64(weights.data(), src.data()); - - // Sanity check: the source unpack reproduces the original weights. - std::array via_src{}; - sq2::UnpackSourceBlock_BlkLen64_Reference(src.data(), via_src.data()); - for (size_t i = 0; i < weights.size(); ++i) { - ASSERT_EQ(via_src[i], weights[i]) << "Source-unpack disagrees at i=" << i; - } - - // Pack into the new layout, then unpack via the reference inverse. - std::array packed{}; - sq2::PackBlock_BlkLen64(src.data(), packed.data()); - - std::array recovered{}; - sq2::UnpackBlock_BlkLen64_Reference(packed.data(), recovered.data()); - - for (size_t i = 0; i < weights.size(); ++i) { - ASSERT_EQ(recovered[i], weights[i]) - << "Round-trip mismatch at i=" << i - << ": expected " << static_cast(weights[i]) - << ", got " << static_cast(recovered[i]); - } -} - -// -// Same round-trip but with pseudo-random weights, repeated across many -// blocks and seeds. Catches accidental position-dependent bugs that the -// deterministic pattern above might mask. -// -TEST(MlasSq2BitTest, PackUnpackRoundTrip_BlkLen64_Randomized) -{ - constexpr size_t kBlockCount = 17; // arbitrary, > 1, not a SIMD-friendly number - constexpr unsigned kSeeds = 8; - - for (unsigned seed = 0; seed < kSeeds; ++seed) { - std::mt19937 rng(seed * 7919u + 1u); - std::uniform_int_distribution dist(0u, 3u); - - std::vector weights(kBlockCount * sq2::kBlkLen); - for (auto& w : weights) { - w = static_cast(dist(rng)); - } - - std::vector src(kBlockCount * sq2::kBlkBytes); - for (size_t blk = 0; blk < kBlockCount; ++blk) { - PackSourceBlock_BlkLen64(weights.data() + blk * sq2::kBlkLen, - src.data() + blk * sq2::kBlkBytes); - } - - std::vector packed(kBlockCount * sq2::kPackedBlkBytes); - for (size_t blk = 0; blk < kBlockCount; ++blk) { - sq2::PackBlock_BlkLen64(src.data() + blk * sq2::kBlkBytes, - packed.data() + blk * sq2::kPackedBlkBytes); - } - - for (size_t blk = 0; blk < kBlockCount; ++blk) { - std::array recovered{}; - sq2::UnpackBlock_BlkLen64_Reference(packed.data() + blk * sq2::kPackedBlkBytes, - recovered.data()); - for (size_t i = 0; i < sq2::kBlkLen; ++i) { - ASSERT_EQ(recovered[i], weights[blk * sq2::kBlkLen + i]) - << "Random round-trip mismatch seed=" << seed - << " blk=" << blk - << " i=" << i; - } - } - } -} - -// -// Bit-position invariant: writing only value v into every weight slot must -// yield a packed buffer where every byte is 0x55 * v (== v repeated at -// positions 0,2,4,6). This catches confusion between low/high nibbles or -// reversed-bit packing. -// -TEST(MlasSq2BitTest, PackUnpackRoundTrip_BlkLen64_ConstantValues) -{ - for (uint8_t v = 0; v < 4; ++v) { - std::array weights{}; - weights.fill(v); - - std::array src{}; - PackSourceBlock_BlkLen64(weights.data(), src.data()); - - std::array packed{}; - sq2::PackBlock_BlkLen64(src.data(), packed.data()); - - const uint8_t expected_byte = static_cast(v * 0x55u); // v at bits {0..1,2..3,4..5,6..7} - for (size_t i = 0; i < sq2::kPackedBlkBytes; ++i) { - ASSERT_EQ(static_cast(packed[i]), expected_byte) - << "Constant-fill v=" << static_cast(v) - << " byte_i=" << i; - } - - std::array recovered{}; - sq2::UnpackBlock_BlkLen64_Reference(packed.data(), recovered.data()); - for (size_t i = 0; i < sq2::kBlkLen; ++i) { - ASSERT_EQ(recovered[i], v) << "Constant-fill v=" << static_cast(v) << " i=" << i; - } - } -} - // ----------------------------------------------------------------------------- -// EXPERIMENTAL: super-block (4-K-block) round-trip tests +// block-group (4-K-block) round-trip tests // ----------------------------------------------------------------------------- // -// These exercise PackSuperBlock4_BlkLen64 / UnpackSuperBlock4_BlkLen64_Reference, -// which underpin the fast-unpack prototype (single 64-byte load + 4 fixed -// shift-and-mask producing 4 block ZMMs vs the current broadcast + variable -// shift per block). The tests live alongside the per-block tests above so a -// regression in either layout is caught by the same test target. +// These exercise PackBlockGroup_BlkLen64 / UnPackBlockGroup_BlkLen64_Reference, +// which underpin the fast-unpack path (single 64-byte load + 4 fixed +// shift-and-mask producing 4 block ZMMs). // -// Deterministic-pattern super-block round-trip. Block k assigns weight i the +// Deterministic-pattern block-group round-trip. Block k assigns weight i the // value ((i + k) % 4), giving every (block_index, position, value) a unique // fingerprint that pinpoints a layout swap if any. // -TEST(MlasSq2BitTest, PackUnpackRoundTrip_SuperBlock4_BlkLen64_DeterministicPattern) +TEST(MlasSq2BitTest, PackUnpackRoundTrip_BlockGroup_BlkLen64_DeterministicPattern) { - std::array, sq2::kSuperBlockBlks> weights{}; - for (size_t k = 0; k < sq2::kSuperBlockBlks; ++k) { + std::array, sq2::kBlockGroupBlks> weights{}; + for (size_t k = 0; k < sq2::kBlockGroupBlks; ++k) { for (size_t i = 0; i < sq2::kBlkLen; ++i) { weights[k][i] = static_cast((i + k) % 4); } } - std::array, sq2::kSuperBlockBlks> src{}; - for (size_t k = 0; k < sq2::kSuperBlockBlks; ++k) { + std::array, sq2::kBlockGroupBlks> src{}; + for (size_t k = 0; k < sq2::kBlockGroupBlks; ++k) { PackSourceBlock_BlkLen64(weights[k].data(), src[k].data()); } - std::array packed{}; - sq2::PackSuperBlock4_BlkLen64(src[0].data(), src[1].data(), src[2].data(), src[3].data(), + std::array packed{}; + sq2::PackBlockGroup_BlkLen64(src[0].data(), src[1].data(), src[2].data(), src[3].data(), packed.data()); - std::array, sq2::kSuperBlockBlks> recovered{}; - sq2::UnpackSuperBlock4_BlkLen64_Reference(packed.data(), + std::array, sq2::kBlockGroupBlks> recovered{}; + sq2::UnPackBlockGroup_BlkLen64_Reference(packed.data(), recovered[0].data(), recovered[1].data(), recovered[2].data(), recovered[3].data()); - for (size_t k = 0; k < sq2::kSuperBlockBlks; ++k) { + for (size_t k = 0; k < sq2::kBlockGroupBlks; ++k) { for (size_t i = 0; i < sq2::kBlkLen; ++i) { ASSERT_EQ(recovered[k][i], weights[k][i]) - << "Super-block round-trip mismatch k=" << k << " i=" << i; + << "block-group round-trip mismatch k=" << k << " i=" << i; } } } // -// Randomized super-block round-trip across several seeds. The four input +// Randomized block-group round-trip across several seeds. The four input // blocks are independent random fills; the test fails fast if any (block, // weight) entry is mis-routed by the packed-byte layout. // -TEST(MlasSq2BitTest, PackUnpackRoundTrip_SuperBlock4_BlkLen64_Randomized) +TEST(MlasSq2BitTest, PackUnpackRoundTrip_BlockGroup_BlkLen64_Randomized) { constexpr unsigned kSeeds = 8; @@ -225,97 +103,97 @@ TEST(MlasSq2BitTest, PackUnpackRoundTrip_SuperBlock4_BlkLen64_Randomized) std::mt19937 rng(seed * 5051u + 13u); std::uniform_int_distribution dist(0u, 3u); - std::array, sq2::kSuperBlockBlks> weights{}; + std::array, sq2::kBlockGroupBlks> weights{}; for (auto& blk : weights) { for (auto& w : blk) { w = static_cast(dist(rng)); } } - std::array, sq2::kSuperBlockBlks> src{}; - for (size_t k = 0; k < sq2::kSuperBlockBlks; ++k) { + std::array, sq2::kBlockGroupBlks> src{}; + for (size_t k = 0; k < sq2::kBlockGroupBlks; ++k) { PackSourceBlock_BlkLen64(weights[k].data(), src[k].data()); } - std::array packed{}; - sq2::PackSuperBlock4_BlkLen64(src[0].data(), src[1].data(), src[2].data(), src[3].data(), + std::array packed{}; + sq2::PackBlockGroup_BlkLen64(src[0].data(), src[1].data(), src[2].data(), src[3].data(), packed.data()); - std::array, sq2::kSuperBlockBlks> recovered{}; - sq2::UnpackSuperBlock4_BlkLen64_Reference(packed.data(), + std::array, sq2::kBlockGroupBlks> recovered{}; + sq2::UnPackBlockGroup_BlkLen64_Reference(packed.data(), recovered[0].data(), recovered[1].data(), recovered[2].data(), recovered[3].data()); - for (size_t k = 0; k < sq2::kSuperBlockBlks; ++k) { + for (size_t k = 0; k < sq2::kBlockGroupBlks; ++k) { for (size_t i = 0; i < sq2::kBlkLen; ++i) { ASSERT_EQ(recovered[k][i], weights[k][i]) - << "Random super-block mismatch seed=" << seed << " k=" << k << " i=" << i; + << "Random block-group mismatch seed=" << seed << " k=" << k << " i=" << i; } } } } // -// Constant-value invariants for the super-block layout: +// Constant-value invariants for the block-group layout: // - All four blocks set to the same value v produces packed bytes equal to // 0x55 * v (v repeated at bit positions {0..1,2..3,4..5,6..7}). // - Block_k set to value v with all other blocks zero produces packed bytes // equal to (v << (2*k)) -- exclusively occupying the k-th bit slot. // -TEST(MlasSq2BitTest, PackUnpackRoundTrip_SuperBlock4_BlkLen64_ConstantValues) +TEST(MlasSq2BitTest, PackUnpackRoundTrip_BlockGroup_BlkLen64_ConstantValues) { // Case 1: every block filled with v. for (uint8_t v = 0; v < 4; ++v) { - std::array, sq2::kSuperBlockBlks> weights{}; + std::array, sq2::kBlockGroupBlks> weights{}; for (auto& blk : weights) { blk.fill(v); } - std::array, sq2::kSuperBlockBlks> src{}; - for (size_t k = 0; k < sq2::kSuperBlockBlks; ++k) { + std::array, sq2::kBlockGroupBlks> src{}; + for (size_t k = 0; k < sq2::kBlockGroupBlks; ++k) { PackSourceBlock_BlkLen64(weights[k].data(), src[k].data()); } - std::array packed{}; - sq2::PackSuperBlock4_BlkLen64(src[0].data(), src[1].data(), src[2].data(), src[3].data(), + std::array packed{}; + sq2::PackBlockGroup_BlkLen64(src[0].data(), src[1].data(), src[2].data(), src[3].data(), packed.data()); const uint8_t expected_byte = static_cast(v * 0x55u); - for (size_t i = 0; i < sq2::kSuperBlockBytes; ++i) { + for (size_t i = 0; i < sq2::kBlockGroupBytes; ++i) { ASSERT_EQ(static_cast(packed[i]), expected_byte) << "Uniform-fill v=" << static_cast(v) << " byte_i=" << i; } } // Case 2: only one block at a time carries a non-zero value. - for (size_t target_k = 0; target_k < sq2::kSuperBlockBlks; ++target_k) { + for (size_t target_k = 0; target_k < sq2::kBlockGroupBlks; ++target_k) { for (uint8_t v = 1; v < 4; ++v) { - std::array, sq2::kSuperBlockBlks> weights{}; + std::array, sq2::kBlockGroupBlks> weights{}; weights[target_k].fill(v); - std::array, sq2::kSuperBlockBlks> src{}; - for (size_t k = 0; k < sq2::kSuperBlockBlks; ++k) { + std::array, sq2::kBlockGroupBlks> src{}; + for (size_t k = 0; k < sq2::kBlockGroupBlks; ++k) { PackSourceBlock_BlkLen64(weights[k].data(), src[k].data()); } - std::array packed{}; - sq2::PackSuperBlock4_BlkLen64(src[0].data(), src[1].data(), src[2].data(), src[3].data(), + std::array packed{}; + sq2::PackBlockGroup_BlkLen64(src[0].data(), src[1].data(), src[2].data(), src[3].data(), packed.data()); const uint8_t expected_byte = static_cast(v << (2 * target_k)); - for (size_t i = 0; i < sq2::kSuperBlockBytes; ++i) { + for (size_t i = 0; i < sq2::kBlockGroupBytes; ++i) { ASSERT_EQ(static_cast(packed[i]), expected_byte) << "Isolated block target_k=" << target_k << " v=" << static_cast(v) << " byte_i=" << i; } - std::array, sq2::kSuperBlockBlks> recovered{}; - sq2::UnpackSuperBlock4_BlkLen64_Reference( + std::array, sq2::kBlockGroupBlks> recovered{}; + sq2::UnPackBlockGroup_BlkLen64_Reference( packed.data(), recovered[0].data(), recovered[1].data(), recovered[2].data(), recovered[3].data()); - for (size_t k = 0; k < sq2::kSuperBlockBlks; ++k) { + for (size_t k = 0; k < sq2::kBlockGroupBlks; ++k) { const uint8_t expect_val = (k == target_k) ? v : uint8_t{0}; for (size_t i = 0; i < sq2::kBlkLen; ++i) { ASSERT_EQ(recovered[k][i], expect_val) diff --git a/onnxruntime/test/mlas/unittest/test_sqnbitgemm_2bit_gemm.cpp b/onnxruntime/test/mlas/unittest/test_sqnbitgemm_2bit_gemm.cpp index f7c55f8eb86f9..aa48f9d6a153a 100644 --- a/onnxruntime/test/mlas/unittest/test_sqnbitgemm_2bit_gemm.cpp +++ b/onnxruntime/test/mlas/unittest/test_sqnbitgemm_2bit_gemm.cpp @@ -6,31 +6,20 @@ Licensed under the MIT License. Module Name: - test_sqnbitgemm_2bit_gemm.cpp + test_sqnbitgemm_2bit.cpp Abstract: - Numerical correctness tests for the 2-bit weight CompInt8 GEMM path on - AVX-512 hosts. Covers three execution routes: + Unit tests for the block-group W2 path + (sqnbitgemm_kernel_avx512_2bit.{h,cpp}). - 1) MlasQNBitGemmBatch (public API) -> platform-selected dispatch -> - AVX-512-VNNI W2 kernel. This is what production callers hit on a - VNNI host. + End-to-end pack + scalar GEMM correctness coverage These + tests do NOT exercise any SIMD path -- they validate the layout, the + pack 3-call sequence, and the scalar oracle kernel that will back the + SIMD kernel. - 2) AVX-512-VNNI W2 kernel via direct test-entry forwarder, bypassing - the platform dispatcher. Same kernel as (1); validates the - forwarder mechanism used by (3). - - 3) AVX-512BW (non-VNNI) W2 kernel via direct test-entry forwarder. - Validates the kernel a non-VNNI AVX-512 host would normally run. - On a VNNI host the platform dispatcher never picks this path, so - the direct-call route is the only way to exercise it. - - All three are compared against a single bit-exact integer-domain - reference (`ReferenceGemm_W2_CompInt8`) that reproduces the same per- - block int8 A quantization MLAS uses internally (amax/127, symmetric) - and the same dequant-free integer dot product against raw 2-bit - weights with an implicit zero-point of 2. + Tests deliberately use the same shapes as the production W2 tests so a + side-by-side comparison is straightforward. --*/ @@ -39,27 +28,24 @@ Module Name: #include #include #include -#include -#include #include #include #include "core/mlas/inc/mlas_qnbit.h" +#include "core/mlas/lib/qnbitgemm.h" +#include "core/mlas/lib/mlasi.h" #include "core/mlas/lib/sqnbitgemm_kernel_avx512_2bit.h" -#include "core/mlas/lib/qnbitgemm.h" // for PackedQuantBDataStruct (test direct-call path) -#include "core/mlas/lib/mlasi.h" // for GetMlasPlatform().Avx512Supported_ namespace { namespace sq2 = onnxruntime::mlas::sq2bit_avx512; -constexpr size_t kBlkLen = sq2::kBlkLen; // 64 -constexpr size_t kBlkBytes = sq2::kBlkBytes; // 16 source bytes per block -constexpr size_t kBlkBitWidth = 2; -constexpr MLAS_QNBIT_GEMM_COMPUTE_TYPE kComputeType = SQNBIT_CompInt8; +constexpr size_t kBlkLen = sq2::kBlkLen; // 64 +constexpr size_t kBlkBytes = sq2::kBlkBytes; // 16 +constexpr size_t kBlockGroupBlks = sq2::kBlockGroupBlks; // 4 -// Standard ONNX 2-bit packing: byte_i = w[4i] | w[4i+1]<<2 | w[4i+2]<<4 | w[4i+3]<<6. -inline void +// Standard ONNX 2-bit source packing (1 byte = 4 weights). +void PackSourceBlock_BlkLen64(const uint8_t weights[kBlkLen], std::byte* src_out) { for (size_t i = 0; i < kBlkBytes; ++i) { @@ -73,84 +59,49 @@ PackSourceBlock_BlkLen64(const uint8_t weights[kBlkLen], std::byte* src_out) } } -// -// Mirror the MLAS per-row int8 block quantizer used by the CompInt8 path: -// per-block symmetric scale = amax / 127, round-to-nearest, clamp to [-127, 127]. -// +// Bit-exact mirror of MlasQNBitGemm's per-block int8 quantizer (amax/127, +// round-half-to-even via std::nearbyint, scale_recip = 127/amax). void -QuantizeA_Reference(size_t M, - size_t K, - const float* A, - int8_t* QuantAData, - float* QuantAScale) +QuantizeA_Reference(size_t M, size_t K, const float* A, + int8_t* QuantAData, float* QuantAScale) { const size_t BlockCountK = (K + kBlkLen - 1) / kBlkLen; for (size_t m = 0; m < M; ++m) { for (size_t k = 0, k_blk = 0; k < K; k += kBlkLen, ++k_blk) { const size_t local_len = std::min(K - k, kBlkLen); - float amax = 0.0f; for (size_t kk = 0; kk < local_len; ++kk) { amax = std::max(amax, std::fabs(A[m * K + k + kk])); } - - constexpr float range_max = static_cast((1 << 7) - 1); + constexpr float range_max = 127.0f; const float scale = amax / range_max; - // Match MLAS's QuantizeARow_CompInt8_avx512 bit-for-bit: it computes - // `inverse_scale = 127 / amax` directly from amax (single op), not - // `1 / (amax/127)` (two ops). The two are equal in real math but - // differ by 1 ULP in float, which is enough to flip the rounded - // int8 by ±1 for A values that land right on a half-integer. const float scale_recip = amax != 0.0f ? range_max / amax : 0.0f; - QuantAScale[m * BlockCountK + k_blk] = scale; - for (size_t kk = 0; kk < kBlkLen; ++kk) { const float a = (kk < local_len) ? A[m * K + k + kk] : 0.0f; - // `std::nearbyint` uses the current floating-point rounding mode - // (default round-half-to-even) and matches what MLAS's - // `_mm512_roundscale_ps(v, _MM_ROUND_NEAREST)` does at .5 - // boundaries. `std::round` would diverge (rounds half away from - // zero), introducing systematic mismatches in the public-API path - // where MLAS quantises A and the reference computes ABlockSum - // from its own quantization. const float q = std::nearbyint(a * scale_recip); QuantAData[m * BlockCountK * kBlkLen + k + kk] = - static_cast( - std::clamp(q, - static_cast(std::numeric_limits::min()), - static_cast(std::numeric_limits::max()))); + static_cast(std::clamp(q, -127.0f, 127.0f)); } } } } // -// Reference GEMM that exactly mirrors the math performed by the MLAS W2 -// CompInt8 path: -// -// C[m,n] = bias[n] -// + sum_blk( scale_a[m,blk] * scale_b[n,blk] -// * dot(qa[m,blk,:], (qb[n,blk,:] - zp[n,blk])) ) -// -// When BZeroPoints is null the symmetric default ZP = 2 is used for every -// block (matches the kernel's behavior when no zero-point tensor is supplied). -// When non-null, BZeroPoints[n * BlockCountK + blk] gives the per-block ZP in -// [0, 3]. +// Integer-domain GEMM oracle: bit-exact match to the math the MLAS W2 path +// performs (kernel int8 GEMM + SGEMM zero-point correction collapsed into a +// single direct dot of (qa * (qb - zp))). // void -ReferenceGemm_W2_CompInt8(size_t M, - size_t N, - size_t K, +ReferenceGemm_W2_CompInt8(size_t M, size_t N, size_t K, const float* A, - const std::vector& BWeights, // [N * K] in [0,3] - const float* QuantBScale, - const uint8_t* BZeroPoints, // [N * BlockCountK] in [0,3] or nullptr + const std::vector& BWeights, // [N * K] in [0, 3] + const float* QuantBScale, // [N * BlockCountK] + const uint8_t* BZeroPoints, // [N * BlockCountK] or nullptr const float* Bias, float* C) { const size_t BlockCountK = (K + kBlkLen - 1) / kBlkLen; - std::vector QuantAData(M * BlockCountK * kBlkLen, int8_t{0}); std::vector QuantAScale(M * BlockCountK, 0.0f); QuantizeA_Reference(M, K, A, QuantAData.data(), QuantAScale.data()); @@ -165,11 +116,11 @@ ReferenceGemm_W2_CompInt8(size_t M, const int32_t zp = BZeroPoints != nullptr ? static_cast(BZeroPoints[n * BlockCountK + blk]) : static_cast(sq2::kDefaultSymmetricZeroPoint2Bit); - int32_t dot = 0; for (size_t kk = 0; kk < local_len; ++kk) { const int8_t qa = QuantAData[m * BlockCountK * kBlkLen + k + kk]; - const int32_t qb = static_cast(BWeights[n * K + k + kk]) - zp; + const int32_t qb = + static_cast(BWeights[n * K + k + kk]) - zp; dot += static_cast(qa) * qb; } acc += static_cast(dot) * a_scale * b_scale; @@ -180,11 +131,10 @@ ReferenceGemm_W2_CompInt8(size_t M, } // -// Pack per-block 2-bit zero points into the standard ONNX MatMulNBits W2 -// layout: 4 ZPs per byte along K, row-major in N. Row stride is -// ceil(BlockCountK / 4) bytes. +// Pack per-block W2 zero points into the standard ONNX byte stream +// (4 zp per byte along K, row-major in N). // -inline std::vector +std::vector PackW2ZeroPoints(size_t N, size_t BlockCountK, const std::vector& BZeroPoints) { const size_t ZPCountK = (BlockCountK + 3) / 4; @@ -201,238 +151,376 @@ PackW2ZeroPoints(size_t N, size_t BlockCountK, const std::vector& BZero return packed; } -class MlasSQ2BitGemmTest { - public: - static void Run(size_t M, size_t N, size_t K, bool WithBias, uint32_t seed, - bool WithZeroPoints = false) - { - const size_t BlockCountK = (K + kBlkLen - 1) / kBlkLen; - ASSERT_EQ(K % kBlkLen, 0u) << "Test K must be a multiple of BlkLen=64"; - - std::mt19937 rng(seed); - std::uniform_real_distribution a_dist(-1.0f, 1.0f); - std::uniform_int_distribution w_dist(0, 3); - std::uniform_real_distribution s_dist(0.05f, 0.5f); - - std::vector A(M * K); - for (auto& v : A) v = a_dist(rng); - - // Raw weights in [0,3], natural [n, k] order; the test owns this oracle - // copy and is the source of truth for the reference math. - std::vector BWeights(N * K); - for (auto& v : BWeights) v = static_cast(w_dist(rng)); - - // Source-packed B in the layout that MlasQNBitGemmPackQuantBData consumes: - // column-major in N, kBlkBytes per block, standard ONNX 4-weights-per-byte. - std::vector QuantBData(N * BlockCountK * kBlkBytes, std::byte{0}); - for (size_t n = 0; n < N; ++n) { - for (size_t blk = 0; blk < BlockCountK; ++blk) { - uint8_t blk_weights[kBlkLen]; - for (size_t kk = 0; kk < kBlkLen; ++kk) { - blk_weights[kk] = BWeights[n * K + blk * kBlkLen + kk]; - } - PackSourceBlock_BlkLen64( - blk_weights, - QuantBData.data() + (n * BlockCountK + blk) * kBlkBytes); +// +// Test harness that drives the block-group path directly (without going through +// MlasQNBitGemmBatch). Builds the same packed buffer the dispatcher would +// construct, runs the chosen block-group kernel (scalar / AVX-512BW / VNNI), +// and compares to ReferenceGemm_W2_CompInt8. +// +// `KernelFn` matches the SQ4BitGemmKernel_BlkSum_CompInt8_Fn signature, which +// every block-group kernel variant honors via direct-call forwarders declared +// in sqnbitgemm_kernel_avx512_2bit.h. +// +using W2KernelFn = size_t (MLASCALL*)( + size_t, const std::byte*, const float*, const std::byte*, const float*, + const std::byte*, float*, size_t, size_t, size_t, size_t, + const float*, size_t, const float*, const float*); + +void +RunW2Case(size_t M, size_t N, size_t K, bool WithBias, uint32_t seed, + bool WithZeroPoints, W2KernelFn kernel, + const char* kernel_name) +{ + const size_t BlockCountK = (K + kBlkLen - 1) / kBlkLen; + ASSERT_EQ(K % kBlkLen, 0u) << "Test K must be a multiple of BlkLen=64"; + // BlockCountK no longer required to be a multiple of kBlockGroupBlks -- + // the K-tail handler picks up the trailing 1-3 blocks. + + std::mt19937 rng(seed); + std::uniform_real_distribution a_dist(-1.0f, 1.0f); + std::uniform_int_distribution w_dist(0, 3); + std::uniform_real_distribution s_dist(0.05f, 0.5f); + + std::vector A(M * K); + for (auto& v : A) v = a_dist(rng); + + std::vector BWeights(N * K); + for (auto& v : BWeights) v = static_cast(w_dist(rng)); + + // Source-packed B (standard ONNX layout) -- the input to the pack helper. + std::vector QuantBData(N * BlockCountK * kBlkBytes, std::byte{0}); + for (size_t n = 0; n < N; ++n) { + for (size_t blk = 0; blk < BlockCountK; ++blk) { + uint8_t blk_weights[kBlkLen]; + for (size_t kk = 0; kk < kBlkLen; ++kk) { + blk_weights[kk] = BWeights[n * K + blk * kBlkLen + kk]; } + PackSourceBlock_BlkLen64(blk_weights, + QuantBData.data() + (n * BlockCountK + blk) * kBlkBytes); } + } - std::vector QuantBScale(N * BlockCountK); - for (auto& v : QuantBScale) v = s_dist(rng); - - // Per-block zero points (W2: each ZP in [0,3]). The reference path - // uses the raw [N * BlockCountK] uint8 vector; MLAS gets the standard - // ONNX-packed byte stream produced by PackW2ZeroPoints. - std::vector BZeroPoints; - std::vector BZeroPointsPacked; - const uint8_t* BZeroPointsRef = nullptr; - const std::byte* BZeroPointsMlas = nullptr; - if (WithZeroPoints) { - BZeroPoints.resize(N * BlockCountK); - for (auto& v : BZeroPoints) v = static_cast(w_dist(rng)); - BZeroPointsRef = BZeroPoints.data(); - BZeroPointsPacked = PackW2ZeroPoints(N, BlockCountK, BZeroPoints); - BZeroPointsMlas = BZeroPointsPacked.data(); - } + std::vector QuantBScale(N * BlockCountK); + for (auto& v : QuantBScale) v = s_dist(rng); + + std::vector BZeroPoints; + std::vector BZeroPointsPacked; + const uint8_t* BZeroPointsRef = nullptr; + const std::byte* BZeroPointsMlas = nullptr; + if (WithZeroPoints) { + BZeroPoints.resize(N * BlockCountK); + for (auto& v : BZeroPoints) v = static_cast(w_dist(rng)); + BZeroPointsRef = BZeroPoints.data(); + BZeroPointsPacked = PackW2ZeroPoints(N, BlockCountK, BZeroPoints); + BZeroPointsMlas = BZeroPointsPacked.data(); + } - std::vector Bias; - const float* BiasPtr = nullptr; - if (WithBias) { - Bias.resize(N); - for (auto& v : Bias) v = a_dist(rng); - BiasPtr = Bias.data(); - } + std::vector Bias; + const float* BiasPtr = nullptr; + if (WithBias) { + Bias.resize(N); + for (auto& v : Bias) v = a_dist(rng); + BiasPtr = Bias.data(); + } - // Pack B through the public API. - const size_t PackedSize = MlasQNBitGemmPackQuantBDataSize( - N, K, kBlkBitWidth, kBlkLen, WithZeroPoints, kComputeType, nullptr); - ASSERT_GT(PackedSize, 0u); - std::vector PackedQuantB(PackedSize, std::byte{0}); - - // Mirror the matmul_nbits.cc prepack flow on AMD64: three separate - // calls, one per input (B data, scales, zero_points). Each call passes - // only its own input and nullptr for the others. The pack function - // must update the BlkSum buffer on the zero_points call by reading - // scales from the already-packed buffer. - MlasQNBitGemmPackQuantBData( - N, K, kBlkBitWidth, kBlkLen, kComputeType, - QuantBData.data(), PackedQuantB.data(), - /*QuantBScale=*/nullptr, WithZeroPoints, /*QuantBZeroPoint=*/nullptr, - nullptr, nullptr); - MlasQNBitGemmPackQuantBData( - N, K, kBlkBitWidth, kBlkLen, kComputeType, - /*QuantBData=*/nullptr, PackedQuantB.data(), - QuantBScale.data(), WithZeroPoints, /*QuantBZeroPoint=*/nullptr, - nullptr, nullptr); - if (WithZeroPoints) { - MlasQNBitGemmPackQuantBData( - N, K, kBlkBitWidth, kBlkLen, kComputeType, - /*QuantBData=*/nullptr, PackedQuantB.data(), - /*QuantBScale=*/nullptr, WithZeroPoints, BZeroPointsMlas, - nullptr, nullptr); - } + // Allocate the packed-B buffer (same total size as the production path). + const size_t PackedSize = sq2::Q2BitGemmPackQuantBDataSize_Avx512( + N, K, kBlkLen, WithZeroPoints, SQNBIT_CompInt8, nullptr); + ASSERT_GT(PackedSize, 0u) << "block-group pack size unsupported for the chosen shape"; + + std::vector PackedQuantBBuf(PackedSize, std::byte{0}); + // The W2 PackedQuantBDataStruct constructor pads BlockCountK to a multiple + // of 4 internally (see qnbitgemm.h) so the slab layout matches what the + // block-group pack helper writes regardless of whether the caller passes + // the logical or padded BlockCountK. We pass the logical value to mirror + // exactly what matmul_nbits.cc does in production. + PackedQuantBDataStruct packed_b( + PackedQuantBBuf.data(), N, BlockCountK, kBlkLen, /*QuantAUnsigned=*/false); + + // Mirror the matmul_nbits.cc prepack 3-call pattern (B, scales, ZP) so the + // pack code path is exercised exactly as the production dispatcher would. + sq2::SQ2BitGemmPackQuantBDataAndBlkSum_Scalar( + N, K, kBlkLen, SQNBIT_CompInt8, + QuantBData.data(), /*scales=*/nullptr, + WithZeroPoints, /*zp=*/nullptr, + packed_b, /*tp=*/nullptr, /*cfg=*/nullptr); + sq2::SQ2BitGemmPackQuantBDataAndBlkSum_Scalar( + N, K, kBlkLen, SQNBIT_CompInt8, + /*B=*/nullptr, QuantBScale.data(), + WithZeroPoints, /*zp=*/nullptr, + packed_b, /*tp=*/nullptr, /*cfg=*/nullptr); + if (WithZeroPoints) { + sq2::SQ2BitGemmPackQuantBDataAndBlkSum_Scalar( + N, K, kBlkLen, SQNBIT_CompInt8, + /*B=*/nullptr, /*scales=*/nullptr, + WithZeroPoints, BZeroPointsMlas, + packed_b, /*tp=*/nullptr, /*cfg=*/nullptr); + } + + // Quantize A the same way MLAS would (per-block amax/127, banker rounding). + std::vector QuantAData(M * BlockCountK * kBlkLen, int8_t{0}); + std::vector QuantAScale(M * BlockCountK, 0.0f); + QuantizeA_Reference(M, K, A.data(), QuantAData.data(), QuantAScale.data()); - const size_t WorkspaceSize = MlasQNBitGemmBatchWorkspaceSize( - M, N, K, 1, kBlkBitWidth, kBlkLen, WithZeroPoints, kComputeType, nullptr); - std::vector Workspace(std::max(WorkspaceSize, 1), std::byte{0}); - - std::vector C(M * N, 0.0f); - - MLAS_QNBIT_GEMM_DATA_PARAMS params{}; - params.A = A.data(); - params.lda = K; - params.QuantBDataWorkspace = PackedQuantB.data(); - params.PackedQuantBData = PackedQuantB.data(); - params.QuantBScale = QuantBScale.data(); - params.QuantBZeroPoint = BZeroPointsMlas; - params.Bias = BiasPtr; - params.C = C.data(); - params.ldc = N; - params.PostProcessor = nullptr; - - MlasQNBitGemmBatch(M, N, K, 1, kBlkBitWidth, kBlkLen, kComputeType, - ¶ms, Workspace.data(), nullptr, nullptr); - - std::vector CRef(M * N, 0.0f); - ReferenceGemm_W2_CompInt8(M, N, K, A.data(), BWeights, QuantBScale.data(), - BZeroPointsRef, BiasPtr, CRef.data()); - - // Both paths perform the identical integer-domain dot product followed - // by the same float multiply-add chain, so the result should agree to - // a small relative tolerance driven only by float accumulation order. - const float abs_tol = 1e-4f; - const float rel_tol = 1e-4f; - for (size_t i = 0; i < M * N; ++i) { - const float diff = std::fabs(C[i] - CRef[i]); - const float bound = abs_tol + rel_tol * std::fabs(CRef[i]); - ASSERT_LE(diff, bound) - << "Mismatch at i=" << i - << " (m=" << (i / N) << ", n=" << (i % N) << ")" - << " MLAS=" << C[i] << " Ref=" << CRef[i] - << " M=" << M << " N=" << N << " K=" << K - << " WithBias=" << WithBias - << " WithZeroPoints=" << WithZeroPoints; + std::vector ABlockSum(M * BlockCountK, 0.0f); + for (size_t m = 0; m < M; ++m) { + for (size_t blk = 0; blk < BlockCountK; ++blk) { + int32_t sum = 0; + for (size_t kk = 0; kk < kBlkLen; ++kk) { + sum += static_cast( + QuantAData[m * BlockCountK * kBlkLen + blk * kBlkLen + kk]); + } + ABlockSum[m * BlockCountK + blk] = + QuantAScale[m * BlockCountK + blk] * static_cast(sum); } } -}; + + std::vector C(M * N, 0.0f); + kernel( + kBlkLen, + reinterpret_cast(QuantAData.data()), + QuantAScale.data(), + packed_b.PackedQuantBData, + packed_b.PackedQuantBScale, + /*QuantBZeroPoint=*/nullptr, + C.data(), + M, N, /*CountK=*/K, BlockCountK, + BiasPtr, + /*ldc=*/N, + ABlockSum.data(), + packed_b.QuantBBlkSum); + + std::vector CRef(M * N, 0.0f); + ReferenceGemm_W2_CompInt8(M, N, K, A.data(), BWeights, QuantBScale.data(), + BZeroPointsRef, BiasPtr, CRef.data()); + + const float abs_tol = 1e-4f; + const float rel_tol = 1e-4f; + for (size_t i = 0; i < M * N; ++i) { + const float diff = std::fabs(C[i] - CRef[i]); + const float bound = abs_tol + rel_tol * std::fabs(CRef[i]); + ASSERT_LE(diff, bound) + << "block-group " << kernel_name << " mismatch at i=" << i + << " (m=" << (i / N) << ", n=" << (i % N) << ")" + << " out=" << C[i] << " ref=" << CRef[i] + << " M=" << M << " N=" << N << " K=" << K + << " WithBias=" << WithBias + << " WithZeroPoints=" << WithZeroPoints; + } +} } // namespace // -// Public-API correctness test. On a VNNI host the platform dispatcher -// resolves to the AVX-512-VNNI W2 kernel. Skips if no W2 path is available. +// Scalar block-group test, no zero-points. Covers the same small synthetic +// shapes + customer prefill sizes used by the production W2 tests. All shapes +// have K as a multiple of (kBlkLen * kBlockGroupBlks) = 256. Customer K=384 +// is NOT a multiple of 256 so it's excluded; that shape will need a tail +// handler in a follow-up. // -TEST(MlasSq2BitTest, GemmCompInt8_BlkLen64_PublicApi) +TEST(MlasSq2BitTest, Scalar_BlkLen64) { - if (!MlasIsQNBitGemmAvailable(kBlkBitWidth, kBlkLen, kComputeType)) { - GTEST_SKIP() << "MlasQNBitGemm W2/BlkLen=64/CompInt8 not available on this host"; - } - struct Shape { size_t M, N, K; }; constexpr Shape shapes[] = { - {1, 16, 64}, - {1, 32, 128}, - {1, 64, 256}, - {4, 16, 64}, - {4, 33, 192}, - {7, 17, 128}, - {16, 64, 512}, - {32, 128, 256}, - // Customer model shapes routed through the full MlasQNBitGemmBatch - // dispatcher (threading, PerGemmQuantAWorkspace, packed-B). Both decode - // (M=1) and prefill (M=128) sizes; M=128 forces the multi-threaded - // path that the direct-call test cannot reach. - { 1, 1024, 384}, { 1, 192, 1024}, { 1, 384, 1024}, { 1, 4096, 1024}, { 1, 1024, 4096}, - {128, 1024, 384}, {128, 192, 1024}, {128, 384, 1024}, {128, 4096, 1024}, {128, 1024, 4096}, + {1, 16, 256}, + {1, 32, 256}, + {1, 64, 512}, + {4, 16, 256}, + {4, 33, 256}, + {7, 17, 256}, + {16, 64, 512}, + {32, 128, 256}, + // Customer prefill (only the K values that are multiples of 256). + { 1, 1024, 1024}, { 1, 192, 1024}, { 1, 384, 1024}, { 1, 4096, 1024}, { 1, 1024, 4096}, + {128, 1024, 1024}, {128, 192, 1024}, {128, 384, 1024}, {128, 4096, 1024}, {128, 1024, 4096}, }; for (uint32_t seed : {0xC0FFEEu, 0xBADC0DEu}) { for (const Shape& s : shapes) { for (bool bias : {false, true}) { - MlasSQ2BitGemmTest::Run(s.M, s.N, s.K, bias, seed + (bias ? 1u : 0u)); + RunW2Case(s.M, s.N, s.K, bias, seed + (bias ? 1u : 0u), + /*WithZeroPoints=*/false, + sq2::SQ2BitGemmKernel_BlkSum_CompInt8_Scalar, + "scalar"); } } } } // -// Same coverage as GemmCompInt8_BlkLen64_PublicApi but with per-block -// non-default zero points (random ZP in [0, 3] per block). This is the -// configuration the customer model uses: 4-input MatMulNBits nodes with an -// explicit zero_points initializer. +// Same coverage with per-block non-default zero points. // -TEST(MlasSq2BitTest, GemmCompInt8_BlkLen64_PublicApi_WithZeroPoints) +TEST(MlasSq2BitTest, Scalar_BlkLen64_WithZeroPoints) { - if (!MlasIsQNBitGemmAvailable(kBlkBitWidth, kBlkLen, kComputeType)) { - GTEST_SKIP() << "MlasQNBitGemm W2/BlkLen=64/CompInt8 not available on this host"; - } - struct Shape { size_t M, N, K; }; constexpr Shape shapes[] = { - {1, 16, 64}, - {1, 32, 128}, - {1, 64, 256}, - {4, 16, 64}, - {4, 33, 192}, - {7, 17, 128}, - {16, 64, 512}, - {32, 128, 256}, - // Customer model shapes (BlkLen=64, asymmetric ZP). - { 1, 1024, 384}, { 1, 192, 1024}, { 1, 384, 1024}, { 1, 4096, 1024}, { 1, 1024, 4096}, - {128, 1024, 384}, {128, 192, 1024}, {128, 384, 1024}, {128, 4096, 1024}, {128, 1024, 4096}, + {1, 16, 256}, + {1, 32, 256}, + {1, 64, 512}, + {4, 16, 256}, + {4, 33, 256}, + {7, 17, 256}, + {16, 64, 512}, + {32, 128, 256}, + { 1, 1024, 1024}, { 1, 192, 1024}, { 1, 384, 1024}, { 1, 4096, 1024}, { 1, 1024, 4096}, + {128, 1024, 1024}, {128, 192, 1024}, {128, 384, 1024}, {128, 4096, 1024}, {128, 1024, 4096}, }; for (uint32_t seed : {0xC0FFEEu, 0xBADC0DEu}) { for (const Shape& s : shapes) { for (bool bias : {false, true}) { - MlasSQ2BitGemmTest::Run(s.M, s.N, s.K, bias, seed + (bias ? 1u : 0u), - /*WithZeroPoints=*/true); + RunW2Case(s.M, s.N, s.K, bias, seed + (bias ? 1u : 0u), + /*WithZeroPoints=*/true, + sq2::SQ2BitGemmKernel_BlkSum_CompInt8_Scalar, + "scalar"); } } } } // -// ----------------------------------------------------------------------------- -// W2-v1 direct-call kernel tests (removed). +// SIMD block-group shape coverage. Phase-3+K-tail kernel requires: +// * BlkLen == 64 +// * CountN a multiple of kNCols4 (=4) +// CountM and BlockCountK have NO alignment requirements: +// - R2xC4 handles the M-aligned head; a single R1xC4 picks up the optional +// trailing odd row. CountM == 1 dispatches directly to R1xC4. +// - The K-loop iterates `BlockCountK / 4` full block-groups plus a partial +// "tail group" of 1-3 trailing K-blocks. The pack helpers zero-pad the +// trailing slots so they contribute 0 to the dot product; the tile loads +// only valid A blocks (zero ZMM for missing ones) to avoid OOB. +// +// Customer prefill shapes (M in {1, 128}, N in {192, 384, 1024, 4096}, K in +// {1024, 4096}) are all covered. M=3 and M=5 exercise the M-tail path. +// +// K-tail handler (BlockCountK not a multiple of kBlockGroupBlks=4): the +// pack helper zero-pads the trailing 1-3 K-block slots; the SIMD K-loop +// processes them via the 4-block accumulator with zero ZMM for the missing +// A blocks. GroupStride uses BlockCountKPadded so N-group advances land on +// the right packed-B address regardless of K % 4. Customer K=384 and the +// synthetic K=320, K=448 shapes exercise this path. +// +constexpr struct { size_t M, N, K; } kSimdShapes[] = { + {1, 16, 256}, // R1 only + {1, 192, 1024}, // R1 only, customer N + {1, 1024, 4096}, // R1 only, customer N + {2, 16, 256}, + {2, 32, 256}, + {2, 64, 512}, + {3, 16, 256}, // R2 head (1 pair) + R1 tail + {3, 384, 1024}, + {4, 16, 256}, + {4, 32, 256}, + {5, 64, 512}, // R2 head (2 pairs) + R1 tail + {16, 64, 512}, + {32, 128, 256}, + // Customer prefill (M=128) at all (K, N) pairs, including K=384. + {128, 1024, 384}, // K-tail: BlockCountK=6, 1 full group + tail of 2 blocks + {128, 1024, 1024}, {128, 192, 1024}, {128, 384, 1024}, + {128, 4096, 1024}, {128, 1024, 4096}, + // Customer decode (M=1) at K=384 (the case the K%4 gate previously blocked). + { 1, 1024, 384}, + // Synthetic K-tail stress shapes covering all (TailBlocks in {1, 2, 3}). + { 2, 16, 320}, // tail=1 + { 4, 16, 320}, + {128, 1024, 320}, + { 2, 16, 448}, // tail=3 + { 4, 16, 448}, + {128, 1024, 448}, + // N-tail stress (CountN % 4 != 0). The R2/R1 main tiles handle the + // NMain = floor(CountN/4)*4 cols; the per-1-col tail tile picks up + // the trailing 1-3 cols against the column-major tail region of the + // packed buffer. NMain = 0 cases (N in {1,2,3}) exercise the tail + // tile in isolation. + { 1, 1, 256}, // NMain=0, NTail=1, single-column decode + { 1, 3, 256}, // NMain=0, NTail=3 + { 4, 3, 256}, // NMain=0, NTail=3, R2+R1 head still empty + { 1, 17, 256}, // NMain=16, NTail=1, decode + { 4, 17, 256}, + {128, 17, 256}, + { 1, 33, 256}, // NMain=32, NTail=1 + { 4, 33, 256}, // exact shape that failed the dispatch swap + {128, 33, 256}, + { 1, 18, 256}, // NMain=16, NTail=2 + { 4, 18, 256}, + {128, 19, 256}, // NMain=16, NTail=3 + // N-tail combined with K-tail (the most generic case). + { 1, 17, 384}, + { 4, 33, 384}, + {128, 19, 448}, +}; + +// +// AVX-512BW (non-VNNI) SIMD block-group kernel. // -// Earlier revisions of this file exercised the AVX-512 / AVX-512-VNNI W2-v1 -// kernels (`sq2bit_avx512::SQ2BitGemmKernel_BlkSum_CompInt8_Avx512[Vnni]_TestEntry`) -// directly via test-entry forwarders, bypassing the platform dispatcher. That -// arrangement made sense when the production dispatch used the W2-v1 packed-B -// layout: the public pack API and the direct kernel call shared a buffer ABI. +TEST(MlasSq2BitTest, BlkLen64_Avx512) +{ + if (!GetMlasPlatform().Avx512Supported_) { + GTEST_SKIP() << "AVX-512BW (DQ/VL) not available on this host"; + } + for (uint32_t seed : {0xC0FFEEu, 0xBADC0DEu}) { + for (const auto& s : kSimdShapes) { + for (bool bias : {false, true}) { + RunW2Case(s.M, s.N, s.K, bias, seed + (bias ? 1u : 0u), + /*WithZeroPoints=*/false, + sq2::SQ2BitGemmKernel_BlkSum_CompInt8_Avx512_TestEntry, + "AVX-512BW"); + } + } + } +} + +TEST(MlasSq2BitTest, BlkLen64_Avx512_WithZeroPoints) +{ + if (!GetMlasPlatform().Avx512Supported_) { + GTEST_SKIP() << "AVX-512BW (DQ/VL) not available on this host"; + } + for (uint32_t seed : {0xC0FFEEu, 0xBADC0DEu}) { + for (const auto& s : kSimdShapes) { + for (bool bias : {false, true}) { + RunW2Case(s.M, s.N, s.K, bias, seed + (bias ? 1u : 0u), + /*WithZeroPoints=*/true, + sq2::SQ2BitGemmKernel_BlkSum_CompInt8_Avx512_TestEntry, + "AVX-512BW"); + } + } + } +} + // -// The production dispatch now routes through the W2-v2 super-block kernel -// (`sq2bit_avx512_super::*`) which uses a different packed-B layout (4-N-col -// grouped + super-block K stride; see PackedQuantBOffsetBytes_W2_SuperBlock). -// Driving the W2-v1 kernel with a W2-v2 pack would compare apples to oranges -// and produce spurious failures, so the W2-v1 direct-call harness has been -// retired. W2-v1 sources are kept in the build for now as a fallback while -// customers validate W2-v2 in production. +// AVX-512-VNNI SIMD block-group kernel. Gated on the platform having selected +// the VNNI dispatch table (the SIMD path uses `_mm512_dpbusd_epi32`). // -// Coverage is preserved by: -// * `GemmCompInt8_BlkLen64_PublicApi[_WithZeroPoints]` -- end-to-end through -// `MlasQNBitGemmBatch`; this is what production code paths invoke. -// * `SuperBlock*` (test_sqnbitgemm_2bit_superblock.cpp) -- direct-call -// coverage of the new default W2-v2 kernel using the matching pack helper. -// ----------------------------------------------------------------------------- +TEST(MlasSq2BitTest, BlkLen64_Avx512Vnni) +{ + if (GetMlasPlatform().QNBitGemmDispatch != &MlasSQNBitGemmDispatchAvx512vnni) { + GTEST_SKIP() << "AVX-512-VNNI not selected as the active dispatch on this host"; + } + for (uint32_t seed : {0xC0FFEEu, 0xBADC0DEu}) { + for (const auto& s : kSimdShapes) { + for (bool bias : {false, true}) { + RunW2Case(s.M, s.N, s.K, bias, seed + (bias ? 1u : 0u), + /*WithZeroPoints=*/false, + sq2::SQ2BitGemmKernel_BlkSum_CompInt8_Avx512Vnni_TestEntry, + "AVX-512-VNNI"); + } + } + } +} + +TEST(MlasSq2BitTest, BlkLen64_Avx512Vnni_WithZeroPoints) +{ + if (GetMlasPlatform().QNBitGemmDispatch != &MlasSQNBitGemmDispatchAvx512vnni) { + GTEST_SKIP() << "AVX-512-VNNI not selected as the active dispatch on this host"; + } + for (uint32_t seed : {0xC0FFEEu, 0xBADC0DEu}) { + for (const auto& s : kSimdShapes) { + for (bool bias : {false, true}) { + RunW2Case(s.M, s.N, s.K, bias, seed + (bias ? 1u : 0u), + /*WithZeroPoints=*/true, + sq2::SQ2BitGemmKernel_BlkSum_CompInt8_Avx512Vnni_TestEntry, + "AVX-512-VNNI"); + } + } + } +} diff --git a/onnxruntime/test/mlas/unittest/test_sqnbitgemm_2bit_superblock.cpp b/onnxruntime/test/mlas/unittest/test_sqnbitgemm_2bit_superblock.cpp deleted file mode 100644 index 5117149f00548..0000000000000 --- a/onnxruntime/test/mlas/unittest/test_sqnbitgemm_2bit_superblock.cpp +++ /dev/null @@ -1,529 +0,0 @@ -/*++ - -Copyright (c) Microsoft Corporation. All rights reserved. - -Licensed under the MIT License. - -Module Name: - - test_sqnbitgemm_2bit_superblock.cpp - -Abstract: - - Unit tests for the EXPERIMENTAL super-block W2 path - (sqnbitgemm_kernel_avx512_2bit_superblock.{h,cpp}). - - Phase 2 coverage: end-to-end pack + scalar GEMM correctness against the - same integer-domain reference used by the production W2 tests. These - tests do NOT exercise any SIMD path -- they validate the layout, the - pack 3-call sequence, and the scalar oracle kernel that will back the - Phase 3 SIMD work. - - Tests deliberately use the same shapes as the production W2 tests so a - side-by-side comparison is straightforward. - ---*/ - -#include "gtest/gtest.h" - -#include -#include -#include -#include -#include - -#include "core/mlas/inc/mlas_qnbit.h" -#include "core/mlas/lib/qnbitgemm.h" -#include "core/mlas/lib/mlasi.h" -#include "core/mlas/lib/sqnbitgemm_kernel_avx512_2bit.h" -#include "core/mlas/lib/sqnbitgemm_kernel_avx512_2bit_superblock.h" - -namespace { - -namespace sq2 = onnxruntime::mlas::sq2bit_avx512; -namespace sq2sb = onnxruntime::mlas::sq2bit_avx512_super; - -constexpr size_t kBlkLen = sq2::kBlkLen; // 64 -constexpr size_t kBlkBytes = sq2::kBlkBytes; // 16 -constexpr size_t kSuperBlockBlks = sq2::kSuperBlockBlks; // 4 - -// Standard ONNX 2-bit source packing (1 byte = 4 weights). -void -PackSourceBlock_BlkLen64(const uint8_t weights[kBlkLen], std::byte* src_out) -{ - for (size_t i = 0; i < kBlkBytes; ++i) { - const uint8_t v0 = weights[4 * i + 0] & 0x03u; - const uint8_t v1 = weights[4 * i + 1] & 0x03u; - const uint8_t v2 = weights[4 * i + 2] & 0x03u; - const uint8_t v3 = weights[4 * i + 3] & 0x03u; - src_out[i] = static_cast( - static_cast(v0 | (v1 << 2) | (v2 << 4) | (v3 << 6)) - ); - } -} - -// Bit-exact mirror of MlasQNBitGemm's per-block int8 quantizer (amax/127, -// round-half-to-even via std::nearbyint, scale_recip = 127/amax). -void -QuantizeA_Reference(size_t M, size_t K, const float* A, - int8_t* QuantAData, float* QuantAScale) -{ - const size_t BlockCountK = (K + kBlkLen - 1) / kBlkLen; - for (size_t m = 0; m < M; ++m) { - for (size_t k = 0, k_blk = 0; k < K; k += kBlkLen, ++k_blk) { - const size_t local_len = std::min(K - k, kBlkLen); - float amax = 0.0f; - for (size_t kk = 0; kk < local_len; ++kk) { - amax = std::max(amax, std::fabs(A[m * K + k + kk])); - } - constexpr float range_max = 127.0f; - const float scale = amax / range_max; - const float scale_recip = amax != 0.0f ? range_max / amax : 0.0f; - QuantAScale[m * BlockCountK + k_blk] = scale; - for (size_t kk = 0; kk < kBlkLen; ++kk) { - const float a = (kk < local_len) ? A[m * K + k + kk] : 0.0f; - const float q = std::nearbyint(a * scale_recip); - QuantAData[m * BlockCountK * kBlkLen + k + kk] = - static_cast(std::clamp(q, -127.0f, 127.0f)); - } - } - } -} - -// -// Integer-domain GEMM oracle: bit-exact match to the math the MLAS W2 path -// performs (kernel int8 GEMM + SGEMM zero-point correction collapsed into a -// single direct dot of (qa * (qb - zp))). -// -void -ReferenceGemm_W2_CompInt8(size_t M, size_t N, size_t K, - const float* A, - const std::vector& BWeights, // [N * K] in [0, 3] - const float* QuantBScale, // [N * BlockCountK] - const uint8_t* BZeroPoints, // [N * BlockCountK] or nullptr - const float* Bias, - float* C) -{ - const size_t BlockCountK = (K + kBlkLen - 1) / kBlkLen; - std::vector QuantAData(M * BlockCountK * kBlkLen, int8_t{0}); - std::vector QuantAScale(M * BlockCountK, 0.0f); - QuantizeA_Reference(M, K, A, QuantAData.data(), QuantAScale.data()); - - for (size_t m = 0; m < M; ++m) { - for (size_t n = 0; n < N; ++n) { - float acc = (Bias != nullptr) ? Bias[n] : 0.0f; - for (size_t k = 0, blk = 0; k < K; k += kBlkLen, ++blk) { - const size_t local_len = std::min(K - k, kBlkLen); - const float a_scale = QuantAScale[m * BlockCountK + blk]; - const float b_scale = QuantBScale[n * BlockCountK + blk]; - const int32_t zp = BZeroPoints != nullptr - ? static_cast(BZeroPoints[n * BlockCountK + blk]) - : static_cast(sq2::kDefaultSymmetricZeroPoint2Bit); - int32_t dot = 0; - for (size_t kk = 0; kk < local_len; ++kk) { - const int8_t qa = QuantAData[m * BlockCountK * kBlkLen + k + kk]; - const int32_t qb = - static_cast(BWeights[n * K + k + kk]) - zp; - dot += static_cast(qa) * qb; - } - acc += static_cast(dot) * a_scale * b_scale; - } - C[m * N + n] = acc; - } - } -} - -// -// Pack per-block W2 zero points into the standard ONNX byte stream -// (4 zp per byte along K, row-major in N). -// -std::vector -PackW2ZeroPoints(size_t N, size_t BlockCountK, const std::vector& BZeroPoints) -{ - const size_t ZPCountK = (BlockCountK + 3) / 4; - std::vector packed(N * ZPCountK, std::byte{0}); - for (size_t n = 0; n < N; ++n) { - for (size_t blk = 0; blk < BlockCountK; ++blk) { - const uint8_t zp = BZeroPoints[n * BlockCountK + blk] & 0x03u; - const size_t byte_idx = n * ZPCountK + (blk / 4); - const size_t bit_off = (blk % 4) * 2; - packed[byte_idx] = static_cast( - static_cast(packed[byte_idx]) | (zp << bit_off)); - } - } - return packed; -} - -// -// Test harness that drives the super-block path directly (without going through -// MlasQNBitGemmBatch). Builds the same packed buffer the dispatcher would -// construct, runs the chosen super-block kernel (scalar / AVX-512BW / VNNI), -// and compares to ReferenceGemm_W2_CompInt8. -// -// `KernelFn` matches the SQ4BitGemmKernel_BlkSum_CompInt8_Fn signature, which -// every super-block kernel variant honors via direct-call forwarders declared -// in sqnbitgemm_kernel_avx512_2bit_superblock.h. -// -using SuperBlockKernelFn = size_t (MLASCALL*)( - size_t, const std::byte*, const float*, const std::byte*, const float*, - const std::byte*, float*, size_t, size_t, size_t, size_t, - const float*, size_t, const float*, const float*); - -void -RunSuperBlockCase(size_t M, size_t N, size_t K, bool WithBias, uint32_t seed, - bool WithZeroPoints, SuperBlockKernelFn kernel, - const char* kernel_name) -{ - const size_t BlockCountK = (K + kBlkLen - 1) / kBlkLen; - ASSERT_EQ(K % kBlkLen, 0u) << "Test K must be a multiple of BlkLen=64"; - // BlockCountK no longer required to be a multiple of kSuperBlockBlks -- - // the K-tail handler picks up the trailing 1-3 blocks. - - std::mt19937 rng(seed); - std::uniform_real_distribution a_dist(-1.0f, 1.0f); - std::uniform_int_distribution w_dist(0, 3); - std::uniform_real_distribution s_dist(0.05f, 0.5f); - - std::vector A(M * K); - for (auto& v : A) v = a_dist(rng); - - std::vector BWeights(N * K); - for (auto& v : BWeights) v = static_cast(w_dist(rng)); - - // Source-packed B (standard ONNX layout) -- the input to the pack helper. - std::vector QuantBData(N * BlockCountK * kBlkBytes, std::byte{0}); - for (size_t n = 0; n < N; ++n) { - for (size_t blk = 0; blk < BlockCountK; ++blk) { - uint8_t blk_weights[kBlkLen]; - for (size_t kk = 0; kk < kBlkLen; ++kk) { - blk_weights[kk] = BWeights[n * K + blk * kBlkLen + kk]; - } - PackSourceBlock_BlkLen64(blk_weights, - QuantBData.data() + (n * BlockCountK + blk) * kBlkBytes); - } - } - - std::vector QuantBScale(N * BlockCountK); - for (auto& v : QuantBScale) v = s_dist(rng); - - std::vector BZeroPoints; - std::vector BZeroPointsPacked; - const uint8_t* BZeroPointsRef = nullptr; - const std::byte* BZeroPointsMlas = nullptr; - if (WithZeroPoints) { - BZeroPoints.resize(N * BlockCountK); - for (auto& v : BZeroPoints) v = static_cast(w_dist(rng)); - BZeroPointsRef = BZeroPoints.data(); - BZeroPointsPacked = PackW2ZeroPoints(N, BlockCountK, BZeroPoints); - BZeroPointsMlas = BZeroPointsPacked.data(); - } - - std::vector Bias; - const float* BiasPtr = nullptr; - if (WithBias) { - Bias.resize(N); - for (auto& v : Bias) v = a_dist(rng); - BiasPtr = Bias.data(); - } - - // Allocate the packed-B buffer (same total size as the production path). - const size_t PackedSize = sq2sb::Q2BitGemmPackQuantBDataSize_SuperBlock( - N, K, kBlkLen, WithZeroPoints, SQNBIT_CompInt8, nullptr); - ASSERT_GT(PackedSize, 0u) << "Super-block pack size unsupported for the chosen shape"; - - std::vector PackedQuantBBuf(PackedSize, std::byte{0}); - // The W2 PackedQuantBDataStruct constructor pads BlockCountK to a multiple - // of 4 internally (see qnbitgemm.h) so the slab layout matches what the - // super-block pack helper writes regardless of whether the caller passes - // the logical or padded BlockCountK. We pass the logical value to mirror - // exactly what matmul_nbits.cc does in production. - PackedQuantBDataStruct packed_b( - PackedQuantBBuf.data(), N, BlockCountK, kBlkLen, /*QuantAUnsigned=*/false); - - // Mirror the matmul_nbits.cc prepack 3-call pattern (B, scales, ZP) so the - // pack code path is exercised exactly as the production dispatcher would. - sq2sb::SQ2BitGemmPackQuantBDataAndBlkSum_SuperBlockScalar( - N, K, kBlkLen, SQNBIT_CompInt8, - QuantBData.data(), /*scales=*/nullptr, - WithZeroPoints, /*zp=*/nullptr, - packed_b, /*tp=*/nullptr, /*cfg=*/nullptr); - sq2sb::SQ2BitGemmPackQuantBDataAndBlkSum_SuperBlockScalar( - N, K, kBlkLen, SQNBIT_CompInt8, - /*B=*/nullptr, QuantBScale.data(), - WithZeroPoints, /*zp=*/nullptr, - packed_b, /*tp=*/nullptr, /*cfg=*/nullptr); - if (WithZeroPoints) { - sq2sb::SQ2BitGemmPackQuantBDataAndBlkSum_SuperBlockScalar( - N, K, kBlkLen, SQNBIT_CompInt8, - /*B=*/nullptr, /*scales=*/nullptr, - WithZeroPoints, BZeroPointsMlas, - packed_b, /*tp=*/nullptr, /*cfg=*/nullptr); - } - - // Quantize A the same way MLAS would (per-block amax/127, banker rounding). - std::vector QuantAData(M * BlockCountK * kBlkLen, int8_t{0}); - std::vector QuantAScale(M * BlockCountK, 0.0f); - QuantizeA_Reference(M, K, A.data(), QuantAData.data(), QuantAScale.data()); - - std::vector ABlockSum(M * BlockCountK, 0.0f); - for (size_t m = 0; m < M; ++m) { - for (size_t blk = 0; blk < BlockCountK; ++blk) { - int32_t sum = 0; - for (size_t kk = 0; kk < kBlkLen; ++kk) { - sum += static_cast( - QuantAData[m * BlockCountK * kBlkLen + blk * kBlkLen + kk]); - } - ABlockSum[m * BlockCountK + blk] = - QuantAScale[m * BlockCountK + blk] * static_cast(sum); - } - } - - std::vector C(M * N, 0.0f); - kernel( - kBlkLen, - reinterpret_cast(QuantAData.data()), - QuantAScale.data(), - packed_b.PackedQuantBData, - packed_b.PackedQuantBScale, - /*QuantBZeroPoint=*/nullptr, - C.data(), - M, N, /*CountK=*/K, BlockCountK, - BiasPtr, - /*ldc=*/N, - ABlockSum.data(), - packed_b.QuantBBlkSum); - - std::vector CRef(M * N, 0.0f); - ReferenceGemm_W2_CompInt8(M, N, K, A.data(), BWeights, QuantBScale.data(), - BZeroPointsRef, BiasPtr, CRef.data()); - - const float abs_tol = 1e-4f; - const float rel_tol = 1e-4f; - for (size_t i = 0; i < M * N; ++i) { - const float diff = std::fabs(C[i] - CRef[i]); - const float bound = abs_tol + rel_tol * std::fabs(CRef[i]); - ASSERT_LE(diff, bound) - << "Super-block " << kernel_name << " mismatch at i=" << i - << " (m=" << (i / N) << ", n=" << (i % N) << ")" - << " out=" << C[i] << " ref=" << CRef[i] - << " M=" << M << " N=" << N << " K=" << K - << " WithBias=" << WithBias - << " WithZeroPoints=" << WithZeroPoints; - } -} - -} // namespace - -// -// Scalar super-block test, no zero-points. Covers the same small synthetic -// shapes + customer prefill sizes used by the production W2 tests. All shapes -// have K as a multiple of (kBlkLen * kSuperBlockBlks) = 256. Customer K=384 -// is NOT a multiple of 256 so it's excluded; that shape will need a tail -// handler in a follow-up. -// -TEST(MlasSq2BitTest, SuperBlockScalar_BlkLen64) -{ - struct Shape { size_t M, N, K; }; - constexpr Shape shapes[] = { - {1, 16, 256}, - {1, 32, 256}, - {1, 64, 512}, - {4, 16, 256}, - {4, 33, 256}, - {7, 17, 256}, - {16, 64, 512}, - {32, 128, 256}, - // Customer prefill (only the K values that are multiples of 256). - { 1, 1024, 1024}, { 1, 192, 1024}, { 1, 384, 1024}, { 1, 4096, 1024}, { 1, 1024, 4096}, - {128, 1024, 1024}, {128, 192, 1024}, {128, 384, 1024}, {128, 4096, 1024}, {128, 1024, 4096}, - }; - - for (uint32_t seed : {0xC0FFEEu, 0xBADC0DEu}) { - for (const Shape& s : shapes) { - for (bool bias : {false, true}) { - RunSuperBlockCase(s.M, s.N, s.K, bias, seed + (bias ? 1u : 0u), - /*WithZeroPoints=*/false, - sq2sb::SQ2BitGemmKernel_BlkSum_CompInt8_SuperBlockScalar, - "scalar"); - } - } - } -} - -// -// Same coverage with per-block non-default zero points. -// -TEST(MlasSq2BitTest, SuperBlockScalar_BlkLen64_WithZeroPoints) -{ - struct Shape { size_t M, N, K; }; - constexpr Shape shapes[] = { - {1, 16, 256}, - {1, 32, 256}, - {1, 64, 512}, - {4, 16, 256}, - {4, 33, 256}, - {7, 17, 256}, - {16, 64, 512}, - {32, 128, 256}, - { 1, 1024, 1024}, { 1, 192, 1024}, { 1, 384, 1024}, { 1, 4096, 1024}, { 1, 1024, 4096}, - {128, 1024, 1024}, {128, 192, 1024}, {128, 384, 1024}, {128, 4096, 1024}, {128, 1024, 4096}, - }; - - for (uint32_t seed : {0xC0FFEEu, 0xBADC0DEu}) { - for (const Shape& s : shapes) { - for (bool bias : {false, true}) { - RunSuperBlockCase(s.M, s.N, s.K, bias, seed + (bias ? 1u : 0u), - /*WithZeroPoints=*/true, - sq2sb::SQ2BitGemmKernel_BlkSum_CompInt8_SuperBlockScalar, - "scalar"); - } - } - } -} - -// -// SIMD super-block shape coverage. Phase-3+K-tail kernel requires: -// * BlkLen == 64 -// * CountN a multiple of kNCols4 (=4) -// CountM and BlockCountK have NO alignment requirements: -// - R2xC4 handles the M-aligned head; a single R1xC4 picks up the optional -// trailing odd row. CountM == 1 dispatches directly to R1xC4. -// - The K-loop iterates `BlockCountK / 4` full super-blocks plus a partial -// "tail super" of 1-3 trailing K-blocks. The pack helpers zero-pad the -// trailing slots so they contribute 0 to the dot product; the tile loads -// only valid A blocks (zero ZMM for missing ones) to avoid OOB. -// -// Customer prefill shapes (M in {1, 128}, N in {192, 384, 1024, 4096}, K in -// {1024, 4096}) are all covered. M=3 and M=5 exercise the M-tail path. -// -// K-tail handler (BlockCountK not a multiple of kSuperBlockBlks=4): the -// pack helper zero-pads the trailing 1-3 K-block slots; the SIMD K-loop -// processes them via the 4-block accumulator with zero ZMM for the missing -// A blocks. GroupStride uses BlockCountKPadded so N-group advances land on -// the right packed-B address regardless of K % 4. Customer K=384 and the -// synthetic K=320, K=448 shapes exercise this path. -// -constexpr struct { size_t M, N, K; } kSimdShapes[] = { - {1, 16, 256}, // R1 only - {1, 192, 1024}, // R1 only, customer N - {1, 1024, 4096}, // R1 only, customer N - {2, 16, 256}, - {2, 32, 256}, - {2, 64, 512}, - {3, 16, 256}, // R2 head (1 pair) + R1 tail - {3, 384, 1024}, - {4, 16, 256}, - {4, 32, 256}, - {5, 64, 512}, // R2 head (2 pairs) + R1 tail - {16, 64, 512}, - {32, 128, 256}, - // Customer prefill (M=128) at all (K, N) pairs, including K=384. - {128, 1024, 384}, // K-tail: BlockCountK=6, 1 full super + tail of 2 blocks - {128, 1024, 1024}, {128, 192, 1024}, {128, 384, 1024}, - {128, 4096, 1024}, {128, 1024, 4096}, - // Customer decode (M=1) at K=384 (the case the K%4 gate previously blocked). - { 1, 1024, 384}, - // Synthetic K-tail stress shapes covering all (TailBlocks in {1, 2, 3}). - { 2, 16, 320}, // tail=1 - { 4, 16, 320}, - {128, 1024, 320}, - { 2, 16, 448}, // tail=3 - { 4, 16, 448}, - {128, 1024, 448}, - // N-tail stress (CountN % 4 != 0). The R2/R1 main tiles handle the - // NMain = floor(CountN/4)*4 cols; the per-1-col tail tile picks up - // the trailing 1-3 cols against the column-major tail region of the - // packed buffer. NMain = 0 cases (N in {1,2,3}) exercise the tail - // tile in isolation. - { 1, 1, 256}, // NMain=0, NTail=1, single-column decode - { 1, 3, 256}, // NMain=0, NTail=3 - { 4, 3, 256}, // NMain=0, NTail=3, R2+R1 head still empty - { 1, 17, 256}, // NMain=16, NTail=1, decode - { 4, 17, 256}, - {128, 17, 256}, - { 1, 33, 256}, // NMain=32, NTail=1 - { 4, 33, 256}, // exact shape that failed the dispatch swap - {128, 33, 256}, - { 1, 18, 256}, // NMain=16, NTail=2 - { 4, 18, 256}, - {128, 19, 256}, // NMain=16, NTail=3 - // N-tail combined with K-tail (the most generic case). - { 1, 17, 384}, - { 4, 33, 384}, - {128, 19, 448}, -}; - -// -// AVX-512BW (non-VNNI) SIMD super-block kernel. -// -TEST(MlasSq2BitTest, SuperBlock_BlkLen64_Avx512) -{ - if (!GetMlasPlatform().Avx512Supported_) { - GTEST_SKIP() << "AVX-512BW (DQ/VL) not available on this host"; - } - for (uint32_t seed : {0xC0FFEEu, 0xBADC0DEu}) { - for (const auto& s : kSimdShapes) { - for (bool bias : {false, true}) { - RunSuperBlockCase(s.M, s.N, s.K, bias, seed + (bias ? 1u : 0u), - /*WithZeroPoints=*/false, - sq2sb::SQ2BitGemmKernel_BlkSum_CompInt8_Super_Avx512_TestEntry, - "AVX-512BW"); - } - } - } -} - -TEST(MlasSq2BitTest, SuperBlock_BlkLen64_Avx512_WithZeroPoints) -{ - if (!GetMlasPlatform().Avx512Supported_) { - GTEST_SKIP() << "AVX-512BW (DQ/VL) not available on this host"; - } - for (uint32_t seed : {0xC0FFEEu, 0xBADC0DEu}) { - for (const auto& s : kSimdShapes) { - for (bool bias : {false, true}) { - RunSuperBlockCase(s.M, s.N, s.K, bias, seed + (bias ? 1u : 0u), - /*WithZeroPoints=*/true, - sq2sb::SQ2BitGemmKernel_BlkSum_CompInt8_Super_Avx512_TestEntry, - "AVX-512BW"); - } - } - } -} - -// -// AVX-512-VNNI SIMD super-block kernel. Gated on the platform having selected -// the VNNI dispatch table (the SIMD path uses `_mm512_dpbusd_epi32`). -// -TEST(MlasSq2BitTest, SuperBlock_BlkLen64_Avx512Vnni) -{ - if (GetMlasPlatform().QNBitGemmDispatch != &MlasSQNBitGemmDispatchAvx512vnni) { - GTEST_SKIP() << "AVX-512-VNNI not selected as the active dispatch on this host"; - } - for (uint32_t seed : {0xC0FFEEu, 0xBADC0DEu}) { - for (const auto& s : kSimdShapes) { - for (bool bias : {false, true}) { - RunSuperBlockCase(s.M, s.N, s.K, bias, seed + (bias ? 1u : 0u), - /*WithZeroPoints=*/false, - sq2sb::SQ2BitGemmKernel_BlkSum_CompInt8_Super_Avx512Vnni_TestEntry, - "AVX-512-VNNI"); - } - } - } -} - -TEST(MlasSq2BitTest, SuperBlock_BlkLen64_Avx512Vnni_WithZeroPoints) -{ - if (GetMlasPlatform().QNBitGemmDispatch != &MlasSQNBitGemmDispatchAvx512vnni) { - GTEST_SKIP() << "AVX-512-VNNI not selected as the active dispatch on this host"; - } - for (uint32_t seed : {0xC0FFEEu, 0xBADC0DEu}) { - for (const auto& s : kSimdShapes) { - for (bool bias : {false, true}) { - RunSuperBlockCase(s.M, s.N, s.K, bias, seed + (bias ? 1u : 0u), - /*WithZeroPoints=*/true, - sq2sb::SQ2BitGemmKernel_BlkSum_CompInt8_Super_Avx512Vnni_TestEntry, - "AVX-512-VNNI"); - } - } - } -} From b8ef59d0df689b623bafb2f5002701b353c792ae Mon Sep 17 00:00:00 2001 From: Hari Seshadri Date: Thu, 11 Jun 2026 19:14:21 -0700 Subject: [PATCH 09/17] Support BlkLen 32 and 128 --- .../mlas/lib/sqnbitgemm_kernel_avx512.cpp | 12 + .../lib/sqnbitgemm_kernel_avx512_2bit.cpp | 450 ++++++++- .../mlas/lib/sqnbitgemm_kernel_avx512_2bit.h | 149 +++ .../sqnbitgemm_kernel_avx512_2bit_blklen128.h | 859 ++++++++++++++++++ .../sqnbitgemm_kernel_avx512_2bit_blklen32.h | 696 ++++++++++++++ .../mlas/lib/sqnbitgemm_kernel_avx512vnni.cpp | 12 + .../test/contrib_ops/matmul_2bits_test.cc | 101 ++ .../test/mlas/bench/bench_qnbitgemm.cpp | 6 +- .../mlas/unittest/test_sqnbitgemm_2bit.cpp | 311 +++++++ .../unittest/test_sqnbitgemm_2bit_gemm.cpp | 735 +++++++++++++++ 10 files changed, 3325 insertions(+), 6 deletions(-) create mode 100644 onnxruntime/core/mlas/lib/sqnbitgemm_kernel_avx512_2bit_blklen128.h create mode 100644 onnxruntime/core/mlas/lib/sqnbitgemm_kernel_avx512_2bit_blklen32.h diff --git a/onnxruntime/core/mlas/lib/sqnbitgemm_kernel_avx512.cpp b/onnxruntime/core/mlas/lib/sqnbitgemm_kernel_avx512.cpp index c979841e20d24..0c946814cda83 100644 --- a/onnxruntime/core/mlas/lib/sqnbitgemm_kernel_avx512.cpp +++ b/onnxruntime/core/mlas/lib/sqnbitgemm_kernel_avx512.cpp @@ -28,6 +28,8 @@ Module Name: #include "sqnbitgemm_kernel_avx512_int8_blklen128.h" #include "sqnbitgemm_kernel_avx512_2bit.h" #include "sqnbitgemm_kernel_avx512_2bit_blklen64.h" +#include "sqnbitgemm_kernel_avx512_2bit_blklen128.h" +#include "sqnbitgemm_kernel_avx512_2bit_blklen32.h" // // SQNBIT_CompFp32 kernel implementation. @@ -500,6 +502,16 @@ SQ2BitGemmKernel_BlkSum_CompInt8_Avx512_TestEntry( const float* ABlockSum, const float* QuantBBlkSum) { + if (BlkLen == 128) { + return SQ2BitGemmKernel_BlkSum_CompInt8_BlkLen128_Avx512( + QuantA, QuantAScale, QuantBData, QuantBScale, + C, CountM, CountN, BlockCountK, Bias, ldc, ABlockSum, QuantBBlkSum); + } + if (BlkLen == 32) { + return SQ2BitGemmKernel_BlkSum_CompInt8_BlkLen32_Avx512( + QuantA, QuantAScale, QuantBData, QuantBScale, + C, CountM, CountN, BlockCountK, Bias, ldc, ABlockSum, QuantBBlkSum); + } return SQ2BitGemmKernel_BlkSum_CompInt8_Avx512( BlkLen, QuantA, QuantAScale, QuantBData, QuantBScale, QuantBZeroPoint, C, CountM, CountN, CountK, BlockCountK, Bias, ldc, ABlockSum, QuantBBlkSum); diff --git a/onnxruntime/core/mlas/lib/sqnbitgemm_kernel_avx512_2bit.cpp b/onnxruntime/core/mlas/lib/sqnbitgemm_kernel_avx512_2bit.cpp index 19fc022c62ab9..cac3149cf294c 100644 --- a/onnxruntime/core/mlas/lib/sqnbitgemm_kernel_avx512_2bit.cpp +++ b/onnxruntime/core/mlas/lib/sqnbitgemm_kernel_avx512_2bit.cpp @@ -72,7 +72,10 @@ Q2BitGemmPackQuantBDataSize_Avx512( const MLAS_BACKEND_KERNEL_SELECTOR_CONFIG* /* BackendKernelSelectorConfig */ ) { - if (BlkLen != kBlkLen || ComputeType != SQNBIT_CompInt8) { + if (ComputeType != SQNBIT_CompInt8) { + return 0; + } + if (BlkLen != kBlkLen && BlkLen != kBlkLen128 && BlkLen != kBlkLen32) { return 0; } const size_t BlockCountK = MlasDivRoundup(K, BlkLen); @@ -82,6 +85,12 @@ Q2BitGemmPackQuantBDataSize_Avx512( const size_t BlockCountKPadded = MlasDivRoundup(BlockCountK, kBlockGroupBlks) * kBlockGroupBlks; + // Per-block packed-data byte count: BlkLen / 4 (= kBlkBytes for BlkLen=64, + // = kBlkBytes128 for BlkLen=128). Use the runtime BlkLen so the same + // pack-size helper covers both kernel variants -- the BlkLen=64 storage + // total when BlkLen=64 is bit-identical to the original computation. + const size_t BlkBytes = BlkLen / kWeightsPerByte; + // Use BlockCountKPadded for BlkSum sizing too. The actual SGEMM-correction // step only reads LOGICAL BlockCountK entries, but PackedQuantBDataStruct // is constructed by the caller with a single BlockCountK value that @@ -90,7 +99,7 @@ Q2BitGemmPackQuantBDataSize_Avx512( // the struct's BlkSum pointer would land inside the packed-B region // (because the caller's struct uses one BlockCountK consistently). The // extra storage from padding the BlkSum is ~16 floats per N -- trivial. - size_t PackedQuantBDataSize = N * BlockCountKPadded * kBlkBytes; + size_t PackedQuantBDataSize = N * BlockCountKPadded * BlkBytes; const size_t ScaleSize = N * BlockCountKPadded * sizeof(float); size_t BlkSumSize = MlasDivRoundup(N, 16) * BlockCountKPadded * 16 * sizeof(float); @@ -124,6 +133,27 @@ Q2BitGemmPackQuantBDataSize_Avx512( // We write scales when scales arrive, then re-derive BlkSum whenever either // scales or zero-points arrive, reading scales from the already-packed buffer. // + +// Forward declaration: BlkLen=128 variant lives further down in this TU. +static void +SQ2BitGemmPackQuantBDataAndBlkSum_BlkLen128_Scalar( + size_t N, size_t K, + const std::byte* QuantBDataBegin, + const float* QuantBScaleBegin, + const std::byte* QuantBZPBegin, + PackedQuantBDataStruct& PackedQuantB, + MLAS_THREADPOOL* ThreadPool); + +// Forward declaration: BlkLen=32 variant lives further down in this TU. +static void +SQ2BitGemmPackQuantBDataAndBlkSum_BlkLen32_Scalar( + size_t N, size_t K, + const std::byte* QuantBDataBegin, + const float* QuantBScaleBegin, + const std::byte* QuantBZPBegin, + PackedQuantBDataStruct& PackedQuantB, + MLAS_THREADPOOL* ThreadPool); + void MLASCALL SQ2BitGemmPackQuantBDataAndBlkSum_Scalar( size_t N, @@ -139,6 +169,20 @@ SQ2BitGemmPackQuantBDataAndBlkSum_Scalar( const MLAS_BACKEND_KERNEL_SELECTOR_CONFIG* /* BackendKernelSelectorConfig */ ) { + // BlkLen=128 dispatches to its own parallel implementation below; the + // BlkLen=64 path remains exactly as before for full bit-for-bit parity. + if (BlkLen == kBlkLen128) { + SQ2BitGemmPackQuantBDataAndBlkSum_BlkLen128_Scalar( + N, K, QuantBDataBegin, QuantBScaleBegin, QuantBZPBegin, + PackedQuantB, ThreadPool); + return; + } + if (BlkLen == kBlkLen32) { + SQ2BitGemmPackQuantBDataAndBlkSum_BlkLen32_Scalar( + N, K, QuantBDataBegin, QuantBScaleBegin, QuantBZPBegin, + PackedQuantB, ThreadPool); + return; + } assert(BlkLen == kBlkLen); if (BlkLen != kBlkLen) { return; @@ -258,6 +302,39 @@ SQ2BitGemmPackQuantBDataAndBlkSum_Scalar( // SIMD path. It also lets us validate the pack layout end-to-end via the // existing MlasQNBitGemmBatch dispatch path once we wire it up. // + +// Forward declaration: BlkLen=128 variant lives further down in this TU. +static size_t +SQ2BitGemmKernel_BlkSum_CompInt8_BlkLen128_Scalar( + const std::byte* QuantA, + const float* QuantAScale, + const std::byte* QuantBData, + const float* QuantBScale, + float* C, + size_t CountM, + size_t CountN, + size_t BlockCountK, + const float* Bias, + size_t ldc, + const float* ABlockSum, + const float* QuantBBlkSum); + +// Forward declaration: BlkLen=32 variant lives further down in this TU. +static size_t +SQ2BitGemmKernel_BlkSum_CompInt8_BlkLen32_Scalar( + const std::byte* QuantA, + const float* QuantAScale, + const std::byte* QuantBData, + const float* QuantBScale, + float* C, + size_t CountM, + size_t CountN, + size_t BlockCountK, + const float* Bias, + size_t ldc, + const float* ABlockSum, + const float* QuantBBlkSum); + size_t MLASCALL SQ2BitGemmKernel_BlkSum_CompInt8_Scalar( const size_t BlkLen, @@ -277,6 +354,16 @@ SQ2BitGemmKernel_BlkSum_CompInt8_Scalar( const float* QuantBBlkSum ) { + if (BlkLen == kBlkLen128) { + return SQ2BitGemmKernel_BlkSum_CompInt8_BlkLen128_Scalar( + QuantA, QuantAScale, QuantBData, QuantBScale, + C, CountM, CountN, BlockCountK, Bias, ldc, ABlockSum, QuantBBlkSum); + } + if (BlkLen == kBlkLen32) { + return SQ2BitGemmKernel_BlkSum_CompInt8_BlkLen32_Scalar( + QuantA, QuantAScale, QuantBData, QuantBScale, + C, CountM, CountN, BlockCountK, Bias, ldc, ABlockSum, QuantBBlkSum); + } if (BlkLen != kBlkLen) { return 0; } @@ -353,6 +440,365 @@ SQ2BitGemmKernel_BlkSum_CompInt8_Scalar( return CountM; } +// ----------------------------------------------------------------------------- +// BlkLen=128 variants of pack-and-blksum and the scalar oracle. +// +// These are direct ports of the BlkLen=64 variants with three substitutions: +// * kBlkBytes -> kBlkBytes128 (32 bytes per block) +// * kBlkLen -> kBlkLen128 (128 weights per block) +// * PackBlockGroup_BlkLen64 -> PackBlockGroup_BlkLen128 +// * PackedQuantBOffsetBytes_W2 -> PackedQuantBOffsetBytes_W2_BlkLen128 +// +// PackedQuantBScaleOffset_W2 is reused as-is (scale layout is BlkLen-invariant). +// The N-major / 4-col-grouped layout, the K-tail rounding rule, and the +// SGEMM correction step (width-16 BlkSum) are identical to BlkLen=64. +// ----------------------------------------------------------------------------- + +static void +SQ2BitGemmPackQuantBDataAndBlkSum_BlkLen128_Scalar( + size_t N, + size_t K, + const std::byte* QuantBDataBegin, + const float* QuantBScaleBegin, + const std::byte* QuantBZPBegin, + PackedQuantBDataStruct& PackedQuantB, + MLAS_THREADPOOL* ThreadPool) +{ + const size_t BlockCountK = MlasDivRoundup(K, kBlkLen128); + if (BlockCountK == 0) { + return; + } + + const size_t BlockGroupCountKPadded = + MlasDivRoundup(BlockCountK, kBlockGroupBlks); + const size_t BlockCountKPadded = BlockGroupCountKPadded * kBlockGroupBlks; + const size_t NMain = (N / kNCols4) * kNCols4; + + static const std::byte kZeroBlock128[kBlkBytes128] = {}; + + // ----- B-data pack ----- + if (QuantBDataBegin != nullptr) { + std::byte* PackedQuantBData = PackedQuantB.PackedQuantBData; + const size_t Iterations = N * BlockGroupCountKPadded; + MlasTrySimpleParallel( + ThreadPool, static_cast(Iterations), + [&](ptrdiff_t tid) { + const size_t n = static_cast(tid) / BlockGroupCountKPadded; + const size_t blk_group = static_cast(tid) % BlockGroupCountKPadded; + const size_t blk0 = blk_group * kBlockGroupBlks; + + auto src_for = [&](size_t blk) -> const std::byte* { + if (blk < BlockCountK) { + return QuantBDataBegin + (n * BlockCountK + blk) * kBlkBytes128; + } + return kZeroBlock128; + }; + const std::byte* src_blk_0 = src_for(blk0 + 0); + const std::byte* src_blk_1 = src_for(blk0 + 1); + const std::byte* src_blk_2 = src_for(blk0 + 2); + const std::byte* src_blk_3 = src_for(blk0 + 3); + + const size_t dst_offset = + PackedQuantBOffsetBytes_W2_BlkLen128(n, blk_group, BlockGroupCountKPadded, NMain); + PackBlockGroup_BlkLen128(src_blk_0, src_blk_1, src_blk_2, src_blk_3, + PackedQuantBData + dst_offset); + } + ); + } + + // ----- Scales (same layout as BlkLen=64; reuse PackedQuantBScaleOffset_W2) ----- + if (QuantBScaleBegin != nullptr) { + float* PackedScales = PackedQuantB.PackedQuantBScale; + const size_t Iterations = N * BlockCountKPadded; + MlasTrySimpleParallel( + ThreadPool, static_cast(Iterations), + [&](ptrdiff_t tid) { + const size_t n = static_cast(tid) / BlockCountKPadded; + const size_t blk = static_cast(tid) % BlockCountKPadded; + const float scale = (blk < BlockCountK) + ? QuantBScaleBegin[n * BlockCountK + blk] + : 0.0f; + PackedScales[PackedQuantBScaleOffset_W2(n, blk, BlockCountKPadded, NMain)] = scale; + } + ); + } + + // ----- BlkSum (recomputed whenever scales or ZPs arrive) ----- + if (QuantBScaleBegin != nullptr || QuantBZPBegin != nullptr) { + const float* PackedScales = PackedQuantB.PackedQuantBScale; + float* BlkSum = PackedQuantB.QuantBBlkSum; + const size_t ZPCountK = MlasDivRoundup(BlockCountK, 4); + const size_t Iterations = N * BlockCountK; + MlasTrySimpleParallel( + ThreadPool, static_cast(Iterations), + [&](ptrdiff_t tid) { + const size_t n = static_cast(tid) / BlockCountK; + const size_t blk = static_cast(tid) % BlockCountK; + const float scale = + PackedScales[PackedQuantBScaleOffset_W2(n, blk, BlockCountKPadded, NMain)]; + + uint8_t zp = kDefaultSymmetricZeroPoint2Bit; + if (QuantBZPBegin != nullptr) { + const size_t zp_byte_idx = n * ZPCountK + (blk / 4); + const size_t zp_bit_off = (blk % 4) * 2; + zp = static_cast( + (static_cast(QuantBZPBegin[zp_byte_idx]) >> zp_bit_off) & 0x03u); + } + + const size_t blksum_offset = ((n / 16) * BlockCountK + blk) * 16 + (n % 16); + BlkSum[blksum_offset] = -scale * static_cast(zp); + } + ); + } +} + +static size_t +SQ2BitGemmKernel_BlkSum_CompInt8_BlkLen128_Scalar( + const std::byte* QuantA, + const float* QuantAScale, + const std::byte* QuantBData, + const float* QuantBScale, + float* C, + size_t CountM, + size_t CountN, + size_t BlockCountK, + const float* Bias, + size_t ldc, + const float* ABlockSum, + const float* QuantBBlkSum) +{ + if (BlockCountK == 0) { + return 0; + } + + const size_t BlockGroupCountKPadded = + MlasDivRoundup(BlockCountK, kBlockGroupBlks); + const size_t BlockCountKPadded = BlockGroupCountKPadded * kBlockGroupBlks; + + const size_t NMainLocal = (CountN / kNCols4) * kNCols4; + + const size_t lda = BlockCountK * kBlkLen128; // bytes per A row (int8) + const size_t lda_scale = BlockCountK; // floats per A scale row + + for (size_t m = 0; m < CountM; ++m) { + const int8_t* a_row = reinterpret_cast(QuantA + m * lda); + const float* a_scale_row = QuantAScale + m * lda_scale; + const float* a_blksum_row = ABlockSum + m * lda_scale; + float* c_row = C + m * ldc; + + for (size_t n = 0; n < CountN; ++n) { + float acc = (Bias != nullptr) ? Bias[n] : 0.0f; + + for (size_t blk = 0; blk < BlockCountK; ++blk) { + const size_t blk_group = blk / kBlockGroupBlks; + const size_t blk_in_group = blk % kBlockGroupBlks; + const size_t block_group_offset = + PackedQuantBOffsetBytes_W2_BlkLen128(n, blk_group, BlockGroupCountKPadded, NMainLocal); + const std::byte* block_group = QuantBData + block_group_offset; + + uint8_t b_unpacked[kBlkLen128]; + for (size_t i = 0; i < kBlkLen128; ++i) { + const uint8_t byte = static_cast(block_group[i]); + b_unpacked[i] = static_cast((byte >> (2 * blk_in_group)) & 0x03u); + } + + const int8_t* a_blk = a_row + blk * kBlkLen128; + int32_t dot = 0; + for (size_t i = 0; i < kBlkLen128; ++i) { + dot += static_cast(a_blk[i]) * static_cast(b_unpacked[i]); + } + + const float b_scale = + QuantBScale[PackedQuantBScaleOffset_W2(n, blk, BlockCountKPadded, NMainLocal)]; + acc += a_scale_row[blk] * b_scale * static_cast(dot); + + const size_t blksum_offset = ((n / 16) * BlockCountK + blk) * 16 + (n % 16); + acc += a_blksum_row[blk] * QuantBBlkSum[blksum_offset]; + } + + c_row[n] = acc; + } + } + + return CountM; +} + +// ----------------------------------------------------------------------------- +// BlkLen=32 variants of pack-and-blksum and the scalar oracle. Same pattern +// as the BlkLen=128 variants above (which are themselves direct ports of the +// BlkLen=64 implementations). Only the BlkLen-specific constants and the +// BlkLen=32 pack/offset helpers differ. +// ----------------------------------------------------------------------------- + +static void +SQ2BitGemmPackQuantBDataAndBlkSum_BlkLen32_Scalar( + size_t N, + size_t K, + const std::byte* QuantBDataBegin, + const float* QuantBScaleBegin, + const std::byte* QuantBZPBegin, + PackedQuantBDataStruct& PackedQuantB, + MLAS_THREADPOOL* ThreadPool) +{ + const size_t BlockCountK = MlasDivRoundup(K, kBlkLen32); + if (BlockCountK == 0) { + return; + } + + const size_t BlockGroupCountKPadded = + MlasDivRoundup(BlockCountK, kBlockGroupBlks); + const size_t BlockCountKPadded = BlockGroupCountKPadded * kBlockGroupBlks; + const size_t NMain = (N / kNCols4) * kNCols4; + + static const std::byte kZeroBlock32[kBlkBytes32] = {}; + + // ----- B-data pack ----- + if (QuantBDataBegin != nullptr) { + std::byte* PackedQuantBData = PackedQuantB.PackedQuantBData; + const size_t Iterations = N * BlockGroupCountKPadded; + MlasTrySimpleParallel( + ThreadPool, static_cast(Iterations), + [&](ptrdiff_t tid) { + const size_t n = static_cast(tid) / BlockGroupCountKPadded; + const size_t blk_group = static_cast(tid) % BlockGroupCountKPadded; + const size_t blk0 = blk_group * kBlockGroupBlks; + + auto src_for = [&](size_t blk) -> const std::byte* { + if (blk < BlockCountK) { + return QuantBDataBegin + (n * BlockCountK + blk) * kBlkBytes32; + } + return kZeroBlock32; + }; + const std::byte* src_blk_0 = src_for(blk0 + 0); + const std::byte* src_blk_1 = src_for(blk0 + 1); + const std::byte* src_blk_2 = src_for(blk0 + 2); + const std::byte* src_blk_3 = src_for(blk0 + 3); + + const size_t dst_offset = + PackedQuantBOffsetBytes_W2_BlkLen32(n, blk_group, BlockGroupCountKPadded, NMain); + PackBlockGroup_BlkLen32(src_blk_0, src_blk_1, src_blk_2, src_blk_3, + PackedQuantBData + dst_offset); + } + ); + } + + // ----- Scales (same layout as BlkLen=64; reuse PackedQuantBScaleOffset_W2) ----- + if (QuantBScaleBegin != nullptr) { + float* PackedScales = PackedQuantB.PackedQuantBScale; + const size_t Iterations = N * BlockCountKPadded; + MlasTrySimpleParallel( + ThreadPool, static_cast(Iterations), + [&](ptrdiff_t tid) { + const size_t n = static_cast(tid) / BlockCountKPadded; + const size_t blk = static_cast(tid) % BlockCountKPadded; + const float scale = (blk < BlockCountK) + ? QuantBScaleBegin[n * BlockCountK + blk] + : 0.0f; + PackedScales[PackedQuantBScaleOffset_W2(n, blk, BlockCountKPadded, NMain)] = scale; + } + ); + } + + // ----- BlkSum (recomputed whenever scales or ZPs arrive) ----- + if (QuantBScaleBegin != nullptr || QuantBZPBegin != nullptr) { + const float* PackedScales = PackedQuantB.PackedQuantBScale; + float* BlkSum = PackedQuantB.QuantBBlkSum; + const size_t ZPCountK = MlasDivRoundup(BlockCountK, 4); + const size_t Iterations = N * BlockCountK; + MlasTrySimpleParallel( + ThreadPool, static_cast(Iterations), + [&](ptrdiff_t tid) { + const size_t n = static_cast(tid) / BlockCountK; + const size_t blk = static_cast(tid) % BlockCountK; + const float scale = + PackedScales[PackedQuantBScaleOffset_W2(n, blk, BlockCountKPadded, NMain)]; + + uint8_t zp = kDefaultSymmetricZeroPoint2Bit; + if (QuantBZPBegin != nullptr) { + const size_t zp_byte_idx = n * ZPCountK + (blk / 4); + const size_t zp_bit_off = (blk % 4) * 2; + zp = static_cast( + (static_cast(QuantBZPBegin[zp_byte_idx]) >> zp_bit_off) & 0x03u); + } + + const size_t blksum_offset = ((n / 16) * BlockCountK + blk) * 16 + (n % 16); + BlkSum[blksum_offset] = -scale * static_cast(zp); + } + ); + } +} + +static size_t +SQ2BitGemmKernel_BlkSum_CompInt8_BlkLen32_Scalar( + const std::byte* QuantA, + const float* QuantAScale, + const std::byte* QuantBData, + const float* QuantBScale, + float* C, + size_t CountM, + size_t CountN, + size_t BlockCountK, + const float* Bias, + size_t ldc, + const float* ABlockSum, + const float* QuantBBlkSum) +{ + if (BlockCountK == 0) { + return 0; + } + + const size_t BlockGroupCountKPadded = + MlasDivRoundup(BlockCountK, kBlockGroupBlks); + const size_t BlockCountKPadded = BlockGroupCountKPadded * kBlockGroupBlks; + + const size_t NMainLocal = (CountN / kNCols4) * kNCols4; + + const size_t lda = BlockCountK * kBlkLen32; // bytes per A row (int8) + const size_t lda_scale = BlockCountK; // floats per A scale row + + for (size_t m = 0; m < CountM; ++m) { + const int8_t* a_row = reinterpret_cast(QuantA + m * lda); + const float* a_scale_row = QuantAScale + m * lda_scale; + const float* a_blksum_row = ABlockSum + m * lda_scale; + float* c_row = C + m * ldc; + + for (size_t n = 0; n < CountN; ++n) { + float acc = (Bias != nullptr) ? Bias[n] : 0.0f; + + for (size_t blk = 0; blk < BlockCountK; ++blk) { + const size_t blk_group = blk / kBlockGroupBlks; + const size_t blk_in_group = blk % kBlockGroupBlks; + const size_t block_group_offset = + PackedQuantBOffsetBytes_W2_BlkLen32(n, blk_group, BlockGroupCountKPadded, NMainLocal); + const std::byte* block_group = QuantBData + block_group_offset; + + uint8_t b_unpacked[kBlkLen32]; + for (size_t i = 0; i < kBlkLen32; ++i) { + const uint8_t byte = static_cast(block_group[i]); + b_unpacked[i] = static_cast((byte >> (2 * blk_in_group)) & 0x03u); + } + + const int8_t* a_blk = a_row + blk * kBlkLen32; + int32_t dot = 0; + for (size_t i = 0; i < kBlkLen32; ++i) { + dot += static_cast(a_blk[i]) * static_cast(b_unpacked[i]); + } + + const float b_scale = + QuantBScale[PackedQuantBScaleOffset_W2(n, blk, BlockCountKPadded, NMainLocal)]; + acc += a_scale_row[blk] * b_scale * static_cast(dot); + + const size_t blksum_offset = ((n / 16) * BlockCountK + blk) * 16 + (n % 16); + acc += a_blksum_row[blk] * QuantBBlkSum[blksum_offset]; + } + + c_row[n] = acc; + } + } + + return CountM; +} + } // namespace sq2bit_avx512 } // namespace mlas } // namespace onnxruntime diff --git a/onnxruntime/core/mlas/lib/sqnbitgemm_kernel_avx512_2bit.h b/onnxruntime/core/mlas/lib/sqnbitgemm_kernel_avx512_2bit.h index 34b08771802ca..152f33d3f3228 100644 --- a/onnxruntime/core/mlas/lib/sqnbitgemm_kernel_avx512_2bit.h +++ b/onnxruntime/core/mlas/lib/sqnbitgemm_kernel_avx512_2bit.h @@ -158,6 +158,155 @@ UnpackSourceBlock_BlkLen64_Reference(const std::byte* src, uint8_t out[kBlkLen]) } } +// ----------------------------------------------------------------------------- +// BlkLen=128 constants and pack helpers (parallel to the BlkLen=64 set above). +// The block-group still aggregates 4 K-blocks; only the per-block width +// (and hence the per-group byte count) changes. The N-major / 4-col-grouped +// layout structure, the K-tail rounding rule, and the SGEMM correction step +// are identical to BlkLen=64. +// ----------------------------------------------------------------------------- + +constexpr size_t kBlkLen128 = 128; +constexpr size_t kBlkBytes128 = kBlkLen128 / kWeightsPerByte; // 32 packed src bytes per block +constexpr size_t kBlockGroupBytes128 = kBlockGroupBlks * kBlkBytes128; // 128 bytes per block-group +constexpr size_t kBlockGroupWeights128 = kBlockGroupBlks * kBlkLen128; // 512 weights per block-group + +// +// Pack 4 consecutive K-blocks (4 * 32 = 128 source bytes in standard ONNX +// layout) into a 128-byte block-group. Identical bit-layout rule as the +// BlkLen=64 variant: byte b of the destination holds bits[0..1] of block_0's +// weight[b], bits[2..3] of block_1's weight[b], etc. Only the byte count +// (= kBlkLen128 = 128) differs. +// +inline void +PackBlockGroup_BlkLen128(const std::byte* src_block_0, + const std::byte* src_block_1, + const std::byte* src_block_2, + const std::byte* src_block_3, + std::byte* dst) +{ + for (size_t i = 0; i < kBlkLen128; ++i) { + const uint8_t v0 = ExtractSrcWeight(src_block_0, i); + const uint8_t v1 = ExtractSrcWeight(src_block_1, i); + const uint8_t v2 = ExtractSrcWeight(src_block_2, i); + const uint8_t v3 = ExtractSrcWeight(src_block_3, i); + dst[i] = static_cast( + static_cast(v0 | (v1 << 2) | (v2 << 4) | (v3 << 6)) + ); + } +} + +// +// Reference unpack of one 128-byte block-group back into 4 K-blocks worth of +// natural-order uint8 weights ([0, 3]). Independent of PackBlockGroup_BlkLen128 +// so it can serve as a round-trip oracle. +// +inline void +UnPackBlockGroup_BlkLen128_Reference(const std::byte* packed, + uint8_t out_block_0[kBlkLen128], + uint8_t out_block_1[kBlkLen128], + uint8_t out_block_2[kBlkLen128], + uint8_t out_block_3[kBlkLen128]) +{ + for (size_t i = 0; i < kBlkLen128; ++i) { + const uint8_t b = static_cast(packed[i]); + out_block_0[i] = static_cast((b >> 0) & 0x03u); + out_block_1[i] = static_cast((b >> 2) & 0x03u); + out_block_2[i] = static_cast((b >> 4) & 0x03u); + out_block_3[i] = static_cast((b >> 6) & 0x03u); + } +} + +// +// Byte offset for the BlkLen=128 packed B-data buffer. Same shape as +// PackedQuantBOffsetBytes_W2 (the BlkLen=64 variant just above) but using +// kBlockGroupBytes128 (= 128) per slot instead of kBlockGroupBytes (= 64). +// The N-major / 4-col-grouped layout is identical so the dispatcher's +// per-N-tile pointer arithmetic and the N-tail handling are reusable. +// +inline size_t +PackedQuantBOffsetBytes_W2_BlkLen128(size_t n, size_t blk_group, + size_t BlockGroupCountKPadded, size_t NMain) +{ + if (n < NMain) { + const size_t g = n / kNCols4; + const size_t c = n % kNCols4; + const size_t per_group_bytes = BlockGroupCountKPadded * kNCols4 * kBlockGroupBytes128; + return g * per_group_bytes + + blk_group * (kNCols4 * kBlockGroupBytes128) + + c * kBlockGroupBytes128; + } + return (n * BlockGroupCountKPadded + blk_group) * kBlockGroupBytes128; +} + +// +// NOTE on scale offsets: +// PackedQuantBScaleOffset_W2 is BlkLen-independent (one float per K-block, +// the layout depends only on kBlockGroupBlks and kNCols4). The BlkLen=128 +// path reuses it verbatim; no PackedQuantBScaleOffset_W2_BlkLen128 needed. +// + +// ----------------------------------------------------------------------------- +// BlkLen=32 constants and pack helpers (parallel to the BlkLen=64/128 sets). +// 4 K-blocks * 32 weights = 128 weights per group = 32 packed bytes per group +// (fits in one YMM). Same N-major / 4-col-grouped layout. Same K-tail rounding +// rule (round BlockCountK up to a multiple of kBlockGroupBlks=4). +// ----------------------------------------------------------------------------- + +constexpr size_t kBlkLen32 = 32; +constexpr size_t kBlkBytes32 = kBlkLen32 / kWeightsPerByte; // 8 packed src bytes per block +constexpr size_t kBlockGroupBytes32 = kBlockGroupBlks * kBlkBytes32; // 32 bytes per block-group +constexpr size_t kBlockGroupWeights32 = kBlockGroupBlks * kBlkLen32; // 128 weights per block-group + +inline void +PackBlockGroup_BlkLen32(const std::byte* src_block_0, + const std::byte* src_block_1, + const std::byte* src_block_2, + const std::byte* src_block_3, + std::byte* dst) +{ + for (size_t i = 0; i < kBlkLen32; ++i) { + const uint8_t v0 = ExtractSrcWeight(src_block_0, i); + const uint8_t v1 = ExtractSrcWeight(src_block_1, i); + const uint8_t v2 = ExtractSrcWeight(src_block_2, i); + const uint8_t v3 = ExtractSrcWeight(src_block_3, i); + dst[i] = static_cast( + static_cast(v0 | (v1 << 2) | (v2 << 4) | (v3 << 6)) + ); + } +} + +inline void +UnPackBlockGroup_BlkLen32_Reference(const std::byte* packed, + uint8_t out_block_0[kBlkLen32], + uint8_t out_block_1[kBlkLen32], + uint8_t out_block_2[kBlkLen32], + uint8_t out_block_3[kBlkLen32]) +{ + for (size_t i = 0; i < kBlkLen32; ++i) { + const uint8_t b = static_cast(packed[i]); + out_block_0[i] = static_cast((b >> 0) & 0x03u); + out_block_1[i] = static_cast((b >> 2) & 0x03u); + out_block_2[i] = static_cast((b >> 4) & 0x03u); + out_block_3[i] = static_cast((b >> 6) & 0x03u); + } +} + +inline size_t +PackedQuantBOffsetBytes_W2_BlkLen32(size_t n, size_t blk_group, + size_t BlockGroupCountKPadded, size_t NMain) +{ + if (n < NMain) { + const size_t g = n / kNCols4; + const size_t c = n % kNCols4; + const size_t per_group_bytes = BlockGroupCountKPadded * kNCols4 * kBlockGroupBytes32; + return g * per_group_bytes + + blk_group * (kNCols4 * kBlockGroupBytes32) + + c * kBlockGroupBytes32; + } + return (n * BlockGroupCountKPadded + blk_group) * kBlockGroupBytes32; +} + // ----------------------------------------------------------------------------- // block-group packed-data layout. // diff --git a/onnxruntime/core/mlas/lib/sqnbitgemm_kernel_avx512_2bit_blklen128.h b/onnxruntime/core/mlas/lib/sqnbitgemm_kernel_avx512_2bit_blklen128.h new file mode 100644 index 0000000000000..1065c32634c62 --- /dev/null +++ b/onnxruntime/core/mlas/lib/sqnbitgemm_kernel_avx512_2bit_blklen128.h @@ -0,0 +1,859 @@ +/*++ + +Copyright (c) Microsoft Corporation. All rights reserved. + +Licensed under the MIT License. + +Module Name: + + sqnbitgemm_kernel_avx512_2bit_blklen128.h + +Abstract: + + AVX-512 (-VNNI) W2 kernel for BlkLen=128. Sibling of + sqnbitgemm_kernel_avx512_2bit_blklen64.h: same R2xC4 tile, same SGEMM + correction step, same N-tail handling, same K-tail handling. The only + differences are: + * Each K-block is 128 weights = 32 bytes packed instead of 64 = 16. + * Each block-group is 4 K-blocks = 128 bytes = 2 ZMMs (vs 1 ZMM for + BlkLen=64). + * Unpack: 2 ZMM loads + 4 fixed shift/AND per load -> 4 PAIRS of ZMMs, + each pair holding 128 weights of one K-block (low half + high half). + * MAC per K-block: 2 dpbusds (low+high halves) instead of 1. + + Templated on `` like the BlkLen=64 sibling. + + Constraints: + * BlkLen == 128 only (other BlkLens go through their own SIMD headers). + * BlockCountK has no alignment requirement: the main loop iterates + full block-groups; a K-tail handler picks up 1-3 trailing K-blocks. + * Production: CountM and CountN handled via the same R2/R1 + N-tail + split as the BlkLen=64 kernel. + + Register pressure note: + * R2xC4 needs 8 acc + 16 A halves (preloaded per block-group iter) + + 8 B halves (loaded once per N-col within the iter) = 32 ZMMs at + peak, exactly at the AVX-512 register count. We do NOT preload all + 16 A halves at once; instead we hold 16 A halves at the top of the + block-group iter (one per K-block per M-row, low+high), let the + compiler spill the FMA temporaries if needed, and rely on the L1 + being warm for any spilled A halves. This matches the strategy the + BlkLen=64 R2xC4 tile uses (it preloads 8 A vecs and the compiler + manages the rest). + +--*/ + +#pragma once + +#include +#include +#include + +#include + +#include "mlasi.h" +#include "qnbitgemm.h" +#include "sqnbitgemm_kernel_avx512_2bit.h" + +namespace onnxruntime { +namespace mlas { +namespace sq2bit_avx512 { + +inline constexpr size_t kNRows2_BlkLen128 = 2; // R2 tile shape, same as BlkLen=64 + +// +// Cheap block-group unpack for BlkLen=128: 2x ZMM loads + 4x (fixed-shift + AND) +// per load = 8 unpacked half-ZMMs. +// +// Bit layout of each byte b of the 128-byte packed block-group (b in [0, 127]): +// bits[0..1] = block_0.weight[b] +// bits[2..3] = block_1.weight[b] +// bits[4..5] = block_2.weight[b] +// bits[6..7] = block_3.weight[b] +// +// For each K-block k in [0,3], the 128 unpacked weights are split into two +// 64-byte ZMMs: bv_k_lo holds weights[0..63], bv_k_hi holds weights[64..127]. +// +static MLAS_FORCEINLINE void +load_unpack_w2_block_group_blklen128( + const std::byte* packed, + __m512i& bv0_lo, __m512i& bv1_lo, __m512i& bv2_lo, __m512i& bv3_lo, + __m512i& bv0_hi, __m512i& bv1_hi, __m512i& bv2_hi, __m512i& bv3_hi) +{ + const __m512i group_lo = _mm512_loadu_si512(reinterpret_cast(packed)); + const __m512i group_hi = _mm512_loadu_si512(reinterpret_cast(packed + 64)); + const __m512i mask03 = _mm512_set1_epi8(0x03); + + bv0_lo = _mm512_and_si512(group_lo, mask03); + bv1_lo = _mm512_and_si512(_mm512_srli_epi16(group_lo, 2), mask03); + bv2_lo = _mm512_and_si512(_mm512_srli_epi16(group_lo, 4), mask03); + bv3_lo = _mm512_and_si512(_mm512_srli_epi16(group_lo, 6), mask03); + + bv0_hi = _mm512_and_si512(group_hi, mask03); + bv1_hi = _mm512_and_si512(_mm512_srli_epi16(group_hi, 2), mask03); + bv2_hi = _mm512_and_si512(_mm512_srli_epi16(group_hi, 4), mask03); + bv3_hi = _mm512_and_si512(_mm512_srli_epi16(group_hi, 6), mask03); +} + +// +// Per single M-row, 4-K-block dot-and-accumulate for BlkLen=128. +// Each K-block is 128 weights = 2 dpbusds (low half + high half) summed +// before being scaled and added to acc. +// +// Math per K-block: +// int32 d = dpbusd(0, bv_lo, av_lo); d = dpbusd(d, bv_hi, av_hi); +// acc += scale_a[blk] * scale_b[blk] * cvtepi32_ps(d) +// +// scale_a and scale_b each point to 4 consecutive floats (one per K-block). +// +// We use the same two-sub-accumulator trick as the BlkLen=64 kernel: lo +// accumulator gets blocks {0,2}, hi accumulator gets {1,3}, summed at end. +// Critical path is ~2 FMA latencies per chain (vs 4 if single-chained). +// +template +static MLAS_FORCEINLINE void +dot_accumulate_4blk_w2_blklen128( + const __m512i& av0_lo, const __m512i& av0_hi, + const __m512i& av1_lo, const __m512i& av1_hi, + const __m512i& av2_lo, const __m512i& av2_hi, + const __m512i& av3_lo, const __m512i& av3_hi, + const __m512i& bv0_lo, const __m512i& bv0_hi, + const __m512i& bv1_lo, const __m512i& bv1_hi, + const __m512i& bv2_lo, const __m512i& bv2_hi, + const __m512i& bv3_lo, const __m512i& bv3_hi, + const float* scale_a, + const float* scale_b, + __m512& acc) +{ + __m512i d0, d1, d2, d3; + if constexpr (kVnni) { + d0 = _mm512_dpbusd_epi32(_mm512_setzero_si512(), bv0_lo, av0_lo); + d0 = _mm512_dpbusd_epi32(d0, bv0_hi, av0_hi); + d1 = _mm512_dpbusd_epi32(_mm512_setzero_si512(), bv1_lo, av1_lo); + d1 = _mm512_dpbusd_epi32(d1, bv1_hi, av1_hi); + d2 = _mm512_dpbusd_epi32(_mm512_setzero_si512(), bv2_lo, av2_lo); + d2 = _mm512_dpbusd_epi32(d2, bv2_hi, av2_hi); + d3 = _mm512_dpbusd_epi32(_mm512_setzero_si512(), bv3_lo, av3_lo); + d3 = _mm512_dpbusd_epi32(d3, bv3_hi, av3_hi); + } else { + // Non-VNNI: vpmaddubsw -> vpmaddwd chain, low and high halves summed + // before adding into the K-block accumulator. + const __m512i ones = _mm512_set1_epi16(1); + const __m512i t0_lo = _mm512_maddubs_epi16(bv0_lo, av0_lo); + const __m512i t0_hi = _mm512_maddubs_epi16(bv0_hi, av0_hi); + const __m512i t1_lo = _mm512_maddubs_epi16(bv1_lo, av1_lo); + const __m512i t1_hi = _mm512_maddubs_epi16(bv1_hi, av1_hi); + const __m512i t2_lo = _mm512_maddubs_epi16(bv2_lo, av2_lo); + const __m512i t2_hi = _mm512_maddubs_epi16(bv2_hi, av2_hi); + const __m512i t3_lo = _mm512_maddubs_epi16(bv3_lo, av3_lo); + const __m512i t3_hi = _mm512_maddubs_epi16(bv3_hi, av3_hi); + d0 = _mm512_add_epi32(_mm512_madd_epi16(t0_lo, ones), + _mm512_madd_epi16(t0_hi, ones)); + d1 = _mm512_add_epi32(_mm512_madd_epi16(t1_lo, ones), + _mm512_madd_epi16(t1_hi, ones)); + d2 = _mm512_add_epi32(_mm512_madd_epi16(t2_lo, ones), + _mm512_madd_epi16(t2_hi, ones)); + d3 = _mm512_add_epi32(_mm512_madd_epi16(t3_lo, ones), + _mm512_madd_epi16(t3_hi, ones)); + } + + const __m512 s0 = _mm512_set1_ps(scale_a[0] * scale_b[0]); + const __m512 s1 = _mm512_set1_ps(scale_a[1] * scale_b[1]); + const __m512 s2 = _mm512_set1_ps(scale_a[2] * scale_b[2]); + const __m512 s3 = _mm512_set1_ps(scale_a[3] * scale_b[3]); + + __m512 acc_lo = _mm512_fmadd_ps(_mm512_cvtepi32_ps(d0), s0, _mm512_setzero_ps()); + __m512 acc_hi = _mm512_fmadd_ps(_mm512_cvtepi32_ps(d1), s1, _mm512_setzero_ps()); + acc_lo = _mm512_fmadd_ps(_mm512_cvtepi32_ps(d2), s2, acc_lo); + acc_hi = _mm512_fmadd_ps(_mm512_cvtepi32_ps(d3), s3, acc_hi); + acc = _mm512_add_ps(acc, _mm512_add_ps(acc_lo, acc_hi)); +} + +// +// 2 M-rows x 1 N-col x 4 K-blocks (one block-group) accumulator for BlkLen=128. +// The block-group B load + unpack is shared across the 2 M-rows. +// +template +static MLAS_FORCEINLINE void +accumulate_w2_blklen128_r2c1blk4( + const __m512i& av00_lo, const __m512i& av00_hi, + const __m512i& av01_lo, const __m512i& av01_hi, + const __m512i& av02_lo, const __m512i& av02_hi, + const __m512i& av03_lo, const __m512i& av03_hi, + const __m512i& av10_lo, const __m512i& av10_hi, + const __m512i& av11_lo, const __m512i& av11_hi, + const __m512i& av12_lo, const __m512i& av12_hi, + const __m512i& av13_lo, const __m512i& av13_hi, + const std::byte* QuantBDataPtr, + const float* scale_a0, + const float* scale_a1, + const float* scale_b, + __m512& acc0, + __m512& acc1) +{ + __m512i bv0_lo, bv1_lo, bv2_lo, bv3_lo; + __m512i bv0_hi, bv1_hi, bv2_hi, bv3_hi; + load_unpack_w2_block_group_blklen128(QuantBDataPtr, + bv0_lo, bv1_lo, bv2_lo, bv3_lo, + bv0_hi, bv1_hi, bv2_hi, bv3_hi); + + dot_accumulate_4blk_w2_blklen128( + av00_lo, av00_hi, av01_lo, av01_hi, av02_lo, av02_hi, av03_lo, av03_hi, + bv0_lo, bv0_hi, bv1_lo, bv1_hi, bv2_lo, bv2_hi, bv3_lo, bv3_hi, + scale_a0, scale_b, acc0); + dot_accumulate_4blk_w2_blklen128( + av10_lo, av10_hi, av11_lo, av11_hi, av12_lo, av12_hi, av13_lo, av13_hi, + bv0_lo, bv0_hi, bv1_lo, bv1_hi, bv2_lo, bv2_hi, bv3_lo, bv3_hi, + scale_a1, scale_b, acc1); +} + +// +// 1 M-row x 1 N-col x 4 K-blocks (one block-group) accumulator for BlkLen=128. +// +template +static MLAS_FORCEINLINE void +accumulate_w2_blklen128_r1c1blk4( + const __m512i& av00_lo, const __m512i& av00_hi, + const __m512i& av01_lo, const __m512i& av01_hi, + const __m512i& av02_lo, const __m512i& av02_hi, + const __m512i& av03_lo, const __m512i& av03_hi, + const std::byte* QuantBDataPtr, + const float* scale_a0, + const float* scale_b, + __m512& acc0) +{ + __m512i bv0_lo, bv1_lo, bv2_lo, bv3_lo; + __m512i bv0_hi, bv1_hi, bv2_hi, bv3_hi; + load_unpack_w2_block_group_blklen128(QuantBDataPtr, + bv0_lo, bv1_lo, bv2_lo, bv3_lo, + bv0_hi, bv1_hi, bv2_hi, bv3_hi); + + dot_accumulate_4blk_w2_blklen128( + av00_lo, av00_hi, av01_lo, av01_hi, av02_lo, av02_hi, av03_lo, av03_hi, + bv0_lo, bv0_hi, bv1_lo, bv1_hi, bv2_lo, bv2_hi, bv3_lo, bv3_hi, + scale_a0, scale_b, acc0); +} + +// +// R1 x C4 tile (M=1 decode or trailing odd row of R2xC4). +// +template +MLAS_FORCEINLINE void +Q2Int8GemmR1xC4BlkLen128Avx512( + const std::byte* QuantA, + const float* QuantAScale, + const std::byte* QuantBData, + const float* QuantBScale, + float* C, + size_t CountM, // expected to be 1 (caller-enforced) + size_t CountN, + size_t BlockCountK, + const float* Bias, + size_t ldc) +{ + const size_t lda = BlockCountK * kBlkLen128; + constexpr size_t PerColGroupBytes = kBlockGroupBytes128; // 128 B per col per group + constexpr size_t PerColGroupScale = kBlockGroupBlks; // 4 scales per col per group + constexpr size_t PerKGroupAdvanceBytes = kNCols4 * PerColGroupBytes; // 512 B per K-group iter + constexpr size_t PerKGroupAdvanceScale = kNCols4 * PerColGroupScale; // 16 scales per K-group iter + const size_t BlockGroupCountKPadded = + MlasDivRoundup(BlockCountK, kBlockGroupBlks); + const size_t BlockCountKPadded = BlockGroupCountKPadded * kBlockGroupBlks; + const size_t GroupStrideBytes = BlockCountKPadded * kNCols4 * kBlkBytes128; + const size_t GroupStrideScale = BlockCountKPadded * kNCols4; + + assert(CountN % kNCols4 == 0); + const size_t FullGroups = BlockCountK / kBlockGroupBlks; + const size_t TailBlocks = BlockCountK % kBlockGroupBlks; + + for (size_t m = 0; m < CountM; ++m) { + const std::byte* QuantBDataColPtr = QuantBData; + const float* QuantBScaleColPtr = QuantBScale; + const float* BiasPtr = Bias; + float* SumPtr = C + m * ldc; + + for (size_t n = 0; n < CountN; n += kNCols4) { + const std::byte* QuantAPtr = QuantA + m * lda; + const float* QuantAScalePtr = QuantAScale + m * BlockCountK; + + const std::byte* QuantBDataPtr = QuantBDataColPtr; + const float* QuantBScalePtr = QuantBScaleColPtr; + + __m512 acc[kNCols4] = { + _mm512_setzero_ps(), _mm512_setzero_ps(), + _mm512_setzero_ps(), _mm512_setzero_ps() + }; + + for (size_t sb = 0; sb < FullGroups; ++sb) { + // Load 4 K-blocks of A, each split into low+high 64-byte halves. + const __m512i av00_lo = _mm512_loadu_si512( + reinterpret_cast(QuantAPtr + 0 * kBlkLen128)); + const __m512i av00_hi = _mm512_loadu_si512( + reinterpret_cast(QuantAPtr + 0 * kBlkLen128 + 64)); + const __m512i av01_lo = _mm512_loadu_si512( + reinterpret_cast(QuantAPtr + 1 * kBlkLen128)); + const __m512i av01_hi = _mm512_loadu_si512( + reinterpret_cast(QuantAPtr + 1 * kBlkLen128 + 64)); + const __m512i av02_lo = _mm512_loadu_si512( + reinterpret_cast(QuantAPtr + 2 * kBlkLen128)); + const __m512i av02_hi = _mm512_loadu_si512( + reinterpret_cast(QuantAPtr + 2 * kBlkLen128 + 64)); + const __m512i av03_lo = _mm512_loadu_si512( + reinterpret_cast(QuantAPtr + 3 * kBlkLen128)); + const __m512i av03_hi = _mm512_loadu_si512( + reinterpret_cast(QuantAPtr + 3 * kBlkLen128 + 64)); + + accumulate_w2_blklen128_r1c1blk4( + av00_lo, av00_hi, av01_lo, av01_hi, av02_lo, av02_hi, av03_lo, av03_hi, + QuantBDataPtr + 0 * PerColGroupBytes, + QuantAScalePtr, QuantBScalePtr + 0 * PerColGroupScale, acc[0]); + accumulate_w2_blklen128_r1c1blk4( + av00_lo, av00_hi, av01_lo, av01_hi, av02_lo, av02_hi, av03_lo, av03_hi, + QuantBDataPtr + 1 * PerColGroupBytes, + QuantAScalePtr, QuantBScalePtr + 1 * PerColGroupScale, acc[1]); + accumulate_w2_blklen128_r1c1blk4( + av00_lo, av00_hi, av01_lo, av01_hi, av02_lo, av02_hi, av03_lo, av03_hi, + QuantBDataPtr + 2 * PerColGroupBytes, + QuantAScalePtr, QuantBScalePtr + 2 * PerColGroupScale, acc[2]); + accumulate_w2_blklen128_r1c1blk4( + av00_lo, av00_hi, av01_lo, av01_hi, av02_lo, av02_hi, av03_lo, av03_hi, + QuantBDataPtr + 3 * PerColGroupBytes, + QuantAScalePtr, QuantBScalePtr + 3 * PerColGroupScale, acc[3]); + + QuantAPtr += kBlkLen128 * kBlockGroupBlks; + QuantAScalePtr += kBlockGroupBlks; + QuantBDataPtr += PerKGroupAdvanceBytes; + QuantBScalePtr += PerKGroupAdvanceScale; + } + + // K-tail: 1-3 trailing real K-blocks. Zero-fill missing A halves. + if (TailBlocks > 0) { + const __m512i zero = _mm512_setzero_si512(); + auto load_lo = [&](size_t k) -> __m512i { + return (k < TailBlocks) + ? _mm512_loadu_si512(reinterpret_cast( + QuantAPtr + k * kBlkLen128)) + : zero; + }; + auto load_hi = [&](size_t k) -> __m512i { + return (k < TailBlocks) + ? _mm512_loadu_si512(reinterpret_cast( + QuantAPtr + k * kBlkLen128 + 64)) + : zero; + }; + const __m512i av00_lo = load_lo(0); + const __m512i av00_hi = load_hi(0); + const __m512i av01_lo = load_lo(1); + const __m512i av01_hi = load_hi(1); + const __m512i av02_lo = load_lo(2); + const __m512i av02_hi = load_hi(2); + const __m512i av03_lo = zero; + const __m512i av03_hi = zero; + + float scale_a0_safe[kBlockGroupBlks] = {0.0f, 0.0f, 0.0f, 0.0f}; + for (size_t i = 0; i < TailBlocks; ++i) { + scale_a0_safe[i] = QuantAScalePtr[i]; + } + + accumulate_w2_blklen128_r1c1blk4( + av00_lo, av00_hi, av01_lo, av01_hi, av02_lo, av02_hi, av03_lo, av03_hi, + QuantBDataPtr + 0 * PerColGroupBytes, + scale_a0_safe, QuantBScalePtr + 0 * PerColGroupScale, acc[0]); + accumulate_w2_blklen128_r1c1blk4( + av00_lo, av00_hi, av01_lo, av01_hi, av02_lo, av02_hi, av03_lo, av03_hi, + QuantBDataPtr + 1 * PerColGroupBytes, + scale_a0_safe, QuantBScalePtr + 1 * PerColGroupScale, acc[1]); + accumulate_w2_blklen128_r1c1blk4( + av00_lo, av00_hi, av01_lo, av01_hi, av02_lo, av02_hi, av03_lo, av03_hi, + QuantBDataPtr + 2 * PerColGroupBytes, + scale_a0_safe, QuantBScalePtr + 2 * PerColGroupScale, acc[2]); + accumulate_w2_blklen128_r1c1blk4( + av00_lo, av00_hi, av01_lo, av01_hi, av02_lo, av02_hi, av03_lo, av03_hi, + QuantBDataPtr + 3 * PerColGroupBytes, + scale_a0_safe, QuantBScalePtr + 3 * PerColGroupScale, acc[3]); + } + + SumPtr[0] = _mm512_reduce_add_ps(acc[0]); + SumPtr[1] = _mm512_reduce_add_ps(acc[1]); + SumPtr[2] = _mm512_reduce_add_ps(acc[2]); + SumPtr[3] = _mm512_reduce_add_ps(acc[3]); + if (BiasPtr != nullptr) { + SumPtr[0] += BiasPtr[0]; + SumPtr[1] += BiasPtr[1]; + SumPtr[2] += BiasPtr[2]; + SumPtr[3] += BiasPtr[3]; + } + + QuantBDataColPtr += GroupStrideBytes; + QuantBScaleColPtr += GroupStrideScale; + BiasPtr += BiasPtr != nullptr ? kNCols4 : 0; + SumPtr += kNCols4; + } + } +} + +// +// R2 x C4 tile (CountM >= 2 even). Mirrors the BlkLen=64 R2xC4 tile but with +// BlkLen=128-specific A loads (split into low+high halves) and the +// BlkLen=128 accumulator. +// +template +MLAS_FORCEINLINE void +Q2Int8GemmR2xC4BlkLen128Avx512( + const std::byte* QuantA, + const float* QuantAScale, + const std::byte* QuantBData, + const float* QuantBScale, + float* C, + size_t CountM, + size_t CountN, + size_t BlockCountK, + const float* Bias, + size_t ldc) +{ + const size_t lda = BlockCountK * kBlkLen128; + constexpr size_t PerColGroupBytes = kBlockGroupBytes128; + constexpr size_t PerColGroupScale = kBlockGroupBlks; + constexpr size_t PerKGroupAdvanceBytes = kNCols4 * PerColGroupBytes; + constexpr size_t PerKGroupAdvanceScale = kNCols4 * PerColGroupScale; + const size_t BlockGroupCountKPadded = + MlasDivRoundup(BlockCountK, kBlockGroupBlks); + const size_t BlockCountKPadded = BlockGroupCountKPadded * kBlockGroupBlks; + const size_t GroupStrideBytes = BlockCountKPadded * kNCols4 * kBlkBytes128; + const size_t GroupStrideScale = BlockCountKPadded * kNCols4; + + assert(CountM % kNRows2_BlkLen128 == 0); + assert(CountN % kNCols4 == 0); + const size_t FullGroups = BlockCountK / kBlockGroupBlks; + const size_t TailBlocks = BlockCountK % kBlockGroupBlks; + + for (size_t m = 0; m < CountM; m += kNRows2_BlkLen128) { + const std::byte* QuantBDataColPtr = QuantBData; + const float* QuantBScaleColPtr = QuantBScale; + const float* BiasPtr = Bias; + float* SumPtr = C + m * ldc; + + for (size_t n = 0; n < CountN; n += kNCols4) { + const std::byte* QuantAPtr = QuantA + m * lda; + const float* QuantAScalePtr = QuantAScale + m * BlockCountK; + + const std::byte* QuantBDataPtr = QuantBDataColPtr; + const float* QuantBScalePtr = QuantBScaleColPtr; + + __m512 acc[kNCols4 * kNRows2_BlkLen128] = { + _mm512_setzero_ps(), _mm512_setzero_ps(), _mm512_setzero_ps(), _mm512_setzero_ps(), + _mm512_setzero_ps(), _mm512_setzero_ps(), _mm512_setzero_ps(), _mm512_setzero_ps() + }; + + for (size_t sb = 0; sb < FullGroups; ++sb) { + // M-row 0: 4 K-blocks * 2 halves = 8 A vecs + const __m512i av00_lo = _mm512_loadu_si512( + reinterpret_cast(QuantAPtr + 0 * kBlkLen128)); + const __m512i av00_hi = _mm512_loadu_si512( + reinterpret_cast(QuantAPtr + 0 * kBlkLen128 + 64)); + const __m512i av01_lo = _mm512_loadu_si512( + reinterpret_cast(QuantAPtr + 1 * kBlkLen128)); + const __m512i av01_hi = _mm512_loadu_si512( + reinterpret_cast(QuantAPtr + 1 * kBlkLen128 + 64)); + const __m512i av02_lo = _mm512_loadu_si512( + reinterpret_cast(QuantAPtr + 2 * kBlkLen128)); + const __m512i av02_hi = _mm512_loadu_si512( + reinterpret_cast(QuantAPtr + 2 * kBlkLen128 + 64)); + const __m512i av03_lo = _mm512_loadu_si512( + reinterpret_cast(QuantAPtr + 3 * kBlkLen128)); + const __m512i av03_hi = _mm512_loadu_si512( + reinterpret_cast(QuantAPtr + 3 * kBlkLen128 + 64)); + + // M-row 1: 4 K-blocks * 2 halves = 8 A vecs + const __m512i av10_lo = _mm512_loadu_si512( + reinterpret_cast(QuantAPtr + lda + 0 * kBlkLen128)); + const __m512i av10_hi = _mm512_loadu_si512( + reinterpret_cast(QuantAPtr + lda + 0 * kBlkLen128 + 64)); + const __m512i av11_lo = _mm512_loadu_si512( + reinterpret_cast(QuantAPtr + lda + 1 * kBlkLen128)); + const __m512i av11_hi = _mm512_loadu_si512( + reinterpret_cast(QuantAPtr + lda + 1 * kBlkLen128 + 64)); + const __m512i av12_lo = _mm512_loadu_si512( + reinterpret_cast(QuantAPtr + lda + 2 * kBlkLen128)); + const __m512i av12_hi = _mm512_loadu_si512( + reinterpret_cast(QuantAPtr + lda + 2 * kBlkLen128 + 64)); + const __m512i av13_lo = _mm512_loadu_si512( + reinterpret_cast(QuantAPtr + lda + 3 * kBlkLen128)); + const __m512i av13_hi = _mm512_loadu_si512( + reinterpret_cast(QuantAPtr + lda + 3 * kBlkLen128 + 64)); + + accumulate_w2_blklen128_r2c1blk4( + av00_lo, av00_hi, av01_lo, av01_hi, av02_lo, av02_hi, av03_lo, av03_hi, + av10_lo, av10_hi, av11_lo, av11_hi, av12_lo, av12_hi, av13_lo, av13_hi, + QuantBDataPtr + 0 * PerColGroupBytes, + QuantAScalePtr, QuantAScalePtr + BlockCountK, + QuantBScalePtr + 0 * PerColGroupScale, + acc[0], acc[kNCols4 + 0]); + accumulate_w2_blklen128_r2c1blk4( + av00_lo, av00_hi, av01_lo, av01_hi, av02_lo, av02_hi, av03_lo, av03_hi, + av10_lo, av10_hi, av11_lo, av11_hi, av12_lo, av12_hi, av13_lo, av13_hi, + QuantBDataPtr + 1 * PerColGroupBytes, + QuantAScalePtr, QuantAScalePtr + BlockCountK, + QuantBScalePtr + 1 * PerColGroupScale, + acc[1], acc[kNCols4 + 1]); + accumulate_w2_blklen128_r2c1blk4( + av00_lo, av00_hi, av01_lo, av01_hi, av02_lo, av02_hi, av03_lo, av03_hi, + av10_lo, av10_hi, av11_lo, av11_hi, av12_lo, av12_hi, av13_lo, av13_hi, + QuantBDataPtr + 2 * PerColGroupBytes, + QuantAScalePtr, QuantAScalePtr + BlockCountK, + QuantBScalePtr + 2 * PerColGroupScale, + acc[2], acc[kNCols4 + 2]); + accumulate_w2_blklen128_r2c1blk4( + av00_lo, av00_hi, av01_lo, av01_hi, av02_lo, av02_hi, av03_lo, av03_hi, + av10_lo, av10_hi, av11_lo, av11_hi, av12_lo, av12_hi, av13_lo, av13_hi, + QuantBDataPtr + 3 * PerColGroupBytes, + QuantAScalePtr, QuantAScalePtr + BlockCountK, + QuantBScalePtr + 3 * PerColGroupScale, + acc[3], acc[kNCols4 + 3]); + + QuantAPtr += kBlkLen128 * kBlockGroupBlks; + QuantAScalePtr += kBlockGroupBlks; + QuantBDataPtr += PerKGroupAdvanceBytes; + QuantBScalePtr += PerKGroupAdvanceScale; + } + + // K-tail + if (TailBlocks > 0) { + const __m512i zero = _mm512_setzero_si512(); + auto load_a_lo = [&](size_t row_off, size_t k) -> __m512i { + return (k < TailBlocks) + ? _mm512_loadu_si512(reinterpret_cast( + QuantAPtr + row_off + k * kBlkLen128)) + : zero; + }; + auto load_a_hi = [&](size_t row_off, size_t k) -> __m512i { + return (k < TailBlocks) + ? _mm512_loadu_si512(reinterpret_cast( + QuantAPtr + row_off + k * kBlkLen128 + 64)) + : zero; + }; + const __m512i av00_lo = load_a_lo(0, 0); + const __m512i av00_hi = load_a_hi(0, 0); + const __m512i av01_lo = load_a_lo(0, 1); + const __m512i av01_hi = load_a_hi(0, 1); + const __m512i av02_lo = load_a_lo(0, 2); + const __m512i av02_hi = load_a_hi(0, 2); + const __m512i av03_lo = zero; + const __m512i av03_hi = zero; + + const __m512i av10_lo = load_a_lo(lda, 0); + const __m512i av10_hi = load_a_hi(lda, 0); + const __m512i av11_lo = load_a_lo(lda, 1); + const __m512i av11_hi = load_a_hi(lda, 1); + const __m512i av12_lo = load_a_lo(lda, 2); + const __m512i av12_hi = load_a_hi(lda, 2); + const __m512i av13_lo = zero; + const __m512i av13_hi = zero; + + float scale_a0_safe[kBlockGroupBlks] = {0.0f, 0.0f, 0.0f, 0.0f}; + float scale_a1_safe[kBlockGroupBlks] = {0.0f, 0.0f, 0.0f, 0.0f}; + for (size_t i = 0; i < TailBlocks; ++i) { + scale_a0_safe[i] = QuantAScalePtr[i]; + scale_a1_safe[i] = QuantAScalePtr[BlockCountK + i]; + } + + accumulate_w2_blklen128_r2c1blk4( + av00_lo, av00_hi, av01_lo, av01_hi, av02_lo, av02_hi, av03_lo, av03_hi, + av10_lo, av10_hi, av11_lo, av11_hi, av12_lo, av12_hi, av13_lo, av13_hi, + QuantBDataPtr + 0 * PerColGroupBytes, + scale_a0_safe, scale_a1_safe, + QuantBScalePtr + 0 * PerColGroupScale, + acc[0], acc[kNCols4 + 0]); + accumulate_w2_blklen128_r2c1blk4( + av00_lo, av00_hi, av01_lo, av01_hi, av02_lo, av02_hi, av03_lo, av03_hi, + av10_lo, av10_hi, av11_lo, av11_hi, av12_lo, av12_hi, av13_lo, av13_hi, + QuantBDataPtr + 1 * PerColGroupBytes, + scale_a0_safe, scale_a1_safe, + QuantBScalePtr + 1 * PerColGroupScale, + acc[1], acc[kNCols4 + 1]); + accumulate_w2_blklen128_r2c1blk4( + av00_lo, av00_hi, av01_lo, av01_hi, av02_lo, av02_hi, av03_lo, av03_hi, + av10_lo, av10_hi, av11_lo, av11_hi, av12_lo, av12_hi, av13_lo, av13_hi, + QuantBDataPtr + 2 * PerColGroupBytes, + scale_a0_safe, scale_a1_safe, + QuantBScalePtr + 2 * PerColGroupScale, + acc[2], acc[kNCols4 + 2]); + accumulate_w2_blklen128_r2c1blk4( + av00_lo, av00_hi, av01_lo, av01_hi, av02_lo, av02_hi, av03_lo, av03_hi, + av10_lo, av10_hi, av11_lo, av11_hi, av12_lo, av12_hi, av13_lo, av13_hi, + QuantBDataPtr + 3 * PerColGroupBytes, + scale_a0_safe, scale_a1_safe, + QuantBScalePtr + 3 * PerColGroupScale, + acc[3], acc[kNCols4 + 3]); + } + + SumPtr[0] = _mm512_reduce_add_ps(acc[0]); + SumPtr[1] = _mm512_reduce_add_ps(acc[1]); + SumPtr[2] = _mm512_reduce_add_ps(acc[2]); + SumPtr[3] = _mm512_reduce_add_ps(acc[3]); + SumPtr[ldc + 0] = _mm512_reduce_add_ps(acc[kNCols4 + 0]); + SumPtr[ldc + 1] = _mm512_reduce_add_ps(acc[kNCols4 + 1]); + SumPtr[ldc + 2] = _mm512_reduce_add_ps(acc[kNCols4 + 2]); + SumPtr[ldc + 3] = _mm512_reduce_add_ps(acc[kNCols4 + 3]); + if (BiasPtr != nullptr) { + SumPtr[0] += BiasPtr[0]; + SumPtr[1] += BiasPtr[1]; + SumPtr[2] += BiasPtr[2]; + SumPtr[3] += BiasPtr[3]; + SumPtr[ldc + 0] += BiasPtr[0]; + SumPtr[ldc + 1] += BiasPtr[1]; + SumPtr[ldc + 2] += BiasPtr[2]; + SumPtr[ldc + 3] += BiasPtr[3]; + } + + QuantBDataColPtr += GroupStrideBytes; + QuantBScaleColPtr += GroupStrideScale; + BiasPtr += BiasPtr != nullptr ? kNCols4 : 0; + SumPtr += kNCols4; + } + } +} + +// +// N-tail tile (1-3 trailing N-cols when CountN is not a multiple of kNCols4). +// Column-major packed B; walks one column at a time using the R1xC4 helper +// but with CountN = 1 per iteration. Slower than the main tile but bounded +// to at most 3 N-cols per call. +// +template +MLAS_FORCEINLINE void +Q2Int8GemmRMxC_Tail_BlkLen128Avx512( + const std::byte* QuantA, + const float* QuantAScale, + const std::byte* QuantBData, + const float* QuantBScale, + float* C, + size_t CountM, + size_t CountN, + size_t BlockCountK, + const float* Bias, + size_t ldc) +{ + const size_t lda = BlockCountK * kBlkLen128; + const size_t BlockGroupCountKPadded = + MlasDivRoundup(BlockCountK, kBlockGroupBlks); + const size_t BlockCountKPadded = BlockGroupCountKPadded * kBlockGroupBlks; + // Column-major in this tail region: each N-col is BlockCountKPadded * + // kBlkBytes128 bytes for B and BlockCountKPadded floats for scales. + const size_t FullGroups = BlockCountK / kBlockGroupBlks; + const size_t TailBlocks = BlockCountK % kBlockGroupBlks; + + for (size_t m = 0; m < CountM; ++m) { + const std::byte* a_row = QuantA + m * lda; + const float* a_scale_row = QuantAScale + m * BlockCountK; + float* c_row = C + m * ldc; + + for (size_t n = 0; n < CountN; ++n) { + __m512 acc = _mm512_setzero_ps(); + + const std::byte* b_col = QuantBData + n * BlockCountKPadded * kBlkBytes128; + const float* b_scale_col = QuantBScale + n * BlockCountKPadded; + const std::byte* QuantAPtr = a_row; + const float* QuantAScalePtr = a_scale_row; + const std::byte* QuantBDataPtr = b_col; + const float* QuantBScalePtr = b_scale_col; + + for (size_t sb = 0; sb < FullGroups; ++sb) { + const __m512i av00_lo = _mm512_loadu_si512( + reinterpret_cast(QuantAPtr + 0 * kBlkLen128)); + const __m512i av00_hi = _mm512_loadu_si512( + reinterpret_cast(QuantAPtr + 0 * kBlkLen128 + 64)); + const __m512i av01_lo = _mm512_loadu_si512( + reinterpret_cast(QuantAPtr + 1 * kBlkLen128)); + const __m512i av01_hi = _mm512_loadu_si512( + reinterpret_cast(QuantAPtr + 1 * kBlkLen128 + 64)); + const __m512i av02_lo = _mm512_loadu_si512( + reinterpret_cast(QuantAPtr + 2 * kBlkLen128)); + const __m512i av02_hi = _mm512_loadu_si512( + reinterpret_cast(QuantAPtr + 2 * kBlkLen128 + 64)); + const __m512i av03_lo = _mm512_loadu_si512( + reinterpret_cast(QuantAPtr + 3 * kBlkLen128)); + const __m512i av03_hi = _mm512_loadu_si512( + reinterpret_cast(QuantAPtr + 3 * kBlkLen128 + 64)); + + accumulate_w2_blklen128_r1c1blk4( + av00_lo, av00_hi, av01_lo, av01_hi, av02_lo, av02_hi, av03_lo, av03_hi, + QuantBDataPtr, QuantAScalePtr, QuantBScalePtr, acc); + + QuantAPtr += kBlkLen128 * kBlockGroupBlks; + QuantAScalePtr += kBlockGroupBlks; + QuantBDataPtr += kBlockGroupBytes128; + QuantBScalePtr += kBlockGroupBlks; + } + + if (TailBlocks > 0) { + const __m512i zero = _mm512_setzero_si512(); + auto load_lo = [&](size_t k) -> __m512i { + return (k < TailBlocks) + ? _mm512_loadu_si512(reinterpret_cast( + QuantAPtr + k * kBlkLen128)) + : zero; + }; + auto load_hi = [&](size_t k) -> __m512i { + return (k < TailBlocks) + ? _mm512_loadu_si512(reinterpret_cast( + QuantAPtr + k * kBlkLen128 + 64)) + : zero; + }; + const __m512i av00_lo = load_lo(0); + const __m512i av00_hi = load_hi(0); + const __m512i av01_lo = load_lo(1); + const __m512i av01_hi = load_hi(1); + const __m512i av02_lo = load_lo(2); + const __m512i av02_hi = load_hi(2); + const __m512i av03_lo = zero; + const __m512i av03_hi = zero; + + float scale_a0_safe[kBlockGroupBlks] = {0.0f, 0.0f, 0.0f, 0.0f}; + for (size_t i = 0; i < TailBlocks; ++i) { + scale_a0_safe[i] = QuantAScalePtr[i]; + } + + accumulate_w2_blklen128_r1c1blk4( + av00_lo, av00_hi, av01_lo, av01_hi, av02_lo, av02_hi, av03_lo, av03_hi, + QuantBDataPtr, scale_a0_safe, QuantBScalePtr, acc); + } + + const float sum = _mm512_reduce_add_ps(acc); + c_row[n] = (Bias != nullptr) ? (sum + Bias[n]) : sum; + } + } +} + +// +// Top-level BlkLen=128 kernel templated on . Same SGEMM BlkSum +// correction step as the BlkLen=64 kernel. Same R2/R1/N-tail split. +// +template +static MLAS_FORCEINLINE size_t +SQ2BitGemmKernel_BlkSum_CompInt8_BlkLen128_Impl( + const std::byte* QuantA, + const float* QuantAScale, + const std::byte* QuantBData, + const float* QuantBScale, + float* C, + size_t CountM, + size_t CountN, + size_t BlockCountK, + const float* Bias, + size_t ldc, + const float* ABlockSum, + const float* QuantBBlkSum) +{ + if (BlockCountK == 0 || CountM == 0 || CountN == 0) { + return 0; + } + + const size_t NMain = (CountN / kNCols4) * kNCols4; + const size_t NTail = CountN - NMain; + + const size_t M_pairs = CountM / kNRows2_BlkLen128; + const size_t M_main = M_pairs * kNRows2_BlkLen128; + const size_t M_tail = CountM - M_main; + const size_t lda = BlockCountK * kBlkLen128; + + if (NMain > 0) { + if (M_main > 0) { + Q2Int8GemmR2xC4BlkLen128Avx512( + QuantA, QuantAScale, QuantBData, QuantBScale, + C, M_main, NMain, BlockCountK, Bias, ldc); + } + if (M_tail > 0) { + Q2Int8GemmR1xC4BlkLen128Avx512( + QuantA + M_main * lda, + QuantAScale + M_main * BlockCountK, + QuantBData, QuantBScale, + C + M_main * ldc, + /*CountM=*/M_tail, + NMain, BlockCountK, Bias, ldc); + } + } + + if (NTail > 0) { + const size_t BlockGroupCountKPadded = + MlasDivRoundup(BlockCountK, kBlockGroupBlks); + const std::byte* QuantBDataTail = + QuantBData + NMain * BlockGroupCountKPadded * kBlockGroupBytes128; + const float* QuantBScaleTail = + QuantBScale + NMain * BlockGroupCountKPadded * kBlockGroupBlks; + const float* BiasTail = (Bias != nullptr) ? Bias + NMain : nullptr; + + Q2Int8GemmRMxC_Tail_BlkLen128Avx512( + QuantA, QuantAScale, + QuantBDataTail, QuantBScaleTail, + C + NMain, + CountM, NTail, BlockCountK, BiasTail, ldc); + } + + // BlkSum correction (width-16 chunked layout, same as BlkLen=64). + float* c_blk = C; + const float* b_blk_sum = QuantBBlkSum; + size_t RowsRemaining = CountM; + const float* a_blksum_row = ABlockSum; + while (RowsRemaining > 0) { + const auto RowsHandled = GetMlasPlatform().GemmFloatKernel( + a_blksum_row, b_blk_sum, c_blk, + BlockCountK, RowsRemaining, CountN, BlockCountK, ldc, + 1.0f, /*ZeroMode=*/false); + + c_blk += ldc * RowsHandled; + a_blksum_row += BlockCountK * RowsHandled; + RowsRemaining -= RowsHandled; + } + return CountM; +} + +// +// Top-level VNNI variant. Compiled into AVX-512-VNNI sources only. +// +static MLAS_FORCEINLINE size_t MLASCALL +SQ2BitGemmKernel_BlkSum_CompInt8_BlkLen128_Avx512Vnni( + const std::byte* QuantA, + const float* QuantAScale, + const std::byte* QuantBData, + const float* QuantBScale, + float* C, + size_t CountM, + size_t CountN, + size_t BlockCountK, + const float* Bias, + size_t ldc, + const float* ABlockSum, + const float* QuantBBlkSum) +{ + return SQ2BitGemmKernel_BlkSum_CompInt8_BlkLen128_Impl( + QuantA, QuantAScale, QuantBData, QuantBScale, + C, CountM, CountN, BlockCountK, Bias, ldc, ABlockSum, QuantBBlkSum); +} + +// +// Top-level non-VNNI variant. +// +static MLAS_FORCEINLINE size_t MLASCALL +SQ2BitGemmKernel_BlkSum_CompInt8_BlkLen128_Avx512( + const std::byte* QuantA, + const float* QuantAScale, + const std::byte* QuantBData, + const float* QuantBScale, + float* C, + size_t CountM, + size_t CountN, + size_t BlockCountK, + const float* Bias, + size_t ldc, + const float* ABlockSum, + const float* QuantBBlkSum) +{ + return SQ2BitGemmKernel_BlkSum_CompInt8_BlkLen128_Impl( + QuantA, QuantAScale, QuantBData, QuantBScale, + C, CountM, CountN, BlockCountK, Bias, ldc, ABlockSum, QuantBBlkSum); +} + +} // namespace sq2bit_avx512 +} // namespace mlas +} // namespace onnxruntime diff --git a/onnxruntime/core/mlas/lib/sqnbitgemm_kernel_avx512_2bit_blklen32.h b/onnxruntime/core/mlas/lib/sqnbitgemm_kernel_avx512_2bit_blklen32.h new file mode 100644 index 0000000000000..505c3e216184e --- /dev/null +++ b/onnxruntime/core/mlas/lib/sqnbitgemm_kernel_avx512_2bit_blklen32.h @@ -0,0 +1,696 @@ +/*++ + +Copyright (c) Microsoft Corporation. All rights reserved. + +Licensed under the MIT License. + +Module Name: + + sqnbitgemm_kernel_avx512_2bit_blklen32.h + +Abstract: + + AVX-512 (-VNNI) W2 kernel for BlkLen=32. Sibling of + sqnbitgemm_kernel_avx512_2bit_blklen64.h: same R2xC4 tile, same SGEMM + correction step, same N-tail handling, same K-tail handling. The only + differences are: + * Each K-block is 32 weights = 8 bytes packed (vs 64 = 16 for BlkLen=64). + * Each block-group is 4 K-blocks = 32 bytes = 1 YMM (vs 64 bytes = 1 ZMM). + * Unpack: 1 YMM load + 4 fixed shift/AND -> 4 YMMs of 32 unpacked + weights each. We zero-extend those YMMs into ZMM lower halves so the + downstream MAC can stay full-width `dpbusd` (which is what AVX-512 + provides). Wasted ZMM upper half is paid for by avoiding two separate + narrow-MAC paths -- a wash in instruction count and simpler to verify. + * MAC per K-block: 1 dpbusd (operating on 32 of 64 lanes; upper half + zero-padded). Same 2-sub-accumulator FMA chain as BlkLen=64. + + Templated on `` like the BlkLen=64/128 siblings. + + Constraints: + * BlkLen == 32 only. + * BlockCountK has no alignment requirement: the main K-loop iterates + full block-groups; a K-tail handler picks up 1-3 trailing K-blocks. + +--*/ + +#pragma once + +#include +#include +#include + +#include + +#include "mlasi.h" +#include "qnbitgemm.h" +#include "sqnbitgemm_kernel_avx512_2bit.h" + +namespace onnxruntime { +namespace mlas { +namespace sq2bit_avx512 { + +inline constexpr size_t kNRows2_BlkLen32 = 2; + +// +// Cheap block-group unpack for BlkLen=32: 1 YMM load + 4 fixed shift/AND -> +// 4 ZMMs holding 32 active bytes each in the LOW half (upper half zero). +// +// Bit layout of each byte b of the 32-byte packed block-group (b in [0, 31]): +// bits[0..1] = block_0.weight[b] +// bits[2..3] = block_1.weight[b] +// bits[4..5] = block_2.weight[b] +// bits[6..7] = block_3.weight[b] +// +// Each output ZMM contains 32 unpacked weights in lanes [0, 31] and zeros in +// lanes [32, 63]. The downstream dpbusd works correctly on this layout because +// the A vectors we feed it are similarly zero-padded in the upper half. +// +static MLAS_FORCEINLINE void +load_unpack_w2_block_group_blklen32( + const std::byte* packed, + __m512i& bv0_zext, __m512i& bv1_zext, __m512i& bv2_zext, __m512i& bv3_zext) +{ + const __m256i group_ymm = _mm256_loadu_si256(reinterpret_cast(packed)); + const __m256i mask03 = _mm256_set1_epi8(0x03); + + const __m256i bv0_ymm = _mm256_and_si256(group_ymm, mask03); + const __m256i bv1_ymm = _mm256_and_si256(_mm256_srli_epi16(group_ymm, 2), mask03); + const __m256i bv2_ymm = _mm256_and_si256(_mm256_srli_epi16(group_ymm, 4), mask03); + const __m256i bv3_ymm = _mm256_and_si256(_mm256_srli_epi16(group_ymm, 6), mask03); + + // Zero-extend each YMM into the lower half of a ZMM; upper half = 0. + bv0_zext = _mm512_zextsi256_si512(bv0_ymm); + bv1_zext = _mm512_zextsi256_si512(bv1_ymm); + bv2_zext = _mm512_zextsi256_si512(bv2_ymm); + bv3_zext = _mm512_zextsi256_si512(bv3_ymm); +} + +// +// Load one BlkLen=32 A block (32 int8 bytes) zero-extended into a ZMM. +// The lower half holds the 32 active bytes; upper half is zero. This pairs +// with the zero-extended B vectors produced by load_unpack_w2_block_group_blklen32 +// so dpbusd produces the correct int32 partial sums in the low half (and +// 0 in the high half, harmless to the subsequent FP reduction). +// +static MLAS_FORCEINLINE __m512i +load_a_blklen32_zext(const std::byte* a_block) +{ + const __m256i a_ymm = _mm256_loadu_si256(reinterpret_cast(a_block)); + return _mm512_zextsi256_si512(a_ymm); +} + +// +// Per single M-row, 4-K-block dot-and-accumulate for BlkLen=32. Each K-block +// is 32 lanes wide; dpbusd produces 16 int32 partial sums in the low half +// (upper half is zero, contributes nothing). Same 2-sub-accumulator FMA +// strategy as BlkLen=64: blocks {0,2} chain into acc_lo, blocks {1,3} chain +// into acc_hi. +// +// Math per K-block: acc += scale_a[blk] * scale_b[blk] * dot(av[blk], bv[blk]) +// +template +static MLAS_FORCEINLINE void +dot_accumulate_4blk_w2_blklen32( + const __m512i& av0, const __m512i& av1, const __m512i& av2, const __m512i& av3, + const __m512i& bv0, const __m512i& bv1, const __m512i& bv2, const __m512i& bv3, + const float* scale_a, + const float* scale_b, + __m512& acc) +{ + __m512i d0, d1, d2, d3; + if constexpr (kVnni) { + d0 = _mm512_dpbusd_epi32(_mm512_setzero_si512(), bv0, av0); + d1 = _mm512_dpbusd_epi32(_mm512_setzero_si512(), bv1, av1); + d2 = _mm512_dpbusd_epi32(_mm512_setzero_si512(), bv2, av2); + d3 = _mm512_dpbusd_epi32(_mm512_setzero_si512(), bv3, av3); + } else { + const __m512i ones = _mm512_set1_epi16(1); + const __m512i t0 = _mm512_maddubs_epi16(bv0, av0); + const __m512i t1 = _mm512_maddubs_epi16(bv1, av1); + const __m512i t2 = _mm512_maddubs_epi16(bv2, av2); + const __m512i t3 = _mm512_maddubs_epi16(bv3, av3); + d0 = _mm512_madd_epi16(t0, ones); + d1 = _mm512_madd_epi16(t1, ones); + d2 = _mm512_madd_epi16(t2, ones); + d3 = _mm512_madd_epi16(t3, ones); + } + + const __m512 s0 = _mm512_set1_ps(scale_a[0] * scale_b[0]); + const __m512 s1 = _mm512_set1_ps(scale_a[1] * scale_b[1]); + const __m512 s2 = _mm512_set1_ps(scale_a[2] * scale_b[2]); + const __m512 s3 = _mm512_set1_ps(scale_a[3] * scale_b[3]); + + __m512 acc_lo = _mm512_fmadd_ps(_mm512_cvtepi32_ps(d0), s0, _mm512_setzero_ps()); + __m512 acc_hi = _mm512_fmadd_ps(_mm512_cvtepi32_ps(d1), s1, _mm512_setzero_ps()); + acc_lo = _mm512_fmadd_ps(_mm512_cvtepi32_ps(d2), s2, acc_lo); + acc_hi = _mm512_fmadd_ps(_mm512_cvtepi32_ps(d3), s3, acc_hi); + acc = _mm512_add_ps(acc, _mm512_add_ps(acc_lo, acc_hi)); +} + +template +static MLAS_FORCEINLINE void +accumulate_w2_blklen32_r2c1blk4( + const __m512i& av00, const __m512i& av01, const __m512i& av02, const __m512i& av03, + const __m512i& av10, const __m512i& av11, const __m512i& av12, const __m512i& av13, + const std::byte* QuantBDataPtr, + const float* scale_a0, + const float* scale_a1, + const float* scale_b, + __m512& acc0, + __m512& acc1) +{ + __m512i bv0, bv1, bv2, bv3; + load_unpack_w2_block_group_blklen32(QuantBDataPtr, bv0, bv1, bv2, bv3); + + dot_accumulate_4blk_w2_blklen32( + av00, av01, av02, av03, bv0, bv1, bv2, bv3, scale_a0, scale_b, acc0); + dot_accumulate_4blk_w2_blklen32( + av10, av11, av12, av13, bv0, bv1, bv2, bv3, scale_a1, scale_b, acc1); +} + +template +static MLAS_FORCEINLINE void +accumulate_w2_blklen32_r1c1blk4( + const __m512i& av00, const __m512i& av01, const __m512i& av02, const __m512i& av03, + const std::byte* QuantBDataPtr, + const float* scale_a0, + const float* scale_b, + __m512& acc0) +{ + __m512i bv0, bv1, bv2, bv3; + load_unpack_w2_block_group_blklen32(QuantBDataPtr, bv0, bv1, bv2, bv3); + + dot_accumulate_4blk_w2_blklen32( + av00, av01, av02, av03, bv0, bv1, bv2, bv3, scale_a0, scale_b, acc0); +} + +// +// R1 x C4 tile (M=1 decode or trailing odd row of R2xC4). +// +template +MLAS_FORCEINLINE void +Q2Int8GemmR1xC4BlkLen32Avx512( + const std::byte* QuantA, + const float* QuantAScale, + const std::byte* QuantBData, + const float* QuantBScale, + float* C, + size_t CountM, + size_t CountN, + size_t BlockCountK, + const float* Bias, + size_t ldc) +{ + const size_t lda = BlockCountK * kBlkLen32; + constexpr size_t PerColGroupBytes = kBlockGroupBytes32; // 32 B per col per group + constexpr size_t PerColGroupScale = kBlockGroupBlks; // 4 scales per col per group + constexpr size_t PerKGroupAdvanceBytes = kNCols4 * PerColGroupBytes; // 128 B per K-group iter + constexpr size_t PerKGroupAdvanceScale = kNCols4 * PerColGroupScale; // 16 scales per K-group iter + const size_t BlockGroupCountKPadded = + MlasDivRoundup(BlockCountK, kBlockGroupBlks); + const size_t BlockCountKPadded = BlockGroupCountKPadded * kBlockGroupBlks; + const size_t GroupStrideBytes = BlockCountKPadded * kNCols4 * kBlkBytes32; + const size_t GroupStrideScale = BlockCountKPadded * kNCols4; + + assert(CountN % kNCols4 == 0); + const size_t FullGroups = BlockCountK / kBlockGroupBlks; + const size_t TailBlocks = BlockCountK % kBlockGroupBlks; + + for (size_t m = 0; m < CountM; ++m) { + const std::byte* QuantBDataColPtr = QuantBData; + const float* QuantBScaleColPtr = QuantBScale; + const float* BiasPtr = Bias; + float* SumPtr = C + m * ldc; + + for (size_t n = 0; n < CountN; n += kNCols4) { + const std::byte* QuantAPtr = QuantA + m * lda; + const float* QuantAScalePtr = QuantAScale + m * BlockCountK; + + const std::byte* QuantBDataPtr = QuantBDataColPtr; + const float* QuantBScalePtr = QuantBScaleColPtr; + + __m512 acc[kNCols4] = { + _mm512_setzero_ps(), _mm512_setzero_ps(), + _mm512_setzero_ps(), _mm512_setzero_ps() + }; + + for (size_t sb = 0; sb < FullGroups; ++sb) { + const __m512i av00 = load_a_blklen32_zext(QuantAPtr + 0 * kBlkLen32); + const __m512i av01 = load_a_blklen32_zext(QuantAPtr + 1 * kBlkLen32); + const __m512i av02 = load_a_blklen32_zext(QuantAPtr + 2 * kBlkLen32); + const __m512i av03 = load_a_blklen32_zext(QuantAPtr + 3 * kBlkLen32); + + accumulate_w2_blklen32_r1c1blk4( + av00, av01, av02, av03, + QuantBDataPtr + 0 * PerColGroupBytes, + QuantAScalePtr, QuantBScalePtr + 0 * PerColGroupScale, acc[0]); + accumulate_w2_blklen32_r1c1blk4( + av00, av01, av02, av03, + QuantBDataPtr + 1 * PerColGroupBytes, + QuantAScalePtr, QuantBScalePtr + 1 * PerColGroupScale, acc[1]); + accumulate_w2_blklen32_r1c1blk4( + av00, av01, av02, av03, + QuantBDataPtr + 2 * PerColGroupBytes, + QuantAScalePtr, QuantBScalePtr + 2 * PerColGroupScale, acc[2]); + accumulate_w2_blklen32_r1c1blk4( + av00, av01, av02, av03, + QuantBDataPtr + 3 * PerColGroupBytes, + QuantAScalePtr, QuantBScalePtr + 3 * PerColGroupScale, acc[3]); + + QuantAPtr += kBlkLen32 * kBlockGroupBlks; + QuantAScalePtr += kBlockGroupBlks; + QuantBDataPtr += PerKGroupAdvanceBytes; + QuantBScalePtr += PerKGroupAdvanceScale; + } + + if (TailBlocks > 0) { + const __m512i zero = _mm512_setzero_si512(); + auto load_a = [&](size_t k) -> __m512i { + return (k < TailBlocks) + ? load_a_blklen32_zext(QuantAPtr + k * kBlkLen32) + : zero; + }; + const __m512i av00 = load_a(0); + const __m512i av01 = load_a(1); + const __m512i av02 = load_a(2); + const __m512i av03 = zero; + + float scale_a0_safe[kBlockGroupBlks] = {0.0f, 0.0f, 0.0f, 0.0f}; + for (size_t i = 0; i < TailBlocks; ++i) { + scale_a0_safe[i] = QuantAScalePtr[i]; + } + + accumulate_w2_blklen32_r1c1blk4( + av00, av01, av02, av03, + QuantBDataPtr + 0 * PerColGroupBytes, + scale_a0_safe, QuantBScalePtr + 0 * PerColGroupScale, acc[0]); + accumulate_w2_blklen32_r1c1blk4( + av00, av01, av02, av03, + QuantBDataPtr + 1 * PerColGroupBytes, + scale_a0_safe, QuantBScalePtr + 1 * PerColGroupScale, acc[1]); + accumulate_w2_blklen32_r1c1blk4( + av00, av01, av02, av03, + QuantBDataPtr + 2 * PerColGroupBytes, + scale_a0_safe, QuantBScalePtr + 2 * PerColGroupScale, acc[2]); + accumulate_w2_blklen32_r1c1blk4( + av00, av01, av02, av03, + QuantBDataPtr + 3 * PerColGroupBytes, + scale_a0_safe, QuantBScalePtr + 3 * PerColGroupScale, acc[3]); + } + + SumPtr[0] = _mm512_reduce_add_ps(acc[0]); + SumPtr[1] = _mm512_reduce_add_ps(acc[1]); + SumPtr[2] = _mm512_reduce_add_ps(acc[2]); + SumPtr[3] = _mm512_reduce_add_ps(acc[3]); + if (BiasPtr != nullptr) { + SumPtr[0] += BiasPtr[0]; + SumPtr[1] += BiasPtr[1]; + SumPtr[2] += BiasPtr[2]; + SumPtr[3] += BiasPtr[3]; + } + + QuantBDataColPtr += GroupStrideBytes; + QuantBScaleColPtr += GroupStrideScale; + BiasPtr += BiasPtr != nullptr ? kNCols4 : 0; + SumPtr += kNCols4; + } + } +} + +// +// R2 x C4 tile (CountM >= 2 even). +// +template +MLAS_FORCEINLINE void +Q2Int8GemmR2xC4BlkLen32Avx512( + const std::byte* QuantA, + const float* QuantAScale, + const std::byte* QuantBData, + const float* QuantBScale, + float* C, + size_t CountM, + size_t CountN, + size_t BlockCountK, + const float* Bias, + size_t ldc) +{ + const size_t lda = BlockCountK * kBlkLen32; + constexpr size_t PerColGroupBytes = kBlockGroupBytes32; + constexpr size_t PerColGroupScale = kBlockGroupBlks; + constexpr size_t PerKGroupAdvanceBytes = kNCols4 * PerColGroupBytes; + constexpr size_t PerKGroupAdvanceScale = kNCols4 * PerColGroupScale; + const size_t BlockGroupCountKPadded = + MlasDivRoundup(BlockCountK, kBlockGroupBlks); + const size_t BlockCountKPadded = BlockGroupCountKPadded * kBlockGroupBlks; + const size_t GroupStrideBytes = BlockCountKPadded * kNCols4 * kBlkBytes32; + const size_t GroupStrideScale = BlockCountKPadded * kNCols4; + + assert(CountM % kNRows2_BlkLen32 == 0); + assert(CountN % kNCols4 == 0); + const size_t FullGroups = BlockCountK / kBlockGroupBlks; + const size_t TailBlocks = BlockCountK % kBlockGroupBlks; + + for (size_t m = 0; m < CountM; m += kNRows2_BlkLen32) { + const std::byte* QuantBDataColPtr = QuantBData; + const float* QuantBScaleColPtr = QuantBScale; + const float* BiasPtr = Bias; + float* SumPtr = C + m * ldc; + + for (size_t n = 0; n < CountN; n += kNCols4) { + const std::byte* QuantAPtr = QuantA + m * lda; + const float* QuantAScalePtr = QuantAScale + m * BlockCountK; + + const std::byte* QuantBDataPtr = QuantBDataColPtr; + const float* QuantBScalePtr = QuantBScaleColPtr; + + __m512 acc[kNCols4 * kNRows2_BlkLen32] = { + _mm512_setzero_ps(), _mm512_setzero_ps(), _mm512_setzero_ps(), _mm512_setzero_ps(), + _mm512_setzero_ps(), _mm512_setzero_ps(), _mm512_setzero_ps(), _mm512_setzero_ps() + }; + + for (size_t sb = 0; sb < FullGroups; ++sb) { + const __m512i av00 = load_a_blklen32_zext(QuantAPtr + 0 * kBlkLen32); + const __m512i av01 = load_a_blklen32_zext(QuantAPtr + 1 * kBlkLen32); + const __m512i av02 = load_a_blklen32_zext(QuantAPtr + 2 * kBlkLen32); + const __m512i av03 = load_a_blklen32_zext(QuantAPtr + 3 * kBlkLen32); + const __m512i av10 = load_a_blklen32_zext(QuantAPtr + lda + 0 * kBlkLen32); + const __m512i av11 = load_a_blklen32_zext(QuantAPtr + lda + 1 * kBlkLen32); + const __m512i av12 = load_a_blklen32_zext(QuantAPtr + lda + 2 * kBlkLen32); + const __m512i av13 = load_a_blklen32_zext(QuantAPtr + lda + 3 * kBlkLen32); + + accumulate_w2_blklen32_r2c1blk4( + av00, av01, av02, av03, av10, av11, av12, av13, + QuantBDataPtr + 0 * PerColGroupBytes, + QuantAScalePtr, QuantAScalePtr + BlockCountK, + QuantBScalePtr + 0 * PerColGroupScale, + acc[0], acc[kNCols4 + 0]); + accumulate_w2_blklen32_r2c1blk4( + av00, av01, av02, av03, av10, av11, av12, av13, + QuantBDataPtr + 1 * PerColGroupBytes, + QuantAScalePtr, QuantAScalePtr + BlockCountK, + QuantBScalePtr + 1 * PerColGroupScale, + acc[1], acc[kNCols4 + 1]); + accumulate_w2_blklen32_r2c1blk4( + av00, av01, av02, av03, av10, av11, av12, av13, + QuantBDataPtr + 2 * PerColGroupBytes, + QuantAScalePtr, QuantAScalePtr + BlockCountK, + QuantBScalePtr + 2 * PerColGroupScale, + acc[2], acc[kNCols4 + 2]); + accumulate_w2_blklen32_r2c1blk4( + av00, av01, av02, av03, av10, av11, av12, av13, + QuantBDataPtr + 3 * PerColGroupBytes, + QuantAScalePtr, QuantAScalePtr + BlockCountK, + QuantBScalePtr + 3 * PerColGroupScale, + acc[3], acc[kNCols4 + 3]); + + QuantAPtr += kBlkLen32 * kBlockGroupBlks; + QuantAScalePtr += kBlockGroupBlks; + QuantBDataPtr += PerKGroupAdvanceBytes; + QuantBScalePtr += PerKGroupAdvanceScale; + } + + if (TailBlocks > 0) { + const __m512i zero = _mm512_setzero_si512(); + auto load_a = [&](size_t row_off, size_t k) -> __m512i { + return (k < TailBlocks) + ? load_a_blklen32_zext(QuantAPtr + row_off + k * kBlkLen32) + : zero; + }; + const __m512i av00 = load_a(0, 0); + const __m512i av01 = load_a(0, 1); + const __m512i av02 = load_a(0, 2); + const __m512i av03 = zero; + const __m512i av10 = load_a(lda, 0); + const __m512i av11 = load_a(lda, 1); + const __m512i av12 = load_a(lda, 2); + const __m512i av13 = zero; + + float scale_a0_safe[kBlockGroupBlks] = {0.0f, 0.0f, 0.0f, 0.0f}; + float scale_a1_safe[kBlockGroupBlks] = {0.0f, 0.0f, 0.0f, 0.0f}; + for (size_t i = 0; i < TailBlocks; ++i) { + scale_a0_safe[i] = QuantAScalePtr[i]; + scale_a1_safe[i] = QuantAScalePtr[BlockCountK + i]; + } + + accumulate_w2_blklen32_r2c1blk4( + av00, av01, av02, av03, av10, av11, av12, av13, + QuantBDataPtr + 0 * PerColGroupBytes, + scale_a0_safe, scale_a1_safe, + QuantBScalePtr + 0 * PerColGroupScale, + acc[0], acc[kNCols4 + 0]); + accumulate_w2_blklen32_r2c1blk4( + av00, av01, av02, av03, av10, av11, av12, av13, + QuantBDataPtr + 1 * PerColGroupBytes, + scale_a0_safe, scale_a1_safe, + QuantBScalePtr + 1 * PerColGroupScale, + acc[1], acc[kNCols4 + 1]); + accumulate_w2_blklen32_r2c1blk4( + av00, av01, av02, av03, av10, av11, av12, av13, + QuantBDataPtr + 2 * PerColGroupBytes, + scale_a0_safe, scale_a1_safe, + QuantBScalePtr + 2 * PerColGroupScale, + acc[2], acc[kNCols4 + 2]); + accumulate_w2_blklen32_r2c1blk4( + av00, av01, av02, av03, av10, av11, av12, av13, + QuantBDataPtr + 3 * PerColGroupBytes, + scale_a0_safe, scale_a1_safe, + QuantBScalePtr + 3 * PerColGroupScale, + acc[3], acc[kNCols4 + 3]); + } + + SumPtr[0] = _mm512_reduce_add_ps(acc[0]); + SumPtr[1] = _mm512_reduce_add_ps(acc[1]); + SumPtr[2] = _mm512_reduce_add_ps(acc[2]); + SumPtr[3] = _mm512_reduce_add_ps(acc[3]); + SumPtr[ldc + 0] = _mm512_reduce_add_ps(acc[kNCols4 + 0]); + SumPtr[ldc + 1] = _mm512_reduce_add_ps(acc[kNCols4 + 1]); + SumPtr[ldc + 2] = _mm512_reduce_add_ps(acc[kNCols4 + 2]); + SumPtr[ldc + 3] = _mm512_reduce_add_ps(acc[kNCols4 + 3]); + if (BiasPtr != nullptr) { + SumPtr[0] += BiasPtr[0]; + SumPtr[1] += BiasPtr[1]; + SumPtr[2] += BiasPtr[2]; + SumPtr[3] += BiasPtr[3]; + SumPtr[ldc + 0] += BiasPtr[0]; + SumPtr[ldc + 1] += BiasPtr[1]; + SumPtr[ldc + 2] += BiasPtr[2]; + SumPtr[ldc + 3] += BiasPtr[3]; + } + + QuantBDataColPtr += GroupStrideBytes; + QuantBScaleColPtr += GroupStrideScale; + BiasPtr += BiasPtr != nullptr ? kNCols4 : 0; + SumPtr += kNCols4; + } + } +} + +// +// N-tail tile (1-3 trailing N-cols when CountN is not a multiple of kNCols4). +// +template +MLAS_FORCEINLINE void +Q2Int8GemmRMxC_Tail_BlkLen32Avx512( + const std::byte* QuantA, + const float* QuantAScale, + const std::byte* QuantBData, + const float* QuantBScale, + float* C, + size_t CountM, + size_t CountN, + size_t BlockCountK, + const float* Bias, + size_t ldc) +{ + const size_t lda = BlockCountK * kBlkLen32; + const size_t BlockGroupCountKPadded = + MlasDivRoundup(BlockCountK, kBlockGroupBlks); + const size_t BlockCountKPadded = BlockGroupCountKPadded * kBlockGroupBlks; + const size_t FullGroups = BlockCountK / kBlockGroupBlks; + const size_t TailBlocks = BlockCountK % kBlockGroupBlks; + + for (size_t m = 0; m < CountM; ++m) { + const std::byte* a_row = QuantA + m * lda; + const float* a_scale_row = QuantAScale + m * BlockCountK; + float* c_row = C + m * ldc; + + for (size_t n = 0; n < CountN; ++n) { + __m512 acc = _mm512_setzero_ps(); + + const std::byte* b_col = QuantBData + n * BlockCountKPadded * kBlkBytes32; + const float* b_scale_col = QuantBScale + n * BlockCountKPadded; + const std::byte* QuantAPtr = a_row; + const float* QuantAScalePtr = a_scale_row; + const std::byte* QuantBDataPtr = b_col; + const float* QuantBScalePtr = b_scale_col; + + for (size_t sb = 0; sb < FullGroups; ++sb) { + const __m512i av00 = load_a_blklen32_zext(QuantAPtr + 0 * kBlkLen32); + const __m512i av01 = load_a_blklen32_zext(QuantAPtr + 1 * kBlkLen32); + const __m512i av02 = load_a_blklen32_zext(QuantAPtr + 2 * kBlkLen32); + const __m512i av03 = load_a_blklen32_zext(QuantAPtr + 3 * kBlkLen32); + + accumulate_w2_blklen32_r1c1blk4( + av00, av01, av02, av03, + QuantBDataPtr, QuantAScalePtr, QuantBScalePtr, acc); + + QuantAPtr += kBlkLen32 * kBlockGroupBlks; + QuantAScalePtr += kBlockGroupBlks; + QuantBDataPtr += kBlockGroupBytes32; + QuantBScalePtr += kBlockGroupBlks; + } + + if (TailBlocks > 0) { + const __m512i zero = _mm512_setzero_si512(); + auto load_a = [&](size_t k) -> __m512i { + return (k < TailBlocks) + ? load_a_blklen32_zext(QuantAPtr + k * kBlkLen32) + : zero; + }; + const __m512i av00 = load_a(0); + const __m512i av01 = load_a(1); + const __m512i av02 = load_a(2); + const __m512i av03 = zero; + + float scale_a0_safe[kBlockGroupBlks] = {0.0f, 0.0f, 0.0f, 0.0f}; + for (size_t i = 0; i < TailBlocks; ++i) { + scale_a0_safe[i] = QuantAScalePtr[i]; + } + + accumulate_w2_blklen32_r1c1blk4( + av00, av01, av02, av03, + QuantBDataPtr, scale_a0_safe, QuantBScalePtr, acc); + } + + const float sum = _mm512_reduce_add_ps(acc); + c_row[n] = (Bias != nullptr) ? (sum + Bias[n]) : sum; + } + } +} + +// +// Top-level BlkLen=32 kernel templated on . Same SGEMM BlkSum +// correction step as the BlkLen=64/128 kernels. +// +template +static MLAS_FORCEINLINE size_t +SQ2BitGemmKernel_BlkSum_CompInt8_BlkLen32_Impl( + const std::byte* QuantA, + const float* QuantAScale, + const std::byte* QuantBData, + const float* QuantBScale, + float* C, + size_t CountM, + size_t CountN, + size_t BlockCountK, + const float* Bias, + size_t ldc, + const float* ABlockSum, + const float* QuantBBlkSum) +{ + if (BlockCountK == 0 || CountM == 0 || CountN == 0) { + return 0; + } + + const size_t NMain = (CountN / kNCols4) * kNCols4; + const size_t NTail = CountN - NMain; + + const size_t M_pairs = CountM / kNRows2_BlkLen32; + const size_t M_main = M_pairs * kNRows2_BlkLen32; + const size_t M_tail = CountM - M_main; + const size_t lda = BlockCountK * kBlkLen32; + + if (NMain > 0) { + if (M_main > 0) { + Q2Int8GemmR2xC4BlkLen32Avx512( + QuantA, QuantAScale, QuantBData, QuantBScale, + C, M_main, NMain, BlockCountK, Bias, ldc); + } + if (M_tail > 0) { + Q2Int8GemmR1xC4BlkLen32Avx512( + QuantA + M_main * lda, + QuantAScale + M_main * BlockCountK, + QuantBData, QuantBScale, + C + M_main * ldc, + /*CountM=*/M_tail, + NMain, BlockCountK, Bias, ldc); + } + } + + if (NTail > 0) { + const size_t BlockGroupCountKPadded = + MlasDivRoundup(BlockCountK, kBlockGroupBlks); + const std::byte* QuantBDataTail = + QuantBData + NMain * BlockGroupCountKPadded * kBlockGroupBytes32; + const float* QuantBScaleTail = + QuantBScale + NMain * BlockGroupCountKPadded * kBlockGroupBlks; + const float* BiasTail = (Bias != nullptr) ? Bias + NMain : nullptr; + + Q2Int8GemmRMxC_Tail_BlkLen32Avx512( + QuantA, QuantAScale, + QuantBDataTail, QuantBScaleTail, + C + NMain, + CountM, NTail, BlockCountK, BiasTail, ldc); + } + + // BlkSum correction (width-16 chunked layout, same as other BlkLens). + float* c_blk = C; + const float* b_blk_sum = QuantBBlkSum; + size_t RowsRemaining = CountM; + const float* a_blksum_row = ABlockSum; + while (RowsRemaining > 0) { + const auto RowsHandled = GetMlasPlatform().GemmFloatKernel( + a_blksum_row, b_blk_sum, c_blk, + BlockCountK, RowsRemaining, CountN, BlockCountK, ldc, + 1.0f, /*ZeroMode=*/false); + + c_blk += ldc * RowsHandled; + a_blksum_row += BlockCountK * RowsHandled; + RowsRemaining -= RowsHandled; + } + return CountM; +} + +static MLAS_FORCEINLINE size_t MLASCALL +SQ2BitGemmKernel_BlkSum_CompInt8_BlkLen32_Avx512Vnni( + const std::byte* QuantA, + const float* QuantAScale, + const std::byte* QuantBData, + const float* QuantBScale, + float* C, + size_t CountM, + size_t CountN, + size_t BlockCountK, + const float* Bias, + size_t ldc, + const float* ABlockSum, + const float* QuantBBlkSum) +{ + return SQ2BitGemmKernel_BlkSum_CompInt8_BlkLen32_Impl( + QuantA, QuantAScale, QuantBData, QuantBScale, + C, CountM, CountN, BlockCountK, Bias, ldc, ABlockSum, QuantBBlkSum); +} + +static MLAS_FORCEINLINE size_t MLASCALL +SQ2BitGemmKernel_BlkSum_CompInt8_BlkLen32_Avx512( + const std::byte* QuantA, + const float* QuantAScale, + const std::byte* QuantBData, + const float* QuantBScale, + float* C, + size_t CountM, + size_t CountN, + size_t BlockCountK, + const float* Bias, + size_t ldc, + const float* ABlockSum, + const float* QuantBBlkSum) +{ + return SQ2BitGemmKernel_BlkSum_CompInt8_BlkLen32_Impl( + QuantA, QuantAScale, QuantBData, QuantBScale, + C, CountM, CountN, BlockCountK, Bias, ldc, ABlockSum, QuantBBlkSum); +} + +} // namespace sq2bit_avx512 +} // namespace mlas +} // namespace onnxruntime diff --git a/onnxruntime/core/mlas/lib/sqnbitgemm_kernel_avx512vnni.cpp b/onnxruntime/core/mlas/lib/sqnbitgemm_kernel_avx512vnni.cpp index 8e4d5a18d517c..f617eac4b5211 100644 --- a/onnxruntime/core/mlas/lib/sqnbitgemm_kernel_avx512vnni.cpp +++ b/onnxruntime/core/mlas/lib/sqnbitgemm_kernel_avx512vnni.cpp @@ -29,6 +29,8 @@ Module Name: #include "sqnbitgemm_kernel_avx512_int8_blklen128.h" #include "sqnbitgemm_kernel_avx512_2bit.h" #include "sqnbitgemm_kernel_avx512_2bit_blklen64.h" +#include "sqnbitgemm_kernel_avx512_2bit_blklen128.h" +#include "sqnbitgemm_kernel_avx512_2bit_blklen32.h" MLAS_FORCEINLINE void SQ4BitGemmM1Kernel_CompFp32( @@ -484,6 +486,16 @@ SQ2BitGemmKernel_BlkSum_CompInt8_Avx512Vnni_TestEntry( const float* ABlockSum, const float* QuantBBlkSum) { + if (BlkLen == 128) { + return SQ2BitGemmKernel_BlkSum_CompInt8_BlkLen128_Avx512Vnni( + QuantA, QuantAScale, QuantBData, QuantBScale, + C, CountM, CountN, BlockCountK, Bias, ldc, ABlockSum, QuantBBlkSum); + } + if (BlkLen == 32) { + return SQ2BitGemmKernel_BlkSum_CompInt8_BlkLen32_Avx512Vnni( + QuantA, QuantAScale, QuantBData, QuantBScale, + C, CountM, CountN, BlockCountK, Bias, ldc, ABlockSum, QuantBBlkSum); + } return SQ2BitGemmKernel_BlkSum_CompInt8_Avx512Vnni( BlkLen, QuantA, QuantAScale, QuantBData, QuantBScale, QuantBZeroPoint, C, CountM, CountN, CountK, BlockCountK, Bias, ldc, ABlockSum, QuantBBlkSum); diff --git a/onnxruntime/test/contrib_ops/matmul_2bits_test.cc b/onnxruntime/test/contrib_ops/matmul_2bits_test.cc index d45ba429c0490..2a6b461c63a7a 100644 --- a/onnxruntime/test/contrib_ops/matmul_2bits_test.cc +++ b/onnxruntime/test/contrib_ops/matmul_2bits_test.cc @@ -1095,6 +1095,107 @@ TEST(MatMul2Bits, Float32_2b_BlkLen64_Accuracy0) { TestMatMul2BitsTyped(); } +// MatMulNBits op-level coverage for the BlkLen=128 path. Same coverage matrix +// as the BlkLen=64 tests above, adjusted so that K is always a multiple of +// 128 (BlkLen=128 constraint). Exercises: +// * Single-block K (BlockCountK=1) -> K=128 +// * BlockCountK=2,3 not multiple of 4 -> K=256, K=384 (K-tail handler) +// * Exact block-group multiples -> K=512, K=1024, K=2048 +// * BlockCountK=6 (1 full + 2 tail) -> K=768 +// and the same M/N combinations as BlkLen=64. +TEST(MatMul2Bits, Float32_2b_BlkLen128_Accuracy4) { + // Single-block K (K = BlkLen = 128). + TestMatMul2BitsTyped(); + TestMatMul2BitsTyped(); + TestMatMul2BitsTyped(); + TestMatMul2BitsTyped(); + + // BlockCountK=2,3 -- exercises K-tail handler. + TestMatMul2BitsTyped(); + TestMatMul2BitsTyped(); + TestMatMul2BitsTyped(); + + // Exact block-group multiples (BlockCountK = 4, 8, 16). + TestMatMul2BitsTyped(); + TestMatMul2BitsTyped(); + TestMatMul2BitsTyped(); + + // BlockCountK=6 = 1 full block-group + 2-block tail. + TestMatMul2BitsTyped(); + + // Larger M (multi-iter R2 tile) at customer-shape proportions. + TestMatMul2BitsTyped(); + TestMatMul2BitsTyped(); +} + +TEST(MatMul2Bits, Float32_2b_BlkLen128_Accuracy0) { + TestMatMul2BitsTyped(); + TestMatMul2BitsTyped(); + TestMatMul2BitsTyped(); + TestMatMul2BitsTyped(); + + TestMatMul2BitsTyped(); + TestMatMul2BitsTyped(); + TestMatMul2BitsTyped(); + + TestMatMul2BitsTyped(); + TestMatMul2BitsTyped(); + TestMatMul2BitsTyped(); + + TestMatMul2BitsTyped(); + + TestMatMul2BitsTyped(); + TestMatMul2BitsTyped(); +} + +// MatMulNBits op-level coverage for the BlkLen=32 path. K multiples of 32. +// Same coverage matrix as BlkLen=64/128 (single-block, K-tail, exact group +// multiples, customer-shape proportions, single-row decode, M=100 prefill). +TEST(MatMul2Bits, Float32_2b_BlkLen32_Accuracy4) { + // Single-block K (K = BlkLen = 32). + TestMatMul2BitsTyped(); + TestMatMul2BitsTyped(); + TestMatMul2BitsTyped(); + TestMatMul2BitsTyped(); + + // BlockCountK=2,3 -- K-tail handler (BlockCountK not a multiple of 4). + TestMatMul2BitsTyped(); + TestMatMul2BitsTyped(); + TestMatMul2BitsTyped(); + + // Exact block-group multiples (BlockCountK = 4, 8, 16). + TestMatMul2BitsTyped(); + TestMatMul2BitsTyped(); + TestMatMul2BitsTyped(); + + // BlockCountK=6 = 1 full block-group + 2-block tail. + TestMatMul2BitsTyped(); + + // Larger M. + TestMatMul2BitsTyped(); + TestMatMul2BitsTyped(); +} + +TEST(MatMul2Bits, Float32_2b_BlkLen32_Accuracy0) { + TestMatMul2BitsTyped(); + TestMatMul2BitsTyped(); + TestMatMul2BitsTyped(); + TestMatMul2BitsTyped(); + + TestMatMul2BitsTyped(); + TestMatMul2BitsTyped(); + TestMatMul2BitsTyped(); + + TestMatMul2BitsTyped(); + TestMatMul2BitsTyped(); + TestMatMul2BitsTyped(); + + TestMatMul2BitsTyped(); + + TestMatMul2BitsTyped(); + TestMatMul2BitsTyped(); +} + #if defined(USE_WEBGPU) && !defined(ORT_USE_EP_API_ADAPTERS) namespace { diff --git a/onnxruntime/test/mlas/bench/bench_qnbitgemm.cpp b/onnxruntime/test/mlas/bench/bench_qnbitgemm.cpp index 4dba8b916bb75..96d2dcef3a6d9 100644 --- a/onnxruntime/test/mlas/bench/bench_qnbitgemm.cpp +++ b/onnxruntime/test/mlas/bench/bench_qnbitgemm.cpp @@ -141,7 +141,7 @@ static void QNBit2BitArgs(benchmark::internal::Benchmark* b) { b->ArgNames({"BlkLen", "M", "N", "K", "Threads", "Symmetric", "HasBias", "ComputeType"}); b->ArgsProduct({ - {64}, // BlkLen (W2 native kernel constraint) + {32, 64, 128}, // BlkLen (W2 native kernel supports all three) {1, 4096}, // M (decode + prefill) {4096, 11008}, // N {4096, 11008}, // K @@ -150,9 +150,7 @@ static void QNBit2BitArgs(benchmark::internal::Benchmark* b) { {int64_t{false}, int64_t{true}}, // HasBias {int64_t{SQNBIT_CompInt8}}, // ComputeType (W2 native kernel constraint) }); -} - -BENCHMARK(QNBITGEMM)->Apply(QNBitGemmArgs)->UseRealTime(); +}BENCHMARK(QNBITGEMM)->Apply(QNBitGemmArgs)->UseRealTime(); BENCHMARK(QNBITGEMM)->Apply(QNBitGemmArgs)->UseRealTime(); BENCHMARK(QNBITGEMM)->Apply(QNBitGemmArgs)->UseRealTime(); BENCHMARK(QNBITGEMM)->Apply(QNBit2BitArgs)->UseRealTime(); diff --git a/onnxruntime/test/mlas/unittest/test_sqnbitgemm_2bit.cpp b/onnxruntime/test/mlas/unittest/test_sqnbitgemm_2bit.cpp index 159bb8c4fa8f9..bb77139e1ff36 100644 --- a/onnxruntime/test/mlas/unittest/test_sqnbitgemm_2bit.cpp +++ b/onnxruntime/test/mlas/unittest/test_sqnbitgemm_2bit.cpp @@ -204,3 +204,314 @@ TEST(MlasSq2BitTest, PackUnpackRoundTrip_BlockGroup_BlkLen64_ConstantValues) } } } + +// ----------------------------------------------------------------------------- +// Block-group round-trip tests for BlkLen=128. +// Mirror the BlkLen=64 tests above. Each block-group still aggregates 4 +// K-blocks; only the per-block byte width doubles (16 -> 32) and the per- +// group byte total doubles (64 -> 128). The packing rule is identical so +// the same kinds of invariants hold. +// ----------------------------------------------------------------------------- + +namespace { +// Standard ONNX 2-bit source packing for a BlkLen=128 block (32 bytes). +void +PackSourceBlock_BlkLen128(const uint8_t weights[sq2::kBlkLen128], std::byte* src_out) +{ + for (size_t i = 0; i < sq2::kBlkBytes128; ++i) { + const uint8_t v0 = weights[4 * i + 0] & 0x03u; + const uint8_t v1 = weights[4 * i + 1] & 0x03u; + const uint8_t v2 = weights[4 * i + 2] & 0x03u; + const uint8_t v3 = weights[4 * i + 3] & 0x03u; + src_out[i] = static_cast( + static_cast(v0 | (v1 << 2) | (v2 << 4) | (v3 << 6)) + ); + } +} +} // namespace + +TEST(MlasSq2BitTest, PackUnpackRoundTrip_BlockGroup_BlkLen128_DeterministicPattern) +{ + std::array, sq2::kBlockGroupBlks> weights{}; + for (size_t k = 0; k < sq2::kBlockGroupBlks; ++k) { + for (size_t i = 0; i < sq2::kBlkLen128; ++i) { + weights[k][i] = static_cast((i + k) % 4); + } + } + + std::array, sq2::kBlockGroupBlks> src{}; + for (size_t k = 0; k < sq2::kBlockGroupBlks; ++k) { + PackSourceBlock_BlkLen128(weights[k].data(), src[k].data()); + } + + std::array packed{}; + sq2::PackBlockGroup_BlkLen128(src[0].data(), src[1].data(), src[2].data(), src[3].data(), + packed.data()); + + std::array, sq2::kBlockGroupBlks> recovered{}; + sq2::UnPackBlockGroup_BlkLen128_Reference(packed.data(), + recovered[0].data(), recovered[1].data(), + recovered[2].data(), recovered[3].data()); + + for (size_t k = 0; k < sq2::kBlockGroupBlks; ++k) { + for (size_t i = 0; i < sq2::kBlkLen128; ++i) { + ASSERT_EQ(recovered[k][i], weights[k][i]) + << "BlkLen128 round-trip mismatch k=" << k << " i=" << i; + } + } +} + +TEST(MlasSq2BitTest, PackUnpackRoundTrip_BlockGroup_BlkLen128_Randomized) +{ + constexpr unsigned kSeeds = 8; + + for (unsigned seed = 0; seed < kSeeds; ++seed) { + std::mt19937 rng(seed * 5051u + 13u); + std::uniform_int_distribution dist(0u, 3u); + + std::array, sq2::kBlockGroupBlks> weights{}; + for (auto& blk : weights) { + for (auto& w : blk) { + w = static_cast(dist(rng)); + } + } + + std::array, sq2::kBlockGroupBlks> src{}; + for (size_t k = 0; k < sq2::kBlockGroupBlks; ++k) { + PackSourceBlock_BlkLen128(weights[k].data(), src[k].data()); + } + + std::array packed{}; + sq2::PackBlockGroup_BlkLen128(src[0].data(), src[1].data(), src[2].data(), src[3].data(), + packed.data()); + + std::array, sq2::kBlockGroupBlks> recovered{}; + sq2::UnPackBlockGroup_BlkLen128_Reference(packed.data(), + recovered[0].data(), recovered[1].data(), + recovered[2].data(), recovered[3].data()); + + for (size_t k = 0; k < sq2::kBlockGroupBlks; ++k) { + for (size_t i = 0; i < sq2::kBlkLen128; ++i) { + ASSERT_EQ(recovered[k][i], weights[k][i]) + << "BlkLen128 random round-trip seed=" << seed << " k=" << k << " i=" << i; + } + } + } +} + +TEST(MlasSq2BitTest, PackUnpackRoundTrip_BlockGroup_BlkLen128_ConstantValues) +{ + // Case 1: every block filled with v. + for (uint8_t v = 0; v < 4; ++v) { + std::array, sq2::kBlockGroupBlks> weights{}; + for (auto& blk : weights) { + blk.fill(v); + } + + std::array, sq2::kBlockGroupBlks> src{}; + for (size_t k = 0; k < sq2::kBlockGroupBlks; ++k) { + PackSourceBlock_BlkLen128(weights[k].data(), src[k].data()); + } + + std::array packed{}; + sq2::PackBlockGroup_BlkLen128(src[0].data(), src[1].data(), src[2].data(), src[3].data(), + packed.data()); + + const uint8_t expected_byte = static_cast(v * 0x55u); + for (size_t i = 0; i < sq2::kBlockGroupBytes128; ++i) { + ASSERT_EQ(static_cast(packed[i]), expected_byte) + << "BlkLen128 uniform-fill v=" << static_cast(v) << " byte_i=" << i; + } + } + + // Case 2: only one block at a time carries a non-zero value. + for (size_t target_k = 0; target_k < sq2::kBlockGroupBlks; ++target_k) { + for (uint8_t v = 1; v < 4; ++v) { + std::array, sq2::kBlockGroupBlks> weights{}; + weights[target_k].fill(v); + + std::array, sq2::kBlockGroupBlks> src{}; + for (size_t k = 0; k < sq2::kBlockGroupBlks; ++k) { + PackSourceBlock_BlkLen128(weights[k].data(), src[k].data()); + } + + std::array packed{}; + sq2::PackBlockGroup_BlkLen128(src[0].data(), src[1].data(), src[2].data(), src[3].data(), + packed.data()); + + const uint8_t expected_byte = static_cast(v << (2 * target_k)); + for (size_t i = 0; i < sq2::kBlockGroupBytes128; ++i) { + ASSERT_EQ(static_cast(packed[i]), expected_byte) + << "BlkLen128 isolated block target_k=" << target_k + << " v=" << static_cast(v) << " byte_i=" << i; + } + + std::array, sq2::kBlockGroupBlks> recovered{}; + sq2::UnPackBlockGroup_BlkLen128_Reference( + packed.data(), + recovered[0].data(), recovered[1].data(), + recovered[2].data(), recovered[3].data()); + for (size_t k = 0; k < sq2::kBlockGroupBlks; ++k) { + const uint8_t expect_val = (k == target_k) ? v : uint8_t{0}; + for (size_t i = 0; i < sq2::kBlkLen128; ++i) { + ASSERT_EQ(recovered[k][i], expect_val) + << "BlkLen128 isolated round-trip target_k=" << target_k + << " k=" << k << " v=" << static_cast(v) << " i=" << i; + } + } + } + } +} + +// ----------------------------------------------------------------------------- +// Block-group round-trip tests for BlkLen=32. +// Each block-group is still 4 K-blocks; per-block byte width is 8 (vs 16 for +// BlkLen=64 and 32 for BlkLen=128), per-group bytes is 32. +// ----------------------------------------------------------------------------- + +namespace { +void +PackSourceBlock_BlkLen32(const uint8_t weights[sq2::kBlkLen32], std::byte* src_out) +{ + for (size_t i = 0; i < sq2::kBlkBytes32; ++i) { + const uint8_t v0 = weights[4 * i + 0] & 0x03u; + const uint8_t v1 = weights[4 * i + 1] & 0x03u; + const uint8_t v2 = weights[4 * i + 2] & 0x03u; + const uint8_t v3 = weights[4 * i + 3] & 0x03u; + src_out[i] = static_cast( + static_cast(v0 | (v1 << 2) | (v2 << 4) | (v3 << 6)) + ); + } +} +} // namespace + +TEST(MlasSq2BitTest, PackUnpackRoundTrip_BlockGroup_BlkLen32_DeterministicPattern) +{ + std::array, sq2::kBlockGroupBlks> weights{}; + for (size_t k = 0; k < sq2::kBlockGroupBlks; ++k) { + for (size_t i = 0; i < sq2::kBlkLen32; ++i) { + weights[k][i] = static_cast((i + k) % 4); + } + } + + std::array, sq2::kBlockGroupBlks> src{}; + for (size_t k = 0; k < sq2::kBlockGroupBlks; ++k) { + PackSourceBlock_BlkLen32(weights[k].data(), src[k].data()); + } + + std::array packed{}; + sq2::PackBlockGroup_BlkLen32(src[0].data(), src[1].data(), src[2].data(), src[3].data(), + packed.data()); + + std::array, sq2::kBlockGroupBlks> recovered{}; + sq2::UnPackBlockGroup_BlkLen32_Reference(packed.data(), + recovered[0].data(), recovered[1].data(), + recovered[2].data(), recovered[3].data()); + + for (size_t k = 0; k < sq2::kBlockGroupBlks; ++k) { + for (size_t i = 0; i < sq2::kBlkLen32; ++i) { + ASSERT_EQ(recovered[k][i], weights[k][i]) + << "BlkLen32 round-trip mismatch k=" << k << " i=" << i; + } + } +} + +TEST(MlasSq2BitTest, PackUnpackRoundTrip_BlockGroup_BlkLen32_Randomized) +{ + constexpr unsigned kSeeds = 8; + + for (unsigned seed = 0; seed < kSeeds; ++seed) { + std::mt19937 rng(seed * 5051u + 13u); + std::uniform_int_distribution dist(0u, 3u); + + std::array, sq2::kBlockGroupBlks> weights{}; + for (auto& blk : weights) { + for (auto& w : blk) { + w = static_cast(dist(rng)); + } + } + + std::array, sq2::kBlockGroupBlks> src{}; + for (size_t k = 0; k < sq2::kBlockGroupBlks; ++k) { + PackSourceBlock_BlkLen32(weights[k].data(), src[k].data()); + } + + std::array packed{}; + sq2::PackBlockGroup_BlkLen32(src[0].data(), src[1].data(), src[2].data(), src[3].data(), + packed.data()); + + std::array, sq2::kBlockGroupBlks> recovered{}; + sq2::UnPackBlockGroup_BlkLen32_Reference(packed.data(), + recovered[0].data(), recovered[1].data(), + recovered[2].data(), recovered[3].data()); + + for (size_t k = 0; k < sq2::kBlockGroupBlks; ++k) { + for (size_t i = 0; i < sq2::kBlkLen32; ++i) { + ASSERT_EQ(recovered[k][i], weights[k][i]) + << "BlkLen32 random round-trip seed=" << seed << " k=" << k << " i=" << i; + } + } + } +} + +TEST(MlasSq2BitTest, PackUnpackRoundTrip_BlockGroup_BlkLen32_ConstantValues) +{ + for (uint8_t v = 0; v < 4; ++v) { + std::array, sq2::kBlockGroupBlks> weights{}; + for (auto& blk : weights) { + blk.fill(v); + } + + std::array, sq2::kBlockGroupBlks> src{}; + for (size_t k = 0; k < sq2::kBlockGroupBlks; ++k) { + PackSourceBlock_BlkLen32(weights[k].data(), src[k].data()); + } + + std::array packed{}; + sq2::PackBlockGroup_BlkLen32(src[0].data(), src[1].data(), src[2].data(), src[3].data(), + packed.data()); + + const uint8_t expected_byte = static_cast(v * 0x55u); + for (size_t i = 0; i < sq2::kBlockGroupBytes32; ++i) { + ASSERT_EQ(static_cast(packed[i]), expected_byte) + << "BlkLen32 uniform-fill v=" << static_cast(v) << " byte_i=" << i; + } + } + + for (size_t target_k = 0; target_k < sq2::kBlockGroupBlks; ++target_k) { + for (uint8_t v = 1; v < 4; ++v) { + std::array, sq2::kBlockGroupBlks> weights{}; + weights[target_k].fill(v); + + std::array, sq2::kBlockGroupBlks> src{}; + for (size_t k = 0; k < sq2::kBlockGroupBlks; ++k) { + PackSourceBlock_BlkLen32(weights[k].data(), src[k].data()); + } + + std::array packed{}; + sq2::PackBlockGroup_BlkLen32(src[0].data(), src[1].data(), src[2].data(), src[3].data(), + packed.data()); + + const uint8_t expected_byte = static_cast(v << (2 * target_k)); + for (size_t i = 0; i < sq2::kBlockGroupBytes32; ++i) { + ASSERT_EQ(static_cast(packed[i]), expected_byte) + << "BlkLen32 isolated block target_k=" << target_k + << " v=" << static_cast(v) << " byte_i=" << i; + } + + std::array, sq2::kBlockGroupBlks> recovered{}; + sq2::UnPackBlockGroup_BlkLen32_Reference( + packed.data(), + recovered[0].data(), recovered[1].data(), + recovered[2].data(), recovered[3].data()); + for (size_t k = 0; k < sq2::kBlockGroupBlks; ++k) { + const uint8_t expect_val = (k == target_k) ? v : uint8_t{0}; + for (size_t i = 0; i < sq2::kBlkLen32; ++i) { + ASSERT_EQ(recovered[k][i], expect_val) + << "BlkLen32 isolated round-trip target_k=" << target_k + << " k=" << k << " v=" << static_cast(v) << " i=" << i; + } + } + } + } +} diff --git a/onnxruntime/test/mlas/unittest/test_sqnbitgemm_2bit_gemm.cpp b/onnxruntime/test/mlas/unittest/test_sqnbitgemm_2bit_gemm.cpp index aa48f9d6a153a..80a28a51e85a4 100644 --- a/onnxruntime/test/mlas/unittest/test_sqnbitgemm_2bit_gemm.cpp +++ b/onnxruntime/test/mlas/unittest/test_sqnbitgemm_2bit_gemm.cpp @@ -524,3 +524,738 @@ TEST(MlasSq2BitTest, BlkLen64_Avx512Vnni_WithZeroPoints) } } } + +// ============================================================================= +// BlkLen=128 coverage. Mirrors the BlkLen=64 tests above. The helpers are +// duplicated rather than templated to keep the BlkLen=64 path bit-identical; +// the diff is purely additive. +// ============================================================================= + +namespace { + +constexpr size_t kBlkLen128 = sq2::kBlkLen128; // 128 +constexpr size_t kBlkBytes128 = sq2::kBlkBytes128; // 32 + +void +PackSourceBlock_BlkLen128(const uint8_t weights[kBlkLen128], std::byte* src_out) +{ + for (size_t i = 0; i < kBlkBytes128; ++i) { + const uint8_t v0 = weights[4 * i + 0] & 0x03u; + const uint8_t v1 = weights[4 * i + 1] & 0x03u; + const uint8_t v2 = weights[4 * i + 2] & 0x03u; + const uint8_t v3 = weights[4 * i + 3] & 0x03u; + src_out[i] = static_cast( + static_cast(v0 | (v1 << 2) | (v2 << 4) | (v3 << 6)) + ); + } +} + +void +QuantizeA_Reference_BlkLen128(size_t M, size_t K, const float* A, + int8_t* QuantAData, float* QuantAScale) +{ + const size_t BlockCountK = (K + kBlkLen128 - 1) / kBlkLen128; + for (size_t m = 0; m < M; ++m) { + for (size_t k = 0, k_blk = 0; k < K; k += kBlkLen128, ++k_blk) { + const size_t local_len = std::min(K - k, kBlkLen128); + float amax = 0.0f; + for (size_t kk = 0; kk < local_len; ++kk) { + amax = std::max(amax, std::fabs(A[m * K + k + kk])); + } + constexpr float range_max = 127.0f; + const float scale = amax / range_max; + const float scale_recip = amax != 0.0f ? range_max / amax : 0.0f; + QuantAScale[m * BlockCountK + k_blk] = scale; + for (size_t kk = 0; kk < kBlkLen128; ++kk) { + const float a = (kk < local_len) ? A[m * K + k + kk] : 0.0f; + const float q = std::nearbyint(a * scale_recip); + QuantAData[m * BlockCountK * kBlkLen128 + k + kk] = + static_cast(std::clamp(q, -127.0f, 127.0f)); + } + } + } +} + +void +ReferenceGemm_W2_CompInt8_BlkLen128(size_t M, size_t N, size_t K, + const float* A, + const std::vector& BWeights, + const float* QuantBScale, + const uint8_t* BZeroPoints, + const float* Bias, + float* C) +{ + const size_t BlockCountK = (K + kBlkLen128 - 1) / kBlkLen128; + std::vector QuantAData(M * BlockCountK * kBlkLen128, int8_t{0}); + std::vector QuantAScale(M * BlockCountK, 0.0f); + QuantizeA_Reference_BlkLen128(M, K, A, QuantAData.data(), QuantAScale.data()); + + for (size_t m = 0; m < M; ++m) { + for (size_t n = 0; n < N; ++n) { + float acc = (Bias != nullptr) ? Bias[n] : 0.0f; + for (size_t k = 0, blk = 0; k < K; k += kBlkLen128, ++blk) { + const size_t local_len = std::min(K - k, kBlkLen128); + const float a_scale = QuantAScale[m * BlockCountK + blk]; + const float b_scale = QuantBScale[n * BlockCountK + blk]; + const int32_t zp = BZeroPoints != nullptr + ? static_cast(BZeroPoints[n * BlockCountK + blk]) + : static_cast(sq2::kDefaultSymmetricZeroPoint2Bit); + int32_t dot = 0; + for (size_t kk = 0; kk < local_len; ++kk) { + const int8_t qa = QuantAData[m * BlockCountK * kBlkLen128 + k + kk]; + const int32_t qb = + static_cast(BWeights[n * K + k + kk]) - zp; + dot += static_cast(qa) * qb; + } + acc += static_cast(dot) * a_scale * b_scale; + } + C[m * N + n] = acc; + } + } +} + +void +RunW2Case_BlkLen128(size_t M, size_t N, size_t K, bool WithBias, uint32_t seed, + bool WithZeroPoints, W2KernelFn kernel, + const char* kernel_name) +{ + const size_t BlockCountK = (K + kBlkLen128 - 1) / kBlkLen128; + ASSERT_EQ(K % kBlkLen128, 0u) << "BlkLen128 test K must be a multiple of 128"; + + std::mt19937 rng(seed); + std::uniform_real_distribution a_dist(-1.0f, 1.0f); + std::uniform_int_distribution w_dist(0, 3); + std::uniform_real_distribution s_dist(0.05f, 0.5f); + + std::vector A(M * K); + for (auto& v : A) v = a_dist(rng); + + std::vector BWeights(N * K); + for (auto& v : BWeights) v = static_cast(w_dist(rng)); + + // Source-packed B in standard ONNX layout (32 bytes per block at BlkLen=128). + std::vector QuantBData(N * BlockCountK * kBlkBytes128, std::byte{0}); + for (size_t n = 0; n < N; ++n) { + for (size_t blk = 0; blk < BlockCountK; ++blk) { + uint8_t blk_weights[kBlkLen128]; + for (size_t kk = 0; kk < kBlkLen128; ++kk) { + blk_weights[kk] = BWeights[n * K + blk * kBlkLen128 + kk]; + } + PackSourceBlock_BlkLen128(blk_weights, + QuantBData.data() + (n * BlockCountK + blk) * kBlkBytes128); + } + } + + std::vector QuantBScale(N * BlockCountK); + for (auto& v : QuantBScale) v = s_dist(rng); + + std::vector BZeroPoints; + std::vector BZeroPointsPacked; + const uint8_t* BZeroPointsRef = nullptr; + const std::byte* BZeroPointsMlas = nullptr; + if (WithZeroPoints) { + BZeroPoints.resize(N * BlockCountK); + for (auto& v : BZeroPoints) v = static_cast(w_dist(rng)); + BZeroPointsRef = BZeroPoints.data(); + BZeroPointsPacked = PackW2ZeroPoints(N, BlockCountK, BZeroPoints); + BZeroPointsMlas = BZeroPointsPacked.data(); + } + + std::vector Bias; + const float* BiasPtr = nullptr; + if (WithBias) { + Bias.resize(N); + for (auto& v : Bias) v = a_dist(rng); + BiasPtr = Bias.data(); + } + + const size_t PackedSize = sq2::Q2BitGemmPackQuantBDataSize_Avx512( + N, K, kBlkLen128, WithZeroPoints, SQNBIT_CompInt8, nullptr); + ASSERT_GT(PackedSize, 0u) << "BlkLen128 block-group pack size unsupported for shape"; + + std::vector PackedQuantBBuf(PackedSize, std::byte{0}); + PackedQuantBDataStruct packed_b( + PackedQuantBBuf.data(), N, BlockCountK, kBlkLen128, /*QuantAUnsigned=*/false); + + sq2::SQ2BitGemmPackQuantBDataAndBlkSum_Scalar( + N, K, kBlkLen128, SQNBIT_CompInt8, + QuantBData.data(), /*scales=*/nullptr, + WithZeroPoints, /*zp=*/nullptr, + packed_b, /*tp=*/nullptr, /*cfg=*/nullptr); + sq2::SQ2BitGemmPackQuantBDataAndBlkSum_Scalar( + N, K, kBlkLen128, SQNBIT_CompInt8, + /*B=*/nullptr, QuantBScale.data(), + WithZeroPoints, /*zp=*/nullptr, + packed_b, /*tp=*/nullptr, /*cfg=*/nullptr); + if (WithZeroPoints) { + sq2::SQ2BitGemmPackQuantBDataAndBlkSum_Scalar( + N, K, kBlkLen128, SQNBIT_CompInt8, + /*B=*/nullptr, /*scales=*/nullptr, + WithZeroPoints, BZeroPointsMlas, + packed_b, /*tp=*/nullptr, /*cfg=*/nullptr); + } + + std::vector QuantAData(M * BlockCountK * kBlkLen128, int8_t{0}); + std::vector QuantAScale(M * BlockCountK, 0.0f); + QuantizeA_Reference_BlkLen128(M, K, A.data(), QuantAData.data(), QuantAScale.data()); + + std::vector ABlockSum(M * BlockCountK, 0.0f); + for (size_t m = 0; m < M; ++m) { + for (size_t blk = 0; blk < BlockCountK; ++blk) { + int32_t sum = 0; + for (size_t kk = 0; kk < kBlkLen128; ++kk) { + sum += static_cast( + QuantAData[m * BlockCountK * kBlkLen128 + blk * kBlkLen128 + kk]); + } + ABlockSum[m * BlockCountK + blk] = + QuantAScale[m * BlockCountK + blk] * static_cast(sum); + } + } + + std::vector C(M * N, 0.0f); + kernel( + kBlkLen128, + reinterpret_cast(QuantAData.data()), + QuantAScale.data(), + packed_b.PackedQuantBData, + packed_b.PackedQuantBScale, + /*QuantBZeroPoint=*/nullptr, + C.data(), + M, N, /*CountK=*/K, BlockCountK, + BiasPtr, + /*ldc=*/N, + ABlockSum.data(), + packed_b.QuantBBlkSum); + + std::vector CRef(M * N, 0.0f); + ReferenceGemm_W2_CompInt8_BlkLen128(M, N, K, A.data(), BWeights, QuantBScale.data(), + BZeroPointsRef, BiasPtr, CRef.data()); + + const float abs_tol = 1e-4f; + const float rel_tol = 1e-4f; + for (size_t i = 0; i < M * N; ++i) { + const float diff = std::fabs(C[i] - CRef[i]); + const float bound = abs_tol + rel_tol * std::fabs(CRef[i]); + ASSERT_LE(diff, bound) + << "BlkLen128 " << kernel_name << " mismatch at i=" << i + << " (m=" << (i / N) << ", n=" << (i % N) << ")" + << " out=" << C[i] << " ref=" << CRef[i] + << " M=" << M << " N=" << N << " K=" << K + << " WithBias=" << WithBias + << " WithZeroPoints=" << WithZeroPoints; + } +} + +// +// Shape set for BlkLen=128. K must be a multiple of 128 (BlkLen=128 constraint). +// Covers the same regimes as the BlkLen=64 shape set: +// * R1 / R2 tiles, M=1 decode + larger M prefill +// * BlockCountK in {1, 2, 3, 4, 8, 16, 32} -- full + K-tail variants +// * N-tail (NMain=0 and various NTail) combined with K-tail +// +constexpr struct { size_t M, N, K; } kSimdShapes_BlkLen128[] = { + {1, 16, 128}, // R1, BlockCountK=1 + {1, 32, 256}, // R1, BlockCountK=2 (K-tail, no full group) + {1, 1024, 1024}, // R1, BlockCountK=8 + {1, 1024, 4096}, // R1, BlockCountK=32 (customer N) + {2, 16, 128}, + {2, 32, 256}, + {2, 64, 512}, // BlockCountK=4 (one full group) + {3, 16, 256}, // R2 head + R1 tail + {3, 384, 1024}, + {4, 16, 256}, + {4, 32, 256}, + {16, 64, 512}, + {32, 128, 256}, + // M=128 prefill at customer-like shapes (K multiples of 128) + {128, 1024, 1024}, {128, 1024, 4096}, + {128, 192, 1024}, {128, 384, 1024}, + // K-tail stress (BlockCountK not a multiple of 4) + { 2, 16, 384}, // tail=3 + { 4, 16, 384}, + {128, 1024, 384}, // tail=3 at customer M + { 2, 16, 640}, // tail=1 + { 4, 16, 640}, + { 2, 16, 768}, // tail=2 + { 4, 16, 768}, + // N-tail stress + { 1, 1, 256}, + { 1, 3, 256}, + { 4, 3, 256}, + { 1, 17, 256}, + { 4, 17, 256}, + {128, 17, 256}, + { 1, 33, 256}, + { 4, 33, 256}, + {128, 33, 256}, + { 1, 18, 256}, + { 4, 18, 256}, + {128, 19, 256}, + // N-tail combined with K-tail (most generic) + { 1, 17, 384}, + { 4, 33, 384}, + {128, 19, 640}, +}; + +} // namespace + +TEST(MlasSq2BitTest, Scalar_BlkLen128) +{ + for (uint32_t seed : {0xC0FFEEu, 0xBADC0DEu}) { + for (const auto& s : kSimdShapes_BlkLen128) { + for (bool bias : {false, true}) { + RunW2Case_BlkLen128(s.M, s.N, s.K, bias, seed + (bias ? 1u : 0u), + /*WithZeroPoints=*/false, + sq2::SQ2BitGemmKernel_BlkSum_CompInt8_Scalar, + "Scalar"); + } + } + } +} + +TEST(MlasSq2BitTest, Scalar_BlkLen128_WithZeroPoints) +{ + for (uint32_t seed : {0xC0FFEEu, 0xBADC0DEu}) { + for (const auto& s : kSimdShapes_BlkLen128) { + for (bool bias : {false, true}) { + RunW2Case_BlkLen128(s.M, s.N, s.K, bias, seed + (bias ? 1u : 0u), + /*WithZeroPoints=*/true, + sq2::SQ2BitGemmKernel_BlkSum_CompInt8_Scalar, + "Scalar"); + } + } + } +} + +TEST(MlasSq2BitTest, BlkLen128_Avx512) +{ + if (!GetMlasPlatform().Avx512Supported_) { + GTEST_SKIP() << "AVX-512BW (DQ/VL) not available on this host"; + } + for (uint32_t seed : {0xC0FFEEu, 0xBADC0DEu}) { + for (const auto& s : kSimdShapes_BlkLen128) { + for (bool bias : {false, true}) { + RunW2Case_BlkLen128(s.M, s.N, s.K, bias, seed + (bias ? 1u : 0u), + /*WithZeroPoints=*/false, + sq2::SQ2BitGemmKernel_BlkSum_CompInt8_Avx512_TestEntry, + "AVX-512BW"); + } + } + } +} + +TEST(MlasSq2BitTest, BlkLen128_Avx512_WithZeroPoints) +{ + if (!GetMlasPlatform().Avx512Supported_) { + GTEST_SKIP() << "AVX-512BW (DQ/VL) not available on this host"; + } + for (uint32_t seed : {0xC0FFEEu, 0xBADC0DEu}) { + for (const auto& s : kSimdShapes_BlkLen128) { + for (bool bias : {false, true}) { + RunW2Case_BlkLen128(s.M, s.N, s.K, bias, seed + (bias ? 1u : 0u), + /*WithZeroPoints=*/true, + sq2::SQ2BitGemmKernel_BlkSum_CompInt8_Avx512_TestEntry, + "AVX-512BW"); + } + } + } +} + +TEST(MlasSq2BitTest, BlkLen128_Avx512Vnni) +{ + if (GetMlasPlatform().QNBitGemmDispatch != &MlasSQNBitGemmDispatchAvx512vnni) { + GTEST_SKIP() << "AVX-512-VNNI not selected as the active dispatch on this host"; + } + for (uint32_t seed : {0xC0FFEEu, 0xBADC0DEu}) { + for (const auto& s : kSimdShapes_BlkLen128) { + for (bool bias : {false, true}) { + RunW2Case_BlkLen128(s.M, s.N, s.K, bias, seed + (bias ? 1u : 0u), + /*WithZeroPoints=*/false, + sq2::SQ2BitGemmKernel_BlkSum_CompInt8_Avx512Vnni_TestEntry, + "AVX-512-VNNI"); + } + } + } +} + +TEST(MlasSq2BitTest, BlkLen128_Avx512Vnni_WithZeroPoints) +{ + if (GetMlasPlatform().QNBitGemmDispatch != &MlasSQNBitGemmDispatchAvx512vnni) { + GTEST_SKIP() << "AVX-512-VNNI not selected as the active dispatch on this host"; + } + for (uint32_t seed : {0xC0FFEEu, 0xBADC0DEu}) { + for (const auto& s : kSimdShapes_BlkLen128) { + for (bool bias : {false, true}) { + RunW2Case_BlkLen128(s.M, s.N, s.K, bias, seed + (bias ? 1u : 0u), + /*WithZeroPoints=*/true, + sq2::SQ2BitGemmKernel_BlkSum_CompInt8_Avx512Vnni_TestEntry, + "AVX-512-VNNI"); + } + } + } +} + +// ============================================================================= +// BlkLen=32 coverage. Mirrors BlkLen=128 above with per-block byte width 8 +// and per-group bytes 32. Helpers duplicated (rather than templated) to keep +// the BlkLen=64 hot path bit-identical. +// ============================================================================= + +namespace { + +constexpr size_t kBlkLen32 = sq2::kBlkLen32; // 32 +constexpr size_t kBlkBytes32 = sq2::kBlkBytes32; // 8 + +void +PackSourceBlock_BlkLen32(const uint8_t weights[kBlkLen32], std::byte* src_out) +{ + for (size_t i = 0; i < kBlkBytes32; ++i) { + const uint8_t v0 = weights[4 * i + 0] & 0x03u; + const uint8_t v1 = weights[4 * i + 1] & 0x03u; + const uint8_t v2 = weights[4 * i + 2] & 0x03u; + const uint8_t v3 = weights[4 * i + 3] & 0x03u; + src_out[i] = static_cast( + static_cast(v0 | (v1 << 2) | (v2 << 4) | (v3 << 6)) + ); + } +} + +void +QuantizeA_Reference_BlkLen32(size_t M, size_t K, const float* A, + int8_t* QuantAData, float* QuantAScale) +{ + const size_t BlockCountK = (K + kBlkLen32 - 1) / kBlkLen32; + for (size_t m = 0; m < M; ++m) { + for (size_t k = 0, k_blk = 0; k < K; k += kBlkLen32, ++k_blk) { + const size_t local_len = std::min(K - k, kBlkLen32); + float amax = 0.0f; + for (size_t kk = 0; kk < local_len; ++kk) { + amax = std::max(amax, std::fabs(A[m * K + k + kk])); + } + constexpr float range_max = 127.0f; + const float scale = amax / range_max; + const float scale_recip = amax != 0.0f ? range_max / amax : 0.0f; + QuantAScale[m * BlockCountK + k_blk] = scale; + for (size_t kk = 0; kk < kBlkLen32; ++kk) { + const float a = (kk < local_len) ? A[m * K + k + kk] : 0.0f; + const float q = std::nearbyint(a * scale_recip); + QuantAData[m * BlockCountK * kBlkLen32 + k + kk] = + static_cast(std::clamp(q, -127.0f, 127.0f)); + } + } + } +} + +void +ReferenceGemm_W2_CompInt8_BlkLen32(size_t M, size_t N, size_t K, + const float* A, + const std::vector& BWeights, + const float* QuantBScale, + const uint8_t* BZeroPoints, + const float* Bias, + float* C) +{ + const size_t BlockCountK = (K + kBlkLen32 - 1) / kBlkLen32; + std::vector QuantAData(M * BlockCountK * kBlkLen32, int8_t{0}); + std::vector QuantAScale(M * BlockCountK, 0.0f); + QuantizeA_Reference_BlkLen32(M, K, A, QuantAData.data(), QuantAScale.data()); + + for (size_t m = 0; m < M; ++m) { + for (size_t n = 0; n < N; ++n) { + float acc = (Bias != nullptr) ? Bias[n] : 0.0f; + for (size_t k = 0, blk = 0; k < K; k += kBlkLen32, ++blk) { + const size_t local_len = std::min(K - k, kBlkLen32); + const float a_scale = QuantAScale[m * BlockCountK + blk]; + const float b_scale = QuantBScale[n * BlockCountK + blk]; + const int32_t zp = BZeroPoints != nullptr + ? static_cast(BZeroPoints[n * BlockCountK + blk]) + : static_cast(sq2::kDefaultSymmetricZeroPoint2Bit); + int32_t dot = 0; + for (size_t kk = 0; kk < local_len; ++kk) { + const int8_t qa = QuantAData[m * BlockCountK * kBlkLen32 + k + kk]; + const int32_t qb = + static_cast(BWeights[n * K + k + kk]) - zp; + dot += static_cast(qa) * qb; + } + acc += static_cast(dot) * a_scale * b_scale; + } + C[m * N + n] = acc; + } + } +} + +void +RunW2Case_BlkLen32(size_t M, size_t N, size_t K, bool WithBias, uint32_t seed, + bool WithZeroPoints, W2KernelFn kernel, + const char* kernel_name) +{ + const size_t BlockCountK = (K + kBlkLen32 - 1) / kBlkLen32; + ASSERT_EQ(K % kBlkLen32, 0u) << "BlkLen32 test K must be a multiple of 32"; + + std::mt19937 rng(seed); + std::uniform_real_distribution a_dist(-1.0f, 1.0f); + std::uniform_int_distribution w_dist(0, 3); + std::uniform_real_distribution s_dist(0.05f, 0.5f); + + std::vector A(M * K); + for (auto& v : A) v = a_dist(rng); + + std::vector BWeights(N * K); + for (auto& v : BWeights) v = static_cast(w_dist(rng)); + + std::vector QuantBData(N * BlockCountK * kBlkBytes32, std::byte{0}); + for (size_t n = 0; n < N; ++n) { + for (size_t blk = 0; blk < BlockCountK; ++blk) { + uint8_t blk_weights[kBlkLen32]; + for (size_t kk = 0; kk < kBlkLen32; ++kk) { + blk_weights[kk] = BWeights[n * K + blk * kBlkLen32 + kk]; + } + PackSourceBlock_BlkLen32(blk_weights, + QuantBData.data() + (n * BlockCountK + blk) * kBlkBytes32); + } + } + + std::vector QuantBScale(N * BlockCountK); + for (auto& v : QuantBScale) v = s_dist(rng); + + std::vector BZeroPoints; + std::vector BZeroPointsPacked; + const uint8_t* BZeroPointsRef = nullptr; + const std::byte* BZeroPointsMlas = nullptr; + if (WithZeroPoints) { + BZeroPoints.resize(N * BlockCountK); + for (auto& v : BZeroPoints) v = static_cast(w_dist(rng)); + BZeroPointsRef = BZeroPoints.data(); + BZeroPointsPacked = PackW2ZeroPoints(N, BlockCountK, BZeroPoints); + BZeroPointsMlas = BZeroPointsPacked.data(); + } + + std::vector Bias; + const float* BiasPtr = nullptr; + if (WithBias) { + Bias.resize(N); + for (auto& v : Bias) v = a_dist(rng); + BiasPtr = Bias.data(); + } + + const size_t PackedSize = sq2::Q2BitGemmPackQuantBDataSize_Avx512( + N, K, kBlkLen32, WithZeroPoints, SQNBIT_CompInt8, nullptr); + ASSERT_GT(PackedSize, 0u) << "BlkLen32 pack size unsupported for shape"; + + std::vector PackedQuantBBuf(PackedSize, std::byte{0}); + PackedQuantBDataStruct packed_b( + PackedQuantBBuf.data(), N, BlockCountK, kBlkLen32, /*QuantAUnsigned=*/false); + + sq2::SQ2BitGemmPackQuantBDataAndBlkSum_Scalar( + N, K, kBlkLen32, SQNBIT_CompInt8, + QuantBData.data(), /*scales=*/nullptr, + WithZeroPoints, /*zp=*/nullptr, + packed_b, /*tp=*/nullptr, /*cfg=*/nullptr); + sq2::SQ2BitGemmPackQuantBDataAndBlkSum_Scalar( + N, K, kBlkLen32, SQNBIT_CompInt8, + /*B=*/nullptr, QuantBScale.data(), + WithZeroPoints, /*zp=*/nullptr, + packed_b, /*tp=*/nullptr, /*cfg=*/nullptr); + if (WithZeroPoints) { + sq2::SQ2BitGemmPackQuantBDataAndBlkSum_Scalar( + N, K, kBlkLen32, SQNBIT_CompInt8, + /*B=*/nullptr, /*scales=*/nullptr, + WithZeroPoints, BZeroPointsMlas, + packed_b, /*tp=*/nullptr, /*cfg=*/nullptr); + } + + std::vector QuantAData(M * BlockCountK * kBlkLen32, int8_t{0}); + std::vector QuantAScale(M * BlockCountK, 0.0f); + QuantizeA_Reference_BlkLen32(M, K, A.data(), QuantAData.data(), QuantAScale.data()); + + std::vector ABlockSum(M * BlockCountK, 0.0f); + for (size_t m = 0; m < M; ++m) { + for (size_t blk = 0; blk < BlockCountK; ++blk) { + int32_t sum = 0; + for (size_t kk = 0; kk < kBlkLen32; ++kk) { + sum += static_cast( + QuantAData[m * BlockCountK * kBlkLen32 + blk * kBlkLen32 + kk]); + } + ABlockSum[m * BlockCountK + blk] = + QuantAScale[m * BlockCountK + blk] * static_cast(sum); + } + } + + std::vector C(M * N, 0.0f); + kernel( + kBlkLen32, + reinterpret_cast(QuantAData.data()), + QuantAScale.data(), + packed_b.PackedQuantBData, + packed_b.PackedQuantBScale, + /*QuantBZeroPoint=*/nullptr, + C.data(), + M, N, /*CountK=*/K, BlockCountK, + BiasPtr, + /*ldc=*/N, + ABlockSum.data(), + packed_b.QuantBBlkSum); + + std::vector CRef(M * N, 0.0f); + ReferenceGemm_W2_CompInt8_BlkLen32(M, N, K, A.data(), BWeights, QuantBScale.data(), + BZeroPointsRef, BiasPtr, CRef.data()); + + const float abs_tol = 1e-4f; + const float rel_tol = 1e-4f; + for (size_t i = 0; i < M * N; ++i) { + const float diff = std::fabs(C[i] - CRef[i]); + const float bound = abs_tol + rel_tol * std::fabs(CRef[i]); + ASSERT_LE(diff, bound) + << "BlkLen32 " << kernel_name << " mismatch at i=" << i + << " (m=" << (i / N) << ", n=" << (i % N) << ")" + << " out=" << C[i] << " ref=" << CRef[i] + << " M=" << M << " N=" << N << " K=" << K + << " WithBias=" << WithBias + << " WithZeroPoints=" << WithZeroPoints; + } +} + +// K-shape constraint: K multiple of 32 (BlkLen=32). Covers BlockCountK in +// {1, 2, 3, 4, 8, 16, 32, 64} -- both K-tail variants (BlockCountK not a +// multiple of 4) and exact block-group multiples. +constexpr struct { size_t M, N, K; } kSimdShapes_BlkLen32[] = { + {1, 16, 32}, // R1, BlockCountK=1 + {1, 32, 64}, // R1, BlockCountK=2 (K-tail, no full group) + {1, 1024, 256}, // R1, BlockCountK=8 + {1, 1024, 1024}, // R1, BlockCountK=32 + {2, 16, 32}, + {2, 32, 64}, + {2, 64, 128}, // BlockCountK=4 (one full group) + {3, 16, 64}, // R2 head + R1 tail + {3, 384, 256}, + {4, 16, 64}, + {4, 32, 64}, + {16, 64, 128}, + {32, 128, 128}, + // M=128 prefill + {128, 1024, 256}, {128, 1024, 1024}, + {128, 192, 256}, {128, 384, 256}, + // K-tail (BlockCountK not multiple of 4) + { 2, 16, 96}, // tail=3 + { 4, 16, 96}, + {128, 1024, 96}, + { 2, 16, 160}, // tail=1 + { 4, 16, 160}, + { 2, 16, 192}, // tail=2 (BlockCountK=6) + { 4, 16, 192}, + // N-tail + { 1, 1, 64}, + { 1, 3, 64}, + { 4, 3, 64}, + { 1, 17, 64}, + { 4, 17, 64}, + {128, 17, 64}, + { 1, 33, 64}, + { 4, 33, 64}, + {128, 33, 64}, + { 1, 18, 64}, + { 4, 18, 64}, + {128, 19, 64}, + // N-tail + K-tail + { 1, 17, 96}, + { 4, 33, 96}, + {128, 19, 160}, +}; + +} // namespace + +TEST(MlasSq2BitTest, Scalar_BlkLen32) +{ + for (uint32_t seed : {0xC0FFEEu, 0xBADC0DEu}) { + for (const auto& s : kSimdShapes_BlkLen32) { + for (bool bias : {false, true}) { + RunW2Case_BlkLen32(s.M, s.N, s.K, bias, seed + (bias ? 1u : 0u), + /*WithZeroPoints=*/false, + sq2::SQ2BitGemmKernel_BlkSum_CompInt8_Scalar, + "Scalar"); + } + } + } +} + +TEST(MlasSq2BitTest, Scalar_BlkLen32_WithZeroPoints) +{ + for (uint32_t seed : {0xC0FFEEu, 0xBADC0DEu}) { + for (const auto& s : kSimdShapes_BlkLen32) { + for (bool bias : {false, true}) { + RunW2Case_BlkLen32(s.M, s.N, s.K, bias, seed + (bias ? 1u : 0u), + /*WithZeroPoints=*/true, + sq2::SQ2BitGemmKernel_BlkSum_CompInt8_Scalar, + "Scalar"); + } + } + } +} + +TEST(MlasSq2BitTest, BlkLen32_Avx512) +{ + if (!GetMlasPlatform().Avx512Supported_) { + GTEST_SKIP() << "AVX-512BW (DQ/VL) not available on this host"; + } + for (uint32_t seed : {0xC0FFEEu, 0xBADC0DEu}) { + for (const auto& s : kSimdShapes_BlkLen32) { + for (bool bias : {false, true}) { + RunW2Case_BlkLen32(s.M, s.N, s.K, bias, seed + (bias ? 1u : 0u), + /*WithZeroPoints=*/false, + sq2::SQ2BitGemmKernel_BlkSum_CompInt8_Avx512_TestEntry, + "AVX-512BW"); + } + } + } +} + +TEST(MlasSq2BitTest, BlkLen32_Avx512_WithZeroPoints) +{ + if (!GetMlasPlatform().Avx512Supported_) { + GTEST_SKIP() << "AVX-512BW (DQ/VL) not available on this host"; + } + for (uint32_t seed : {0xC0FFEEu, 0xBADC0DEu}) { + for (const auto& s : kSimdShapes_BlkLen32) { + for (bool bias : {false, true}) { + RunW2Case_BlkLen32(s.M, s.N, s.K, bias, seed + (bias ? 1u : 0u), + /*WithZeroPoints=*/true, + sq2::SQ2BitGemmKernel_BlkSum_CompInt8_Avx512_TestEntry, + "AVX-512BW"); + } + } + } +} + +TEST(MlasSq2BitTest, BlkLen32_Avx512Vnni) +{ + if (GetMlasPlatform().QNBitGemmDispatch != &MlasSQNBitGemmDispatchAvx512vnni) { + GTEST_SKIP() << "AVX-512-VNNI not selected as the active dispatch on this host"; + } + for (uint32_t seed : {0xC0FFEEu, 0xBADC0DEu}) { + for (const auto& s : kSimdShapes_BlkLen32) { + for (bool bias : {false, true}) { + RunW2Case_BlkLen32(s.M, s.N, s.K, bias, seed + (bias ? 1u : 0u), + /*WithZeroPoints=*/false, + sq2::SQ2BitGemmKernel_BlkSum_CompInt8_Avx512Vnni_TestEntry, + "AVX-512-VNNI"); + } + } + } +} + +TEST(MlasSq2BitTest, BlkLen32_Avx512Vnni_WithZeroPoints) +{ + if (GetMlasPlatform().QNBitGemmDispatch != &MlasSQNBitGemmDispatchAvx512vnni) { + GTEST_SKIP() << "AVX-512-VNNI not selected as the active dispatch on this host"; + } + for (uint32_t seed : {0xC0FFEEu, 0xBADC0DEu}) { + for (const auto& s : kSimdShapes_BlkLen32) { + for (bool bias : {false, true}) { + RunW2Case_BlkLen32(s.M, s.N, s.K, bias, seed + (bias ? 1u : 0u), + /*WithZeroPoints=*/true, + sq2::SQ2BitGemmKernel_BlkSum_CompInt8_Avx512Vnni_TestEntry, + "AVX-512-VNNI"); + } + } + } +} From 268b97e1e427a2f05b54cf66a5d07cbc9e5bde42 Mon Sep 17 00:00:00 2001 From: Hari Seshadri Date: Thu, 11 Jun 2026 19:15:26 -0700 Subject: [PATCH 10/17] Missed cmake file --- cmake/onnxruntime_mlas.cmake | 4 ++++ 1 file changed, 4 insertions(+) diff --git a/cmake/onnxruntime_mlas.cmake b/cmake/onnxruntime_mlas.cmake index f700a662becba..1a5eba9de4759 100644 --- a/cmake/onnxruntime_mlas.cmake +++ b/cmake/onnxruntime_mlas.cmake @@ -244,6 +244,8 @@ function(setup_mlas_source_for_windows) ${MLAS_SRC_DIR}/sqnbitgemm_kernel_avx512_2bit.h ${MLAS_SRC_DIR}/sqnbitgemm_kernel_avx512_2bit.cpp ${MLAS_SRC_DIR}/sqnbitgemm_kernel_avx512_2bit_blklen64.h + ${MLAS_SRC_DIR}/sqnbitgemm_kernel_avx512_2bit_blklen128.h + ${MLAS_SRC_DIR}/sqnbitgemm_kernel_avx512_2bit_blklen32.h ${MLAS_SRC_DIR}/sqnbitgemm_kernel_avx512.cpp ${MLAS_SRC_DIR}/sqnbitgemm_kernel_avx512vnni.cpp ${MLAS_SRC_DIR}/qkv_quant_kernel_avx512vnni.cpp @@ -797,6 +799,8 @@ else() ${MLAS_SRC_DIR}/sqnbitgemm_kernel_avx512_2bit.h ${MLAS_SRC_DIR}/sqnbitgemm_kernel_avx512_2bit.cpp ${MLAS_SRC_DIR}/sqnbitgemm_kernel_avx512_2bit_blklen64.h + ${MLAS_SRC_DIR}/sqnbitgemm_kernel_avx512_2bit_blklen128.h + ${MLAS_SRC_DIR}/sqnbitgemm_kernel_avx512_2bit_blklen32.h ${MLAS_SRC_DIR}/sqnbitgemm_lut_kernel_avx2.h ${MLAS_SRC_DIR}/sqnbitgemm_lut_kernel_avx2.cpp ${MLAS_SRC_DIR}/rotary_embedding_kernel_avx2.h From 60132f37b2973487a578358cdae7307c12a206c3 Mon Sep 17 00:00:00 2001 From: Hari Seshadri Date: Mon, 15 Jun 2026 18:18:58 -0700 Subject: [PATCH 11/17] Fixes --- .../test/contrib_ops/matmul_2bits_test.cc | 71 + .../test/mlas/bench/bench_qnbitgemm.cpp | 19 +- .../mlas/unittest/test_sqnbitgemm_2bit.cpp | 688 +++--- .../unittest/test_sqnbitgemm_2bit_gemm.cpp | 1976 ++++++++--------- 4 files changed, 1399 insertions(+), 1355 deletions(-) diff --git a/onnxruntime/test/contrib_ops/matmul_2bits_test.cc b/onnxruntime/test/contrib_ops/matmul_2bits_test.cc index 2a6b461c63a7a..1cdb6d4b247e4 100644 --- a/onnxruntime/test/contrib_ops/matmul_2bits_test.cc +++ b/onnxruntime/test/contrib_ops/matmul_2bits_test.cc @@ -981,6 +981,77 @@ TEST(MatMul2Bits, MLFloat16_2b_MLFloat16ZP_Fallback) { .RunWithConfig(); } +// MLFloat16 activation + uint8 zero point fallback test. +// Exercises the nbits_ == 2 arm of MatMulNBits::ComputeBUnpacked +// (MlasDequantizeBlockwise) for the uint8-ZP path. +TEST(MatMul2Bits, MLFloat16_2b_Uint8ZP_Fallback) { + RandomValueGenerator random{1234}; + const int64_t M = 1, N = 32, K = 32, block_size = 16; + std::vector input0_fp32_vals(random.Gaussian(AsSpan({M, K}), 0.0f, 0.25f)); + std::vector input1_fp32_vals(random.Gaussian(AsSpan({K, N}), 0.0f, 0.25f)); + + int q_rows, q_cols; + MlasBlockwiseQuantizedShape(static_cast(block_size), true, + static_cast(K), static_cast(N), + q_rows, q_cols); + size_t q_data_size_in_bytes, q_scale_size, q_zp_size_in_bytes; + MlasBlockwiseQuantizedBufferSizes(static_cast(block_size), true, + static_cast(K), static_cast(N), + q_data_size_in_bytes, q_scale_size, &q_zp_size_in_bytes); + + std::vector input1_vals(q_data_size_in_bytes); + std::vector scales(q_scale_size); + std::vector zero_points(q_zp_size_in_bytes); + + auto& ortenv = **ort_env.get(); + onnxruntime::concurrency::ThreadPool* tp = ortenv.GetEnvironment().GetIntraOpThreadPool(); + + MlasQuantizeBlockwise( + input1_vals.data(), scales.data(), zero_points.data(), + input1_fp32_vals.data(), static_cast(block_size), + true, static_cast(K), static_cast(N), + static_cast(N), tp); + + // Reference dequant via MLAS (matches the kernel's MlasDequantizeBlockwise arm). + MlasDequantizeBlockwise( + input1_fp32_vals.data(), input1_vals.data(), scales.data(), zero_points.data(), + static_cast(block_size), true, + static_cast(K), static_cast(N), tp); + + std::vector expected_vals(M * N); + for (int64_t m = 0; m < M; m++) { + for (int64_t n = 0; n < N; n++) { + float sum = 0.0f; + for (int64_t k = 0; k < K; k++) { + sum += input0_fp32_vals[m * K + k] * input1_fp32_vals[n * K + k]; + } + expected_vals[m * N + n] = sum; + } + } + + int64_t k_blocks = (K + block_size - 1) / block_size; + + OpTester test("MatMulNBits", 1, kMSDomain); + test.AddAttribute("K", K); + test.AddAttribute("N", N); + test.AddAttribute("block_size", block_size); + test.AddAttribute("bits", QBits); + test.AddAttribute("accuracy_level", static_cast(0)); + + test.AddInput("A", {M, K}, FloatsToMLFloat16s(input0_fp32_vals), false); + test.AddInput("B", {q_cols, k_blocks, q_rows / k_blocks}, input1_vals, true); + test.AddInput("scales", {N, k_blocks}, FloatsToMLFloat16s(scales), true); + test.AddInput("zero_points", + {N, static_cast(q_zp_size_in_bytes) / N}, zero_points, true); + + test.AddOutput("Y", {M, N}, FloatsToMLFloat16s(expected_vals)); + test.SetOutputAbsErr("Y", 0.1f); + test.SetOutputRelErr("Y", 0.02f); + + test.ConfigEp(DefaultCpuExecutionProvider()) + .RunWithConfig(); +} + TEST(MatMul2Bits, Float32_2b_Accuracy0) { TestMatMul2BitsTyped(); TestMatMul2BitsTyped(); diff --git a/onnxruntime/test/mlas/bench/bench_qnbitgemm.cpp b/onnxruntime/test/mlas/bench/bench_qnbitgemm.cpp index 96d2dcef3a6d9..0b143d41cb29f 100644 --- a/onnxruntime/test/mlas/bench/bench_qnbitgemm.cpp +++ b/onnxruntime/test/mlas/bench/bench_qnbitgemm.cpp @@ -141,16 +141,17 @@ static void QNBit2BitArgs(benchmark::internal::Benchmark* b) { b->ArgNames({"BlkLen", "M", "N", "K", "Threads", "Symmetric", "HasBias", "ComputeType"}); b->ArgsProduct({ - {32, 64, 128}, // BlkLen (W2 native kernel supports all three) - {1, 4096}, // M (decode + prefill) - {4096, 11008}, // N - {4096, 11008}, // K - {1, 8}, // Threads - {int64_t{true}}, // Symmetric (W2 native kernel constraint) - {int64_t{false}, int64_t{true}}, // HasBias - {int64_t{SQNBIT_CompInt8}}, // ComputeType (W2 native kernel constraint) + {32, 64, 128}, // BlkLen (W2 native kernel supports all three) + {1, 4096}, // M (decode + prefill) + {4096, 11008}, // N + {4096, 11008}, // K + {1, 8}, // Threads + {int64_t{true}}, // Symmetric (W2 native kernel constraint) + {int64_t{false}, int64_t{true}}, // HasBias + {int64_t{SQNBIT_CompInt8}}, // ComputeType (W2 native kernel constraint) }); -}BENCHMARK(QNBITGEMM)->Apply(QNBitGemmArgs)->UseRealTime(); +} +BENCHMARK(QNBITGEMM)->Apply(QNBitGemmArgs)->UseRealTime(); BENCHMARK(QNBITGEMM)->Apply(QNBitGemmArgs)->UseRealTime(); BENCHMARK(QNBITGEMM)->Apply(QNBitGemmArgs)->UseRealTime(); BENCHMARK(QNBITGEMM)->Apply(QNBit2BitArgs)->UseRealTime(); diff --git a/onnxruntime/test/mlas/unittest/test_sqnbitgemm_2bit.cpp b/onnxruntime/test/mlas/unittest/test_sqnbitgemm_2bit.cpp index bb77139e1ff36..d510378b9e287 100644 --- a/onnxruntime/test/mlas/unittest/test_sqnbitgemm_2bit.cpp +++ b/onnxruntime/test/mlas/unittest/test_sqnbitgemm_2bit.cpp @@ -30,18 +30,15 @@ namespace { namespace sq2 = onnxruntime::mlas::sq2bit_avx512; // Standard ONNX 2-bit packing: byte_i = w[4i] | w[4i+1]<<2 | w[4i+2]<<4 | w[4i+3]<<6. -void -PackSourceBlock_BlkLen64(const uint8_t weights[sq2::kBlkLen], std::byte* src_out) -{ - for (size_t i = 0; i < sq2::kBlkBytes; ++i) { - const uint8_t v0 = weights[4 * i + 0] & 0x03u; - const uint8_t v1 = weights[4 * i + 1] & 0x03u; - const uint8_t v2 = weights[4 * i + 2] & 0x03u; - const uint8_t v3 = weights[4 * i + 3] & 0x03u; - src_out[i] = static_cast( - static_cast(v0 | (v1 << 2) | (v2 << 4) | (v3 << 6)) - ); - } +void PackSourceBlock_BlkLen64(const uint8_t weights[sq2::kBlkLen], std::byte* src_out) { + for (size_t i = 0; i < sq2::kBlkBytes; ++i) { + const uint8_t v0 = weights[4 * i + 0] & 0x03u; + const uint8_t v1 = weights[4 * i + 1] & 0x03u; + const uint8_t v2 = weights[4 * i + 2] & 0x03u; + const uint8_t v3 = weights[4 * i + 3] & 0x03u; + src_out[i] = static_cast( + static_cast(v0 | (v1 << 2) | (v2 << 4) | (v3 << 6))); + } } } // namespace @@ -59,78 +56,76 @@ PackSourceBlock_BlkLen64(const uint8_t weights[sq2::kBlkLen], std::byte* src_out // value ((i + k) % 4), giving every (block_index, position, value) a unique // fingerprint that pinpoints a layout swap if any. // -TEST(MlasSq2BitTest, PackUnpackRoundTrip_BlockGroup_BlkLen64_DeterministicPattern) -{ +TEST(MlasSq2BitTest, PackUnpackRoundTrip_BlockGroup_BlkLen64_DeterministicPattern) { + std::array, sq2::kBlockGroupBlks> weights{}; + for (size_t k = 0; k < sq2::kBlockGroupBlks; ++k) { + for (size_t i = 0; i < sq2::kBlkLen; ++i) { + weights[k][i] = static_cast((i + k) % 4); + } + } + + std::array, sq2::kBlockGroupBlks> src{}; + for (size_t k = 0; k < sq2::kBlockGroupBlks; ++k) { + PackSourceBlock_BlkLen64(weights[k].data(), src[k].data()); + } + + std::array packed{}; + sq2::PackBlockGroup_BlkLen64(src[0].data(), src[1].data(), src[2].data(), src[3].data(), + packed.data()); + + std::array, sq2::kBlockGroupBlks> recovered{}; + sq2::UnPackBlockGroup_BlkLen64_Reference(packed.data(), + recovered[0].data(), recovered[1].data(), + recovered[2].data(), recovered[3].data()); + + for (size_t k = 0; k < sq2::kBlockGroupBlks; ++k) { + for (size_t i = 0; i < sq2::kBlkLen; ++i) { + ASSERT_EQ(recovered[k][i], weights[k][i]) + << "block-group round-trip mismatch k=" << k << " i=" << i; + } + } +} + +// +// Randomized block-group round-trip across several seeds. The four input +// blocks are independent random fills; the test fails fast if any (block, +// weight) entry is mis-routed by the packed-byte layout. +// +TEST(MlasSq2BitTest, PackUnpackRoundTrip_BlockGroup_BlkLen64_Randomized) { + constexpr unsigned kSeeds = 8; + + for (unsigned seed = 0; seed < kSeeds; ++seed) { + std::mt19937 rng(seed * 5051u + 13u); + std::uniform_int_distribution dist(0u, 3u); + std::array, sq2::kBlockGroupBlks> weights{}; - for (size_t k = 0; k < sq2::kBlockGroupBlks; ++k) { - for (size_t i = 0; i < sq2::kBlkLen; ++i) { - weights[k][i] = static_cast((i + k) % 4); - } + for (auto& blk : weights) { + for (auto& w : blk) { + w = static_cast(dist(rng)); + } } std::array, sq2::kBlockGroupBlks> src{}; for (size_t k = 0; k < sq2::kBlockGroupBlks; ++k) { - PackSourceBlock_BlkLen64(weights[k].data(), src[k].data()); + PackSourceBlock_BlkLen64(weights[k].data(), src[k].data()); } std::array packed{}; sq2::PackBlockGroup_BlkLen64(src[0].data(), src[1].data(), src[2].data(), src[3].data(), - packed.data()); + packed.data()); std::array, sq2::kBlockGroupBlks> recovered{}; sq2::UnPackBlockGroup_BlkLen64_Reference(packed.data(), - recovered[0].data(), recovered[1].data(), - recovered[2].data(), recovered[3].data()); + recovered[0].data(), recovered[1].data(), + recovered[2].data(), recovered[3].data()); for (size_t k = 0; k < sq2::kBlockGroupBlks; ++k) { - for (size_t i = 0; i < sq2::kBlkLen; ++i) { - ASSERT_EQ(recovered[k][i], weights[k][i]) - << "block-group round-trip mismatch k=" << k << " i=" << i; - } - } -} - -// -// Randomized block-group round-trip across several seeds. The four input -// blocks are independent random fills; the test fails fast if any (block, -// weight) entry is mis-routed by the packed-byte layout. -// -TEST(MlasSq2BitTest, PackUnpackRoundTrip_BlockGroup_BlkLen64_Randomized) -{ - constexpr unsigned kSeeds = 8; - - for (unsigned seed = 0; seed < kSeeds; ++seed) { - std::mt19937 rng(seed * 5051u + 13u); - std::uniform_int_distribution dist(0u, 3u); - - std::array, sq2::kBlockGroupBlks> weights{}; - for (auto& blk : weights) { - for (auto& w : blk) { - w = static_cast(dist(rng)); - } - } - - std::array, sq2::kBlockGroupBlks> src{}; - for (size_t k = 0; k < sq2::kBlockGroupBlks; ++k) { - PackSourceBlock_BlkLen64(weights[k].data(), src[k].data()); - } - - std::array packed{}; - sq2::PackBlockGroup_BlkLen64(src[0].data(), src[1].data(), src[2].data(), src[3].data(), - packed.data()); - - std::array, sq2::kBlockGroupBlks> recovered{}; - sq2::UnPackBlockGroup_BlkLen64_Reference(packed.data(), - recovered[0].data(), recovered[1].data(), - recovered[2].data(), recovered[3].data()); - - for (size_t k = 0; k < sq2::kBlockGroupBlks; ++k) { - for (size_t i = 0; i < sq2::kBlkLen; ++i) { - ASSERT_EQ(recovered[k][i], weights[k][i]) - << "Random block-group mismatch seed=" << seed << " k=" << k << " i=" << i; - } - } + for (size_t i = 0; i < sq2::kBlkLen; ++i) { + ASSERT_EQ(recovered[k][i], weights[k][i]) + << "Random block-group mismatch seed=" << seed << " k=" << k << " i=" << i; + } } + } } // @@ -140,69 +135,68 @@ TEST(MlasSq2BitTest, PackUnpackRoundTrip_BlockGroup_BlkLen64_Randomized) // - Block_k set to value v with all other blocks zero produces packed bytes // equal to (v << (2*k)) -- exclusively occupying the k-th bit slot. // -TEST(MlasSq2BitTest, PackUnpackRoundTrip_BlockGroup_BlkLen64_ConstantValues) -{ - // Case 1: every block filled with v. - for (uint8_t v = 0; v < 4; ++v) { - std::array, sq2::kBlockGroupBlks> weights{}; - for (auto& blk : weights) { - blk.fill(v); - } +TEST(MlasSq2BitTest, PackUnpackRoundTrip_BlockGroup_BlkLen64_ConstantValues) { + // Case 1: every block filled with v. + for (uint8_t v = 0; v < 4; ++v) { + std::array, sq2::kBlockGroupBlks> weights{}; + for (auto& blk : weights) { + blk.fill(v); + } - std::array, sq2::kBlockGroupBlks> src{}; - for (size_t k = 0; k < sq2::kBlockGroupBlks; ++k) { - PackSourceBlock_BlkLen64(weights[k].data(), src[k].data()); - } + std::array, sq2::kBlockGroupBlks> src{}; + for (size_t k = 0; k < sq2::kBlockGroupBlks; ++k) { + PackSourceBlock_BlkLen64(weights[k].data(), src[k].data()); + } - std::array packed{}; - sq2::PackBlockGroup_BlkLen64(src[0].data(), src[1].data(), src[2].data(), src[3].data(), - packed.data()); + std::array packed{}; + sq2::PackBlockGroup_BlkLen64(src[0].data(), src[1].data(), src[2].data(), src[3].data(), + packed.data()); - const uint8_t expected_byte = static_cast(v * 0x55u); - for (size_t i = 0; i < sq2::kBlockGroupBytes; ++i) { - ASSERT_EQ(static_cast(packed[i]), expected_byte) - << "Uniform-fill v=" << static_cast(v) << " byte_i=" << i; - } + const uint8_t expected_byte = static_cast(v * 0x55u); + for (size_t i = 0; i < sq2::kBlockGroupBytes; ++i) { + ASSERT_EQ(static_cast(packed[i]), expected_byte) + << "Uniform-fill v=" << static_cast(v) << " byte_i=" << i; } + } - // Case 2: only one block at a time carries a non-zero value. - for (size_t target_k = 0; target_k < sq2::kBlockGroupBlks; ++target_k) { - for (uint8_t v = 1; v < 4; ++v) { - std::array, sq2::kBlockGroupBlks> weights{}; - weights[target_k].fill(v); - - std::array, sq2::kBlockGroupBlks> src{}; - for (size_t k = 0; k < sq2::kBlockGroupBlks; ++k) { - PackSourceBlock_BlkLen64(weights[k].data(), src[k].data()); - } - - std::array packed{}; - sq2::PackBlockGroup_BlkLen64(src[0].data(), src[1].data(), src[2].data(), src[3].data(), - packed.data()); - - const uint8_t expected_byte = static_cast(v << (2 * target_k)); - for (size_t i = 0; i < sq2::kBlockGroupBytes; ++i) { - ASSERT_EQ(static_cast(packed[i]), expected_byte) - << "Isolated block target_k=" << target_k - << " v=" << static_cast(v) - << " byte_i=" << i; - } - - std::array, sq2::kBlockGroupBlks> recovered{}; - sq2::UnPackBlockGroup_BlkLen64_Reference( - packed.data(), - recovered[0].data(), recovered[1].data(), - recovered[2].data(), recovered[3].data()); - for (size_t k = 0; k < sq2::kBlockGroupBlks; ++k) { - const uint8_t expect_val = (k == target_k) ? v : uint8_t{0}; - for (size_t i = 0; i < sq2::kBlkLen; ++i) { - ASSERT_EQ(recovered[k][i], expect_val) - << "Isolated round-trip target_k=" << target_k - << " k=" << k << " v=" << static_cast(v) << " i=" << i; - } - } + // Case 2: only one block at a time carries a non-zero value. + for (size_t target_k = 0; target_k < sq2::kBlockGroupBlks; ++target_k) { + for (uint8_t v = 1; v < 4; ++v) { + std::array, sq2::kBlockGroupBlks> weights{}; + weights[target_k].fill(v); + + std::array, sq2::kBlockGroupBlks> src{}; + for (size_t k = 0; k < sq2::kBlockGroupBlks; ++k) { + PackSourceBlock_BlkLen64(weights[k].data(), src[k].data()); + } + + std::array packed{}; + sq2::PackBlockGroup_BlkLen64(src[0].data(), src[1].data(), src[2].data(), src[3].data(), + packed.data()); + + const uint8_t expected_byte = static_cast(v << (2 * target_k)); + for (size_t i = 0; i < sq2::kBlockGroupBytes; ++i) { + ASSERT_EQ(static_cast(packed[i]), expected_byte) + << "Isolated block target_k=" << target_k + << " v=" << static_cast(v) + << " byte_i=" << i; + } + + std::array, sq2::kBlockGroupBlks> recovered{}; + sq2::UnPackBlockGroup_BlkLen64_Reference( + packed.data(), + recovered[0].data(), recovered[1].data(), + recovered[2].data(), recovered[3].data()); + for (size_t k = 0; k < sq2::kBlockGroupBlks; ++k) { + const uint8_t expect_val = (k == target_k) ? v : uint8_t{0}; + for (size_t i = 0; i < sq2::kBlkLen; ++i) { + ASSERT_EQ(recovered[k][i], expect_val) + << "Isolated round-trip target_k=" << target_k + << " k=" << k << " v=" << static_cast(v) << " i=" << i; } + } } + } } // ----------------------------------------------------------------------------- @@ -215,33 +209,65 @@ TEST(MlasSq2BitTest, PackUnpackRoundTrip_BlockGroup_BlkLen64_ConstantValues) namespace { // Standard ONNX 2-bit source packing for a BlkLen=128 block (32 bytes). -void -PackSourceBlock_BlkLen128(const uint8_t weights[sq2::kBlkLen128], std::byte* src_out) -{ - for (size_t i = 0; i < sq2::kBlkBytes128; ++i) { - const uint8_t v0 = weights[4 * i + 0] & 0x03u; - const uint8_t v1 = weights[4 * i + 1] & 0x03u; - const uint8_t v2 = weights[4 * i + 2] & 0x03u; - const uint8_t v3 = weights[4 * i + 3] & 0x03u; - src_out[i] = static_cast( - static_cast(v0 | (v1 << 2) | (v2 << 4) | (v3 << 6)) - ); - } +void PackSourceBlock_BlkLen128(const uint8_t weights[sq2::kBlkLen128], std::byte* src_out) { + for (size_t i = 0; i < sq2::kBlkBytes128; ++i) { + const uint8_t v0 = weights[4 * i + 0] & 0x03u; + const uint8_t v1 = weights[4 * i + 1] & 0x03u; + const uint8_t v2 = weights[4 * i + 2] & 0x03u; + const uint8_t v3 = weights[4 * i + 3] & 0x03u; + src_out[i] = static_cast( + static_cast(v0 | (v1 << 2) | (v2 << 4) | (v3 << 6))); + } } } // namespace -TEST(MlasSq2BitTest, PackUnpackRoundTrip_BlockGroup_BlkLen128_DeterministicPattern) -{ +TEST(MlasSq2BitTest, PackUnpackRoundTrip_BlockGroup_BlkLen128_DeterministicPattern) { + std::array, sq2::kBlockGroupBlks> weights{}; + for (size_t k = 0; k < sq2::kBlockGroupBlks; ++k) { + for (size_t i = 0; i < sq2::kBlkLen128; ++i) { + weights[k][i] = static_cast((i + k) % 4); + } + } + + std::array, sq2::kBlockGroupBlks> src{}; + for (size_t k = 0; k < sq2::kBlockGroupBlks; ++k) { + PackSourceBlock_BlkLen128(weights[k].data(), src[k].data()); + } + + std::array packed{}; + sq2::PackBlockGroup_BlkLen128(src[0].data(), src[1].data(), src[2].data(), src[3].data(), + packed.data()); + + std::array, sq2::kBlockGroupBlks> recovered{}; + sq2::UnPackBlockGroup_BlkLen128_Reference(packed.data(), + recovered[0].data(), recovered[1].data(), + recovered[2].data(), recovered[3].data()); + + for (size_t k = 0; k < sq2::kBlockGroupBlks; ++k) { + for (size_t i = 0; i < sq2::kBlkLen128; ++i) { + ASSERT_EQ(recovered[k][i], weights[k][i]) + << "BlkLen128 round-trip mismatch k=" << k << " i=" << i; + } + } +} + +TEST(MlasSq2BitTest, PackUnpackRoundTrip_BlockGroup_BlkLen128_Randomized) { + constexpr unsigned kSeeds = 8; + + for (unsigned seed = 0; seed < kSeeds; ++seed) { + std::mt19937 rng(seed * 5051u + 13u); + std::uniform_int_distribution dist(0u, 3u); + std::array, sq2::kBlockGroupBlks> weights{}; - for (size_t k = 0; k < sq2::kBlockGroupBlks; ++k) { - for (size_t i = 0; i < sq2::kBlkLen128; ++i) { - weights[k][i] = static_cast((i + k) % 4); - } + for (auto& blk : weights) { + for (auto& w : blk) { + w = static_cast(dist(rng)); + } } std::array, sq2::kBlockGroupBlks> src{}; for (size_t k = 0; k < sq2::kBlockGroupBlks; ++k) { - PackSourceBlock_BlkLen128(weights[k].data(), src[k].data()); + PackSourceBlock_BlkLen128(weights[k].data(), src[k].data()); } std::array packed{}; @@ -254,113 +280,75 @@ TEST(MlasSq2BitTest, PackUnpackRoundTrip_BlockGroup_BlkLen128_DeterministicPatte recovered[2].data(), recovered[3].data()); for (size_t k = 0; k < sq2::kBlockGroupBlks; ++k) { - for (size_t i = 0; i < sq2::kBlkLen128; ++i) { - ASSERT_EQ(recovered[k][i], weights[k][i]) - << "BlkLen128 round-trip mismatch k=" << k << " i=" << i; - } + for (size_t i = 0; i < sq2::kBlkLen128; ++i) { + ASSERT_EQ(recovered[k][i], weights[k][i]) + << "BlkLen128 random round-trip seed=" << seed << " k=" << k << " i=" << i; + } } + } } -TEST(MlasSq2BitTest, PackUnpackRoundTrip_BlockGroup_BlkLen128_Randomized) -{ - constexpr unsigned kSeeds = 8; - - for (unsigned seed = 0; seed < kSeeds; ++seed) { - std::mt19937 rng(seed * 5051u + 13u); - std::uniform_int_distribution dist(0u, 3u); - - std::array, sq2::kBlockGroupBlks> weights{}; - for (auto& blk : weights) { - for (auto& w : blk) { - w = static_cast(dist(rng)); - } - } - - std::array, sq2::kBlockGroupBlks> src{}; - for (size_t k = 0; k < sq2::kBlockGroupBlks; ++k) { - PackSourceBlock_BlkLen128(weights[k].data(), src[k].data()); - } - - std::array packed{}; - sq2::PackBlockGroup_BlkLen128(src[0].data(), src[1].data(), src[2].data(), src[3].data(), - packed.data()); - - std::array, sq2::kBlockGroupBlks> recovered{}; - sq2::UnPackBlockGroup_BlkLen128_Reference(packed.data(), - recovered[0].data(), recovered[1].data(), - recovered[2].data(), recovered[3].data()); - - for (size_t k = 0; k < sq2::kBlockGroupBlks; ++k) { - for (size_t i = 0; i < sq2::kBlkLen128; ++i) { - ASSERT_EQ(recovered[k][i], weights[k][i]) - << "BlkLen128 random round-trip seed=" << seed << " k=" << k << " i=" << i; - } - } +TEST(MlasSq2BitTest, PackUnpackRoundTrip_BlockGroup_BlkLen128_ConstantValues) { + // Case 1: every block filled with v. + for (uint8_t v = 0; v < 4; ++v) { + std::array, sq2::kBlockGroupBlks> weights{}; + for (auto& blk : weights) { + blk.fill(v); } -} - -TEST(MlasSq2BitTest, PackUnpackRoundTrip_BlockGroup_BlkLen128_ConstantValues) -{ - // Case 1: every block filled with v. - for (uint8_t v = 0; v < 4; ++v) { - std::array, sq2::kBlockGroupBlks> weights{}; - for (auto& blk : weights) { - blk.fill(v); - } - std::array, sq2::kBlockGroupBlks> src{}; - for (size_t k = 0; k < sq2::kBlockGroupBlks; ++k) { - PackSourceBlock_BlkLen128(weights[k].data(), src[k].data()); - } + std::array, sq2::kBlockGroupBlks> src{}; + for (size_t k = 0; k < sq2::kBlockGroupBlks; ++k) { + PackSourceBlock_BlkLen128(weights[k].data(), src[k].data()); + } - std::array packed{}; - sq2::PackBlockGroup_BlkLen128(src[0].data(), src[1].data(), src[2].data(), src[3].data(), - packed.data()); + std::array packed{}; + sq2::PackBlockGroup_BlkLen128(src[0].data(), src[1].data(), src[2].data(), src[3].data(), + packed.data()); - const uint8_t expected_byte = static_cast(v * 0x55u); - for (size_t i = 0; i < sq2::kBlockGroupBytes128; ++i) { - ASSERT_EQ(static_cast(packed[i]), expected_byte) - << "BlkLen128 uniform-fill v=" << static_cast(v) << " byte_i=" << i; - } + const uint8_t expected_byte = static_cast(v * 0x55u); + for (size_t i = 0; i < sq2::kBlockGroupBytes128; ++i) { + ASSERT_EQ(static_cast(packed[i]), expected_byte) + << "BlkLen128 uniform-fill v=" << static_cast(v) << " byte_i=" << i; } + } + + // Case 2: only one block at a time carries a non-zero value. + for (size_t target_k = 0; target_k < sq2::kBlockGroupBlks; ++target_k) { + for (uint8_t v = 1; v < 4; ++v) { + std::array, sq2::kBlockGroupBlks> weights{}; + weights[target_k].fill(v); - // Case 2: only one block at a time carries a non-zero value. - for (size_t target_k = 0; target_k < sq2::kBlockGroupBlks; ++target_k) { - for (uint8_t v = 1; v < 4; ++v) { - std::array, sq2::kBlockGroupBlks> weights{}; - weights[target_k].fill(v); - - std::array, sq2::kBlockGroupBlks> src{}; - for (size_t k = 0; k < sq2::kBlockGroupBlks; ++k) { - PackSourceBlock_BlkLen128(weights[k].data(), src[k].data()); - } - - std::array packed{}; - sq2::PackBlockGroup_BlkLen128(src[0].data(), src[1].data(), src[2].data(), src[3].data(), - packed.data()); - - const uint8_t expected_byte = static_cast(v << (2 * target_k)); - for (size_t i = 0; i < sq2::kBlockGroupBytes128; ++i) { - ASSERT_EQ(static_cast(packed[i]), expected_byte) - << "BlkLen128 isolated block target_k=" << target_k - << " v=" << static_cast(v) << " byte_i=" << i; - } - - std::array, sq2::kBlockGroupBlks> recovered{}; - sq2::UnPackBlockGroup_BlkLen128_Reference( - packed.data(), - recovered[0].data(), recovered[1].data(), - recovered[2].data(), recovered[3].data()); - for (size_t k = 0; k < sq2::kBlockGroupBlks; ++k) { - const uint8_t expect_val = (k == target_k) ? v : uint8_t{0}; - for (size_t i = 0; i < sq2::kBlkLen128; ++i) { - ASSERT_EQ(recovered[k][i], expect_val) - << "BlkLen128 isolated round-trip target_k=" << target_k - << " k=" << k << " v=" << static_cast(v) << " i=" << i; - } - } + std::array, sq2::kBlockGroupBlks> src{}; + for (size_t k = 0; k < sq2::kBlockGroupBlks; ++k) { + PackSourceBlock_BlkLen128(weights[k].data(), src[k].data()); + } + + std::array packed{}; + sq2::PackBlockGroup_BlkLen128(src[0].data(), src[1].data(), src[2].data(), src[3].data(), + packed.data()); + + const uint8_t expected_byte = static_cast(v << (2 * target_k)); + for (size_t i = 0; i < sq2::kBlockGroupBytes128; ++i) { + ASSERT_EQ(static_cast(packed[i]), expected_byte) + << "BlkLen128 isolated block target_k=" << target_k + << " v=" << static_cast(v) << " byte_i=" << i; + } + + std::array, sq2::kBlockGroupBlks> recovered{}; + sq2::UnPackBlockGroup_BlkLen128_Reference( + packed.data(), + recovered[0].data(), recovered[1].data(), + recovered[2].data(), recovered[3].data()); + for (size_t k = 0; k < sq2::kBlockGroupBlks; ++k) { + const uint8_t expect_val = (k == target_k) ? v : uint8_t{0}; + for (size_t i = 0; i < sq2::kBlkLen128; ++i) { + ASSERT_EQ(recovered[k][i], expect_val) + << "BlkLen128 isolated round-trip target_k=" << target_k + << " k=" << k << " v=" << static_cast(v) << " i=" << i; } + } } + } } // ----------------------------------------------------------------------------- @@ -370,33 +358,65 @@ TEST(MlasSq2BitTest, PackUnpackRoundTrip_BlockGroup_BlkLen128_ConstantValues) // ----------------------------------------------------------------------------- namespace { -void -PackSourceBlock_BlkLen32(const uint8_t weights[sq2::kBlkLen32], std::byte* src_out) -{ - for (size_t i = 0; i < sq2::kBlkBytes32; ++i) { - const uint8_t v0 = weights[4 * i + 0] & 0x03u; - const uint8_t v1 = weights[4 * i + 1] & 0x03u; - const uint8_t v2 = weights[4 * i + 2] & 0x03u; - const uint8_t v3 = weights[4 * i + 3] & 0x03u; - src_out[i] = static_cast( - static_cast(v0 | (v1 << 2) | (v2 << 4) | (v3 << 6)) - ); - } +void PackSourceBlock_BlkLen32(const uint8_t weights[sq2::kBlkLen32], std::byte* src_out) { + for (size_t i = 0; i < sq2::kBlkBytes32; ++i) { + const uint8_t v0 = weights[4 * i + 0] & 0x03u; + const uint8_t v1 = weights[4 * i + 1] & 0x03u; + const uint8_t v2 = weights[4 * i + 2] & 0x03u; + const uint8_t v3 = weights[4 * i + 3] & 0x03u; + src_out[i] = static_cast( + static_cast(v0 | (v1 << 2) | (v2 << 4) | (v3 << 6))); + } } } // namespace -TEST(MlasSq2BitTest, PackUnpackRoundTrip_BlockGroup_BlkLen32_DeterministicPattern) -{ +TEST(MlasSq2BitTest, PackUnpackRoundTrip_BlockGroup_BlkLen32_DeterministicPattern) { + std::array, sq2::kBlockGroupBlks> weights{}; + for (size_t k = 0; k < sq2::kBlockGroupBlks; ++k) { + for (size_t i = 0; i < sq2::kBlkLen32; ++i) { + weights[k][i] = static_cast((i + k) % 4); + } + } + + std::array, sq2::kBlockGroupBlks> src{}; + for (size_t k = 0; k < sq2::kBlockGroupBlks; ++k) { + PackSourceBlock_BlkLen32(weights[k].data(), src[k].data()); + } + + std::array packed{}; + sq2::PackBlockGroup_BlkLen32(src[0].data(), src[1].data(), src[2].data(), src[3].data(), + packed.data()); + + std::array, sq2::kBlockGroupBlks> recovered{}; + sq2::UnPackBlockGroup_BlkLen32_Reference(packed.data(), + recovered[0].data(), recovered[1].data(), + recovered[2].data(), recovered[3].data()); + + for (size_t k = 0; k < sq2::kBlockGroupBlks; ++k) { + for (size_t i = 0; i < sq2::kBlkLen32; ++i) { + ASSERT_EQ(recovered[k][i], weights[k][i]) + << "BlkLen32 round-trip mismatch k=" << k << " i=" << i; + } + } +} + +TEST(MlasSq2BitTest, PackUnpackRoundTrip_BlockGroup_BlkLen32_Randomized) { + constexpr unsigned kSeeds = 8; + + for (unsigned seed = 0; seed < kSeeds; ++seed) { + std::mt19937 rng(seed * 5051u + 13u); + std::uniform_int_distribution dist(0u, 3u); + std::array, sq2::kBlockGroupBlks> weights{}; - for (size_t k = 0; k < sq2::kBlockGroupBlks; ++k) { - for (size_t i = 0; i < sq2::kBlkLen32; ++i) { - weights[k][i] = static_cast((i + k) % 4); - } + for (auto& blk : weights) { + for (auto& w : blk) { + w = static_cast(dist(rng)); + } } std::array, sq2::kBlockGroupBlks> src{}; for (size_t k = 0; k < sq2::kBlockGroupBlks; ++k) { - PackSourceBlock_BlkLen32(weights[k].data(), src[k].data()); + PackSourceBlock_BlkLen32(weights[k].data(), src[k].data()); } std::array packed{}; @@ -409,109 +429,71 @@ TEST(MlasSq2BitTest, PackUnpackRoundTrip_BlockGroup_BlkLen32_DeterministicPatter recovered[2].data(), recovered[3].data()); for (size_t k = 0; k < sq2::kBlockGroupBlks; ++k) { - for (size_t i = 0; i < sq2::kBlkLen32; ++i) { - ASSERT_EQ(recovered[k][i], weights[k][i]) - << "BlkLen32 round-trip mismatch k=" << k << " i=" << i; - } + for (size_t i = 0; i < sq2::kBlkLen32; ++i) { + ASSERT_EQ(recovered[k][i], weights[k][i]) + << "BlkLen32 random round-trip seed=" << seed << " k=" << k << " i=" << i; + } } + } } -TEST(MlasSq2BitTest, PackUnpackRoundTrip_BlockGroup_BlkLen32_Randomized) -{ - constexpr unsigned kSeeds = 8; - - for (unsigned seed = 0; seed < kSeeds; ++seed) { - std::mt19937 rng(seed * 5051u + 13u); - std::uniform_int_distribution dist(0u, 3u); - - std::array, sq2::kBlockGroupBlks> weights{}; - for (auto& blk : weights) { - for (auto& w : blk) { - w = static_cast(dist(rng)); - } - } - - std::array, sq2::kBlockGroupBlks> src{}; - for (size_t k = 0; k < sq2::kBlockGroupBlks; ++k) { - PackSourceBlock_BlkLen32(weights[k].data(), src[k].data()); - } - - std::array packed{}; - sq2::PackBlockGroup_BlkLen32(src[0].data(), src[1].data(), src[2].data(), src[3].data(), - packed.data()); - - std::array, sq2::kBlockGroupBlks> recovered{}; - sq2::UnPackBlockGroup_BlkLen32_Reference(packed.data(), - recovered[0].data(), recovered[1].data(), - recovered[2].data(), recovered[3].data()); - - for (size_t k = 0; k < sq2::kBlockGroupBlks; ++k) { - for (size_t i = 0; i < sq2::kBlkLen32; ++i) { - ASSERT_EQ(recovered[k][i], weights[k][i]) - << "BlkLen32 random round-trip seed=" << seed << " k=" << k << " i=" << i; - } - } +TEST(MlasSq2BitTest, PackUnpackRoundTrip_BlockGroup_BlkLen32_ConstantValues) { + for (uint8_t v = 0; v < 4; ++v) { + std::array, sq2::kBlockGroupBlks> weights{}; + for (auto& blk : weights) { + blk.fill(v); } -} - -TEST(MlasSq2BitTest, PackUnpackRoundTrip_BlockGroup_BlkLen32_ConstantValues) -{ - for (uint8_t v = 0; v < 4; ++v) { - std::array, sq2::kBlockGroupBlks> weights{}; - for (auto& blk : weights) { - blk.fill(v); - } - std::array, sq2::kBlockGroupBlks> src{}; - for (size_t k = 0; k < sq2::kBlockGroupBlks; ++k) { - PackSourceBlock_BlkLen32(weights[k].data(), src[k].data()); - } + std::array, sq2::kBlockGroupBlks> src{}; + for (size_t k = 0; k < sq2::kBlockGroupBlks; ++k) { + PackSourceBlock_BlkLen32(weights[k].data(), src[k].data()); + } - std::array packed{}; - sq2::PackBlockGroup_BlkLen32(src[0].data(), src[1].data(), src[2].data(), src[3].data(), - packed.data()); + std::array packed{}; + sq2::PackBlockGroup_BlkLen32(src[0].data(), src[1].data(), src[2].data(), src[3].data(), + packed.data()); - const uint8_t expected_byte = static_cast(v * 0x55u); - for (size_t i = 0; i < sq2::kBlockGroupBytes32; ++i) { - ASSERT_EQ(static_cast(packed[i]), expected_byte) - << "BlkLen32 uniform-fill v=" << static_cast(v) << " byte_i=" << i; - } + const uint8_t expected_byte = static_cast(v * 0x55u); + for (size_t i = 0; i < sq2::kBlockGroupBytes32; ++i) { + ASSERT_EQ(static_cast(packed[i]), expected_byte) + << "BlkLen32 uniform-fill v=" << static_cast(v) << " byte_i=" << i; } + } - for (size_t target_k = 0; target_k < sq2::kBlockGroupBlks; ++target_k) { - for (uint8_t v = 1; v < 4; ++v) { - std::array, sq2::kBlockGroupBlks> weights{}; - weights[target_k].fill(v); - - std::array, sq2::kBlockGroupBlks> src{}; - for (size_t k = 0; k < sq2::kBlockGroupBlks; ++k) { - PackSourceBlock_BlkLen32(weights[k].data(), src[k].data()); - } - - std::array packed{}; - sq2::PackBlockGroup_BlkLen32(src[0].data(), src[1].data(), src[2].data(), src[3].data(), - packed.data()); - - const uint8_t expected_byte = static_cast(v << (2 * target_k)); - for (size_t i = 0; i < sq2::kBlockGroupBytes32; ++i) { - ASSERT_EQ(static_cast(packed[i]), expected_byte) - << "BlkLen32 isolated block target_k=" << target_k - << " v=" << static_cast(v) << " byte_i=" << i; - } - - std::array, sq2::kBlockGroupBlks> recovered{}; - sq2::UnPackBlockGroup_BlkLen32_Reference( - packed.data(), - recovered[0].data(), recovered[1].data(), - recovered[2].data(), recovered[3].data()); - for (size_t k = 0; k < sq2::kBlockGroupBlks; ++k) { - const uint8_t expect_val = (k == target_k) ? v : uint8_t{0}; - for (size_t i = 0; i < sq2::kBlkLen32; ++i) { - ASSERT_EQ(recovered[k][i], expect_val) - << "BlkLen32 isolated round-trip target_k=" << target_k - << " k=" << k << " v=" << static_cast(v) << " i=" << i; - } - } + for (size_t target_k = 0; target_k < sq2::kBlockGroupBlks; ++target_k) { + for (uint8_t v = 1; v < 4; ++v) { + std::array, sq2::kBlockGroupBlks> weights{}; + weights[target_k].fill(v); + + std::array, sq2::kBlockGroupBlks> src{}; + for (size_t k = 0; k < sq2::kBlockGroupBlks; ++k) { + PackSourceBlock_BlkLen32(weights[k].data(), src[k].data()); + } + + std::array packed{}; + sq2::PackBlockGroup_BlkLen32(src[0].data(), src[1].data(), src[2].data(), src[3].data(), + packed.data()); + + const uint8_t expected_byte = static_cast(v << (2 * target_k)); + for (size_t i = 0; i < sq2::kBlockGroupBytes32; ++i) { + ASSERT_EQ(static_cast(packed[i]), expected_byte) + << "BlkLen32 isolated block target_k=" << target_k + << " v=" << static_cast(v) << " byte_i=" << i; + } + + std::array, sq2::kBlockGroupBlks> recovered{}; + sq2::UnPackBlockGroup_BlkLen32_Reference( + packed.data(), + recovered[0].data(), recovered[1].data(), + recovered[2].data(), recovered[3].data()); + for (size_t k = 0; k < sq2::kBlockGroupBlks; ++k) { + const uint8_t expect_val = (k == target_k) ? v : uint8_t{0}; + for (size_t i = 0; i < sq2::kBlkLen32; ++i) { + ASSERT_EQ(recovered[k][i], expect_val) + << "BlkLen32 isolated round-trip target_k=" << target_k + << " k=" << k << " v=" << static_cast(v) << " i=" << i; } + } } + } } diff --git a/onnxruntime/test/mlas/unittest/test_sqnbitgemm_2bit_gemm.cpp b/onnxruntime/test/mlas/unittest/test_sqnbitgemm_2bit_gemm.cpp index 80a28a51e85a4..20892d2985049 100644 --- a/onnxruntime/test/mlas/unittest/test_sqnbitgemm_2bit_gemm.cpp +++ b/onnxruntime/test/mlas/unittest/test_sqnbitgemm_2bit_gemm.cpp @@ -36,55 +36,51 @@ Module Name: #include "core/mlas/lib/mlasi.h" #include "core/mlas/lib/sqnbitgemm_kernel_avx512_2bit.h" +#if defined(MLAS_TARGET_AMD64) + namespace { namespace sq2 = onnxruntime::mlas::sq2bit_avx512; -constexpr size_t kBlkLen = sq2::kBlkLen; // 64 -constexpr size_t kBlkBytes = sq2::kBlkBytes; // 16 -constexpr size_t kBlockGroupBlks = sq2::kBlockGroupBlks; // 4 +constexpr size_t kBlkLen = sq2::kBlkLen; // 64 +constexpr size_t kBlkBytes = sq2::kBlkBytes; // 16 // Standard ONNX 2-bit source packing (1 byte = 4 weights). -void -PackSourceBlock_BlkLen64(const uint8_t weights[kBlkLen], std::byte* src_out) -{ - for (size_t i = 0; i < kBlkBytes; ++i) { - const uint8_t v0 = weights[4 * i + 0] & 0x03u; - const uint8_t v1 = weights[4 * i + 1] & 0x03u; - const uint8_t v2 = weights[4 * i + 2] & 0x03u; - const uint8_t v3 = weights[4 * i + 3] & 0x03u; - src_out[i] = static_cast( - static_cast(v0 | (v1 << 2) | (v2 << 4) | (v3 << 6)) - ); - } +void PackSourceBlock_BlkLen64(const uint8_t weights[kBlkLen], std::byte* src_out) { + for (size_t i = 0; i < kBlkBytes; ++i) { + const uint8_t v0 = weights[4 * i + 0] & 0x03u; + const uint8_t v1 = weights[4 * i + 1] & 0x03u; + const uint8_t v2 = weights[4 * i + 2] & 0x03u; + const uint8_t v3 = weights[4 * i + 3] & 0x03u; + src_out[i] = static_cast( + static_cast(v0 | (v1 << 2) | (v2 << 4) | (v3 << 6))); + } } // Bit-exact mirror of MlasQNBitGemm's per-block int8 quantizer (amax/127, // round-half-to-even via std::nearbyint, scale_recip = 127/amax). -void -QuantizeA_Reference(size_t M, size_t K, const float* A, - int8_t* QuantAData, float* QuantAScale) -{ - const size_t BlockCountK = (K + kBlkLen - 1) / kBlkLen; - for (size_t m = 0; m < M; ++m) { - for (size_t k = 0, k_blk = 0; k < K; k += kBlkLen, ++k_blk) { - const size_t local_len = std::min(K - k, kBlkLen); - float amax = 0.0f; - for (size_t kk = 0; kk < local_len; ++kk) { - amax = std::max(amax, std::fabs(A[m * K + k + kk])); - } - constexpr float range_max = 127.0f; - const float scale = amax / range_max; - const float scale_recip = amax != 0.0f ? range_max / amax : 0.0f; - QuantAScale[m * BlockCountK + k_blk] = scale; - for (size_t kk = 0; kk < kBlkLen; ++kk) { - const float a = (kk < local_len) ? A[m * K + k + kk] : 0.0f; - const float q = std::nearbyint(a * scale_recip); - QuantAData[m * BlockCountK * kBlkLen + k + kk] = - static_cast(std::clamp(q, -127.0f, 127.0f)); - } - } +void QuantizeA_Reference(size_t M, size_t K, const float* A, + int8_t* QuantAData, float* QuantAScale) { + const size_t BlockCountK = (K + kBlkLen - 1) / kBlkLen; + for (size_t m = 0; m < M; ++m) { + for (size_t k = 0, k_blk = 0; k < K; k += kBlkLen, ++k_blk) { + const size_t local_len = std::min(K - k, kBlkLen); + float amax = 0.0f; + for (size_t kk = 0; kk < local_len; ++kk) { + amax = std::max(amax, std::fabs(A[m * K + k + kk])); + } + constexpr float range_max = 127.0f; + const float scale = amax / range_max; + const float scale_recip = amax != 0.0f ? range_max / amax : 0.0f; + QuantAScale[m * BlockCountK + k_blk] = scale; + for (size_t kk = 0; kk < kBlkLen; ++kk) { + const float a = (kk < local_len) ? A[m * K + k + kk] : 0.0f; + const float q = std::nearbyint(a * scale_recip); + QuantAData[m * BlockCountK * kBlkLen + k + kk] = + static_cast(std::clamp(q, -127.0f, 127.0f)); + } } + } } // @@ -92,42 +88,40 @@ QuantizeA_Reference(size_t M, size_t K, const float* A, // performs (kernel int8 GEMM + SGEMM zero-point correction collapsed into a // single direct dot of (qa * (qb - zp))). // -void -ReferenceGemm_W2_CompInt8(size_t M, size_t N, size_t K, - const float* A, - const std::vector& BWeights, // [N * K] in [0, 3] - const float* QuantBScale, // [N * BlockCountK] - const uint8_t* BZeroPoints, // [N * BlockCountK] or nullptr - const float* Bias, - float* C) -{ - const size_t BlockCountK = (K + kBlkLen - 1) / kBlkLen; - std::vector QuantAData(M * BlockCountK * kBlkLen, int8_t{0}); - std::vector QuantAScale(M * BlockCountK, 0.0f); - QuantizeA_Reference(M, K, A, QuantAData.data(), QuantAScale.data()); - - for (size_t m = 0; m < M; ++m) { - for (size_t n = 0; n < N; ++n) { - float acc = (Bias != nullptr) ? Bias[n] : 0.0f; - for (size_t k = 0, blk = 0; k < K; k += kBlkLen, ++blk) { - const size_t local_len = std::min(K - k, kBlkLen); - const float a_scale = QuantAScale[m * BlockCountK + blk]; - const float b_scale = QuantBScale[n * BlockCountK + blk]; - const int32_t zp = BZeroPoints != nullptr - ? static_cast(BZeroPoints[n * BlockCountK + blk]) - : static_cast(sq2::kDefaultSymmetricZeroPoint2Bit); - int32_t dot = 0; - for (size_t kk = 0; kk < local_len; ++kk) { - const int8_t qa = QuantAData[m * BlockCountK * kBlkLen + k + kk]; - const int32_t qb = - static_cast(BWeights[n * K + k + kk]) - zp; - dot += static_cast(qa) * qb; - } - acc += static_cast(dot) * a_scale * b_scale; - } - C[m * N + n] = acc; +void ReferenceGemm_W2_CompInt8(size_t M, size_t N, size_t K, + const float* A, + const std::vector& BWeights, // [N * K] in [0, 3] + const float* QuantBScale, // [N * BlockCountK] + const uint8_t* BZeroPoints, // [N * BlockCountK] or nullptr + const float* Bias, + float* C) { + const size_t BlockCountK = (K + kBlkLen - 1) / kBlkLen; + std::vector QuantAData(M * BlockCountK * kBlkLen, int8_t{0}); + std::vector QuantAScale(M * BlockCountK, 0.0f); + QuantizeA_Reference(M, K, A, QuantAData.data(), QuantAScale.data()); + + for (size_t m = 0; m < M; ++m) { + for (size_t n = 0; n < N; ++n) { + float acc = (Bias != nullptr) ? Bias[n] : 0.0f; + for (size_t k = 0, blk = 0; k < K; k += kBlkLen, ++blk) { + const size_t local_len = std::min(K - k, kBlkLen); + const float a_scale = QuantAScale[m * BlockCountK + blk]; + const float b_scale = QuantBScale[n * BlockCountK + blk]; + const int32_t zp = BZeroPoints != nullptr + ? static_cast(BZeroPoints[n * BlockCountK + blk]) + : static_cast(sq2::kDefaultSymmetricZeroPoint2Bit); + int32_t dot = 0; + for (size_t kk = 0; kk < local_len; ++kk) { + const int8_t qa = QuantAData[m * BlockCountK * kBlkLen + k + kk]; + const int32_t qb = + static_cast(BWeights[n * K + k + kk]) - zp; + dot += static_cast(qa) * qb; } + acc += static_cast(dot) * a_scale * b_scale; + } + C[m * N + n] = acc; } + } } // @@ -135,20 +129,19 @@ ReferenceGemm_W2_CompInt8(size_t M, size_t N, size_t K, // (4 zp per byte along K, row-major in N). // std::vector -PackW2ZeroPoints(size_t N, size_t BlockCountK, const std::vector& BZeroPoints) -{ - const size_t ZPCountK = (BlockCountK + 3) / 4; - std::vector packed(N * ZPCountK, std::byte{0}); - for (size_t n = 0; n < N; ++n) { - for (size_t blk = 0; blk < BlockCountK; ++blk) { - const uint8_t zp = BZeroPoints[n * BlockCountK + blk] & 0x03u; - const size_t byte_idx = n * ZPCountK + (blk / 4); - const size_t bit_off = (blk % 4) * 2; - packed[byte_idx] = static_cast( - static_cast(packed[byte_idx]) | (zp << bit_off)); - } +PackW2ZeroPoints(size_t N, size_t BlockCountK, const std::vector& BZeroPoints) { + const size_t ZPCountK = (BlockCountK + 3) / 4; + std::vector packed(N * ZPCountK, std::byte{0}); + for (size_t n = 0; n < N; ++n) { + for (size_t blk = 0; blk < BlockCountK; ++blk) { + const uint8_t zp = BZeroPoints[n * BlockCountK + blk] & 0x03u; + const size_t byte_idx = n * ZPCountK + (blk / 4); + const size_t bit_off = (blk % 4) * 2; + packed[byte_idx] = static_cast( + static_cast(packed[byte_idx]) | (zp << bit_off)); } - return packed; + } + return packed; } // @@ -161,152 +154,150 @@ PackW2ZeroPoints(size_t N, size_t BlockCountK, const std::vector& BZero // every block-group kernel variant honors via direct-call forwarders declared // in sqnbitgemm_kernel_avx512_2bit.h. // -using W2KernelFn = size_t (MLASCALL*)( +using W2KernelFn = size_t(MLASCALL*)( size_t, const std::byte*, const float*, const std::byte*, const float*, const std::byte*, float*, size_t, size_t, size_t, size_t, const float*, size_t, const float*, const float*); -void -RunW2Case(size_t M, size_t N, size_t K, bool WithBias, uint32_t seed, - bool WithZeroPoints, W2KernelFn kernel, - const char* kernel_name) -{ - const size_t BlockCountK = (K + kBlkLen - 1) / kBlkLen; - ASSERT_EQ(K % kBlkLen, 0u) << "Test K must be a multiple of BlkLen=64"; - // BlockCountK no longer required to be a multiple of kBlockGroupBlks -- - // the K-tail handler picks up the trailing 1-3 blocks. - - std::mt19937 rng(seed); - std::uniform_real_distribution a_dist(-1.0f, 1.0f); - std::uniform_int_distribution w_dist(0, 3); - std::uniform_real_distribution s_dist(0.05f, 0.5f); - - std::vector A(M * K); - for (auto& v : A) v = a_dist(rng); - - std::vector BWeights(N * K); - for (auto& v : BWeights) v = static_cast(w_dist(rng)); - - // Source-packed B (standard ONNX layout) -- the input to the pack helper. - std::vector QuantBData(N * BlockCountK * kBlkBytes, std::byte{0}); - for (size_t n = 0; n < N; ++n) { - for (size_t blk = 0; blk < BlockCountK; ++blk) { - uint8_t blk_weights[kBlkLen]; - for (size_t kk = 0; kk < kBlkLen; ++kk) { - blk_weights[kk] = BWeights[n * K + blk * kBlkLen + kk]; - } - PackSourceBlock_BlkLen64(blk_weights, - QuantBData.data() + (n * BlockCountK + blk) * kBlkBytes); - } - } - - std::vector QuantBScale(N * BlockCountK); - for (auto& v : QuantBScale) v = s_dist(rng); - - std::vector BZeroPoints; - std::vector BZeroPointsPacked; - const uint8_t* BZeroPointsRef = nullptr; - const std::byte* BZeroPointsMlas = nullptr; - if (WithZeroPoints) { - BZeroPoints.resize(N * BlockCountK); - for (auto& v : BZeroPoints) v = static_cast(w_dist(rng)); - BZeroPointsRef = BZeroPoints.data(); - BZeroPointsPacked = PackW2ZeroPoints(N, BlockCountK, BZeroPoints); - BZeroPointsMlas = BZeroPointsPacked.data(); +void RunW2Case(size_t M, size_t N, size_t K, bool WithBias, uint32_t seed, + bool WithZeroPoints, W2KernelFn kernel, + const char* kernel_name) { + const size_t BlockCountK = (K + kBlkLen - 1) / kBlkLen; + ASSERT_EQ(K % kBlkLen, 0u) << "Test K must be a multiple of BlkLen=64"; + // BlockCountK no longer required to be a multiple of kBlockGroupBlks -- + // the K-tail handler picks up the trailing 1-3 blocks. + + std::mt19937 rng(seed); + std::uniform_real_distribution a_dist(-1.0f, 1.0f); + std::uniform_int_distribution w_dist(0, 3); + std::uniform_real_distribution s_dist(0.05f, 0.5f); + + std::vector A(M * K); + for (auto& v : A) v = a_dist(rng); + + std::vector BWeights(N * K); + for (auto& v : BWeights) v = static_cast(w_dist(rng)); + + // Source-packed B (standard ONNX layout) -- the input to the pack helper. + std::vector QuantBData(N * BlockCountK * kBlkBytes, std::byte{0}); + for (size_t n = 0; n < N; ++n) { + for (size_t blk = 0; blk < BlockCountK; ++blk) { + uint8_t blk_weights[kBlkLen]; + for (size_t kk = 0; kk < kBlkLen; ++kk) { + blk_weights[kk] = BWeights[n * K + blk * kBlkLen + kk]; + } + PackSourceBlock_BlkLen64(blk_weights, + QuantBData.data() + (n * BlockCountK + blk) * kBlkBytes); } - - std::vector Bias; - const float* BiasPtr = nullptr; - if (WithBias) { - Bias.resize(N); - for (auto& v : Bias) v = a_dist(rng); - BiasPtr = Bias.data(); - } - - // Allocate the packed-B buffer (same total size as the production path). - const size_t PackedSize = sq2::Q2BitGemmPackQuantBDataSize_Avx512( - N, K, kBlkLen, WithZeroPoints, SQNBIT_CompInt8, nullptr); - ASSERT_GT(PackedSize, 0u) << "block-group pack size unsupported for the chosen shape"; - - std::vector PackedQuantBBuf(PackedSize, std::byte{0}); - // The W2 PackedQuantBDataStruct constructor pads BlockCountK to a multiple - // of 4 internally (see qnbitgemm.h) so the slab layout matches what the - // block-group pack helper writes regardless of whether the caller passes - // the logical or padded BlockCountK. We pass the logical value to mirror - // exactly what matmul_nbits.cc does in production. - PackedQuantBDataStruct packed_b( - PackedQuantBBuf.data(), N, BlockCountK, kBlkLen, /*QuantAUnsigned=*/false); - - // Mirror the matmul_nbits.cc prepack 3-call pattern (B, scales, ZP) so the - // pack code path is exercised exactly as the production dispatcher would. + } + + std::vector QuantBScale(N * BlockCountK); + for (auto& v : QuantBScale) v = s_dist(rng); + + std::vector BZeroPoints; + std::vector BZeroPointsPacked; + const uint8_t* BZeroPointsRef = nullptr; + const std::byte* BZeroPointsMlas = nullptr; + if (WithZeroPoints) { + BZeroPoints.resize(N * BlockCountK); + for (auto& v : BZeroPoints) v = static_cast(w_dist(rng)); + BZeroPointsRef = BZeroPoints.data(); + BZeroPointsPacked = PackW2ZeroPoints(N, BlockCountK, BZeroPoints); + BZeroPointsMlas = BZeroPointsPacked.data(); + } + + std::vector Bias; + const float* BiasPtr = nullptr; + if (WithBias) { + Bias.resize(N); + for (auto& v : Bias) v = a_dist(rng); + BiasPtr = Bias.data(); + } + + // Allocate the packed-B buffer (same total size as the production path). + const size_t PackedSize = sq2::Q2BitGemmPackQuantBDataSize_Avx512( + N, K, kBlkLen, WithZeroPoints, SQNBIT_CompInt8, nullptr); + ASSERT_GT(PackedSize, 0u) << "block-group pack size unsupported for the chosen shape"; + + std::vector PackedQuantBBuf(PackedSize, std::byte{0}); + // The W2 PackedQuantBDataStruct constructor pads BlockCountK to a multiple + // of 4 internally (see qnbitgemm.h) so the slab layout matches what the + // block-group pack helper writes regardless of whether the caller passes + // the logical or padded BlockCountK. We pass the logical value to mirror + // exactly what matmul_nbits.cc does in production. + PackedQuantBDataStruct packed_b( + PackedQuantBBuf.data(), N, BlockCountK, kBlkLen, /*QuantAUnsigned=*/false); + + // Mirror the matmul_nbits.cc prepack 3-call pattern (B, scales, ZP) so the + // pack code path is exercised exactly as the production dispatcher would. + sq2::SQ2BitGemmPackQuantBDataAndBlkSum_Scalar( + N, K, kBlkLen, SQNBIT_CompInt8, + QuantBData.data(), /*scales=*/nullptr, + WithZeroPoints, /*zp=*/nullptr, + packed_b, /*tp=*/nullptr, /*cfg=*/nullptr); + sq2::SQ2BitGemmPackQuantBDataAndBlkSum_Scalar( + N, K, kBlkLen, SQNBIT_CompInt8, + /*B=*/nullptr, QuantBScale.data(), + WithZeroPoints, /*zp=*/nullptr, + packed_b, /*tp=*/nullptr, /*cfg=*/nullptr); + if (WithZeroPoints) { sq2::SQ2BitGemmPackQuantBDataAndBlkSum_Scalar( N, K, kBlkLen, SQNBIT_CompInt8, - QuantBData.data(), /*scales=*/nullptr, - WithZeroPoints, /*zp=*/nullptr, + /*B=*/nullptr, /*scales=*/nullptr, + WithZeroPoints, BZeroPointsMlas, packed_b, /*tp=*/nullptr, /*cfg=*/nullptr); - sq2::SQ2BitGemmPackQuantBDataAndBlkSum_Scalar( - N, K, kBlkLen, SQNBIT_CompInt8, - /*B=*/nullptr, QuantBScale.data(), - WithZeroPoints, /*zp=*/nullptr, - packed_b, /*tp=*/nullptr, /*cfg=*/nullptr); - if (WithZeroPoints) { - sq2::SQ2BitGemmPackQuantBDataAndBlkSum_Scalar( - N, K, kBlkLen, SQNBIT_CompInt8, - /*B=*/nullptr, /*scales=*/nullptr, - WithZeroPoints, BZeroPointsMlas, - packed_b, /*tp=*/nullptr, /*cfg=*/nullptr); - } - - // Quantize A the same way MLAS would (per-block amax/127, banker rounding). - std::vector QuantAData(M * BlockCountK * kBlkLen, int8_t{0}); - std::vector QuantAScale(M * BlockCountK, 0.0f); - QuantizeA_Reference(M, K, A.data(), QuantAData.data(), QuantAScale.data()); - - std::vector ABlockSum(M * BlockCountK, 0.0f); - for (size_t m = 0; m < M; ++m) { - for (size_t blk = 0; blk < BlockCountK; ++blk) { - int32_t sum = 0; - for (size_t kk = 0; kk < kBlkLen; ++kk) { - sum += static_cast( - QuantAData[m * BlockCountK * kBlkLen + blk * kBlkLen + kk]); - } - ABlockSum[m * BlockCountK + blk] = - QuantAScale[m * BlockCountK + blk] * static_cast(sum); - } - } - - std::vector C(M * N, 0.0f); - kernel( - kBlkLen, - reinterpret_cast(QuantAData.data()), - QuantAScale.data(), - packed_b.PackedQuantBData, - packed_b.PackedQuantBScale, - /*QuantBZeroPoint=*/nullptr, - C.data(), - M, N, /*CountK=*/K, BlockCountK, - BiasPtr, - /*ldc=*/N, - ABlockSum.data(), - packed_b.QuantBBlkSum); - - std::vector CRef(M * N, 0.0f); - ReferenceGemm_W2_CompInt8(M, N, K, A.data(), BWeights, QuantBScale.data(), - BZeroPointsRef, BiasPtr, CRef.data()); - - const float abs_tol = 1e-4f; - const float rel_tol = 1e-4f; - for (size_t i = 0; i < M * N; ++i) { - const float diff = std::fabs(C[i] - CRef[i]); - const float bound = abs_tol + rel_tol * std::fabs(CRef[i]); - ASSERT_LE(diff, bound) - << "block-group " << kernel_name << " mismatch at i=" << i - << " (m=" << (i / N) << ", n=" << (i % N) << ")" - << " out=" << C[i] << " ref=" << CRef[i] - << " M=" << M << " N=" << N << " K=" << K - << " WithBias=" << WithBias - << " WithZeroPoints=" << WithZeroPoints; + } + + // Quantize A the same way MLAS would (per-block amax/127, banker rounding). + std::vector QuantAData(M * BlockCountK * kBlkLen, int8_t{0}); + std::vector QuantAScale(M * BlockCountK, 0.0f); + QuantizeA_Reference(M, K, A.data(), QuantAData.data(), QuantAScale.data()); + + std::vector ABlockSum(M * BlockCountK, 0.0f); + for (size_t m = 0; m < M; ++m) { + for (size_t blk = 0; blk < BlockCountK; ++blk) { + int32_t sum = 0; + for (size_t kk = 0; kk < kBlkLen; ++kk) { + sum += static_cast( + QuantAData[m * BlockCountK * kBlkLen + blk * kBlkLen + kk]); + } + ABlockSum[m * BlockCountK + blk] = + QuantAScale[m * BlockCountK + blk] * static_cast(sum); } + } + + std::vector C(M * N, 0.0f); + kernel( + kBlkLen, + reinterpret_cast(QuantAData.data()), + QuantAScale.data(), + packed_b.PackedQuantBData, + packed_b.PackedQuantBScale, + /*QuantBZeroPoint=*/nullptr, + C.data(), + M, N, /*CountK=*/K, BlockCountK, + BiasPtr, + /*ldc=*/N, + ABlockSum.data(), + packed_b.QuantBBlkSum); + + std::vector CRef(M * N, 0.0f); + ReferenceGemm_W2_CompInt8(M, N, K, A.data(), BWeights, QuantBScale.data(), + BZeroPointsRef, BiasPtr, CRef.data()); + + const float abs_tol = 1e-4f; + const float rel_tol = 1e-4f; + for (size_t i = 0; i < M * N; ++i) { + const float diff = std::fabs(C[i] - CRef[i]); + const float bound = abs_tol + rel_tol * std::fabs(CRef[i]); + ASSERT_LE(diff, bound) + << "block-group " << kernel_name << " mismatch at i=" << i + << " (m=" << (i / N) << ", n=" << (i % N) << ")" + << " out=" << C[i] << " ref=" << CRef[i] + << " M=" << M << " N=" << N << " K=" << K + << " WithBias=" << WithBias + << " WithZeroPoints=" << WithZeroPoints; + } } } // namespace @@ -318,64 +309,82 @@ RunW2Case(size_t M, size_t N, size_t K, bool WithBias, uint32_t seed, // is NOT a multiple of 256 so it's excluded; that shape will need a tail // handler in a follow-up. // -TEST(MlasSq2BitTest, Scalar_BlkLen64) -{ - struct Shape { size_t M, N, K; }; - constexpr Shape shapes[] = { - {1, 16, 256}, - {1, 32, 256}, - {1, 64, 512}, - {4, 16, 256}, - {4, 33, 256}, - {7, 17, 256}, - {16, 64, 512}, - {32, 128, 256}, - // Customer prefill (only the K values that are multiples of 256). - { 1, 1024, 1024}, { 1, 192, 1024}, { 1, 384, 1024}, { 1, 4096, 1024}, { 1, 1024, 4096}, - {128, 1024, 1024}, {128, 192, 1024}, {128, 384, 1024}, {128, 4096, 1024}, {128, 1024, 4096}, - }; - - for (uint32_t seed : {0xC0FFEEu, 0xBADC0DEu}) { - for (const Shape& s : shapes) { - for (bool bias : {false, true}) { - RunW2Case(s.M, s.N, s.K, bias, seed + (bias ? 1u : 0u), - /*WithZeroPoints=*/false, - sq2::SQ2BitGemmKernel_BlkSum_CompInt8_Scalar, - "scalar"); - } - } +TEST(MlasSq2BitTest, Scalar_BlkLen64) { + struct Shape { + size_t M, N, K; + }; + constexpr Shape shapes[] = { + {1, 16, 256}, + {1, 32, 256}, + {1, 64, 512}, + {4, 16, 256}, + {4, 33, 256}, + {7, 17, 256}, + {16, 64, 512}, + {32, 128, 256}, + // Customer prefill (only the K values that are multiples of 256). + {1, 1024, 1024}, + {1, 192, 1024}, + {1, 384, 1024}, + {1, 4096, 1024}, + {1, 1024, 4096}, + {128, 1024, 1024}, + {128, 192, 1024}, + {128, 384, 1024}, + {128, 4096, 1024}, + {128, 1024, 4096}, + }; + + for (uint32_t seed : {0xC0FFEEu, 0xBADC0DEu}) { + for (const Shape& s : shapes) { + for (bool bias : {false, true}) { + RunW2Case(s.M, s.N, s.K, bias, seed + (bias ? 1u : 0u), + /*WithZeroPoints=*/false, + sq2::SQ2BitGemmKernel_BlkSum_CompInt8_Scalar, + "scalar"); + } } + } } // // Same coverage with per-block non-default zero points. // -TEST(MlasSq2BitTest, Scalar_BlkLen64_WithZeroPoints) -{ - struct Shape { size_t M, N, K; }; - constexpr Shape shapes[] = { - {1, 16, 256}, - {1, 32, 256}, - {1, 64, 512}, - {4, 16, 256}, - {4, 33, 256}, - {7, 17, 256}, - {16, 64, 512}, - {32, 128, 256}, - { 1, 1024, 1024}, { 1, 192, 1024}, { 1, 384, 1024}, { 1, 4096, 1024}, { 1, 1024, 4096}, - {128, 1024, 1024}, {128, 192, 1024}, {128, 384, 1024}, {128, 4096, 1024}, {128, 1024, 4096}, - }; - - for (uint32_t seed : {0xC0FFEEu, 0xBADC0DEu}) { - for (const Shape& s : shapes) { - for (bool bias : {false, true}) { - RunW2Case(s.M, s.N, s.K, bias, seed + (bias ? 1u : 0u), - /*WithZeroPoints=*/true, - sq2::SQ2BitGemmKernel_BlkSum_CompInt8_Scalar, - "scalar"); - } - } +TEST(MlasSq2BitTest, Scalar_BlkLen64_WithZeroPoints) { + struct Shape { + size_t M, N, K; + }; + constexpr Shape shapes[] = { + {1, 16, 256}, + {1, 32, 256}, + {1, 64, 512}, + {4, 16, 256}, + {4, 33, 256}, + {7, 17, 256}, + {16, 64, 512}, + {32, 128, 256}, + {1, 1024, 1024}, + {1, 192, 1024}, + {1, 384, 1024}, + {1, 4096, 1024}, + {1, 1024, 4096}, + {128, 1024, 1024}, + {128, 192, 1024}, + {128, 384, 1024}, + {128, 4096, 1024}, + {128, 1024, 4096}, + }; + + for (uint32_t seed : {0xC0FFEEu, 0xBADC0DEu}) { + for (const Shape& s : shapes) { + for (bool bias : {false, true}) { + RunW2Case(s.M, s.N, s.K, bias, seed + (bias ? 1u : 0u), + /*WithZeroPoints=*/true, + sq2::SQ2BitGemmKernel_BlkSum_CompInt8_Scalar, + "scalar"); + } } + } } // @@ -400,129 +409,130 @@ TEST(MlasSq2BitTest, Scalar_BlkLen64_WithZeroPoints) // the right packed-B address regardless of K % 4. Customer K=384 and the // synthetic K=320, K=448 shapes exercise this path. // -constexpr struct { size_t M, N, K; } kSimdShapes[] = { - {1, 16, 256}, // R1 only - {1, 192, 1024}, // R1 only, customer N - {1, 1024, 4096}, // R1 only, customer N - {2, 16, 256}, - {2, 32, 256}, - {2, 64, 512}, - {3, 16, 256}, // R2 head (1 pair) + R1 tail - {3, 384, 1024}, - {4, 16, 256}, - {4, 32, 256}, - {5, 64, 512}, // R2 head (2 pairs) + R1 tail - {16, 64, 512}, - {32, 128, 256}, +constexpr struct { + size_t M, N, K; +} kSimdShapes[] = { + {1, 16, 256}, // R1 only + {1, 192, 1024}, // R1 only, customer N + {1, 1024, 4096}, // R1 only, customer N + {2, 16, 256}, + {2, 32, 256}, + {2, 64, 512}, + {3, 16, 256}, // R2 head (1 pair) + R1 tail + {3, 384, 1024}, + {4, 16, 256}, + {4, 32, 256}, + {5, 64, 512}, // R2 head (2 pairs) + R1 tail + {16, 64, 512}, + {32, 128, 256}, // Customer prefill (M=128) at all (K, N) pairs, including K=384. - {128, 1024, 384}, // K-tail: BlockCountK=6, 1 full group + tail of 2 blocks - {128, 1024, 1024}, {128, 192, 1024}, {128, 384, 1024}, - {128, 4096, 1024}, {128, 1024, 4096}, + {128, 1024, 384}, // K-tail: BlockCountK=6, 1 full group + tail of 2 blocks + {128, 1024, 1024}, + {128, 192, 1024}, + {128, 384, 1024}, + {128, 4096, 1024}, + {128, 1024, 4096}, // Customer decode (M=1) at K=384 (the case the K%4 gate previously blocked). - { 1, 1024, 384}, + {1, 1024, 384}, // Synthetic K-tail stress shapes covering all (TailBlocks in {1, 2, 3}). - { 2, 16, 320}, // tail=1 - { 4, 16, 320}, - {128, 1024, 320}, - { 2, 16, 448}, // tail=3 - { 4, 16, 448}, - {128, 1024, 448}, + {2, 16, 320}, // tail=1 + {4, 16, 320}, + {128, 1024, 320}, + {2, 16, 448}, // tail=3 + {4, 16, 448}, + {128, 1024, 448}, // N-tail stress (CountN % 4 != 0). The R2/R1 main tiles handle the // NMain = floor(CountN/4)*4 cols; the per-1-col tail tile picks up // the trailing 1-3 cols against the column-major tail region of the // packed buffer. NMain = 0 cases (N in {1,2,3}) exercise the tail // tile in isolation. - { 1, 1, 256}, // NMain=0, NTail=1, single-column decode - { 1, 3, 256}, // NMain=0, NTail=3 - { 4, 3, 256}, // NMain=0, NTail=3, R2+R1 head still empty - { 1, 17, 256}, // NMain=16, NTail=1, decode - { 4, 17, 256}, - {128, 17, 256}, - { 1, 33, 256}, // NMain=32, NTail=1 - { 4, 33, 256}, // exact shape that failed the dispatch swap - {128, 33, 256}, - { 1, 18, 256}, // NMain=16, NTail=2 - { 4, 18, 256}, - {128, 19, 256}, // NMain=16, NTail=3 + {1, 1, 256}, // NMain=0, NTail=1, single-column decode + {1, 3, 256}, // NMain=0, NTail=3 + {4, 3, 256}, // NMain=0, NTail=3, R2+R1 head still empty + {1, 17, 256}, // NMain=16, NTail=1, decode + {4, 17, 256}, + {128, 17, 256}, + {1, 33, 256}, // NMain=32, NTail=1 + {4, 33, 256}, // exact shape that failed the dispatch swap + {128, 33, 256}, + {1, 18, 256}, // NMain=16, NTail=2 + {4, 18, 256}, + {128, 19, 256}, // NMain=16, NTail=3 // N-tail combined with K-tail (the most generic case). - { 1, 17, 384}, - { 4, 33, 384}, - {128, 19, 448}, + {1, 17, 384}, + {4, 33, 384}, + {128, 19, 448}, }; // // AVX-512BW (non-VNNI) SIMD block-group kernel. // -TEST(MlasSq2BitTest, BlkLen64_Avx512) -{ - if (!GetMlasPlatform().Avx512Supported_) { - GTEST_SKIP() << "AVX-512BW (DQ/VL) not available on this host"; - } - for (uint32_t seed : {0xC0FFEEu, 0xBADC0DEu}) { - for (const auto& s : kSimdShapes) { - for (bool bias : {false, true}) { - RunW2Case(s.M, s.N, s.K, bias, seed + (bias ? 1u : 0u), - /*WithZeroPoints=*/false, - sq2::SQ2BitGemmKernel_BlkSum_CompInt8_Avx512_TestEntry, - "AVX-512BW"); - } - } +TEST(MlasSq2BitTest, BlkLen64_Avx512) { + if (!GetMlasPlatform().Avx512Supported_) { + GTEST_SKIP() << "AVX-512BW (DQ/VL) not available on this host"; + } + for (uint32_t seed : {0xC0FFEEu, 0xBADC0DEu}) { + for (const auto& s : kSimdShapes) { + for (bool bias : {false, true}) { + RunW2Case(s.M, s.N, s.K, bias, seed + (bias ? 1u : 0u), + /*WithZeroPoints=*/false, + sq2::SQ2BitGemmKernel_BlkSum_CompInt8_Avx512_TestEntry, + "AVX-512BW"); + } } + } } -TEST(MlasSq2BitTest, BlkLen64_Avx512_WithZeroPoints) -{ - if (!GetMlasPlatform().Avx512Supported_) { - GTEST_SKIP() << "AVX-512BW (DQ/VL) not available on this host"; - } - for (uint32_t seed : {0xC0FFEEu, 0xBADC0DEu}) { - for (const auto& s : kSimdShapes) { - for (bool bias : {false, true}) { - RunW2Case(s.M, s.N, s.K, bias, seed + (bias ? 1u : 0u), - /*WithZeroPoints=*/true, - sq2::SQ2BitGemmKernel_BlkSum_CompInt8_Avx512_TestEntry, - "AVX-512BW"); - } - } +TEST(MlasSq2BitTest, BlkLen64_Avx512_WithZeroPoints) { + if (!GetMlasPlatform().Avx512Supported_) { + GTEST_SKIP() << "AVX-512BW (DQ/VL) not available on this host"; + } + for (uint32_t seed : {0xC0FFEEu, 0xBADC0DEu}) { + for (const auto& s : kSimdShapes) { + for (bool bias : {false, true}) { + RunW2Case(s.M, s.N, s.K, bias, seed + (bias ? 1u : 0u), + /*WithZeroPoints=*/true, + sq2::SQ2BitGemmKernel_BlkSum_CompInt8_Avx512_TestEntry, + "AVX-512BW"); + } } + } } // // AVX-512-VNNI SIMD block-group kernel. Gated on the platform having selected // the VNNI dispatch table (the SIMD path uses `_mm512_dpbusd_epi32`). // -TEST(MlasSq2BitTest, BlkLen64_Avx512Vnni) -{ - if (GetMlasPlatform().QNBitGemmDispatch != &MlasSQNBitGemmDispatchAvx512vnni) { - GTEST_SKIP() << "AVX-512-VNNI not selected as the active dispatch on this host"; - } - for (uint32_t seed : {0xC0FFEEu, 0xBADC0DEu}) { - for (const auto& s : kSimdShapes) { - for (bool bias : {false, true}) { - RunW2Case(s.M, s.N, s.K, bias, seed + (bias ? 1u : 0u), - /*WithZeroPoints=*/false, - sq2::SQ2BitGemmKernel_BlkSum_CompInt8_Avx512Vnni_TestEntry, - "AVX-512-VNNI"); - } - } +TEST(MlasSq2BitTest, BlkLen64_Avx512Vnni) { + if (GetMlasPlatform().QNBitGemmDispatch != &MlasSQNBitGemmDispatchAvx512vnni) { + GTEST_SKIP() << "AVX-512-VNNI not selected as the active dispatch on this host"; + } + for (uint32_t seed : {0xC0FFEEu, 0xBADC0DEu}) { + for (const auto& s : kSimdShapes) { + for (bool bias : {false, true}) { + RunW2Case(s.M, s.N, s.K, bias, seed + (bias ? 1u : 0u), + /*WithZeroPoints=*/false, + sq2::SQ2BitGemmKernel_BlkSum_CompInt8_Avx512Vnni_TestEntry, + "AVX-512-VNNI"); + } } + } } -TEST(MlasSq2BitTest, BlkLen64_Avx512Vnni_WithZeroPoints) -{ - if (GetMlasPlatform().QNBitGemmDispatch != &MlasSQNBitGemmDispatchAvx512vnni) { - GTEST_SKIP() << "AVX-512-VNNI not selected as the active dispatch on this host"; - } - for (uint32_t seed : {0xC0FFEEu, 0xBADC0DEu}) { - for (const auto& s : kSimdShapes) { - for (bool bias : {false, true}) { - RunW2Case(s.M, s.N, s.K, bias, seed + (bias ? 1u : 0u), - /*WithZeroPoints=*/true, - sq2::SQ2BitGemmKernel_BlkSum_CompInt8_Avx512Vnni_TestEntry, - "AVX-512-VNNI"); - } - } +TEST(MlasSq2BitTest, BlkLen64_Avx512Vnni_WithZeroPoints) { + if (GetMlasPlatform().QNBitGemmDispatch != &MlasSQNBitGemmDispatchAvx512vnni) { + GTEST_SKIP() << "AVX-512-VNNI not selected as the active dispatch on this host"; + } + for (uint32_t seed : {0xC0FFEEu, 0xBADC0DEu}) { + for (const auto& s : kSimdShapes) { + for (bool bias : {false, true}) { + RunW2Case(s.M, s.N, s.K, bias, seed + (bias ? 1u : 0u), + /*WithZeroPoints=*/true, + sq2::SQ2BitGemmKernel_BlkSum_CompInt8_Avx512Vnni_TestEntry, + "AVX-512-VNNI"); + } } + } } // ============================================================================= @@ -533,217 +543,208 @@ TEST(MlasSq2BitTest, BlkLen64_Avx512Vnni_WithZeroPoints) namespace { -constexpr size_t kBlkLen128 = sq2::kBlkLen128; // 128 -constexpr size_t kBlkBytes128 = sq2::kBlkBytes128; // 32 - -void -PackSourceBlock_BlkLen128(const uint8_t weights[kBlkLen128], std::byte* src_out) -{ - for (size_t i = 0; i < kBlkBytes128; ++i) { - const uint8_t v0 = weights[4 * i + 0] & 0x03u; - const uint8_t v1 = weights[4 * i + 1] & 0x03u; - const uint8_t v2 = weights[4 * i + 2] & 0x03u; - const uint8_t v3 = weights[4 * i + 3] & 0x03u; - src_out[i] = static_cast( - static_cast(v0 | (v1 << 2) | (v2 << 4) | (v3 << 6)) - ); - } +constexpr size_t kBlkLen128 = sq2::kBlkLen128; // 128 +constexpr size_t kBlkBytes128 = sq2::kBlkBytes128; // 32 + +void PackSourceBlock_BlkLen128(const uint8_t weights[kBlkLen128], std::byte* src_out) { + for (size_t i = 0; i < kBlkBytes128; ++i) { + const uint8_t v0 = weights[4 * i + 0] & 0x03u; + const uint8_t v1 = weights[4 * i + 1] & 0x03u; + const uint8_t v2 = weights[4 * i + 2] & 0x03u; + const uint8_t v3 = weights[4 * i + 3] & 0x03u; + src_out[i] = static_cast( + static_cast(v0 | (v1 << 2) | (v2 << 4) | (v3 << 6))); + } } -void -QuantizeA_Reference_BlkLen128(size_t M, size_t K, const float* A, - int8_t* QuantAData, float* QuantAScale) -{ - const size_t BlockCountK = (K + kBlkLen128 - 1) / kBlkLen128; - for (size_t m = 0; m < M; ++m) { - for (size_t k = 0, k_blk = 0; k < K; k += kBlkLen128, ++k_blk) { - const size_t local_len = std::min(K - k, kBlkLen128); - float amax = 0.0f; - for (size_t kk = 0; kk < local_len; ++kk) { - amax = std::max(amax, std::fabs(A[m * K + k + kk])); - } - constexpr float range_max = 127.0f; - const float scale = amax / range_max; - const float scale_recip = amax != 0.0f ? range_max / amax : 0.0f; - QuantAScale[m * BlockCountK + k_blk] = scale; - for (size_t kk = 0; kk < kBlkLen128; ++kk) { - const float a = (kk < local_len) ? A[m * K + k + kk] : 0.0f; - const float q = std::nearbyint(a * scale_recip); - QuantAData[m * BlockCountK * kBlkLen128 + k + kk] = - static_cast(std::clamp(q, -127.0f, 127.0f)); - } - } - } -} - -void -ReferenceGemm_W2_CompInt8_BlkLen128(size_t M, size_t N, size_t K, - const float* A, - const std::vector& BWeights, - const float* QuantBScale, - const uint8_t* BZeroPoints, - const float* Bias, - float* C) -{ - const size_t BlockCountK = (K + kBlkLen128 - 1) / kBlkLen128; - std::vector QuantAData(M * BlockCountK * kBlkLen128, int8_t{0}); - std::vector QuantAScale(M * BlockCountK, 0.0f); - QuantizeA_Reference_BlkLen128(M, K, A, QuantAData.data(), QuantAScale.data()); - - for (size_t m = 0; m < M; ++m) { - for (size_t n = 0; n < N; ++n) { - float acc = (Bias != nullptr) ? Bias[n] : 0.0f; - for (size_t k = 0, blk = 0; k < K; k += kBlkLen128, ++blk) { - const size_t local_len = std::min(K - k, kBlkLen128); - const float a_scale = QuantAScale[m * BlockCountK + blk]; - const float b_scale = QuantBScale[n * BlockCountK + blk]; - const int32_t zp = BZeroPoints != nullptr - ? static_cast(BZeroPoints[n * BlockCountK + blk]) - : static_cast(sq2::kDefaultSymmetricZeroPoint2Bit); - int32_t dot = 0; - for (size_t kk = 0; kk < local_len; ++kk) { - const int8_t qa = QuantAData[m * BlockCountK * kBlkLen128 + k + kk]; - const int32_t qb = - static_cast(BWeights[n * K + k + kk]) - zp; - dot += static_cast(qa) * qb; - } - acc += static_cast(dot) * a_scale * b_scale; - } - C[m * N + n] = acc; - } +void QuantizeA_Reference_BlkLen128(size_t M, size_t K, const float* A, + int8_t* QuantAData, float* QuantAScale) { + const size_t BlockCountK = (K + kBlkLen128 - 1) / kBlkLen128; + for (size_t m = 0; m < M; ++m) { + for (size_t k = 0, k_blk = 0; k < K; k += kBlkLen128, ++k_blk) { + const size_t local_len = std::min(K - k, kBlkLen128); + float amax = 0.0f; + for (size_t kk = 0; kk < local_len; ++kk) { + amax = std::max(amax, std::fabs(A[m * K + k + kk])); + } + constexpr float range_max = 127.0f; + const float scale = amax / range_max; + const float scale_recip = amax != 0.0f ? range_max / amax : 0.0f; + QuantAScale[m * BlockCountK + k_blk] = scale; + for (size_t kk = 0; kk < kBlkLen128; ++kk) { + const float a = (kk < local_len) ? A[m * K + k + kk] : 0.0f; + const float q = std::nearbyint(a * scale_recip); + QuantAData[m * BlockCountK * kBlkLen128 + k + kk] = + static_cast(std::clamp(q, -127.0f, 127.0f)); + } } + } } -void -RunW2Case_BlkLen128(size_t M, size_t N, size_t K, bool WithBias, uint32_t seed, - bool WithZeroPoints, W2KernelFn kernel, - const char* kernel_name) -{ - const size_t BlockCountK = (K + kBlkLen128 - 1) / kBlkLen128; - ASSERT_EQ(K % kBlkLen128, 0u) << "BlkLen128 test K must be a multiple of 128"; - - std::mt19937 rng(seed); - std::uniform_real_distribution a_dist(-1.0f, 1.0f); - std::uniform_int_distribution w_dist(0, 3); - std::uniform_real_distribution s_dist(0.05f, 0.5f); - - std::vector A(M * K); - for (auto& v : A) v = a_dist(rng); - - std::vector BWeights(N * K); - for (auto& v : BWeights) v = static_cast(w_dist(rng)); - - // Source-packed B in standard ONNX layout (32 bytes per block at BlkLen=128). - std::vector QuantBData(N * BlockCountK * kBlkBytes128, std::byte{0}); +void ReferenceGemm_W2_CompInt8_BlkLen128(size_t M, size_t N, size_t K, + const float* A, + const std::vector& BWeights, + const float* QuantBScale, + const uint8_t* BZeroPoints, + const float* Bias, + float* C) { + const size_t BlockCountK = (K + kBlkLen128 - 1) / kBlkLen128; + std::vector QuantAData(M * BlockCountK * kBlkLen128, int8_t{0}); + std::vector QuantAScale(M * BlockCountK, 0.0f); + QuantizeA_Reference_BlkLen128(M, K, A, QuantAData.data(), QuantAScale.data()); + + for (size_t m = 0; m < M; ++m) { for (size_t n = 0; n < N; ++n) { - for (size_t blk = 0; blk < BlockCountK; ++blk) { - uint8_t blk_weights[kBlkLen128]; - for (size_t kk = 0; kk < kBlkLen128; ++kk) { - blk_weights[kk] = BWeights[n * K + blk * kBlkLen128 + kk]; - } - PackSourceBlock_BlkLen128(blk_weights, - QuantBData.data() + (n * BlockCountK + blk) * kBlkBytes128); + float acc = (Bias != nullptr) ? Bias[n] : 0.0f; + for (size_t k = 0, blk = 0; k < K; k += kBlkLen128, ++blk) { + const size_t local_len = std::min(K - k, kBlkLen128); + const float a_scale = QuantAScale[m * BlockCountK + blk]; + const float b_scale = QuantBScale[n * BlockCountK + blk]; + const int32_t zp = BZeroPoints != nullptr + ? static_cast(BZeroPoints[n * BlockCountK + blk]) + : static_cast(sq2::kDefaultSymmetricZeroPoint2Bit); + int32_t dot = 0; + for (size_t kk = 0; kk < local_len; ++kk) { + const int8_t qa = QuantAData[m * BlockCountK * kBlkLen128 + k + kk]; + const int32_t qb = + static_cast(BWeights[n * K + k + kk]) - zp; + dot += static_cast(qa) * qb; } + acc += static_cast(dot) * a_scale * b_scale; + } + C[m * N + n] = acc; } + } +} - std::vector QuantBScale(N * BlockCountK); - for (auto& v : QuantBScale) v = s_dist(rng); - - std::vector BZeroPoints; - std::vector BZeroPointsPacked; - const uint8_t* BZeroPointsRef = nullptr; - const std::byte* BZeroPointsMlas = nullptr; - if (WithZeroPoints) { - BZeroPoints.resize(N * BlockCountK); - for (auto& v : BZeroPoints) v = static_cast(w_dist(rng)); - BZeroPointsRef = BZeroPoints.data(); - BZeroPointsPacked = PackW2ZeroPoints(N, BlockCountK, BZeroPoints); - BZeroPointsMlas = BZeroPointsPacked.data(); - } - - std::vector Bias; - const float* BiasPtr = nullptr; - if (WithBias) { - Bias.resize(N); - for (auto& v : Bias) v = a_dist(rng); - BiasPtr = Bias.data(); +void RunW2Case_BlkLen128(size_t M, size_t N, size_t K, bool WithBias, uint32_t seed, + bool WithZeroPoints, W2KernelFn kernel, + const char* kernel_name) { + const size_t BlockCountK = (K + kBlkLen128 - 1) / kBlkLen128; + ASSERT_EQ(K % kBlkLen128, 0u) << "BlkLen128 test K must be a multiple of 128"; + + std::mt19937 rng(seed); + std::uniform_real_distribution a_dist(-1.0f, 1.0f); + std::uniform_int_distribution w_dist(0, 3); + std::uniform_real_distribution s_dist(0.05f, 0.5f); + + std::vector A(M * K); + for (auto& v : A) v = a_dist(rng); + + std::vector BWeights(N * K); + for (auto& v : BWeights) v = static_cast(w_dist(rng)); + + // Source-packed B in standard ONNX layout (32 bytes per block at BlkLen=128). + std::vector QuantBData(N * BlockCountK * kBlkBytes128, std::byte{0}); + for (size_t n = 0; n < N; ++n) { + for (size_t blk = 0; blk < BlockCountK; ++blk) { + uint8_t blk_weights[kBlkLen128]; + for (size_t kk = 0; kk < kBlkLen128; ++kk) { + blk_weights[kk] = BWeights[n * K + blk * kBlkLen128 + kk]; + } + PackSourceBlock_BlkLen128(blk_weights, + QuantBData.data() + (n * BlockCountK + blk) * kBlkBytes128); } - - const size_t PackedSize = sq2::Q2BitGemmPackQuantBDataSize_Avx512( - N, K, kBlkLen128, WithZeroPoints, SQNBIT_CompInt8, nullptr); - ASSERT_GT(PackedSize, 0u) << "BlkLen128 block-group pack size unsupported for shape"; - - std::vector PackedQuantBBuf(PackedSize, std::byte{0}); - PackedQuantBDataStruct packed_b( - PackedQuantBBuf.data(), N, BlockCountK, kBlkLen128, /*QuantAUnsigned=*/false); - - sq2::SQ2BitGemmPackQuantBDataAndBlkSum_Scalar( - N, K, kBlkLen128, SQNBIT_CompInt8, - QuantBData.data(), /*scales=*/nullptr, - WithZeroPoints, /*zp=*/nullptr, - packed_b, /*tp=*/nullptr, /*cfg=*/nullptr); + } + + std::vector QuantBScale(N * BlockCountK); + for (auto& v : QuantBScale) v = s_dist(rng); + + std::vector BZeroPoints; + std::vector BZeroPointsPacked; + const uint8_t* BZeroPointsRef = nullptr; + const std::byte* BZeroPointsMlas = nullptr; + if (WithZeroPoints) { + BZeroPoints.resize(N * BlockCountK); + for (auto& v : BZeroPoints) v = static_cast(w_dist(rng)); + BZeroPointsRef = BZeroPoints.data(); + BZeroPointsPacked = PackW2ZeroPoints(N, BlockCountK, BZeroPoints); + BZeroPointsMlas = BZeroPointsPacked.data(); + } + + std::vector Bias; + const float* BiasPtr = nullptr; + if (WithBias) { + Bias.resize(N); + for (auto& v : Bias) v = a_dist(rng); + BiasPtr = Bias.data(); + } + + const size_t PackedSize = sq2::Q2BitGemmPackQuantBDataSize_Avx512( + N, K, kBlkLen128, WithZeroPoints, SQNBIT_CompInt8, nullptr); + ASSERT_GT(PackedSize, 0u) << "BlkLen128 block-group pack size unsupported for shape"; + + std::vector PackedQuantBBuf(PackedSize, std::byte{0}); + PackedQuantBDataStruct packed_b( + PackedQuantBBuf.data(), N, BlockCountK, kBlkLen128, /*QuantAUnsigned=*/false); + + sq2::SQ2BitGemmPackQuantBDataAndBlkSum_Scalar( + N, K, kBlkLen128, SQNBIT_CompInt8, + QuantBData.data(), /*scales=*/nullptr, + WithZeroPoints, /*zp=*/nullptr, + packed_b, /*tp=*/nullptr, /*cfg=*/nullptr); + sq2::SQ2BitGemmPackQuantBDataAndBlkSum_Scalar( + N, K, kBlkLen128, SQNBIT_CompInt8, + /*B=*/nullptr, QuantBScale.data(), + WithZeroPoints, /*zp=*/nullptr, + packed_b, /*tp=*/nullptr, /*cfg=*/nullptr); + if (WithZeroPoints) { sq2::SQ2BitGemmPackQuantBDataAndBlkSum_Scalar( N, K, kBlkLen128, SQNBIT_CompInt8, - /*B=*/nullptr, QuantBScale.data(), - WithZeroPoints, /*zp=*/nullptr, + /*B=*/nullptr, /*scales=*/nullptr, + WithZeroPoints, BZeroPointsMlas, packed_b, /*tp=*/nullptr, /*cfg=*/nullptr); - if (WithZeroPoints) { - sq2::SQ2BitGemmPackQuantBDataAndBlkSum_Scalar( - N, K, kBlkLen128, SQNBIT_CompInt8, - /*B=*/nullptr, /*scales=*/nullptr, - WithZeroPoints, BZeroPointsMlas, - packed_b, /*tp=*/nullptr, /*cfg=*/nullptr); - } - - std::vector QuantAData(M * BlockCountK * kBlkLen128, int8_t{0}); - std::vector QuantAScale(M * BlockCountK, 0.0f); - QuantizeA_Reference_BlkLen128(M, K, A.data(), QuantAData.data(), QuantAScale.data()); - - std::vector ABlockSum(M * BlockCountK, 0.0f); - for (size_t m = 0; m < M; ++m) { - for (size_t blk = 0; blk < BlockCountK; ++blk) { - int32_t sum = 0; - for (size_t kk = 0; kk < kBlkLen128; ++kk) { - sum += static_cast( - QuantAData[m * BlockCountK * kBlkLen128 + blk * kBlkLen128 + kk]); - } - ABlockSum[m * BlockCountK + blk] = - QuantAScale[m * BlockCountK + blk] * static_cast(sum); - } - } - - std::vector C(M * N, 0.0f); - kernel( - kBlkLen128, - reinterpret_cast(QuantAData.data()), - QuantAScale.data(), - packed_b.PackedQuantBData, - packed_b.PackedQuantBScale, - /*QuantBZeroPoint=*/nullptr, - C.data(), - M, N, /*CountK=*/K, BlockCountK, - BiasPtr, - /*ldc=*/N, - ABlockSum.data(), - packed_b.QuantBBlkSum); - - std::vector CRef(M * N, 0.0f); - ReferenceGemm_W2_CompInt8_BlkLen128(M, N, K, A.data(), BWeights, QuantBScale.data(), - BZeroPointsRef, BiasPtr, CRef.data()); - - const float abs_tol = 1e-4f; - const float rel_tol = 1e-4f; - for (size_t i = 0; i < M * N; ++i) { - const float diff = std::fabs(C[i] - CRef[i]); - const float bound = abs_tol + rel_tol * std::fabs(CRef[i]); - ASSERT_LE(diff, bound) - << "BlkLen128 " << kernel_name << " mismatch at i=" << i - << " (m=" << (i / N) << ", n=" << (i % N) << ")" - << " out=" << C[i] << " ref=" << CRef[i] - << " M=" << M << " N=" << N << " K=" << K - << " WithBias=" << WithBias - << " WithZeroPoints=" << WithZeroPoints; + } + + std::vector QuantAData(M * BlockCountK * kBlkLen128, int8_t{0}); + std::vector QuantAScale(M * BlockCountK, 0.0f); + QuantizeA_Reference_BlkLen128(M, K, A.data(), QuantAData.data(), QuantAScale.data()); + + std::vector ABlockSum(M * BlockCountK, 0.0f); + for (size_t m = 0; m < M; ++m) { + for (size_t blk = 0; blk < BlockCountK; ++blk) { + int32_t sum = 0; + for (size_t kk = 0; kk < kBlkLen128; ++kk) { + sum += static_cast( + QuantAData[m * BlockCountK * kBlkLen128 + blk * kBlkLen128 + kk]); + } + ABlockSum[m * BlockCountK + blk] = + QuantAScale[m * BlockCountK + blk] * static_cast(sum); } + } + + std::vector C(M * N, 0.0f); + kernel( + kBlkLen128, + reinterpret_cast(QuantAData.data()), + QuantAScale.data(), + packed_b.PackedQuantBData, + packed_b.PackedQuantBScale, + /*QuantBZeroPoint=*/nullptr, + C.data(), + M, N, /*CountK=*/K, BlockCountK, + BiasPtr, + /*ldc=*/N, + ABlockSum.data(), + packed_b.QuantBBlkSum); + + std::vector CRef(M * N, 0.0f); + ReferenceGemm_W2_CompInt8_BlkLen128(M, N, K, A.data(), BWeights, QuantBScale.data(), + BZeroPointsRef, BiasPtr, CRef.data()); + + const float abs_tol = 1e-4f; + const float rel_tol = 1e-4f; + for (size_t i = 0; i < M * N; ++i) { + const float diff = std::fabs(C[i] - CRef[i]); + const float bound = abs_tol + rel_tol * std::fabs(CRef[i]); + ASSERT_LE(diff, bound) + << "BlkLen128 " << kernel_name << " mismatch at i=" << i + << " (m=" << (i / N) << ", n=" << (i % N) << ")" + << " out=" << C[i] << " ref=" << CRef[i] + << " M=" << M << " N=" << N << " K=" << K + << " WithBias=" << WithBias + << " WithZeroPoints=" << WithZeroPoints; + } } // @@ -753,146 +754,144 @@ RunW2Case_BlkLen128(size_t M, size_t N, size_t K, bool WithBias, uint32_t seed, // * BlockCountK in {1, 2, 3, 4, 8, 16, 32} -- full + K-tail variants // * N-tail (NMain=0 and various NTail) combined with K-tail // -constexpr struct { size_t M, N, K; } kSimdShapes_BlkLen128[] = { - {1, 16, 128}, // R1, BlockCountK=1 - {1, 32, 256}, // R1, BlockCountK=2 (K-tail, no full group) - {1, 1024, 1024}, // R1, BlockCountK=8 - {1, 1024, 4096}, // R1, BlockCountK=32 (customer N) - {2, 16, 128}, - {2, 32, 256}, - {2, 64, 512}, // BlockCountK=4 (one full group) - {3, 16, 256}, // R2 head + R1 tail - {3, 384, 1024}, - {4, 16, 256}, - {4, 32, 256}, - {16, 64, 512}, - {32, 128, 256}, +constexpr struct { + size_t M, N, K; +} kSimdShapes_BlkLen128[] = { + {1, 16, 128}, // R1, BlockCountK=1 + {1, 32, 256}, // R1, BlockCountK=2 (K-tail, no full group) + {1, 1024, 1024}, // R1, BlockCountK=8 + {1, 1024, 4096}, // R1, BlockCountK=32 (customer N) + {2, 16, 128}, + {2, 32, 256}, + {2, 64, 512}, // BlockCountK=4 (one full group) + {3, 16, 256}, // R2 head + R1 tail + {3, 384, 1024}, + {4, 16, 256}, + {4, 32, 256}, + {16, 64, 512}, + {32, 128, 256}, // M=128 prefill at customer-like shapes (K multiples of 128) - {128, 1024, 1024}, {128, 1024, 4096}, - {128, 192, 1024}, {128, 384, 1024}, + {128, 1024, 1024}, + {128, 1024, 4096}, + {128, 192, 1024}, + {128, 384, 1024}, // K-tail stress (BlockCountK not a multiple of 4) - { 2, 16, 384}, // tail=3 - { 4, 16, 384}, - {128, 1024, 384}, // tail=3 at customer M - { 2, 16, 640}, // tail=1 - { 4, 16, 640}, - { 2, 16, 768}, // tail=2 - { 4, 16, 768}, + {2, 16, 384}, // tail=3 + {4, 16, 384}, + {128, 1024, 384}, // tail=3 at customer M + {2, 16, 640}, // tail=1 + {4, 16, 640}, + {2, 16, 768}, // tail=2 + {4, 16, 768}, // N-tail stress - { 1, 1, 256}, - { 1, 3, 256}, - { 4, 3, 256}, - { 1, 17, 256}, - { 4, 17, 256}, - {128, 17, 256}, - { 1, 33, 256}, - { 4, 33, 256}, - {128, 33, 256}, - { 1, 18, 256}, - { 4, 18, 256}, - {128, 19, 256}, + {1, 1, 256}, + {1, 3, 256}, + {4, 3, 256}, + {1, 17, 256}, + {4, 17, 256}, + {128, 17, 256}, + {1, 33, 256}, + {4, 33, 256}, + {128, 33, 256}, + {1, 18, 256}, + {4, 18, 256}, + {128, 19, 256}, // N-tail combined with K-tail (most generic) - { 1, 17, 384}, - { 4, 33, 384}, - {128, 19, 640}, + {1, 17, 384}, + {4, 33, 384}, + {128, 19, 640}, }; } // namespace -TEST(MlasSq2BitTest, Scalar_BlkLen128) -{ - for (uint32_t seed : {0xC0FFEEu, 0xBADC0DEu}) { - for (const auto& s : kSimdShapes_BlkLen128) { - for (bool bias : {false, true}) { - RunW2Case_BlkLen128(s.M, s.N, s.K, bias, seed + (bias ? 1u : 0u), - /*WithZeroPoints=*/false, - sq2::SQ2BitGemmKernel_BlkSum_CompInt8_Scalar, - "Scalar"); - } - } +TEST(MlasSq2BitTest, Scalar_BlkLen128) { + for (uint32_t seed : {0xC0FFEEu, 0xBADC0DEu}) { + for (const auto& s : kSimdShapes_BlkLen128) { + for (bool bias : {false, true}) { + RunW2Case_BlkLen128(s.M, s.N, s.K, bias, seed + (bias ? 1u : 0u), + /*WithZeroPoints=*/false, + sq2::SQ2BitGemmKernel_BlkSum_CompInt8_Scalar, + "Scalar"); + } } + } } -TEST(MlasSq2BitTest, Scalar_BlkLen128_WithZeroPoints) -{ - for (uint32_t seed : {0xC0FFEEu, 0xBADC0DEu}) { - for (const auto& s : kSimdShapes_BlkLen128) { - for (bool bias : {false, true}) { - RunW2Case_BlkLen128(s.M, s.N, s.K, bias, seed + (bias ? 1u : 0u), - /*WithZeroPoints=*/true, - sq2::SQ2BitGemmKernel_BlkSum_CompInt8_Scalar, - "Scalar"); - } - } +TEST(MlasSq2BitTest, Scalar_BlkLen128_WithZeroPoints) { + for (uint32_t seed : {0xC0FFEEu, 0xBADC0DEu}) { + for (const auto& s : kSimdShapes_BlkLen128) { + for (bool bias : {false, true}) { + RunW2Case_BlkLen128(s.M, s.N, s.K, bias, seed + (bias ? 1u : 0u), + /*WithZeroPoints=*/true, + sq2::SQ2BitGemmKernel_BlkSum_CompInt8_Scalar, + "Scalar"); + } } + } } -TEST(MlasSq2BitTest, BlkLen128_Avx512) -{ - if (!GetMlasPlatform().Avx512Supported_) { - GTEST_SKIP() << "AVX-512BW (DQ/VL) not available on this host"; - } - for (uint32_t seed : {0xC0FFEEu, 0xBADC0DEu}) { - for (const auto& s : kSimdShapes_BlkLen128) { - for (bool bias : {false, true}) { - RunW2Case_BlkLen128(s.M, s.N, s.K, bias, seed + (bias ? 1u : 0u), - /*WithZeroPoints=*/false, - sq2::SQ2BitGemmKernel_BlkSum_CompInt8_Avx512_TestEntry, - "AVX-512BW"); - } - } +TEST(MlasSq2BitTest, BlkLen128_Avx512) { + if (!GetMlasPlatform().Avx512Supported_) { + GTEST_SKIP() << "AVX-512BW (DQ/VL) not available on this host"; + } + for (uint32_t seed : {0xC0FFEEu, 0xBADC0DEu}) { + for (const auto& s : kSimdShapes_BlkLen128) { + for (bool bias : {false, true}) { + RunW2Case_BlkLen128(s.M, s.N, s.K, bias, seed + (bias ? 1u : 0u), + /*WithZeroPoints=*/false, + sq2::SQ2BitGemmKernel_BlkSum_CompInt8_Avx512_TestEntry, + "AVX-512BW"); + } } + } } -TEST(MlasSq2BitTest, BlkLen128_Avx512_WithZeroPoints) -{ - if (!GetMlasPlatform().Avx512Supported_) { - GTEST_SKIP() << "AVX-512BW (DQ/VL) not available on this host"; - } - for (uint32_t seed : {0xC0FFEEu, 0xBADC0DEu}) { - for (const auto& s : kSimdShapes_BlkLen128) { - for (bool bias : {false, true}) { - RunW2Case_BlkLen128(s.M, s.N, s.K, bias, seed + (bias ? 1u : 0u), - /*WithZeroPoints=*/true, - sq2::SQ2BitGemmKernel_BlkSum_CompInt8_Avx512_TestEntry, - "AVX-512BW"); - } - } +TEST(MlasSq2BitTest, BlkLen128_Avx512_WithZeroPoints) { + if (!GetMlasPlatform().Avx512Supported_) { + GTEST_SKIP() << "AVX-512BW (DQ/VL) not available on this host"; + } + for (uint32_t seed : {0xC0FFEEu, 0xBADC0DEu}) { + for (const auto& s : kSimdShapes_BlkLen128) { + for (bool bias : {false, true}) { + RunW2Case_BlkLen128(s.M, s.N, s.K, bias, seed + (bias ? 1u : 0u), + /*WithZeroPoints=*/true, + sq2::SQ2BitGemmKernel_BlkSum_CompInt8_Avx512_TestEntry, + "AVX-512BW"); + } } + } } -TEST(MlasSq2BitTest, BlkLen128_Avx512Vnni) -{ - if (GetMlasPlatform().QNBitGemmDispatch != &MlasSQNBitGemmDispatchAvx512vnni) { - GTEST_SKIP() << "AVX-512-VNNI not selected as the active dispatch on this host"; - } - for (uint32_t seed : {0xC0FFEEu, 0xBADC0DEu}) { - for (const auto& s : kSimdShapes_BlkLen128) { - for (bool bias : {false, true}) { - RunW2Case_BlkLen128(s.M, s.N, s.K, bias, seed + (bias ? 1u : 0u), - /*WithZeroPoints=*/false, - sq2::SQ2BitGemmKernel_BlkSum_CompInt8_Avx512Vnni_TestEntry, - "AVX-512-VNNI"); - } - } +TEST(MlasSq2BitTest, BlkLen128_Avx512Vnni) { + if (GetMlasPlatform().QNBitGemmDispatch != &MlasSQNBitGemmDispatchAvx512vnni) { + GTEST_SKIP() << "AVX-512-VNNI not selected as the active dispatch on this host"; + } + for (uint32_t seed : {0xC0FFEEu, 0xBADC0DEu}) { + for (const auto& s : kSimdShapes_BlkLen128) { + for (bool bias : {false, true}) { + RunW2Case_BlkLen128(s.M, s.N, s.K, bias, seed + (bias ? 1u : 0u), + /*WithZeroPoints=*/false, + sq2::SQ2BitGemmKernel_BlkSum_CompInt8_Avx512Vnni_TestEntry, + "AVX-512-VNNI"); + } } + } } -TEST(MlasSq2BitTest, BlkLen128_Avx512Vnni_WithZeroPoints) -{ - if (GetMlasPlatform().QNBitGemmDispatch != &MlasSQNBitGemmDispatchAvx512vnni) { - GTEST_SKIP() << "AVX-512-VNNI not selected as the active dispatch on this host"; - } - for (uint32_t seed : {0xC0FFEEu, 0xBADC0DEu}) { - for (const auto& s : kSimdShapes_BlkLen128) { - for (bool bias : {false, true}) { - RunW2Case_BlkLen128(s.M, s.N, s.K, bias, seed + (bias ? 1u : 0u), - /*WithZeroPoints=*/true, - sq2::SQ2BitGemmKernel_BlkSum_CompInt8_Avx512Vnni_TestEntry, - "AVX-512-VNNI"); - } - } +TEST(MlasSq2BitTest, BlkLen128_Avx512Vnni_WithZeroPoints) { + if (GetMlasPlatform().QNBitGemmDispatch != &MlasSQNBitGemmDispatchAvx512vnni) { + GTEST_SKIP() << "AVX-512-VNNI not selected as the active dispatch on this host"; + } + for (uint32_t seed : {0xC0FFEEu, 0xBADC0DEu}) { + for (const auto& s : kSimdShapes_BlkLen128) { + for (bool bias : {false, true}) { + RunW2Case_BlkLen128(s.M, s.N, s.K, bias, seed + (bias ? 1u : 0u), + /*WithZeroPoints=*/true, + sq2::SQ2BitGemmKernel_BlkSum_CompInt8_Avx512Vnni_TestEntry, + "AVX-512-VNNI"); + } } + } } // ============================================================================= @@ -903,359 +902,350 @@ TEST(MlasSq2BitTest, BlkLen128_Avx512Vnni_WithZeroPoints) namespace { -constexpr size_t kBlkLen32 = sq2::kBlkLen32; // 32 -constexpr size_t kBlkBytes32 = sq2::kBlkBytes32; // 8 - -void -PackSourceBlock_BlkLen32(const uint8_t weights[kBlkLen32], std::byte* src_out) -{ - for (size_t i = 0; i < kBlkBytes32; ++i) { - const uint8_t v0 = weights[4 * i + 0] & 0x03u; - const uint8_t v1 = weights[4 * i + 1] & 0x03u; - const uint8_t v2 = weights[4 * i + 2] & 0x03u; - const uint8_t v3 = weights[4 * i + 3] & 0x03u; - src_out[i] = static_cast( - static_cast(v0 | (v1 << 2) | (v2 << 4) | (v3 << 6)) - ); - } -} - -void -QuantizeA_Reference_BlkLen32(size_t M, size_t K, const float* A, - int8_t* QuantAData, float* QuantAScale) -{ - const size_t BlockCountK = (K + kBlkLen32 - 1) / kBlkLen32; - for (size_t m = 0; m < M; ++m) { - for (size_t k = 0, k_blk = 0; k < K; k += kBlkLen32, ++k_blk) { - const size_t local_len = std::min(K - k, kBlkLen32); - float amax = 0.0f; - for (size_t kk = 0; kk < local_len; ++kk) { - amax = std::max(amax, std::fabs(A[m * K + k + kk])); - } - constexpr float range_max = 127.0f; - const float scale = amax / range_max; - const float scale_recip = amax != 0.0f ? range_max / amax : 0.0f; - QuantAScale[m * BlockCountK + k_blk] = scale; - for (size_t kk = 0; kk < kBlkLen32; ++kk) { - const float a = (kk < local_len) ? A[m * K + k + kk] : 0.0f; - const float q = std::nearbyint(a * scale_recip); - QuantAData[m * BlockCountK * kBlkLen32 + k + kk] = - static_cast(std::clamp(q, -127.0f, 127.0f)); - } - } - } +constexpr size_t kBlkLen32 = sq2::kBlkLen32; // 32 +constexpr size_t kBlkBytes32 = sq2::kBlkBytes32; // 8 + +void PackSourceBlock_BlkLen32(const uint8_t weights[kBlkLen32], std::byte* src_out) { + for (size_t i = 0; i < kBlkBytes32; ++i) { + const uint8_t v0 = weights[4 * i + 0] & 0x03u; + const uint8_t v1 = weights[4 * i + 1] & 0x03u; + const uint8_t v2 = weights[4 * i + 2] & 0x03u; + const uint8_t v3 = weights[4 * i + 3] & 0x03u; + src_out[i] = static_cast( + static_cast(v0 | (v1 << 2) | (v2 << 4) | (v3 << 6))); + } } -void -ReferenceGemm_W2_CompInt8_BlkLen32(size_t M, size_t N, size_t K, - const float* A, - const std::vector& BWeights, - const float* QuantBScale, - const uint8_t* BZeroPoints, - const float* Bias, - float* C) -{ - const size_t BlockCountK = (K + kBlkLen32 - 1) / kBlkLen32; - std::vector QuantAData(M * BlockCountK * kBlkLen32, int8_t{0}); - std::vector QuantAScale(M * BlockCountK, 0.0f); - QuantizeA_Reference_BlkLen32(M, K, A, QuantAData.data(), QuantAScale.data()); - - for (size_t m = 0; m < M; ++m) { - for (size_t n = 0; n < N; ++n) { - float acc = (Bias != nullptr) ? Bias[n] : 0.0f; - for (size_t k = 0, blk = 0; k < K; k += kBlkLen32, ++blk) { - const size_t local_len = std::min(K - k, kBlkLen32); - const float a_scale = QuantAScale[m * BlockCountK + blk]; - const float b_scale = QuantBScale[n * BlockCountK + blk]; - const int32_t zp = BZeroPoints != nullptr - ? static_cast(BZeroPoints[n * BlockCountK + blk]) - : static_cast(sq2::kDefaultSymmetricZeroPoint2Bit); - int32_t dot = 0; - for (size_t kk = 0; kk < local_len; ++kk) { - const int8_t qa = QuantAData[m * BlockCountK * kBlkLen32 + k + kk]; - const int32_t qb = - static_cast(BWeights[n * K + k + kk]) - zp; - dot += static_cast(qa) * qb; - } - acc += static_cast(dot) * a_scale * b_scale; - } - C[m * N + n] = acc; - } +void QuantizeA_Reference_BlkLen32(size_t M, size_t K, const float* A, + int8_t* QuantAData, float* QuantAScale) { + const size_t BlockCountK = (K + kBlkLen32 - 1) / kBlkLen32; + for (size_t m = 0; m < M; ++m) { + for (size_t k = 0, k_blk = 0; k < K; k += kBlkLen32, ++k_blk) { + const size_t local_len = std::min(K - k, kBlkLen32); + float amax = 0.0f; + for (size_t kk = 0; kk < local_len; ++kk) { + amax = std::max(amax, std::fabs(A[m * K + k + kk])); + } + constexpr float range_max = 127.0f; + const float scale = amax / range_max; + const float scale_recip = amax != 0.0f ? range_max / amax : 0.0f; + QuantAScale[m * BlockCountK + k_blk] = scale; + for (size_t kk = 0; kk < kBlkLen32; ++kk) { + const float a = (kk < local_len) ? A[m * K + k + kk] : 0.0f; + const float q = std::nearbyint(a * scale_recip); + QuantAData[m * BlockCountK * kBlkLen32 + k + kk] = + static_cast(std::clamp(q, -127.0f, 127.0f)); + } } + } } -void -RunW2Case_BlkLen32(size_t M, size_t N, size_t K, bool WithBias, uint32_t seed, - bool WithZeroPoints, W2KernelFn kernel, - const char* kernel_name) -{ - const size_t BlockCountK = (K + kBlkLen32 - 1) / kBlkLen32; - ASSERT_EQ(K % kBlkLen32, 0u) << "BlkLen32 test K must be a multiple of 32"; - - std::mt19937 rng(seed); - std::uniform_real_distribution a_dist(-1.0f, 1.0f); - std::uniform_int_distribution w_dist(0, 3); - std::uniform_real_distribution s_dist(0.05f, 0.5f); - - std::vector A(M * K); - for (auto& v : A) v = a_dist(rng); - - std::vector BWeights(N * K); - for (auto& v : BWeights) v = static_cast(w_dist(rng)); - - std::vector QuantBData(N * BlockCountK * kBlkBytes32, std::byte{0}); +void ReferenceGemm_W2_CompInt8_BlkLen32(size_t M, size_t N, size_t K, + const float* A, + const std::vector& BWeights, + const float* QuantBScale, + const uint8_t* BZeroPoints, + const float* Bias, + float* C) { + const size_t BlockCountK = (K + kBlkLen32 - 1) / kBlkLen32; + std::vector QuantAData(M * BlockCountK * kBlkLen32, int8_t{0}); + std::vector QuantAScale(M * BlockCountK, 0.0f); + QuantizeA_Reference_BlkLen32(M, K, A, QuantAData.data(), QuantAScale.data()); + + for (size_t m = 0; m < M; ++m) { for (size_t n = 0; n < N; ++n) { - for (size_t blk = 0; blk < BlockCountK; ++blk) { - uint8_t blk_weights[kBlkLen32]; - for (size_t kk = 0; kk < kBlkLen32; ++kk) { - blk_weights[kk] = BWeights[n * K + blk * kBlkLen32 + kk]; - } - PackSourceBlock_BlkLen32(blk_weights, - QuantBData.data() + (n * BlockCountK + blk) * kBlkBytes32); + float acc = (Bias != nullptr) ? Bias[n] : 0.0f; + for (size_t k = 0, blk = 0; k < K; k += kBlkLen32, ++blk) { + const size_t local_len = std::min(K - k, kBlkLen32); + const float a_scale = QuantAScale[m * BlockCountK + blk]; + const float b_scale = QuantBScale[n * BlockCountK + blk]; + const int32_t zp = BZeroPoints != nullptr + ? static_cast(BZeroPoints[n * BlockCountK + blk]) + : static_cast(sq2::kDefaultSymmetricZeroPoint2Bit); + int32_t dot = 0; + for (size_t kk = 0; kk < local_len; ++kk) { + const int8_t qa = QuantAData[m * BlockCountK * kBlkLen32 + k + kk]; + const int32_t qb = + static_cast(BWeights[n * K + k + kk]) - zp; + dot += static_cast(qa) * qb; } + acc += static_cast(dot) * a_scale * b_scale; + } + C[m * N + n] = acc; } + } +} - std::vector QuantBScale(N * BlockCountK); - for (auto& v : QuantBScale) v = s_dist(rng); - - std::vector BZeroPoints; - std::vector BZeroPointsPacked; - const uint8_t* BZeroPointsRef = nullptr; - const std::byte* BZeroPointsMlas = nullptr; - if (WithZeroPoints) { - BZeroPoints.resize(N * BlockCountK); - for (auto& v : BZeroPoints) v = static_cast(w_dist(rng)); - BZeroPointsRef = BZeroPoints.data(); - BZeroPointsPacked = PackW2ZeroPoints(N, BlockCountK, BZeroPoints); - BZeroPointsMlas = BZeroPointsPacked.data(); - } - - std::vector Bias; - const float* BiasPtr = nullptr; - if (WithBias) { - Bias.resize(N); - for (auto& v : Bias) v = a_dist(rng); - BiasPtr = Bias.data(); +void RunW2Case_BlkLen32(size_t M, size_t N, size_t K, bool WithBias, uint32_t seed, + bool WithZeroPoints, W2KernelFn kernel, + const char* kernel_name) { + const size_t BlockCountK = (K + kBlkLen32 - 1) / kBlkLen32; + ASSERT_EQ(K % kBlkLen32, 0u) << "BlkLen32 test K must be a multiple of 32"; + + std::mt19937 rng(seed); + std::uniform_real_distribution a_dist(-1.0f, 1.0f); + std::uniform_int_distribution w_dist(0, 3); + std::uniform_real_distribution s_dist(0.05f, 0.5f); + + std::vector A(M * K); + for (auto& v : A) v = a_dist(rng); + + std::vector BWeights(N * K); + for (auto& v : BWeights) v = static_cast(w_dist(rng)); + + std::vector QuantBData(N * BlockCountK * kBlkBytes32, std::byte{0}); + for (size_t n = 0; n < N; ++n) { + for (size_t blk = 0; blk < BlockCountK; ++blk) { + uint8_t blk_weights[kBlkLen32]; + for (size_t kk = 0; kk < kBlkLen32; ++kk) { + blk_weights[kk] = BWeights[n * K + blk * kBlkLen32 + kk]; + } + PackSourceBlock_BlkLen32(blk_weights, + QuantBData.data() + (n * BlockCountK + blk) * kBlkBytes32); } - - const size_t PackedSize = sq2::Q2BitGemmPackQuantBDataSize_Avx512( - N, K, kBlkLen32, WithZeroPoints, SQNBIT_CompInt8, nullptr); - ASSERT_GT(PackedSize, 0u) << "BlkLen32 pack size unsupported for shape"; - - std::vector PackedQuantBBuf(PackedSize, std::byte{0}); - PackedQuantBDataStruct packed_b( - PackedQuantBBuf.data(), N, BlockCountK, kBlkLen32, /*QuantAUnsigned=*/false); - + } + + std::vector QuantBScale(N * BlockCountK); + for (auto& v : QuantBScale) v = s_dist(rng); + + std::vector BZeroPoints; + std::vector BZeroPointsPacked; + const uint8_t* BZeroPointsRef = nullptr; + const std::byte* BZeroPointsMlas = nullptr; + if (WithZeroPoints) { + BZeroPoints.resize(N * BlockCountK); + for (auto& v : BZeroPoints) v = static_cast(w_dist(rng)); + BZeroPointsRef = BZeroPoints.data(); + BZeroPointsPacked = PackW2ZeroPoints(N, BlockCountK, BZeroPoints); + BZeroPointsMlas = BZeroPointsPacked.data(); + } + + std::vector Bias; + const float* BiasPtr = nullptr; + if (WithBias) { + Bias.resize(N); + for (auto& v : Bias) v = a_dist(rng); + BiasPtr = Bias.data(); + } + + const size_t PackedSize = sq2::Q2BitGemmPackQuantBDataSize_Avx512( + N, K, kBlkLen32, WithZeroPoints, SQNBIT_CompInt8, nullptr); + ASSERT_GT(PackedSize, 0u) << "BlkLen32 pack size unsupported for shape"; + + std::vector PackedQuantBBuf(PackedSize, std::byte{0}); + PackedQuantBDataStruct packed_b( + PackedQuantBBuf.data(), N, BlockCountK, kBlkLen32, /*QuantAUnsigned=*/false); + + sq2::SQ2BitGemmPackQuantBDataAndBlkSum_Scalar( + N, K, kBlkLen32, SQNBIT_CompInt8, + QuantBData.data(), /*scales=*/nullptr, + WithZeroPoints, /*zp=*/nullptr, + packed_b, /*tp=*/nullptr, /*cfg=*/nullptr); + sq2::SQ2BitGemmPackQuantBDataAndBlkSum_Scalar( + N, K, kBlkLen32, SQNBIT_CompInt8, + /*B=*/nullptr, QuantBScale.data(), + WithZeroPoints, /*zp=*/nullptr, + packed_b, /*tp=*/nullptr, /*cfg=*/nullptr); + if (WithZeroPoints) { sq2::SQ2BitGemmPackQuantBDataAndBlkSum_Scalar( N, K, kBlkLen32, SQNBIT_CompInt8, - QuantBData.data(), /*scales=*/nullptr, - WithZeroPoints, /*zp=*/nullptr, + /*B=*/nullptr, /*scales=*/nullptr, + WithZeroPoints, BZeroPointsMlas, packed_b, /*tp=*/nullptr, /*cfg=*/nullptr); - sq2::SQ2BitGemmPackQuantBDataAndBlkSum_Scalar( - N, K, kBlkLen32, SQNBIT_CompInt8, - /*B=*/nullptr, QuantBScale.data(), - WithZeroPoints, /*zp=*/nullptr, - packed_b, /*tp=*/nullptr, /*cfg=*/nullptr); - if (WithZeroPoints) { - sq2::SQ2BitGemmPackQuantBDataAndBlkSum_Scalar( - N, K, kBlkLen32, SQNBIT_CompInt8, - /*B=*/nullptr, /*scales=*/nullptr, - WithZeroPoints, BZeroPointsMlas, - packed_b, /*tp=*/nullptr, /*cfg=*/nullptr); - } - - std::vector QuantAData(M * BlockCountK * kBlkLen32, int8_t{0}); - std::vector QuantAScale(M * BlockCountK, 0.0f); - QuantizeA_Reference_BlkLen32(M, K, A.data(), QuantAData.data(), QuantAScale.data()); - - std::vector ABlockSum(M * BlockCountK, 0.0f); - for (size_t m = 0; m < M; ++m) { - for (size_t blk = 0; blk < BlockCountK; ++blk) { - int32_t sum = 0; - for (size_t kk = 0; kk < kBlkLen32; ++kk) { - sum += static_cast( - QuantAData[m * BlockCountK * kBlkLen32 + blk * kBlkLen32 + kk]); - } - ABlockSum[m * BlockCountK + blk] = - QuantAScale[m * BlockCountK + blk] * static_cast(sum); - } - } - - std::vector C(M * N, 0.0f); - kernel( - kBlkLen32, - reinterpret_cast(QuantAData.data()), - QuantAScale.data(), - packed_b.PackedQuantBData, - packed_b.PackedQuantBScale, - /*QuantBZeroPoint=*/nullptr, - C.data(), - M, N, /*CountK=*/K, BlockCountK, - BiasPtr, - /*ldc=*/N, - ABlockSum.data(), - packed_b.QuantBBlkSum); - - std::vector CRef(M * N, 0.0f); - ReferenceGemm_W2_CompInt8_BlkLen32(M, N, K, A.data(), BWeights, QuantBScale.data(), - BZeroPointsRef, BiasPtr, CRef.data()); - - const float abs_tol = 1e-4f; - const float rel_tol = 1e-4f; - for (size_t i = 0; i < M * N; ++i) { - const float diff = std::fabs(C[i] - CRef[i]); - const float bound = abs_tol + rel_tol * std::fabs(CRef[i]); - ASSERT_LE(diff, bound) - << "BlkLen32 " << kernel_name << " mismatch at i=" << i - << " (m=" << (i / N) << ", n=" << (i % N) << ")" - << " out=" << C[i] << " ref=" << CRef[i] - << " M=" << M << " N=" << N << " K=" << K - << " WithBias=" << WithBias - << " WithZeroPoints=" << WithZeroPoints; + } + + std::vector QuantAData(M * BlockCountK * kBlkLen32, int8_t{0}); + std::vector QuantAScale(M * BlockCountK, 0.0f); + QuantizeA_Reference_BlkLen32(M, K, A.data(), QuantAData.data(), QuantAScale.data()); + + std::vector ABlockSum(M * BlockCountK, 0.0f); + for (size_t m = 0; m < M; ++m) { + for (size_t blk = 0; blk < BlockCountK; ++blk) { + int32_t sum = 0; + for (size_t kk = 0; kk < kBlkLen32; ++kk) { + sum += static_cast( + QuantAData[m * BlockCountK * kBlkLen32 + blk * kBlkLen32 + kk]); + } + ABlockSum[m * BlockCountK + blk] = + QuantAScale[m * BlockCountK + blk] * static_cast(sum); } + } + + std::vector C(M * N, 0.0f); + kernel( + kBlkLen32, + reinterpret_cast(QuantAData.data()), + QuantAScale.data(), + packed_b.PackedQuantBData, + packed_b.PackedQuantBScale, + /*QuantBZeroPoint=*/nullptr, + C.data(), + M, N, /*CountK=*/K, BlockCountK, + BiasPtr, + /*ldc=*/N, + ABlockSum.data(), + packed_b.QuantBBlkSum); + + std::vector CRef(M * N, 0.0f); + ReferenceGemm_W2_CompInt8_BlkLen32(M, N, K, A.data(), BWeights, QuantBScale.data(), + BZeroPointsRef, BiasPtr, CRef.data()); + + const float abs_tol = 1e-4f; + const float rel_tol = 1e-4f; + for (size_t i = 0; i < M * N; ++i) { + const float diff = std::fabs(C[i] - CRef[i]); + const float bound = abs_tol + rel_tol * std::fabs(CRef[i]); + ASSERT_LE(diff, bound) + << "BlkLen32 " << kernel_name << " mismatch at i=" << i + << " (m=" << (i / N) << ", n=" << (i % N) << ")" + << " out=" << C[i] << " ref=" << CRef[i] + << " M=" << M << " N=" << N << " K=" << K + << " WithBias=" << WithBias + << " WithZeroPoints=" << WithZeroPoints; + } } // K-shape constraint: K multiple of 32 (BlkLen=32). Covers BlockCountK in // {1, 2, 3, 4, 8, 16, 32, 64} -- both K-tail variants (BlockCountK not a // multiple of 4) and exact block-group multiples. -constexpr struct { size_t M, N, K; } kSimdShapes_BlkLen32[] = { - {1, 16, 32}, // R1, BlockCountK=1 - {1, 32, 64}, // R1, BlockCountK=2 (K-tail, no full group) - {1, 1024, 256}, // R1, BlockCountK=8 - {1, 1024, 1024}, // R1, BlockCountK=32 - {2, 16, 32}, - {2, 32, 64}, - {2, 64, 128}, // BlockCountK=4 (one full group) - {3, 16, 64}, // R2 head + R1 tail - {3, 384, 256}, - {4, 16, 64}, - {4, 32, 64}, - {16, 64, 128}, - {32, 128, 128}, +constexpr struct { + size_t M, N, K; +} kSimdShapes_BlkLen32[] = { + {1, 16, 32}, // R1, BlockCountK=1 + {1, 32, 64}, // R1, BlockCountK=2 (K-tail, no full group) + {1, 1024, 256}, // R1, BlockCountK=8 + {1, 1024, 1024}, // R1, BlockCountK=32 + {2, 16, 32}, + {2, 32, 64}, + {2, 64, 128}, // BlockCountK=4 (one full group) + {3, 16, 64}, // R2 head + R1 tail + {3, 384, 256}, + {4, 16, 64}, + {4, 32, 64}, + {16, 64, 128}, + {32, 128, 128}, // M=128 prefill - {128, 1024, 256}, {128, 1024, 1024}, - {128, 192, 256}, {128, 384, 256}, + {128, 1024, 256}, + {128, 1024, 1024}, + {128, 192, 256}, + {128, 384, 256}, // K-tail (BlockCountK not multiple of 4) - { 2, 16, 96}, // tail=3 - { 4, 16, 96}, - {128, 1024, 96}, - { 2, 16, 160}, // tail=1 - { 4, 16, 160}, - { 2, 16, 192}, // tail=2 (BlockCountK=6) - { 4, 16, 192}, + {2, 16, 96}, // tail=3 + {4, 16, 96}, + {128, 1024, 96}, + {2, 16, 160}, // tail=1 + {4, 16, 160}, + {2, 16, 192}, // tail=2 (BlockCountK=6) + {4, 16, 192}, // N-tail - { 1, 1, 64}, - { 1, 3, 64}, - { 4, 3, 64}, - { 1, 17, 64}, - { 4, 17, 64}, - {128, 17, 64}, - { 1, 33, 64}, - { 4, 33, 64}, - {128, 33, 64}, - { 1, 18, 64}, - { 4, 18, 64}, - {128, 19, 64}, + {1, 1, 64}, + {1, 3, 64}, + {4, 3, 64}, + {1, 17, 64}, + {4, 17, 64}, + {128, 17, 64}, + {1, 33, 64}, + {4, 33, 64}, + {128, 33, 64}, + {1, 18, 64}, + {4, 18, 64}, + {128, 19, 64}, // N-tail + K-tail - { 1, 17, 96}, - { 4, 33, 96}, - {128, 19, 160}, + {1, 17, 96}, + {4, 33, 96}, + {128, 19, 160}, }; } // namespace -TEST(MlasSq2BitTest, Scalar_BlkLen32) -{ - for (uint32_t seed : {0xC0FFEEu, 0xBADC0DEu}) { - for (const auto& s : kSimdShapes_BlkLen32) { - for (bool bias : {false, true}) { - RunW2Case_BlkLen32(s.M, s.N, s.K, bias, seed + (bias ? 1u : 0u), - /*WithZeroPoints=*/false, - sq2::SQ2BitGemmKernel_BlkSum_CompInt8_Scalar, - "Scalar"); - } - } +TEST(MlasSq2BitTest, Scalar_BlkLen32) { + for (uint32_t seed : {0xC0FFEEu, 0xBADC0DEu}) { + for (const auto& s : kSimdShapes_BlkLen32) { + for (bool bias : {false, true}) { + RunW2Case_BlkLen32(s.M, s.N, s.K, bias, seed + (bias ? 1u : 0u), + /*WithZeroPoints=*/false, + sq2::SQ2BitGemmKernel_BlkSum_CompInt8_Scalar, + "Scalar"); + } } + } } -TEST(MlasSq2BitTest, Scalar_BlkLen32_WithZeroPoints) -{ - for (uint32_t seed : {0xC0FFEEu, 0xBADC0DEu}) { - for (const auto& s : kSimdShapes_BlkLen32) { - for (bool bias : {false, true}) { - RunW2Case_BlkLen32(s.M, s.N, s.K, bias, seed + (bias ? 1u : 0u), - /*WithZeroPoints=*/true, - sq2::SQ2BitGemmKernel_BlkSum_CompInt8_Scalar, - "Scalar"); - } - } +TEST(MlasSq2BitTest, Scalar_BlkLen32_WithZeroPoints) { + for (uint32_t seed : {0xC0FFEEu, 0xBADC0DEu}) { + for (const auto& s : kSimdShapes_BlkLen32) { + for (bool bias : {false, true}) { + RunW2Case_BlkLen32(s.M, s.N, s.K, bias, seed + (bias ? 1u : 0u), + /*WithZeroPoints=*/true, + sq2::SQ2BitGemmKernel_BlkSum_CompInt8_Scalar, + "Scalar"); + } } + } } -TEST(MlasSq2BitTest, BlkLen32_Avx512) -{ - if (!GetMlasPlatform().Avx512Supported_) { - GTEST_SKIP() << "AVX-512BW (DQ/VL) not available on this host"; - } - for (uint32_t seed : {0xC0FFEEu, 0xBADC0DEu}) { - for (const auto& s : kSimdShapes_BlkLen32) { - for (bool bias : {false, true}) { - RunW2Case_BlkLen32(s.M, s.N, s.K, bias, seed + (bias ? 1u : 0u), - /*WithZeroPoints=*/false, - sq2::SQ2BitGemmKernel_BlkSum_CompInt8_Avx512_TestEntry, - "AVX-512BW"); - } - } +TEST(MlasSq2BitTest, BlkLen32_Avx512) { + if (!GetMlasPlatform().Avx512Supported_) { + GTEST_SKIP() << "AVX-512BW (DQ/VL) not available on this host"; + } + for (uint32_t seed : {0xC0FFEEu, 0xBADC0DEu}) { + for (const auto& s : kSimdShapes_BlkLen32) { + for (bool bias : {false, true}) { + RunW2Case_BlkLen32(s.M, s.N, s.K, bias, seed + (bias ? 1u : 0u), + /*WithZeroPoints=*/false, + sq2::SQ2BitGemmKernel_BlkSum_CompInt8_Avx512_TestEntry, + "AVX-512BW"); + } } + } } -TEST(MlasSq2BitTest, BlkLen32_Avx512_WithZeroPoints) -{ - if (!GetMlasPlatform().Avx512Supported_) { - GTEST_SKIP() << "AVX-512BW (DQ/VL) not available on this host"; - } - for (uint32_t seed : {0xC0FFEEu, 0xBADC0DEu}) { - for (const auto& s : kSimdShapes_BlkLen32) { - for (bool bias : {false, true}) { - RunW2Case_BlkLen32(s.M, s.N, s.K, bias, seed + (bias ? 1u : 0u), - /*WithZeroPoints=*/true, - sq2::SQ2BitGemmKernel_BlkSum_CompInt8_Avx512_TestEntry, - "AVX-512BW"); - } - } +TEST(MlasSq2BitTest, BlkLen32_Avx512_WithZeroPoints) { + if (!GetMlasPlatform().Avx512Supported_) { + GTEST_SKIP() << "AVX-512BW (DQ/VL) not available on this host"; + } + for (uint32_t seed : {0xC0FFEEu, 0xBADC0DEu}) { + for (const auto& s : kSimdShapes_BlkLen32) { + for (bool bias : {false, true}) { + RunW2Case_BlkLen32(s.M, s.N, s.K, bias, seed + (bias ? 1u : 0u), + /*WithZeroPoints=*/true, + sq2::SQ2BitGemmKernel_BlkSum_CompInt8_Avx512_TestEntry, + "AVX-512BW"); + } } + } } -TEST(MlasSq2BitTest, BlkLen32_Avx512Vnni) -{ - if (GetMlasPlatform().QNBitGemmDispatch != &MlasSQNBitGemmDispatchAvx512vnni) { - GTEST_SKIP() << "AVX-512-VNNI not selected as the active dispatch on this host"; - } - for (uint32_t seed : {0xC0FFEEu, 0xBADC0DEu}) { - for (const auto& s : kSimdShapes_BlkLen32) { - for (bool bias : {false, true}) { - RunW2Case_BlkLen32(s.M, s.N, s.K, bias, seed + (bias ? 1u : 0u), - /*WithZeroPoints=*/false, - sq2::SQ2BitGemmKernel_BlkSum_CompInt8_Avx512Vnni_TestEntry, - "AVX-512-VNNI"); - } - } +TEST(MlasSq2BitTest, BlkLen32_Avx512Vnni) { + if (GetMlasPlatform().QNBitGemmDispatch != &MlasSQNBitGemmDispatchAvx512vnni) { + GTEST_SKIP() << "AVX-512-VNNI not selected as the active dispatch on this host"; + } + for (uint32_t seed : {0xC0FFEEu, 0xBADC0DEu}) { + for (const auto& s : kSimdShapes_BlkLen32) { + for (bool bias : {false, true}) { + RunW2Case_BlkLen32(s.M, s.N, s.K, bias, seed + (bias ? 1u : 0u), + /*WithZeroPoints=*/false, + sq2::SQ2BitGemmKernel_BlkSum_CompInt8_Avx512Vnni_TestEntry, + "AVX-512-VNNI"); + } } + } } -TEST(MlasSq2BitTest, BlkLen32_Avx512Vnni_WithZeroPoints) -{ - if (GetMlasPlatform().QNBitGemmDispatch != &MlasSQNBitGemmDispatchAvx512vnni) { - GTEST_SKIP() << "AVX-512-VNNI not selected as the active dispatch on this host"; - } - for (uint32_t seed : {0xC0FFEEu, 0xBADC0DEu}) { - for (const auto& s : kSimdShapes_BlkLen32) { - for (bool bias : {false, true}) { - RunW2Case_BlkLen32(s.M, s.N, s.K, bias, seed + (bias ? 1u : 0u), - /*WithZeroPoints=*/true, - sq2::SQ2BitGemmKernel_BlkSum_CompInt8_Avx512Vnni_TestEntry, - "AVX-512-VNNI"); - } - } +TEST(MlasSq2BitTest, BlkLen32_Avx512Vnni_WithZeroPoints) { + if (GetMlasPlatform().QNBitGemmDispatch != &MlasSQNBitGemmDispatchAvx512vnni) { + GTEST_SKIP() << "AVX-512-VNNI not selected as the active dispatch on this host"; + } + for (uint32_t seed : {0xC0FFEEu, 0xBADC0DEu}) { + for (const auto& s : kSimdShapes_BlkLen32) { + for (bool bias : {false, true}) { + RunW2Case_BlkLen32(s.M, s.N, s.K, bias, seed + (bias ? 1u : 0u), + /*WithZeroPoints=*/true, + sq2::SQ2BitGemmKernel_BlkSum_CompInt8_Avx512Vnni_TestEntry, + "AVX-512-VNNI"); + } } + } } + +#endif // defined(MLAS_TARGET_AMD64) From 1587c2b2d8ecaaf10a80d3730141b8b4c8bc4fc2 Mon Sep 17 00:00:00 2001 From: Hari Seshadri Date: Mon, 15 Jun 2026 18:45:10 -0700 Subject: [PATCH 12/17] Cosmetic cleanups --- .../cpu/quantization/matmul_nbits.cc | 7 ++- onnxruntime/core/mlas/lib/qnbitgemm.cpp | 2 +- onnxruntime/core/mlas/lib/qnbitgemm.h | 31 +++++++---- .../mlas/lib/sqnbitgemm_kernel_avx512.cpp | 15 ++++-- .../mlas/lib/sqnbitgemm_kernel_avx512_2bit.h | 15 +++--- .../mlas/lib/sqnbitgemm_kernel_avx512vnni.cpp | 19 ++++--- .../test/contrib_ops/matmul_2bits_test.cc | 19 +++++-- onnxruntime/test/mlas/bench/bench_lutgemm.cpp | 13 ++--- .../test/mlas/bench/bench_qnbitgemm.cpp | 24 +++++---- .../unittest/test_sqnbitgemm_2bit_gemm.cpp | 51 ++++++++++--------- 10 files changed, 118 insertions(+), 78 deletions(-) diff --git a/onnxruntime/contrib_ops/cpu/quantization/matmul_nbits.cc b/onnxruntime/contrib_ops/cpu/quantization/matmul_nbits.cc index 6f33b34a049fb..6bd1690fca815 100644 --- a/onnxruntime/contrib_ops/cpu/quantization/matmul_nbits.cc +++ b/onnxruntime/contrib_ops/cpu/quantization/matmul_nbits.cc @@ -364,10 +364,9 @@ Status MatMulNBits::PrePack(const Tensor& tensor, int input_idx, /*out*/ All // Guard: for LUT-eligible nodes, scales/ZP are already packed inside // packed_b_ by the LUT branch above (or by the LUT scale-pack path at // the bottom of this function). Re-running the non-LUT pack here would - // overwrite the LUT-packed buffer with W2 layout bytes, causing heap - // corruption when the LUT compute path later reads back the packed-B - // contents. prefer_lut_gemm_ is gated to T1==float (see ctor), so - // checking it here is sufficient. + // corrupt the LUT-packed buffer (overwrite the LUT layout with W2 layout + // bytes), so the LUT compute path would then read garbage. prefer_lut_gemm_ + // is gated to T1==float (see ctor), so checking it here is sufficient. bool should_pack_scale_and_zp_inputs = [&]() { #if defined(MLAS_TARGET_AMD64_IX86) return true; diff --git a/onnxruntime/core/mlas/lib/qnbitgemm.cpp b/onnxruntime/core/mlas/lib/qnbitgemm.cpp index 5f1357020accd..92949f97ab343 100644 --- a/onnxruntime/core/mlas/lib/qnbitgemm.cpp +++ b/onnxruntime/core/mlas/lib/qnbitgemm.cpp @@ -1003,7 +1003,7 @@ SQ2BitGemm_CompInt8( // The packed-B-data and per-block-scale buffers may use a layout that // addresses each N-col at a stride larger than `k_blks` (e.g. the AVX-512 - // W2 kernel rounds up to a multiple of 4 to amortise unpack across 4 + // W2 kernel rounds up to a multiple of 4 to amortize unpack across 4 // K-blocks). Ask the dispatch for the effective per-N-col block count; // default to logical k_blks if the dispatch doesn't override. // BlkSum keeps the logical stride because the SGEMM correction step diff --git a/onnxruntime/core/mlas/lib/qnbitgemm.h b/onnxruntime/core/mlas/lib/qnbitgemm.h index 0e56e20579d8b..b9c88ab1c7ede 100644 --- a/onnxruntime/core/mlas/lib/qnbitgemm.h +++ b/onnxruntime/core/mlas/lib/qnbitgemm.h @@ -46,19 +46,30 @@ MlasAlignAddress(void* addr, const size_t alignment) return addr; } +// Number of consecutive K-blocks bundled into a single packed slot by the +// AVX-512 W2 native kernel. The pack helper rounds BlockCountK up to a +// multiple of this value so the kernel can iterate the padded count without +// bounds checks. This must match onnxruntime::mlas::sq2bit_avx512::kBlockGroupBlks +// in sqnbitgemm_kernel_avx512_2bit.h; a static_assert at each dispatch wiring +// site keeps the two definitions in sync. +inline constexpr size_t kSq2BitAvx512WeightKBlockGroup = 4; + template struct PackedQuantBDataStruct { PackedQuantBDataStruct(void* PackedQuantBWorkspace, size_t N, size_t BlockCountK, size_t BlkLen, bool QuantAUnsigned) : QuantBWorkspace_(PackedQuantBWorkspace), N_(N), BlockCountK_(BlockCountK), BlkLen_(BlkLen) { - // For 2-bit weights, the AVX-512 W2 packed layout groups 4 consecutive - // K-blocks into a single 64-byte slot so the SIMD unpack is one ZMM - // load + four fixed shift/mask pairs. The pack-size dispatch - // (Q2BitGemmPackQuantBDataSize_Avx512) rounds BlockCountK up to a - // multiple of 4 internally; we mirror that rounding here so the - // allocated buffer always matches the slab layout computed below. + // For 2-bit weights, the AVX-512 W2 packed layout groups + // kSq2BitAvx512WeightKBlockGroup (=4) consecutive K-blocks into a + // single 64-byte slot so the SIMD unpack is one ZMM load + four fixed + // shift/mask pairs. The pack-size dispatch + // (Q2BitGemmPackQuantBDataSize_Avx512) rounds BlockCountK up to that + // multiple internally; we mirror that rounding here so the allocated + // buffer always matches the slab layout computed below. const size_t EffectiveBlockCountK = - (BlkBitWidth == 2) ? ((BlockCountK + 3) / 4) * 4 : BlockCountK; + (BlkBitWidth == 2) + ? MlasDivRoundup(BlockCountK, kSq2BitAvx512WeightKBlockGroup) * kSq2BitAvx512WeightKBlockGroup + : BlockCountK; const size_t PackedQuantBDataSize = N * EffectiveBlockCountK * MlasQNBitBlkDataSizeInBytes(BlkBitWidth, BlkLen); size_t BlkSumSize = MlasDivRoundup(N, 16) * EffectiveBlockCountK * 16 * sizeof(T); #if defined(MLAS_TARGET_AMD64_IX86) @@ -475,7 +486,9 @@ struct MLAS_QNBIT_GEMM_DISPATCH { // /** Gets size of packed quantized B data containing 2-bit integers. See MlasQNBitGemmPackQuantBDataSize(). */ - Q4BitGemmPackQuantBDataSize_Fn* Q2BitGemmPackQuantBDataSize = nullptr; + // The W2 pack-size signature matches the W4 one; alias for self-documentation. + using Q2BitGemmPackQuantBDataSize_Fn = Q4BitGemmPackQuantBDataSize_Fn; + Q2BitGemmPackQuantBDataSize_Fn* Q2BitGemmPackQuantBDataSize = nullptr; /** Packs quantized B data + per-block sums for the 2-bit CompInt8 kernel. */ typedef void(SQ2BitGemmPackQuantBDataAndSumBlk_Fn)( @@ -504,7 +517,7 @@ struct MLAS_QNBIT_GEMM_DISPATCH { * stride of `effective_block_count * `, * and so does the dispatcher's per-N-tile pointer arithmetic. The * AVX-512 W2 kernel rounds BlockCountK up to a multiple of 4 - * internally to amortise unpack cost across 4 consecutive K-blocks; + * internally to amortize unpack cost across 4 consecutive K-blocks; * the buffer is sized accordingly (see PackedQuantBDataStruct, which * always pads for BlkBitWidth==2) so the dispatcher must use the * matching stride or it will step past the data when n != 0. diff --git a/onnxruntime/core/mlas/lib/sqnbitgemm_kernel_avx512.cpp b/onnxruntime/core/mlas/lib/sqnbitgemm_kernel_avx512.cpp index 0c946814cda83..ada1075bcc176 100644 --- a/onnxruntime/core/mlas/lib/sqnbitgemm_kernel_avx512.cpp +++ b/onnxruntime/core/mlas/lib/sqnbitgemm_kernel_avx512.cpp @@ -480,12 +480,14 @@ SQ8BitGemmPackQuantBDataAndBlkSum512( } // -// Unit-test entry point for the AVX-512BW (non-VNNI) W2 kernel. +// BlkLen-routing wrapper for the W2 CompInt8 AVX-512BW (non-VNNI) dispatch +// entry. Production code reaches this via the MLAS dispatch table; tests +// call it directly via the namespace. // Sibling of the VNNI variant in sqnbitgemm_kernel_avx512vnni.cpp. // namespace onnxruntime::mlas::sq2bit_avx512 { size_t MLASCALL -SQ2BitGemmKernel_BlkSum_CompInt8_Avx512_TestEntry( +SQ2BitGemmKernel_BlkSum_CompInt8_Avx512_Dispatch( size_t BlkLen, const std::byte* QuantA, const float* QuantAScale, @@ -542,10 +544,15 @@ const MLAS_QNBIT_GEMM_DISPATCH MlasSQNBitGemmDispatchAvx512 = []() { // once. Packs B with the K dimension rounded up to a multiple of // kBlockGroupBlks (= 4); the kernel iterates the padded count, so the // dispatcher reports the rounded-up stride via Q2BitGemmEffectiveBlockCountK. + static_assert( + onnxruntime::mlas::sq2bit_avx512::kBlockGroupBlks == kSq2BitAvx512WeightKBlockGroup, + "kBlockGroupBlks (kernel-internal) must match kSq2BitAvx512WeightKBlockGroup (qnbitgemm.h)."); d.Q2BitGemmPackQuantBDataSize = onnxruntime::mlas::sq2bit_avx512::Q2BitGemmPackQuantBDataSize_Avx512; d.SQ2BitGemmPackQuantBDataAndBlkSum = onnxruntime::mlas::sq2bit_avx512::SQ2BitGemmPackQuantBDataAndBlkSum_Scalar; - d.SQ2BitGemmKernel_BlkSum_CompInt8 = onnxruntime::mlas::sq2bit_avx512::SQ2BitGemmKernel_BlkSum_CompInt8_Avx512_TestEntry; - d.Q2BitGemmEffectiveBlockCountK = [](size_t BlockCountK) { return ((BlockCountK + 3) / 4) * 4; }; + d.SQ2BitGemmKernel_BlkSum_CompInt8 = onnxruntime::mlas::sq2bit_avx512::SQ2BitGemmKernel_BlkSum_CompInt8_Avx512_Dispatch; + d.Q2BitGemmEffectiveBlockCountK = [](size_t BlockCountK) { + return MlasDivRoundup(BlockCountK, kSq2BitAvx512WeightKBlockGroup) * kSq2BitAvx512WeightKBlockGroup; + }; return d; }(); diff --git a/onnxruntime/core/mlas/lib/sqnbitgemm_kernel_avx512_2bit.h b/onnxruntime/core/mlas/lib/sqnbitgemm_kernel_avx512_2bit.h index 152f33d3f3228..67c8fdc2734c9 100644 --- a/onnxruntime/core/mlas/lib/sqnbitgemm_kernel_avx512_2bit.h +++ b/onnxruntime/core/mlas/lib/sqnbitgemm_kernel_avx512_2bit.h @@ -35,8 +35,9 @@ Module Name: * BlockCountK must be a multiple of kBlockGroupBlks = 4. The path returns 0 from the pack-size helper for non-multiples; the caller falls back to the existing W2 path. - * The customer model's K dimensions (384, 1024, 4096) are all multiples - of 256 (= 4 * 64), so all customer shapes satisfy this constraint. + * Typical W2 production K dimensions (e.g. 384, 1024, 4096) are all + multiples of 256 (= kBlockGroupBlks * 64) and therefore satisfy this + constraint. --*/ @@ -438,13 +439,15 @@ SQ2BitGemmKernel_BlkSum_CompInt8_Scalar( ); // -// Unit-test forwarders for the AVX-512 SIMD block-group kernels. Same gating -// rules as the existing W2 test entries: the caller MUST verify +// BlkLen-routing wrappers for the W2 CompInt8 dispatch entries. +// Production code calls these via the MLAS dispatch table +// (MlasSQNBitGemmDispatchAvx512 / Avx512vnni); tests call them directly +// via the namespace. The caller MUST verify // GetMlasPlatform().Avx512Supported_ (and, for the VNNI variant, that the // active dispatch is the AVX-512-VNNI one) before invoking these symbols. // size_t MLASCALL -SQ2BitGemmKernel_BlkSum_CompInt8_Avx512_TestEntry( +SQ2BitGemmKernel_BlkSum_CompInt8_Avx512_Dispatch( size_t BlkLen, const std::byte* QuantA, const float* QuantAScale, @@ -463,7 +466,7 @@ SQ2BitGemmKernel_BlkSum_CompInt8_Avx512_TestEntry( ); size_t MLASCALL -SQ2BitGemmKernel_BlkSum_CompInt8_Avx512Vnni_TestEntry( +SQ2BitGemmKernel_BlkSum_CompInt8_Avx512Vnni_Dispatch( size_t BlkLen, const std::byte* QuantA, const float* QuantAScale, diff --git a/onnxruntime/core/mlas/lib/sqnbitgemm_kernel_avx512vnni.cpp b/onnxruntime/core/mlas/lib/sqnbitgemm_kernel_avx512vnni.cpp index f617eac4b5211..662a953de2de4 100644 --- a/onnxruntime/core/mlas/lib/sqnbitgemm_kernel_avx512vnni.cpp +++ b/onnxruntime/core/mlas/lib/sqnbitgemm_kernel_avx512vnni.cpp @@ -462,14 +462,14 @@ SQ8BitGemmPackQuantBDataAndBlkSum512vnni( } // -// Unit-test entry point for the AVX-512-VNNI W2 kernel -// (sqnbitgemm_kernel_avx512_2bit_blklen64.h). Exposed for -// direct invocation from tests; production wiring (Phase 4) will add the -// runtime dispatcher integration. +// BlkLen-routing wrapper for the W2 CompInt8 AVX-512-VNNI dispatch entry +// (sqnbitgemm_kernel_avx512_2bit_blklen64.h and friends). Production code +// reaches this via the MLAS dispatch table; tests call it directly via the +// namespace. // namespace onnxruntime::mlas::sq2bit_avx512 { size_t MLASCALL -SQ2BitGemmKernel_BlkSum_CompInt8_Avx512Vnni_TestEntry( +SQ2BitGemmKernel_BlkSum_CompInt8_Avx512Vnni_Dispatch( size_t BlkLen, const std::byte* QuantA, const float* QuantAScale, @@ -528,10 +528,15 @@ const MLAS_QNBIT_GEMM_DISPATCH MlasSQNBitGemmDispatchAvx512vnni = []() { // 64-byte ZMM load + four fixed shift/mask pairs to unpack 4 K-blocks at // once, with VNNI's `vpdpbusd` for the integer MAC. See the matching // comment in sqnbitgemm_kernel_avx512.cpp for the K-padding contract. + static_assert( + onnxruntime::mlas::sq2bit_avx512::kBlockGroupBlks == kSq2BitAvx512WeightKBlockGroup, + "kBlockGroupBlks (kernel-internal) must match kSq2BitAvx512WeightKBlockGroup (qnbitgemm.h)."); d.Q2BitGemmPackQuantBDataSize = onnxruntime::mlas::sq2bit_avx512::Q2BitGemmPackQuantBDataSize_Avx512; d.SQ2BitGemmPackQuantBDataAndBlkSum = onnxruntime::mlas::sq2bit_avx512::SQ2BitGemmPackQuantBDataAndBlkSum_Scalar; - d.SQ2BitGemmKernel_BlkSum_CompInt8 = onnxruntime::mlas::sq2bit_avx512::SQ2BitGemmKernel_BlkSum_CompInt8_Avx512Vnni_TestEntry; - d.Q2BitGemmEffectiveBlockCountK = [](size_t BlockCountK) { return ((BlockCountK + 3) / 4) * 4; }; + d.SQ2BitGemmKernel_BlkSum_CompInt8 = onnxruntime::mlas::sq2bit_avx512::SQ2BitGemmKernel_BlkSum_CompInt8_Avx512Vnni_Dispatch; + d.Q2BitGemmEffectiveBlockCountK = [](size_t BlockCountK) { + return MlasDivRoundup(BlockCountK, kSq2BitAvx512WeightKBlockGroup) * kSq2BitAvx512WeightKBlockGroup; + }; return d; }(); diff --git a/onnxruntime/test/contrib_ops/matmul_2bits_test.cc b/onnxruntime/test/contrib_ops/matmul_2bits_test.cc index 1cdb6d4b247e4..9deb064a90853 100644 --- a/onnxruntime/test/contrib_ops/matmul_2bits_test.cc +++ b/onnxruntime/test/contrib_ops/matmul_2bits_test.cc @@ -1098,6 +1098,15 @@ TEST(MatMul2Bits, Float32_2b_Accuracy4) { TestMatMul2BitsTyped(); } +// NOTE on host-coverage of the new W2 op-level tests below. +// These tests run unconditionally on DefaultCpuExecutionProvider() and compute +// the expected output via dequant + matmul (see TestOptions2Bits / RunTest2Bits). +// The new native AVX-512(+VNNI) W2 kernel is dispatched only on hosts where +// GetMlasPlatform().Avx512Supported_ is true; on other hosts the op falls back +// to dequant + SGEMM and the same correctness check still passes. So these +// tests guard correctness everywhere and exercise the new SIMD path on +// AVX-512(+VNNI) CI machines and dev boxes. + // MatMulNBits operator-level coverage for the native AVX-512 W2 CPU kernel. // The kernel is gated on (BlkBitWidth=2, BlkLen=64, ComputeType=SQNBIT_CompInt8) // and is reachable through the public MatMulNBits op when accuracy_level=4 and @@ -1110,7 +1119,7 @@ TEST(MatMul2Bits, Float32_2b_Accuracy4) { // * Single block (BlockCountK=1) -> K=64 // * BlockCountK=2,3 not multiple of 4 -> K=128, K=192 (K-tail handler) // * Exact block-group multiples -> K=256, K=512, K=1024 -// * BlockCountK=6 (1 full + 2 tail) -> K=384 (matches customer model) +// * BlockCountK=6 (1 full + 2 tail) -> K=384 (non-aligned tail case) // and combinations of M ∈ {1 (decode), 2, 4, 100 (prefill)} and // N ∈ {16 (N-tail R1xC1/R2xC1), 32, 288, 1024 (kNCols4 multiples)}. // @@ -1134,10 +1143,10 @@ TEST(MatMul2Bits, Float32_2b_BlkLen64_Accuracy4) { TestMatMul2BitsTyped(); TestMatMul2BitsTyped(); - // Customer-model proportion: BlockCountK=6 = 1 full block-group + 2-block tail. + // Realistic-shape proportion: BlockCountK=6 = 1 full block-group + 2-block tail. TestMatMul2BitsTyped(); - // Larger M (multiple R2 tile iterations) and customer-shape K=1024. + // Larger M (multiple R2 tile iterations) and representative K=1024. TestMatMul2BitsTyped(); TestMatMul2BitsTyped(); } @@ -1194,7 +1203,7 @@ TEST(MatMul2Bits, Float32_2b_BlkLen128_Accuracy4) { // BlockCountK=6 = 1 full block-group + 2-block tail. TestMatMul2BitsTyped(); - // Larger M (multi-iter R2 tile) at customer-shape proportions. + // Larger M (multi-iter R2 tile) at representative-shape proportions. TestMatMul2BitsTyped(); TestMatMul2BitsTyped(); } @@ -1221,7 +1230,7 @@ TEST(MatMul2Bits, Float32_2b_BlkLen128_Accuracy0) { // MatMulNBits op-level coverage for the BlkLen=32 path. K multiples of 32. // Same coverage matrix as BlkLen=64/128 (single-block, K-tail, exact group -// multiples, customer-shape proportions, single-row decode, M=100 prefill). +// multiples, representative-shape proportions, single-row decode, M=100 prefill). TEST(MatMul2Bits, Float32_2b_BlkLen32_Accuracy4) { // Single-block K (K = BlkLen = 32). TestMatMul2BitsTyped(); diff --git a/onnxruntime/test/mlas/bench/bench_lutgemm.cpp b/onnxruntime/test/mlas/bench/bench_lutgemm.cpp index 422df80e9d247..dec40a9124772 100644 --- a/onnxruntime/test/mlas/bench/bench_lutgemm.cpp +++ b/onnxruntime/test/mlas/bench/bench_lutgemm.cpp @@ -233,7 +233,8 @@ static void LutGemmComputeArgs(benchmark::internal::Benchmark* b) { }); } -// Customer 2-bit MatMulNBits shapes (BlkLen=64). Five distinct (K, N) pairs: +// Representative 2-bit MatMulNBits shapes (BlkLen=64) drawn from real W2 +// production models. Five distinct (K, N) pairs: // (K=384, N=1024): 20 nodes // (K=1024, N=192): 40 nodes // (K=1024, N=384): 20 nodes @@ -241,12 +242,12 @@ static void LutGemmComputeArgs(benchmark::internal::Benchmark* b) { // (K=4096, N=1024): 20 nodes // Covers both M=1 (decode) and M=128 (prefill) so the LUT path can be // compared apples-to-apples against the W4 CompInt8 and W2 kernels -// (QNBITGEMM/QNBitGemmCustomerArgs and -// QNBITGEMM/QNBit2BitCustomerArgs). -static void LutGemmCustomerArgs(benchmark::internal::Benchmark* b) { +// (QNBITGEMM/QNBitGemmRealisticShapesArgs and +// QNBITGEMM/QNBit2BitRealisticShapesArgs). +static void LutGemmRealisticShapesArgs(benchmark::internal::Benchmark* b) { b->ArgNames(lutgemm_compute_arg_names); // Separate Args() entries so we only run the exact (M, K, N) tuples that - // appear in the customer model. + // appear in the representative production model. const int64_t BlkLen = 64; const int64_t Threads = 8; const int64_t HasZP = 0; @@ -267,7 +268,7 @@ static void LutGemmCustomerArgs(benchmark::internal::Benchmark* b) { if (is_lutgemm_supported) { BENCHMARK(LUTGEMM_PACK<2>)->Apply(LutGemmPackArgs)->UseRealTime(); BENCHMARK(LUTGEMM_COMPUTE<2>)->Apply(LutGemmComputeArgs)->UseRealTime(); - BENCHMARK(LUTGEMM_COMPUTE<2>)->Apply(LutGemmCustomerArgs)->UseRealTime(); + BENCHMARK(LUTGEMM_COMPUTE<2>)->Apply(LutGemmRealisticShapesArgs)->UseRealTime(); return true; } return false; diff --git a/onnxruntime/test/mlas/bench/bench_qnbitgemm.cpp b/onnxruntime/test/mlas/bench/bench_qnbitgemm.cpp index 0b143d41cb29f..7fb7198bd8729 100644 --- a/onnxruntime/test/mlas/bench/bench_qnbitgemm.cpp +++ b/onnxruntime/test/mlas/bench/bench_qnbitgemm.cpp @@ -156,9 +156,10 @@ BENCHMARK(QNBITGEMM)->Apply(QNBitGemmArgs)->UseRealTime(); BENCHMARK(QNBITGEMM)->Apply(QNBitGemmArgs)->UseRealTime(); BENCHMARK(QNBITGEMM)->Apply(QNBit2BitArgs)->UseRealTime(); -// Customer MatMulNBits shapes mirrored at 4-bit for a head-to-head comparison -// vs the W2 LUT path (LUTGEMM_COMPUTE/CUSTOMER). Customer model uses BlkLen=64. -// Five distinct (K, N) pairs: +// Representative MatMulNBits shapes mirrored at 4-bit for a head-to-head +// comparison vs the W2 LUT path (LUTGEMM_COMPUTE/RealisticShapes). These +// shapes use BlkLen=64 and reflect proportions that appear in real W2 +// production models. Five distinct (K, N) pairs: // (K=384, N=1024): 20 nodes // (K=1024, N=192): 40 nodes // (K=1024, N=384): 20 nodes @@ -166,7 +167,7 @@ BENCHMARK(QNBITGEMM)->Apply(QNBit2BitArgs)->UseRealTime(); // (K=4096, N=1024): 20 nodes // Both M=1 (decode) and M=128 (prefill) are exercised — paired with the W2 // rows below so we get a 3-way (W2 / W4 / W8) comparison at each M. -static void QNBitGemmCustomerArgs(benchmark::internal::Benchmark* b) { +static void QNBitGemmRealisticShapesArgs(benchmark::internal::Benchmark* b) { b->ArgNames({"BlkLen", "M", "N", "K", "Threads", "Symmetric", "HasBias", "ComputeType"}); const int64_t BlkLen = 64; const int64_t Threads = 8; @@ -185,13 +186,14 @@ static void QNBitGemmCustomerArgs(benchmark::internal::Benchmark* b) { } } -BENCHMARK(QNBITGEMM)->Apply(QNBitGemmCustomerArgs)->UseRealTime(); +BENCHMARK(QNBITGEMM)->Apply(QNBitGemmRealisticShapesArgs)->UseRealTime(); -// 2-bit weight rows for the customer shapes. Exercises the AVX-512 W2 native -// path (VNNI variant on AVX-512-VNNI hosts; non-VNNI variant on AVX-512BW -// hosts). W2 is registered only for SQNBIT_CompInt8 and BlkLen=64, so we -// emit just that one ComputeType. Covers both M=1 (decode) and M=128 (prefill). -static void QNBit2BitCustomerArgs(benchmark::internal::Benchmark* b) { +// 2-bit weight rows for the same representative shapes. Exercises the AVX-512 +// W2 native path (VNNI variant on AVX-512-VNNI hosts; non-VNNI variant on +// AVX-512BW hosts). W2 is registered only for SQNBIT_CompInt8 and BlkLen=64, +// so we emit just that one ComputeType. Covers both M=1 (decode) and M=128 +// (prefill). +static void QNBit2BitRealisticShapesArgs(benchmark::internal::Benchmark* b) { b->ArgNames({"BlkLen", "M", "N", "K", "Threads", "Symmetric", "HasBias", "ComputeType"}); const int64_t BlkLen = 64; const int64_t Threads = 8; @@ -208,7 +210,7 @@ static void QNBit2BitCustomerArgs(benchmark::internal::Benchmark* b) { } } -BENCHMARK(QNBITGEMM)->Apply(QNBit2BitCustomerArgs)->UseRealTime(); +BENCHMARK(QNBITGEMM)->Apply(QNBit2BitRealisticShapesArgs)->UseRealTime(); // This test gets benchmark arguments from environment variables. template diff --git a/onnxruntime/test/mlas/unittest/test_sqnbitgemm_2bit_gemm.cpp b/onnxruntime/test/mlas/unittest/test_sqnbitgemm_2bit_gemm.cpp index 20892d2985049..5826837f94b4f 100644 --- a/onnxruntime/test/mlas/unittest/test_sqnbitgemm_2bit_gemm.cpp +++ b/onnxruntime/test/mlas/unittest/test_sqnbitgemm_2bit_gemm.cpp @@ -304,9 +304,9 @@ void RunW2Case(size_t M, size_t N, size_t K, bool WithBias, uint32_t seed, // // Scalar block-group test, no zero-points. Covers the same small synthetic -// shapes + customer prefill sizes used by the production W2 tests. All shapes -// have K as a multiple of (kBlkLen * kBlockGroupBlks) = 256. Customer K=384 -// is NOT a multiple of 256 so it's excluded; that shape will need a tail +// shapes + representative prefill sizes used by the production W2 tests. All +// shapes have K as a multiple of (kBlkLen * kBlockGroupBlks) = 256. K=384 is +// NOT a multiple of 256 so it's excluded; that shape will need a tail // handler in a follow-up. // TEST(MlasSq2BitTest, Scalar_BlkLen64) { @@ -322,7 +322,7 @@ TEST(MlasSq2BitTest, Scalar_BlkLen64) { {7, 17, 256}, {16, 64, 512}, {32, 128, 256}, - // Customer prefill (only the K values that are multiples of 256). + // Representative prefill (only the K values that are multiples of 256). {1, 1024, 1024}, {1, 192, 1024}, {1, 384, 1024}, @@ -399,22 +399,23 @@ TEST(MlasSq2BitTest, Scalar_BlkLen64_WithZeroPoints) { // trailing slots so they contribute 0 to the dot product; the tile loads // only valid A blocks (zero ZMM for missing ones) to avoid OOB. // -// Customer prefill shapes (M in {1, 128}, N in {192, 384, 1024, 4096}, K in +// Representative prefill shapes (M in {1, 128}, N in {192, 384, 1024, 4096}, +// K in // {1024, 4096}) are all covered. M=3 and M=5 exercise the M-tail path. // // K-tail handler (BlockCountK not a multiple of kBlockGroupBlks=4): the // pack helper zero-pads the trailing 1-3 K-block slots; the SIMD K-loop // processes them via the 4-block accumulator with zero ZMM for the missing // A blocks. GroupStride uses BlockCountKPadded so N-group advances land on -// the right packed-B address regardless of K % 4. Customer K=384 and the +// the right packed-B address regardless of K % 4. K=384 and the // synthetic K=320, K=448 shapes exercise this path. // constexpr struct { size_t M, N, K; } kSimdShapes[] = { {1, 16, 256}, // R1 only - {1, 192, 1024}, // R1 only, customer N - {1, 1024, 4096}, // R1 only, customer N + {1, 192, 1024}, // R1 only, representative N + {1, 1024, 4096}, // R1 only, representative N {2, 16, 256}, {2, 32, 256}, {2, 64, 512}, @@ -425,14 +426,14 @@ constexpr struct { {5, 64, 512}, // R2 head (2 pairs) + R1 tail {16, 64, 512}, {32, 128, 256}, - // Customer prefill (M=128) at all (K, N) pairs, including K=384. + // Representative prefill (M=128) at all (K, N) pairs, including K=384. {128, 1024, 384}, // K-tail: BlockCountK=6, 1 full group + tail of 2 blocks {128, 1024, 1024}, {128, 192, 1024}, {128, 384, 1024}, {128, 4096, 1024}, {128, 1024, 4096}, - // Customer decode (M=1) at K=384 (the case the K%4 gate previously blocked). + // Representative decode (M=1) at K=384 (the case the K%4 gate previously blocked). {1, 1024, 384}, // Synthetic K-tail stress shapes covering all (TailBlocks in {1, 2, 3}). {2, 16, 320}, // tail=1 @@ -476,7 +477,7 @@ TEST(MlasSq2BitTest, BlkLen64_Avx512) { for (bool bias : {false, true}) { RunW2Case(s.M, s.N, s.K, bias, seed + (bias ? 1u : 0u), /*WithZeroPoints=*/false, - sq2::SQ2BitGemmKernel_BlkSum_CompInt8_Avx512_TestEntry, + sq2::SQ2BitGemmKernel_BlkSum_CompInt8_Avx512_Dispatch, "AVX-512BW"); } } @@ -492,7 +493,7 @@ TEST(MlasSq2BitTest, BlkLen64_Avx512_WithZeroPoints) { for (bool bias : {false, true}) { RunW2Case(s.M, s.N, s.K, bias, seed + (bias ? 1u : 0u), /*WithZeroPoints=*/true, - sq2::SQ2BitGemmKernel_BlkSum_CompInt8_Avx512_TestEntry, + sq2::SQ2BitGemmKernel_BlkSum_CompInt8_Avx512_Dispatch, "AVX-512BW"); } } @@ -512,7 +513,7 @@ TEST(MlasSq2BitTest, BlkLen64_Avx512Vnni) { for (bool bias : {false, true}) { RunW2Case(s.M, s.N, s.K, bias, seed + (bias ? 1u : 0u), /*WithZeroPoints=*/false, - sq2::SQ2BitGemmKernel_BlkSum_CompInt8_Avx512Vnni_TestEntry, + sq2::SQ2BitGemmKernel_BlkSum_CompInt8_Avx512Vnni_Dispatch, "AVX-512-VNNI"); } } @@ -528,7 +529,7 @@ TEST(MlasSq2BitTest, BlkLen64_Avx512Vnni_WithZeroPoints) { for (bool bias : {false, true}) { RunW2Case(s.M, s.N, s.K, bias, seed + (bias ? 1u : 0u), /*WithZeroPoints=*/true, - sq2::SQ2BitGemmKernel_BlkSum_CompInt8_Avx512Vnni_TestEntry, + sq2::SQ2BitGemmKernel_BlkSum_CompInt8_Avx512Vnni_Dispatch, "AVX-512-VNNI"); } } @@ -760,7 +761,7 @@ constexpr struct { {1, 16, 128}, // R1, BlockCountK=1 {1, 32, 256}, // R1, BlockCountK=2 (K-tail, no full group) {1, 1024, 1024}, // R1, BlockCountK=8 - {1, 1024, 4096}, // R1, BlockCountK=32 (customer N) + {1, 1024, 4096}, // R1, BlockCountK=32 (representative N) {2, 16, 128}, {2, 32, 256}, {2, 64, 512}, // BlockCountK=4 (one full group) @@ -770,7 +771,7 @@ constexpr struct { {4, 32, 256}, {16, 64, 512}, {32, 128, 256}, - // M=128 prefill at customer-like shapes (K multiples of 128) + // M=128 prefill at representative shapes (K multiples of 128) {128, 1024, 1024}, {128, 1024, 4096}, {128, 192, 1024}, @@ -778,7 +779,7 @@ constexpr struct { // K-tail stress (BlockCountK not a multiple of 4) {2, 16, 384}, // tail=3 {4, 16, 384}, - {128, 1024, 384}, // tail=3 at customer M + {128, 1024, 384}, // tail=3 at representative M {2, 16, 640}, // tail=1 {4, 16, 640}, {2, 16, 768}, // tail=2 @@ -839,7 +840,7 @@ TEST(MlasSq2BitTest, BlkLen128_Avx512) { for (bool bias : {false, true}) { RunW2Case_BlkLen128(s.M, s.N, s.K, bias, seed + (bias ? 1u : 0u), /*WithZeroPoints=*/false, - sq2::SQ2BitGemmKernel_BlkSum_CompInt8_Avx512_TestEntry, + sq2::SQ2BitGemmKernel_BlkSum_CompInt8_Avx512_Dispatch, "AVX-512BW"); } } @@ -855,7 +856,7 @@ TEST(MlasSq2BitTest, BlkLen128_Avx512_WithZeroPoints) { for (bool bias : {false, true}) { RunW2Case_BlkLen128(s.M, s.N, s.K, bias, seed + (bias ? 1u : 0u), /*WithZeroPoints=*/true, - sq2::SQ2BitGemmKernel_BlkSum_CompInt8_Avx512_TestEntry, + sq2::SQ2BitGemmKernel_BlkSum_CompInt8_Avx512_Dispatch, "AVX-512BW"); } } @@ -871,7 +872,7 @@ TEST(MlasSq2BitTest, BlkLen128_Avx512Vnni) { for (bool bias : {false, true}) { RunW2Case_BlkLen128(s.M, s.N, s.K, bias, seed + (bias ? 1u : 0u), /*WithZeroPoints=*/false, - sq2::SQ2BitGemmKernel_BlkSum_CompInt8_Avx512Vnni_TestEntry, + sq2::SQ2BitGemmKernel_BlkSum_CompInt8_Avx512Vnni_Dispatch, "AVX-512-VNNI"); } } @@ -887,7 +888,7 @@ TEST(MlasSq2BitTest, BlkLen128_Avx512Vnni_WithZeroPoints) { for (bool bias : {false, true}) { RunW2Case_BlkLen128(s.M, s.N, s.K, bias, seed + (bias ? 1u : 0u), /*WithZeroPoints=*/true, - sq2::SQ2BitGemmKernel_BlkSum_CompInt8_Avx512Vnni_TestEntry, + sq2::SQ2BitGemmKernel_BlkSum_CompInt8_Avx512Vnni_Dispatch, "AVX-512-VNNI"); } } @@ -1193,7 +1194,7 @@ TEST(MlasSq2BitTest, BlkLen32_Avx512) { for (bool bias : {false, true}) { RunW2Case_BlkLen32(s.M, s.N, s.K, bias, seed + (bias ? 1u : 0u), /*WithZeroPoints=*/false, - sq2::SQ2BitGemmKernel_BlkSum_CompInt8_Avx512_TestEntry, + sq2::SQ2BitGemmKernel_BlkSum_CompInt8_Avx512_Dispatch, "AVX-512BW"); } } @@ -1209,7 +1210,7 @@ TEST(MlasSq2BitTest, BlkLen32_Avx512_WithZeroPoints) { for (bool bias : {false, true}) { RunW2Case_BlkLen32(s.M, s.N, s.K, bias, seed + (bias ? 1u : 0u), /*WithZeroPoints=*/true, - sq2::SQ2BitGemmKernel_BlkSum_CompInt8_Avx512_TestEntry, + sq2::SQ2BitGemmKernel_BlkSum_CompInt8_Avx512_Dispatch, "AVX-512BW"); } } @@ -1225,7 +1226,7 @@ TEST(MlasSq2BitTest, BlkLen32_Avx512Vnni) { for (bool bias : {false, true}) { RunW2Case_BlkLen32(s.M, s.N, s.K, bias, seed + (bias ? 1u : 0u), /*WithZeroPoints=*/false, - sq2::SQ2BitGemmKernel_BlkSum_CompInt8_Avx512Vnni_TestEntry, + sq2::SQ2BitGemmKernel_BlkSum_CompInt8_Avx512Vnni_Dispatch, "AVX-512-VNNI"); } } @@ -1241,7 +1242,7 @@ TEST(MlasSq2BitTest, BlkLen32_Avx512Vnni_WithZeroPoints) { for (bool bias : {false, true}) { RunW2Case_BlkLen32(s.M, s.N, s.K, bias, seed + (bias ? 1u : 0u), /*WithZeroPoints=*/true, - sq2::SQ2BitGemmKernel_BlkSum_CompInt8_Avx512Vnni_TestEntry, + sq2::SQ2BitGemmKernel_BlkSum_CompInt8_Avx512Vnni_Dispatch, "AVX-512-VNNI"); } } From 2701ba268d55512f1a3cf586caaa612d1aceb352 Mon Sep 17 00:00:00 2001 From: Hari Seshadri Date: Tue, 16 Jun 2026 14:53:34 -0700 Subject: [PATCH 13/17] Copilot comments --- onnxruntime/core/mlas/lib/qnbitgemm.cpp | 13 ++++++-- .../mlas/lib/sqnbitgemm_kernel_avx512_2bit.h | 30 +++++++++++-------- .../sqnbitgemm_kernel_avx512_2bit_blklen64.h | 15 ++++++---- .../test/mlas/bench/bench_qnbitgemm.cpp | 6 ++-- 4 files changed, 42 insertions(+), 22 deletions(-) diff --git a/onnxruntime/core/mlas/lib/qnbitgemm.cpp b/onnxruntime/core/mlas/lib/qnbitgemm.cpp index 92949f97ab343..364220e04ea61 100644 --- a/onnxruntime/core/mlas/lib/qnbitgemm.cpp +++ b/onnxruntime/core/mlas/lib/qnbitgemm.cpp @@ -973,8 +973,11 @@ SQ8BitGemm_CompInt8( // // 2-bit weight CompInt8 wrapper. Mirrors SQ4BitGemm_CompInt8 but specialised -// for BlkBitWidth=2 (kBlkBytes = BlkLen/4) and uses the simple column-major -// QuantBBlkSum layout produced by SQ2BitGemmPackQuantBDataAndBlkSum_*. +// for BlkBitWidth=2 (kBlkBytes = BlkLen/4). QuantBBlkSum uses the existing +// SGEMM-style width-16 chunked layout produced by the PackQuantBDataAndBlkSum +// helpers (same as the W4 CompInt8 path); see the comment on +// Q2BitGemmEffectiveBlockCountK in qnbitgemm.h for how the W2 packed-B +// stride differs from the BlkSum stride. // void SQ2BitGemm_CompInt8( @@ -1020,6 +1023,12 @@ SQ2BitGemm_CompInt8( const std::byte* QuantA = per_gemm_quant_a_workspace->QuantData + RangeStartM * lda; const float* QuantAScale = per_gemm_quant_a_workspace->QuantScale + RangeStartM * k_blks; + // The packed-B and BlkSum layouts both group N-cols in multiples of 4 + // (kNCols4 in the W2 kernel; same convention as SQ4BitGemm_CompInt8). + // The work partitioner produces aligned RangeStartN values, so this + // assert is invariant; it documents the contract and would catch a + // future partitioner regression. + assert(RangeStartN % 4 == 0); const std::byte* QuantBData = static_cast(DataParams->PackedQuantBData) + RangeStartN * ldb; const float* QuantBScale = DataParams->QuantBScale + RangeStartN * k_blks_eff; const float* ABlockSum = per_gemm_quant_a_workspace->BlockSum + RangeStartM * k_blks; diff --git a/onnxruntime/core/mlas/lib/sqnbitgemm_kernel_avx512_2bit.h b/onnxruntime/core/mlas/lib/sqnbitgemm_kernel_avx512_2bit.h index 67c8fdc2734c9..16e23b3993d57 100644 --- a/onnxruntime/core/mlas/lib/sqnbitgemm_kernel_avx512_2bit.h +++ b/onnxruntime/core/mlas/lib/sqnbitgemm_kernel_avx512_2bit.h @@ -15,11 +15,14 @@ Module Name: four fixed `vpsrlw+vpand` pairs (instead of the current per-block broadcast + variable shift). - Layout summary (BlkLen=64 only): - - * Each "block-group" packs FOUR consecutive K-blocks (256 weights total). - * Total storage per block-group = kBlkBytes * 4 = 64 bytes (identical to - 4 separately-packed blocks under the current scheme). + Layout summary (BlkLen ∈ {32, 64, 128}): + + * Each "block-group" packs FOUR consecutive K-blocks. + * Total storage per block-group = kBlkBytes * 4 bytes (kBlkBytes is + BlkLen-dependent: 8 bytes at BlkLen=32, 16 bytes at BlkLen=64, 32 + bytes at BlkLen=128). Identical to 4 separately-packed blocks under + the per-block scheme, just relaid out so all four blocks are reachable + with one ZMM load. * Byte b of the block-group holds: bits[0..1] = block_0.weight[b] bits[2..3] = block_1.weight[b] @@ -31,13 +34,16 @@ Module Name: Restrictions: - * BlkLen == 64 only. - * BlockCountK must be a multiple of kBlockGroupBlks = 4. The - path returns 0 from the pack-size helper for non-multiples; the - caller falls back to the existing W2 path. - * Typical W2 production K dimensions (e.g. 384, 1024, 4096) are all - multiples of 256 (= kBlockGroupBlks * 64) and therefore satisfy this - constraint. + * BlkLen ∈ {32, 64, 128}. + * No constraint on BlockCountK. The pack-size helper rounds BlockCountK + up to a multiple of kBlockGroupBlks = 4 internally, and the SIMD + K-loop processes the padded blocks via the 4-block accumulator with + zero ZMM for the missing A blocks (K-tail handler). + * Typical W2 production K dimensions (e.g. 384, 1024, 4096) at + BlkLen=64 are all multiples of 256 (= kBlockGroupBlks * 64) and so + do not exercise the K-tail handler; smaller K (or BlkLen=32 / 128 + shapes whose K is not a multiple of kBlockGroupBlks * BlkLen) take + the K-tail path and are exercised by the unit tests. --*/ diff --git a/onnxruntime/core/mlas/lib/sqnbitgemm_kernel_avx512_2bit_blklen64.h b/onnxruntime/core/mlas/lib/sqnbitgemm_kernel_avx512_2bit_blklen64.h index 53ac8b3d18f5b..6ec0c38660e4a 100644 --- a/onnxruntime/core/mlas/lib/sqnbitgemm_kernel_avx512_2bit_blklen64.h +++ b/onnxruntime/core/mlas/lib/sqnbitgemm_kernel_avx512_2bit_blklen64.h @@ -22,11 +22,16 @@ Module Name: Constraints (SIMD): * BlkLen == 64 only. - * BlockCountK must be a multiple of kBlockGroupBlks (= 4). - * CountM must be a multiple of 2 (R2 tile); CountN must be a multiple - of 4 (C4 tile). Tail handling is provided by the same R2xC1/R1xC4/ - R1xC1 helpers as the existing W2 kernel via a downgrade path the - dispatcher will pick when these constraints don't hold. + * BlockCountK has no alignment requirement. The pack helper rounds + BlockCountK up to a multiple of kBlockGroupBlks (= 4); the SIMD + K-loop processes the padded blocks via the 4-block accumulator with + zero ZMM for the missing A blocks (K-tail handler). + * CountM has no alignment requirement (R2xC4 head + optional R1xC4 + tail picks up the trailing odd row). + * CountN has no alignment requirement (R2/R1 xC4 main covers + NMain = floor(CountN/4)*4; a per-1-col tail tile picks up the + trailing 1-3 N-cols, including the NMain=0 case where N in + {1, 2, 3}). Layout reference: * 64-byte block-group: byte b holds 2-bit weight b from each of 4 diff --git a/onnxruntime/test/mlas/bench/bench_qnbitgemm.cpp b/onnxruntime/test/mlas/bench/bench_qnbitgemm.cpp index 7fb7198bd8729..2432f35128d5f 100644 --- a/onnxruntime/test/mlas/bench/bench_qnbitgemm.cpp +++ b/onnxruntime/test/mlas/bench/bench_qnbitgemm.cpp @@ -134,9 +134,9 @@ static void QNBitGemmArgs(benchmark::internal::Benchmark* b) { }); } -// Standard sweep for the native W2 kernel. W2 has fewer free dimensions than -// W4 (symmetric-only, BlkLen=64 only, SQNBIT_CompInt8 only), so the grid -// uses fixed values for those axes and sweeps the rest like QNBitGemmArgs. +// Standard sweep for the native W2 kernel. W2 has fewer free dimensions +// than W4 (symmetric-only, SQNBIT_CompInt8 only), so the grid uses fixed +// values for those axes and sweeps the rest like QNBitGemmArgs. static void QNBit2BitArgs(benchmark::internal::Benchmark* b) { b->ArgNames({"BlkLen", "M", "N", "K", "Threads", "Symmetric", "HasBias", "ComputeType"}); From 61c844ac1cf4f19bf96162fa81e722e74bf52381 Mon Sep 17 00:00:00 2001 From: Hari Seshadri Date: Tue, 16 Jun 2026 15:17:59 -0700 Subject: [PATCH 14/17] Copilot comments - 2 --- .../lib/sqnbitgemm_kernel_avx512_2bit.cpp | 27 +++++++++++-------- .../sqnbitgemm_kernel_avx512_2bit_blklen64.h | 5 ++-- onnxruntime/test/mlas/bench/bench_lutgemm.cpp | 14 +++++++--- 3 files changed, 30 insertions(+), 16 deletions(-) diff --git a/onnxruntime/core/mlas/lib/sqnbitgemm_kernel_avx512_2bit.cpp b/onnxruntime/core/mlas/lib/sqnbitgemm_kernel_avx512_2bit.cpp index cac3149cf294c..31ae663b379f2 100644 --- a/onnxruntime/core/mlas/lib/sqnbitgemm_kernel_avx512_2bit.cpp +++ b/onnxruntime/core/mlas/lib/sqnbitgemm_kernel_avx512_2bit.cpp @@ -17,10 +17,10 @@ Module Name: per-block broadcast + variable shift unpack with a single 64-byte load + four fixed-shift+mask pairs). - This translation unit is scalar / portable. The vectorized inner loop - that consumes the block-group layout lives in a separate header - (sqnbitgemm_kernel_avx512_2bit_blklen64.h) and is wired into the AVX-512 - and AVX-512-VNNI dispatch tables. + This translation unit is scalar / portable. The vectorized inner loops + that consume the block-group layout live in separate headers + (sqnbitgemm_kernel_avx512_2bit_blklen{32,64,128}.h) and are wired into the + AVX-512 and AVX-512-VNNI dispatch tables. --*/ @@ -40,19 +40,23 @@ namespace sq2bit_avx512 { // // Workspace / pack-buffer size for the block-group W2 path. Returns 0 if any -// of the configuration constraints is violated; the caller (MlasQNBitGemmPackQuantBDataSize) -// treats that as "unsupported" and falls back to the original W2 path. +// of the configuration constraints is violated; the caller +// (MlasQNBitGemmPackQuantBDataSize) treats that as "unsupported" and falls +// back to LUT (if opted in and the shape is LUT-eligible) or to the fp32 +// dequant + SGEMM path in MatMulNBits::ComputeBUnpacked. // // Constraints: -// * BlkLen == 64 +// * BlkLen ∈ {32, 64, 128} // * ComputeType == SQNBIT_CompInt8 // // K-tail handling: BlockCountK is rounded UP to a multiple of kBlockGroupBlks // for the storage that the inner K-loop walks (PackedQuantBData, // PackedQuantBScale). Padding slots hold zeroed weights and scales, so they -// contribute exactly 0 to the dot product. The BlkSum buffer is kept at the -// LOGICAL BlockCountK because it is consumed by the SGEMM correction step, -// not by the inner K-loop. +// contribute exactly 0 to the dot product. The BlkSum buffer is *physically* +// sized with the padded BlockCountK so its offset within the combined +// workspace is consistent with the caller's PackedQuantBDataStruct layout +// (which uses one BlockCountK value for the whole struct); the SGEMM +// correction step still reads only the logical BlockCountK entries from it. // // Storage matches the original W2 layout total bytes when BlockCountK is a // multiple of 4. When not a multiple of 4, storage grows by at most 3 K-blocks @@ -60,7 +64,8 @@ namespace sq2bit_avx512 { // // [PackedQuantBData] N * BlockCountKPadded * kBlkBytes // [PackedQuantBScale] N * BlockCountKPadded * sizeof(float) -// [QuantBBlkSum] roundup_16(N) * BlockCountK (logical) * 16 floats +// [QuantBBlkSum] roundup_16(N) * BlockCountKPadded * 16 floats (only +// the first BlockCountK entries per N are populated) // size_t MLASCALL Q2BitGemmPackQuantBDataSize_Avx512( diff --git a/onnxruntime/core/mlas/lib/sqnbitgemm_kernel_avx512_2bit_blklen64.h b/onnxruntime/core/mlas/lib/sqnbitgemm_kernel_avx512_2bit_blklen64.h index 6ec0c38660e4a..f3a46dcd6200b 100644 --- a/onnxruntime/core/mlas/lib/sqnbitgemm_kernel_avx512_2bit_blklen64.h +++ b/onnxruntime/core/mlas/lib/sqnbitgemm_kernel_avx512_2bit_blklen64.h @@ -382,8 +382,9 @@ Q2Int8GemmR1xC4BlkLen64Avx512( // // R2 x C4 tile -- the main hot path for prefill (M >= 2). Iterates the K // dimension in block-group strides of kBlockGroupBlks (= 4) K-blocks at a -// time. Assumes BlockCountK is a multiple of kBlockGroupBlks; the dispatcher -// must verify this before selecting the block-group kernel. +// time, and handles K-tails (BlockCountK % kBlockGroupBlks != 0) via the +// shared 4-block accumulator with zero-padded A/scale slots for the +// missing trailing K-blocks. // template MLAS_FORCEINLINE void diff --git a/onnxruntime/test/mlas/bench/bench_lutgemm.cpp b/onnxruntime/test/mlas/bench/bench_lutgemm.cpp index dec40a9124772..b710cb18d85bc 100644 --- a/onnxruntime/test/mlas/bench/bench_lutgemm.cpp +++ b/onnxruntime/test/mlas/bench/bench_lutgemm.cpp @@ -169,7 +169,11 @@ void LUTGEMM_COMPUTE(benchmark::State& state) { // when LUT is gated out for the given shape (e.g. N % n_div != 0). state.SetLabel("path=Dequant+SGEMM"); - std::vector DequantB(K * N); // [K, N] row-major + // MlasDequantizeBlockwise(columnwise=true, K, N) emits B in [N, K] row-major + // layout (equivalently [K, N] column-major) -- this is "B^T" from the + // MatMul perspective and matches what matmul_nbits.cc::ComputeBUnpacked + // feeds to MlasGemmBatch. + std::vector DequantB(K * N); // Time dequant+SGEMM as a unit (this is what the runtime pays per call, // since the dequantized buffer is not cached across MatMulNBits calls). @@ -180,11 +184,15 @@ void LUTGEMM_COMPUTE(benchmark::State& state) { static_cast(BlkLen), /*columnwise*/ true, static_cast(K), static_cast(N), tp.get()); - MlasGemm(CblasNoTrans, CblasNoTrans, + // DequantB is [N, K] row-major; use CblasTrans + ldb=K to match + // matmul_nbits.cc::ComputeBUnpacked (which calls + // MlasGemmBatch(CblasNoTrans, CblasTrans, ..., ldb=K, ...) on the + // same dequantized buffer). + MlasGemm(CblasNoTrans, CblasTrans, M, N, K, 1.0f, A.data(), K, - DequantB.data(), N, + DequantB.data(), K, 0.0f, C.data(), N, tp.get(), nullptr); From 847222f1eb869e85293389889bf9c4fc9effdbd0 Mon Sep 17 00:00:00 2001 From: Hari Seshadri Date: Tue, 16 Jun 2026 15:38:24 -0700 Subject: [PATCH 15/17] Copilot comments - 3 --- onnxruntime/core/mlas/lib/qnbitgemm.cpp | 14 ++++++++------ 1 file changed, 8 insertions(+), 6 deletions(-) diff --git a/onnxruntime/core/mlas/lib/qnbitgemm.cpp b/onnxruntime/core/mlas/lib/qnbitgemm.cpp index 364220e04ea61..5710df2a06aae 100644 --- a/onnxruntime/core/mlas/lib/qnbitgemm.cpp +++ b/onnxruntime/core/mlas/lib/qnbitgemm.cpp @@ -1023,12 +1023,14 @@ SQ2BitGemm_CompInt8( const std::byte* QuantA = per_gemm_quant_a_workspace->QuantData + RangeStartM * lda; const float* QuantAScale = per_gemm_quant_a_workspace->QuantScale + RangeStartM * k_blks; - // The packed-B and BlkSum layouts both group N-cols in multiples of 4 - // (kNCols4 in the W2 kernel; same convention as SQ4BitGemm_CompInt8). - // The work partitioner produces aligned RangeStartN values, so this - // assert is invariant; it documents the contract and would catch a - // future partitioner regression. - assert(RangeStartN % 4 == 0); + // QuantBBlkSum uses the width-16 chunked layout produced by + // ComputePackBlkSum (BlockSum[n] lives at chunk (n/16) at intra-chunk + // offset n%16), so the flat `QuantBBlkSum + RangeStartN * k_blks` + // addressing below is only correct when RangeStartN is a multiple of 16. + // The work partitioner enforces this via MLAS_QGEMM_STRIDEN_THREAD_ALIGN + // (= 16); the assert documents the contract and would catch a future + // partitioner regression. + assert(RangeStartN % 16 == 0); const std::byte* QuantBData = static_cast(DataParams->PackedQuantBData) + RangeStartN * ldb; const float* QuantBScale = DataParams->QuantBScale + RangeStartN * k_blks_eff; const float* ABlockSum = per_gemm_quant_a_workspace->BlockSum + RangeStartM * k_blks; From 4d442ce0b1606716ec72ee0612a6fac4bbce3e0d Mon Sep 17 00:00:00 2001 From: Hari Seshadri Date: Tue, 16 Jun 2026 16:35:11 -0700 Subject: [PATCH 16/17] Fix availability contract --- onnxruntime/core/mlas/lib/qnbitgemm.cpp | 18 ++++++++---- .../unittest/test_sqnbitgemm_2bit_gemm.cpp | 29 +++++++++++++++++++ 2 files changed, 42 insertions(+), 5 deletions(-) diff --git a/onnxruntime/core/mlas/lib/qnbitgemm.cpp b/onnxruntime/core/mlas/lib/qnbitgemm.cpp index 5710df2a06aae..1337b798d3a8c 100644 --- a/onnxruntime/core/mlas/lib/qnbitgemm.cpp +++ b/onnxruntime/core/mlas/lib/qnbitgemm.cpp @@ -49,7 +49,13 @@ GetQNBitGemmVariant( MLAS_QNBIT_GEMM_COMPUTE_TYPE ComputeType ) { - if ((BlkLen == 16 || BlkLen == 32 || BlkLen == 64 || BlkLen == 128 || BlkLen == 256)) { + // BlkLen support varies by BlkBitWidth: + // W4 / W8 : {16, 32, 64, 128, 256} + // W2 : {32, 64, 128} (the native non-LUT kernel does not implement 16 or 256; + // the pack-size helper returns 0 for those, and we must report + // this truthfully via MlasIsQNBitGemmAvailable for direct MLAS + // callers who rely on availability as the support contract.) + if (BlkLen == 16 || BlkLen == 32 || BlkLen == 64 || BlkLen == 128 || BlkLen == 256) { if (BlkBitWidth == 4) { if (ComputeType == SQNBIT_CompFp32) { return SQ4BitGemmVariant_CompFp32; @@ -64,10 +70,12 @@ GetQNBitGemmVariant( } else if (ComputeType == HQNBIT_CompFp16) { return HQ8BitGemmVariant_CompFp16; } - } else if (BlkBitWidth == 2) { - if (ComputeType == SQNBIT_CompInt8) { - return SQ2BitGemmVariant_CompInt8; - } + } + } + + if (BlkBitWidth == 2 && (BlkLen == 32 || BlkLen == 64 || BlkLen == 128)) { + if (ComputeType == SQNBIT_CompInt8) { + return SQ2BitGemmVariant_CompInt8; } } diff --git a/onnxruntime/test/mlas/unittest/test_sqnbitgemm_2bit_gemm.cpp b/onnxruntime/test/mlas/unittest/test_sqnbitgemm_2bit_gemm.cpp index 5826837f94b4f..7690b32c3c6a8 100644 --- a/onnxruntime/test/mlas/unittest/test_sqnbitgemm_2bit_gemm.cpp +++ b/onnxruntime/test/mlas/unittest/test_sqnbitgemm_2bit_gemm.cpp @@ -1249,4 +1249,33 @@ TEST(MlasSq2BitTest, BlkLen32_Avx512Vnni_WithZeroPoints) { } } +// +// Availability contract test for W2 + SQNBIT_CompInt8. +// +// The W2 native kernel only implements BlkLen ∈ {32, 64, 128}. +// MlasIsQNBitGemmAvailable must report this truthfully so direct MLAS callers +// can rely on it as the support contract. Previously the variant gate also +// admitted BlkLen 16 and 256 (since they're valid for W4 / W8), which made +// availability return true for shapes that Q2BitGemmPackQuantBDataSize_Avx512 +// would refuse to size (returning 0). +// +TEST(MlasSq2BitTest, AvailabilityContract_BlkLens) { + if (!GetMlasPlatform().Avx512Supported_) { + GTEST_SKIP() << "W2 native dispatch is AVX-512-only on x86_64"; + } + + // Supported BlkLens. + EXPECT_TRUE(MlasIsQNBitGemmAvailable(2, 32, SQNBIT_CompInt8)); + EXPECT_TRUE(MlasIsQNBitGemmAvailable(2, 64, SQNBIT_CompInt8)); + EXPECT_TRUE(MlasIsQNBitGemmAvailable(2, 128, SQNBIT_CompInt8)); + + // Unsupported BlkLens for W2 (valid for W4/W8 but not implemented for W2). + EXPECT_FALSE(MlasIsQNBitGemmAvailable(2, 16, SQNBIT_CompInt8)); + EXPECT_FALSE(MlasIsQNBitGemmAvailable(2, 256, SQNBIT_CompInt8)); + + // Compute types not implemented for W2. + EXPECT_FALSE(MlasIsQNBitGemmAvailable(2, 64, SQNBIT_CompFp32)); + EXPECT_FALSE(MlasIsQNBitGemmAvailable(2, 64, HQNBIT_CompFp16)); +} + #endif // defined(MLAS_TARGET_AMD64) From b40dfb20823fecb7bea4a5606b9af38734ff46f8 Mon Sep 17 00:00:00 2001 From: Hari Seshadri Date: Tue, 16 Jun 2026 21:52:53 -0700 Subject: [PATCH 17/17] Tianlei PR comments --- .../sqnbitgemm_kernel_avx512_2bit_blklen128.h | 14 ++++++++++++-- .../lib/sqnbitgemm_kernel_avx512_2bit_blklen32.h | 14 ++++++++++++-- .../lib/sqnbitgemm_kernel_avx512_2bit_blklen64.h | 16 +++++++++++++--- 3 files changed, 37 insertions(+), 7 deletions(-) diff --git a/onnxruntime/core/mlas/lib/sqnbitgemm_kernel_avx512_2bit_blklen128.h b/onnxruntime/core/mlas/lib/sqnbitgemm_kernel_avx512_2bit_blklen128.h index 1065c32634c62..febf7f75b18c9 100644 --- a/onnxruntime/core/mlas/lib/sqnbitgemm_kernel_avx512_2bit_blklen128.h +++ b/onnxruntime/core/mlas/lib/sqnbitgemm_kernel_avx512_2bit_blklen128.h @@ -127,6 +127,7 @@ dot_accumulate_4blk_w2_blklen128( { __m512i d0, d1, d2, d3; if constexpr (kVnni) { + // dpbusd: 2nd operand (bv=unsigned [0,3]) x 3rd operand (av=signed int8); consistent with maddubs(bv, av) below d0 = _mm512_dpbusd_epi32(_mm512_setzero_si512(), bv0_lo, av0_lo); d0 = _mm512_dpbusd_epi32(d0, bv0_hi, av0_hi); d1 = _mm512_dpbusd_epi32(_mm512_setzero_si512(), bv1_lo, av1_lo); @@ -162,8 +163,8 @@ dot_accumulate_4blk_w2_blklen128( const __m512 s2 = _mm512_set1_ps(scale_a[2] * scale_b[2]); const __m512 s3 = _mm512_set1_ps(scale_a[3] * scale_b[3]); - __m512 acc_lo = _mm512_fmadd_ps(_mm512_cvtepi32_ps(d0), s0, _mm512_setzero_ps()); - __m512 acc_hi = _mm512_fmadd_ps(_mm512_cvtepi32_ps(d1), s1, _mm512_setzero_ps()); + __m512 acc_lo = _mm512_mul_ps(_mm512_cvtepi32_ps(d0), s0); + __m512 acc_hi = _mm512_mul_ps(_mm512_cvtepi32_ps(d1), s1); acc_lo = _mm512_fmadd_ps(_mm512_cvtepi32_ps(d2), s2, acc_lo); acc_hi = _mm512_fmadd_ps(_mm512_cvtepi32_ps(d3), s3, acc_hi); acc = _mm512_add_ps(acc, _mm512_add_ps(acc_lo, acc_hi)); @@ -347,6 +348,9 @@ Q2Int8GemmR1xC4BlkLen128Avx512( const __m512i av01_hi = load_hi(1); const __m512i av02_lo = load_lo(2); const __m512i av02_hi = load_hi(2); + // TailBlocks = BlockCountK % kBlockGroupBlks in [1,3], so block 3 is never a real tail + // block. Keep av03 hardcoded to zero. The unpacked bv3 can still be non-zero, + // but its contribution is zeroed twice: by av03==0 and scale_a0_safe[3]==0. const __m512i av03_lo = zero; const __m512i av03_hi = zero; @@ -538,6 +542,9 @@ Q2Int8GemmR2xC4BlkLen128Avx512( const __m512i av01_hi = load_a_hi(0, 1); const __m512i av02_lo = load_a_lo(0, 2); const __m512i av02_hi = load_a_hi(0, 2); + // TailBlocks = BlockCountK % kBlockGroupBlks in [1,3], so block 3 is never a real tail + // block. Keep av03 hardcoded to zero. The unpacked bv3 can still be non-zero, + // but its contribution is zeroed twice: by av03==0 and scale_a0_safe[3]==0. const __m512i av03_lo = zero; const __m512i av03_hi = zero; @@ -706,6 +713,9 @@ Q2Int8GemmRMxC_Tail_BlkLen128Avx512( const __m512i av01_hi = load_hi(1); const __m512i av02_lo = load_lo(2); const __m512i av02_hi = load_hi(2); + // TailBlocks = BlockCountK % kBlockGroupBlks in [1,3], so block 3 is never a real tail + // block. Keep av03 hardcoded to zero. The unpacked bv3 can still be non-zero, + // but its contribution is zeroed twice: by av03==0 and scale_a0_safe[3]==0. const __m512i av03_lo = zero; const __m512i av03_hi = zero; diff --git a/onnxruntime/core/mlas/lib/sqnbitgemm_kernel_avx512_2bit_blklen32.h b/onnxruntime/core/mlas/lib/sqnbitgemm_kernel_avx512_2bit_blklen32.h index 505c3e216184e..d7869c6406df3 100644 --- a/onnxruntime/core/mlas/lib/sqnbitgemm_kernel_avx512_2bit_blklen32.h +++ b/onnxruntime/core/mlas/lib/sqnbitgemm_kernel_avx512_2bit_blklen32.h @@ -119,6 +119,7 @@ dot_accumulate_4blk_w2_blklen32( { __m512i d0, d1, d2, d3; if constexpr (kVnni) { + // dpbusd: 2nd operand (bv=unsigned [0,3]) × 3rd operand (av=signed int8); consistent with maddubs(bv, av) below d0 = _mm512_dpbusd_epi32(_mm512_setzero_si512(), bv0, av0); d1 = _mm512_dpbusd_epi32(_mm512_setzero_si512(), bv1, av1); d2 = _mm512_dpbusd_epi32(_mm512_setzero_si512(), bv2, av2); @@ -140,8 +141,8 @@ dot_accumulate_4blk_w2_blklen32( const __m512 s2 = _mm512_set1_ps(scale_a[2] * scale_b[2]); const __m512 s3 = _mm512_set1_ps(scale_a[3] * scale_b[3]); - __m512 acc_lo = _mm512_fmadd_ps(_mm512_cvtepi32_ps(d0), s0, _mm512_setzero_ps()); - __m512 acc_hi = _mm512_fmadd_ps(_mm512_cvtepi32_ps(d1), s1, _mm512_setzero_ps()); + __m512 acc_lo = _mm512_mul_ps(_mm512_cvtepi32_ps(d0), s0); + __m512 acc_hi = _mm512_mul_ps(_mm512_cvtepi32_ps(d1), s1); acc_lo = _mm512_fmadd_ps(_mm512_cvtepi32_ps(d2), s2, acc_lo); acc_hi = _mm512_fmadd_ps(_mm512_cvtepi32_ps(d3), s3, acc_hi); acc = _mm512_add_ps(acc, _mm512_add_ps(acc_lo, acc_hi)); @@ -273,6 +274,9 @@ Q2Int8GemmR1xC4BlkLen32Avx512( const __m512i av00 = load_a(0); const __m512i av01 = load_a(1); const __m512i av02 = load_a(2); + // TailBlocks = BlockCountK % kBlockGroupBlks ∈ [1,3], so block 3 is never a real tail + // block — hardcode av03 = zero. The unpacked bv3 is still non-zero, but its + // contribution is zeroed twice: by av03==0 here and scale_a0_safe[3]==0 below. const __m512i av03 = zero; float scale_a0_safe[kBlockGroupBlks] = {0.0f, 0.0f, 0.0f, 0.0f}; @@ -419,6 +423,9 @@ Q2Int8GemmR2xC4BlkLen32Avx512( const __m512i av00 = load_a(0, 0); const __m512i av01 = load_a(0, 1); const __m512i av02 = load_a(0, 2); + // TailBlocks = BlockCountK % kBlockGroupBlks ∈ [1,3], so block 3 is never a real tail + // block — hardcode av03 = zero. The unpacked bv3 is still non-zero, but its + // contribution is zeroed twice: by av03==0 here and scale_a0_safe[3]==0 below. const __m512i av03 = zero; const __m512i av10 = load_a(lda, 0); const __m512i av11 = load_a(lda, 1); @@ -550,6 +557,9 @@ Q2Int8GemmRMxC_Tail_BlkLen32Avx512( const __m512i av00 = load_a(0); const __m512i av01 = load_a(1); const __m512i av02 = load_a(2); + // TailBlocks = BlockCountK % kBlockGroupBlks ∈ [1,3], so block 3 is never a real tail + // block — hardcode av03 = zero. The unpacked bv3 is still non-zero, but its + // contribution is zeroed twice: by av03==0 here and scale_a0_safe[3]==0 below. const __m512i av03 = zero; float scale_a0_safe[kBlockGroupBlks] = {0.0f, 0.0f, 0.0f, 0.0f}; diff --git a/onnxruntime/core/mlas/lib/sqnbitgemm_kernel_avx512_2bit_blklen64.h b/onnxruntime/core/mlas/lib/sqnbitgemm_kernel_avx512_2bit_blklen64.h index f3a46dcd6200b..a48599ebcecf1 100644 --- a/onnxruntime/core/mlas/lib/sqnbitgemm_kernel_avx512_2bit_blklen64.h +++ b/onnxruntime/core/mlas/lib/sqnbitgemm_kernel_avx512_2bit_blklen64.h @@ -121,6 +121,7 @@ dot_accumulate_4blk_w2(const __m512i& av0_64_epi8, const __m512i& av1_64_epi8, { __m512i d0, d1, d2, d3; if constexpr (kVnni) { + // dpbusd: 2nd operand (bv=unsigned [0,3]) x 3rd operand (av=signed int8); consistent with maddubs(bv, av) below d0 = _mm512_dpbusd_epi32(_mm512_setzero_si512(), bv0_64_epi8, av0_64_epi8); d1 = _mm512_dpbusd_epi32(_mm512_setzero_si512(), bv1_64_epi8, av1_64_epi8); d2 = _mm512_dpbusd_epi32(_mm512_setzero_si512(), bv2_64_epi8, av2_64_epi8); @@ -148,8 +149,8 @@ dot_accumulate_4blk_w2(const __m512i& av0_64_epi8, const __m512i& av1_64_epi8, // Two interleaved sub-accumulators: lo gets blocks {0, 2}, hi gets {1, 3}. // Each sub-accumulator chain is 2 FMAs deep (~8c) vs the 4-FMA single chain. - __m512 acc_lo = _mm512_fmadd_ps(_mm512_cvtepi32_ps(d0), s0, _mm512_setzero_ps()); - __m512 acc_hi = _mm512_fmadd_ps(_mm512_cvtepi32_ps(d1), s1, _mm512_setzero_ps()); + __m512 acc_lo = _mm512_mul_ps(_mm512_cvtepi32_ps(d0), s0); + __m512 acc_hi = _mm512_mul_ps(_mm512_cvtepi32_ps(d1), s1); acc_lo = _mm512_fmadd_ps(_mm512_cvtepi32_ps(d2), s2, acc_lo); acc_hi = _mm512_fmadd_ps(_mm512_cvtepi32_ps(d3), s3, acc_hi); acc = _mm512_add_ps(acc, _mm512_add_ps(acc_lo, acc_hi)); @@ -326,7 +327,10 @@ Q2Int8GemmR1xC4BlkLen64Avx512( const __m512i av02 = (TailBlocks > 2) ? _mm512_loadu_si512(reinterpret_cast(QuantAPtr + 2 * kBlkLen)) : zero; - const __m512i av03 = zero; // TailBlocks at most 3 + // TailBlocks = BlockCountK % kBlockGroupBlks in [1,3], so block 3 is never a real tail + // block. Keep av03 hardcoded to zero. The unpacked bv3 can still be non-zero, + // but its contribution is zeroed twice: by av03==0 and scale_a0_safe[3]==0. + const __m512i av03 = zero; // Bounded scale_a copy. float scale_a0_safe[kBlockGroupBlks] = {0.0f, 0.0f, 0.0f, 0.0f}; @@ -503,6 +507,9 @@ Q2Int8GemmR2xC4BlkLen64Avx512( const __m512i av02 = (TailBlocks > 2) ? _mm512_loadu_si512(reinterpret_cast(QuantAPtr + 2 * kBlkLen)) : zero; + // TailBlocks = BlockCountK % kBlockGroupBlks in [1,3], so block 3 is never a real tail + // block. Keep av03 hardcoded to zero. The unpacked bv3 can still be non-zero, + // but its contribution is zeroed twice: by av03==0 and scale_a0_safe[3]==0. const __m512i av03 = zero; const __m512i av10 = _mm512_loadu_si512( reinterpret_cast(QuantAPtr + lda + 0 * kBlkLen)); @@ -667,6 +674,9 @@ Q2Int8GemmRMxC_Tail_BlkLen64Avx512( const __m512i av02 = (TailBlocks > 2) ? _mm512_loadu_si512(reinterpret_cast(QuantAPtr + 2 * kBlkLen)) : zero; + // TailBlocks = BlockCountK % kBlockGroupBlks in [1,3], so block 3 is never a real tail + // block. Keep av03 hardcoded to zero. The unpacked bv3 can still be non-zero, + // but its contribution is zeroed twice: by av03==0 and scale_a0_safe[3]==0. const __m512i av03 = zero; float scale_a0_safe[kBlockGroupBlks] = {0.0f, 0.0f, 0.0f, 0.0f};