mirror of
https://github.com/saymrwulf/onnxruntime.git
synced 2026-07-30 20:18:08 +00:00
Matmul_nbits kernel for mlas sqnbits to support Fp16 inputs (#21807)
This commit is contained in:
parent
7e2c722459
commit
a89bddd5c2
11 changed files with 348 additions and 123 deletions
|
|
@ -580,10 +580,10 @@ message(STATUS "CMAKE_CXX_COMPILER_VERSION: ${CMAKE_CXX_COMPILER_VERSION}")
|
|||
|
||||
if(NOT "${CMAKE_CXX_COMPILER_ID}" STREQUAL "GNU" OR CMAKE_CXX_COMPILER_VERSION VERSION_GREATER "11")
|
||||
message(STATUS "Using -mavx2 -mfma -mavxvnni flags")
|
||||
set_source_files_properties(${mlas_platform_srcs_avx2} PROPERTIES COMPILE_FLAGS "-mavx2 -mfma -mavxvnni")
|
||||
set_source_files_properties(${mlas_platform_srcs_avx2} PROPERTIES COMPILE_FLAGS "-mavx2 -mfma -mf16c -mavxvnni")
|
||||
else()
|
||||
message(STATUS "Using -mavx2 -mfma flags")
|
||||
set_source_files_properties(${mlas_platform_srcs_avx2} PROPERTIES COMPILE_FLAGS "-mavx2 -mfma")
|
||||
set_source_files_properties(${mlas_platform_srcs_avx2} PROPERTIES COMPILE_FLAGS "-mavx2 -mfma -mf16c")
|
||||
endif()
|
||||
set(mlas_platform_srcs_avx512f
|
||||
${MLAS_SRC_DIR}/x86_64/DgemmKernelAvx512F.S
|
||||
|
|
|
|||
|
|
@ -488,7 +488,7 @@ Do not modify directly.*
|
|||
|MatMulFpQ4|*in* A:**T1**<br> *in* B:**T2**<br> *in* B_shape:**T3**<br> *out* Y:**T1**|1+|**T1** = tensor(float)<br/> **T2** = tensor(uint8)<br/> **T3** = tensor(int64)|
|
||||
|MatMulInteger16|*in* A:**T1**<br> *in* B:**T2**<br> *out* Y:**T3**|1+|**T1** = tensor(int16)<br/> **T2** = tensor(int16)<br/> **T3** = tensor(int32)|
|
||||
|MatMulIntegerToFloat|*in* A:**T1**<br> *in* B:**T2**<br> *in* a_scale:**T3**<br> *in* b_scale:**T3**<br> *in* a_zero_point:**T1**<br> *in* b_zero_point:**T2**<br> *in* bias:**T3**<br> *out* Y:**T3**|1+|**T1** = tensor(int8), tensor(uint8)<br/> **T2** = tensor(int8), tensor(uint8)<br/> **T3** = tensor(float)|
|
||||
|MatMulNBits|*in* A:**T1**<br> *in* B:**T2**<br> *in* scales:**T1**<br> *in* zero_points:**T3**<br> *in* g_idx:**T4**<br> *in* bias:**T1**<br> *out* Y:**T1**|1+|**T1** = tensor(float)<br/> **T2** = tensor(uint8)<br/> **T3** = tensor(float), tensor(uint8)<br/> **T4** = tensor(int32)|
|
||||
|MatMulNBits|*in* A:**T1**<br> *in* B:**T2**<br> *in* scales:**T1**<br> *in* zero_points:**T3**<br> *in* g_idx:**T4**<br> *in* bias:**T1**<br> *out* Y:**T1**|1+|**T1** = tensor(float), tensor(float16)<br/> **T2** = tensor(uint8)<br/> **T3** = tensor(float), tensor(float16), tensor(uint8)<br/> **T4** = tensor(int32)|
|
||||
|MaxpoolWithMask|*in* X:**T**<br> *in* M:**tensor(int32)**<br> *out* Y:**T**|1+|**T** = tensor(float)|
|
||||
|MultiHeadAttention|*in* query:**T**<br> *in* key:**T**<br> *in* value:**T**<br> *in* bias:**T**<br> *in* key_padding_mask:**M**<br> *in* attention_bias:**T**<br> *in* past_key:**T**<br> *in* past_value:**T**<br> *out* output:**T**<br> *out* present_key:**T**<br> *out* present_value:**T**|1+|**T** = tensor(float)|
|
||||
|MurmurHash3|*in* X:**T1**<br> *out* Y:**T2**|1+|**T1** = tensor(double), tensor(float), tensor(int32), tensor(int64), tensor(string), tensor(uint32), tensor(uint64)<br/> **T2** = tensor(int32), tensor(uint32)|
|
||||
|
|
|
|||
|
|
@ -146,8 +146,15 @@ class MatMulNBits final : public OpKernel {
|
|||
bool all_constant_{false};
|
||||
|
||||
#endif // defined(ORT_NEURAL_SPEED)
|
||||
|
||||
template <typename AType>
|
||||
Status ComputeTyped(OpKernelContext* ctx) const;
|
||||
};
|
||||
|
||||
bool IsATypeFloat16(const Tensor& tensor) {
|
||||
return tensor.GetElementType() == ONNX_NAMESPACE::TensorProto_DataType_FLOAT16;
|
||||
}
|
||||
|
||||
Status MatMulNBits::PrePack(const Tensor& tensor, int input_idx, /*out*/ AllocatorPtr alloc,
|
||||
/*out*/ bool& is_packed,
|
||||
/*out*/ PrePackedWeights* prepacked_weights) {
|
||||
|
|
@ -211,10 +218,10 @@ Status MatMulNBits::PrePack(const Tensor& tensor, int input_idx, /*out*/ Allocat
|
|||
#else // defined(ORT_NEURAL_SPEED)
|
||||
ORT_UNUSED_PARAMETER(prepacked_weights);
|
||||
const auto compute_type = static_cast<MLAS_SQNBIT_GEMM_COMPUTE_TYPE>(accuracy_level_);
|
||||
if (!MlasIsSQNBitGemmAvailable(nbits_, block_size_, compute_type)) {
|
||||
return Status::OK();
|
||||
}
|
||||
if (input_idx == InputIndex::B) {
|
||||
if (!MlasIsSQNBitGemmAvailable(nbits_, block_size_, compute_type)) {
|
||||
return Status::OK();
|
||||
}
|
||||
packed_b_size_ = MlasSQNBitGemmPackQuantBDataSize(N_, K_, nbits_, block_size_, compute_type);
|
||||
if (packed_b_size_ == 0) {
|
||||
return Status::OK();
|
||||
|
|
@ -226,8 +233,15 @@ Status MatMulNBits::PrePack(const Tensor& tensor, int input_idx, /*out*/ Allocat
|
|||
} else if (compute_type == CompInt8) {
|
||||
#ifdef MLAS_TARGET_AMD64_IX86
|
||||
if (input_idx == InputIndex::scales && packed_b_ != nullptr) {
|
||||
auto sptr = tensor.Data<float>();
|
||||
MlasSQNBitGemmPackQuantBData(N_, K_, nbits_, block_size_, compute_type, nullptr, packed_b_.get(), sptr, has_zp_input_, nullptr, nullptr);
|
||||
if (IsATypeFloat16(tensor)) {
|
||||
auto sptr = tensor.Data<MLFloat16>();
|
||||
std::vector<float> scales_v(static_cast<unsigned int>(tensor.Shape().Size()));
|
||||
MlasConvertHalfToFloatBuffer(sptr, &scales_v[0], scales_v.size());
|
||||
MlasSQNBitGemmPackQuantBData(N_, K_, nbits_, block_size_, compute_type, nullptr, packed_b_.get(), &scales_v[0], has_zp_input_, nullptr, nullptr);
|
||||
} else {
|
||||
auto sptr = tensor.Data<float>();
|
||||
MlasSQNBitGemmPackQuantBData(N_, K_, nbits_, block_size_, compute_type, nullptr, packed_b_.get(), sptr, has_zp_input_, nullptr, nullptr);
|
||||
}
|
||||
is_packed = false;
|
||||
} else if (input_idx == InputIndex::zero_points && packed_b_ != nullptr) {
|
||||
auto zptr = tensor.Data<uint8_t>();
|
||||
|
|
@ -274,9 +288,20 @@ Status MatMulNBits::UseSharedPrePackedBuffers(std::vector<BufferUniquePtr>& prep
|
|||
}
|
||||
|
||||
Status MatMulNBits::Compute(OpKernelContext* ctx) const {
|
||||
const Tensor* a = ctx->Input<Tensor>(InputIndex::A);
|
||||
|
||||
if (IsATypeFloat16(*a)) {
|
||||
return ComputeTyped<MLFloat16>(ctx);
|
||||
} else {
|
||||
return ComputeTyped<float>(ctx);
|
||||
}
|
||||
}
|
||||
|
||||
template <typename AType>
|
||||
Status MatMulNBits::ComputeTyped(OpKernelContext* ctx) const {
|
||||
concurrency::ThreadPool* thread_pool = ctx->GetOperatorThreadPool();
|
||||
const Tensor* a = ctx->Input<Tensor>(InputIndex::A);
|
||||
const auto* a_data = a->Data<float>();
|
||||
const auto* a_data = a->Data<AType>();
|
||||
|
||||
TensorShape b_shape({static_cast<int64_t>(N_), static_cast<int64_t>(K_)});
|
||||
MatMulComputeHelper helper;
|
||||
|
|
@ -289,7 +314,7 @@ Status MatMulNBits::Compute(OpKernelContext* ctx) const {
|
|||
return Status::OK();
|
||||
}
|
||||
|
||||
auto* y_data = y->MutableData<float>();
|
||||
auto* y_data = y->MutableData<AType>();
|
||||
|
||||
const size_t batch_count = helper.OutputOffsets().size();
|
||||
const size_t M = static_cast<size_t>(helper.M());
|
||||
|
|
@ -297,9 +322,12 @@ Status MatMulNBits::Compute(OpKernelContext* ctx) const {
|
|||
const size_t K = static_cast<size_t>(helper.K());
|
||||
const size_t lda = helper.Lda(false);
|
||||
|
||||
const bool has_single_b_matrix = std::all_of(helper.RightOffsets().begin(),
|
||||
helper.RightOffsets().end(),
|
||||
[](size_t offset) { return offset == 0; });
|
||||
// clang-format off
|
||||
const bool has_single_b_matrix = std::all_of(
|
||||
helper.RightOffsets().begin(),
|
||||
helper.RightOffsets().end(),
|
||||
[](size_t offset) { return offset == 0; });
|
||||
// clang-format on
|
||||
|
||||
#if defined(ORT_NEURAL_SPEED)
|
||||
|
||||
|
|
@ -336,9 +364,9 @@ Status MatMulNBits::Compute(OpKernelContext* ctx) const {
|
|||
const Tensor* zero_points = ctx->Input<Tensor>(InputIndex::zero_points);
|
||||
const Tensor* bias = ctx->Input<Tensor>(InputIndex::bias);
|
||||
|
||||
const auto* scales_data = scales->Data<float>();
|
||||
const auto* scales_data = scales->Data<AType>();
|
||||
const auto* zero_points_data = zero_points == nullptr ? nullptr : zero_points->DataRaw();
|
||||
const auto* bias_data = bias == nullptr ? nullptr : bias->Data<float>();
|
||||
const auto* bias_data = bias == nullptr ? nullptr : bias->Data<AType>();
|
||||
|
||||
IAllocatorUniquePtr<std::byte> workspace{};
|
||||
const size_t workspace_size = MlasSQNBitGemmBatchWorkspaceSize(
|
||||
|
|
@ -349,26 +377,64 @@ Status MatMulNBits::Compute(OpKernelContext* ctx) const {
|
|||
workspace = IAllocator::MakeUniquePtr<std::byte>(allocator, workspace_size);
|
||||
}
|
||||
|
||||
InlinedVector<MLAS_SQNBIT_GEMM_DATA_PARAMS> data(batch_count);
|
||||
for (size_t i = 0; i < batch_count; ++i) {
|
||||
data[i].A = a_data + helper.LeftOffsets()[i];
|
||||
data[i].lda = lda;
|
||||
#ifdef MLAS_TARGET_AMD64_IX86
|
||||
if (compute_type == CompInt8) {
|
||||
data[i].QuantBDataWorkspace = packed_b_.get();
|
||||
}
|
||||
#endif
|
||||
data[i].PackedQuantBData = static_cast<std::byte*>(packed_b_.get());
|
||||
data[i].QuantBScale = scales_data;
|
||||
data[i].QuantBZeroPoint = zero_points_data;
|
||||
data[i].Bias = bias_data;
|
||||
data[i].C = y_data + helper.OutputOffsets()[i];
|
||||
data[i].ldc = N;
|
||||
}
|
||||
MlasSQNBitGemmBatch(M, N, K, batch_count, nbits_, block_size_, compute_type, data.data(), workspace.get(),
|
||||
thread_pool);
|
||||
if constexpr (std::is_same<AType, MLFloat16>::value) {
|
||||
InlinedVector<MLAS_SQNBIT_GEMM_DATA_PARAMS> data(batch_count);
|
||||
|
||||
return Status::OK();
|
||||
AllocatorPtr allocator;
|
||||
ORT_RETURN_IF_ERROR(ctx->GetTempSpaceAllocator(&allocator));
|
||||
|
||||
auto tmp_a_data_ptr = IAllocator::MakeUniquePtr<float>(allocator, (size_t)(a->Shape().Size()));
|
||||
MlasConvertHalfToFloatBuffer(a_data, tmp_a_data_ptr.get(), static_cast<size_t>(a->Shape().Size()));
|
||||
|
||||
auto tmp_scales_data_ptr = IAllocator::MakeUniquePtr<float>(allocator, (size_t)(scales->Shape().Size()));
|
||||
MlasConvertHalfToFloatBuffer(scales_data, tmp_scales_data_ptr.get(), static_cast<size_t>(scales->Shape().Size()));
|
||||
|
||||
std::vector<float> bias_data_v;
|
||||
if (bias_data != nullptr) {
|
||||
bias_data_v.resize((const unsigned int)(bias->Shape().Size()));
|
||||
MlasConvertHalfToFloatBuffer(bias_data, &bias_data_v[0], bias_data_v.size());
|
||||
}
|
||||
std::vector<float> C_v((const unsigned int)(y->Shape().Size()));
|
||||
for (size_t i = 0; i < batch_count; ++i) {
|
||||
data[i].A = tmp_a_data_ptr.get() + helper.LeftOffsets()[i];
|
||||
data[i].lda = lda;
|
||||
#ifdef MLAS_TARGET_AMD64_IX86
|
||||
if (compute_type == CompInt8) {
|
||||
data[i].QuantBDataWorkspace = packed_b_.get();
|
||||
}
|
||||
#endif
|
||||
data[i].PackedQuantBData = static_cast<std::byte*>(packed_b_.get());
|
||||
data[i].QuantBScale = tmp_scales_data_ptr.get();
|
||||
data[i].QuantBZeroPoint = zero_points_data;
|
||||
data[i].Bias = bias_data != nullptr ? &bias_data_v[0] : nullptr;
|
||||
data[i].C = &C_v[0] + helper.OutputOffsets()[i];
|
||||
data[i].ldc = N;
|
||||
}
|
||||
MlasSQNBitGemmBatch(M, N, K, batch_count, nbits_, block_size_, compute_type, data.data(), workspace.get(),
|
||||
thread_pool);
|
||||
MlasConvertFloatToHalfBuffer(&C_v[0], y_data, C_v.size());
|
||||
return Status::OK();
|
||||
} else {
|
||||
InlinedVector<MLAS_SQNBIT_GEMM_DATA_PARAMS> data(batch_count);
|
||||
for (size_t i = 0; i < batch_count; ++i) {
|
||||
data[i].A = a_data + helper.LeftOffsets()[i];
|
||||
data[i].lda = lda;
|
||||
#ifdef MLAS_TARGET_AMD64_IX86
|
||||
if (compute_type == CompInt8) {
|
||||
data[i].QuantBDataWorkspace = packed_b_.get();
|
||||
}
|
||||
#endif
|
||||
data[i].PackedQuantBData = static_cast<std::byte*>(packed_b_.get());
|
||||
data[i].QuantBScale = scales_data;
|
||||
data[i].QuantBZeroPoint = zero_points_data;
|
||||
data[i].Bias = bias_data;
|
||||
data[i].C = y_data + helper.OutputOffsets()[i];
|
||||
data[i].ldc = N;
|
||||
}
|
||||
MlasSQNBitGemmBatch(M, N, K, batch_count, nbits_, block_size_, compute_type, data.data(), workspace.get(),
|
||||
thread_pool);
|
||||
return Status::OK();
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
|
|
@ -380,7 +446,17 @@ Status MatMulNBits::Compute(OpKernelContext* ctx) const {
|
|||
const Tensor* zero_points = ctx->Input<Tensor>(InputIndex::zero_points);
|
||||
const Tensor* reorder_idx = ctx->Input<Tensor>(InputIndex::g_idx);
|
||||
|
||||
const auto* scales_data = scales->Data<float>();
|
||||
const auto* scales_data = scales->Data<AType>();
|
||||
const float* scales_data_;
|
||||
std::vector<float> scales_data_v;
|
||||
if constexpr (std::is_same<AType, MLFloat16>::value) {
|
||||
scales_data_v.resize((const unsigned int)scales->Shape().Size());
|
||||
MlasConvertHalfToFloatBuffer(scales_data, &scales_data_v[0], scales_data_v.size());
|
||||
scales_data_ = &scales_data_v[0];
|
||||
} else {
|
||||
scales_data_ = scales_data;
|
||||
}
|
||||
|
||||
const auto* zero_points_data = zero_points == nullptr ? nullptr : zero_points->DataRaw();
|
||||
const auto* reorder_idx_data = reorder_idx == nullptr ? nullptr : reorder_idx->Data<int32_t>();
|
||||
|
||||
|
|
@ -391,12 +467,12 @@ Status MatMulNBits::Compute(OpKernelContext* ctx) const {
|
|||
AllocatorPtr allocator;
|
||||
ORT_RETURN_IF_ERROR(ctx->GetTempSpaceAllocator(&allocator));
|
||||
auto tmp_b_data_ptr = IAllocator::MakeUniquePtr<float>(allocator, SafeInt<size_t>(K_) * N_);
|
||||
if ((reorder_idx_data == nullptr) && (!zero_points || !zero_points->IsDataType<float>())) {
|
||||
if ((reorder_idx_data == nullptr) && (!zero_points || !zero_points->IsDataType<AType>())) {
|
||||
// dequantize b, only 4b quantization is supported for now
|
||||
MlasDequantizeBlockwise<float, 4>(
|
||||
tmp_b_data_ptr.get(), // dequantized output
|
||||
b_data, // quantized input
|
||||
scales_data, // quantization scales
|
||||
scales_data_, // 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
|
||||
|
|
@ -406,12 +482,12 @@ Status MatMulNBits::Compute(OpKernelContext* ctx) const {
|
|||
} else {
|
||||
ORT_ENFORCE(column_wise_quant_, "Row-wise quantization is not supported for now");
|
||||
// !!!!!!!!!!!!!! naive implementation, need to be optimized !!!!!!!!!!!!!!
|
||||
if ((zero_points && zero_points->IsDataType<float>())) {
|
||||
DequantizeBlockwise<float, float>(
|
||||
if ((zero_points && zero_points->IsDataType<AType>())) {
|
||||
DequantizeBlockwise<float, AType>(
|
||||
tmp_b_data_ptr.get(), // dequantized output
|
||||
b_data, // quantized input
|
||||
scales_data, // quantization scales
|
||||
static_cast<const float*>(zero_points_data), // quantization zero points
|
||||
scales_data_, // quantization scales
|
||||
static_cast<const AType*>(zero_points_data), // quantization zero points
|
||||
reorder_idx_data,
|
||||
static_cast<int32_t>(block_size_), // quantization block size
|
||||
column_wise_quant_, // columnwise quantization or row-wise
|
||||
|
|
@ -422,7 +498,7 @@ Status MatMulNBits::Compute(OpKernelContext* ctx) const {
|
|||
DequantizeBlockwise<float, uint8_t>(
|
||||
tmp_b_data_ptr.get(), // dequantized output
|
||||
b_data, // quantized input
|
||||
scales_data, // quantization scales
|
||||
scales_data_, // quantization scales
|
||||
static_cast<const uint8_t*>(zero_points_data), // quantization zero points
|
||||
reorder_idx_data,
|
||||
static_cast<int32_t>(block_size_), // quantization block size
|
||||
|
|
@ -436,40 +512,80 @@ Status MatMulNBits::Compute(OpKernelContext* ctx) const {
|
|||
auto tm_b_data_ptr_trans = IAllocator::MakeUniquePtr<float>(allocator, SafeInt<size_t>(K_) * N_);
|
||||
MlasTranspose(tmp_b_data_ptr.get(), tm_b_data_ptr_trans.get(), N_, K_);
|
||||
#endif
|
||||
if constexpr (std::is_same<AType, MLFloat16>::value) {
|
||||
std::vector<MLAS_SGEMM_DATA_PARAMS> data(batch_count);
|
||||
|
||||
std::vector<MLAS_SGEMM_DATA_PARAMS> data(batch_count);
|
||||
for (size_t i = 0; i < batch_count; i++) {
|
||||
data[i].BIsPacked = false;
|
||||
data[i].A = a_data + helper.LeftOffsets()[i];
|
||||
data[i].lda = lda;
|
||||
data[i].B = tmp_b_data_ptr.get() + helper.RightOffsets()[i];
|
||||
data[i].ldb = ldb;
|
||||
data[i].C = y_data + helper.OutputOffsets()[i];
|
||||
data[i].ldc = N;
|
||||
data[i].alpha = 1.f;
|
||||
data[i].beta = 0.0f;
|
||||
}
|
||||
auto tmp_a_data_ptr = IAllocator::MakeUniquePtr<float>(allocator, (size_t)(a->Shape().Size()));
|
||||
MlasConvertHalfToFloatBuffer(a_data, tmp_a_data_ptr.get(), static_cast<size_t>(a->Shape().Size()));
|
||||
|
||||
// if there is a bias input, copy bias values into C and set beta to 1.0f
|
||||
if (const Tensor* bias = ctx->Input<Tensor>(InputIndex::bias);
|
||||
bias != nullptr) {
|
||||
gsl::span<const float> bias_span = bias->DataAsSpan<float>();
|
||||
for (size_t i = 0; i < batch_count; ++i) {
|
||||
float* C_row = data[i].C;
|
||||
const size_t ldc = data[i].ldc;
|
||||
for (size_t m = 0; m < M; ++m) {
|
||||
memcpy(C_row, bias_span.data(), bias_span.size_bytes());
|
||||
C_row += ldc;
|
||||
}
|
||||
|
||||
data[i].beta = 1.0f;
|
||||
auto tmp_c_ptr = IAllocator::MakeUniquePtr<float>(allocator, (size_t)(y->Shape().Size()));
|
||||
for (size_t i = 0; i < batch_count; i++) {
|
||||
data[i].BIsPacked = false;
|
||||
data[i].A = tmp_a_data_ptr.get() + helper.LeftOffsets()[i];
|
||||
data[i].lda = lda;
|
||||
data[i].B = tmp_b_data_ptr.get() + helper.RightOffsets()[i];
|
||||
data[i].ldb = ldb;
|
||||
data[i].C = tmp_c_ptr.get() + helper.OutputOffsets()[i];
|
||||
data[i].ldc = N;
|
||||
data[i].alpha = 1.f;
|
||||
data[i].beta = 0.0f;
|
||||
}
|
||||
|
||||
// if there is a bias input, copy bias values into C and set beta to 1.0f
|
||||
if (const Tensor* bias = ctx->Input<Tensor>(InputIndex::bias);
|
||||
bias != nullptr) {
|
||||
auto tmp_bias_data_ptr = IAllocator::MakeUniquePtr<float>(allocator, (size_t)(bias->Shape().Size()));
|
||||
MlasConvertHalfToFloatBuffer(bias->Data<AType>(), tmp_bias_data_ptr.get(), static_cast<size_t>(bias->Shape().Size()));
|
||||
for (size_t i = 0; i < batch_count; ++i) {
|
||||
float* C_row = data[i].C;
|
||||
const size_t ldc = data[i].ldc;
|
||||
for (size_t m = 0; m < M; ++m) {
|
||||
std::copy(tmp_bias_data_ptr.get(), tmp_bias_data_ptr.get() + bias->Shape().Size(), C_row);
|
||||
C_row += ldc;
|
||||
}
|
||||
data[i].beta = 1.0f;
|
||||
}
|
||||
}
|
||||
|
||||
MlasGemmBatch(CblasNoTrans, CblasTrans,
|
||||
M, N, K, data.data(), batch_count, thread_pool);
|
||||
MlasConvertFloatToHalfBuffer(tmp_c_ptr.get(), y_data, static_cast<size_t>(y->Shape().Size()));
|
||||
return Status::OK();
|
||||
} else {
|
||||
std::vector<MLAS_SGEMM_DATA_PARAMS> data(batch_count);
|
||||
for (size_t i = 0; i < batch_count; i++) {
|
||||
data[i].BIsPacked = false;
|
||||
data[i].A = a_data + helper.LeftOffsets()[i];
|
||||
data[i].lda = lda;
|
||||
data[i].B = tmp_b_data_ptr.get() + helper.RightOffsets()[i];
|
||||
data[i].ldb = ldb;
|
||||
data[i].C = y_data + helper.OutputOffsets()[i];
|
||||
data[i].ldc = N;
|
||||
data[i].alpha = 1.f;
|
||||
data[i].beta = 0.0f;
|
||||
}
|
||||
|
||||
// if there is a bias input, copy bias values into C and set beta to 1.0f
|
||||
if (const Tensor* bias = ctx->Input<Tensor>(InputIndex::bias);
|
||||
bias != nullptr) {
|
||||
gsl::span<const float> bias_span = bias->DataAsSpan<float>();
|
||||
for (size_t i = 0; i < batch_count; ++i) {
|
||||
float* C_row = data[i].C;
|
||||
const size_t ldc = data[i].ldc;
|
||||
for (size_t m = 0; m < M; ++m) {
|
||||
memcpy(C_row, bias_span.data(), bias_span.size_bytes());
|
||||
C_row += ldc;
|
||||
}
|
||||
|
||||
data[i].beta = 1.0f;
|
||||
}
|
||||
}
|
||||
|
||||
MlasGemmBatch(CblasNoTrans, CblasTrans,
|
||||
M, N, K, data.data(), batch_count, thread_pool);
|
||||
|
||||
return Status::OK();
|
||||
}
|
||||
|
||||
MlasGemmBatch(CblasNoTrans, CblasTrans,
|
||||
M, N, K, data.data(), batch_count, thread_pool);
|
||||
|
||||
return Status::OK();
|
||||
}
|
||||
|
||||
ONNX_OPERATOR_KERNEL_EX(
|
||||
|
|
@ -478,9 +594,9 @@ ONNX_OPERATOR_KERNEL_EX(
|
|||
1,
|
||||
kCpuExecutionProvider,
|
||||
KernelDefBuilder()
|
||||
.TypeConstraint("T1", DataTypeImpl::GetTensorType<float>())
|
||||
.TypeConstraint("T1", {DataTypeImpl::GetTensorType<float>(), DataTypeImpl::GetTensorType<MLFloat16>()})
|
||||
.TypeConstraint("T2", DataTypeImpl::GetTensorType<uint8_t>())
|
||||
.TypeConstraint("T3", {DataTypeImpl::GetTensorType<uint8_t>(), DataTypeImpl::GetTensorType<float>()})
|
||||
.TypeConstraint("T3", {DataTypeImpl::GetTensorType<uint8_t>(), DataTypeImpl::GetTensorType<float>(), DataTypeImpl::GetTensorType<MLFloat16>()})
|
||||
.TypeConstraint("T4", DataTypeImpl::GetTensorType<int32_t>()),
|
||||
MatMulNBits);
|
||||
|
||||
|
|
|
|||
|
|
@ -54,12 +54,12 @@ void Dequantize4BitsKernelReOrder(
|
|||
T scale = *(scale_data + n_idx * scales_shape_x + rid);
|
||||
float zp_f = 8;
|
||||
if (zero_points) {
|
||||
if constexpr (std::is_same_v<zeroT, T>) {
|
||||
zp_f = *(zero_points + n_idx * scales_shape_x + rid);
|
||||
} else {
|
||||
if constexpr (std::is_same_v<zeroT, uint8_t>) {
|
||||
uint8_t zp = 8;
|
||||
zp = zero_points[n_idx * zero_point_shape_x + rid / 2];
|
||||
zp = (rid & 0x01) ? (zp >> 4) : (zp & 0x0f);
|
||||
} else {
|
||||
zp_f = *(zero_points + static_cast<uint64_t>(n_idx) * static_cast<uint64_t>(scales_shape_x) + static_cast<uint64_t>(rid));
|
||||
}
|
||||
}
|
||||
|
||||
|
|
@ -112,5 +112,10 @@ template void DequantizeBlockwise<float, float>(
|
|||
const float* zero_points, const int32_t* reorder_idx, int32_t block_size,
|
||||
bool columnwise, int32_t K, int32_t N, onnxruntime::concurrency::ThreadPool* thread_pool);
|
||||
|
||||
template void DequantizeBlockwise<float, MLFloat16>(
|
||||
float* output, const uint8_t* quant_data, const float* scales_data,
|
||||
const MLFloat16* zero_points, const int32_t* reorder_idx, int32_t block_size,
|
||||
bool columnwise, int32_t K, int32_t N, onnxruntime::concurrency::ThreadPool* thread_pool);
|
||||
|
||||
} // namespace contrib
|
||||
} // namespace onnxruntime
|
||||
|
|
|
|||
|
|
@ -20,6 +20,7 @@ Abstract:
|
|||
#include <cstddef>
|
||||
#include <cstdlib>
|
||||
#include <cstdint>
|
||||
#include <stdexcept>
|
||||
|
||||
//
|
||||
// Define the calling convention for Windows targets.
|
||||
|
|
@ -1025,18 +1026,6 @@ MlasComputeTanh(
|
|||
size_t N
|
||||
);
|
||||
|
||||
//
|
||||
// Half-precision floating-point routines.
|
||||
//
|
||||
|
||||
void
|
||||
MLASCALL
|
||||
MlasConvertHalfToFloatBuffer(
|
||||
const unsigned short* Source,
|
||||
float* Destination,
|
||||
size_t Count
|
||||
);
|
||||
|
||||
//
|
||||
// Transpose routines.
|
||||
//
|
||||
|
|
@ -1426,7 +1415,27 @@ using MLAS_FP16 = onnxruntime::MLFloat16;
|
|||
|
||||
constexpr size_t FP16_SIZE = sizeof(uint16_t);
|
||||
|
||||
/**
|
||||
//
|
||||
// Half-precision floating-point routines.
|
||||
//
|
||||
|
||||
void
|
||||
MLASCALL
|
||||
MlasConvertHalfToFloatBuffer(
|
||||
const MLAS_FP16* Source,
|
||||
float* Destination,
|
||||
size_t Count
|
||||
);
|
||||
|
||||
void
|
||||
MLASCALL
|
||||
MlasConvertFloatToHalfBuffer(
|
||||
const float* Source,
|
||||
MLAS_FP16* Destination,
|
||||
size_t Count
|
||||
);
|
||||
|
||||
/**
|
||||
* @brief Whether current CPU supports FP16 acceleration.
|
||||
*/
|
||||
bool MLASCALL
|
||||
|
|
@ -1787,6 +1796,7 @@ MlasTranspose(
|
|||
M, N);
|
||||
}
|
||||
|
||||
|
||||
#ifdef MLAS_F16VEC_INTRINSICS_SUPPORTED
|
||||
/**
|
||||
* @brief Max Pooling for fp16 NHWC
|
||||
|
|
|
|||
|
|
@ -23,37 +23,35 @@ union fp32_bits {
|
|||
void
|
||||
MLASCALL
|
||||
MlasConvertHalfToFloatBuffer(
|
||||
const unsigned short* Source,
|
||||
const MLAS_FP16* Source,
|
||||
float* Destination,
|
||||
size_t Count
|
||||
)
|
||||
{
|
||||
|
||||
if (GetMlasPlatform().CastF16ToF32Kernel == nullptr) {
|
||||
// If there is no kernel use the reference implementation, adapted from mlas_float16.h.
|
||||
constexpr fp32_bits magic = {113 << 23};
|
||||
constexpr uint32_t shifted_exp = 0x7c00 << 13; // exponent mask after shift
|
||||
|
||||
for (size_t i = 0; i < Count; ++i) {
|
||||
fp32_bits o;
|
||||
o.u = (Source[i] & 0x7fff) << 13; // exponent/mantissa bits
|
||||
uint32_t exp = shifted_exp & o.u; // just the exponent
|
||||
o.u += (127 - 15) << 23; // exponent adjust
|
||||
|
||||
// handle exponent special cases
|
||||
if (exp == shifted_exp) { // Inf/NaN?
|
||||
o.u += (128 - 16) << 23; // extra exp adjust
|
||||
} else if (exp == 0) { // Zero/Denormal?
|
||||
o.u += 1 << 23; // extra exp adjust
|
||||
o.f -= magic.f; // renormalize
|
||||
}
|
||||
|
||||
o.u |= (Source[i] & 0x8000) << 16; // sign bit
|
||||
Destination[i] = o.f;
|
||||
Destination[i] = Source[i].ToFloat();
|
||||
}
|
||||
|
||||
} else {
|
||||
// If the kernel is available, use it to perform the conversion.
|
||||
GetMlasPlatform().CastF16ToF32Kernel(Source, Destination, Count);
|
||||
GetMlasPlatform().CastF16ToF32Kernel(reinterpret_cast<const unsigned short*>(Source), Destination, Count);
|
||||
}
|
||||
}
|
||||
|
||||
void
|
||||
MLASCALL
|
||||
MlasConvertFloatToHalfBuffer(
|
||||
const float* Source,
|
||||
MLAS_FP16* Destination,
|
||||
size_t Count
|
||||
)
|
||||
{
|
||||
if (GetMlasPlatform().CastF32ToF16Kernel == nullptr) {
|
||||
for (size_t i = 0; i < Count; ++i) {
|
||||
Destination[i] = MLAS_FP16(Source[i]);
|
||||
}
|
||||
} else {
|
||||
// If the kernel is available, use it to perform the conversion.
|
||||
GetMlasPlatform().CastF32ToF16Kernel(Source, reinterpret_cast<unsigned short*>(Destination), Count);
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -610,13 +610,19 @@ void
|
|||
size_t N
|
||||
);
|
||||
|
||||
typedef
|
||||
typedef
|
||||
void(MLASCALL MLAS_CAST_F16_TO_F32_KERNEL)(
|
||||
const unsigned short* Source,
|
||||
float* Destination,
|
||||
size_t Count
|
||||
);
|
||||
|
||||
typedef void(MLASCALL MLAS_CAST_F32_TO_F16_KERNEL)(
|
||||
const float* Source,
|
||||
unsigned short* Destination,
|
||||
size_t Count
|
||||
);
|
||||
|
||||
typedef
|
||||
void
|
||||
(MLASCALL MLAS_QLINEAR_BINARY_OP_S8_KERNEL)(
|
||||
|
|
@ -880,6 +886,8 @@ extern "C" {
|
|||
#if defined(MLAS_TARGET_AMD64)
|
||||
MLAS_CAST_F16_TO_F32_KERNEL MlasCastF16ToF32KernelSse;
|
||||
MLAS_CAST_F16_TO_F32_KERNEL MlasCastF16ToF32KernelAvx;
|
||||
MLAS_CAST_F16_TO_F32_KERNEL MlasCastF16ToF32KernelAvx2;
|
||||
MLAS_CAST_F32_TO_F16_KERNEL MlasCastF32ToF16KernelAvx2;
|
||||
#endif
|
||||
|
||||
}
|
||||
|
|
@ -1165,6 +1173,7 @@ struct MLAS_PLATFORM {
|
|||
const MLAS_SQNBIT_GEMM_DISPATCH* SQNBitGemmDispatch{nullptr};
|
||||
|
||||
MLAS_CAST_F16_TO_F32_KERNEL* CastF16ToF32Kernel;
|
||||
MLAS_CAST_F32_TO_F16_KERNEL* CastF32ToF16Kernel;
|
||||
};
|
||||
|
||||
inline
|
||||
|
|
|
|||
|
|
@ -245,6 +245,7 @@ Return Value:
|
|||
this->ConvDepthwiseS8S8Kernel = MlasConvDepthwiseKernel<int8_t, int8_t>;
|
||||
this->ConvDepthwiseS8U8Kernel = MlasConvDepthwiseKernel<int8_t, uint8_t>;
|
||||
this->CastF16ToF32Kernel = nullptr;
|
||||
this->CastF32ToF16Kernel = nullptr;
|
||||
|
||||
#if defined(MLAS_TARGET_AMD64_IX86)
|
||||
|
||||
|
|
@ -387,6 +388,9 @@ Return Value:
|
|||
this->ConvDepthwiseS8U8Kernel = MlasConvDepthwiseKernelAvx2<int8_t, uint8_t>;
|
||||
this->ComputeSumExpF32Kernel = MlasComputeSumExpF32KernelFma3;
|
||||
this->SQNBitGemmDispatch = &MlasSQNBitGemmDispatchAvx2;
|
||||
this->CastF16ToF32Kernel = &MlasCastF16ToF32KernelAvx2;
|
||||
this->CastF32ToF16Kernel = &MlasCastF32ToF16KernelAvx2;
|
||||
|
||||
|
||||
//
|
||||
// Check if the processor supports Hybrid core architecture.
|
||||
|
|
|
|||
|
|
@ -29,6 +29,51 @@ Abstract:
|
|||
#include "sqnbitgemm_m1_sym_kernel_avx2_int8_blklen32.h"
|
||||
#include "sqnbitgemm_m1_sym_kernel_avx2_int8_blklen64.h"
|
||||
|
||||
void
|
||||
MlasCastF16ToF32KernelAvx2(const unsigned short* src_fp16, float* dst_fp32, size_t size)
|
||||
{
|
||||
size_t i = 0;
|
||||
|
||||
// Process 16 elements at a time using AVX2
|
||||
for (; i + 15 < size; i += 16) {
|
||||
// Load 16 FP16 values into an AVX2 register
|
||||
__m256i fp16_values = _mm256_loadu_si256(reinterpret_cast<const __m256i*>(src_fp16 + i));
|
||||
|
||||
// Convert FP16 values to FP32
|
||||
__m256 fp32_values1 = _mm256_cvtph_ps(_mm256_castsi256_si128(fp16_values));
|
||||
__m256 fp32_values2 = _mm256_cvtph_ps(_mm256_extracti128_si256(fp16_values, 1));
|
||||
|
||||
// Store the converted FP32 values into the output vector
|
||||
_mm256_storeu_ps(dst_fp32 + i, fp32_values1);
|
||||
_mm256_storeu_ps(dst_fp32 + i + 8, fp32_values2);
|
||||
}
|
||||
|
||||
// Process any remaining elements
|
||||
const MLAS_FP16* fp16 = reinterpret_cast<const MLAS_FP16*>(src_fp16);
|
||||
for (; i < size; ++i) {
|
||||
dst_fp32[i] = fp16[i].ToFloat();
|
||||
}
|
||||
}
|
||||
|
||||
void
|
||||
MlasCastF32ToF16KernelAvx2(const float* src_fp32, unsigned short* dst_fp16, size_t size)
|
||||
{
|
||||
size_t i = 0;
|
||||
|
||||
// Process 8 elements at a time using AVX2
|
||||
for (; i + 8 <= size; i += 8) {
|
||||
__m256 fp32_chunk = _mm256_loadu_ps(&src_fp32[i]);
|
||||
__m128i fp16_chunk = _mm256_cvtps_ph(fp32_chunk, _MM_FROUND_TO_NEAREST_INT);
|
||||
_mm_storeu_si128(reinterpret_cast<__m128i*>(&dst_fp16[i]), fp16_chunk);
|
||||
}
|
||||
|
||||
// Process any remaining elements
|
||||
for (; i < size; ++i) {
|
||||
MLAS_FP16 fp16(src_fp32[i]);
|
||||
dst_fp16[i] = fp16.val;
|
||||
}
|
||||
}
|
||||
|
||||
MLAS_FORCEINLINE
|
||||
__m256
|
||||
load_float_n_avx2(const float* data, int n)
|
||||
|
|
|
|||
|
|
@ -258,7 +258,7 @@ struct TensorCaster<MLFloat16, float> {
|
|||
auto out_data = out.MutableData<float>();
|
||||
auto in_data = in.Data<MLFloat16>();
|
||||
const size_t shape_size = narrow<size_t>(shape.Size());
|
||||
MlasConvertHalfToFloatBuffer(&in_data[0].val, out_data, shape_size);
|
||||
MlasConvertHalfToFloatBuffer(in_data, out_data, shape_size);
|
||||
}
|
||||
};
|
||||
|
||||
|
|
|
|||
|
|
@ -262,8 +262,8 @@ void RunTest(const TestOptions& opts,
|
|||
|
||||
} // namespace
|
||||
|
||||
TEST(MatMulNBits, Float32) {
|
||||
// onnxruntime::profiling::Profiler::Profiler::Instance().StartProfiling<char>("profile.json");
|
||||
template <typename AType>
|
||||
void TestMatMulNBitsTyped() {
|
||||
for (auto M : {1, 2, 100}) {
|
||||
for (auto N : {/*2560, */ 1, 2, 32, 288}) {
|
||||
for (auto K : {/*2560, */ 16, 32, 64, 128, 256, 1024, 93, 1234}) {
|
||||
|
|
@ -276,30 +276,53 @@ TEST(MatMulNBits, Float32) {
|
|||
|
||||
if (base_opts.accuracy_level == 4) {
|
||||
base_opts.output_abs_error = 0.1f;
|
||||
} else {
|
||||
if constexpr (std::is_same<AType, MLFloat16>::value) {
|
||||
base_opts.output_abs_error = 0.01f;
|
||||
}
|
||||
}
|
||||
|
||||
{
|
||||
TestOptions opts = base_opts;
|
||||
RunTest<float>(opts);
|
||||
RunTest<AType>(opts);
|
||||
}
|
||||
|
||||
{
|
||||
TestOptions opts = base_opts;
|
||||
opts.has_zero_point = true;
|
||||
RunTest<float>(opts);
|
||||
RunTest<AType>(opts);
|
||||
}
|
||||
|
||||
#if !defined(ORT_NEURAL_SPEED) && !defined(USE_DML)
|
||||
{
|
||||
TestOptions opts = base_opts;
|
||||
opts.has_g_idx = true;
|
||||
RunTest<float>(opts);
|
||||
RunTest<AType>(opts);
|
||||
}
|
||||
|
||||
{
|
||||
TestOptions opts = base_opts;
|
||||
opts.has_g_idx = true;
|
||||
opts.has_bias = true;
|
||||
if constexpr (std::is_same<AType, float>::value) {
|
||||
if (opts.accuracy_level == 0 || opts.accuracy_level == 1) {
|
||||
// CI failure (not able to repro on either local machines):
|
||||
// M:100, N:288, K:1234, block_size:16, accuracy_level:0, has_zero_point:0, zp_is_4bit:1, has_g_idx:1, has_bias:1
|
||||
// The difference between cur_expected[i] and cur_actual[i] is 1.0401010513305664e-05, which exceeds tolerance,
|
||||
// tolerance evaluates to 1.006456386676291e-05.
|
||||
opts.output_abs_error = 0.0001f;
|
||||
}
|
||||
}
|
||||
// only enabled for CPU EP for now
|
||||
std::vector<std::unique_ptr<IExecutionProvider>> explicit_eps;
|
||||
explicit_eps.emplace_back(DefaultCpuExecutionProvider());
|
||||
RunTest<AType>(opts, std::move(explicit_eps));
|
||||
}
|
||||
|
||||
{
|
||||
TestOptions opts = base_opts;
|
||||
opts.has_zero_point = true, opts.zp_is_4bit = false;
|
||||
RunTest<float>(opts);
|
||||
RunTest<AType>(opts);
|
||||
}
|
||||
#endif // !defined(ORT_NEURAL_SPEED) && !defined(USE_DML)
|
||||
|
||||
|
|
@ -311,7 +334,7 @@ TEST(MatMulNBits, Float32) {
|
|||
std::vector<std::unique_ptr<IExecutionProvider>> explicit_eps;
|
||||
explicit_eps.emplace_back(DefaultCpuExecutionProvider());
|
||||
|
||||
RunTest<float>(opts, std::move(explicit_eps));
|
||||
RunTest<AType>(opts, std::move(explicit_eps));
|
||||
}
|
||||
}
|
||||
}
|
||||
|
|
@ -320,6 +343,21 @@ TEST(MatMulNBits, Float32) {
|
|||
}
|
||||
}
|
||||
|
||||
TEST(MatMulNBits, Float32) {
|
||||
// onnxruntime::profiling::Profiler::Profiler::Instance().StartProfiling<char>("profile.json");
|
||||
TestMatMulNBitsTyped<float>();
|
||||
}
|
||||
|
||||
#ifdef MLAS_TARGET_AMD64_IX86
|
||||
#if !defined(ORT_NEURAL_SPEED) && !defined(USE_DML)
|
||||
// Actual and expected difference is over 0.01 with DmlExecutionProvider.
|
||||
// Skip the tests instead of raising the tolerance to make is pass.
|
||||
TEST(MatMulNBits, Float16) {
|
||||
TestMatMulNBitsTyped<MLFloat16>();
|
||||
}
|
||||
#endif
|
||||
#endif
|
||||
|
||||
#if defined(USE_CUDA) || defined(USE_ROCM) || defined(USE_DML)
|
||||
|
||||
namespace {
|
||||
|
|
@ -367,7 +405,7 @@ void RunTest(int64_t M, int64_t N, int64_t K, int64_t block_size, int64_t accura
|
|||
}
|
||||
} // namespace
|
||||
|
||||
TEST(MatMulNBits, Float16) {
|
||||
TEST(MatMulNBits, Float16Cuda) {
|
||||
#if defined(USE_CUDA) || defined(USE_ROCM)
|
||||
auto has_gidx_options = {true, false};
|
||||
#else
|
||||
|
|
|
|||
Loading…
Reference in a new issue