MLAS: misc cleanup (#7013)

Miscellaneous changes to synchronize the style used over time:

Remove unneeded PFN types in favor of FN*.
Switch more functions over to using the common FUNCTION_ENTRY macro.
Switch logistic/tanh kernels over to the style used in TransKernelFma3.asm.
This commit is contained in:
Tracy Sharpe 2021-03-15 18:24:18 -07:00 committed by GitHub
parent 4e670f7ab1
commit 5480f8dd1d
No known key found for this signature in database
GPG key ID: 4AEE18F83AFDEB23
30 changed files with 246 additions and 376 deletions

View file

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

View file

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

View file

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

View file

@ -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<typename FilterType>
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<int8_t>::DepthwiseKernel* ConvDepthwiseU8S8Kernel;
MLAS_U8X8_KERNEL<uint8_t>::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

View file

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

View file

@ -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<MLAS_MAXIMUM_POOLING>,
@ -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<MLAS_MAXIMUM_POOLING>,
MlasPoolGlobalKernel<MLAS_AVERAGE_POOLING>,
MlasPoolGlobalKernel<MLAS_AVERAGE_POOLING>,
};
static const PMLAS_POOL_KERNEL_ROUTINE MlasPoolVectorKernels[][2] =
static MLAS_POOL_KERNEL_ROUTINE* const MlasPoolVectorKernels[][2] =
{
{
MlasPool2DVectorKernel<MLAS_MAXIMUM_POOLING>,
@ -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) {

View file

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

View file

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

View file

@ -18,7 +18,7 @@ Abstract:
void
MlasExecuteThreaded(
PMLAS_THREADED_ROUTINE ThreadedRoutine,
MLAS_THREADED_ROUTINE* ThreadedRoutine,
void* Context,
int32_t Iterations,
MLAS_THREADPOOL* ThreadPool

View file

@ -29,6 +29,6 @@ Abstract:
// Generate the GEMM kernel.
//
FgemmKernelAvxFunction C_UNDERSCORE(MlasGemmDoubleKernelAvx)
FgemmKernelAvxFunction MlasGemmDoubleKernelAvx
.end

View file

@ -29,6 +29,6 @@ Abstract:
// Generate the GEMM kernel.
//
FgemmKernelAvx512FFunction C_UNDERSCORE(MlasGemmDoubleKernelAvx512F)
FgemmKernelAvx512FFunction MlasGemmDoubleKernelAvx512F
.end

View file

@ -29,6 +29,6 @@ Abstract:
// Generate the GEMM kernel.
//
FgemmKernelFma3Function C_UNDERSCORE(MlasGemmDoubleKernelFma3)
FgemmKernelFma3Function MlasGemmDoubleKernelFma3
.end

View file

@ -225,6 +225,6 @@ Implicit Arguments:
// Generate the GEMM kernel.
//
FgemmKernelSse2Function C_UNDERSCORE(MlasGemmDoubleKernelSse)
FgemmKernelSse2Function MlasGemmDoubleKernelSse
.end

View file

@ -453,8 +453,7 @@ Return Value:
--*/
.globl \FunctionName\()
\FunctionName\():
FUNCTION_ENTRY \FunctionName\()
push rbp
push rbx

View file

@ -402,8 +402,7 @@ Return Value:
--*/
.globl \FunctionName\()
\FunctionName\():
FUNCTION_ENTRY \FunctionName\()
push rbp
push rbx

View file

@ -449,8 +449,7 @@ Return Value:
--*/
.globl \FunctionName\()
\FunctionName\():
FUNCTION_ENTRY \FunctionName\()
push rbp
push rbx

View file

@ -130,8 +130,7 @@ Return Value:
--*/
.globl \FunctionName\()
\FunctionName\():
FUNCTION_ENTRY \FunctionName\()
push rbp
push rbx

View file

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

View file

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

View file

@ -29,6 +29,6 @@ Abstract:
// Generate the GEMM kernel.
//
FgemmKernelAvxFunction C_UNDERSCORE(MlasGemmFloatKernelAvx)
FgemmKernelAvxFunction MlasGemmFloatKernelAvx
.end

View file

@ -29,6 +29,6 @@ Abstract:
// Generate the GEMM kernel.
//
FgemmKernelAvx512FFunction C_UNDERSCORE(MlasGemmFloatKernelAvx512F)
FgemmKernelAvx512FFunction MlasGemmFloatKernelAvx512F
.end

View file

@ -29,6 +29,6 @@ Abstract:
// Generate the GEMM kernel.
//
FgemmKernelFma3Function C_UNDERSCORE(MlasGemmFloatKernelFma3)
FgemmKernelFma3Function MlasGemmFloatKernelFma3
.end

View file

@ -268,6 +268,6 @@ Implicit Arguments:
// Generate the GEMM kernel.
//
FgemmKernelSse2Function C_UNDERSCORE(MlasGemmFloatKernelSse)
FgemmKernelSse2Function MlasGemmFloatKernelSse
.end

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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