diff --git a/onnxruntime/core/mlas/lib/qnbitgemm.cpp b/onnxruntime/core/mlas/lib/qnbitgemm.cpp index 7e7baf137f..096c795b4e 100644 --- a/onnxruntime/core/mlas/lib/qnbitgemm.cpp +++ b/onnxruntime/core/mlas/lib/qnbitgemm.cpp @@ -554,6 +554,7 @@ SQ2BitGemm_CompInt8( const size_t /*RangeCountN*/ ) { + // TODO: implement this to call 2bit t-mac kernel } void @@ -920,7 +921,7 @@ MlasQNBitGemmBatch( const auto* Data = &DataParams[gemm_i]; void* PerGemmWorkspace = reinterpret_cast(Workspace) + gemm_i * PerGemmWorkspaceStride; - if (ComputeType == SQNBIT_CompInt8 && GetMlasPlatform().QNBitGemmDispatch->SQ4BitGemmPackQuantBDataAndBlkSum != nullptr) { + if (BlkBitWidth == 4 && ComputeType == SQNBIT_CompInt8 && GetMlasPlatform().QNBitGemmDispatch->SQ4BitGemmPackQuantBDataAndBlkSum != nullptr) { PackedQuantBDataStruct packed_quant_b(const_cast(Data->QuantBDataWorkspace), N, BlockCountK, BlkLen); const_cast*>(Data)->PackedQuantBData = packed_quant_b.PackedQuantBData; const_cast*>(Data)->QuantBBlkSum = packed_quant_b.QuantBBlkSum; @@ -991,7 +992,7 @@ MlasQNBitGemmBatch( void* PerGemmWorkspace = reinterpret_cast(Workspace) + gemm_i * PerGemmWorkspaceStride; - if (ComputeType == SQNBIT_CompInt8 && GetMlasPlatform().QNBitGemmDispatch->SQ4BitGemmPackQuantBDataAndBlkSum != nullptr) { + if (BlkBitWidth == 4 && ComputeType == SQNBIT_CompInt8 && GetMlasPlatform().QNBitGemmDispatch->SQ4BitGemmPackQuantBDataAndBlkSum != nullptr) { PackedQuantBDataStruct packed_quant_b(const_cast(Data->QuantBDataWorkspace), N, BlockCountK, BlkLen); const_cast*>(Data)->PackedQuantBData = packed_quant_b.PackedQuantBData; const_cast*>(Data)->QuantBBlkSum = packed_quant_b.QuantBBlkSum; diff --git a/onnxruntime/core/mlas/lib/sqnbitgemm_bitnet_kernel_avx2.cpp b/onnxruntime/core/mlas/lib/sqnbitgemm_bitnet_kernel_avx2.cpp index 6c1a133609..1d7a1ce73e 100644 --- a/onnxruntime/core/mlas/lib/sqnbitgemm_bitnet_kernel_avx2.cpp +++ b/onnxruntime/core/mlas/lib/sqnbitgemm_bitnet_kernel_avx2.cpp @@ -15,45 +15,63 @@ Abstract: --*/ -#include -#include -#include - #include "qnbitgemm.h" +#include "sqnbitgemm_q8_block.h" size_t Q2BitGemmPackQuantBDataSize( - size_t /*N*/, - size_t /*K*/, - size_t /*BlkLen*/, - MLAS_QNBIT_GEMM_COMPUTE_TYPE /*ComputeType*/ + size_t N, + size_t K, + size_t BlkLen, + MLAS_QNBIT_GEMM_COMPUTE_TYPE ComputeType ) { - return 0; + // TODO: This code shall change according to T-Mac. + MLAS_UNREFERENCED_PARAMETER(ComputeType); // same size regardless of ComputeType + + constexpr size_t BlkBitWidth = 2; + + const size_t BlockCountK = MlasDivRoundup(K, BlkLen); + const size_t PackedQuantBDataSize = N * BlockCountK * MlasQNBitBlkDataSizeInBytes(BlkBitWidth, BlkLen); + return PackedQuantBDataSize; } void SQ2BitGemmPackQuantBData( size_t /*N*/, size_t /*K*/, size_t /*BlkLen*/, - MLAS_QNBIT_GEMM_COMPUTE_TYPE /* ComputeType*/, + MLAS_QNBIT_GEMM_COMPUTE_TYPE /*ComputeType*/, const std::byte* /*QuantBDataBegin*/, std::byte* /*PackedQuantBDataBegin*/, MLAS_THREADPOOL* /*ThreadPool*/ ) { + // TODO: need implementation } size_t Q2BitGemmPerGemmWorkspaceSize( - size_t /*M*/, - size_t /*N*/, - size_t /*K*/, - size_t /*BlkLen*/, - MLAS_QNBIT_GEMM_COMPUTE_TYPE /*ComputeType*/ + size_t M, + size_t N, + size_t K, + size_t BlkLen, + MLAS_QNBIT_GEMM_COMPUTE_TYPE ComputeType ) { - return 0; + MLAS_UNREFERENCED_PARAMETER(N); + + switch (ComputeType) { + case SQNBIT_CompInt8: { + // workspace buffer is used for block quantization of A to int8 + const size_t BlockCountK = MlasDivRoundup(K, BlkLen); + // QuantData + Scale + const size_t PerGemmWorkspaceSize = M * BlockCountK * Q8BlkSize(BlkLen); + return PerGemmWorkspaceSize; + } + default: { + return 0; + } + } } size_t diff --git a/onnxruntime/test/contrib_ops/matmul_4bits_test.cc b/onnxruntime/test/contrib_ops/matmul_4bits_test.cc index d6940dc2cf..bfd682ae39 100644 --- a/onnxruntime/test/contrib_ops/matmul_4bits_test.cc +++ b/onnxruntime/test/contrib_ops/matmul_4bits_test.cc @@ -97,7 +97,7 @@ std::ostream& operator<<(std::ostream& os, const TestOptions& opts) { << ", has_bias:" << opts.has_bias; } -template +template void RunTest(const TestOptions& opts, std::vector>&& explicit_eps = {}) { SCOPED_TRACE(opts); @@ -284,8 +284,7 @@ void TestMatMulNBitsTyped() { base_opts.output_rel_error = 0.02f; } - if constexpr (qbits == 4) - { + if constexpr (qbits == 4) { TestOptions opts = base_opts; RunTest(opts); } @@ -297,15 +296,13 @@ void TestMatMulNBitsTyped() { } #if !defined(USE_DML) && !defined(USE_WEBGPU) - if constexpr (qbits == 4) - { + if constexpr (qbits == 4) { TestOptions opts = base_opts; opts.has_g_idx = true; RunTest(opts); } - if constexpr (qbits == 4) - { + if constexpr (qbits == 4) { TestOptions opts = base_opts; opts.has_g_idx = true; opts.has_bias = true; diff --git a/onnxruntime/test/mlas/unittest/test_sqnbitgemm.cpp b/onnxruntime/test/mlas/unittest/test_sqnbitgemm.cpp index 365137d466..26f02466be 100644 --- a/onnxruntime/test/mlas/unittest/test_sqnbitgemm.cpp +++ b/onnxruntime/test/mlas/unittest/test_sqnbitgemm.cpp @@ -146,17 +146,18 @@ class MlasSQNBitGemmTest : public MlasTestBase { if constexpr (BlkBitWidth == 4) { b_zp = 8; } else if constexpr (BlkBitWidth == 2) { - assert(QuantBZeroPoint && "zero point input is needed for BlkBitWidth == 2"); + b_zp = 2; } else { static_assert(false, "only implemented for 2- and 4-bit quantized B"); } int pack_size = 8 / BlkBitWidth; if (QuantBZeroPoint != nullptr) { - const uint8_t b_zp_byte = QuantBZeroPoint[n * ((BlockCountK + 1) / pack_size) + k_blk / pack_size]; if constexpr (BlkBitWidth == 4) { + const uint8_t b_zp_byte = QuantBZeroPoint[n * ((BlockCountK + 1) / pack_size) + k_blk / pack_size]; b_zp = (k_blk & 1) ? (b_zp_byte >> 4) : (b_zp_byte & 0x0F); } else if constexpr (BlkBitWidth == 2) { + const uint8_t b_zp_byte = QuantBZeroPoint[n * ((BlockCountK + 3) / pack_size) + k_blk / pack_size]; int shift = (k_blk & 3) * 2; b_zp = (b_zp_byte >> shift) & 0x03; } @@ -396,6 +397,11 @@ class SQNBitGemmShortExecuteTest : public MlasTestFixture::RegisterShortExecuteTests(); - //count += SQNBitGemmShortExecuteTest<2, 32>::RegisterShortExecuteTests(); + count += SQNBitGemmShortExecuteTest<2, 32>::RegisterShortExecuteTests(); //count += SQNBitGemmShortExecuteTest<2, 64>::RegisterShortExecuteTests(); //count += SQNBitGemmShortExecuteTest<2, 128>::RegisterShortExecuteTests(); //count += SQNBitGemmShortExecuteTest<2, 256>::RegisterShortExecuteTests();