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

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
10 changes: 10 additions & 0 deletions cmake/onnxruntime_mlas.cmake
Original file line number Diff line number Diff line change
Expand Up @@ -241,6 +241,11 @@ 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_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
Expand Down Expand Up @@ -791,6 +796,11 @@ 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_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
Expand Down
23 changes: 20 additions & 3 deletions onnxruntime/contrib_ops/cpu/quantization/matmul_nbits.cc
Original file line number Diff line number Diff line change
Expand Up @@ -359,8 +359,14 @@ Status MatMulNBits<T1>::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
// 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;
Expand Down Expand Up @@ -1074,7 +1080,18 @@ Status MatMulNBits<MLFloat16>::ComputeBUnpacked(const Tensor* a,
auto tmp_b_data_ptr = IAllocator::MakeUniquePtr<float>(allocator, SafeInt<size_t>(K_) * N_, true);

if ((reorder_idx_data == nullptr) && (!zero_points || !zero_points->IsDataType<MLFloat16>())) {
if (nbits_ == 4) {
if (nbits_ == 2) {
MlasDequantizeBlockwise<float, 2>(
tmp_b_data_ptr.get(), // dequantized output
b_data, // quantized input
scales_ptr, // quantization scales
static_cast<const uint8_t*>(zero_points_data), // quantization zero points
static_cast<int32_t>(block_size_), // quantization block size
column_wise_quant_, // columnwise quantization or row-wise
static_cast<int32_t>(K_), // number of rows in quantized input
static_cast<int32_t>(N_), // number of columns in quantized input
thread_pool);
} else if (nbits_ == 4) {
MlasDequantizeBlockwise<float, 4>(
tmp_b_data_ptr.get(), // dequantized output
b_data, // quantized input
Expand All @@ -1085,7 +1102,7 @@ Status MatMulNBits<MLFloat16>::ComputeBUnpacked(const Tensor* a,
static_cast<int32_t>(K_), // number of rows in quantized input
static_cast<int32_t>(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<float, 8>(
tmp_b_data_ptr.get(), // dequantized output
Expand Down
175 changes: 171 additions & 4 deletions onnxruntime/core/mlas/lib/qnbitgemm.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -33,6 +33,7 @@ enum QNBitGemmVariant {
HQ4BitGemmVariant_CompFp16,
SQ8BitGemmVariant_CompInt8,
HQ8BitGemmVariant_CompFp16,
SQ2BitGemmVariant_CompInt8,

// End of valid variants

Expand All @@ -48,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;
Expand All @@ -66,6 +73,12 @@ GetQNBitGemmVariant(
}
}

if (BlkBitWidth == 2 && (BlkLen == 32 || BlkLen == 64 || BlkLen == 128)) {
if (ComputeType == SQNBIT_CompInt8) {
return SQ2BitGemmVariant_CompInt8;
}
}

return SQNBitGemmVariantInvalid;
}

Expand Down Expand Up @@ -117,6 +130,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;
}
Expand All @@ -143,7 +162,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);
}

Expand All @@ -162,7 +181,7 @@ QNBitGemmPerGemmWorkspaceAlignment(
return 1;
}

if (BlkBitWidth == 4 || BlkBitWidth == 8) {
if (BlkBitWidth == 4 || BlkBitWidth == 8 || BlkBitWidth == 2) {
return Dispatch->QNBitGemmPerGemmWorkspaceAlignment(BlkLen, ComputeType);
}

Expand Down Expand Up @@ -240,6 +259,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;
Expand Down Expand Up @@ -355,6 +379,24 @@ MlasQNBitGemmPackQuantBData(
BackendKernelSelectorConfig
);
}
} else if (BlkBitWidth == 2) {
if (ComputeType == SQNBIT_CompInt8 && Dispatch->SQ2BitGemmPackQuantBDataAndBlkSum != nullptr) {
const size_t BlockCountK = MlasDivRoundup(K, BlkLen);
PackedQuantBDataStruct<float, 2> packed_quant_b(PackedQuantBDataAndOrBlkSumWorkspace, N, BlockCountK, BlkLen, false);
Dispatch->SQ2BitGemmPackQuantBDataAndBlkSum(
N,
K,
BlkLen,
ComputeType,
static_cast<const std::byte*>(QuantBData),
static_cast<const float*>(QuantBScale),
HasZeroPoint,
static_cast<const std::byte*>(QuantBZeroPoint),
packed_quant_b,
ThreadPool,
BackendKernelSelectorConfig
);
}
}
}

Expand Down Expand Up @@ -937,6 +979,111 @@ SQ8BitGemm_CompInt8(
}
}

//
// 2-bit weight CompInt8 wrapper. Mirrors SQ4BitGemm_CompInt8 but specialised
// 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(
const size_t BlkLen,
const size_t K,
const MLAS_QNBIT_GEMM_DATA_PARAMS<float>* 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<PerGemmQuantAWorkspace*>(PerGemmWorkspace);

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
// 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
// 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;

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_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;

// 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<const std::byte*>(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;
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_eff;
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 <typename T>
void
InitializeWorkspace_CompInt8(
Expand Down Expand Up @@ -1007,7 +1154,7 @@ InitializeWorkspace_CompInt8<float>(
});
} 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];
Expand Down Expand Up @@ -1090,6 +1237,7 @@ GetInitializeWorkspace(QNBitGemmVariant variant)
switch (variant) {
case SQ4BitGemmVariant_CompInt8:
case SQ8BitGemmVariant_CompInt8:
case SQ2BitGemmVariant_CompInt8:
return InitializeWorkspace_CompInt8<float>;
default:
return nullptr;
Expand Down Expand Up @@ -1132,6 +1280,10 @@ GetQNBitGemm(QNBitGemmVariant variant)
return SQ4BitGemm_CompInt8;
case SQ8BitGemmVariant_CompInt8:
return SQ8BitGemm_CompInt8;
case SQ2BitGemmVariant_CompInt8:
// W2 CompInt8 compute kernel registered by the AVX-512 /
// AVX-512-VNNI dispatch tables.
return SQ2BitGemm_CompInt8;
default:
return nullptr;
}
Expand Down Expand Up @@ -1255,6 +1407,13 @@ MlasQNBitGemmBatch(
const_cast<MLAS_QNBIT_GEMM_DATA_PARAMS<T>*>(Data)->QuantBScale = packed_quant_b.PackedQuantBScale;
const_cast<MLAS_QNBIT_GEMM_DATA_PARAMS<T>*>(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<T, 2> packed_quant_b(const_cast<void*>(Data->QuantBDataWorkspace), N, BlockCountK, BlkLen, false);
const_cast<MLAS_QNBIT_GEMM_DATA_PARAMS<T>*>(Data)->PackedQuantBData = packed_quant_b.PackedQuantBData;
const_cast<MLAS_QNBIT_GEMM_DATA_PARAMS<T>*>(Data)->QuantBBlkSum = packed_quant_b.QuantBBlkSum;
const_cast<MLAS_QNBIT_GEMM_DATA_PARAMS<T>*>(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 {
Expand Down Expand Up @@ -1336,6 +1495,14 @@ MlasQNBitGemmBatch(
const_cast<MLAS_QNBIT_GEMM_DATA_PARAMS<T>*>(Data)->QuantBScale = packed_quant_b.PackedQuantBScale;
const_cast<MLAS_QNBIT_GEMM_DATA_PARAMS<T>*>(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<T, 2> packed_quant_b(const_cast<void*>(Data->QuantBDataWorkspace), N, BlockCountK, BlkLen, false);
const_cast<MLAS_QNBIT_GEMM_DATA_PARAMS<T>*>(Data)->PackedQuantBData = packed_quant_b.PackedQuantBData;
const_cast<MLAS_QNBIT_GEMM_DATA_PARAMS<T>*>(Data)->QuantBBlkSum = packed_quant_b.QuantBBlkSum;
const_cast<MLAS_QNBIT_GEMM_DATA_PARAMS<T>*>(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 {
Expand Down
Loading
Loading