diff --git a/cmake/onnxruntime_mlas.cmake b/cmake/onnxruntime_mlas.cmake
index b612b3ead4..e35c83ba45 100644
--- a/cmake/onnxruntime_mlas.cmake
+++ b/cmake/onnxruntime_mlas.cmake
@@ -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
diff --git a/docs/OperatorKernels.md b/docs/OperatorKernels.md
index d57394b3e7..121240e6e1 100644
--- a/docs/OperatorKernels.md
+++ b/docs/OperatorKernels.md
@@ -488,7 +488,7 @@ Do not modify directly.*
|MatMulFpQ4|*in* A:**T1**
*in* B:**T2**
*in* B_shape:**T3**
*out* Y:**T1**|1+|**T1** = tensor(float)
**T2** = tensor(uint8)
**T3** = tensor(int64)|
|MatMulInteger16|*in* A:**T1**
*in* B:**T2**
*out* Y:**T3**|1+|**T1** = tensor(int16)
**T2** = tensor(int16)
**T3** = tensor(int32)|
|MatMulIntegerToFloat|*in* A:**T1**
*in* B:**T2**
*in* a_scale:**T3**
*in* b_scale:**T3**
*in* a_zero_point:**T1**
*in* b_zero_point:**T2**
*in* bias:**T3**
*out* Y:**T3**|1+|**T1** = tensor(int8), tensor(uint8)
**T2** = tensor(int8), tensor(uint8)
**T3** = tensor(float)|
-|MatMulNBits|*in* A:**T1**
*in* B:**T2**
*in* scales:**T1**
*in* zero_points:**T3**
*in* g_idx:**T4**
*in* bias:**T1**
*out* Y:**T1**|1+|**T1** = tensor(float)
**T2** = tensor(uint8)
**T3** = tensor(float), tensor(uint8)
**T4** = tensor(int32)|
+|MatMulNBits|*in* A:**T1**
*in* B:**T2**
*in* scales:**T1**
*in* zero_points:**T3**
*in* g_idx:**T4**
*in* bias:**T1**
*out* Y:**T1**|1+|**T1** = tensor(float), tensor(float16)
**T2** = tensor(uint8)
**T3** = tensor(float), tensor(float16), tensor(uint8)
**T4** = tensor(int32)|
|MaxpoolWithMask|*in* X:**T**
*in* M:**tensor(int32)**
*out* Y:**T**|1+|**T** = tensor(float)|
|MultiHeadAttention|*in* query:**T**
*in* key:**T**
*in* value:**T**
*in* bias:**T**
*in* key_padding_mask:**M**
*in* attention_bias:**T**
*in* past_key:**T**
*in* past_value:**T**
*out* output:**T**
*out* present_key:**T**
*out* present_value:**T**|1+|**T** = tensor(float)|
|MurmurHash3|*in* X:**T1**
*out* Y:**T2**|1+|**T1** = tensor(double), tensor(float), tensor(int32), tensor(int64), tensor(string), tensor(uint32), tensor(uint64)
**T2** = tensor(int32), tensor(uint32)|
diff --git a/onnxruntime/contrib_ops/cpu/quantization/matmul_nbits.cc b/onnxruntime/contrib_ops/cpu/quantization/matmul_nbits.cc
index bf43aca73e..ccb779721d 100644
--- a/onnxruntime/contrib_ops/cpu/quantization/matmul_nbits.cc
+++ b/onnxruntime/contrib_ops/cpu/quantization/matmul_nbits.cc
@@ -146,8 +146,15 @@ class MatMulNBits final : public OpKernel {
bool all_constant_{false};
#endif // defined(ORT_NEURAL_SPEED)
+
+ template
+ 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(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();
- 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();
+ std::vector scales_v(static_cast(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();
+ 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();
@@ -274,9 +288,20 @@ Status MatMulNBits::UseSharedPrePackedBuffers(std::vector& prep
}
Status MatMulNBits::Compute(OpKernelContext* ctx) const {
+ const Tensor* a = ctx->Input(InputIndex::A);
+
+ if (IsATypeFloat16(*a)) {
+ return ComputeTyped(ctx);
+ } else {
+ return ComputeTyped(ctx);
+ }
+}
+
+template
+Status MatMulNBits::ComputeTyped(OpKernelContext* ctx) const {
concurrency::ThreadPool* thread_pool = ctx->GetOperatorThreadPool();
const Tensor* a = ctx->Input(InputIndex::A);
- const auto* a_data = a->Data();
+ const auto* a_data = a->Data();
TensorShape b_shape({static_cast(N_), static_cast(K_)});
MatMulComputeHelper helper;
@@ -289,7 +314,7 @@ Status MatMulNBits::Compute(OpKernelContext* ctx) const {
return Status::OK();
}
- auto* y_data = y->MutableData();
+ auto* y_data = y->MutableData();
const size_t batch_count = helper.OutputOffsets().size();
const size_t M = static_cast(helper.M());
@@ -297,9 +322,12 @@ Status MatMulNBits::Compute(OpKernelContext* ctx) const {
const size_t K = static_cast(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(InputIndex::zero_points);
const Tensor* bias = ctx->Input(InputIndex::bias);
- const auto* scales_data = scales->Data();
+ const auto* scales_data = scales->Data();
const auto* zero_points_data = zero_points == nullptr ? nullptr : zero_points->DataRaw();
- const auto* bias_data = bias == nullptr ? nullptr : bias->Data();
+ const auto* bias_data = bias == nullptr ? nullptr : bias->Data();
IAllocatorUniquePtr workspace{};
const size_t workspace_size = MlasSQNBitGemmBatchWorkspaceSize(
@@ -349,26 +377,64 @@ Status MatMulNBits::Compute(OpKernelContext* ctx) const {
workspace = IAllocator::MakeUniquePtr(allocator, workspace_size);
}
- InlinedVector 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(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::value) {
+ InlinedVector data(batch_count);
- return Status::OK();
+ AllocatorPtr allocator;
+ ORT_RETURN_IF_ERROR(ctx->GetTempSpaceAllocator(&allocator));
+
+ auto tmp_a_data_ptr = IAllocator::MakeUniquePtr(allocator, (size_t)(a->Shape().Size()));
+ MlasConvertHalfToFloatBuffer(a_data, tmp_a_data_ptr.get(), static_cast(a->Shape().Size()));
+
+ auto tmp_scales_data_ptr = IAllocator::MakeUniquePtr(allocator, (size_t)(scales->Shape().Size()));
+ MlasConvertHalfToFloatBuffer(scales_data, tmp_scales_data_ptr.get(), static_cast(scales->Shape().Size()));
+
+ std::vector 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 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(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 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(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(InputIndex::zero_points);
const Tensor* reorder_idx = ctx->Input(InputIndex::g_idx);
- const auto* scales_data = scales->Data();
+ const auto* scales_data = scales->Data();
+ const float* scales_data_;
+ std::vector scales_data_v;
+ if constexpr (std::is_same::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();
@@ -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(allocator, SafeInt(K_) * N_);
- if ((reorder_idx_data == nullptr) && (!zero_points || !zero_points->IsDataType())) {
+ if ((reorder_idx_data == nullptr) && (!zero_points || !zero_points->IsDataType())) {
// dequantize b, only 4b quantization is supported for now
MlasDequantizeBlockwise(
tmp_b_data_ptr.get(), // dequantized output
b_data, // quantized input
- scales_data, // quantization scales
+ scales_data_, // quantization scales
static_cast(zero_points_data), // quantization zero points
static_cast(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())) {
- DequantizeBlockwise(
+ if ((zero_points && zero_points->IsDataType())) {
+ DequantizeBlockwise(
tmp_b_data_ptr.get(), // dequantized output
b_data, // quantized input
- scales_data, // quantization scales
- static_cast(zero_points_data), // quantization zero points
+ scales_data_, // quantization scales
+ static_cast(zero_points_data), // quantization zero points
reorder_idx_data,
static_cast(block_size_), // quantization block size
column_wise_quant_, // columnwise quantization or row-wise
@@ -422,7 +498,7 @@ Status MatMulNBits::Compute(OpKernelContext* ctx) const {
DequantizeBlockwise(
tmp_b_data_ptr.get(), // dequantized output
b_data, // quantized input
- scales_data, // quantization scales
+ scales_data_, // quantization scales
static_cast(zero_points_data), // quantization zero points
reorder_idx_data,
static_cast(block_size_), // quantization block size
@@ -436,40 +512,80 @@ Status MatMulNBits::Compute(OpKernelContext* ctx) const {
auto tm_b_data_ptr_trans = IAllocator::MakeUniquePtr(allocator, SafeInt(K_) * N_);
MlasTranspose(tmp_b_data_ptr.get(), tm_b_data_ptr_trans.get(), N_, K_);
#endif
+ if constexpr (std::is_same::value) {
+ std::vector data(batch_count);
- std::vector 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(allocator, (size_t)(a->Shape().Size()));
+ MlasConvertHalfToFloatBuffer(a_data, tmp_a_data_ptr.get(), static_cast(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(InputIndex::bias);
- bias != nullptr) {
- gsl::span bias_span = bias->DataAsSpan();
- 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(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(InputIndex::bias);
+ bias != nullptr) {
+ auto tmp_bias_data_ptr = IAllocator::MakeUniquePtr(allocator, (size_t)(bias->Shape().Size()));
+ MlasConvertHalfToFloatBuffer(bias->Data(), tmp_bias_data_ptr.get(), static_cast(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(y->Shape().Size()));
+ return Status::OK();
+ } else {
+ std::vector 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(InputIndex::bias);
+ bias != nullptr) {
+ gsl::span bias_span = bias->DataAsSpan();
+ 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())
+ .TypeConstraint("T1", {DataTypeImpl::GetTensorType(), DataTypeImpl::GetTensorType()})
.TypeConstraint("T2", DataTypeImpl::GetTensorType())
- .TypeConstraint("T3", {DataTypeImpl::GetTensorType(), DataTypeImpl::GetTensorType()})
+ .TypeConstraint("T3", {DataTypeImpl::GetTensorType(), DataTypeImpl::GetTensorType(), DataTypeImpl::GetTensorType()})
.TypeConstraint("T4", DataTypeImpl::GetTensorType()),
MatMulNBits);
diff --git a/onnxruntime/contrib_ops/cpu/quantization/matmul_nbits_impl.cc b/onnxruntime/contrib_ops/cpu/quantization/matmul_nbits_impl.cc
index b28f3758f8..6a19a741c3 100644
--- a/onnxruntime/contrib_ops/cpu/quantization/matmul_nbits_impl.cc
+++ b/onnxruntime/contrib_ops/cpu/quantization/matmul_nbits_impl.cc
@@ -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) {
- zp_f = *(zero_points + n_idx * scales_shape_x + rid);
- } else {
+ if constexpr (std::is_same_v) {
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(n_idx) * static_cast(scales_shape_x) + static_cast(rid));
}
}
@@ -112,5 +112,10 @@ template void DequantizeBlockwise(
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* 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
diff --git a/onnxruntime/core/mlas/inc/mlas.h b/onnxruntime/core/mlas/inc/mlas.h
index 8b3156d77e..28ae64c4d5 100644
--- a/onnxruntime/core/mlas/inc/mlas.h
+++ b/onnxruntime/core/mlas/inc/mlas.h
@@ -20,6 +20,7 @@ Abstract:
#include
#include
#include
+#include
//
// 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
diff --git a/onnxruntime/core/mlas/lib/cast.cpp b/onnxruntime/core/mlas/lib/cast.cpp
index 24af4064bb..a6138e29bd 100644
--- a/onnxruntime/core/mlas/lib/cast.cpp
+++ b/onnxruntime/core/mlas/lib/cast.cpp
@@ -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(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(Destination), Count);
}
}
diff --git a/onnxruntime/core/mlas/lib/mlasi.h b/onnxruntime/core/mlas/lib/mlasi.h
index 6f5db766b7..8e8f46b8a1 100644
--- a/onnxruntime/core/mlas/lib/mlasi.h
+++ b/onnxruntime/core/mlas/lib/mlasi.h
@@ -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
diff --git a/onnxruntime/core/mlas/lib/platform.cpp b/onnxruntime/core/mlas/lib/platform.cpp
index 4cd7faaa9e..2b4d99800c 100644
--- a/onnxruntime/core/mlas/lib/platform.cpp
+++ b/onnxruntime/core/mlas/lib/platform.cpp
@@ -245,6 +245,7 @@ Return Value:
this->ConvDepthwiseS8S8Kernel = MlasConvDepthwiseKernel;
this->ConvDepthwiseS8U8Kernel = MlasConvDepthwiseKernel;
this->CastF16ToF32Kernel = nullptr;
+ this->CastF32ToF16Kernel = nullptr;
#if defined(MLAS_TARGET_AMD64_IX86)
@@ -387,6 +388,9 @@ Return Value:
this->ConvDepthwiseS8U8Kernel = MlasConvDepthwiseKernelAvx2;
this->ComputeSumExpF32Kernel = MlasComputeSumExpF32KernelFma3;
this->SQNBitGemmDispatch = &MlasSQNBitGemmDispatchAvx2;
+ this->CastF16ToF32Kernel = &MlasCastF16ToF32KernelAvx2;
+ this->CastF32ToF16Kernel = &MlasCastF32ToF16KernelAvx2;
+
//
// Check if the processor supports Hybrid core architecture.
diff --git a/onnxruntime/core/mlas/lib/sqnbitgemm_kernel_avx2.cpp b/onnxruntime/core/mlas/lib/sqnbitgemm_kernel_avx2.cpp
index 55d86bb9cc..baaa4ba1a3 100644
--- a/onnxruntime/core/mlas/lib/sqnbitgemm_kernel_avx2.cpp
+++ b/onnxruntime/core/mlas/lib/sqnbitgemm_kernel_avx2.cpp
@@ -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(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(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)
diff --git a/onnxruntime/core/providers/cpu/tensor/cast_op.cc b/onnxruntime/core/providers/cpu/tensor/cast_op.cc
index f2aaa75cad..35f3b12aeb 100644
--- a/onnxruntime/core/providers/cpu/tensor/cast_op.cc
+++ b/onnxruntime/core/providers/cpu/tensor/cast_op.cc
@@ -258,7 +258,7 @@ struct TensorCaster {
auto out_data = out.MutableData();
auto in_data = in.Data();
const size_t shape_size = narrow(shape.Size());
- MlasConvertHalfToFloatBuffer(&in_data[0].val, out_data, shape_size);
+ MlasConvertHalfToFloatBuffer(in_data, out_data, shape_size);
}
};
diff --git a/onnxruntime/test/contrib_ops/matmul_4bits_test.cc b/onnxruntime/test/contrib_ops/matmul_4bits_test.cc
index 548f24e8ac..fa7c6bce7c 100644
--- a/onnxruntime/test/contrib_ops/matmul_4bits_test.cc
+++ b/onnxruntime/test/contrib_ops/matmul_4bits_test.cc
@@ -262,8 +262,8 @@ void RunTest(const TestOptions& opts,
} // namespace
-TEST(MatMulNBits, Float32) {
- // onnxruntime::profiling::Profiler::Profiler::Instance().StartProfiling("profile.json");
+template
+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::value) {
+ base_opts.output_abs_error = 0.01f;
+ }
}
{
TestOptions opts = base_opts;
- RunTest(opts);
+ RunTest(opts);
}
{
TestOptions opts = base_opts;
opts.has_zero_point = true;
- RunTest(opts);
+ RunTest(opts);
}
#if !defined(ORT_NEURAL_SPEED) && !defined(USE_DML)
{
TestOptions opts = base_opts;
opts.has_g_idx = true;
- RunTest(opts);
+ RunTest(opts);
+ }
+
+ {
+ TestOptions opts = base_opts;
+ opts.has_g_idx = true;
+ opts.has_bias = true;
+ if constexpr (std::is_same::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> explicit_eps;
+ explicit_eps.emplace_back(DefaultCpuExecutionProvider());
+ RunTest(opts, std::move(explicit_eps));
}
{
TestOptions opts = base_opts;
opts.has_zero_point = true, opts.zp_is_4bit = false;
- RunTest(opts);
+ RunTest(opts);
}
#endif // !defined(ORT_NEURAL_SPEED) && !defined(USE_DML)
@@ -311,7 +334,7 @@ TEST(MatMulNBits, Float32) {
std::vector> explicit_eps;
explicit_eps.emplace_back(DefaultCpuExecutionProvider());
- RunTest(opts, std::move(explicit_eps));
+ RunTest(opts, std::move(explicit_eps));
}
}
}
@@ -320,6 +343,21 @@ TEST(MatMulNBits, Float32) {
}
}
+TEST(MatMulNBits, Float32) {
+ // onnxruntime::profiling::Profiler::Profiler::Instance().StartProfiling("profile.json");
+ TestMatMulNBitsTyped();
+}
+
+#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();
+}
+#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