mirror of
https://github.com/saymrwulf/onnxruntime.git
synced 2026-07-29 20:14:01 +00:00
some fixes
Signed-off-by: Liqun Fu <liqun.fu@microsoft.com>
This commit is contained in:
parent
8c1cfe11d3
commit
f6f22e30d5
4 changed files with 50 additions and 28 deletions
|
|
@ -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;
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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;
|
||||
|
|
|
|||
|
|
@ -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();
|
||||
|
|
|
|||
Loading…
Reference in a new issue