diff --git a/onnxruntime/contrib_ops/cpu/quantization/attention_quant.cc b/onnxruntime/contrib_ops/cpu/quantization/attention_quant.cc index 145276c73c..c6f313322e 100644 --- a/onnxruntime/contrib_ops/cpu/quantization/attention_quant.cc +++ b/onnxruntime/contrib_ops/cpu/quantization/attention_quant.cc @@ -110,14 +110,6 @@ Status QAttention::Compute(OpKernelContext* context) const { auto V = K + batch_size * sequence_length * hidden_size; T* QKV[3] = {Q, K, V}; - auto gemm_data_quant = allocator->Alloc(SafeInt(batch_size) * sequence_length * 3 * hidden_size * sizeof(int32_t)); - BufferUniquePtr gemm_buffer_quant(gemm_data_quant, BufferDeleter(allocator)); - - auto Q_quant = reinterpret_cast(gemm_data_quant); - auto K_quant = Q_quant + batch_size * sequence_length * hidden_size; - auto V_quant = K_quant + batch_size * sequence_length * hidden_size; - int32_t* QKV_quant[3] = {Q_quant, K_quant, V_quant}; - { const int loop_len = 3 * batch_size * num_heads_; const auto input_data = input->template Data(); @@ -134,7 +126,7 @@ Status QAttention::Compute(OpKernelContext* context) const { int input_offset = batch_index * sequence_length * hidden_size; int weights_offset = qkv_index * hidden_size + head_index * head_size; - int32_t* qkv_dest = QKV_quant[qkv_index]; + float* qkv_dest = QKV[qkv_index]; int qkv_offset = (batch_index * num_heads_ + head_index) * (sequence_length * head_size); // original transposed iteration @@ -152,19 +144,10 @@ Status QAttention::Compute(OpKernelContext* context) const { weight_zero_point, // weight zero point qkv_dest + qkv_offset, // C head_size, // ldc + &dequant_scale, // output scale + bias_data + weights_offset, // bias nullptr // use single-thread ); - - // dequantize and add bias - // broadcast 3NH -> (3.B.N.S.H) - const T* bias_src = bias_data + weights_offset; - int32_t* gemm_quant_src = QKV_quant[qkv_index] + qkv_offset; - T* data_dest = QKV[qkv_index] + qkv_offset; - for (int seq_index = 0; seq_index < sequence_length; seq_index++) { - for (int head_idx = 0; head_idx < head_size; head_idx++) { - *data_dest++ = *gemm_quant_src++ * dequant_scale + bias_src[head_idx]; - } - } } }); } diff --git a/onnxruntime/contrib_ops/cpu/quantization/dynamic_quantize_matmul.cc b/onnxruntime/contrib_ops/cpu/quantization/dynamic_quantize_matmul.cc index 2e7566e27b..6fbf3eed0a 100644 --- a/onnxruntime/contrib_ops/cpu/quantization/dynamic_quantize_matmul.cc +++ b/onnxruntime/contrib_ops/cpu/quantization/dynamic_quantize_matmul.cc @@ -69,7 +69,7 @@ Status DynamicQuantizeMatMul::Compute(OpKernelContext* ctx) const { } // calculate quantization parameter of a - const float* a_data = a->template Data(); + const auto* a_data = a->template Data(); int64_t num_of_elements = a->Shape().Size(); float a_scale; @@ -86,8 +86,12 @@ Status DynamicQuantizeMatMul::Compute(OpKernelContext* ctx) const { MatMulComputeHelper helper; ORT_RETURN_IF_ERROR(helper.Compute(a->Shape(), b->Shape())); - int32_t* matmul_output = static_cast(allocator->Alloc(SafeInt(helper.OutputShape().Size()) * sizeof(int32_t))); - BufferUniquePtr matmul_output_holder(matmul_output, BufferDeleter(allocator)); + const auto* b_data = b->template Data(); + + Tensor* y = ctx->Output(0, helper.OutputShape()); + auto* y_data = y->template MutableData(); + + const float multiplier = a_scale * b_scale; concurrency::ThreadPool* thread_pool = ctx->GetOperatorThreadPool(); for (size_t i = 0; i < helper.OutputOffsets().size(); i++) { @@ -97,21 +101,16 @@ Status DynamicQuantizeMatMul::Compute(OpKernelContext* ctx) const { a_data_quant + helper.LeftOffsets()[i], static_cast(helper.K()), a_zp, - b->template Data() + helper.RightOffsets()[i], + b_data + helper.RightOffsets()[i], static_cast(helper.N()), b_zp, - matmul_output + helper.OutputOffsets()[i], + y_data + helper.OutputOffsets()[i], static_cast(helper.N()), + &multiplier, + nullptr, thread_pool); } - Tensor* y = ctx->Output(0, helper.OutputShape()); - float* y_data = y->template MutableData(); - - float multiplier = a_scale * b_scale; - for (int64_t i = 0; i < helper.OutputShape().Size(); i++) { - y_data[i] = matmul_output[i] * multiplier; - } return Status::OK(); } diff --git a/onnxruntime/contrib_ops/cpu/quantization/dynamic_quantize_matmul.h b/onnxruntime/contrib_ops/cpu/quantization/dynamic_quantize_matmul.h index 28cd21bc00..ba2fce147b 100644 --- a/onnxruntime/contrib_ops/cpu/quantization/dynamic_quantize_matmul.h +++ b/onnxruntime/contrib_ops/cpu/quantization/dynamic_quantize_matmul.h @@ -12,9 +12,10 @@ namespace contrib { template class DynamicQuantizeMatMul final : public OpKernel { public: - DynamicQuantizeMatMul(const OpKernelInfo& info) : OpKernel(info) {} + DynamicQuantizeMatMul(const OpKernelInfo& info) : OpKernel(info) {} Status Compute(OpKernelContext* context) const override; }; + } // namespace contrib } // namespace onnxruntime diff --git a/onnxruntime/core/mlas/inc/mlas.h b/onnxruntime/core/mlas/inc/mlas.h index b342b712fe..fd5e62da1a 100644 --- a/onnxruntime/core/mlas/inc/mlas.h +++ b/onnxruntime/core/mlas/inc/mlas.h @@ -165,6 +165,26 @@ MlasGemm( MLAS_THREADPOOL* ThreadPool ); +template +void +MLASCALL +MlasGemm( + size_t M, + size_t N, + size_t K, + const AType* A, + size_t lda, + AType offa, + const BType* B, + size_t ldb, + BType offb, + float* C, + size_t ldc, + const float* Scale, + const float* Bias, + MLAS_THREADPOOL* ThreadPool + ); + // // Convolution routines. // diff --git a/onnxruntime/core/mlas/lib/mlasi.h b/onnxruntime/core/mlas/lib/mlasi.h index 323414eb85..f136debc8e 100644 --- a/onnxruntime/core/mlas/lib/mlasi.h +++ b/onnxruntime/core/mlas/lib/mlasi.h @@ -254,12 +254,7 @@ typedef MLAS_SGEMM_TRANSPOSE_PACKB_BLOCK_ROUTINE* PMLAS_SGEMM_TRANSPOSE_PACKB_BL typedef void (MLASCALL MLAS_GEMM_U8X8_OPERATION)( - const struct MLAS_GEMM_U8X8_WORK_BLOCK* WorkBlock, - size_t M, - size_t N, - const uint8_t* A, - const uint8_t* B, - int32_t* C + const struct MLAS_GEMM_U8X8_WORK_BLOCK* WorkBlock ); typedef MLAS_GEMM_U8X8_OPERATION* PMLAS_GEMM_U8X8_OPERATION; @@ -334,7 +329,7 @@ void size_t OutputCount, size_t OutputCountRightPad, const float* Bias, - unsigned Flags + unsigned KernelFlags ); typedef MLAS_CONV_FLOAT_KERNEL* PMLAS_CONV_FLOAT_KERNEL; @@ -357,7 +352,7 @@ void size_t OutputCount, size_t OutputCountRightPad, const float* Bias, - unsigned Flags + unsigned KernelFlags ); typedef MLAS_CONV_DEPTHWISE_FLOAT_KERNEL* PMLAS_CONV_DEPTHWISE_FLOAT_KERNEL; @@ -376,7 +371,7 @@ void size_t OutputStride, size_t OutputCount, const float* Bias, - unsigned Flags + unsigned KernelFlags ); typedef MLAS_CONV_POINTWISE_FLOAT_KERNEL* PMLAS_CONV_POINTWISE_FLOAT_KERNEL; diff --git a/onnxruntime/core/mlas/lib/qgemm.cpp b/onnxruntime/core/mlas/lib/qgemm.cpp index 6ed2d89a19..8c553728bb 100644 --- a/onnxruntime/core/mlas/lib/qgemm.cpp +++ b/onnxruntime/core/mlas/lib/qgemm.cpp @@ -34,21 +34,19 @@ struct MLAS_GEMM_U8X8_WORK_BLOCK { size_t ldb; int32_t* C; size_t ldc; + const float* Scale; + const float* BiasFloat; uint8_t offa; uint8_t offb; bool BTypeIsSigned; + bool CTypeIsFloat; }; template MLAS_FORCEINLINE void MlasGemmU8X8Operation( - const MLAS_GEMM_U8X8_WORK_BLOCK* WorkBlock, - size_t M, - size_t N, - const uint8_t* A, - const uint8_t* B, - int32_t* C + const MLAS_GEMM_U8X8_WORK_BLOCK* WorkBlock ) /*++ @@ -61,16 +59,6 @@ Arguments: WorkBlock - Supplies the structure containing the GEMM parameters. - M - Supplies the number of rows of matrix A and matrix C. - - N - Supplies the number of columns of matrix B and matrix C. - - A - Supplies the address of matrix A. - - B - Supplies the address of matrix B. - - C - Supplies the address of matrix C. - Return Value: None. @@ -83,6 +71,10 @@ Return Value: MLAS_DECLSPEC_ALIGN(int32_t RowSumVector[KernelType::StrideM], 64); MLAS_DECLSPEC_ALIGN(int32_t ColumnSumVector[KernelType::StrideN], 64); + const uint8_t* A = WorkBlock->A; + const uint8_t* B = WorkBlock->B; + int32_t* C = WorkBlock->C; + const size_t lda = WorkBlock->lda; const size_t ldb = WorkBlock->ldb; const size_t ldc = WorkBlock->ldc; @@ -103,6 +95,8 @@ Return Value: // Step through each slice of matrix B along the K dimension. // + const size_t M = WorkBlock->M; + const size_t N = WorkBlock->N; const size_t K = WorkBlock->K; size_t CountK; @@ -124,7 +118,7 @@ Return Value: // Copy a panel of matrix B to a local packed buffer. // - KernelType::CopyPackB(PanelB, B + n + k * ldb, ldb, CountN, CountK, + KernelType::CopyPackB(PanelB, B + n, ldb, CountN, CountK, ColumnSumVector, -offa, WorkBlock->BTypeIsSigned); // @@ -146,8 +140,8 @@ Return Value: // Copy a panel of matrix A to a local packed buffer. // - KernelType::CopyPackA(PanelA, A + k + m * lda, lda, CountM, - CountK, RowSumVector, -offb); + KernelType::CopyPackA(PanelA, A + m * lda, lda, CountM, CountK, + RowSumVector, -offb); // // Step through the rows of the local packed buffer. @@ -157,13 +151,20 @@ Return Value: int32_t* RowSums = RowSumVector; size_t RowsRemaining = CountM; + bool ZeroMode = (k == 0); + bool PostProcess = (k + CountK == K); + while (RowsRemaining > 0) { size_t RowsHandled; RowsHandled = KernelType::Kernel(pa, PanelB, c, PackedCountK, RowsRemaining, CountN, ldc, RowSums, ColumnSumVector, - DepthValue, k == 0); + DepthValue, ZeroMode); + + if (PostProcess && WorkBlock->CTypeIsFloat) { + KernelType::OutputFloat(WorkBlock, c, n, RowsHandled, CountN); + } c += ldc * RowsHandled; pa += KernelType::PackedK * PackedCountK * RowsHandled; @@ -172,6 +173,9 @@ Return Value: } } } + + A += CountK; + B += CountK * ldb; } } @@ -702,6 +706,112 @@ Return Value: } } +void +MlasGemmU8X8OutputFloatSse( + const MLAS_GEMM_U8X8_WORK_BLOCK* WorkBlock, + int32_t* C, + size_t StartN, + size_t CountM, + size_t CountN + ) +/*++ + +Routine Description: + + This routine is an inner kernel to compute matrix multiplication for a + single row. + +Arguments: + + WorkBlock - Supplies the structure containing the GEMM parameters. + + C - Supplies the address of matrix C. + + StartN - Supplies the starting column offset relative to the base of the + work block. This is used to offset into column vectors accessed via the + work block. + + CountM - Supplies the number of rows of the output matrix to process. + + CountN - Supplies the number of columns of the output matrix to process. + +Return Value: + + None. + +--*/ +{ + const size_t ldc = WorkBlock->ldc; + __m128 ScaleVector = _mm_load_ps1(WorkBlock->Scale); + + // + // Check if the optional bias vector was supplied. + // + + const float* BiasFloat = WorkBlock->BiasFloat; + + if (BiasFloat != nullptr) { + + BiasFloat += StartN; + + while (CountM-- > 0) { + + const float* bias = BiasFloat; + int32_t* c = C; + size_t n = CountN; + + while (n >= 4) { + + __m128 FloatVector = _mm_cvtepi32_ps(_mm_loadu_si128((__m128i*)c)); + FloatVector = _mm_mul_ps(FloatVector, ScaleVector); + FloatVector = _mm_add_ps(FloatVector, _mm_loadu_ps(bias)); + _mm_storeu_ps((float*)c, FloatVector); + + bias += 4; + c += 4; + n -= 4; + } + + for (size_t offset = 0; offset < n; offset++) { + + __m128 FloatVector = _mm_set_ss(float(c[offset])); + FloatVector = _mm_mul_ss(FloatVector, ScaleVector); + FloatVector = _mm_add_ss(FloatVector, _mm_load_ss(&bias[offset])); + _mm_store_ss((float*)&c[offset], FloatVector); + } + + C += ldc; + } + + } else { + + while (CountM-- > 0) { + + int32_t* c = C; + size_t n = CountN; + + while (n >= 4) { + + __m128 FloatVector = _mm_cvtepi32_ps(_mm_loadu_si128((__m128i*)c)); + FloatVector = _mm_mul_ps(FloatVector, ScaleVector); + _mm_storeu_ps((float*)c, FloatVector); + + c += 4; + n -= 4; + } + + for (size_t offset = 0; offset < n; offset++) { + + __m128 FloatVector = _mm_set_ss((float)c[offset]); + FloatVector = _mm_mul_ss(FloatVector, ScaleVector); + _mm_store_ss((float*)&c[offset], FloatVector); + } + + C += ldc; + } + } +} + struct MLAS_GEMM_U8X8_KERNEL_SSE { typedef int16_t PackedAType; @@ -772,6 +882,20 @@ struct MLAS_GEMM_U8X8_KERNEL_SSE return 1; } + + MLAS_FORCEINLINE + static + void + OutputFloat( + const MLAS_GEMM_U8X8_WORK_BLOCK* WorkBlock, + int32_t* C, + size_t StartN, + size_t CountM, + size_t CountN + ) + { + MlasGemmU8X8OutputFloatSse(WorkBlock, C, StartN, CountM, CountN); + } }; constexpr size_t MLAS_GEMM_U8X8_KERNEL_SSE::PackedK; @@ -782,12 +906,7 @@ constexpr size_t MLAS_GEMM_U8X8_KERNEL_SSE::StrideK; void MLASCALL MlasGemmU8X8OperationSse( - const MLAS_GEMM_U8X8_WORK_BLOCK* WorkBlock, - size_t M, - size_t N, - const uint8_t* A, - const uint8_t* B, - int32_t* C + const MLAS_GEMM_U8X8_WORK_BLOCK* WorkBlock ) /*++ @@ -802,23 +921,13 @@ Arguments: WorkBlock - Supplies the structure containing the GEMM parameters. - M - Supplies the number of rows of matrix A and matrix C. - - N - Supplies the number of columns of matrix B and matrix C. - - A - Supplies the address of matrix A. - - B - Supplies the address of matrix B. - - C - Supplies the address of matrix C. - Return Value: None. --*/ { - return MlasGemmU8X8Operation(WorkBlock, M, N, A, B, C); + return MlasGemmU8X8Operation(WorkBlock); } #endif @@ -953,6 +1062,20 @@ struct MLAS_GEMM_U8S8_KERNEL_AVX2 return MlasPlatform.GemmU8S8Kernel(A, B, C, PackedCountK, CountM, CountN, ldc, RowSumVector, ColumnSumVector, DepthValue, ZeroMode); } + + MLAS_FORCEINLINE + static + void + OutputFloat( + const MLAS_GEMM_U8X8_WORK_BLOCK* WorkBlock, + int32_t* C, + size_t StartN, + size_t CountM, + size_t CountN + ) + { + MlasGemmU8X8OutputFloatSse(WorkBlock, C, StartN, CountM, CountN); + } }; constexpr size_t MLAS_GEMM_U8S8_KERNEL_AVX2::PackedK; @@ -963,12 +1086,7 @@ constexpr size_t MLAS_GEMM_U8S8_KERNEL_AVX2::StrideK; void MLASCALL MlasGemmU8S8OperationAvx2( - const MLAS_GEMM_U8X8_WORK_BLOCK* WorkBlock, - size_t M, - size_t N, - const uint8_t* A, - const uint8_t* B, - int32_t* C + const MLAS_GEMM_U8X8_WORK_BLOCK* WorkBlock ) /*++ @@ -983,31 +1101,23 @@ Arguments: WorkBlock - Supplies the structure containing the GEMM parameters. - M - Supplies the number of rows of matrix A and matrix C. - - N - Supplies the number of columns of matrix B and matrix C. - - A - Supplies the address of matrix A. - - B - Supplies the address of matrix B. - - C - Supplies the address of matrix C. - Return Value: None. --*/ { - if (M == 1 && WorkBlock->offa == 0 && WorkBlock->offb == 0) { + if ((WorkBlock->M == 1) && WorkBlock->BTypeIsSigned && !WorkBlock->CTypeIsFloat && + (WorkBlock->offa == 0) && (WorkBlock->offb == 0)) { if (MlasPlatform.GemvU8S8Kernel != nullptr) { - MlasPlatform.GemvU8S8Kernel(A, B, C, WorkBlock->K, N, WorkBlock->ldb); + MlasPlatform.GemvU8S8Kernel(WorkBlock->A, WorkBlock->B, WorkBlock->C, + WorkBlock->K, WorkBlock->N, WorkBlock->ldb); return; } } - return MlasGemmU8X8Operation(WorkBlock, M, N, A, B, C); + return MlasGemmU8X8Operation(WorkBlock); } struct MLAS_GEMM_U8U8_KERNEL_AVX2 @@ -1076,6 +1186,20 @@ struct MLAS_GEMM_U8U8_KERNEL_AVX2 return MlasPlatform.GemmU8U8Kernel(A, B, C, PackedCountK, CountM, CountN, ldc, RowSumVector, ColumnSumVector, DepthValue, ZeroMode); } + + MLAS_FORCEINLINE + static + void + OutputFloat( + const MLAS_GEMM_U8X8_WORK_BLOCK* WorkBlock, + int32_t* C, + size_t StartN, + size_t CountM, + size_t CountN + ) + { + MlasGemmU8X8OutputFloatSse(WorkBlock, C, StartN, CountM, CountN); + } }; constexpr size_t MLAS_GEMM_U8U8_KERNEL_AVX2::PackedK; @@ -1086,12 +1210,7 @@ constexpr size_t MLAS_GEMM_U8U8_KERNEL_AVX2::StrideK; void MLASCALL MlasGemmU8U8OperationAvx2( - const MLAS_GEMM_U8X8_WORK_BLOCK* WorkBlock, - size_t M, - size_t N, - const uint8_t* A, - const uint8_t* B, - int32_t* C + const MLAS_GEMM_U8X8_WORK_BLOCK* WorkBlock ) /*++ @@ -1106,23 +1225,13 @@ Arguments: WorkBlock - Supplies the structure containing the GEMM parameters. - M - Supplies the number of rows of matrix A and matrix C. - - N - Supplies the number of columns of matrix B and matrix C. - - A - Supplies the address of matrix A. - - B - Supplies the address of matrix B. - - C - Supplies the address of matrix C. - Return Value: None. --*/ { - return MlasGemmU8X8Operation(WorkBlock, M, N, A, B, C); + return MlasGemmU8X8Operation(WorkBlock); } #endif @@ -1195,23 +1304,29 @@ Return Value: // Dispatch the partitioned operation. // - const uint8_t* a = WorkBlock->A + m * WorkBlock->lda; - const uint8_t* b = WorkBlock->B + n; - int32_t* c = WorkBlock->C + n + m * WorkBlock->ldc; + MLAS_GEMM_U8X8_WORK_BLOCK LocalWorkBlock; - PMLAS_GEMM_U8X8_OPERATION GemmU8X8Operation; + memcpy(&LocalWorkBlock, WorkBlock, sizeof(MLAS_GEMM_U8X8_WORK_BLOCK)); + + LocalWorkBlock.M = CountM; + LocalWorkBlock.N = CountN; + LocalWorkBlock.A += m * LocalWorkBlock.lda; + LocalWorkBlock.B += n; + LocalWorkBlock.C += m * LocalWorkBlock.ldc + n; + + if (LocalWorkBlock.BiasFloat != nullptr) { + LocalWorkBlock.BiasFloat += n; + } #if defined(MLAS_TARGET_AMD64) if (WorkBlock->BTypeIsSigned) { - GemmU8X8Operation = MlasPlatform.GemmU8S8Operation; + MlasPlatform.GemmU8S8Operation(&LocalWorkBlock); } else { - GemmU8X8Operation = MlasPlatform.GemmU8U8Operation; + MlasPlatform.GemmU8U8Operation(&LocalWorkBlock); } #else - GemmU8X8Operation = MlasGemmU8X8OperationSse; + MlasGemmU8X8OperationSse(&LocalWorkBlock); #endif - - GemmU8X8Operation(WorkBlock, CountM, CountN, a, b, c); } void @@ -1360,6 +1475,8 @@ Return Value: // Capture the GEMM parameters to the work block. // + memset(&WorkBlock, 0, sizeof(MLAS_GEMM_U8X8_WORK_BLOCK)); + WorkBlock.M = M; WorkBlock.N = N; WorkBlock.K = K; @@ -1416,4 +1533,135 @@ MlasGemm( MLAS_THREADPOOL* ThreadPool ); +template +void +MLASCALL +MlasGemm( + size_t M, + size_t N, + size_t K, + const AType* A, + size_t lda, + AType offa, + const BType* B, + size_t ldb, + BType offb, + float* C, + size_t ldc, + const float* Scale, + const float* Bias, + MLAS_THREADPOOL* ThreadPool + ) +/*++ + +Routine Description: + + This module implements the quantized integer matrix/matrix multiply + operation (QGEMM). + +Arguments: + + M - Supplies the number of rows of matrix A and matrix C. + + N - Supplies the number of columns of matrix B and matrix C. + + K - Supplies the number of columns of matrix A and the number of rows of + matrix B. + + A - Supplies the address of matrix A. + + lda - Supplies the first dimension of matrix A. + + offa - Supplies the zero point offset of matrix A. + + B - Supplies the address of matrix B. + + ldb - Supplies the first dimension of matrix B. + + offb - Supplies the zero point offset of matrix B. + + C - Supplies the address of matrix C. + + ldc - Supplies the first dimension of matrix C. + + ThreadPool - Supplies the thread pool object to use, else nullptr if the + base library threading support should be used. + +Return Value: + + None. + +--*/ +{ + MLAS_GEMM_U8X8_WORK_BLOCK WorkBlock; + + // + // Capture the GEMM parameters to the work block. + // + + memset(&WorkBlock, 0, sizeof(MLAS_GEMM_U8X8_WORK_BLOCK)); + + WorkBlock.M = M; + WorkBlock.N = N; + WorkBlock.K = K; + WorkBlock.A = A; + WorkBlock.lda = lda; + WorkBlock.B = (const uint8_t*)B; + WorkBlock.ldb = ldb; + WorkBlock.C = (int32_t*)C; + WorkBlock.ldc = ldc; + WorkBlock.Scale = Scale; + WorkBlock.BiasFloat = Bias; + WorkBlock.offa = offa; + WorkBlock.offb = offb; + WorkBlock.BTypeIsSigned = std::is_signed::value; + WorkBlock.CTypeIsFloat = true; + + // + // Schedule the operation across a set of worker threads. + // + + MlasGemmU8X8Schedule(&WorkBlock, ThreadPool); +} + +template +void +MLASCALL +MlasGemm( + size_t M, + size_t N, + size_t K, + const uint8_t* A, + size_t lda, + uint8_t offa, + const int8_t* B, + size_t ldb, + int8_t offb, + float* C, + size_t ldc, + const float* Scale, + const float* Bias, + MLAS_THREADPOOL* ThreadPool + ); + +template +void +MLASCALL +MlasGemm( + size_t M, + size_t N, + size_t K, + const uint8_t* A, + size_t lda, + uint8_t offa, + const uint8_t* B, + size_t ldb, + uint8_t offb, + float* C, + size_t ldc, + const float* Scale, + const float* Bias, + MLAS_THREADPOOL* ThreadPool + ); + #endif diff --git a/onnxruntime/core/util/qmath.cc b/onnxruntime/core/util/qmath.cc index 166793486d..67f84b32b4 100644 --- a/onnxruntime/core/util/qmath.cc +++ b/onnxruntime/core/util/qmath.cc @@ -69,4 +69,73 @@ void QGemm( #endif } +template +void QGemm( + int M, + int N, + int K, + const LeftScalar* lhs_data, + int lda, + const LeftScalar lhs_offset, + const RightScalar* rhs_data, + int ldb, + const RightScalar rhs_offset, + float* result_data, + int ldc, + const float* result_scale, + const float* bias, + concurrency::ThreadPool* thread_pool) { +#ifdef MLAS_SUPPORTS_GEMM_U8X8 + MlasGemm(M, N, K, lhs_data, lda, lhs_offset, rhs_data, ldb, rhs_offset, result_data, ldc, result_scale, bias, thread_pool); +#else + QGemm(M, N, K, lhs_data, lda, lhs_offset, rhs_data, ldb, rhs_offset, reinterpret_cast(result_data), ldc, thread_pool); + for (int m = 0; m < M; m++) { + if (bias != nullptr) { + for (int n = 0; n < N; n++) { + result_data[n] = static_cast(reinterpret_cast(result_data)[n]) * result_scale[0] + bias[n]; + } + } else { + for (int n = 0; n < N; n++) { + result_data[n] = static_cast(reinterpret_cast(result_data)[n]) * result_scale[0]; + } + } + result_data += ldc; + } +#endif +} + +template +void QGemm( + int M, + int N, + int K, + const uint8_t* lhs_data, + int lda, + const uint8_t lhs_offset, + const int8_t* rhs_data, + int ldb, + const int8_t rhs_offset, + float* result_data, + int ldc, + const float* result_scale, + const float* bias, + concurrency::ThreadPool* thread_pool); + +template +void QGemm( + int M, + int N, + int K, + const uint8_t* lhs_data, + int lda, + const uint8_t lhs_offset, + const uint8_t* rhs_data, + int ldb, + const uint8_t rhs_offset, + float* result_data, + int ldc, + const float* result_scale, + const float* bias, + concurrency::ThreadPool* thread_pool); + } // namespace onnxruntime diff --git a/onnxruntime/core/util/qmath.h b/onnxruntime/core/util/qmath.h index b79775b3f8..9d333023d5 100644 --- a/onnxruntime/core/util/qmath.h +++ b/onnxruntime/core/util/qmath.h @@ -29,6 +29,23 @@ void QGemm( int ldc, concurrency::ThreadPool* thread_pool); +template +void QGemm( + int M, + int N, + int K, + const LeftScalar* lhs_data, + int lda, + const LeftScalar lhs_offset, + const RightScalar* rhs_data, + int ldb, + const RightScalar rhs_offset, + float* result_data, + int ldc, + const float* result_scale, + const float* bias, + concurrency::ThreadPool* thread_pool); + inline float RoundHalfToEven(float input) { std::fesetround(FE_TONEAREST); auto result = std::nearbyintf(input); diff --git a/onnxruntime/test/contrib_ops/dynamic_quantize_matmul_test.cc b/onnxruntime/test/contrib_ops/dynamic_quantize_matmul_test.cc index 2b4be7cc8b..839278b8a7 100644 --- a/onnxruntime/test/contrib_ops/dynamic_quantize_matmul_test.cc +++ b/onnxruntime/test/contrib_ops/dynamic_quantize_matmul_test.cc @@ -7,6 +7,7 @@ #include "test/framework/test_utils.h" #include "test/providers/provider_test_utils.h" #include "test/util/include/default_providers.h" +#include "core/util/qmath.h" #include #include @@ -146,6 +147,7 @@ class DynamicQuantizeMatMulOpTester : public OpTester { }; TEST(DynamicQuantizeMatMul, Int8_test) { +#ifdef MLAS_SUPPORTS_GEMM_U8X8 std::vector A_dims{4, 128}; std::vector B_dims{128, 128}; std::vector Y_dims{4, 128}; @@ -155,6 +157,7 @@ TEST(DynamicQuantizeMatMul, Int8_test) { Y_dims, "testdata/dynamic_quantize_matmul_int8.onnx"); test.Run(); +#endif } TEST(DynamicQuantizeMatMul, UInt8_test) { diff --git a/onnxruntime/test/mlas/unittest.cpp b/onnxruntime/test/mlas/unittest.cpp index ab5677b85e..f145e6eac3 100644 --- a/onnxruntime/test/mlas/unittest.cpp +++ b/onnxruntime/test/mlas/unittest.cpp @@ -47,7 +47,7 @@ Abstract: MLAS_THREADPOOL* threadpool = nullptr; -template +template class MatrixGuardBuffer { public: @@ -202,7 +202,7 @@ public: } }; -template +template class MlasFgemmTest : public MlasTestBase { private: @@ -254,6 +254,7 @@ private: // Sensitive to comparing positive/negative zero. if (C[f] != CReference[f]) { printf("mismatch TransA=%d, TransB=%d, M=%zd, N=%zd, K=%zd, alpha=%f, beta=%f %f %f!\n", TransA, TransB, M, N, K, alpha, beta, float(C[f]), float(CReference[f])); + break; } } } @@ -467,8 +468,11 @@ public: #ifdef MLAS_HAS_QGEMM_U8X8 -template -class MlasQgemmU8X8Test : public MlasTestBase +template +class MlasQgemmU8X8Test; + +template +class MlasQgemmU8X8Test : public MlasTestBase { private: void @@ -644,6 +648,124 @@ public: } }; +template +class MlasQgemmU8X8Test : public MlasTestBase +{ +private: + void + Test( + size_t M, + size_t N, + size_t K, + uint8_t offa, + uint8_t offb + ) + { + const uint8_t* A = BufferA.GetBuffer(K * M); + const xint8_t* B = BufferB.GetBuffer(N * K); + float* C = BufferC.GetBuffer(N * M); + float* CReference = BufferCReference.GetBuffer(N * M); + const float* Bias = BufferBias.GetBuffer(N); + + const float AScale = 0.5f; + float* AFloat = BufferAFloat.GetBuffer(K * M); + DequantizeLinear(A, AFloat, K * M, AScale, offa); + + const float BScale = 0.25f; + float* BFloat = BufferBFloat.GetBuffer(N * K); + DequantizeLinear(B, BFloat, N * K, BScale, xint8_t(offb)); + + const float CScale = AScale * BScale; + + Test(M, N, K, A, AFloat, K, offa, B, BFloat, N, xint8_t(offb), C, CReference, N, CScale, nullptr); + Test(M, N, K, A, AFloat, K, offa, B, BFloat, N, xint8_t(offb), C, CReference, N, CScale, Bias); + } + + void + Test( + size_t M, + size_t N, + size_t K, + const uint8_t* A, + const float* AFloat, + size_t lda, + uint8_t offa, + const xint8_t* B, + const float* BFloat, + size_t ldb, + xint8_t offb, + float* C, + float* CReference, + size_t ldc, + float CScale, + const float* Bias + ) + { + MlasGemm(CblasNoTrans, CblasNoTrans, M, N, K, 1.0f, AFloat, lda, BFloat, ldb, 0.0f, CReference, ldc, threadpool); + + if (Bias != nullptr) { + for (size_t m = 0; m < M; m++) { + for (size_t n = 0; n < N; n++) { + CReference[m * ldc + n] += Bias[n]; + } + } + } + + MlasGemm(M, N, K, A, lda, offa, B, ldb, offb, C, ldc, &CScale, Bias, threadpool); + + for (size_t f = 0; f < M * N; f++) { + // Sensitive to comparing positive/negative zero. + if (C[f] != CReference[f]) { + printf("mismatch M=%zd, N=%zd, K=%zd, offa=%d, offb=%d! %f %f\n", M, N, K, (int)offa, (int)offb, C[f], CReference[f]); + } + } + } + + template + void + DequantizeLinear( + const qint8_t* Input, + float* Output, + size_t N, + float scale, + qint8_t offset + ) + { + for (size_t n = 0; n < N; n++) { + Output[n] = float((int32_t(Input[n]) - offset)) * scale; + } + } + + MatrixGuardBuffer BufferA; + MatrixGuardBuffer BufferB; + MatrixGuardBuffer BufferAFloat; + MatrixGuardBuffer BufferBFloat; + MatrixGuardBuffer BufferC; + MatrixGuardBuffer BufferCReference; + MatrixGuardBuffer BufferBias; + +public: + void + ExecuteShort( + void + ) override + { + for (size_t b = 1; b < 16; b++) { + Test(b, b, b, 34, 46); + } + for (size_t b = 16; b <= 256; b <<= 1) { + Test(b, b, b, 15, 191); + } + for (size_t b = 256; b < 320; b += 32) { + Test(b, b, b, 223, 73); + } + for (size_t b = 1; b < 96; b++) { + Test(1, b, 32, 0, 0); + } + Test(1024, 1024, 256, 13, 15); + } +}; + #endif class MlasConv2DTest : public MlasTestBase @@ -2254,9 +2376,14 @@ RunThreadedTests( #endif #ifdef MLAS_HAS_QGEMM_U8X8 - printf("QGEMM tests.\n"); - onnxruntime::make_unique>()->ExecuteShort(); - onnxruntime::make_unique>()->ExecuteShort(); + printf("QGEMM U8S8=int32_t tests.\n"); + onnxruntime::make_unique>()->ExecuteShort(); + printf("QGEMM U8U8=int32_t tests.\n"); + onnxruntime::make_unique>()->ExecuteShort(); + printf("QGEMM U8S8=float tests.\n"); + onnxruntime::make_unique>()->ExecuteShort(); + printf("QGEMM U8U8=float tests.\n"); + onnxruntime::make_unique>()->ExecuteShort(); #endif printf("Conv2D tests.\n");