Matmul_nbits kernel for mlas sqnbits to support Fp16 inputs (#21807)

This commit is contained in:
liqun Fu 2024-09-13 14:55:08 -07:00 committed by GitHub
parent 7e2c722459
commit a89bddd5c2
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
11 changed files with 348 additions and 123 deletions

View file

@ -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

View file

@ -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)|

View file

@ -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);

View file

@ -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

View file

@ -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

View file

@ -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);
}
}

View file

@ -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

View file

@ -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.

View file

@ -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)

View file

@ -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);
}
};

View file

@ -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