some fixes

Signed-off-by: Liqun Fu <liqun.fu@microsoft.com>
This commit is contained in:
Liqun Fu 2025-01-30 22:41:16 -08:00
parent 8c1cfe11d3
commit f6f22e30d5
4 changed files with 50 additions and 28 deletions

View file

@ -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<std::byte*>(Workspace) + gemm_i * PerGemmWorkspaceStride;
if (ComputeType == SQNBIT_CompInt8 && GetMlasPlatform().QNBitGemmDispatch->SQ4BitGemmPackQuantBDataAndBlkSum != nullptr) {
if (BlkBitWidth == 4 && ComputeType == SQNBIT_CompInt8 && GetMlasPlatform().QNBitGemmDispatch->SQ4BitGemmPackQuantBDataAndBlkSum != nullptr) {
PackedQuantBDataStruct<T> packed_quant_b(const_cast<void*>(Data->QuantBDataWorkspace), N, BlockCountK, BlkLen);
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;
@ -991,7 +992,7 @@ MlasQNBitGemmBatch(
void* PerGemmWorkspace =
reinterpret_cast<std::byte*>(Workspace) + gemm_i * PerGemmWorkspaceStride;
if (ComputeType == SQNBIT_CompInt8 && GetMlasPlatform().QNBitGemmDispatch->SQ4BitGemmPackQuantBDataAndBlkSum != nullptr) {
if (BlkBitWidth == 4 && ComputeType == SQNBIT_CompInt8 && GetMlasPlatform().QNBitGemmDispatch->SQ4BitGemmPackQuantBDataAndBlkSum != nullptr) {
PackedQuantBDataStruct<T> packed_quant_b(const_cast<void*>(Data->QuantBDataWorkspace), N, BlockCountK, BlkLen);
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;

View file

@ -15,45 +15,63 @@ Abstract:
--*/
#include <algorithm>
#include <cassert>
#include <utility>
#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

View file

@ -97,7 +97,7 @@ std::ostream& operator<<(std::ostream& os, const TestOptions& opts) {
<< ", has_bias:" << opts.has_bias;
}
template <typename T1, int qbits>
template <typename T1, int qbits=4>
void RunTest(const TestOptions& opts,
std::vector<std::unique_ptr<IExecutionProvider>>&& 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<AType, qbits>(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<AType, qbits>(opts);
}
if constexpr (qbits == 4)
{
if constexpr (qbits == 4) {
TestOptions opts = base_opts;
opts.has_g_idx = true;
opts.has_bias = true;

View file

@ -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<MlasSQNBitGemmTest<Blk
for (MLAS_QNBIT_GEMM_COMPUTE_TYPE ComputeType : {SQNBIT_CompFp32, SQNBIT_CompInt8}) {
for (bool WithThreadpool : {false, true}) {
for (bool Symmetric : {false, true}) {
if constexpr (BlkBitWidth == 2) {
if (SQNBIT_CompFp32 == ComputeType) {
continue;
}
}
for (size_t b = 1; b < 16; b++) {
tests_registered += RegisterSingleTest(b, b, b, ComputeType, WithThreadpool, Symmetric, false);
tests_registered += RegisterSingleTest(b, b, b, ComputeType, WithThreadpool, Symmetric, true);
@ -440,7 +446,7 @@ class SQNBitGemmShortExecuteTest : public MlasTestFixture<MlasSQNBitGemmTest<Blk
static size_t SQNBitGemmRegisterAllShortExecuteTests() {
size_t count = 0;
//count += SQNBitGemmShortExecuteTest<2, 16>::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();