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