mirror of
https://github.com/saymrwulf/onnxruntime.git
synced 2026-07-30 20:18:08 +00:00
MLAS: qgemm refactoring (#4030)
Treat U8U8 as U8S8 for VNNI for performance and optimize SSE2 kernel.
This commit is contained in:
parent
abcd1576c9
commit
0d8abc1a99
21 changed files with 800 additions and 1448 deletions
|
|
@ -55,7 +55,6 @@ if(MSVC)
|
|||
${ONNXRUNTIME_ROOT}/core/mlas/lib/amd64/QgemvU8S8KernelAvx512Vnni.asm
|
||||
${ONNXRUNTIME_ROOT}/core/mlas/lib/amd64/QgemmU8U8KernelAvx2.asm
|
||||
${ONNXRUNTIME_ROOT}/core/mlas/lib/amd64/QgemmU8U8KernelAvx512Core.asm
|
||||
${ONNXRUNTIME_ROOT}/core/mlas/lib/amd64/QgemmU8U8KernelAvx512Vnni.asm
|
||||
${ONNXRUNTIME_ROOT}/core/mlas/lib/amd64/DgemmKernelSse2.asm
|
||||
${ONNXRUNTIME_ROOT}/core/mlas/lib/amd64/DgemmKernelAvx.asm
|
||||
${ONNXRUNTIME_ROOT}/core/mlas/lib/amd64/DgemmKernelFma3.asm
|
||||
|
|
@ -255,7 +254,6 @@ else()
|
|||
${ONNXRUNTIME_ROOT}/core/mlas/lib/x86_64/QgemmU8S8KernelAvx512Vnni.S
|
||||
${ONNXRUNTIME_ROOT}/core/mlas/lib/x86_64/QgemvU8S8KernelAvx512Vnni.S
|
||||
${ONNXRUNTIME_ROOT}/core/mlas/lib/x86_64/QgemmU8U8KernelAvx512Core.S
|
||||
${ONNXRUNTIME_ROOT}/core/mlas/lib/x86_64/QgemmU8U8KernelAvx512Vnni.S
|
||||
)
|
||||
if(HAS_AVX512CORE)
|
||||
set_source_files_properties(${mlas_platform_srcs_avx512core} PROPERTIES COMPILE_FLAGS "-mavx512bw -mavx512dq -mavx512vl")
|
||||
|
|
|
|||
|
|
@ -147,35 +147,19 @@ MlasGemm(
|
|||
MLAS_THREADPOOL* ThreadPool
|
||||
);
|
||||
|
||||
template<typename AType, typename BType>
|
||||
void
|
||||
MLASCALL
|
||||
MlasGemm(
|
||||
size_t M,
|
||||
size_t N,
|
||||
size_t K,
|
||||
const uint8_t* A,
|
||||
const AType* A,
|
||||
size_t lda,
|
||||
uint8_t offa,
|
||||
const int8_t* B,
|
||||
AType offa,
|
||||
const BType* B,
|
||||
size_t ldb,
|
||||
int8_t offb,
|
||||
int32_t* C,
|
||||
size_t ldc,
|
||||
MLAS_THREADPOOL* ThreadPool
|
||||
);
|
||||
|
||||
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,
|
||||
BType offb,
|
||||
int32_t* C,
|
||||
size_t ldc,
|
||||
MLAS_THREADPOOL* ThreadPool
|
||||
|
|
|
|||
|
|
@ -119,16 +119,16 @@ ErfKernelFrame ENDS
|
|||
|
||||
alloc_stack (ErfKernelFrame.ReturnAddress)
|
||||
|
||||
save_xmm128_avx xmm6,ErfKernelFrame.SavedXmm6
|
||||
save_xmm128_avx xmm7,ErfKernelFrame.SavedXmm7
|
||||
save_xmm128_avx xmm8,ErfKernelFrame.SavedXmm8
|
||||
save_xmm128_avx xmm9,ErfKernelFrame.SavedXmm9
|
||||
save_xmm128_avx xmm10,ErfKernelFrame.SavedXmm10
|
||||
save_xmm128_avx xmm11,ErfKernelFrame.SavedXmm11
|
||||
save_xmm128_avx xmm12,ErfKernelFrame.SavedXmm12
|
||||
save_xmm128_avx xmm13,ErfKernelFrame.SavedXmm13
|
||||
save_xmm128_avx xmm14,ErfKernelFrame.SavedXmm14
|
||||
save_xmm128_avx xmm15,ErfKernelFrame.SavedXmm15
|
||||
save_xmm128 xmm6,ErfKernelFrame.SavedXmm6
|
||||
save_xmm128 xmm7,ErfKernelFrame.SavedXmm7
|
||||
save_xmm128 xmm8,ErfKernelFrame.SavedXmm8
|
||||
save_xmm128 xmm9,ErfKernelFrame.SavedXmm9
|
||||
save_xmm128 xmm10,ErfKernelFrame.SavedXmm10
|
||||
save_xmm128 xmm11,ErfKernelFrame.SavedXmm11
|
||||
save_xmm128 xmm12,ErfKernelFrame.SavedXmm12
|
||||
save_xmm128 xmm13,ErfKernelFrame.SavedXmm13
|
||||
save_xmm128 xmm14,ErfKernelFrame.SavedXmm14
|
||||
save_xmm128 xmm15,ErfKernelFrame.SavedXmm15
|
||||
|
||||
END_PROLOGUE
|
||||
|
||||
|
|
@ -548,16 +548,16 @@ LBiggerNumbersRemaining:
|
|||
|
||||
LErfBatchExp:
|
||||
vzeroupper
|
||||
vmovaps xmm6,ErfKernelFrame.SavedXmm6[rsp]
|
||||
vmovaps xmm7,ErfKernelFrame.SavedXmm7[rsp]
|
||||
vmovaps xmm8,ErfKernelFrame.SavedXmm8[rsp]
|
||||
vmovaps xmm9,ErfKernelFrame.SavedXmm9[rsp]
|
||||
vmovaps xmm10,ErfKernelFrame.SavedXmm10[rsp]
|
||||
vmovaps xmm11,ErfKernelFrame.SavedXmm11[rsp]
|
||||
vmovaps xmm12,ErfKernelFrame.SavedXmm12[rsp]
|
||||
vmovaps xmm13,ErfKernelFrame.SavedXmm13[rsp]
|
||||
vmovaps xmm14,ErfKernelFrame.SavedXmm14[rsp]
|
||||
vmovaps xmm15,ErfKernelFrame.SavedXmm15[rsp]
|
||||
movaps xmm6,ErfKernelFrame.SavedXmm6[rsp]
|
||||
movaps xmm7,ErfKernelFrame.SavedXmm7[rsp]
|
||||
movaps xmm8,ErfKernelFrame.SavedXmm8[rsp]
|
||||
movaps xmm9,ErfKernelFrame.SavedXmm9[rsp]
|
||||
movaps xmm10,ErfKernelFrame.SavedXmm10[rsp]
|
||||
movaps xmm11,ErfKernelFrame.SavedXmm11[rsp]
|
||||
movaps xmm12,ErfKernelFrame.SavedXmm12[rsp]
|
||||
movaps xmm13,ErfKernelFrame.SavedXmm13[rsp]
|
||||
movaps xmm14,ErfKernelFrame.SavedXmm14[rsp]
|
||||
movaps xmm15,ErfKernelFrame.SavedXmm15[rsp]
|
||||
add rsp,(ErfKernelFrame.ReturnAddress)
|
||||
|
||||
BEGIN_EPILOGUE
|
||||
|
|
|
|||
|
|
@ -13,7 +13,7 @@
|
|||
; This module implements the kernels for the floating point matrix/matrix
|
||||
; multiply operation (SGEMM and DGEMM).
|
||||
;
|
||||
; This implementation uses AVX fused multiply/add instructions.
|
||||
; This implementation uses SSE2 instructions.
|
||||
;
|
||||
;--
|
||||
|
||||
|
|
|
|||
|
|
@ -98,16 +98,16 @@ LogisticKernelFrame ENDS
|
|||
|
||||
alloc_stack (LogisticKernelFrame.ReturnAddress)
|
||||
|
||||
save_xmm128_avx xmm6,LogisticKernelFrame.SavedXmm6
|
||||
save_xmm128_avx xmm7,LogisticKernelFrame.SavedXmm7
|
||||
save_xmm128_avx xmm8,LogisticKernelFrame.SavedXmm8
|
||||
save_xmm128_avx xmm9,LogisticKernelFrame.SavedXmm9
|
||||
save_xmm128_avx xmm10,LogisticKernelFrame.SavedXmm10
|
||||
save_xmm128_avx xmm11,LogisticKernelFrame.SavedXmm11
|
||||
save_xmm128_avx xmm12,LogisticKernelFrame.SavedXmm12
|
||||
save_xmm128_avx xmm13,LogisticKernelFrame.SavedXmm13
|
||||
save_xmm128_avx xmm14,LogisticKernelFrame.SavedXmm14
|
||||
save_xmm128_avx xmm15,LogisticKernelFrame.SavedXmm15
|
||||
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
|
||||
|
||||
END_PROLOGUE
|
||||
|
||||
|
|
@ -185,16 +185,16 @@ ProcessRemainingCount:
|
|||
|
||||
ExitKernel:
|
||||
vzeroupper
|
||||
vmovaps xmm6,LogisticKernelFrame.SavedXmm6[rsp]
|
||||
vmovaps xmm7,LogisticKernelFrame.SavedXmm7[rsp]
|
||||
vmovaps xmm8,LogisticKernelFrame.SavedXmm8[rsp]
|
||||
vmovaps xmm9,LogisticKernelFrame.SavedXmm9[rsp]
|
||||
vmovaps xmm10,LogisticKernelFrame.SavedXmm10[rsp]
|
||||
vmovaps xmm11,LogisticKernelFrame.SavedXmm11[rsp]
|
||||
vmovaps xmm12,LogisticKernelFrame.SavedXmm12[rsp]
|
||||
vmovaps xmm13,LogisticKernelFrame.SavedXmm13[rsp]
|
||||
vmovaps xmm14,LogisticKernelFrame.SavedXmm14[rsp]
|
||||
vmovaps xmm15,LogisticKernelFrame.SavedXmm15[rsp]
|
||||
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)
|
||||
|
||||
BEGIN_EPILOGUE
|
||||
|
|
|
|||
|
|
@ -62,6 +62,7 @@ GemmU8S8CopyPackBFrame STRUCT
|
|||
SavedXmm6 OWORD ?
|
||||
SavedXmm7 OWORD ?
|
||||
SavedXmm8 OWORD ?
|
||||
SavedXmm9 OWORD ?
|
||||
Padding QWORD ?
|
||||
SavedRdi QWORD ?
|
||||
SavedRsi QWORD ?
|
||||
|
|
@ -75,6 +76,7 @@ GemmU8S8CopyPackBFrame STRUCT
|
|||
CountK QWORD ?
|
||||
ColumnSumVector QWORD ?
|
||||
offa QWORD ?
|
||||
BTypeIsSigned QWORD ?
|
||||
|
||||
GemmU8S8CopyPackBFrame ENDS
|
||||
|
||||
|
|
@ -481,6 +483,9 @@ ExitRoutine:
|
|||
; offa - Supplies the zero point offset for the other source matrix of the
|
||||
; matrix multiplication.
|
||||
;
|
||||
; BTypeIsSigned - Supplies true if the source matrix is signed data, else
|
||||
; false if the the source matrix is unsigned data.
|
||||
;
|
||||
; Return Value:
|
||||
;
|
||||
; None.
|
||||
|
|
@ -497,6 +502,7 @@ ExitRoutine:
|
|||
save_xmm128 xmm6,GemmU8S8CopyPackBFrame.SavedXmm6
|
||||
save_xmm128 xmm7,GemmU8S8CopyPackBFrame.SavedXmm7
|
||||
save_xmm128 xmm8,GemmU8S8CopyPackBFrame.SavedXmm8
|
||||
save_xmm128 xmm9,GemmU8S8CopyPackBFrame.SavedXmm9
|
||||
|
||||
END_PROLOGUE
|
||||
|
||||
|
|
@ -510,6 +516,17 @@ ExitRoutine:
|
|||
vpsllw ymm0,ymm8,8 ; generate word vector [0x0100]
|
||||
vpor ymm8,ymm8,ymm0 ; generate word vector [0x0101]
|
||||
|
||||
;
|
||||
; Compute the bit flip vector to adjust input from U8 to S8.
|
||||
;
|
||||
|
||||
vpxor xmm9,xmm9,xmm9 ; generate word vector [0x0000]
|
||||
cmp BYTE PTR GemmU8S8CopyPackBFrame.BTypeIsSigned[rsp],0
|
||||
jnz SkipUnsignedBitFlipVector
|
||||
vpsllw ymm9,ymm8,7 ; generate word vector [0x8080]
|
||||
|
||||
SkipUnsignedBitFlipVector:
|
||||
|
||||
;
|
||||
; Process 16 columns of matrix B in a loop.
|
||||
;
|
||||
|
|
@ -544,6 +561,8 @@ InterleaveRowDataN16:
|
|||
vpunpckhwd xmm3,xmm3,xmm5
|
||||
vinserti128 ymm4,ymm4,xmm6,1
|
||||
vinserti128 ymm2,ymm2,xmm3,1
|
||||
vpxor ymm4,ymm4,ymm9 ; optionally adjust unsigned data
|
||||
vpxor ymm2,ymm2,ymm9
|
||||
vmovdqu YMMWORD PTR [rcx],ymm4 ; store interleaved rows
|
||||
vmovdqu YMMWORD PTR [rcx+32],ymm2
|
||||
vpmaddubsw ymm4,ymm8,ymm4 ; horizontal byte+byte=word per row
|
||||
|
|
@ -562,9 +581,9 @@ ProcessRemainingRowsN16:
|
|||
add rbx,4 ; correct for over-subtract above
|
||||
jz ReduceColumnSumVectorN16
|
||||
vmovdqu xmm2,XMMWORD PTR [rdx]
|
||||
vpxor xmm3,xmm3,xmm3
|
||||
vpxor xmm4,xmm4,xmm4
|
||||
vpxor xmm5,xmm5,xmm5
|
||||
vmovaps xmm3,xmm9
|
||||
vmovaps xmm4,xmm9
|
||||
vmovaps xmm5,xmm9
|
||||
xor ebx,ebx ; no more rows remaining
|
||||
test r10b,2 ; (CountK & 2) != 0?
|
||||
jz InterleaveRowDataN16
|
||||
|
|
@ -596,6 +615,7 @@ ExitRoutine:
|
|||
movaps xmm6,GemmU8S8CopyPackBFrame.SavedXmm6[rsp]
|
||||
movaps xmm7,GemmU8S8CopyPackBFrame.SavedXmm7[rsp]
|
||||
movaps xmm8,GemmU8S8CopyPackBFrame.SavedXmm8[rsp]
|
||||
movaps xmm9,GemmU8S8CopyPackBFrame.SavedXmm9[rsp]
|
||||
add rsp,(GemmU8S8CopyPackBFrame.SavedRdi)
|
||||
|
||||
BEGIN_EPILOGUE
|
||||
|
|
@ -613,8 +633,8 @@ ExitRoutine:
|
|||
ProcessColumnNUnaligned:
|
||||
vpxor xmm0,xmm0,xmm0 ; clear column accumulators
|
||||
vpxor xmm1,xmm1,xmm1
|
||||
vmovdqu YMMWORD PTR GemmU8S8CopyPackBFrame.PaddedMatrixBData[rsp],ymm0
|
||||
vmovdqu YMMWORD PTR GemmU8S8CopyPackBFrame.PaddedMatrixBData[rsp+32],ymm0
|
||||
vmovdqu YMMWORD PTR GemmU8S8CopyPackBFrame.PaddedMatrixBData[rsp],ymm9
|
||||
vmovdqu YMMWORD PTR GemmU8S8CopyPackBFrame.PaddedMatrixBData[rsp+32],ymm9
|
||||
sub r10,4
|
||||
jb ProcessRemainingRowsNUnaligned
|
||||
|
||||
|
|
@ -690,6 +710,8 @@ ProcessPaddedMatrixBData:
|
|||
vpunpckhwd xmm3,xmm3,xmm5
|
||||
vinserti128 ymm4,ymm4,xmm6,1
|
||||
vinserti128 ymm2,ymm2,xmm3,1
|
||||
vpxor ymm4,ymm4,ymm9 ; optionally adjust unsigned data
|
||||
vpxor ymm2,ymm2,ymm9
|
||||
vmovdqu YMMWORD PTR [rcx],ymm4 ; store interleaved rows
|
||||
vmovdqu YMMWORD PTR [rcx+32],ymm2
|
||||
vpmaddubsw ymm4,ymm8,ymm4 ; horizontal byte+byte=word per row
|
||||
|
|
@ -711,9 +733,8 @@ ProcessRemainingRowsNUnaligned:
|
|||
|
||||
.errnz GemmU8S8CopyPackBFrame.PaddedMatrixBData
|
||||
mov rbp,rsp ; GemmU8S8CopyPackBFrame.PaddedMatrixBData
|
||||
vpxor xmm6,xmm6,xmm6
|
||||
vmovdqu YMMWORD PTR [rbp],ymm6
|
||||
vmovdqu YMMWORD PTR [rbp+32],ymm6
|
||||
vmovdqu YMMWORD PTR [rbp],ymm9
|
||||
vmovdqu YMMWORD PTR [rbp+32],ymm9
|
||||
|
||||
CopyUnalignedRowLoop:
|
||||
lea rdi,[rbp+16] ; advance next padded buffer by 16 bytes
|
||||
|
|
|
|||
|
|
@ -1,64 +0,0 @@
|
|||
;++
|
||||
;
|
||||
; Copyright (c) Microsoft Corporation. All rights reserved.
|
||||
;
|
||||
; Licensed under the MIT License.
|
||||
;
|
||||
; Module Name:
|
||||
;
|
||||
; QgemmU8U8KernelAvx512Common.inc
|
||||
;
|
||||
; Abstract:
|
||||
;
|
||||
; This module contains common kernel macros and structures for the quantized
|
||||
; integer matrix/matrix multiply operation (QGEMM) for the AVX512 core and
|
||||
; AVX512VNNI kernels.
|
||||
;
|
||||
;--
|
||||
|
||||
INCLUDE QgemmU8X8KernelAvx512Common.inc
|
||||
|
||||
;
|
||||
; Macro Description:
|
||||
;
|
||||
; This macro generates code to execute the block compute macro multiple
|
||||
; times and advancing the matrix A and matrix B data pointers.
|
||||
;
|
||||
; Arguments:
|
||||
;
|
||||
; ColumnCount - Supplies the number of columns to produce.
|
||||
;
|
||||
; RowCount - Supplies the number of rows to produce.
|
||||
;
|
||||
; Implicit Arguments:
|
||||
;
|
||||
; rbx - Supplies the address into the matrix A data plus 3 rows.
|
||||
;
|
||||
; rcx - Supplies the address into the matrix A data.
|
||||
;
|
||||
; rdx - Supplies the address into the matrix B data.
|
||||
;
|
||||
; r9 - Supplies the length in bytes of a row from matrix A.
|
||||
;
|
||||
; r14 - Supplies the stride in bytes of between packed blocks of matrix B.
|
||||
;
|
||||
; zmm14-zmm31 - Supplies the block accumulators.
|
||||
;
|
||||
|
||||
ComputeBlockLoop MACRO ColumnCount, RowCount
|
||||
|
||||
LOCAL ComputeBlockBy1Loop
|
||||
|
||||
mov rsi,r9 ; reload row length remaining
|
||||
|
||||
ComputeBlockBy1Loop:
|
||||
ComputeBlock ColumnCount, RowCount, 0, 0
|
||||
add rcx,4 ; advance matrix A by 1 pair
|
||||
IF RowCount GT 3
|
||||
add rbx,4 ; advance matrix A plus 3 rows by 1 pair
|
||||
ENDIF
|
||||
add rdx,32 ; advance matrix B
|
||||
sub rsi,4
|
||||
jnz ComputeBlockBy1Loop
|
||||
|
||||
ENDM
|
||||
|
|
@ -19,7 +19,7 @@
|
|||
|
||||
.xlist
|
||||
INCLUDE mlasi.inc
|
||||
INCLUDE QgemmU8U8KernelAvx512Common.inc
|
||||
INCLUDE QgemmU8X8KernelAvx512Common.inc
|
||||
.list
|
||||
|
||||
;
|
||||
|
|
@ -118,6 +118,51 @@ ENDIF
|
|||
|
||||
ENDM
|
||||
|
||||
;
|
||||
; Macro Description:
|
||||
;
|
||||
; This macro generates code to execute the block compute macro multiple
|
||||
; times and advancing the matrix A and matrix B data pointers.
|
||||
;
|
||||
; Arguments:
|
||||
;
|
||||
; ColumnCount - Supplies the number of columns to produce.
|
||||
;
|
||||
; RowCount - Supplies the number of rows to produce.
|
||||
;
|
||||
; Implicit Arguments:
|
||||
;
|
||||
; rbx - Supplies the address into the matrix A data plus 3 rows.
|
||||
;
|
||||
; rcx - Supplies the address into the matrix A data.
|
||||
;
|
||||
; rdx - Supplies the address into the matrix B data.
|
||||
;
|
||||
; r9 - Supplies the length in bytes of a row from matrix A.
|
||||
;
|
||||
; r14 - Supplies the stride in bytes of between packed blocks of matrix B.
|
||||
;
|
||||
; zmm14-zmm31 - Supplies the block accumulators.
|
||||
;
|
||||
|
||||
ComputeBlockLoop MACRO ColumnCount, RowCount
|
||||
|
||||
LOCAL ComputeBlockBy1Loop
|
||||
|
||||
mov rsi,r9 ; reload row length remaining
|
||||
|
||||
ComputeBlockBy1Loop:
|
||||
ComputeBlock ColumnCount, RowCount, 0, 0
|
||||
add rcx,4 ; advance matrix A by 1 pair
|
||||
IF RowCount GT 3
|
||||
add rbx,4 ; advance matrix A plus 3 rows by 1 pair
|
||||
ENDIF
|
||||
add rdx,32 ; advance matrix B
|
||||
sub rsi,4
|
||||
jnz ComputeBlockBy1Loop
|
||||
|
||||
ENDM
|
||||
|
||||
;
|
||||
; Generate the GEMM kernel.
|
||||
;
|
||||
|
|
|
|||
|
|
@ -1,110 +0,0 @@
|
|||
;++
|
||||
;
|
||||
; Copyright (c) Microsoft Corporation. All rights reserved.
|
||||
;
|
||||
; Licensed under the MIT License.
|
||||
;
|
||||
; Module Name:
|
||||
;
|
||||
; QgemmU8U8KernelAvx512Vnni.asm
|
||||
;
|
||||
; Abstract:
|
||||
;
|
||||
; This module implements the kernels for the quantized integer matrix/matrix
|
||||
; multiply operation (QGEMM).
|
||||
;
|
||||
; This implementation uses AVX512VNNI instructions.
|
||||
;
|
||||
;--
|
||||
|
||||
.xlist
|
||||
INCLUDE mlasi.inc
|
||||
INCLUDE QgemmU8U8KernelAvx512Common.inc
|
||||
INCLUDE AssembleAvx512Vnni.inc
|
||||
.list
|
||||
|
||||
;
|
||||
; Macro Description:
|
||||
;
|
||||
; This macro generates code to multiply and accumulate each row of the output
|
||||
; block.
|
||||
;
|
||||
; Arguments:
|
||||
;
|
||||
; ColumnCount - Supplies the number of columns to produce.
|
||||
;
|
||||
; RowCount - Supplies the number of rows to produce.
|
||||
;
|
||||
; VectorOffset - Supplies the byte offset from matrix B to fetch elements.
|
||||
;
|
||||
; BroadcastOffset - Supplies the byte offset from matrix A to fetch elements.
|
||||
;
|
||||
; Implicit Arguments:
|
||||
;
|
||||
; rbx - Supplies the address into the matrix A data plus 3 rows.
|
||||
;
|
||||
; rcx - Supplies the address into the matrix A data.
|
||||
;
|
||||
; rdx - Supplies the address into the matrix B data.
|
||||
;
|
||||
; r9 - Supplies the length in bytes of a row from matrix A.
|
||||
;
|
||||
; r14 - Supplies the stride in bytes of between packed blocks of matrix B.
|
||||
;
|
||||
; zmm14-zmm31 - Supplies the block accumulators.
|
||||
;
|
||||
|
||||
ComputeBlock MACRO ColumnCount, RowCount, VectorOffset, BroadcastOffset
|
||||
|
||||
IF ColumnCount GE 32
|
||||
IF ColumnCount GE 48
|
||||
vpmovzxbw zmm0,YMMWORD PTR [rdx+VectorOffset]
|
||||
vpmovzxbw zmm1,YMMWORD PTR [rdx+r14+VectorOffset]
|
||||
vpmovzxbw zmm2,YMMWORD PTR [rdx+r14*2+VectorOffset]
|
||||
ELSE
|
||||
vpmovzxbw zmm1,YMMWORD PTR [rdx+VectorOffset]
|
||||
vpmovzxbw zmm2,YMMWORD PTR [rdx+r14+VectorOffset]
|
||||
ENDIF
|
||||
EmitIfCountGE RowCount, 1, <vpbroadcastd zmm3,DWORD PTR [rcx+BroadcastOffset]>
|
||||
EmitIfCount2GE RowCount, 1, ColumnCount, 48, <VpdpwssdZmmZmmZmm zmm26,zmm3,zmm0>
|
||||
EmitIfCount2GE RowCount, 1, ColumnCount, 32, <VpdpwssdZmmZmmZmm zmm20,zmm3,zmm1>
|
||||
EmitIfCount2GE RowCount, 1, ColumnCount, 16, <VpdpwssdZmmZmmZmm zmm14,zmm3,zmm2>
|
||||
EmitIfCountGE RowCount, 2, <vpbroadcastd zmm3,DWORD PTR [rcx+r9+BroadcastOffset]>
|
||||
EmitIfCount2GE RowCount, 2, ColumnCount, 48, <VpdpwssdZmmZmmZmm zmm27,zmm3,zmm0>
|
||||
EmitIfCount2GE RowCount, 2, ColumnCount, 32, <VpdpwssdZmmZmmZmm zmm21,zmm3,zmm1>
|
||||
EmitIfCount2GE RowCount, 2, ColumnCount, 16, <VpdpwssdZmmZmmZmm zmm15,zmm3,zmm2>
|
||||
EmitIfCountGE RowCount, 3, <vpbroadcastd zmm3,DWORD PTR [rcx+r9*2+BroadcastOffset]>
|
||||
EmitIfCount2GE RowCount, 3, ColumnCount, 48, <VpdpwssdZmmZmmZmm zmm28,zmm3,zmm0>
|
||||
EmitIfCount2GE RowCount, 3, ColumnCount, 32, <VpdpwssdZmmZmmZmm zmm22,zmm3,zmm1>
|
||||
EmitIfCount2GE RowCount, 3, ColumnCount, 16, <VpdpwssdZmmZmmZmm zmm16,zmm3,zmm2>
|
||||
EmitIfCountGE RowCount, 4, <vpbroadcastd zmm3,DWORD PTR [rbx+BroadcastOffset]>
|
||||
EmitIfCount2GE RowCount, 4, ColumnCount, 48, <VpdpwssdZmmZmmZmm zmm29,zmm3,zmm0>
|
||||
EmitIfCount2GE RowCount, 4, ColumnCount, 32, <VpdpwssdZmmZmmZmm zmm23,zmm3,zmm1>
|
||||
EmitIfCount2GE RowCount, 4, ColumnCount, 16, <VpdpwssdZmmZmmZmm zmm17,zmm3,zmm2>
|
||||
EmitIfCountGE RowCount, 5, <vpbroadcastd zmm3,DWORD PTR [rbx+r9+BroadcastOffset]>
|
||||
EmitIfCount2GE RowCount, 5, ColumnCount, 48, <VpdpwssdZmmZmmZmm zmm30,zmm3,zmm0>
|
||||
EmitIfCount2GE RowCount, 5, ColumnCount, 32, <VpdpwssdZmmZmmZmm zmm24,zmm3,zmm1>
|
||||
EmitIfCount2GE RowCount, 5, ColumnCount, 16, <VpdpwssdZmmZmmZmm zmm18,zmm3,zmm2>
|
||||
EmitIfCountGE RowCount, 6, <vpbroadcastd zmm3,DWORD PTR [rbx+r9*2+BroadcastOffset]>
|
||||
EmitIfCount2GE RowCount, 6, ColumnCount, 48, <VpdpwssdZmmZmmZmm zmm31,zmm3,zmm0>
|
||||
EmitIfCount2GE RowCount, 6, ColumnCount, 32, <VpdpwssdZmmZmmZmm zmm25,zmm3,zmm1>
|
||||
EmitIfCount2GE RowCount, 6, ColumnCount, 16, <VpdpwssdZmmZmmZmm zmm19,zmm3,zmm2>
|
||||
ELSE
|
||||
vpmovzxbw zmm2,YMMWORD PTR [rdx+VectorOffset]
|
||||
EmitIfCountGE RowCount, 1, <VpdpwssdZmmZmmBroadcast zmm14,zmm2,rcx,BroadcastOffset>
|
||||
EmitIfCountGE RowCount, 2, <VpdpwssdZmmZmmBroadcast zmm15,zmm2,rcx,BroadcastOffset,r9,1>
|
||||
EmitIfCountGE RowCount, 3, <VpdpwssdZmmZmmBroadcast zmm16,zmm2,rcx,BroadcastOffset,r9,2>
|
||||
EmitIfCountGE RowCount, 4, <VpdpwssdZmmZmmBroadcast zmm17,zmm2,rbx,BroadcastOffset>
|
||||
EmitIfCountGE RowCount, 5, <VpdpwssdZmmZmmBroadcast zmm18,zmm2,rbx,BroadcastOffset,r9,1>
|
||||
EmitIfCountGE RowCount, 6, <VpdpwssdZmmZmmBroadcast zmm19,zmm2,rbx,BroadcastOffset,r9,2>
|
||||
ENDIF
|
||||
|
||||
ENDM
|
||||
|
||||
;
|
||||
; Generate the GEMM kernel.
|
||||
;
|
||||
|
||||
GemmU8X8KernelAvx512Function U8U8, Avx512Vnni
|
||||
|
||||
END
|
||||
|
|
@ -128,8 +128,8 @@ IF RowCount GT 3
|
|||
ENDIF
|
||||
ComputeBlockLoop ColumnCount, RowCount
|
||||
IF RowCount GT 3
|
||||
lea rbx,[r8+rax*2] ; compute matrix C plus 3 rows
|
||||
add rbx,rax
|
||||
lea rbx,[rax*2+rax]
|
||||
add rbx,r8 ; compute matrix C plus 3 rows
|
||||
ENDIF
|
||||
|
||||
ENDM
|
||||
|
|
|
|||
|
|
@ -85,9 +85,9 @@ SgemmKernelM1Frame ENDS
|
|||
push_reg rbx
|
||||
push_reg rsi
|
||||
alloc_stack (SgemmKernelM1Frame.SavedRsi)
|
||||
save_xmm128_avx xmm6,SgemmKernelM1Frame.SavedXmm6
|
||||
save_xmm128_avx xmm7,SgemmKernelM1Frame.SavedXmm7
|
||||
save_xmm128_avx xmm8,SgemmKernelM1Frame.SavedXmm8
|
||||
save_xmm128 xmm6,SgemmKernelM1Frame.SavedXmm6
|
||||
save_xmm128 xmm7,SgemmKernelM1Frame.SavedXmm7
|
||||
save_xmm128 xmm8,SgemmKernelM1Frame.SavedXmm8
|
||||
|
||||
END_PROLOGUE
|
||||
|
||||
|
|
@ -217,9 +217,9 @@ ProcessRemainingCountK:
|
|||
|
||||
ExitKernel:
|
||||
vzeroupper
|
||||
vmovaps xmm6,SgemmKernelM1Frame.SavedXmm6[rsp]
|
||||
vmovaps xmm7,SgemmKernelM1Frame.SavedXmm7[rsp]
|
||||
vmovaps xmm8,SgemmKernelM1Frame.SavedXmm8[rsp]
|
||||
movaps xmm6,SgemmKernelM1Frame.SavedXmm6[rsp]
|
||||
movaps xmm7,SgemmKernelM1Frame.SavedXmm7[rsp]
|
||||
movaps xmm8,SgemmKernelM1Frame.SavedXmm8[rsp]
|
||||
add rsp,(SgemmKernelM1Frame.SavedRsi)
|
||||
|
||||
BEGIN_EPILOGUE
|
||||
|
|
@ -348,8 +348,8 @@ ProcessRemainingCountN1:
|
|||
push_reg rbx
|
||||
push_reg rsi
|
||||
alloc_stack (SgemmKernelM1Frame.SavedRsi)
|
||||
save_xmm128_avx xmm6,SgemmKernelM1Frame.SavedXmm6
|
||||
save_xmm128_avx xmm7,SgemmKernelM1Frame.SavedXmm7
|
||||
save_xmm128 xmm6,SgemmKernelM1Frame.SavedXmm6
|
||||
save_xmm128 xmm7,SgemmKernelM1Frame.SavedXmm7
|
||||
|
||||
END_PROLOGUE
|
||||
|
||||
|
|
@ -465,8 +465,8 @@ ProcessRemainingCountN:
|
|||
|
||||
ExitKernel:
|
||||
vzeroupper
|
||||
vmovaps xmm6,SgemmKernelM1Frame.SavedXmm6[rsp]
|
||||
vmovaps xmm7,SgemmKernelM1Frame.SavedXmm7[rsp]
|
||||
movaps xmm6,SgemmKernelM1Frame.SavedXmm6[rsp]
|
||||
movaps xmm7,SgemmKernelM1Frame.SavedXmm7[rsp]
|
||||
add rsp,(SgemmKernelM1Frame.SavedRsi)
|
||||
|
||||
BEGIN_EPILOGUE
|
||||
|
|
|
|||
|
|
@ -98,16 +98,16 @@ TanhKernelFrame ENDS
|
|||
|
||||
alloc_stack (TanhKernelFrame.ReturnAddress)
|
||||
|
||||
save_xmm128_avx xmm6,TanhKernelFrame.SavedXmm6
|
||||
save_xmm128_avx xmm7,TanhKernelFrame.SavedXmm7
|
||||
save_xmm128_avx xmm8,TanhKernelFrame.SavedXmm8
|
||||
save_xmm128_avx xmm9,TanhKernelFrame.SavedXmm9
|
||||
save_xmm128_avx xmm10,TanhKernelFrame.SavedXmm10
|
||||
save_xmm128_avx xmm11,TanhKernelFrame.SavedXmm11
|
||||
save_xmm128_avx xmm12,TanhKernelFrame.SavedXmm12
|
||||
save_xmm128_avx xmm13,TanhKernelFrame.SavedXmm13
|
||||
save_xmm128_avx xmm14,TanhKernelFrame.SavedXmm14
|
||||
save_xmm128_avx xmm15,TanhKernelFrame.SavedXmm15
|
||||
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
|
||||
|
||||
END_PROLOGUE
|
||||
|
||||
|
|
@ -177,16 +177,16 @@ ProcessRemainingCount:
|
|||
|
||||
ExitKernel:
|
||||
vzeroupper
|
||||
vmovaps xmm6,TanhKernelFrame.SavedXmm6[rsp]
|
||||
vmovaps xmm7,TanhKernelFrame.SavedXmm7[rsp]
|
||||
vmovaps xmm8,TanhKernelFrame.SavedXmm8[rsp]
|
||||
vmovaps xmm9,TanhKernelFrame.SavedXmm9[rsp]
|
||||
vmovaps xmm10,TanhKernelFrame.SavedXmm10[rsp]
|
||||
vmovaps xmm11,TanhKernelFrame.SavedXmm11[rsp]
|
||||
vmovaps xmm12,TanhKernelFrame.SavedXmm12[rsp]
|
||||
vmovaps xmm13,TanhKernelFrame.SavedXmm13[rsp]
|
||||
vmovaps xmm14,TanhKernelFrame.SavedXmm14[rsp]
|
||||
vmovaps xmm15,TanhKernelFrame.SavedXmm15[rsp]
|
||||
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)
|
||||
|
||||
BEGIN_EPILOGUE
|
||||
|
|
|
|||
|
|
@ -19,27 +19,6 @@
|
|||
INCLUDE macamd64.inc
|
||||
.list
|
||||
|
||||
;
|
||||
; Macro Description:
|
||||
;
|
||||
; This macro uses AVX instructions to save a vector register as part of a
|
||||
; function prologue as an alternative to save_xmm128.
|
||||
;
|
||||
; Arguments:
|
||||
;
|
||||
; Register - Supplies the vector register to be saved.
|
||||
;
|
||||
; Offset - Supplies the stack frame offset to store the contents of the
|
||||
; vector register.
|
||||
;
|
||||
|
||||
save_xmm128_avx MACRO Register, Offset
|
||||
|
||||
vmovaps Offset[rsp], Register
|
||||
.savexmm128 Register, Offset
|
||||
|
||||
ENDM
|
||||
|
||||
;
|
||||
; Macro Description:
|
||||
;
|
||||
|
|
|
|||
|
|
@ -22,6 +22,7 @@ Abstract:
|
|||
#include <algorithm>
|
||||
#include <limits>
|
||||
#include <cmath>
|
||||
#include <type_traits>
|
||||
|
||||
#if defined(_WIN32)
|
||||
#include <windows.h>
|
||||
|
|
@ -246,31 +247,22 @@ typedef MLAS_SGEMM_TRANSPOSE_PACKB_BLOCK_ROUTINE* PMLAS_SGEMM_TRANSPOSE_PACKB_BL
|
|||
|
||||
typedef
|
||||
void
|
||||
(MLASCALL MLAS_GEMM_U8S8_COPY_PACKA_ROUTINE)(
|
||||
uint8_t* D,
|
||||
(MLASCALL MLAS_GEMM_U8X8_OPERATION)(
|
||||
const struct MLAS_GEMM_U8X8_WORK_BLOCK* WorkBlock,
|
||||
size_t M,
|
||||
size_t N,
|
||||
size_t K,
|
||||
const uint8_t* A,
|
||||
size_t lda,
|
||||
size_t CountM,
|
||||
size_t CountK,
|
||||
int32_t* RowSumVector,
|
||||
int16_t offb
|
||||
);
|
||||
|
||||
typedef MLAS_GEMM_U8S8_COPY_PACKA_ROUTINE* PMLAS_GEMM_U8S8_COPY_PACKA_ROUTINE;
|
||||
|
||||
typedef
|
||||
void
|
||||
(MLASCALL MLAS_GEMM_U8S8_COPY_PACKB_ROUTINE)(
|
||||
int8_t* D,
|
||||
const int8_t* B,
|
||||
int16_t offa,
|
||||
const uint8_t* B,
|
||||
size_t ldb,
|
||||
size_t CountN,
|
||||
size_t CountK,
|
||||
int32_t* ColumnSumVector,
|
||||
int16_t offa
|
||||
int16_t offb,
|
||||
int32_t* C,
|
||||
size_t ldc
|
||||
);
|
||||
|
||||
typedef MLAS_GEMM_U8S8_COPY_PACKB_ROUTINE* PMLAS_GEMM_U8S8_COPY_PACKB_ROUTINE;
|
||||
typedef MLAS_GEMM_U8X8_OPERATION* PMLAS_GEMM_U8X8_OPERATION;
|
||||
|
||||
typedef
|
||||
size_t
|
||||
|
|
@ -303,34 +295,6 @@ size_t
|
|||
|
||||
typedef MLAS_GEMV_U8S8_KERNEL* PMLAS_GEMV_U8S8_KERNEL;
|
||||
|
||||
typedef
|
||||
void
|
||||
(MLASCALL MLAS_GEMM_U8U8_COPY_PACKA_ROUTINE)(
|
||||
int16_t* D,
|
||||
const uint8_t* A,
|
||||
size_t lda,
|
||||
size_t CountM,
|
||||
size_t CountK,
|
||||
int32_t* RowSumVector,
|
||||
int16_t offb
|
||||
);
|
||||
|
||||
typedef MLAS_GEMM_U8U8_COPY_PACKA_ROUTINE* PMLAS_GEMM_U8U8_COPY_PACKA_ROUTINE;
|
||||
|
||||
typedef
|
||||
void
|
||||
(MLASCALL MLAS_GEMM_U8U8_COPY_PACKB_ROUTINE)(
|
||||
uint8_t* D,
|
||||
const uint8_t* B,
|
||||
size_t ldb,
|
||||
size_t CountN,
|
||||
size_t CountK,
|
||||
int32_t* ColumnSumVector,
|
||||
int16_t offa
|
||||
);
|
||||
|
||||
typedef MLAS_GEMM_U8U8_COPY_PACKB_ROUTINE* PMLAS_GEMM_U8U8_COPY_PACKB_ROUTINE;
|
||||
|
||||
typedef
|
||||
size_t
|
||||
(MLASCALL MLAS_GEMM_U8U8_KERNEL)(
|
||||
|
|
@ -349,24 +313,6 @@ size_t
|
|||
|
||||
typedef MLAS_GEMM_U8U8_KERNEL* PMLAS_GEMM_U8U8_KERNEL;
|
||||
|
||||
typedef
|
||||
void
|
||||
(MLASCALL MLAS_GEMM_X8X8_OPERATION)(
|
||||
size_t M,
|
||||
size_t N,
|
||||
size_t K,
|
||||
const uint8_t* A,
|
||||
size_t lda,
|
||||
int16_t offa,
|
||||
const uint8_t* B,
|
||||
size_t ldb,
|
||||
int16_t offb,
|
||||
int32_t* C,
|
||||
size_t ldc
|
||||
);
|
||||
|
||||
typedef MLAS_GEMM_X8X8_OPERATION* PMLAS_GEMM_X8X8_OPERATION;
|
||||
|
||||
typedef
|
||||
void
|
||||
(MLASCALL MLAS_CONV_FLOAT_KERNEL)(
|
||||
|
|
@ -540,26 +486,18 @@ extern "C" {
|
|||
#endif
|
||||
|
||||
#if defined(MLAS_TARGET_AMD64_IX86)
|
||||
MLAS_GEMM_U8S8_COPY_PACKA_ROUTINE MlasGemmU8S8CopyPackASse;
|
||||
MLAS_GEMM_U8S8_COPY_PACKB_ROUTINE MlasGemmU8S8CopyPackBSse;
|
||||
MLAS_GEMM_U8S8_KERNEL MlasGemmU8S8KernelSse;
|
||||
MLAS_GEMM_U8U8_COPY_PACKA_ROUTINE MlasGemmU8U8CopyPackASse;
|
||||
MLAS_GEMM_U8U8_COPY_PACKB_ROUTINE MlasGemmU8U8CopyPackBSse;
|
||||
MLAS_GEMM_U8U8_KERNEL MlasGemmU8U8KernelSse;
|
||||
MLAS_GEMM_U8X8_OPERATION MlasGemmU8X8OperationSse;
|
||||
MLAS_GEMM_U8X8_OPERATION MlasGemmU8S8OperationAvx2;
|
||||
MLAS_GEMM_U8X8_OPERATION MlasGemmU8U8OperationAvx2;
|
||||
#if defined(MLAS_TARGET_AMD64)
|
||||
MLAS_GEMM_U8S8_COPY_PACKA_ROUTINE MlasGemmU8S8CopyPackAAvx2;
|
||||
MLAS_GEMM_U8S8_COPY_PACKB_ROUTINE MlasGemmU8S8CopyPackBAvx2;
|
||||
MLAS_GEMM_U8S8_KERNEL MlasGemmU8S8KernelAvx2;
|
||||
MLAS_GEMV_U8S8_KERNEL MlasGemvU8S8KernelAvx2;
|
||||
MLAS_GEMM_U8S8_KERNEL MlasGemmU8S8KernelAvx512Core;
|
||||
MLAS_GEMV_U8S8_KERNEL MlasGemvU8S8KernelAvx512Core;
|
||||
MLAS_GEMM_U8S8_KERNEL MlasGemmU8S8KernelAvx512Vnni;
|
||||
MLAS_GEMV_U8S8_KERNEL MlasGemvU8S8KernelAvx512Vnni;
|
||||
MLAS_GEMM_U8U8_COPY_PACKA_ROUTINE MlasGemmU8U8CopyPackAAvx2;
|
||||
MLAS_GEMM_U8U8_COPY_PACKB_ROUTINE MlasGemmU8U8CopyPackBAvx2;
|
||||
MLAS_GEMM_U8U8_KERNEL MlasGemmU8U8KernelAvx2;
|
||||
MLAS_GEMM_U8U8_KERNEL MlasGemmU8U8KernelAvx512Core;
|
||||
MLAS_GEMM_U8U8_KERNEL MlasGemmU8U8KernelAvx512Vnni;
|
||||
#endif
|
||||
#endif
|
||||
|
||||
|
|
@ -682,12 +620,8 @@ struct MLAS_PLATFORM {
|
|||
|
||||
#if defined(MLAS_TARGET_AMD64_IX86)
|
||||
PMLAS_GEMM_FLOAT_KERNEL GemmFloatKernel;
|
||||
PMLAS_GEMM_U8S8_COPY_PACKA_ROUTINE GemmU8S8CopyPackARoutine;
|
||||
PMLAS_GEMM_U8S8_COPY_PACKB_ROUTINE GemmU8S8CopyPackBRoutine;
|
||||
PMLAS_GEMM_U8S8_KERNEL GemmU8S8Kernel;
|
||||
PMLAS_GEMM_U8U8_COPY_PACKA_ROUTINE GemmU8U8CopyPackARoutine;
|
||||
PMLAS_GEMM_U8U8_COPY_PACKB_ROUTINE GemmU8U8CopyPackBRoutine;
|
||||
PMLAS_GEMM_U8U8_KERNEL GemmU8U8Kernel;
|
||||
PMLAS_GEMM_U8X8_OPERATION GemmU8S8Operation;
|
||||
PMLAS_GEMM_U8X8_OPERATION GemmU8U8Operation;
|
||||
#endif
|
||||
|
||||
#if defined(MLAS_TARGET_AMD64)
|
||||
|
|
@ -695,7 +629,9 @@ struct MLAS_PLATFORM {
|
|||
PMLAS_SGEMM_KERNEL_M1_ROUTINE KernelM1TransposeBRoutine;
|
||||
PMLAS_SGEMM_TRANSPOSE_PACKB_BLOCK_ROUTINE TransposePackB16x4Routine;
|
||||
PMLAS_GEMM_DOUBLE_KERNEL GemmDoubleKernel;
|
||||
PMLAS_GEMM_U8S8_KERNEL GemmU8S8Kernel;
|
||||
PMLAS_GEMV_U8S8_KERNEL GemvU8S8Kernel;
|
||||
PMLAS_GEMM_U8U8_KERNEL GemmU8U8Kernel;
|
||||
PMLAS_CONV_FLOAT_KERNEL ConvNchwFloatKernel;
|
||||
PMLAS_CONV_FLOAT_KERNEL ConvNchwcFloatKernel;
|
||||
PMLAS_CONV_DEPTHWISE_FLOAT_KERNEL ConvDepthwiseFloatKernel;
|
||||
|
|
@ -831,6 +767,9 @@ typedef __m128i MLAS_INT32X4;
|
|||
typedef __vector float MLAS_FLOAT32X4;
|
||||
typedef __vector int MLAS_INT32X4;
|
||||
typedef __vector unsigned MLAS_UINT32X4;
|
||||
#else
|
||||
typedef float MLAS_FLOAT32X4 __attribute__ ((vector_size(16)));
|
||||
typedef int32_t MLAS_INT32X4 __attribute__ ((vector_size(16)));
|
||||
#endif
|
||||
|
||||
MLAS_FORCEINLINE
|
||||
|
|
@ -937,19 +876,8 @@ MlasStoreAlignedFloat32x4(float* Buffer, MLAS_FLOAT32X4 Vector)
|
|||
MLAS_UNREFERENCED_PARAMETER(Buffer);
|
||||
MLAS_UNREFERENCED_PARAMETER(Vector);
|
||||
vec_st(Vector, 0, Buffer);
|
||||
#endif
|
||||
}
|
||||
|
||||
MLAS_FORCEINLINE
|
||||
void
|
||||
MlasStoreLowHalfFloat32x4(float* Buffer, MLAS_FLOAT32X4 Vector)
|
||||
{
|
||||
#if defined(MLAS_NEON_INTRINSICS)
|
||||
vst1_f32(Buffer, vget_low_f32(Vector));
|
||||
#elif defined(MLAS_SSE2_INTRINSICS)
|
||||
_mm_storel_pi((__m64*)Buffer, Vector);
|
||||
#elif defined(MLAS_VSX_INTRINSICS)
|
||||
*((int64_t*)Buffer) = ((__vector int64_t)Vector)[0];
|
||||
#else
|
||||
MlasStoreFloat32x4(Buffer, Vector);
|
||||
#endif
|
||||
}
|
||||
|
||||
|
|
@ -969,6 +897,22 @@ MlasStoreLaneFloat32x4(float* Buffer, MLAS_FLOAT32X4 Vector)
|
|||
#endif
|
||||
}
|
||||
|
||||
MLAS_FORCEINLINE
|
||||
void
|
||||
MlasStoreLowHalfFloat32x4(float* Buffer, MLAS_FLOAT32X4 Vector)
|
||||
{
|
||||
#if defined(MLAS_NEON_INTRINSICS)
|
||||
vst1_f32(Buffer, vget_low_f32(Vector));
|
||||
#elif defined(MLAS_SSE2_INTRINSICS)
|
||||
_mm_storel_pi((__m64*)Buffer, Vector);
|
||||
#elif defined(MLAS_VSX_INTRINSICS)
|
||||
*((int64_t*)Buffer) = ((__vector int64_t)Vector)[0];
|
||||
#else
|
||||
MlasStoreLaneFloat32x4<0>(&Buffer[0], Vector);
|
||||
MlasStoreLaneFloat32x4<1>(&Buffer[1], Vector);
|
||||
#endif
|
||||
}
|
||||
|
||||
template<unsigned Lane>
|
||||
MLAS_FORCEINLINE
|
||||
float
|
||||
|
|
@ -1100,8 +1044,10 @@ MlasMaximumFloat32x4(MLAS_FLOAT32X4 Vector1, MLAS_FLOAT32X4 Vector2)
|
|||
return vmaxq_f32(Vector1, Vector2);
|
||||
#elif defined(MLAS_SSE2_INTRINSICS)
|
||||
return _mm_max_ps(Vector1, Vector2);
|
||||
#else
|
||||
#elif defined(MLAS_VSX_INTRINSICS)
|
||||
return vec_sel(Vector2, Vector1, vec_cmpgt(Vector1, Vector2));
|
||||
#else
|
||||
#error Unsupported architecture.
|
||||
#endif
|
||||
}
|
||||
|
||||
|
|
@ -1113,8 +1059,10 @@ MlasMinimumFloat32x4(MLAS_FLOAT32X4 Vector1, MLAS_FLOAT32X4 Vector2)
|
|||
return vminq_f32(Vector1, Vector2);
|
||||
#elif defined(MLAS_SSE2_INTRINSICS)
|
||||
return _mm_min_ps(Vector1, Vector2);
|
||||
#else
|
||||
#elif defined(MLAS_VSX_INTRINSICS)
|
||||
return vec_sel(Vector2, Vector1, vec_cmpgt(Vector2, Vector1));
|
||||
#else
|
||||
#error Unsupported architecture.
|
||||
#endif
|
||||
}
|
||||
|
||||
|
|
@ -1298,10 +1246,8 @@ MlasShiftLeftInt32x4(MLAS_INT32X4 Vector)
|
|||
return vshlq_n_s32(Vector, ShiftCount);
|
||||
#elif defined(MLAS_SSE2_INTRINSICS)
|
||||
return _mm_slli_epi32(Vector, ShiftCount);
|
||||
#elif defined(MLAS_VSX_INTRINSICS)
|
||||
return vec_sl(Vector, MLAS_UINT32X4(MlasBroadcastInt32x4(ShiftCount)));
|
||||
#else
|
||||
#error Unsupported architecture.
|
||||
return Vector << ShiftCount;
|
||||
#endif
|
||||
}
|
||||
|
||||
|
|
|
|||
|
|
@ -115,12 +115,8 @@ Return Value:
|
|||
//
|
||||
|
||||
this->GemmFloatKernel = MlasGemmFloatKernelSse;
|
||||
this->GemmU8S8CopyPackARoutine = MlasGemmU8S8CopyPackASse;
|
||||
this->GemmU8S8CopyPackBRoutine = MlasGemmU8S8CopyPackBSse;
|
||||
this->GemmU8S8Kernel = MlasGemmU8S8KernelSse;
|
||||
this->GemmU8U8CopyPackARoutine = MlasGemmU8U8CopyPackASse;
|
||||
this->GemmU8U8CopyPackBRoutine = MlasGemmU8U8CopyPackBSse;
|
||||
this->GemmU8U8Kernel = MlasGemmU8U8KernelSse;
|
||||
this->GemmU8S8Operation = MlasGemmU8X8OperationSse;
|
||||
this->GemmU8U8Operation = MlasGemmU8X8OperationSse;
|
||||
|
||||
#if defined(MLAS_TARGET_AMD64)
|
||||
|
||||
|
|
@ -199,12 +195,10 @@ Return Value:
|
|||
|
||||
if (((Cpuid1[2] & 0x1000) != 0) && ((Cpuid7[1] & 0x20) != 0)) {
|
||||
|
||||
this->GemmU8S8CopyPackARoutine = MlasGemmU8S8CopyPackAAvx2;
|
||||
this->GemmU8S8CopyPackBRoutine = MlasGemmU8S8CopyPackBAvx2;
|
||||
this->GemmU8S8Operation = MlasGemmU8S8OperationAvx2;
|
||||
this->GemmU8S8Kernel = MlasGemmU8S8KernelAvx2;
|
||||
this->GemvU8S8Kernel = MlasGemvU8S8KernelAvx2;
|
||||
this->GemmU8U8CopyPackARoutine = MlasGemmU8U8CopyPackAAvx2;
|
||||
this->GemmU8U8CopyPackBRoutine = MlasGemmU8U8CopyPackBAvx2;
|
||||
this->GemmU8U8Operation = MlasGemmU8U8OperationAvx2;
|
||||
this->GemmU8U8Kernel = MlasGemmU8U8KernelAvx2;
|
||||
|
||||
this->GemmFloatKernel = MlasGemmFloatKernelFma3;
|
||||
|
|
@ -261,9 +255,9 @@ Return Value:
|
|||
|
||||
if ((Cpuid7[2] & 0x800) != 0) {
|
||||
|
||||
this->GemmU8U8Operation = MlasGemmU8S8OperationAvx2;
|
||||
this->GemmU8S8Kernel = MlasGemmU8S8KernelAvx512Vnni;
|
||||
this->GemvU8S8Kernel = MlasGemvU8S8KernelAvx512Vnni;
|
||||
this->GemmU8U8Kernel = MlasGemmU8U8KernelAvx512Vnni;
|
||||
}
|
||||
}
|
||||
|
||||
|
|
|
|||
File diff suppressed because it is too large
Load diff
|
|
@ -45,6 +45,7 @@ Abstract:
|
|||
.equ .LGemmU8S8CopyPackBFrame_SavedRbp, 8
|
||||
.equ .LGemmU8S8CopyPackBFrame_ReturnAddress, 16
|
||||
.equ .LGemmU8S8CopyPackBFrame_offa, 24
|
||||
.equ .LGemmU8S8CopyPackBFrame_BTypeIsSigned, 32
|
||||
|
||||
.text
|
||||
|
||||
|
|
@ -427,6 +428,9 @@ Arguments:
|
|||
offa - Supplies the zero point offset for the other source matrix of the
|
||||
matrix multiplication.
|
||||
|
||||
BTypeIsSigned - Supplies true if the source matrix is signed data, else
|
||||
false if the the source matrix is unsigned data.
|
||||
|
||||
Return Value:
|
||||
|
||||
None.
|
||||
|
|
@ -447,6 +451,17 @@ C_UNDERSCORE(MlasGemmU8S8CopyPackBAvx2):
|
|||
vpsllw ymm0,ymm8,8 # generate word vector [0x0100]
|
||||
vpor ymm8,ymm8,ymm0 # generate word vector [0x0101]
|
||||
|
||||
//
|
||||
// Compute the bit flip vector to adjust input from U8 to S8.
|
||||
//
|
||||
|
||||
vpxor xmm9,xmm9,xmm9 # generate word vector [0x0000]
|
||||
cmp BYTE PTR .LGemmU8S8CopyPackBFrame_BTypeIsSigned[rsp],0
|
||||
jnz .LCopyPackB.SkipUnsignedBitFlipVector
|
||||
vpsllw ymm9,ymm8,7 # generate word vector [0x8080]
|
||||
|
||||
.LCopyPackB.SkipUnsignedBitFlipVector:
|
||||
|
||||
//
|
||||
// Process 16 columns of matrix B in a loop.
|
||||
//
|
||||
|
|
@ -481,6 +496,8 @@ C_UNDERSCORE(MlasGemmU8S8CopyPackBAvx2):
|
|||
vpunpckhwd xmm3,xmm3,xmm5
|
||||
vinserti128 ymm4,ymm4,xmm6,1
|
||||
vinserti128 ymm2,ymm2,xmm3,1
|
||||
vpxor ymm4,ymm4,ymm9 # optionally adjust unsigned data
|
||||
vpxor ymm2,ymm2,ymm9
|
||||
vmovdqu YMMWORD PTR [rdi],ymm4 # store interleaved rows
|
||||
vmovdqu YMMWORD PTR [rdi+32],ymm2
|
||||
vpmaddubsw ymm4,ymm8,ymm4 # horizontal byte+byte=word per row
|
||||
|
|
@ -499,9 +516,9 @@ C_UNDERSCORE(MlasGemmU8S8CopyPackBAvx2):
|
|||
add rbx,4 # correct for over-subtract above
|
||||
jz .LCopyPackB.ReduceColumnSumVectorN16
|
||||
vmovdqu xmm2,XMMWORD PTR [rdx]
|
||||
vpxor xmm3,xmm3,xmm3
|
||||
vpxor xmm4,xmm4,xmm4
|
||||
vpxor xmm5,xmm5,xmm5
|
||||
vmovaps xmm3,xmm9
|
||||
vmovaps xmm4,xmm9
|
||||
vmovaps xmm5,xmm9
|
||||
xor ebx,ebx # no more rows remaining
|
||||
test r8b,2 # (CountK & 2) != 0?
|
||||
jz .LCopyPackB.InterleaveRowDataN16
|
||||
|
|
@ -542,8 +559,8 @@ C_UNDERSCORE(MlasGemmU8S8CopyPackBAvx2):
|
|||
.LCopyPackB.ProcessColumnNUnaligned:
|
||||
vpxor xmm0,xmm0,xmm0 # clear column accumulators
|
||||
vpxor xmm1,xmm1,xmm1
|
||||
vmovdqu YMMWORD PTR .LGemmU8S8CopyPackBFrame_PaddedMatrixBData[rsp],ymm0
|
||||
vmovdqu YMMWORD PTR .LGemmU8S8CopyPackBFrame_PaddedMatrixBData[rsp+32],ymm0
|
||||
vmovdqu YMMWORD PTR .LGemmU8S8CopyPackBFrame_PaddedMatrixBData[rsp],ymm9
|
||||
vmovdqu YMMWORD PTR .LGemmU8S8CopyPackBFrame_PaddedMatrixBData[rsp+32],ymm9
|
||||
sub r8,4
|
||||
jb .LCopyPackB.ProcessRemainingRowsNUnaligned
|
||||
|
||||
|
|
@ -618,6 +635,8 @@ C_UNDERSCORE(MlasGemmU8S8CopyPackBAvx2):
|
|||
vpunpckhwd xmm3,xmm3,xmm5
|
||||
vinserti128 ymm4,ymm4,xmm6,1
|
||||
vinserti128 ymm2,ymm2,xmm3,1
|
||||
vpxor ymm4,ymm4,ymm9 # optionally adjust unsigned data
|
||||
vpxor ymm2,ymm2,ymm9
|
||||
vmovdqu YMMWORD PTR [rdi],ymm4 # store interleaved rows
|
||||
vmovdqu YMMWORD PTR [rdi+32],ymm2
|
||||
vpmaddubsw ymm4,ymm8,ymm4 # horizontal byte+byte=word per row
|
||||
|
|
@ -638,9 +657,8 @@ C_UNDERSCORE(MlasGemmU8S8CopyPackBAvx2):
|
|||
//
|
||||
|
||||
lea rbp,.LGemmU8S8CopyPackBFrame_PaddedMatrixBData[rsp]
|
||||
vpxor xmm6,xmm6,xmm6
|
||||
vmovdqu YMMWORD PTR [rbp],ymm6
|
||||
vmovdqu YMMWORD PTR [rbp+32],ymm6
|
||||
vmovdqu YMMWORD PTR [rbp],ymm9
|
||||
vmovdqu YMMWORD PTR [rbp+32],ymm9
|
||||
|
||||
.LCopyPackB.CopyUnalignedRowLoop:
|
||||
lea r11,[rbp+16] # advance next padded buffer by 16 bytes
|
||||
|
|
|
|||
|
|
@ -1,64 +0,0 @@
|
|||
/*++
|
||||
|
||||
Copyright (c) Microsoft Corporation. All rights reserved.
|
||||
|
||||
Licensed under the MIT License.
|
||||
|
||||
Module Name:
|
||||
|
||||
QgemmU8U8KernelAvx512Common.h
|
||||
|
||||
Abstract:
|
||||
|
||||
This module contains common kernel macros and structures for the quantized
|
||||
integer matrix/matrix multiply operation (QGEMM) for the AVX512 core and
|
||||
AVX512VNNI kernels.
|
||||
|
||||
--*/
|
||||
|
||||
#include "QgemmU8X8KernelAvx512Common.h"
|
||||
|
||||
/*++
|
||||
|
||||
Macro Description:
|
||||
|
||||
This macro generates code to execute the block compute macro multiple
|
||||
times and advancing the matrix A and matrix B data pointers.
|
||||
|
||||
Arguments:
|
||||
|
||||
ColumnCount - Supplies the number of columns to produce.
|
||||
|
||||
RowCount - Supplies the number of rows to produce.
|
||||
|
||||
Implicit Arguments:
|
||||
|
||||
rbx - Supplies the address into the matrix A data plus 3 rows.
|
||||
|
||||
rdi - Supplies the address into the matrix A data.
|
||||
|
||||
rsi - Supplies the address into the matrix B data.
|
||||
|
||||
rcx - Supplies the length in bytes of a row from matrix A.
|
||||
|
||||
r14 - Supplies the stride in bytes of between packed blocks of matrix B.
|
||||
|
||||
zmm14-zmm31 - Supplies the block accumulators.
|
||||
|
||||
--*/
|
||||
|
||||
.macro ComputeBlockLoop ColumnCount, RowCount
|
||||
|
||||
mov rbp,rcx # reload row length remaining
|
||||
|
||||
.LComputeBlockBy1Loop\@:
|
||||
ComputeBlock \ColumnCount\(), \RowCount\(), 0, 0
|
||||
add rdi,4 # advance matrix A by 1 pair
|
||||
.if \RowCount\() > 3
|
||||
add rbx,4 # advance matrix A plus 3 rows by 1 pair
|
||||
.endif
|
||||
add rsi,32 # advance matrix B
|
||||
sub rbp,4
|
||||
jnz .LComputeBlockBy1Loop\@
|
||||
|
||||
.endm
|
||||
|
|
@ -18,7 +18,7 @@ Abstract:
|
|||
--*/
|
||||
|
||||
#include "asmmacro.h"
|
||||
#include "QgemmU8U8KernelAvx512Common.h"
|
||||
#include "QgemmU8X8KernelAvx512Common.h"
|
||||
|
||||
.intel_syntax noprefix
|
||||
|
||||
|
|
@ -124,6 +124,51 @@ Implicit Arguments:
|
|||
|
||||
.endm
|
||||
|
||||
/*++
|
||||
|
||||
Macro Description:
|
||||
|
||||
This macro generates code to execute the block compute macro multiple
|
||||
times and advancing the matrix A and matrix B data pointers.
|
||||
|
||||
Arguments:
|
||||
|
||||
ColumnCount - Supplies the number of columns to produce.
|
||||
|
||||
RowCount - Supplies the number of rows to produce.
|
||||
|
||||
Implicit Arguments:
|
||||
|
||||
rbx - Supplies the address into the matrix A data plus 3 rows.
|
||||
|
||||
rdi - Supplies the address into the matrix A data.
|
||||
|
||||
rsi - Supplies the address into the matrix B data.
|
||||
|
||||
rcx - Supplies the length in bytes of a row from matrix A.
|
||||
|
||||
r14 - Supplies the stride in bytes of between packed blocks of matrix B.
|
||||
|
||||
zmm14-zmm31 - Supplies the block accumulators.
|
||||
|
||||
--*/
|
||||
|
||||
.macro ComputeBlockLoop ColumnCount, RowCount
|
||||
|
||||
mov rbp,rcx # reload row length remaining
|
||||
|
||||
.LComputeBlockBy1Loop\@:
|
||||
ComputeBlock \ColumnCount\(), \RowCount\(), 0, 0
|
||||
add rdi,4 # advance matrix A by 1 pair
|
||||
.if \RowCount\() > 3
|
||||
add rbx,4 # advance matrix A plus 3 rows by 1 pair
|
||||
.endif
|
||||
add rsi,32 # advance matrix B
|
||||
sub rbp,4
|
||||
jnz .LComputeBlockBy1Loop\@
|
||||
|
||||
.endm
|
||||
|
||||
//
|
||||
// Generate the GEMM kernel.
|
||||
//
|
||||
|
|
|
|||
|
|
@ -1,114 +0,0 @@
|
|||
/*++
|
||||
|
||||
Copyright (c) Microsoft Corporation. All rights reserved.
|
||||
|
||||
Licensed under the MIT License.
|
||||
|
||||
Module Name:
|
||||
|
||||
QgemmU8U8KernelAvx512Vnni.s
|
||||
|
||||
Abstract:
|
||||
|
||||
This module implements the kernels for the quantized integer matrix/matrix
|
||||
multiply operation (QGEMM).
|
||||
|
||||
This implementation uses AVX512VNNI instructions.
|
||||
|
||||
--*/
|
||||
|
||||
#include "asmmacro.h"
|
||||
#include "QgemmU8U8KernelAvx512Common.h"
|
||||
#include "AssembleAvx512Vnni.h"
|
||||
|
||||
.intel_syntax noprefix
|
||||
|
||||
.text
|
||||
|
||||
/*++
|
||||
|
||||
Macro Description:
|
||||
|
||||
This macro generates code to multiply and accumulate each row of the output
|
||||
block.
|
||||
|
||||
Arguments:
|
||||
|
||||
ColumnCount - Supplies the number of columns to produce.
|
||||
|
||||
RowCount - Supplies the number of rows to produce.
|
||||
|
||||
VectorOffset - Supplies the byte offset from matrix B to fetch elements.
|
||||
|
||||
BroadcastOffset - Supplies the byte offset from matrix A to fetch elements.
|
||||
|
||||
Implicit Arguments:
|
||||
|
||||
rbx - Supplies the address into the matrix A data plus 3 rows.
|
||||
|
||||
rdi - Supplies the address into the matrix A data.
|
||||
|
||||
rsi - Supplies the address into the matrix B data.
|
||||
|
||||
rcx - Supplies the length in bytes of a row from matrix A.
|
||||
|
||||
r14 - Supplies the stride in bytes of between packed blocks of matrix B.
|
||||
|
||||
zmm14-zmm31 - Supplies the block accumulators.
|
||||
|
||||
--*/
|
||||
|
||||
.macro ComputeBlock ColumnCount, RowCount, VectorOffset, BroadcastOffset
|
||||
|
||||
.if \ColumnCount\() >= 32
|
||||
.if \ColumnCount\() >= 48
|
||||
vpmovzxbw zmm0,YMMWORD PTR [rsi+\VectorOffset\()]
|
||||
vpmovzxbw zmm1,YMMWORD PTR [rsi+r14+\VectorOffset\()]
|
||||
vpmovzxbw zmm2,YMMWORD PTR [rsi+r14*2+\VectorOffset\()]
|
||||
.else
|
||||
vpmovzxbw zmm1,YMMWORD PTR [rsi+\VectorOffset\()]
|
||||
vpmovzxbw zmm2,YMMWORD PTR [rsi+r14+\VectorOffset\()]
|
||||
.endif
|
||||
EmitIfCountGE \RowCount\(), 1, "vpbroadcastd zmm3,DWORD PTR [rdi+\BroadcastOffset\()]"
|
||||
EmitIfCount2GE \RowCount\(), 1, \ColumnCount\(), 48, "VpdpwssdZmmZmmZmm zmm26,zmm3,zmm0"
|
||||
EmitIfCount2GE \RowCount\(), 1, \ColumnCount\(), 32, "VpdpwssdZmmZmmZmm zmm20,zmm3,zmm1"
|
||||
EmitIfCount2GE \RowCount\(), 1, \ColumnCount\(), 16, "VpdpwssdZmmZmmZmm zmm14,zmm3,zmm2"
|
||||
EmitIfCountGE \RowCount\(), 2, "vpbroadcastd zmm3,DWORD PTR [rdi+rcx+\BroadcastOffset\()]"
|
||||
EmitIfCount2GE \RowCount\(), 2, \ColumnCount\(), 48, "VpdpwssdZmmZmmZmm zmm27,zmm3,zmm0"
|
||||
EmitIfCount2GE \RowCount\(), 2, \ColumnCount\(), 32, "VpdpwssdZmmZmmZmm zmm21,zmm3,zmm1"
|
||||
EmitIfCount2GE \RowCount\(), 2, \ColumnCount\(), 16, "VpdpwssdZmmZmmZmm zmm15,zmm3,zmm2"
|
||||
EmitIfCountGE \RowCount\(), 3, "vpbroadcastd zmm3,DWORD PTR [rdi+rcx*2+\BroadcastOffset\()]"
|
||||
EmitIfCount2GE \RowCount\(), 3, \ColumnCount\(), 48, "VpdpwssdZmmZmmZmm zmm28,zmm3,zmm0"
|
||||
EmitIfCount2GE \RowCount\(), 3, \ColumnCount\(), 32, "VpdpwssdZmmZmmZmm zmm22,zmm3,zmm1"
|
||||
EmitIfCount2GE \RowCount\(), 3, \ColumnCount\(), 16, "VpdpwssdZmmZmmZmm zmm16,zmm3,zmm2"
|
||||
EmitIfCountGE \RowCount\(), 4, "vpbroadcastd zmm3,DWORD PTR [rbx+\BroadcastOffset\()]"
|
||||
EmitIfCount2GE \RowCount\(), 4, \ColumnCount\(), 48, "VpdpwssdZmmZmmZmm zmm29,zmm3,zmm0"
|
||||
EmitIfCount2GE \RowCount\(), 4, \ColumnCount\(), 32, "VpdpwssdZmmZmmZmm zmm23,zmm3,zmm1"
|
||||
EmitIfCount2GE \RowCount\(), 4, \ColumnCount\(), 16, "VpdpwssdZmmZmmZmm zmm17,zmm3,zmm2"
|
||||
EmitIfCountGE \RowCount\(), 5, "vpbroadcastd zmm3,DWORD PTR [rbx+rcx+\BroadcastOffset\()]"
|
||||
EmitIfCount2GE \RowCount\(), 5, \ColumnCount\(), 48, "VpdpwssdZmmZmmZmm zmm30,zmm3,zmm0"
|
||||
EmitIfCount2GE \RowCount\(), 5, \ColumnCount\(), 32, "VpdpwssdZmmZmmZmm zmm24,zmm3,zmm1"
|
||||
EmitIfCount2GE \RowCount\(), 5, \ColumnCount\(), 16, "VpdpwssdZmmZmmZmm zmm18,zmm3,zmm2"
|
||||
EmitIfCountGE \RowCount\(), 6, "vpbroadcastd zmm3,DWORD PTR [rbx+rcx*2+\BroadcastOffset\()]"
|
||||
EmitIfCount2GE \RowCount\(), 6, \ColumnCount\(), 48, "VpdpwssdZmmZmmZmm zmm31,zmm3,zmm0"
|
||||
EmitIfCount2GE \RowCount\(), 6, \ColumnCount\(), 32, "VpdpwssdZmmZmmZmm zmm25,zmm3,zmm1"
|
||||
EmitIfCount2GE \RowCount\(), 6, \ColumnCount\(), 16, "VpdpwssdZmmZmmZmm zmm19,zmm3,zmm2"
|
||||
.else
|
||||
vpmovzxbw zmm2,YMMWORD PTR [rsi+\VectorOffset\()]
|
||||
EmitIfCountGE \RowCount\(), 1, "VpdpwssdZmmZmmBroadcast zmm14,zmm2,rdi,\BroadcastOffset\()"
|
||||
EmitIfCountGE \RowCount\(), 2, "VpdpwssdZmmZmmBroadcast zmm15,zmm2,rdi,\BroadcastOffset\(),rcx,1"
|
||||
EmitIfCountGE \RowCount\(), 3, "VpdpwssdZmmZmmBroadcast zmm16,zmm2,rdi,\BroadcastOffset\(),rcx,2"
|
||||
EmitIfCountGE \RowCount\(), 4, "VpdpwssdZmmZmmBroadcast zmm17,zmm2,rbx,\BroadcastOffset\()"
|
||||
EmitIfCountGE \RowCount\(), 5, "VpdpwssdZmmZmmBroadcast zmm18,zmm2,rbx,\BroadcastOffset\(),rcx,1"
|
||||
EmitIfCountGE \RowCount\(), 6, "VpdpwssdZmmZmmBroadcast zmm19,zmm2,rbx,\BroadcastOffset\(),rcx,2"
|
||||
.endif
|
||||
|
||||
.endm
|
||||
|
||||
//
|
||||
// Generate the GEMM kernel.
|
||||
//
|
||||
|
||||
GemmU8X8KernelAvx512Function U8U8, Avx512Vnni
|
||||
|
||||
.end
|
||||
|
|
@ -106,8 +106,8 @@ Implicit Arguments:
|
|||
.endif
|
||||
ComputeBlockLoop \ColumnCount\(), \RowCount\()
|
||||
.if \RowCount\() > 3
|
||||
lea rbx,[rdx+rax*2] # compute matrix C plus 3 rows
|
||||
add rbx,rax
|
||||
lea rbx,[rax*2+rax]
|
||||
add rbx,rdx # compute matrix C plus 3 rows
|
||||
.endif
|
||||
|
||||
.endm
|
||||
|
|
|
|||
Loading…
Reference in a new issue