MLAS: qgemm refactoring (#4030)

Treat U8U8 as U8S8 for VNNI for performance and optimize SSE2 kernel.
This commit is contained in:
Tracy Sharpe 2020-05-26 17:27:32 -07:00 committed by GitHub
parent abcd1576c9
commit 0d8abc1a99
No known key found for this signature in database
GPG key ID: 4AEE18F83AFDEB23
21 changed files with 800 additions and 1448 deletions

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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