diff --git a/onnxruntime/core/mlas/lib/amd64/LogisticKernelFma3.asm b/onnxruntime/core/mlas/lib/amd64/LogisticKernelFma3.asm index a51be4b7e6..e50a99baf1 100644 --- a/onnxruntime/core/mlas/lib/amd64/LogisticKernelFma3.asm +++ b/onnxruntime/core/mlas/lib/amd64/LogisticKernelFma3.asm @@ -19,61 +19,12 @@ .xlist INCLUDE mlasi.inc +INCLUDE TransKernelCommon.inc .list - EXTERN MlasMaskMoveAvx:NEAR + EXTERN MlasMaskMoveTableAvx:NEAR EXTERN MlasLogisticConstants:NEAR -; -; Structure layout for the logistic constants block. -; - -LogisticConstants STRUCT - - LowerRange DWORD ? - UpperRange DWORD ? - alpha_9 DWORD ? - alpha_7 DWORD ? - alpha_5 DWORD ? - alpha_3 DWORD ? - alpha_1 DWORD ? - beta_10 DWORD ? - beta_8 DWORD ? - beta_6 DWORD ? - beta_4 DWORD ? - beta_2 DWORD ? - beta_0 DWORD ? - one_half DWORD ? - -LogisticConstants ENDS - -; -; Stack frame layout for the logistic kernel. -; - -LogisticKernelFrame STRUCT - - SavedXmm6 OWORD ? - SavedXmm7 OWORD ? - SavedXmm8 OWORD ? - SavedXmm9 OWORD ? - SavedXmm10 OWORD ? - SavedXmm11 OWORD ? - SavedXmm12 OWORD ? - SavedXmm13 OWORD ? - SavedXmm14 OWORD ? - SavedXmm15 OWORD ? - Padding0 QWORD ? - Padding1 QWORD ? - CountN QWORD ? - ReturnAddress QWORD ? - PreviousP1Home QWORD ? - PreviousP2Home QWORD ? - PreviousP3Home QWORD ? - PreviousP4Home QWORD ? - -LogisticKernelFrame ENDS - ;++ ; ; Routine Description: @@ -94,20 +45,20 @@ LogisticKernelFrame ENDS ; ;-- - NESTED_ENTRY MlasLogisticKernelFma3, _TEXT + NESTED_ENTRY MlasComputeLogisticF32KernelFma3, _TEXT - alloc_stack (LogisticKernelFrame.ReturnAddress) + alloc_stack (TransKernelFrame.ReturnAddress) - save_xmm128 xmm6,LogisticKernelFrame.SavedXmm6 - save_xmm128 xmm7,LogisticKernelFrame.SavedXmm7 - save_xmm128 xmm8,LogisticKernelFrame.SavedXmm8 - save_xmm128 xmm9,LogisticKernelFrame.SavedXmm9 - save_xmm128 xmm10,LogisticKernelFrame.SavedXmm10 - save_xmm128 xmm11,LogisticKernelFrame.SavedXmm11 - save_xmm128 xmm12,LogisticKernelFrame.SavedXmm12 - save_xmm128 xmm13,LogisticKernelFrame.SavedXmm13 - save_xmm128 xmm14,LogisticKernelFrame.SavedXmm14 - save_xmm128 xmm15,LogisticKernelFrame.SavedXmm15 + save_xmm128 xmm6,TransKernelFrame.SavedXmm6 + save_xmm128 xmm7,TransKernelFrame.SavedXmm7 + save_xmm128 xmm8,TransKernelFrame.SavedXmm8 + save_xmm128 xmm9,TransKernelFrame.SavedXmm9 + save_xmm128 xmm10,TransKernelFrame.SavedXmm10 + save_xmm128 xmm11,TransKernelFrame.SavedXmm11 + save_xmm128 xmm12,TransKernelFrame.SavedXmm12 + save_xmm128 xmm13,TransKernelFrame.SavedXmm13 + save_xmm128 xmm14,TransKernelFrame.SavedXmm14 + save_xmm128 xmm15,TransKernelFrame.SavedXmm15 END_PROLOGUE @@ -158,9 +109,9 @@ ComputeLogisticBy8Loop: ProcessRemainingCount: add r8,8 ; correct for over-subtract above jz ExitKernel - mov DWORD PTR LogisticKernelFrame.CountN[rsp],r8d - vbroadcastss ymm2,DWORD PTR LogisticKernelFrame.CountN[rsp] - vpcmpgtd ymm2,ymm2,YMMWORD PTR [MlasMaskMoveAvx] + neg r8 + lea r10,MlasMaskMoveTableAvx+8*4 + vmovups ymm2,YMMWORD PTR [r10+r8*4] vmaskmovps ymm0,ymm2,YMMWORD PTR [rcx] vmaxps ymm0,ymm4,ymm0 ; clamp lower bound vminps ymm0,ymm5,ymm0 ; clamp upper bound @@ -185,22 +136,22 @@ ProcessRemainingCount: ExitKernel: vzeroupper - movaps xmm6,LogisticKernelFrame.SavedXmm6[rsp] - movaps xmm7,LogisticKernelFrame.SavedXmm7[rsp] - movaps xmm8,LogisticKernelFrame.SavedXmm8[rsp] - movaps xmm9,LogisticKernelFrame.SavedXmm9[rsp] - movaps xmm10,LogisticKernelFrame.SavedXmm10[rsp] - movaps xmm11,LogisticKernelFrame.SavedXmm11[rsp] - movaps xmm12,LogisticKernelFrame.SavedXmm12[rsp] - movaps xmm13,LogisticKernelFrame.SavedXmm13[rsp] - movaps xmm14,LogisticKernelFrame.SavedXmm14[rsp] - movaps xmm15,LogisticKernelFrame.SavedXmm15[rsp] - add rsp,(LogisticKernelFrame.ReturnAddress) + movaps xmm6,TransKernelFrame.SavedXmm6[rsp] + movaps xmm7,TransKernelFrame.SavedXmm7[rsp] + movaps xmm8,TransKernelFrame.SavedXmm8[rsp] + movaps xmm9,TransKernelFrame.SavedXmm9[rsp] + movaps xmm10,TransKernelFrame.SavedXmm10[rsp] + movaps xmm11,TransKernelFrame.SavedXmm11[rsp] + movaps xmm12,TransKernelFrame.SavedXmm12[rsp] + movaps xmm13,TransKernelFrame.SavedXmm13[rsp] + movaps xmm14,TransKernelFrame.SavedXmm14[rsp] + movaps xmm15,TransKernelFrame.SavedXmm15[rsp] + add rsp,(TransKernelFrame.ReturnAddress) BEGIN_EPILOGUE ret - NESTED_END MlasLogisticKernelFma3, _TEXT + NESTED_END MlasComputeLogisticF32KernelFma3, _TEXT END diff --git a/onnxruntime/core/mlas/lib/amd64/TanhKernelFma3.asm b/onnxruntime/core/mlas/lib/amd64/TanhKernelFma3.asm index 6cb3c55af2..6d94d533d7 100644 --- a/onnxruntime/core/mlas/lib/amd64/TanhKernelFma3.asm +++ b/onnxruntime/core/mlas/lib/amd64/TanhKernelFma3.asm @@ -19,60 +19,12 @@ .xlist INCLUDE mlasi.inc +INCLUDE TransKernelCommon.inc .list - EXTERN MlasMaskMoveAvx:NEAR + EXTERN MlasMaskMoveTableAvx:NEAR EXTERN MlasTanhConstants:NEAR -; -; Structure layout for the tanh constants block. -; - -TanhConstants STRUCT - - LowerRange DWORD ? - UpperRange DWORD ? - alpha_13 DWORD ? - alpha_11 DWORD ? - alpha_9 DWORD ? - alpha_7 DWORD ? - alpha_5 DWORD ? - alpha_3 DWORD ? - alpha_1 DWORD ? - beta_6 DWORD ? - beta_4 DWORD ? - beta_2 DWORD ? - beta_0 DWORD ? - -TanhConstants ENDS - -; -; Stack frame layout for the tanh kernel. -; - -TanhKernelFrame STRUCT - - SavedXmm6 OWORD ? - SavedXmm7 OWORD ? - SavedXmm8 OWORD ? - SavedXmm9 OWORD ? - SavedXmm10 OWORD ? - SavedXmm11 OWORD ? - SavedXmm12 OWORD ? - SavedXmm13 OWORD ? - SavedXmm14 OWORD ? - SavedXmm15 OWORD ? - Padding0 QWORD ? - Padding1 QWORD ? - CountN QWORD ? - ReturnAddress QWORD ? - PreviousP1Home QWORD ? - PreviousP2Home QWORD ? - PreviousP3Home QWORD ? - PreviousP4Home QWORD ? - -TanhKernelFrame ENDS - ;++ ; ; Routine Description: @@ -94,20 +46,20 @@ TanhKernelFrame ENDS ; ;-- - NESTED_ENTRY MlasTanhKernelFma3, _TEXT + NESTED_ENTRY MlasComputeTanhF32KernelFma3, _TEXT - alloc_stack (TanhKernelFrame.ReturnAddress) + alloc_stack (TransKernelFrame.ReturnAddress) - save_xmm128 xmm6,TanhKernelFrame.SavedXmm6 - save_xmm128 xmm7,TanhKernelFrame.SavedXmm7 - save_xmm128 xmm8,TanhKernelFrame.SavedXmm8 - save_xmm128 xmm9,TanhKernelFrame.SavedXmm9 - save_xmm128 xmm10,TanhKernelFrame.SavedXmm10 - save_xmm128 xmm11,TanhKernelFrame.SavedXmm11 - save_xmm128 xmm12,TanhKernelFrame.SavedXmm12 - save_xmm128 xmm13,TanhKernelFrame.SavedXmm13 - save_xmm128 xmm14,TanhKernelFrame.SavedXmm14 - save_xmm128 xmm15,TanhKernelFrame.SavedXmm15 + save_xmm128 xmm6,TransKernelFrame.SavedXmm6 + save_xmm128 xmm7,TransKernelFrame.SavedXmm7 + save_xmm128 xmm8,TransKernelFrame.SavedXmm8 + save_xmm128 xmm9,TransKernelFrame.SavedXmm9 + save_xmm128 xmm10,TransKernelFrame.SavedXmm10 + save_xmm128 xmm11,TransKernelFrame.SavedXmm11 + save_xmm128 xmm12,TransKernelFrame.SavedXmm12 + save_xmm128 xmm13,TransKernelFrame.SavedXmm13 + save_xmm128 xmm14,TransKernelFrame.SavedXmm14 + save_xmm128 xmm15,TransKernelFrame.SavedXmm15 END_PROLOGUE @@ -154,9 +106,9 @@ ComputeTanhBy8Loop: ProcessRemainingCount: add r8,8 ; correct for over-subtract above jz ExitKernel - mov DWORD PTR TanhKernelFrame.CountN[rsp],r8d - vbroadcastss ymm2,DWORD PTR TanhKernelFrame.CountN[rsp] - vpcmpgtd ymm2,ymm2,YMMWORD PTR [MlasMaskMoveAvx] + neg r8 + lea r10,MlasMaskMoveTableAvx+8*4 + vmovups ymm2,YMMWORD PTR [r10+r8*4] vmaskmovps ymm0,ymm2,YMMWORD PTR [rcx] vmaxps ymm0,ymm4,ymm0 ; clamp lower bound vminps ymm0,ymm5,ymm0 ; clamp upper bound @@ -177,22 +129,22 @@ ProcessRemainingCount: ExitKernel: vzeroupper - movaps xmm6,TanhKernelFrame.SavedXmm6[rsp] - movaps xmm7,TanhKernelFrame.SavedXmm7[rsp] - movaps xmm8,TanhKernelFrame.SavedXmm8[rsp] - movaps xmm9,TanhKernelFrame.SavedXmm9[rsp] - movaps xmm10,TanhKernelFrame.SavedXmm10[rsp] - movaps xmm11,TanhKernelFrame.SavedXmm11[rsp] - movaps xmm12,TanhKernelFrame.SavedXmm12[rsp] - movaps xmm13,TanhKernelFrame.SavedXmm13[rsp] - movaps xmm14,TanhKernelFrame.SavedXmm14[rsp] - movaps xmm15,TanhKernelFrame.SavedXmm15[rsp] - add rsp,(TanhKernelFrame.ReturnAddress) + movaps xmm6,TransKernelFrame.SavedXmm6[rsp] + movaps xmm7,TransKernelFrame.SavedXmm7[rsp] + movaps xmm8,TransKernelFrame.SavedXmm8[rsp] + movaps xmm9,TransKernelFrame.SavedXmm9[rsp] + movaps xmm10,TransKernelFrame.SavedXmm10[rsp] + movaps xmm11,TransKernelFrame.SavedXmm11[rsp] + movaps xmm12,TransKernelFrame.SavedXmm12[rsp] + movaps xmm13,TransKernelFrame.SavedXmm13[rsp] + movaps xmm14,TransKernelFrame.SavedXmm14[rsp] + movaps xmm15,TransKernelFrame.SavedXmm15[rsp] + add rsp,(TransKernelFrame.ReturnAddress) BEGIN_EPILOGUE ret - NESTED_END MlasTanhKernelFma3, _TEXT + NESTED_END MlasComputeTanhF32KernelFma3, _TEXT END diff --git a/onnxruntime/core/mlas/lib/amd64/TransKernelCommon.inc b/onnxruntime/core/mlas/lib/amd64/TransKernelCommon.inc index 2e6792e5ad..96582fe46c 100644 --- a/onnxruntime/core/mlas/lib/amd64/TransKernelCommon.inc +++ b/onnxruntime/core/mlas/lib/amd64/TransKernelCommon.inc @@ -40,6 +40,51 @@ ExpConstants STRUCT ExpConstants ENDS +; +; Structure layout for the logistic constants block. +; + +LogisticConstants STRUCT + + LowerRange DWORD ? + UpperRange DWORD ? + alpha_9 DWORD ? + alpha_7 DWORD ? + alpha_5 DWORD ? + alpha_3 DWORD ? + alpha_1 DWORD ? + beta_10 DWORD ? + beta_8 DWORD ? + beta_6 DWORD ? + beta_4 DWORD ? + beta_2 DWORD ? + beta_0 DWORD ? + one_half DWORD ? + +LogisticConstants ENDS + +; +; Structure layout for the tanh constants block. +; + +TanhConstants STRUCT + + LowerRange DWORD ? + UpperRange DWORD ? + alpha_13 DWORD ? + alpha_11 DWORD ? + alpha_9 DWORD ? + alpha_7 DWORD ? + alpha_5 DWORD ? + alpha_3 DWORD ? + alpha_1 DWORD ? + beta_6 DWORD ? + beta_4 DWORD ? + beta_2 DWORD ? + beta_0 DWORD ? + +TanhConstants ENDS + ; ; Stack frame layout for the transcedental functions. ; diff --git a/onnxruntime/core/mlas/lib/mlasi.h b/onnxruntime/core/mlas/lib/mlasi.h index 8f8727bc2c..ad56b2ec22 100644 --- a/onnxruntime/core/mlas/lib/mlasi.h +++ b/onnxruntime/core/mlas/lib/mlasi.h @@ -203,10 +203,6 @@ size_t #endif -typedef MLAS_GEMM_FLOAT_KERNEL* PMLAS_GEMM_FLOAT_KERNEL; - -typedef MLAS_GEMM_DOUBLE_KERNEL* PMLAS_GEMM_DOUBLE_KERNEL; - typedef size_t (MLASCALL MLAS_GEMV_FLOAT_KERNEL)( @@ -219,8 +215,6 @@ size_t bool ZeroMode ); -typedef MLAS_GEMV_FLOAT_KERNEL* PMLAS_GEMV_FLOAT_KERNEL; - typedef void (MLASCALL MLAS_SGEMM_KERNEL_M1_ROUTINE)( @@ -233,8 +227,6 @@ void float beta ); -typedef MLAS_SGEMM_KERNEL_M1_ROUTINE* PMLAS_SGEMM_KERNEL_M1_ROUTINE; - typedef void (MLASCALL MLAS_SGEMM_TRANSPOSE_PACKB_BLOCK_ROUTINE)( @@ -243,8 +235,6 @@ void size_t ldb ); -typedef MLAS_SGEMM_TRANSPOSE_PACKB_BLOCK_ROUTINE* PMLAS_SGEMM_TRANSPOSE_PACKB_BLOCK_ROUTINE; - typedef size_t (MLASCALL MLAS_GEMM_U8S8_KERNEL)( @@ -261,8 +251,6 @@ size_t bool ZeroMode ); -typedef MLAS_GEMM_U8S8_KERNEL* PMLAS_GEMM_U8S8_KERNEL; - typedef size_t (MLASCALL MLAS_GEMV_U8S8_KERNEL)( @@ -274,8 +262,6 @@ size_t size_t ldb ); -typedef MLAS_GEMV_U8S8_KERNEL* PMLAS_GEMV_U8S8_KERNEL; - typedef size_t (MLASCALL MLAS_GEMM_U8U8_KERNEL)( @@ -292,8 +278,6 @@ size_t bool ZeroMode ); -typedef MLAS_GEMM_U8U8_KERNEL* PMLAS_GEMM_U8U8_KERNEL; - typedef void (MLASCALL MLAS_CONV_FLOAT_KERNEL)( @@ -318,8 +302,6 @@ void unsigned KernelFlags ); -typedef MLAS_CONV_FLOAT_KERNEL* PMLAS_CONV_FLOAT_KERNEL; - typedef void (MLASCALL MLAS_CONV_DEPTHWISE_FLOAT_KERNEL)( @@ -341,8 +323,6 @@ void unsigned KernelFlags ); -typedef MLAS_CONV_DEPTHWISE_FLOAT_KERNEL* PMLAS_CONV_DEPTHWISE_FLOAT_KERNEL; - typedef void (MLASCALL MLAS_CONV_POINTWISE_FLOAT_KERNEL)( @@ -360,8 +340,6 @@ void unsigned KernelFlags ); -typedef MLAS_CONV_POINTWISE_FLOAT_KERNEL* PMLAS_CONV_POINTWISE_FLOAT_KERNEL; - typedef void (MLASCALL MLAS_POOL_FLOAT_KERNEL)( @@ -381,8 +359,6 @@ void size_t OutputCountRightPad ); -typedef MLAS_POOL_FLOAT_KERNEL* PMLAS_POOL_FLOAT_KERNEL; - typedef void (MLASCALL MLAS_COMPUTE_UNARY_FLOAT_KERNEL)( @@ -391,8 +367,6 @@ void size_t N ); -typedef MLAS_COMPUTE_UNARY_FLOAT_KERNEL* PMLAS_COMPUTE_UNARY_FLOAT_KERNEL; - typedef float (MLASCALL MLAS_COMPUTE_SUMEXP_FLOAT_KERNEL)( @@ -402,8 +376,6 @@ float const float* NegativeMaximum ); -typedef MLAS_COMPUTE_SUMEXP_FLOAT_KERNEL* PMLAS_COMPUTE_SUMEXP_FLOAT_KERNEL; - typedef void (MLASCALL MLAS_COMPUTE_SOFTMAX_OUTPUT_FLOAT_KERNEL)( @@ -412,8 +384,6 @@ void const float* Parameters ); -typedef MLAS_COMPUTE_SOFTMAX_OUTPUT_FLOAT_KERNEL* PMLAS_COMPUTE_SOFTMAX_OUTPUT_FLOAT_KERNEL; - typedef void (MLASCALL MLAS_COMPUTE_LOGSOFTMAX_OUTPUT_FLOAT_KERNEL)( @@ -423,8 +393,6 @@ void const float* Parameters ); -typedef MLAS_COMPUTE_LOGSOFTMAX_OUTPUT_FLOAT_KERNEL* PMLAS_COMPUTE_LOGSOFTMAX_OUTPUT_FLOAT_KERNEL; - typedef float (MLASCALL MLAS_REDUCE_MAXIMUM_FLOAT_KERNEL)( @@ -432,8 +400,6 @@ float size_t N ); -typedef MLAS_REDUCE_MAXIMUM_FLOAT_KERNEL* PMLAS_REDUCE_MAXIMUM_FLOAT_KERNEL; - typedef void (MLASCALL MLAS_REDUCE_MINIMUM_MAXIMUM_FLOAT_KERNEL)( @@ -443,8 +409,6 @@ void size_t N ); -typedef MLAS_REDUCE_MINIMUM_MAXIMUM_FLOAT_KERNEL* PMLAS_REDUCE_MINIMUM_MAXIMUM_FLOAT_KERNEL; - typedef void (MLASCALL MLAS_QLINEAR_BINARY_OP_S8_KERNEL)( @@ -461,8 +425,6 @@ void bool IsScalarB ); -typedef MLAS_QLINEAR_BINARY_OP_S8_KERNEL* PMLAS_QLINEAR_BINARY_OP_S8_KERNEL; - typedef void (MLASCALL MLAS_QLINEAR_BINARY_OP_U8_KERNEL)( @@ -479,8 +441,6 @@ void bool IsScalarB ); -typedef MLAS_QLINEAR_BINARY_OP_U8_KERNEL* PMLAS_QLINEAR_BINARY_OP_U8_KERNEL; - typedef void (MLASCALL MLAS_QUANTIZE_LINEAR_U8_KERNEL)( @@ -491,8 +451,6 @@ void uint8_t ZeroPoint ); -typedef MLAS_QUANTIZE_LINEAR_U8_KERNEL* PMLAS_QUANTIZE_LINEAR_U8_KERNEL; - typedef void (MLASCALL MLAS_QUANTIZE_LINEAR_S8_KERNEL)( @@ -503,8 +461,6 @@ void int8_t ZeroPoint ); -typedef MLAS_QUANTIZE_LINEAR_S8_KERNEL* PMLAS_QUANTIZE_LINEAR_S8_KERNEL; - template struct MLAS_U8X8_KERNEL { @@ -620,8 +576,8 @@ extern "C" { MLAS_COMPUTE_UNARY_FLOAT_KERNEL MlasErfKernelFma3; MLAS_COMPUTE_UNARY_FLOAT_KERNEL MlasComputeExpF32KernelFma3; MLAS_COMPUTE_UNARY_FLOAT_KERNEL MlasComputeExpF32KernelAvx512F; - MLAS_COMPUTE_UNARY_FLOAT_KERNEL MlasLogisticKernelFma3; - MLAS_COMPUTE_UNARY_FLOAT_KERNEL MlasTanhKernelFma3; + MLAS_COMPUTE_UNARY_FLOAT_KERNEL MlasComputeLogisticF32KernelFma3; + MLAS_COMPUTE_UNARY_FLOAT_KERNEL MlasComputeTanhF32KernelFma3; MLAS_COMPUTE_SUMEXP_FLOAT_KERNEL MlasComputeSumExpF32KernelFma3; MLAS_COMPUTE_SUMEXP_FLOAT_KERNEL MlasComputeSumExpF32KernelAvx512F; MLAS_COMPUTE_SOFTMAX_OUTPUT_FLOAT_KERNEL MlasComputeSoftmaxOutputF32KernelAvx; @@ -742,39 +698,39 @@ struct MLAS_PLATFORM { MLAS_PLATFORM(void); #if defined(MLAS_TARGET_AMD64_IX86) - PMLAS_GEMM_FLOAT_KERNEL GemmFloatKernel; + MLAS_GEMM_FLOAT_KERNEL* GemmFloatKernel; #endif #if defined(MLAS_TARGET_AMD64) - PMLAS_SGEMM_KERNEL_M1_ROUTINE KernelM1Routine; - PMLAS_SGEMM_KERNEL_M1_ROUTINE KernelM1TransposeBRoutine; - PMLAS_SGEMM_TRANSPOSE_PACKB_BLOCK_ROUTINE TransposePackB16x4Routine; - PMLAS_GEMM_DOUBLE_KERNEL GemmDoubleKernel; + MLAS_SGEMM_KERNEL_M1_ROUTINE* KernelM1Routine; + MLAS_SGEMM_KERNEL_M1_ROUTINE* KernelM1TransposeBRoutine; + MLAS_SGEMM_TRANSPOSE_PACKB_BLOCK_ROUTINE* TransposePackB16x4Routine; + MLAS_GEMM_DOUBLE_KERNEL* GemmDoubleKernel; const MLAS_GEMM_U8X8_DISPATCH* GemmU8S8Dispatch; - PMLAS_GEMM_U8S8_KERNEL GemmU8S8Kernel; - PMLAS_GEMV_U8S8_KERNEL GemvU8S8Kernel; + MLAS_GEMM_U8S8_KERNEL* GemmU8S8Kernel; + MLAS_GEMV_U8S8_KERNEL* GemvU8S8Kernel; const MLAS_GEMM_U8X8_DISPATCH* GemmU8U8Dispatch; - PMLAS_GEMM_U8U8_KERNEL GemmU8U8Kernel; - PMLAS_CONV_FLOAT_KERNEL ConvNchwFloatKernel; - PMLAS_CONV_FLOAT_KERNEL ConvNchwcFloatKernel; - PMLAS_CONV_DEPTHWISE_FLOAT_KERNEL ConvDepthwiseFloatKernel; - PMLAS_CONV_POINTWISE_FLOAT_KERNEL ConvPointwiseFloatKernel; - PMLAS_POOL_FLOAT_KERNEL PoolFloatKernel[MlasPoolingKindCount]; - PMLAS_COMPUTE_UNARY_FLOAT_KERNEL ErfKernelRoutine; - PMLAS_QLINEAR_BINARY_OP_S8_KERNEL QLinearAddS8Kernel; - PMLAS_QLINEAR_BINARY_OP_U8_KERNEL QLinearAddU8Kernel; + MLAS_GEMM_U8U8_KERNEL* GemmU8U8Kernel; + MLAS_CONV_FLOAT_KERNEL* ConvNchwFloatKernel; + MLAS_CONV_FLOAT_KERNEL* ConvNchwcFloatKernel; + MLAS_CONV_DEPTHWISE_FLOAT_KERNEL* ConvDepthwiseFloatKernel; + MLAS_CONV_POINTWISE_FLOAT_KERNEL* ConvPointwiseFloatKernel; + MLAS_POOL_FLOAT_KERNEL* PoolFloatKernel[MlasPoolingKindCount]; + MLAS_COMPUTE_UNARY_FLOAT_KERNEL* ErfKernelRoutine; + MLAS_QLINEAR_BINARY_OP_S8_KERNEL* QLinearAddS8Kernel; + MLAS_QLINEAR_BINARY_OP_U8_KERNEL* QLinearAddU8Kernel; MLAS_U8X8_KERNEL::DepthwiseKernel* ConvDepthwiseU8S8Kernel; MLAS_U8X8_KERNEL::DepthwiseKernel* ConvDepthwiseU8U8Kernel; - PMLAS_COMPUTE_UNARY_FLOAT_KERNEL ComputeExpF32Kernel; - PMLAS_COMPUTE_UNARY_FLOAT_KERNEL LogisticKernelRoutine; - PMLAS_COMPUTE_UNARY_FLOAT_KERNEL TanhKernelRoutine; - PMLAS_COMPUTE_SUMEXP_FLOAT_KERNEL ComputeSumExpF32Kernel; - PMLAS_COMPUTE_SOFTMAX_OUTPUT_FLOAT_KERNEL ComputeSoftmaxOutputF32Kernel; - PMLAS_COMPUTE_LOGSOFTMAX_OUTPUT_FLOAT_KERNEL ComputeLogSoftmaxOutputF32Kernel; - PMLAS_REDUCE_MAXIMUM_FLOAT_KERNEL ReduceMaximumF32Kernel; - PMLAS_REDUCE_MINIMUM_MAXIMUM_FLOAT_KERNEL ReduceMinimumMaximumF32Kernel; - PMLAS_QUANTIZE_LINEAR_S8_KERNEL QuantizeLinearS8Kernel; - PMLAS_QUANTIZE_LINEAR_U8_KERNEL QuantizeLinearU8Kernel; + MLAS_COMPUTE_UNARY_FLOAT_KERNEL* ComputeExpF32Kernel; + MLAS_COMPUTE_UNARY_FLOAT_KERNEL* LogisticKernelRoutine; + MLAS_COMPUTE_UNARY_FLOAT_KERNEL* TanhKernelRoutine; + MLAS_COMPUTE_SUMEXP_FLOAT_KERNEL* ComputeSumExpF32Kernel; + MLAS_COMPUTE_SOFTMAX_OUTPUT_FLOAT_KERNEL* ComputeSoftmaxOutputF32Kernel; + MLAS_COMPUTE_LOGSOFTMAX_OUTPUT_FLOAT_KERNEL* ComputeLogSoftmaxOutputF32Kernel; + MLAS_REDUCE_MAXIMUM_FLOAT_KERNEL* ReduceMaximumF32Kernel; + MLAS_REDUCE_MINIMUM_MAXIMUM_FLOAT_KERNEL* ReduceMinimumMaximumF32Kernel; + MLAS_QUANTIZE_LINEAR_S8_KERNEL* QuantizeLinearS8Kernel; + MLAS_QUANTIZE_LINEAR_U8_KERNEL* QuantizeLinearU8Kernel; uint32_t NchwcBlockSize; uint32_t PreferredBufferAlignment; uint32_t MaximumThreadCount; @@ -800,11 +756,9 @@ void int32_t Index ); -typedef MLAS_THREADED_ROUTINE* PMLAS_THREADED_ROUTINE; - void MlasExecuteThreaded( - PMLAS_THREADED_ROUTINE ThreadedRoutine, + MLAS_THREADED_ROUTINE* ThreadedRoutine, void* Context, int32_t Iterations, MLAS_THREADPOOL* ThreadPool diff --git a/onnxruntime/core/mlas/lib/platform.cpp b/onnxruntime/core/mlas/lib/platform.cpp index eea76913c1..43dab620a2 100644 --- a/onnxruntime/core/mlas/lib/platform.cpp +++ b/onnxruntime/core/mlas/lib/platform.cpp @@ -229,8 +229,8 @@ Return Value: this->ConvDepthwiseFloatKernel = MlasConvDepthwiseFloatKernelFma3; this->ConvPointwiseFloatKernel = MlasConvPointwiseFloatKernelFma3; this->ComputeExpF32Kernel = MlasComputeExpF32KernelFma3; - this->LogisticKernelRoutine = MlasLogisticKernelFma3; - this->TanhKernelRoutine = MlasTanhKernelFma3; + this->LogisticKernelRoutine = MlasComputeLogisticF32KernelFma3; + this->TanhKernelRoutine = MlasComputeTanhF32KernelFma3; this->ErfKernelRoutine = MlasErfKernelFma3; this->QLinearAddS8Kernel = MlasQLinearAddS8KernelAvx2; this->QLinearAddU8Kernel = MlasQLinearAddU8KernelAvx2; diff --git a/onnxruntime/core/mlas/lib/pooling.cpp b/onnxruntime/core/mlas/lib/pooling.cpp index a8fe8024df..5151cf58f0 100644 --- a/onnxruntime/core/mlas/lib/pooling.cpp +++ b/onnxruntime/core/mlas/lib/pooling.cpp @@ -45,8 +45,6 @@ void float* Output ); -typedef MLAS_POOL_KERNEL_ROUTINE* PMLAS_POOL_KERNEL_ROUTINE; - // // Define the number of elements to allocate on the stack for the reduction // buffer in the vectorized kernels. @@ -1078,7 +1076,7 @@ Return Value: // Stores pointers to the pooling kernel routines. // -static const PMLAS_POOL_KERNEL_ROUTINE MlasPoolGenericKernels[][3] = +static MLAS_POOL_KERNEL_ROUTINE* const MlasPoolGenericKernels[][3] = { { MlasPool1DKernel, @@ -1097,14 +1095,14 @@ static const PMLAS_POOL_KERNEL_ROUTINE MlasPoolGenericKernels[][3] = }, }; -static const PMLAS_POOL_KERNEL_ROUTINE MlasPoolGlobalKernels[] = +static MLAS_POOL_KERNEL_ROUTINE* const MlasPoolGlobalKernels[] = { MlasPoolGlobalKernel, MlasPoolGlobalKernel, MlasPoolGlobalKernel, }; -static const PMLAS_POOL_KERNEL_ROUTINE MlasPoolVectorKernels[][2] = +static MLAS_POOL_KERNEL_ROUTINE* const MlasPoolVectorKernels[][2] = { { MlasPool2DVectorKernel, @@ -1249,7 +1247,7 @@ Return Value: // in the reduction buffer. // - PMLAS_POOL_KERNEL_ROUTINE PoolKernelRoutine = MlasPoolGenericKernels[PoolingKind][Dimensions - 1]; + MLAS_POOL_KERNEL_ROUTINE* PoolKernelRoutine = MlasPoolGenericKernels[PoolingKind][Dimensions - 1]; if (InputAndKernelShapeMatch && AllStridesAreOne && AllPaddingIsZero) { diff --git a/onnxruntime/core/mlas/lib/sgemm.cpp b/onnxruntime/core/mlas/lib/sgemm.cpp index 17e32ff087..996a18ff9b 100644 --- a/onnxruntime/core/mlas/lib/sgemm.cpp +++ b/onnxruntime/core/mlas/lib/sgemm.cpp @@ -491,7 +491,7 @@ Return Value: #if defined(MLAS_TARGET_AMD64) - PMLAS_SGEMM_TRANSPOSE_PACKB_BLOCK_ROUTINE SgemmTransposePackB16x4Routine = + MLAS_SGEMM_TRANSPOSE_PACKB_BLOCK_ROUTINE* SgemmTransposePackB16x4Routine = MlasPlatform.TransposePackB16x4Routine; while (x >= 4) { @@ -910,7 +910,7 @@ Return Value: #if defined(MLAS_TARGET_AMD64) - PMLAS_SGEMM_KERNEL_M1_ROUTINE SgemmKernelM1Routine; + MLAS_SGEMM_KERNEL_M1_ROUTINE* SgemmKernelM1Routine; if (TransB == CblasNoTrans) { SgemmKernelM1Routine = MlasPlatform.KernelM1Routine; @@ -945,7 +945,7 @@ Return Value: #if defined(MLAS_TARGET_AMD64) - PMLAS_SGEMM_KERNEL_M1_ROUTINE SgemmKernelM1Routine; + MLAS_SGEMM_KERNEL_M1_ROUTINE* SgemmKernelM1Routine; if (TransA == CblasNoTrans) { SgemmKernelM1Routine = MlasPlatform.KernelM1TransposeBRoutine; diff --git a/onnxruntime/core/mlas/lib/snchwc.cpp b/onnxruntime/core/mlas/lib/snchwc.cpp index 3afcd7b451..23240d1909 100644 --- a/onnxruntime/core/mlas/lib/snchwc.cpp +++ b/onnxruntime/core/mlas/lib/snchwc.cpp @@ -1094,7 +1094,7 @@ struct MLAS_NCHWC_CONV_DEPTHWISE_ALGORITHM : MLAS_NCHWC_CONV_ALGORITHM struct MLAS_NCHWC_POOL_ALGORITHM : MLAS_NCHWC_NN_ALGORITHM { #if !defined(MLAS_TARGET_AMD64) - static const PMLAS_POOL_FLOAT_KERNEL PoolKernels[]; + static MLAS_POOL_FLOAT_KERNEL* const PoolKernels[]; #endif const MLAS_NCHWC_POOL_WORK_BLOCK* WorkBlock; @@ -1199,7 +1199,7 @@ struct MLAS_NCHWC_POOL_ALGORITHM : MLAS_NCHWC_NN_ALGORITHM #if !defined(MLAS_TARGET_AMD64) -const PMLAS_POOL_FLOAT_KERNEL MLAS_NCHWC_POOL_ALGORITHM::PoolKernels[] = +MLAS_POOL_FLOAT_KERNEL* const MLAS_NCHWC_POOL_ALGORITHM::PoolKernels[] = { MlasPoolMaximumFloatKernel, MlasPoolAverageExcludePadFloatKernel, @@ -1307,7 +1307,7 @@ Return Value: // reorder the filter tensor in the expected format for the given algorithm. // - PMLAS_THREADED_ROUTINE ThreadedRoutine; + MLAS_THREADED_ROUTINE* ThreadedRoutine; if (WorkBlock.InputChannels >= MlasNchwcGetBlockSize()) { if (WorkBlock.KernelShape[0] == 1 && WorkBlock.KernelShape[1] == 1 && diff --git a/onnxruntime/core/mlas/lib/threading.cpp b/onnxruntime/core/mlas/lib/threading.cpp index 7c442c5f6b..e8cb071982 100644 --- a/onnxruntime/core/mlas/lib/threading.cpp +++ b/onnxruntime/core/mlas/lib/threading.cpp @@ -18,7 +18,7 @@ Abstract: void MlasExecuteThreaded( - PMLAS_THREADED_ROUTINE ThreadedRoutine, + MLAS_THREADED_ROUTINE* ThreadedRoutine, void* Context, int32_t Iterations, MLAS_THREADPOOL* ThreadPool diff --git a/onnxruntime/core/mlas/lib/x86_64/DgemmKernelAvx.S b/onnxruntime/core/mlas/lib/x86_64/DgemmKernelAvx.S index 575c822d01..8f791a734a 100644 --- a/onnxruntime/core/mlas/lib/x86_64/DgemmKernelAvx.S +++ b/onnxruntime/core/mlas/lib/x86_64/DgemmKernelAvx.S @@ -29,6 +29,6 @@ Abstract: // Generate the GEMM kernel. // -FgemmKernelAvxFunction C_UNDERSCORE(MlasGemmDoubleKernelAvx) +FgemmKernelAvxFunction MlasGemmDoubleKernelAvx .end diff --git a/onnxruntime/core/mlas/lib/x86_64/DgemmKernelAvx512F.S b/onnxruntime/core/mlas/lib/x86_64/DgemmKernelAvx512F.S index be015ca457..23f8afcb2b 100644 --- a/onnxruntime/core/mlas/lib/x86_64/DgemmKernelAvx512F.S +++ b/onnxruntime/core/mlas/lib/x86_64/DgemmKernelAvx512F.S @@ -29,6 +29,6 @@ Abstract: // Generate the GEMM kernel. // -FgemmKernelAvx512FFunction C_UNDERSCORE(MlasGemmDoubleKernelAvx512F) +FgemmKernelAvx512FFunction MlasGemmDoubleKernelAvx512F .end diff --git a/onnxruntime/core/mlas/lib/x86_64/DgemmKernelFma3.S b/onnxruntime/core/mlas/lib/x86_64/DgemmKernelFma3.S index 77e04bfc93..707882af78 100644 --- a/onnxruntime/core/mlas/lib/x86_64/DgemmKernelFma3.S +++ b/onnxruntime/core/mlas/lib/x86_64/DgemmKernelFma3.S @@ -29,6 +29,6 @@ Abstract: // Generate the GEMM kernel. // -FgemmKernelFma3Function C_UNDERSCORE(MlasGemmDoubleKernelFma3) +FgemmKernelFma3Function MlasGemmDoubleKernelFma3 .end diff --git a/onnxruntime/core/mlas/lib/x86_64/DgemmKernelSse2.S b/onnxruntime/core/mlas/lib/x86_64/DgemmKernelSse2.S index cf7630545f..929eaf3510 100644 --- a/onnxruntime/core/mlas/lib/x86_64/DgemmKernelSse2.S +++ b/onnxruntime/core/mlas/lib/x86_64/DgemmKernelSse2.S @@ -225,6 +225,6 @@ Implicit Arguments: // Generate the GEMM kernel. // -FgemmKernelSse2Function C_UNDERSCORE(MlasGemmDoubleKernelSse) +FgemmKernelSse2Function MlasGemmDoubleKernelSse .end diff --git a/onnxruntime/core/mlas/lib/x86_64/FgemmKernelAvx512FCommon.h b/onnxruntime/core/mlas/lib/x86_64/FgemmKernelAvx512FCommon.h index 22e4c361d1..9f243ee8c0 100644 --- a/onnxruntime/core/mlas/lib/x86_64/FgemmKernelAvx512FCommon.h +++ b/onnxruntime/core/mlas/lib/x86_64/FgemmKernelAvx512FCommon.h @@ -453,8 +453,7 @@ Return Value: --*/ - .globl \FunctionName\() -\FunctionName\(): + FUNCTION_ENTRY \FunctionName\() push rbp push rbx diff --git a/onnxruntime/core/mlas/lib/x86_64/FgemmKernelAvxCommon.h b/onnxruntime/core/mlas/lib/x86_64/FgemmKernelAvxCommon.h index a5abee9770..d84cc54953 100644 --- a/onnxruntime/core/mlas/lib/x86_64/FgemmKernelAvxCommon.h +++ b/onnxruntime/core/mlas/lib/x86_64/FgemmKernelAvxCommon.h @@ -402,8 +402,7 @@ Return Value: --*/ - .globl \FunctionName\() -\FunctionName\(): + FUNCTION_ENTRY \FunctionName\() push rbp push rbx diff --git a/onnxruntime/core/mlas/lib/x86_64/FgemmKernelFma3Common.h b/onnxruntime/core/mlas/lib/x86_64/FgemmKernelFma3Common.h index f108b5c8a5..3c26e053b7 100644 --- a/onnxruntime/core/mlas/lib/x86_64/FgemmKernelFma3Common.h +++ b/onnxruntime/core/mlas/lib/x86_64/FgemmKernelFma3Common.h @@ -449,8 +449,7 @@ Return Value: --*/ - .globl \FunctionName\() -\FunctionName\(): + FUNCTION_ENTRY \FunctionName\() push rbp push rbx diff --git a/onnxruntime/core/mlas/lib/x86_64/FgemmKernelSse2Common.h b/onnxruntime/core/mlas/lib/x86_64/FgemmKernelSse2Common.h index 88cc1b4fd1..2f71864f00 100644 --- a/onnxruntime/core/mlas/lib/x86_64/FgemmKernelSse2Common.h +++ b/onnxruntime/core/mlas/lib/x86_64/FgemmKernelSse2Common.h @@ -130,8 +130,7 @@ Return Value: --*/ - .globl \FunctionName\() -\FunctionName\(): + FUNCTION_ENTRY \FunctionName\() push rbp push rbx diff --git a/onnxruntime/core/mlas/lib/x86_64/LogisticKernelFma3.S b/onnxruntime/core/mlas/lib/x86_64/LogisticKernelFma3.S index 243b355398..f1ee717363 100644 --- a/onnxruntime/core/mlas/lib/x86_64/LogisticKernelFma3.S +++ b/onnxruntime/core/mlas/lib/x86_64/LogisticKernelFma3.S @@ -18,37 +18,12 @@ Abstract: --*/ #include "asmmacro.h" +#include "TransKernelCommon.h" .intel_syntax noprefix .text -// -// Structure layout for the logistic constants block. -// - - .equ LogisticConstants_LowerRange, 0 - .equ LogisticConstants_UpperRange, 4 - .equ LogisticConstants_alpha_9, 8 - .equ LogisticConstants_alpha_7, 12 - .equ LogisticConstants_alpha_5, 16 - .equ LogisticConstants_alpha_3, 20 - .equ LogisticConstants_alpha_1, 24 - .equ LogisticConstants_beta_10, 28 - .equ LogisticConstants_beta_8, 32 - .equ LogisticConstants_beta_6, 36 - .equ LogisticConstants_beta_4, 40 - .equ LogisticConstants_beta_2, 44 - .equ LogisticConstants_beta_0, 48 - .equ LogisticConstants_one_half, 52 - -// -// Stack frame layout for the logistic kernel. -// - - .equ LogisticKernelFrame_CountN, -8 - .equ LogisticKernelFrame_ReturnAddress, 0 - /*++ Routine Description: @@ -69,22 +44,21 @@ Return Value: --*/ - .globl C_UNDERSCORE(MlasLogisticKernelFma3) -C_UNDERSCORE(MlasLogisticKernelFma3): + FUNCTION_ENTRY MlasComputeLogisticF32KernelFma3 lea rax,C_UNDERSCORE(MlasLogisticConstants)[rip] - vbroadcastss ymm4,LogisticConstants_LowerRange[rax] - vbroadcastss ymm5,LogisticConstants_UpperRange[rax] - vbroadcastss ymm6,LogisticConstants_alpha_9[rax] - vbroadcastss ymm7,LogisticConstants_alpha_7[rax] - vbroadcastss ymm8,LogisticConstants_alpha_5[rax] - vbroadcastss ymm9,LogisticConstants_alpha_3[rax] - vbroadcastss ymm10,LogisticConstants_alpha_1[rax] - vbroadcastss ymm11,LogisticConstants_beta_10[rax] - vbroadcastss ymm12,LogisticConstants_beta_6[rax] - vbroadcastss ymm13,LogisticConstants_beta_4[rax] - vbroadcastss ymm14,LogisticConstants_beta_2[rax] - vbroadcastss ymm15,LogisticConstants_beta_0[rax] + vbroadcastss ymm4,.LLogisticConstants_LowerRange[rax] + vbroadcastss ymm5,.LLogisticConstants_UpperRange[rax] + vbroadcastss ymm6,.LLogisticConstants_alpha_9[rax] + vbroadcastss ymm7,.LLogisticConstants_alpha_7[rax] + vbroadcastss ymm8,.LLogisticConstants_alpha_5[rax] + vbroadcastss ymm9,.LLogisticConstants_alpha_3[rax] + vbroadcastss ymm10,.LLogisticConstants_alpha_1[rax] + vbroadcastss ymm11,.LLogisticConstants_beta_10[rax] + vbroadcastss ymm12,.LLogisticConstants_beta_6[rax] + vbroadcastss ymm13,.LLogisticConstants_beta_4[rax] + vbroadcastss ymm14,.LLogisticConstants_beta_2[rax] + vbroadcastss ymm15,.LLogisticConstants_beta_0[rax] sub rdx,8 jb .LProcessRemainingCount @@ -94,7 +68,7 @@ C_UNDERSCORE(MlasLogisticKernelFma3): vmovaps ymm2,ymm7 vminps ymm0,ymm5,ymm0 # clamp upper bound vmulps ymm1,ymm0,ymm0 # x2 - vbroadcastss ymm3,LogisticConstants_beta_8[rax] + vbroadcastss ymm3,.LLogisticConstants_beta_8[rax] vfmadd231ps ymm2,ymm1,ymm6 # p = x2 * alpha_9 + alpha_7 vfmadd213ps ymm2,ymm1,ymm8 # p = x2 * p + alpha_5 vfmadd213ps ymm2,ymm1,ymm9 # p = x2 * p + alpha_3 @@ -105,7 +79,7 @@ C_UNDERSCORE(MlasLogisticKernelFma3): vfmadd213ps ymm3,ymm1,ymm14 # q = x2 * q + beta_2 vfmadd213ps ymm3,ymm1,ymm15 # q = x2 * q + beta_0 vmulps ymm2,ymm0,ymm2 # p = x * p - vbroadcastss ymm0,LogisticConstants_one_half[rax] + vbroadcastss ymm0,.LLogisticConstants_one_half[rax] vdivps ymm2,ymm2,ymm3 vxorps ymm3,ymm3,ymm3 vaddps ymm0,ymm2,ymm0 # logistic = p / q + 0.5 @@ -119,14 +93,14 @@ C_UNDERSCORE(MlasLogisticKernelFma3): .LProcessRemainingCount: add rdx,8 # correct for over-subtract above jz .LExitKernel - mov DWORD PTR LogisticKernelFrame_CountN[rsp],edx - vbroadcastss ymm2,DWORD PTR LogisticKernelFrame_CountN[rsp] - vpcmpgtd ymm2,ymm2,YMMWORD PTR C_UNDERSCORE(MlasMaskMoveAvx)[rip] + neg rdx + lea r10,C_UNDERSCORE(MlasMaskMoveTableAvx)[rip+8*4] + vmovups ymm2,YMMWORD PTR [r10+rdx*4] vmaskmovps ymm0,ymm2,YMMWORD PTR [rdi] vmaxps ymm0,ymm4,ymm0 # clamp lower bound vminps ymm0,ymm5,ymm0 # clamp upper bound vmulps ymm1,ymm0,ymm0 # x2 - vbroadcastss ymm3,LogisticConstants_beta_8[rax] + vbroadcastss ymm3,.LLogisticConstants_beta_8[rax] vfmadd231ps ymm7,ymm1,ymm6 # p = x2 * alpha_9 + alpha_7 vfmadd213ps ymm7,ymm1,ymm8 # p = x2 * p + alpha_5 vfmadd213ps ymm7,ymm1,ymm9 # p = x2 * p + alpha_3 @@ -137,7 +111,7 @@ C_UNDERSCORE(MlasLogisticKernelFma3): vfmadd213ps ymm3,ymm1,ymm14 # q = x2 * q + beta_2 vfmadd213ps ymm3,ymm1,ymm15 # q = x2 * q + beta_0 vmulps ymm7,ymm0,ymm7 # p = x * p - vbroadcastss ymm0,LogisticConstants_one_half[rax] + vbroadcastss ymm0,.LLogisticConstants_one_half[rax] vdivps ymm7,ymm7,ymm3 vxorps ymm3,ymm3,ymm3 vaddps ymm0,ymm7,ymm0 # logistic = p / q + 0.5 diff --git a/onnxruntime/core/mlas/lib/x86_64/SconvKernelCommon.h b/onnxruntime/core/mlas/lib/x86_64/SconvKernelCommon.h index 2e319f1e14..a2a858d15d 100644 --- a/onnxruntime/core/mlas/lib/x86_64/SconvKernelCommon.h +++ b/onnxruntime/core/mlas/lib/x86_64/SconvKernelCommon.h @@ -344,8 +344,7 @@ Return Value: --*/ - .globl C_UNDERSCORE(MlasConv\KernelType\()FloatKernel\Isa\()) -C_UNDERSCORE(MlasConv\KernelType\()FloatKernel\Isa\()): + FUNCTION_ENTRY MlasConv\KernelType\()FloatKernel\Isa\() push rbp push rbx @@ -511,8 +510,7 @@ Return Value: --*/ - .globl C_UNDERSCORE(MlasConvDepthwiseFloatKernel\Isa\()) -C_UNDERSCORE(MlasConvDepthwiseFloatKernel\Isa\()): + FUNCTION_ENTRY MlasConvDepthwiseFloatKernel\Isa\() push rbp push rbx @@ -707,8 +705,7 @@ Return Value: --*/ - .globl C_UNDERSCORE(MlasConvPointwiseFloatKernel\Isa\()) -C_UNDERSCORE(MlasConvPointwiseFloatKernel\Isa\()): + FUNCTION_ENTRY MlasConvPointwiseFloatKernel\Isa\() push rbp push rbx diff --git a/onnxruntime/core/mlas/lib/x86_64/SgemmKernelAvx.S b/onnxruntime/core/mlas/lib/x86_64/SgemmKernelAvx.S index 9769ee5af1..a0a66f330a 100644 --- a/onnxruntime/core/mlas/lib/x86_64/SgemmKernelAvx.S +++ b/onnxruntime/core/mlas/lib/x86_64/SgemmKernelAvx.S @@ -29,6 +29,6 @@ Abstract: // Generate the GEMM kernel. // -FgemmKernelAvxFunction C_UNDERSCORE(MlasGemmFloatKernelAvx) +FgemmKernelAvxFunction MlasGemmFloatKernelAvx .end diff --git a/onnxruntime/core/mlas/lib/x86_64/SgemmKernelAvx512F.S b/onnxruntime/core/mlas/lib/x86_64/SgemmKernelAvx512F.S index 598732a905..c75df76030 100644 --- a/onnxruntime/core/mlas/lib/x86_64/SgemmKernelAvx512F.S +++ b/onnxruntime/core/mlas/lib/x86_64/SgemmKernelAvx512F.S @@ -29,6 +29,6 @@ Abstract: // Generate the GEMM kernel. // -FgemmKernelAvx512FFunction C_UNDERSCORE(MlasGemmFloatKernelAvx512F) +FgemmKernelAvx512FFunction MlasGemmFloatKernelAvx512F .end diff --git a/onnxruntime/core/mlas/lib/x86_64/SgemmKernelFma3.S b/onnxruntime/core/mlas/lib/x86_64/SgemmKernelFma3.S index 36082201b1..4725459323 100644 --- a/onnxruntime/core/mlas/lib/x86_64/SgemmKernelFma3.S +++ b/onnxruntime/core/mlas/lib/x86_64/SgemmKernelFma3.S @@ -29,6 +29,6 @@ Abstract: // Generate the GEMM kernel. // -FgemmKernelFma3Function C_UNDERSCORE(MlasGemmFloatKernelFma3) +FgemmKernelFma3Function MlasGemmFloatKernelFma3 .end diff --git a/onnxruntime/core/mlas/lib/x86_64/SgemmKernelSse2.S b/onnxruntime/core/mlas/lib/x86_64/SgemmKernelSse2.S index b51a0956bd..e605128537 100644 --- a/onnxruntime/core/mlas/lib/x86_64/SgemmKernelSse2.S +++ b/onnxruntime/core/mlas/lib/x86_64/SgemmKernelSse2.S @@ -268,6 +268,6 @@ Implicit Arguments: // Generate the GEMM kernel. // -FgemmKernelSse2Function C_UNDERSCORE(MlasGemmFloatKernelSse) +FgemmKernelSse2Function MlasGemmFloatKernelSse .end diff --git a/onnxruntime/core/mlas/lib/x86_64/SoftmaxKernelAvx.S b/onnxruntime/core/mlas/lib/x86_64/SoftmaxKernelAvx.S index 7432ff0f92..76247ecf7c 100644 --- a/onnxruntime/core/mlas/lib/x86_64/SoftmaxKernelAvx.S +++ b/onnxruntime/core/mlas/lib/x86_64/SoftmaxKernelAvx.S @@ -42,8 +42,7 @@ Return Value: --*/ - .globl C_UNDERSCORE(MlasReduceMaximumF32KernelAvx) -C_UNDERSCORE(MlasReduceMaximumF32KernelAvx): + FUNCTION_ENTRY MlasReduceMaximumF32KernelAvx vbroadcastss ymm0,DWORD PTR C_UNDERSCORE(MlasMinimumF32Value)[rip] test rsi,rsi @@ -118,8 +117,7 @@ Return Value: --*/ - .globl C_UNDERSCORE(MlasComputeSoftmaxOutputF32KernelAvx) -C_UNDERSCORE(MlasComputeSoftmaxOutputF32KernelAvx): + FUNCTION_ENTRY MlasComputeSoftmaxOutputF32KernelAvx vbroadcastss ymm4,DWORD PTR [rdx] # broadcast scale value cmp rsi,32 @@ -187,8 +185,7 @@ Return Value: --*/ - .globl C_UNDERSCORE(MlasComputeLogSoftmaxOutputF32KernelAvx) -C_UNDERSCORE(MlasComputeLogSoftmaxOutputF32KernelAvx): + FUNCTION_ENTRY MlasComputeLogSoftmaxOutputF32KernelAvx vbroadcastss ymm4,DWORD PTR [rcx] # broadcast negative minimum value vbroadcastss ymm5,DWORD PTR [rcx+4] # broadcast log(SumExp) diff --git a/onnxruntime/core/mlas/lib/x86_64/SpoolKernelAvxCommon.h b/onnxruntime/core/mlas/lib/x86_64/SpoolKernelAvxCommon.h index 08ce505b44..68de6acc7a 100644 --- a/onnxruntime/core/mlas/lib/x86_64/SpoolKernelAvxCommon.h +++ b/onnxruntime/core/mlas/lib/x86_64/SpoolKernelAvxCommon.h @@ -95,8 +95,7 @@ Return Value: --*/ - .globl C_UNDERSCORE(MlasPool\PoolingType\()FloatKernel\Isa\()) -C_UNDERSCORE(MlasPool\PoolingType\()FloatKernel\Isa\()): + FUNCTION_ENTRY MlasPool\PoolingType\()FloatKernel\Isa\() SpoolKernelEntry \PoolingType\() diff --git a/onnxruntime/core/mlas/lib/x86_64/SpoolKernelSse2.S b/onnxruntime/core/mlas/lib/x86_64/SpoolKernelSse2.S index 4dc26a369c..285da3072c 100644 --- a/onnxruntime/core/mlas/lib/x86_64/SpoolKernelSse2.S +++ b/onnxruntime/core/mlas/lib/x86_64/SpoolKernelSse2.S @@ -275,8 +275,7 @@ Return Value: --*/ - .globl C_UNDERSCORE(MlasPool\PoolingType\()FloatKernel\Isa\()) -C_UNDERSCORE(MlasPool\PoolingType\()FloatKernel\Isa\()): + FUNCTION_ENTRY MlasPool\PoolingType\()FloatKernel\Isa\() SpoolKernelEntry \PoolingType\() diff --git a/onnxruntime/core/mlas/lib/x86_64/TanhKernelFma3.S b/onnxruntime/core/mlas/lib/x86_64/TanhKernelFma3.S index dd5584648d..d7c2fd1c6e 100644 --- a/onnxruntime/core/mlas/lib/x86_64/TanhKernelFma3.S +++ b/onnxruntime/core/mlas/lib/x86_64/TanhKernelFma3.S @@ -18,36 +18,12 @@ Abstract: --*/ #include "asmmacro.h" +#include "TransKernelCommon.h" .intel_syntax noprefix .text -// -// Structure layout for the tanh constants block. -// - - .equ TanhConstants_LowerRange, 0 - .equ TanhConstants_UpperRange, 4 - .equ TanhConstants_alpha_13, 8 - .equ TanhConstants_alpha_11, 12 - .equ TanhConstants_alpha_9, 16 - .equ TanhConstants_alpha_7, 20 - .equ TanhConstants_alpha_5, 24 - .equ TanhConstants_alpha_3, 28 - .equ TanhConstants_alpha_1, 32 - .equ TanhConstants_beta_6, 36 - .equ TanhConstants_beta_4, 40 - .equ TanhConstants_beta_2, 44 - .equ TanhConstants_beta_0, 48 - -// -// Stack frame layout for the tanh kernel. -// - - .equ TanhKernelFrame_CountN, -8 - .equ TanhKernelFrame_ReturnAddress, 0 - /*++ Routine Description: @@ -69,22 +45,21 @@ Return Value: --*/ - .globl C_UNDERSCORE(MlasTanhKernelFma3) -C_UNDERSCORE(MlasTanhKernelFma3): + FUNCTION_ENTRY MlasComputeTanhF32KernelFma3 lea rax,C_UNDERSCORE(MlasTanhConstants)[rip] - vbroadcastss ymm4,TanhConstants_LowerRange[rax] - vbroadcastss ymm5,TanhConstants_UpperRange[rax] - vbroadcastss ymm6,TanhConstants_alpha_13[rax] - vbroadcastss ymm7,TanhConstants_alpha_11[rax] - vbroadcastss ymm8,TanhConstants_alpha_9[rax] - vbroadcastss ymm9,TanhConstants_alpha_7[rax] - vbroadcastss ymm10,TanhConstants_alpha_5[rax] - vbroadcastss ymm11,TanhConstants_alpha_3[rax] - vbroadcastss ymm12,TanhConstants_alpha_1[rax] - vbroadcastss ymm13,TanhConstants_beta_6[rax] - vbroadcastss ymm14,TanhConstants_beta_2[rax] - vbroadcastss ymm15,TanhConstants_beta_0[rax] + vbroadcastss ymm4,.LTanhConstants_LowerRange[rax] + vbroadcastss ymm5,.LTanhConstants_UpperRange[rax] + vbroadcastss ymm6,.LTanhConstants_alpha_13[rax] + vbroadcastss ymm7,.LTanhConstants_alpha_11[rax] + vbroadcastss ymm8,.LTanhConstants_alpha_9[rax] + vbroadcastss ymm9,.LTanhConstants_alpha_7[rax] + vbroadcastss ymm10,.LTanhConstants_alpha_5[rax] + vbroadcastss ymm11,.LTanhConstants_alpha_3[rax] + vbroadcastss ymm12,.LTanhConstants_alpha_1[rax] + vbroadcastss ymm13,.LTanhConstants_beta_6[rax] + vbroadcastss ymm14,.LTanhConstants_beta_2[rax] + vbroadcastss ymm15,.LTanhConstants_beta_0[rax] sub rdx,8 jb .LProcessRemainingCount @@ -94,7 +69,7 @@ C_UNDERSCORE(MlasTanhKernelFma3): vmovaps ymm2,ymm7 vminps ymm0,ymm5,ymm0 # clamp upper bound vmulps ymm1,ymm0,ymm0 # x2 - vbroadcastss ymm3,TanhConstants_beta_4[rax] + vbroadcastss ymm3,.LTanhConstants_beta_4[rax] vfmadd231ps ymm2,ymm1,ymm6 # p = x2 * alpha_13 + alpha_11 vfmadd213ps ymm2,ymm1,ymm8 # p = x2 * p + alpha_9 vfmadd213ps ymm2,ymm1,ymm9 # p = x2 * p + alpha_7 @@ -115,14 +90,14 @@ C_UNDERSCORE(MlasTanhKernelFma3): .LProcessRemainingCount: add rdx,8 # correct for over-subtract above jz .LExitKernel - mov DWORD PTR TanhKernelFrame_CountN[rsp],edx - vbroadcastss ymm2,DWORD PTR TanhKernelFrame_CountN[rsp] - vpcmpgtd ymm2,ymm2,YMMWORD PTR C_UNDERSCORE(MlasMaskMoveAvx)[rip] + neg rdx + lea r10,C_UNDERSCORE(MlasMaskMoveTableAvx)[rip+8*4] + vmovups ymm2,YMMWORD PTR [r10+rdx*4] vmaskmovps ymm0,ymm2,YMMWORD PTR [rdi] vmaxps ymm0,ymm4,ymm0 # clamp lower bound vminps ymm0,ymm5,ymm0 # clamp upper bound vmulps ymm1,ymm0,ymm0 # x2 - vbroadcastss ymm3,TanhConstants_beta_4[rax] + vbroadcastss ymm3,.LTanhConstants_beta_4[rax] vfmadd231ps ymm7,ymm1,ymm6 # p = x2 * alpha_13 + alpha_11 vfmadd213ps ymm7,ymm1,ymm8 # p = x2 * p + alpha_9 vfmadd213ps ymm7,ymm1,ymm9 # p = x2 * p + alpha_7 diff --git a/onnxruntime/core/mlas/lib/x86_64/TransKernelAvx512F.S b/onnxruntime/core/mlas/lib/x86_64/TransKernelAvx512F.S index 0eabe85b28..64b6204249 100644 --- a/onnxruntime/core/mlas/lib/x86_64/TransKernelAvx512F.S +++ b/onnxruntime/core/mlas/lib/x86_64/TransKernelAvx512F.S @@ -43,8 +43,7 @@ Return Value: --*/ - .globl C_UNDERSCORE(MlasComputeExpF32KernelAvx512F) -C_UNDERSCORE(MlasComputeExpF32KernelAvx512F): + FUNCTION_ENTRY MlasComputeExpF32KernelAvx512F lea rax,C_UNDERSCORE(MlasExpConstants)[rip] vbroadcastss zmm21,.LExpConstants_LowerRange[rax] @@ -133,8 +132,7 @@ Return Value: --*/ - .globl C_UNDERSCORE(MlasComputeSumExpF32KernelAvx512F) -C_UNDERSCORE(MlasComputeSumExpF32KernelAvx512F): + FUNCTION_ENTRY MlasComputeSumExpF32KernelAvx512F lea rax,C_UNDERSCORE(MlasExpConstants)[rip] vbroadcastss zmm21,.LExpConstants_LowerRange[rax] diff --git a/onnxruntime/core/mlas/lib/x86_64/TransKernelCommon.h b/onnxruntime/core/mlas/lib/x86_64/TransKernelCommon.h index 75b034be0a..f8c76522c7 100644 --- a/onnxruntime/core/mlas/lib/x86_64/TransKernelCommon.h +++ b/onnxruntime/core/mlas/lib/x86_64/TransKernelCommon.h @@ -35,3 +35,40 @@ Abstract: .equ .LExpConstants_poly_56, 52 .equ .LExpConstants_MinimumExponent, 56 .equ .LExpConstants_MaximumExponent, 60 + +// +// Structure layout for the logistic constants block. +// + + .equ .LLogisticConstants_LowerRange, 0 + .equ .LLogisticConstants_UpperRange, 4 + .equ .LLogisticConstants_alpha_9, 8 + .equ .LLogisticConstants_alpha_7, 12 + .equ .LLogisticConstants_alpha_5, 16 + .equ .LLogisticConstants_alpha_3, 20 + .equ .LLogisticConstants_alpha_1, 24 + .equ .LLogisticConstants_beta_10, 28 + .equ .LLogisticConstants_beta_8, 32 + .equ .LLogisticConstants_beta_6, 36 + .equ .LLogisticConstants_beta_4, 40 + .equ .LLogisticConstants_beta_2, 44 + .equ .LLogisticConstants_beta_0, 48 + .equ .LLogisticConstants_one_half, 52 + +// +// Structure layout for the tanh constants block. +// + + .equ .LTanhConstants_LowerRange, 0 + .equ .LTanhConstants_UpperRange, 4 + .equ .LTanhConstants_alpha_13, 8 + .equ .LTanhConstants_alpha_11, 12 + .equ .LTanhConstants_alpha_9, 16 + .equ .LTanhConstants_alpha_7, 20 + .equ .LTanhConstants_alpha_5, 24 + .equ .LTanhConstants_alpha_3, 28 + .equ .LTanhConstants_alpha_1, 32 + .equ .LTanhConstants_beta_6, 36 + .equ .LTanhConstants_beta_4, 40 + .equ .LTanhConstants_beta_2, 44 + .equ .LTanhConstants_beta_0, 48 diff --git a/onnxruntime/core/mlas/lib/x86_64/TransKernelFma3.S b/onnxruntime/core/mlas/lib/x86_64/TransKernelFma3.S index ae3712365c..829c735bac 100644 --- a/onnxruntime/core/mlas/lib/x86_64/TransKernelFma3.S +++ b/onnxruntime/core/mlas/lib/x86_64/TransKernelFma3.S @@ -43,8 +43,7 @@ Return Value: --*/ - .globl C_UNDERSCORE(MlasComputeExpF32KernelFma3) -C_UNDERSCORE(MlasComputeExpF32KernelFma3): + FUNCTION_ENTRY MlasComputeExpF32KernelFma3 lea rax,C_UNDERSCORE(MlasExpConstants)[rip] vbroadcastss ymm4,.LExpConstants_LowerRange[rax] @@ -155,8 +154,7 @@ Return Value: --*/ - .globl C_UNDERSCORE(MlasComputeSumExpF32KernelFma3) -C_UNDERSCORE(MlasComputeSumExpF32KernelFma3): + FUNCTION_ENTRY MlasComputeSumExpF32KernelFma3 lea rax,C_UNDERSCORE(MlasExpConstants)[rip] vbroadcastss ymm9,DWORD PTR [rcx] # broadcast negative maximum value