From d5f6343a4afbb1d3ae7acdc881f79bd93cc7c0b5 Mon Sep 17 00:00:00 2001 From: mguynn-intc Date: Wed, 18 Sep 2024 22:18:23 -0700 Subject: [PATCH] Implementation of AVX-VNNI-INT8 dot product instructions into MLAS GEMM (#21984) ### Description ONNXRuntime implementation of S8S8 was using the default C++ implementation; with this new ISA, all variants of QGemm Int8 can support VNNI dot product and full AVX2 instructions. All signed/unsigned variants support VNNI instructions starting with LNL. Renamed structs and functions to better indicate support of all Int8 vs U8X8 ### Motivation and Context LNL HW implemented new ISA, and this code enables that ISA in QGemm. Speed is improved for S8S8 to match with existing U8S8 code. S8U8 would also match speed if ONNX formally accepted the data type. --- .../core/mlas/lib/amd64/AssembleAvxVnni.inc | 154 +++++++- .../mlas/lib/amd64/QgemmU8S8KernelAvx2.asm | 200 ++++++++--- .../mlas/lib/amd64/QgemmU8X8KernelAvx2.asm | 246 +++++++------ onnxruntime/core/mlas/lib/mlasi.h | 10 + onnxruntime/core/mlas/lib/platform.cpp | 11 + onnxruntime/core/mlas/lib/qgemm.h | 21 +- .../core/mlas/lib/qgemm_kernel_avx2.cpp | 252 +++++++++++++ .../core/mlas/lib/x86_64/AssembleAvxVnni.h | 150 ++++++++ .../mlas/lib/x86_64/QgemmU8S8KernelAvx2.S | 332 +++++++++++------- .../mlas/lib/x86_64/QgemmU8X8KernelAvx2.S | 225 +++++++----- onnxruntime/test/mlas/unittest/test_qgemm.cpp | 22 +- 11 files changed, 1238 insertions(+), 385 deletions(-) diff --git a/onnxruntime/core/mlas/lib/amd64/AssembleAvxVnni.inc b/onnxruntime/core/mlas/lib/amd64/AssembleAvxVnni.inc index 8717528d0d..d5867e8884 100644 --- a/onnxruntime/core/mlas/lib/amd64/AssembleAvxVnni.inc +++ b/onnxruntime/core/mlas/lib/amd64/AssembleAvxVnni.inc @@ -175,4 +175,156 @@ VpdpwssdsXmmXmmXmm MACRO DestReg, Src1Reg, Src2Reg VnniXmmXmmXmm 053h, DestReg, Src1Reg, Src2Reg - ENDM \ No newline at end of file + ENDM + +; +; Macro Description: +; +; This macro builds a VNNI instruction of the form: +; +; instr ymm1,ymm2,ymm3 +; +; Arguments: +; +; Opcode - Specifies the opcode for the VNNI instruction. +; +; Prefix - Specifies the opcode prefix for payload 1 +; +; DestReg - Specifies the destination register. +; +; Src1Reg - Specifies the first source register. +; +; Src2Reg - Specifies the second source register. +; + +Avx2VnniYmmYmmYmm MACRO Opcode, Prefix, DestReg, Src1Reg, Src2Reg + + LOCAL Payload0, Payload1, ModRMByte + + Payload0 = 002h ; "0F 38" prefix + Payload0 = Payload0 + ((((YmmIndex_&DestReg& SHR 3) AND 1) XOR 1) SHL 7) + Payload0 = Payload0 + (1 SHL 6) + Payload0 = Payload0 + ((((YmmIndex_&Src2Reg& SHR 3) AND 1) XOR 1) SHL 5) + + Payload1 = 004h + Prefix ; 256-bit length and opcode prefix + Payload1 = Payload1 + (((YmmIndex_&Src1Reg& AND 15) XOR 15) SHL 3) + + ModRMByte = 0C0h ; register form + ModRMByte = ModRMByte + ((YmmIndex_&DestReg& AND 7) SHL 3) + ModRMByte = ModRMByte + (YmmIndex_&Src2Reg& AND 7) + + db 0C4h, Payload0, Payload1, Opcode, ModRMByte + + ENDM + +VpdpbssdYmmYmmYmm MACRO DestReg, Src1Reg, Src2Reg + + Avx2VnniYmmYmmYmm 050h, 003h, DestReg, Src1Reg, Src2Reg + + ENDM + +VpdpbssdsYmmYmmYmm MACRO DestReg, Src1Reg, Src2Reg + + Avx2VnniYmmYmmYmm 051h, 003h, DestReg, Src1Reg, Src2Reg + + ENDM + +VpdpbsudYmmYmmYmm MACRO DestReg, Src1Reg, Src2Reg + + Avx2VnniYmmYmmYmm 050h, 002h, DestReg, Src1Reg, Src2Reg + + ENDM + +VpdpbsudsYmmYmmYmm MACRO DestReg, Src1Reg, Src2Reg + + Avx2VnniYmmYmmYmm 051h, 002h, DestReg, Src1Reg, Src2Reg + + ENDM + +VpdpbuudYmmYmmYmm MACRO DestReg, Src1Reg, Src2Reg + + Avx2VnniYmmYmmYmm 050h, 000h, DestReg, Src1Reg, Src2Reg + + ENDM + +VpdpbuudsYmmYmmYmm MACRO DestReg, Src1Reg, Src2Reg + + Avx2VnniYmmYmmYmm 051h, 000h, DestReg, Src1Reg, Src2Reg + + ENDM + +; +; Macro Description: +; +; This macro builds a VNNI instruction of the form: +; +; instr xmm1,xmm2,xmm3 +; +; Arguments: +; +; Opcode - Specifies the opcode for the VNNI instruction. +; +; Prefix - Specifies the opcode prefix for payload 1 +; +; DestReg - Specifies the destination register. +; +; Src1Reg - Specifies the first source register. +; +; Src2Reg - Specifies the second source register. +; + +Avx2VnniXmmXmmXmm MACRO Opcode, Prefix, DestReg, Src1Reg, Src2Reg + + LOCAL Payload0, Payload1, ModRMByte + + Payload0 = 002h ; "0F 38" prefix + Payload0 = Payload0 + ((((XmmIndex_&DestReg& SHR 3) AND 1) XOR 1) SHL 7) + Payload0 = Payload0 + (1 SHL 6) + Payload0 = Payload0 + ((((XmmIndex_&Src2Reg& SHR 3) AND 1) XOR 1) SHL 5) + + Payload1 = 000h + Prefix ; 128-bit length and opcode prefix + Payload1 = Payload1 + (((XmmIndex_&Src1Reg& AND 15) XOR 15) SHL 3) + + ModRMByte = 0C0h ; register form + ModRMByte = ModRMByte + ((XmmIndex_&DestReg& AND 7) SHL 3) + ModRMByte = ModRMByte + (XmmIndex_&Src2Reg& AND 7) + + db 0C4h, Payload0, Payload1, Opcode, ModRMByte + + ENDM + +VpdpbssdXmmXmmXmm MACRO DestReg, Src1Reg, Src2Reg + + Avx2VnniXmmXmmXmm 050h, 003h, DestReg, Src1Reg, Src2Reg + + ENDM + +VpdpbssdsXmmXmmXmm MACRO DestReg, Src1Reg, Src2Reg + + Avx2VnniXmmXmmXmm 051h, 003h, DestReg, Src1Reg, Src2Reg + + ENDM + +VpdpbsudXmmXmmXmm MACRO DestReg, Src1Reg, Src2Reg + + Avx2VnniXmmXmmXmm 050h, 002h, DestReg, Src1Reg, Src2Reg + + ENDM + +VpdpbsudsXmmXmmXmm MACRO DestReg, Src1Reg, Src2Reg + + Avx2VnniXmmXmmXmm 051h, 002h, DestReg, Src1Reg, Src2Reg + + ENDM + +VpdpbuudXmmXmmXmm MACRO DestReg, Src1Reg, Src2Reg + + Avx2VnniXmmXmmXmm 050h, 000h, DestReg, Src1Reg, Src2Reg + + ENDM + +VpdpbuudsXmmXmmXmm MACRO DestReg, Src1Reg, Src2Reg + + Avx2VnniXmmXmmXmm 051h, 000h, DestReg, Src1Reg, Src2Reg + + ENDM diff --git a/onnxruntime/core/mlas/lib/amd64/QgemmU8S8KernelAvx2.asm b/onnxruntime/core/mlas/lib/amd64/QgemmU8S8KernelAvx2.asm index cccea70c14..bc5ac69e73 100644 --- a/onnxruntime/core/mlas/lib/amd64/QgemmU8S8KernelAvx2.asm +++ b/onnxruntime/core/mlas/lib/amd64/QgemmU8S8KernelAvx2.asm @@ -14,20 +14,22 @@ ; multiply operation (QGEMM). ; ; This implementation uses AVX2 instructions. +; Support for AVX-VNNI-INT8 for certain code paths. ; ;-- .xlist INCLUDE mlasi.inc +INCLUDE AssembleAvxVnni.inc .list EXTERN MlasMaskMoveTableAvx:NEAR ; -; Stack frame layout for the U8S8 CopyPackA routine. +; Stack frame layout for the Int8 CopyPackA routine. ; -GemmU8S8CopyPackAFrame STRUCT +GemmInt8CopyPackAFrame STRUCT PaddedMatrixAData OWORD 4 DUP (?) SavedXmm6 OWORD ? @@ -50,13 +52,13 @@ GemmU8S8CopyPackAFrame STRUCT CountK QWORD ? RowSumBuffer QWORD ? -GemmU8S8CopyPackAFrame ENDS +GemmInt8CopyPackAFrame ENDS ; -; Stack frame layout for the U8S8 CopyPackB routine. +; Stack frame layout for the Int8 CopyPackB routine. ; -GemmU8S8CopyPackBFrame STRUCT +GemmInt8CopyPackBFrame STRUCT PaddedMatrixBData OWORD 4 DUP (?) SavedXmm6 OWORD ? @@ -77,7 +79,7 @@ GemmU8S8CopyPackBFrame STRUCT ColumnSumBuffer QWORD ? BIsSigned QWORD ? -GemmU8S8CopyPackBFrame ENDS +GemmInt8CopyPackBFrame ENDS ;++ ; @@ -107,7 +109,7 @@ GemmU8S8CopyPackBFrame ENDS ; ;-- - NESTED_ENTRY MlasGemmU8S8CopyPackAAvx2, _TEXT +MlasGemmCopyPackAAvx2 MACRO ASigned rex_push_reg rbp push_reg rbx @@ -115,21 +117,21 @@ GemmU8S8CopyPackBFrame ENDS push_reg rdi push_reg r12 push_reg r13 - alloc_stack (GemmU8S8CopyPackAFrame.SavedR13) - save_xmm128 xmm6,GemmU8S8CopyPackAFrame.SavedXmm6 - save_xmm128 xmm7,GemmU8S8CopyPackAFrame.SavedXmm7 - save_xmm128 xmm8,GemmU8S8CopyPackAFrame.SavedXmm8 - save_xmm128 xmm9,GemmU8S8CopyPackAFrame.SavedXmm9 - save_xmm128 xmm10,GemmU8S8CopyPackAFrame.SavedXmm10 + alloc_stack (GemmInt8CopyPackAFrame.SavedR13) + save_xmm128 xmm6,GemmInt8CopyPackAFrame.SavedXmm6 + save_xmm128 xmm7,GemmInt8CopyPackAFrame.SavedXmm7 + save_xmm128 xmm8,GemmInt8CopyPackAFrame.SavedXmm8 + save_xmm128 xmm9,GemmInt8CopyPackAFrame.SavedXmm9 + save_xmm128 xmm10,GemmInt8CopyPackAFrame.SavedXmm10 END_PROLOGUE mov rdi,rcx mov rsi,rdx - mov r10,GemmU8S8CopyPackAFrame.CountK[rsp] + mov r10,GemmInt8CopyPackAFrame.CountK[rsp] lea r11,[r10+3] and r11,NOT 3 ; align CountK up to quad count - mov r12,GemmU8S8CopyPackAFrame.RowSumBuffer[rsp] + mov r12,GemmInt8CopyPackAFrame.RowSumBuffer[rsp] vpcmpeqw ymm8,ymm8,ymm8 ; generate word vector [0xFFFF] vpsrlw ymm8,ymm8,15 ; generate word vector [0x0001] vpsllw ymm9,ymm8,8 ; generate word vector [0x0100] @@ -152,8 +154,8 @@ GemmU8S8CopyPackBFrame ENDS ; vpxor xmm0,xmm0,xmm0 - vmovdqu YMMWORD PTR GemmU8S8CopyPackAFrame.PaddedMatrixAData[rsp],ymm0 - vmovdqu YMMWORD PTR GemmU8S8CopyPackAFrame.PaddedMatrixAData[rsp+32],ymm0 + vmovdqu YMMWORD PTR GemmInt8CopyPackAFrame.PaddedMatrixAData[rsp],ymm0 + vmovdqu YMMWORD PTR GemmInt8CopyPackAFrame.PaddedMatrixAData[rsp+32],ymm0 ; ; Process 4 rows of matrix A in a loop. @@ -186,6 +188,12 @@ ProcessNextColumnLoopM4: vmovdqu YMMWORD PTR [rcx+r11],ymm5 vmovdqu YMMWORD PTR [rcx+r11*2],ymm6 vmovdqu YMMWORD PTR [rcx+rax],ymm7 +IF ASigned EQ 1 + VpdpbssdYmmYmmYmm ymm0,ymm4,ymm9 + VpdpbssdYmmYmmYmm ymm1,ymm5,ymm9 + VpdpbssdYmmYmmYmm ymm2,ymm6,ymm9 + VpdpbssdYmmYmmYmm ymm3,ymm7,ymm9 +ELSE vpmaddubsw ymm4,ymm4,ymm9 ; horizontal byte+byte=word per row vpaddw ymm0,ymm0,ymm4 ; add words to row accumulators vpmaddubsw ymm5,ymm5,ymm9 @@ -194,6 +202,7 @@ ProcessNextColumnLoopM4: vpaddw ymm2,ymm2,ymm6 vpmaddubsw ymm7,ymm7,ymm9 vpaddw ymm3,ymm3,ymm7 +ENDIF add rdx,32 ; advance matrix A by 32 bytes add rcx,32 ; advance matrix D by 32 bytes sub rbx,32 ; subtract columns remaining @@ -212,6 +221,12 @@ ProcessRemainingColumnsM4: vmovdqu XMMWORD PTR [rcx+r11],xmm5 vmovdqu XMMWORD PTR [rcx+r11*2],xmm6 vmovdqu XMMWORD PTR [rcx+rax],xmm7 +IF ASigned EQ 1 + VpdpbssdYmmYmmYmm ymm0,ymm4,ymm9 + VpdpbssdYmmYmmYmm ymm1,ymm5,ymm9 + VpdpbssdYmmYmmYmm ymm2,ymm6,ymm9 + VpdpbssdYmmYmmYmm ymm3,ymm7,ymm9 +ELSE vpmaddubsw xmm4,xmm4,xmm9 ; horizontal byte+byte=word per row vpaddw ymm0,ymm0,ymm4 ; add words to row accumulators vpmaddubsw xmm5,xmm5,xmm9 @@ -220,6 +235,7 @@ ProcessRemainingColumnsM4: vpaddw ymm2,ymm2,ymm6 vpmaddubsw xmm7,xmm7,xmm9 vpaddw ymm3,ymm3,ymm7 +ENDIF add rdx,16 ; advance matrix A by 16 bytes add rcx,16 ; advance matrix D by 16 bytes test bl,15 ; test for unaligned columns @@ -230,8 +246,8 @@ ProcessRemainingColumnsM4: ; CopyRemainingCountKLessThan16M4: -.errnz GemmU8S8CopyPackAFrame.PaddedMatrixAData - mov rbp,rsp ; GemmU8S8CopyPackAFrame.PaddedMatrixAData +.errnz GemmInt8CopyPackAFrame.PaddedMatrixAData + mov rbp,rsp ; GemmInt8CopyPackAFrame.PaddedMatrixAData test bl,8 ; (CountK & 8) != 0? jz CopyRemainingCountKLessThan8M4 mov rax,QWORD PTR [rdx] @@ -290,15 +306,21 @@ CopyRemainingCountKLessThan2M4: ; ProcessPaddedMatrixADataM4: - vmovdqu xmm4,XMMWORD PTR GemmU8S8CopyPackAFrame.PaddedMatrixAData[rsp] - vmovdqu xmm5,XMMWORD PTR GemmU8S8CopyPackAFrame.PaddedMatrixAData[rsp+16] - vmovdqu xmm6,XMMWORD PTR GemmU8S8CopyPackAFrame.PaddedMatrixAData[rsp+32] - vmovdqu xmm7,XMMWORD PTR GemmU8S8CopyPackAFrame.PaddedMatrixAData[rsp+48] + vmovdqu xmm4,XMMWORD PTR GemmInt8CopyPackAFrame.PaddedMatrixAData[rsp] + vmovdqu xmm5,XMMWORD PTR GemmInt8CopyPackAFrame.PaddedMatrixAData[rsp+16] + vmovdqu xmm6,XMMWORD PTR GemmInt8CopyPackAFrame.PaddedMatrixAData[rsp+32] + vmovdqu xmm7,XMMWORD PTR GemmInt8CopyPackAFrame.PaddedMatrixAData[rsp+48] lea rax,[rcx+r11*2] ; compute matrix D plus 2 rows vpmaskmovd XMMWORD PTR [rcx],xmm10,xmm4 vpmaskmovd XMMWORD PTR [rcx+r11],xmm10,xmm5 vpmaskmovd XMMWORD PTR [rax],xmm10,xmm6 vpmaskmovd XMMWORD PTR [rax+r11],xmm10,xmm7 +IF ASigned EQ 1 + VpdpbssdYmmYmmYmm ymm0,ymm4,ymm9 + VpdpbssdYmmYmmYmm ymm1,ymm5,ymm9 + VpdpbssdYmmYmmYmm ymm2,ymm6,ymm9 + VpdpbssdYmmYmmYmm ymm3,ymm7,ymm9 +ELSE vpmaddubsw xmm4,xmm4,xmm9 ; horizontal byte+byte=word per row vpaddw ymm0,ymm0,ymm4 ; add words to row accumulators vpmaddubsw xmm5,xmm5,xmm9 @@ -307,17 +329,22 @@ ProcessPaddedMatrixADataM4: vpaddw ymm2,ymm2,ymm6 vpmaddubsw xmm7,xmm7,xmm9 vpaddw ymm3,ymm3,ymm7 +ENDIF ; ; Reduce the sums for the four rows of output. ; ReduceRowSumBufferM4: +IF ASigned EQ 1 + vphaddd ymm0,ymm0,ymm1 ; reduce and interleave Sum1/Sum0 +ELSE vpmaddwd ymm0,ymm0,ymm8 ; horizontal word+word=dword per row vpmaddwd ymm1,ymm1,ymm8 vphaddd ymm0,ymm0,ymm1 ; reduce and interleave Sum1/Sum0 vpmaddwd ymm2,ymm2,ymm8 vpmaddwd ymm3,ymm3,ymm8 +ENDIF vphaddd ymm1,ymm2,ymm3 ; reduce and interleave Sum3/Sum2 vphaddd ymm0,ymm0,ymm1 ; reduce and interleave Sum3/Sum2/Sum1/Sum0 vextracti128 xmm1,ymm0,1 ; extract high dwords @@ -348,8 +375,12 @@ ProcessNextRowM1: ProcessNextColumnLoopM1: vmovdqu ymm4,YMMWORD PTR [rdx] vmovdqu YMMWORD PTR [rcx],ymm4 +IF ASigned EQ 1 + VpdpbssdYmmYmmYmm ymm0,ymm4,ymm9 +ELSE vpmaddubsw ymm4,ymm4,ymm9 ; horizontal byte+byte=word per row vpaddw ymm0,ymm0,ymm4 ; add words to row accumulators +ENDIF add rdx,32 ; advance matrix A by 32 bytes add rcx,32 ; advance matrix D by 32 bytes sub rbx,32 ; subtract columns remaining @@ -362,8 +393,12 @@ ProcessRemainingColumnsM1: jz CopyRemainingCountKLessThan16M1 vmovdqu xmm4,XMMWORD PTR [rdx] vmovdqu XMMWORD PTR [rcx],xmm4 +IF ASigned EQ 1 + VpdpbssdYmmYmmYmm ymm0,ymm4,ymm9 +ELSE vpmaddubsw xmm4,xmm4,xmm9 ; horizontal byte+byte=word per row vpaddw ymm0,ymm0,ymm4 ; add words to row accumulators +ENDIF add rdx,16 ; advance matrix A by 16 bytes add rcx,16 ; advance matrix D by 16 bytes test bl,15 ; test for unaligned columns @@ -374,8 +409,8 @@ ProcessRemainingColumnsM1: ; CopyRemainingCountKLessThan16M1: -.errnz GemmU8S8CopyPackAFrame.PaddedMatrixAData - mov rbp,rsp ; GemmU8S8CopyPackAFrame.PaddedMatrixAData +.errnz GemmInt8CopyPackAFrame.PaddedMatrixAData + mov rbp,rsp ; GemmInt8CopyPackAFrame.PaddedMatrixAData test bl,8 ; (CountK & 8) != 0? jz CopyRemainingCountKLessThan8M1 mov rax,QWORD PTR [rdx] @@ -410,17 +445,23 @@ CopyRemainingCountKLessThan2M1: ; ProcessPaddedMatrixADataM1: - vmovdqu xmm4,XMMWORD PTR GemmU8S8CopyPackAFrame.PaddedMatrixAData[rsp] + vmovdqu xmm4,XMMWORD PTR GemmInt8CopyPackAFrame.PaddedMatrixAData[rsp] vpmaskmovd XMMWORD PTR [rcx],xmm10,xmm4 +IF ASigned EQ 1 + VpdpbssdYmmYmmYmm ymm0,ymm4,ymm9 +ELSE vpmaddubsw ymm4,ymm4,ymm9 ; horizontal byte+byte=word per row vpaddw ymm0,ymm0,ymm4 ; add words to row accumulators +ENDIF ; ; Reduce the sum for the single row of output. ; ReduceRowSumBufferM1: +IF ASigned EQ 0 vpmaddwd ymm0,ymm0,ymm8 ; horizontal word+word=dword per row +ENDIF vextracti128 xmm1,ymm0,1 ; extract high dwords vpaddd xmm0,xmm0,xmm1 ; reduction vphaddd xmm0,xmm0,xmm0 @@ -436,12 +477,12 @@ ReduceRowSumBufferM1: ExitRoutine: vzeroupper - movaps xmm6,GemmU8S8CopyPackAFrame.SavedXmm6[rsp] - movaps xmm7,GemmU8S8CopyPackAFrame.SavedXmm7[rsp] - movaps xmm8,GemmU8S8CopyPackAFrame.SavedXmm8[rsp] - movaps xmm9,GemmU8S8CopyPackAFrame.SavedXmm9[rsp] - movaps xmm10,GemmU8S8CopyPackAFrame.SavedXmm10[rsp] - add rsp,(GemmU8S8CopyPackAFrame.SavedR13) + movaps xmm6,GemmInt8CopyPackAFrame.SavedXmm6[rsp] + movaps xmm7,GemmInt8CopyPackAFrame.SavedXmm7[rsp] + movaps xmm8,GemmInt8CopyPackAFrame.SavedXmm8[rsp] + movaps xmm9,GemmInt8CopyPackAFrame.SavedXmm9[rsp] + movaps xmm10,GemmInt8CopyPackAFrame.SavedXmm10[rsp] + add rsp,(GemmInt8CopyPackAFrame.SavedR13) BEGIN_EPILOGUE @@ -453,8 +494,16 @@ ExitRoutine: pop rbp ret + ENDM + + NESTED_ENTRY MlasGemmU8S8CopyPackAAvx2, _TEXT + MlasGemmCopyPackAAvx2 0 NESTED_END MlasGemmU8S8CopyPackAAvx2, _TEXT + NESTED_ENTRY MlasGemmS8CopyPackAAvx2Vnni, _TEXT + MlasGemmCopyPackAAvx2 1 + NESTED_END MlasGemmS8CopyPackAAvx2Vnni, _TEXT + ;++ ; ; Routine Description: @@ -486,24 +535,24 @@ ExitRoutine: ; ;-- - NESTED_ENTRY MlasGemmU8S8CopyPackBAvx2, _TEXT +MlasGemmCopyPackBAvx2 MACRO IsVnni, BSigned rex_push_reg rbp push_reg rbx push_reg rsi push_reg rdi - alloc_stack (GemmU8S8CopyPackBFrame.SavedRdi) - save_xmm128 xmm6,GemmU8S8CopyPackBFrame.SavedXmm6 - save_xmm128 xmm7,GemmU8S8CopyPackBFrame.SavedXmm7 - save_xmm128 xmm8,GemmU8S8CopyPackBFrame.SavedXmm8 - save_xmm128 xmm9,GemmU8S8CopyPackBFrame.SavedXmm9 + alloc_stack (GemmInt8CopyPackBFrame.SavedRdi) + save_xmm128 xmm6,GemmInt8CopyPackBFrame.SavedXmm6 + save_xmm128 xmm7,GemmInt8CopyPackBFrame.SavedXmm7 + save_xmm128 xmm8,GemmInt8CopyPackBFrame.SavedXmm8 + save_xmm128 xmm9,GemmInt8CopyPackBFrame.SavedXmm9 END_PROLOGUE mov rsi,rdx lea rdi,[r8+r8*2] ; compute ldb * 3 - mov r10,GemmU8S8CopyPackBFrame.CountK[rsp] - mov r11,GemmU8S8CopyPackBFrame.ColumnSumBuffer[rsp] + mov r10,GemmInt8CopyPackBFrame.CountK[rsp] + mov r11,GemmInt8CopyPackBFrame.ColumnSumBuffer[rsp] vpcmpeqw ymm7,ymm7,ymm7 ; generate word vector [0xFFFF] vpsrlw ymm7,ymm7,15 ; generate word vector [0x0001] vpsllw ymm8,ymm7,8 ; generate word vector [0x0100] @@ -514,10 +563,11 @@ ExitRoutine: ; vpxor xmm9,xmm9,xmm9 ; generate word vector [0x0000] - cmp BYTE PTR GemmU8S8CopyPackBFrame.BIsSigned[rsp],0 +IF IsVnni EQ 0 + cmp BYTE PTR GemmInt8CopyPackBFrame.BIsSigned[rsp],0 jnz SkipUnsignedBitFlipVector vpsllw ymm9,ymm8,7 ; generate word vector [0x8080] - +ENDIF SkipUnsignedBitFlipVector: ; @@ -554,16 +604,28 @@ InterleaveRowDataN16: vpunpckhwd xmm3,xmm3,xmm5 vinserti128 ymm4,ymm4,xmm6,1 vinserti128 ymm2,ymm2,xmm3,1 +IF IsVnni EQ 0 vpxor ymm4,ymm4,ymm9 ; optionally adjust unsigned data vpxor ymm2,ymm2,ymm9 +ENDIF vmovdqu YMMWORD PTR [rcx],ymm4 ; store interleaved rows vmovdqu YMMWORD PTR [rcx+32],ymm2 +IF IsVnni EQ 1 + IF BSigned EQ 1 + VpdpbssdYmmYmmYmm ymm0,ymm4,ymm8 + VpdpbssdYmmYmmYmm ymm1,ymm2,ymm8 + ELSE + VpdpbuudYmmYmmYmm ymm0,ymm4,ymm8 + VpdpbuudYmmYmmYmm ymm1,ymm2,ymm8 + ENDIF +ELSE vpmaddubsw ymm4,ymm8,ymm4 ; horizontal byte+byte=word per row vpmaddwd ymm4,ymm4,ymm7 ; horizontal word+word=dword per row vpaddd ymm0,ymm0,ymm4 ; accumulate per column vpmaddubsw ymm2,ymm8,ymm2 vpmaddwd ymm2,ymm2,ymm7 vpaddd ymm1,ymm1,ymm2 +ENDIF add rcx,64 ; advance matrix D by 64 bytes sub rbx,4 ; subtract rows remaining jae ProcessNextRowLoopN16 @@ -605,11 +667,11 @@ ProcessRemainingColumns: ExitRoutine: vzeroupper - movaps xmm6,GemmU8S8CopyPackBFrame.SavedXmm6[rsp] - movaps xmm7,GemmU8S8CopyPackBFrame.SavedXmm7[rsp] - movaps xmm8,GemmU8S8CopyPackBFrame.SavedXmm8[rsp] - movaps xmm9,GemmU8S8CopyPackBFrame.SavedXmm9[rsp] - add rsp,(GemmU8S8CopyPackBFrame.SavedRdi) + movaps xmm6,GemmInt8CopyPackBFrame.SavedXmm6[rsp] + movaps xmm7,GemmInt8CopyPackBFrame.SavedXmm7[rsp] + movaps xmm8,GemmInt8CopyPackBFrame.SavedXmm8[rsp] + movaps xmm9,GemmInt8CopyPackBFrame.SavedXmm9[rsp] + add rsp,(GemmInt8CopyPackBFrame.SavedRdi) BEGIN_EPILOGUE @@ -626,15 +688,15 @@ ExitRoutine: ProcessColumnNUnaligned: vpxor xmm0,xmm0,xmm0 ; clear column accumulators vpxor xmm1,xmm1,xmm1 - vmovdqu YMMWORD PTR GemmU8S8CopyPackBFrame.PaddedMatrixBData[rsp],ymm9 - vmovdqu YMMWORD PTR GemmU8S8CopyPackBFrame.PaddedMatrixBData[rsp+32],ymm9 + vmovdqu YMMWORD PTR GemmInt8CopyPackBFrame.PaddedMatrixBData[rsp],ymm9 + vmovdqu YMMWORD PTR GemmInt8CopyPackBFrame.PaddedMatrixBData[rsp+32],ymm9 sub r10,4 jb ProcessRemainingRowsNUnaligned ProcessNextRowLoopNUnaligned: mov rdx,rsi -.errnz GemmU8S8CopyPackBFrame.PaddedMatrixBData - mov rbp,rsp ; GemmU8S8CopyPackBFrame.PaddedMatrixBData +.errnz GemmInt8CopyPackBFrame.PaddedMatrixBData + mov rbp,rsp ; GemmInt8CopyPackBFrame.PaddedMatrixBData test r9b,8 ; (CountN & 8) != 0? jz CopyRemainingCountNLessThan8K4 mov rax,QWORD PTR [rdx] @@ -689,10 +751,10 @@ CopyRemainingCountNLessThan2K4: mov BYTE PTR [rbp+48],al ProcessPaddedMatrixBData: - vmovdqu xmm2,XMMWORD PTR GemmU8S8CopyPackBFrame.PaddedMatrixBData[rsp] - vmovdqu xmm3,XMMWORD PTR GemmU8S8CopyPackBFrame.PaddedMatrixBData[rsp+16] - vmovdqu xmm4,XMMWORD PTR GemmU8S8CopyPackBFrame.PaddedMatrixBData[rsp+32] - vmovdqu xmm5,XMMWORD PTR GemmU8S8CopyPackBFrame.PaddedMatrixBData[rsp+48] + vmovdqu xmm2,XMMWORD PTR GemmInt8CopyPackBFrame.PaddedMatrixBData[rsp] + vmovdqu xmm3,XMMWORD PTR GemmInt8CopyPackBFrame.PaddedMatrixBData[rsp+16] + vmovdqu xmm4,XMMWORD PTR GemmInt8CopyPackBFrame.PaddedMatrixBData[rsp+32] + vmovdqu xmm5,XMMWORD PTR GemmInt8CopyPackBFrame.PaddedMatrixBData[rsp+48] vpunpcklbw xmm6,xmm2,xmm3 ; interleave row data vpunpckhbw xmm3,xmm2,xmm3 vpunpcklbw xmm2,xmm4,xmm5 @@ -703,16 +765,28 @@ ProcessPaddedMatrixBData: vpunpckhwd xmm3,xmm3,xmm5 vinserti128 ymm4,ymm4,xmm6,1 vinserti128 ymm2,ymm2,xmm3,1 +IF IsVnni EQ 0 vpxor ymm4,ymm4,ymm9 ; optionally adjust unsigned data vpxor ymm2,ymm2,ymm9 +ENDIF vmovdqu YMMWORD PTR [rcx],ymm4 ; store interleaved rows vmovdqu YMMWORD PTR [rcx+32],ymm2 +IF IsVnni EQ 1 + IF BSigned EQ 1 + VpdpbssdYmmYmmYmm ymm0,ymm4,ymm8 + VpdpbssdYmmYmmYmm ymm1,ymm2,ymm8 + ELSE + VpdpbuudYmmYmmYmm ymm0,ymm4,ymm8 + VpdpbuudYmmYmmYmm ymm1,ymm2,ymm8 + ENDIF +ELSE vpmaddubsw ymm4,ymm8,ymm4 ; horizontal byte+byte=word per row vpmaddwd ymm4,ymm4,ymm7 ; horizontal word+word=dword per row vpaddd ymm0,ymm0,ymm4 ; accumulate per column vpmaddubsw ymm2,ymm8,ymm2 vpmaddwd ymm2,ymm2,ymm7 vpaddd ymm1,ymm1,ymm2 +ENDIF lea rsi,[rsi+r8*4] ; advance next matrix B by 4 rows add rcx,64 ; advance matrix D by 64 bytes sub r10,4 ; subtract rows remaining @@ -726,8 +800,8 @@ ProcessRemainingRowsNUnaligned: ; Process the less than 4 remaining rows where the row has less than 16 columns. ; -.errnz GemmU8S8CopyPackBFrame.PaddedMatrixBData - mov rbp,rsp ; GemmU8S8CopyPackBFrame.PaddedMatrixBData +.errnz GemmInt8CopyPackBFrame.PaddedMatrixBData + mov rbp,rsp ; GemmInt8CopyPackBFrame.PaddedMatrixBData vmovdqu YMMWORD PTR [rbp],ymm9 vmovdqu YMMWORD PTR [rbp+32],ymm9 @@ -775,6 +849,18 @@ StoreColumnSumBufferNUnaligned: vmovdqu YMMWORD PTR [r11+32],ymm1 jmp ExitRoutine +ENDM + + NESTED_ENTRY MlasGemmU8S8CopyPackBAvx2, _TEXT + MlasGemmCopyPackBAvx2 0 ; sign variable not checked if IsVnni = 0 NESTED_END MlasGemmU8S8CopyPackBAvx2, _TEXT + NESTED_ENTRY MlasGemmU8CopyPackBAvx2Vnni, _TEXT + MlasGemmCopyPackBAvx2 1, 0 + NESTED_END MlasGemmU8CopyPackBAvx2Vnni, _TEXT + + NESTED_ENTRY MlasGemmS8CopyPackBAvx2Vnni, _TEXT + MlasGemmCopyPackBAvx2 1, 1 + NESTED_END MlasGemmS8CopyPackBAvx2Vnni, _TEXT + END diff --git a/onnxruntime/core/mlas/lib/amd64/QgemmU8X8KernelAvx2.asm b/onnxruntime/core/mlas/lib/amd64/QgemmU8X8KernelAvx2.asm index 210f28bdd2..1705a15fa4 100644 --- a/onnxruntime/core/mlas/lib/amd64/QgemmU8X8KernelAvx2.asm +++ b/onnxruntime/core/mlas/lib/amd64/QgemmU8X8KernelAvx2.asm @@ -14,6 +14,7 @@ ; multiply operation (QGEMM). ; ; This implementation uses AVX2 and AVX VNNI instructions. +; AVX-VNNI-INT8 support also included. ; ;-- @@ -25,10 +26,10 @@ INCLUDE AssembleAvxVnni.inc EXTERN MlasMaskMoveTableAvx:NEAR ; -; Stack frame layout for the U8X8 kernel. +; Stack frame layout for the Int8 kernel. ; -GemmU8X8KernelFrame STRUCT +GemmInt8KernelFrame STRUCT SavedXmm6 OWORD ? SavedXmm7 OWORD ? @@ -60,7 +61,7 @@ GemmU8X8KernelFrame STRUCT ZeroPointB QWORD ? ZeroMode QWORD ? -GemmU8X8KernelFrame ENDS +GemmInt8KernelFrame ENDS ; ; Macro Description: @@ -135,7 +136,7 @@ ENDIF ; ymm12 - Supplies a 256-bit with the broadcasted word value 0x0001. ; -ComputeBlockU8S8Avx2 MACRO ColumnCount, RowCount, VectorOffset, BroadcastOffset +ComputeBlockAvx2 MACRO ColumnCount, RowCount, VectorOffset, BroadcastOffset, ASigned, BSigned IF RowCount EQ 1 vpbroadcastd ymm2,DWORD PTR [rcx+BroadcastOffset] @@ -189,13 +190,40 @@ ENDIF ; ymm2 - Supplies the broadcast value loaded from matrix A. ; -MultiplyAccumulateRowU8S8AvxVnni MACRO ColumnCount, Vec1Reg, Vec2Reg +MultiplyAccumulateRowAvxVnni MACRO ColumnCount, Vec1Reg, Vec2Reg, ASigned, BSigned -IF ColumnCount EQ 16 - VpdpbusdsYmmYmmYmm Vec1Reg,ymm2,ymm0 - VpdpbusdsYmmYmmYmm Vec2Reg,ymm2,ymm1 +IF ASigned EQ 1 + IF BSigned EQ 1 + IF ColumnCount EQ 16 + VpdpbssdYmmYmmYmm Vec1Reg,ymm2,ymm0 + VpdpbssdYmmYmmYmm Vec2Reg,ymm2,ymm1 + ELSE + VpdpbssdYmmYmmYmm Vec2Reg,ymm2,ymm0 + ENDIF + ELSE + IF ColumnCount EQ 16 + VpdpbsudYmmYmmYmm Vec1Reg,ymm2,ymm0 + VpdpbsudYmmYmmYmm Vec2Reg,ymm2,ymm1 + ELSE + VpdpbsudYmmYmmYmm Vec2Reg,ymm2,ymm0 + ENDIF + ENDIF ELSE - VpdpbusdsYmmYmmYmm Vec2Reg,ymm2,ymm0 + IF BSigned EQ 1 + IF ColumnCount EQ 16 + VpdpbusdYmmYmmYmm Vec1Reg,ymm2,ymm0 + VpdpbusdYmmYmmYmm Vec2Reg,ymm2,ymm1 + ELSE + VpdpbusdYmmYmmYmm Vec2Reg,ymm2,ymm0 + ENDIF + ELSE + IF ColumnCount EQ 16 + VpdpbuudYmmYmmYmm Vec1Reg,ymm2,ymm0 + VpdpbuudYmmYmmYmm Vec2Reg,ymm2,ymm1 + ELSE + VpdpbuudYmmYmmYmm Vec2Reg,ymm2,ymm0 + ENDIF + ENDIF ENDIF ENDM @@ -229,22 +257,22 @@ ENDIF ; ymm4-ymm15 - Supplies the block accumulators. ; -ComputeBlockU8S8AvxVnni MACRO ColumnCount, RowCount, VectorOffset, BroadcastOffset +ComputeBlockAvxVnni MACRO ColumnCount, RowCount, VectorOffset, BroadcastOffset, ASigned, BSigned vmovdqu ymm0,YMMWORD PTR [rdx+VectorOffset] EmitIfCountGE ColumnCount, 16, EmitIfCountGE RowCount, 1, - EmitIfCountGE RowCount, 1, + EmitIfCountGE RowCount, 1, EmitIfCountGE RowCount, 2, - EmitIfCountGE RowCount, 2, + EmitIfCountGE RowCount, 2, EmitIfCountGE RowCount, 3, - EmitIfCountGE RowCount, 3, + EmitIfCountGE RowCount, 3, EmitIfCountGE RowCount, 4, - EmitIfCountGE RowCount, 4, + EmitIfCountGE RowCount, 4, EmitIfCountGE RowCount, 5, - EmitIfCountGE RowCount, 5, + EmitIfCountGE RowCount, 5, EmitIfCountGE RowCount, 6, - EmitIfCountGE RowCount, 6, + EmitIfCountGE RowCount, 6, ENDM @@ -275,7 +303,7 @@ ComputeBlockU8S8AvxVnni MACRO ColumnCount, RowCount, VectorOffset, BroadcastOffs ; ymm4-ymm11 - Supplies the block accumulators. ; -ComputeBlockLoopU8S8 MACRO Isa, ColumnCount, RowCount +ComputeBlockLoop MACRO Isa, ColumnCount, RowCount, ASigned, BSigned LOCAL ComputeBlockBy4Loop LOCAL ProcessRemainingBlocks @@ -289,10 +317,10 @@ IF (ColumnCount EQ 16) AND (RowCount EQ 1) jb ProcessRemainingBlocks ComputeBlockBy4Loop: - ComputeBlockU8S8&Isa& ColumnCount, RowCount, 0*64, 0 - ComputeBlockU8S8&Isa& ColumnCount, RowCount, 1*64, 4 - ComputeBlockU8S8&Isa& ColumnCount, RowCount, 2*64, 8 - ComputeBlockU8S8&Isa& ColumnCount, RowCount, 3*64, 12 + ComputeBlock&Isa& ColumnCount, RowCount, 0*64, 0, ASigned, BSigned + ComputeBlock&Isa& ColumnCount, RowCount, 1*64, 4, ASigned, BSigned + ComputeBlock&Isa& ColumnCount, RowCount, 2*64, 8, ASigned, BSigned + ComputeBlock&Isa& ColumnCount, RowCount, 3*64, 12, ASigned, BSigned add rcx,4*4 ; advance matrix A by 4 quads add rdx,4*64 ; advance matrix B sub rsi,4*4 @@ -304,7 +332,7 @@ ProcessRemainingBlocks: ENDIF ComputeBlockBy1Loop: - ComputeBlockU8S8&Isa& ColumnCount, RowCount, 0, 0 + ComputeBlock&Isa& ColumnCount, RowCount, 0, 0, ASigned, BSigned add rcx,4 ; advance matrix A by 1 quad IF RowCount GT 3 add rbx,4 ; advance matrix A plus 3 rows by 1 quad @@ -506,11 +534,11 @@ ExitComputeBlockLoop: ; ymm4-ymm15 - Supplies the block accumulators. ; -ProduceOutputBlock MACRO ColumnCount, RowCount +ProduceOutputBlock MACRO ColumnCount, RowCount, ASigned, BSigned LOCAL SkipScaleByZeroPointB LOCAL AccumulatorsInitialized - LOCAL ProduceWithU8S8AvxVnni + LOCAL ProduceWithInt8AvxVnni LOCAL ProduceWithU8U8Avx2 LOCAL ExitProduceOutputBlock @@ -590,16 +618,16 @@ IF RowCount GT 3 lea rbx,[r9*2+r9] add rbx,rcx ; compute matrix A plus 3 rows ENDIF - cmp DWORD PTR GemmU8X8KernelFrame.PreviousP1Home[rsp],0 + cmp DWORD PTR GemmInt8KernelFrame.PreviousP1Home[rsp],0 jg ProduceWithU8U8Avx2 IF RowCount LE 4 - jl ProduceWithU8S8AvxVnni - ComputeBlockLoopU8S8 Avx2, ColumnCount, RowCount + jl ProduceWithInt8AvxVnni + ComputeBlockLoop Avx2, ColumnCount, RowCount, ASigned, BSigned jmp ExitProduceOutputBlock ENDIF -ProduceWithU8S8AvxVnni: - ComputeBlockLoopU8S8 AvxVnni, ColumnCount, RowCount +ProduceWithInt8AvxVnni: + ComputeBlockLoop AvxVnni, ColumnCount, RowCount, ASigned, BSigned jmp ExitProduceOutputBlock ProduceWithU8U8Avx2: @@ -649,7 +677,7 @@ ENDIF ; r13 - Optionally supplies the address of the matrix B zero point buffer. ; -ProcessCountM MACRO RowCount, Fallthrough +ProcessCountM MACRO RowCount, ASigned, BSigned, Fallthrough LOCAL ProcessNextColumnLoop16xN LOCAL SkipAccumulateOutput16xNBlock @@ -665,7 +693,7 @@ ProcessCountM MACRO RowCount, Fallthrough jbe ProcessRemainingCountN ProcessNextColumnLoop16xN: - ProduceOutputBlock 16, RowCount + ProduceOutputBlock 16, RowCount, ASigned, BSigned sub rbp,16 jb OutputMasked16xNBlock test r10b,r10b ; ZeroMode? @@ -708,7 +736,7 @@ ExitProcessCountM: jmp ExitKernel ProcessRemainingCountN: - ProduceOutputBlock 8, RowCount + ProduceOutputBlock 8, RowCount, ASigned, BSigned cmp rbp,8 jb OutputMasked8xNBlock test r10b,r10b ; ZeroMode? @@ -782,31 +810,6 @@ SkipAccumulateOutputMasked8xNBlock: ENDM -; -; Reduce code size for the various types of kernels by sharing the outer logic -; and switching on the selector codes (using sign bit to discriminate). -; - - LEAF_ENTRY MlasGemmU8S8KernelAvxVnni, _TEXT - - mov eax,-1 - jmp MlasGemmU8X8KernelAvx2 - - LEAF_END MlasGemmU8S8KernelAvxVnni, _TEXT - - LEAF_ENTRY MlasGemmU8U8KernelAvx2, _TEXT - - mov eax,1 - jmp MlasGemmU8X8KernelAvx2 - - LEAF_END MlasGemmU8U8KernelAvx2, _TEXT - - LEAF_ENTRY MlasGemmU8S8KernelAvx2, _TEXT - - xor eax,eax - jmp MlasGemmU8X8KernelAvx2 - - LEAF_END MlasGemmU8S8KernelAvx2, _TEXT ;++ ; @@ -818,10 +821,10 @@ SkipAccumulateOutputMasked8xNBlock: ; Arguments: ; ; A (rcx) - Supplies the address of matrix A. The matrix data has been packed -; using MlasGemmU8X8CopyPackAAvx2. +; using MlasGemmCopyPackAAvx2. ; ; B (rdx) - Supplies the address of matrix B. The matrix data has been packed -; using MlasGemmU8X8CopyPackBAvx2. +; using MlasGemmCopyPackBAvx2. ; ; C (r8) - Supplies the address of matrix C. ; @@ -859,7 +862,7 @@ SkipAccumulateOutputMasked8xNBlock: ; ;-- - NESTED_ENTRY MlasGemmU8X8KernelAvx2, _TEXT +MlasGemmInt8KernelAvx2 MACRO ASigned, BSigned rex_push_reg rbp push_reg rbx @@ -867,34 +870,34 @@ SkipAccumulateOutputMasked8xNBlock: push_reg rdi push_reg r12 push_reg r13 - alloc_stack (GemmU8X8KernelFrame.SavedR13) - save_xmm128 xmm6,GemmU8X8KernelFrame.SavedXmm6 - save_xmm128 xmm7,GemmU8X8KernelFrame.SavedXmm7 - save_xmm128 xmm8,GemmU8X8KernelFrame.SavedXmm8 - save_xmm128 xmm9,GemmU8X8KernelFrame.SavedXmm9 - save_xmm128 xmm10,GemmU8X8KernelFrame.SavedXmm10 - save_xmm128 xmm11,GemmU8X8KernelFrame.SavedXmm11 - save_xmm128 xmm12,GemmU8X8KernelFrame.SavedXmm12 - save_xmm128 xmm13,GemmU8X8KernelFrame.SavedXmm13 - save_xmm128 xmm14,GemmU8X8KernelFrame.SavedXmm14 - save_xmm128 xmm15,GemmU8X8KernelFrame.SavedXmm15 + alloc_stack (GemmInt8KernelFrame.SavedR13) + save_xmm128 xmm6,GemmInt8KernelFrame.SavedXmm6 + save_xmm128 xmm7,GemmInt8KernelFrame.SavedXmm7 + save_xmm128 xmm8,GemmInt8KernelFrame.SavedXmm8 + save_xmm128 xmm9,GemmInt8KernelFrame.SavedXmm9 + save_xmm128 xmm10,GemmInt8KernelFrame.SavedXmm10 + save_xmm128 xmm11,GemmInt8KernelFrame.SavedXmm11 + save_xmm128 xmm12,GemmInt8KernelFrame.SavedXmm12 + save_xmm128 xmm13,GemmInt8KernelFrame.SavedXmm13 + save_xmm128 xmm14,GemmInt8KernelFrame.SavedXmm14 + save_xmm128 xmm15,GemmInt8KernelFrame.SavedXmm15 END_PROLOGUE - mov DWORD PTR GemmU8X8KernelFrame.PreviousP1Home[rsp],eax + mov DWORD PTR GemmInt8KernelFrame.PreviousP1Home[rsp],eax mov rdi,rcx - mov rbx,GemmU8X8KernelFrame.CountM[rsp] - mov rbp,GemmU8X8KernelFrame.CountN[rsp] - mov rax,GemmU8X8KernelFrame.ldc[rsp] + mov rbx,GemmInt8KernelFrame.CountM[rsp] + mov rbp,GemmInt8KernelFrame.CountN[rsp] + mov rax,GemmInt8KernelFrame.ldc[rsp] shl rax,2 ; convert ldc to bytes shl r9,2 ; convert to row length - movzx r10,BYTE PTR GemmU8X8KernelFrame.ZeroMode[rsp] - mov r11,GemmU8X8KernelFrame.RowSumBuffer[rsp] - mov r12,GemmU8X8KernelFrame.ColumnSumBuffer[rsp] - mov r13,GemmU8X8KernelFrame.ZeroPointB[rsp] + movzx r10,BYTE PTR GemmInt8KernelFrame.ZeroMode[rsp] + mov r11,GemmInt8KernelFrame.RowSumBuffer[rsp] + mov r12,GemmInt8KernelFrame.ColumnSumBuffer[rsp] + mov r13,GemmInt8KernelFrame.ZeroPointB[rsp] vpcmpeqw ymm12,ymm12,ymm12 ; generate 256-bit word vector [0xFFFF] vpsrlw ymm12,ymm12,15 ; generate 256-bit word vector [0x0001] - cmp DWORD PTR GemmU8X8KernelFrame.PreviousP1Home[rsp],0 + cmp DWORD PTR GemmInt8KernelFrame.PreviousP1Home[rsp],0 je CheckCountM4OrMore ; U8S8 AVX2 kernel requires extra registers ; @@ -914,13 +917,13 @@ CheckCountM4OrMore: je ProcessCountM1 ProcessCountM2: - ProcessCountM 2 + ProcessCountM 2, ASigned, BSigned ProcessCountM4: - ProcessCountM 4 + ProcessCountM 4, ASigned, BSigned ProcessCountM6: - ProcessCountM 6 + ProcessCountM 6, ASigned, BSigned ; ; Restore non-volatile registers and return. @@ -928,17 +931,17 @@ ProcessCountM6: ExitKernel: vzeroupper - movaps xmm6,GemmU8X8KernelFrame.SavedXmm6[rsp] - movaps xmm7,GemmU8X8KernelFrame.SavedXmm7[rsp] - movaps xmm8,GemmU8X8KernelFrame.SavedXmm8[rsp] - movaps xmm9,GemmU8X8KernelFrame.SavedXmm9[rsp] - movaps xmm10,GemmU8X8KernelFrame.SavedXmm10[rsp] - movaps xmm11,GemmU8X8KernelFrame.SavedXmm11[rsp] - movaps xmm12,GemmU8X8KernelFrame.SavedXmm12[rsp] - movaps xmm13,GemmU8X8KernelFrame.SavedXmm13[rsp] - movaps xmm14,GemmU8X8KernelFrame.SavedXmm14[rsp] - movaps xmm15,GemmU8X8KernelFrame.SavedXmm15[rsp] - add rsp,(GemmU8X8KernelFrame.SavedR13) + movaps xmm6,GemmInt8KernelFrame.SavedXmm6[rsp] + movaps xmm7,GemmInt8KernelFrame.SavedXmm7[rsp] + movaps xmm8,GemmInt8KernelFrame.SavedXmm8[rsp] + movaps xmm9,GemmInt8KernelFrame.SavedXmm9[rsp] + movaps xmm10,GemmInt8KernelFrame.SavedXmm10[rsp] + movaps xmm11,GemmInt8KernelFrame.SavedXmm11[rsp] + movaps xmm12,GemmInt8KernelFrame.SavedXmm12[rsp] + movaps xmm13,GemmInt8KernelFrame.SavedXmm13[rsp] + movaps xmm14,GemmInt8KernelFrame.SavedXmm14[rsp] + movaps xmm15,GemmInt8KernelFrame.SavedXmm15[rsp] + add rsp,(GemmInt8KernelFrame.SavedR13) BEGIN_EPILOGUE @@ -951,14 +954,61 @@ ExitKernel: ret ProcessCountM1: - ProcessCountM 1 + ProcessCountM 1, ASigned, BSigned ProcessCountM3: - ProcessCountM 3 + ProcessCountM 3, ASigned, BSigned ProcessCountM5: - ProcessCountM 5 + ProcessCountM 5, ASigned, BSigned - NESTED_END MlasGemmU8X8KernelAvx2, _TEXT + ENDM + +; +; Reduce code size for the various types of kernels by sharing the outer logic +; and switching on the selector codes (using sign bit to discriminate). +; + + NESTED_ENTRY MlasGemmU8S8KernelAvxVnni, _TEXT + + mov eax,-1 + MlasGemmInt8KernelAvx2 0, 1 + + NESTED_END MlasGemmU8S8KernelAvxVnni, _TEXT + + NESTED_ENTRY MlasGemmU8U8KernelAvx2Vnni, _TEXT + + mov eax,-1 + MlasGemmInt8KernelAvx2 0, 0 + + NESTED_END MlasGemmU8U8KernelAvx2Vnni, _TEXT + + NESTED_ENTRY MlasGemmU8U8KernelAvx2, _TEXT + + mov eax,1 + MlasGemmInt8KernelAvx2 0, 0 + + NESTED_END MlasGemmU8U8KernelAvx2, _TEXT + + NESTED_ENTRY MlasGemmU8S8KernelAvx2, _TEXT + + xor eax,eax + MlasGemmInt8KernelAvx2 0, 1 + + NESTED_END MlasGemmU8S8KernelAvx2, _TEXT + + NESTED_ENTRY MlasGemmS8S8KernelAvx2Vnni, _TEXT + + mov eax,-1 + MlasGemmInt8KernelAvx2 1, 1 + + NESTED_END MlasGemmS8S8KernelAvx2Vnni, _TEXT + + NESTED_ENTRY MlasGemmS8U8KernelAvx2Vnni, _TEXT + + mov eax,-1 + MlasGemmInt8KernelAvx2 1, 0 + + NESTED_END MlasGemmS8U8KernelAvx2Vnni, _TEXT END diff --git a/onnxruntime/core/mlas/lib/mlasi.h b/onnxruntime/core/mlas/lib/mlasi.h index 8e8f46b8a1..96ba8c6c92 100644 --- a/onnxruntime/core/mlas/lib/mlasi.h +++ b/onnxruntime/core/mlas/lib/mlasi.h @@ -804,6 +804,9 @@ extern "C" { MLAS_GEMV_U8S8_KERNEL MlasGemvU8S8KernelAvx512Vnni; MLAS_GEMM_U8S8_KERNEL MlasGemmU8S8KernelAvxVnni; MLAS_GEMV_U8S8_KERNEL MlasGemvU8S8KernelAvxVnni; + MLAS_GEMM_U8S8_KERNEL MlasGemmU8U8KernelAvx2Vnni; + MLAS_GEMM_U8S8_KERNEL MlasGemmS8S8KernelAvx2Vnni; + MLAS_GEMM_U8S8_KERNEL MlasGemmS8U8KernelAvx2Vnni; MLAS_GEMM_U8U8_KERNEL MlasGemmU8U8KernelAvx2; MLAS_GEMM_U8U8_KERNEL MlasGemmU8U8KernelAvx512Core; #endif @@ -954,6 +957,9 @@ extern const MLAS_GEMM_QUANT_DISPATCH MlasGemmU8X8DispatchLSX; extern const MLAS_GEMM_QUANT_DISPATCH MlasGemmU8S8DispatchSse41; extern const MLAS_GEMM_QUANT_DISPATCH MlasGemmU8S8DispatchAvx2; extern const MLAS_GEMM_QUANT_DISPATCH MlasGemmU8U8DispatchAvx2; +extern const MLAS_GEMM_QUANT_DISPATCH MlasGemmU8U8DispatchAvx2Vnni; +extern const MLAS_GEMM_QUANT_DISPATCH MlasGemmS8S8DispatchAvx2Vnni; +extern const MLAS_GEMM_QUANT_DISPATCH MlasGemmS8U8DispatchAvx2Vnni; extern const MLAS_GEMM_QUANT_DISPATCH MlasGemmU8S8DispatchAmx; extern const MLAS_GEMM_QUANT_DISPATCH MlasGemmU8X8DispatchNeon; extern const MLAS_GEMM_QUANT_DISPATCH MlasGemmX8S8DispatchNeon; @@ -1103,6 +1109,8 @@ struct MLAS_PLATFORM { #if defined(MLAS_TARGET_AMD64_IX86) const MLAS_GEMM_QUANT_DISPATCH* GemmU8S8Dispatch; const MLAS_GEMM_QUANT_DISPATCH* GemmU8U8Dispatch; + const MLAS_GEMM_QUANT_DISPATCH* GemmS8S8Dispatch{&MlasGemmQuantDispatchDefault}; + const MLAS_GEMM_QUANT_DISPATCH* GemmS8U8Dispatch{&MlasGemmQuantDispatchDefault}; #elif defined(MLAS_TARGET_ARM64) const MLAS_GEMM_QUANT_DISPATCH* GemmU8U8Dispatch; const MLAS_GEMM_QUANT_DISPATCH* GemmU8S8Dispatch; @@ -1134,6 +1142,8 @@ struct MLAS_PLATFORM { MLAS_SGEMM_TRANSPOSE_PACKB_BLOCK_ROUTINE* TransposePackB16x4Routine; MLAS_GEMM_DOUBLE_KERNEL* GemmDoubleKernel; MLAS_GEMM_U8S8_KERNEL* GemmU8S8Kernel; + MLAS_GEMM_U8S8_KERNEL* GemmS8S8Kernel; + MLAS_GEMM_U8S8_KERNEL* GemmS8U8Kernel; MLAS_GEMV_U8S8_KERNEL* GemvU8S8Kernel; MLAS_GEMM_U8U8_KERNEL* GemmU8U8Kernel; MLAS_CONV_FLOAT_KERNEL* ConvNchwFloatKernel; diff --git a/onnxruntime/core/mlas/lib/platform.cpp b/onnxruntime/core/mlas/lib/platform.cpp index 2b4d99800c..102d605227 100644 --- a/onnxruntime/core/mlas/lib/platform.cpp +++ b/onnxruntime/core/mlas/lib/platform.cpp @@ -476,6 +476,17 @@ Return Value: } } + // + // Check if the processor supports AVX-VNNI-INT8 + // + if ((Cpuid7_1[3] & 0x10) != 0) { + this->GemmU8U8Dispatch = &MlasGemmU8U8DispatchAvx2Vnni; + this->GemmS8S8Dispatch = &MlasGemmS8S8DispatchAvx2Vnni; + this->GemmS8S8Kernel = MlasGemmS8S8KernelAvx2Vnni; + this->GemmS8U8Dispatch = &MlasGemmS8U8DispatchAvx2Vnni; + this->GemmS8U8Kernel = MlasGemmS8U8KernelAvx2Vnni; + } + #ifndef __APPLE__ #if (defined(_MSC_VER) && (_MSC_VER >= 1933)) || (defined(__GNUC__) && (__GNUC__ >= 13)) // diff --git a/onnxruntime/core/mlas/lib/qgemm.h b/onnxruntime/core/mlas/lib/qgemm.h index 127aea9029..1ef5b5f741 100644 --- a/onnxruntime/core/mlas/lib/qgemm.h +++ b/onnxruntime/core/mlas/lib/qgemm.h @@ -337,7 +337,7 @@ Return Value: // // Fixup the sign bit of the per-matrix zero point offset of matrix A if the - // kernel requires signed data. + // kernel requires opposite-signed data. // ZeroPointA = MlasGemmQuantFixupZeroPointA(ZeroPointA, Shape->AIsSigned); @@ -865,20 +865,15 @@ MlasGemmQuantGetDispatch( bool BIsSigned ) { - const MLAS_GEMM_QUANT_DISPATCH* GemmQuantDispatch = nullptr; - - if (!AIsSigned || BIsSigned) { - GemmQuantDispatch = &MlasGemmQuantDispatchDefault; - } + const MLAS_GEMM_QUANT_DISPATCH* GemmQuantDispatch = &MlasGemmQuantDispatchDefault; #if defined(MLAS_TARGET_AMD64_IX86) || defined(MLAS_TARGET_LARCH64) - if (!AIsSigned) { - if (BIsSigned) { - GemmQuantDispatch = GetMlasPlatform().GemmU8S8Dispatch; - } - else { - GemmQuantDispatch = GetMlasPlatform().GemmU8U8Dispatch; - } + if (AIsSigned) { + GemmQuantDispatch = + BIsSigned ? GetMlasPlatform().GemmS8S8Dispatch : GetMlasPlatform().GemmS8U8Dispatch; + } else { + GemmQuantDispatch = + BIsSigned ? GetMlasPlatform().GemmU8S8Dispatch : GetMlasPlatform().GemmU8U8Dispatch; } #elif defined(MLAS_TARGET_ARM64) if(BIsSigned) { diff --git a/onnxruntime/core/mlas/lib/qgemm_kernel_avx2.cpp b/onnxruntime/core/mlas/lib/qgemm_kernel_avx2.cpp index deec324d01..a6dbe8defd 100644 --- a/onnxruntime/core/mlas/lib/qgemm_kernel_avx2.cpp +++ b/onnxruntime/core/mlas/lib/qgemm_kernel_avx2.cpp @@ -74,6 +74,39 @@ extern "C" { size_t CountK, int32_t* ColumnSumBuffer ); + + void + MLASCALL + MlasGemmS8CopyPackAAvx2Vnni( + uint8_t* D, + const uint8_t* A, + size_t lda, + size_t CountM, + size_t CountK, + int32_t* RowSumBuffer + ); + + void + MLASCALL + MlasGemmU8CopyPackBAvx2Vnni( + uint8_t* D, + const uint8_t* B, + size_t ldb, + size_t CountN, + size_t CountK, + int32_t* ColumnSumBuffer + ); + + void + MLASCALL + MlasGemmS8CopyPackBAvx2Vnni( + uint8_t* D, + const uint8_t* B, + size_t ldb, + size_t CountN, + size_t CountK, + int32_t* ColumnSumBuffer + ); } struct MLAS_GEMM_U8S8_KERNEL_AVX2 @@ -273,3 +306,222 @@ const MLAS_GEMM_QUANT_DISPATCH MlasGemmU8U8DispatchAvx2 = { MLAS_GEMM_U8U8_KERNEL_AVX2::PackedStrides.K, 6 // assembly kernel M stride }; + +// U8U8 AVX-VNNI-INT8 support +struct MLAS_GEMM_U8U8_KERNEL_AVX2VNNI { + typedef uint8_t PackedAType; + typedef uint8_t PackedBType; + typedef uint8_t OffsetAType; + typedef uint8_t OffsetBType; + + static constexpr size_t PackedK = 4; + static constexpr MLAS_GEMM_QUANT_STRIDES Strides{24, 256, 128}; + static constexpr MLAS_GEMM_QUANT_STRIDES PackedStrides{48, 256, 384}; +}; + +template <> +MLAS_FORCEINLINE void +MlasGemmQuantCopyPackA( + MLAS_GEMM_U8U8_KERNEL_AVX2VNNI::PackedAType* D, + const uint8_t* A, + size_t lda, + size_t CountM, + size_t CountK, + int32_t* RowSumBuffer, + bool AIsSigned +) +{ + MLAS_UNREFERENCED_PARAMETER(AIsSigned); + MlasGemmU8S8CopyPackAAvx2(D, A, lda, CountM, CountK, RowSumBuffer); +} + +template <> +MLAS_FORCEINLINE void +MlasGemmQuantCopyPackB( + MLAS_GEMM_U8U8_KERNEL_AVX2VNNI::PackedBType* D, + const uint8_t* B, + size_t ldb, + size_t CountN, + size_t CountK, + int32_t* ColumnSumBuffer, + bool BIsSigned +) +{ + MLAS_UNREFERENCED_PARAMETER(BIsSigned); + MlasGemmU8CopyPackBAvx2Vnni(D, B, ldb, CountN, CountK, ColumnSumBuffer); +} + +template <> +MLAS_FORCEINLINE + size_t + MlasGemmQuantKernel( + const MLAS_GEMM_U8U8_KERNEL_AVX2VNNI::PackedAType* A, + const MLAS_GEMM_U8U8_KERNEL_AVX2VNNI::PackedBType* B, + int32_t* C, + size_t PackedCountK, + size_t CountM, + size_t CountN, + size_t ldc, + const int32_t* RowSumBuffer, + const int32_t* ColumnSumBuffer, + const int32_t* ZeroPointB, + bool ZeroMode + ) +{ + return MlasGemmU8U8KernelAvx2Vnni(A, B, C, PackedCountK, CountM, CountN, ldc, RowSumBuffer, ColumnSumBuffer, ZeroPointB, ZeroMode); +} + +const MLAS_GEMM_QUANT_DISPATCH MlasGemmU8U8DispatchAvx2Vnni = { + MlasGemmQuantOperation, + MlasGemmQuantPackedOperation, + MlasGemmQuantCopyPackB, + MLAS_GEMM_U8U8_KERNEL_AVX2VNNI::PackedK, + MLAS_GEMM_U8U8_KERNEL_AVX2VNNI::PackedStrides.K, + 6 // assembly kernel M stride +}; + +// S8S8 AVX-VNNI-INT8 support +struct MLAS_GEMM_S8S8_KERNEL_AVX2 { + typedef uint8_t PackedAType; + typedef uint8_t PackedBType; + typedef int8_t OffsetAType; + typedef int8_t OffsetBType; + + static constexpr size_t PackedK = 4; + static constexpr MLAS_GEMM_QUANT_STRIDES Strides{24, 256, 128}; + static constexpr MLAS_GEMM_QUANT_STRIDES PackedStrides{48, 256, 384}; +}; + +template <> +MLAS_FORCEINLINE void +MlasGemmQuantCopyPackA( + MLAS_GEMM_S8S8_KERNEL_AVX2::PackedAType* D, + const uint8_t* A, + size_t lda, + size_t CountM, + size_t CountK, + int32_t* RowSumBuffer, + bool AIsSigned +) +{ + MLAS_UNREFERENCED_PARAMETER(AIsSigned); + MlasGemmS8CopyPackAAvx2Vnni(D, A, lda, CountM, CountK, RowSumBuffer); +} + +template <> +MLAS_FORCEINLINE void +MlasGemmQuantCopyPackB( + MLAS_GEMM_S8S8_KERNEL_AVX2::PackedBType* D, + const uint8_t* B, + size_t ldb, + size_t CountN, + size_t CountK, + int32_t* ColumnSumBuffer, + bool BIsSigned +) +{ + MLAS_UNREFERENCED_PARAMETER(BIsSigned); + MlasGemmS8CopyPackBAvx2Vnni(D, B, ldb, CountN, CountK, ColumnSumBuffer); +} + +template <> +MLAS_FORCEINLINE + size_t + MlasGemmQuantKernel( + const MLAS_GEMM_S8S8_KERNEL_AVX2::PackedAType* A, + const MLAS_GEMM_S8S8_KERNEL_AVX2::PackedBType* B, + int32_t* C, + size_t PackedCountK, + size_t CountM, + size_t CountN, + size_t ldc, + const int32_t* RowSumBuffer, + const int32_t* ColumnSumBuffer, + const int32_t* ZeroPointB, + bool ZeroMode + ) +{ + return GetMlasPlatform().GemmS8S8Kernel(A, B, C, PackedCountK, CountM, CountN, ldc, RowSumBuffer, ColumnSumBuffer, ZeroPointB, ZeroMode); +} + +const MLAS_GEMM_QUANT_DISPATCH MlasGemmS8S8DispatchAvx2Vnni = { + MlasGemmQuantOperation, + MlasGemmQuantPackedOperation, + MlasGemmQuantCopyPackB, + MLAS_GEMM_S8S8_KERNEL_AVX2::PackedK, + MLAS_GEMM_S8S8_KERNEL_AVX2::PackedStrides.K, + 6 // assembly kernel M stride +}; + +// S8U8 AVX-VNNI-INT8 support +struct MLAS_GEMM_S8U8_KERNEL_AVX2 { + typedef uint8_t PackedAType; + typedef uint8_t PackedBType; + typedef int8_t OffsetAType; + typedef uint8_t OffsetBType; + + static constexpr size_t PackedK = 4; + static constexpr MLAS_GEMM_QUANT_STRIDES Strides{24, 256, 128}; + static constexpr MLAS_GEMM_QUANT_STRIDES PackedStrides{48, 256, 384}; +}; + +template <> +MLAS_FORCEINLINE void +MlasGemmQuantCopyPackA( + MLAS_GEMM_S8U8_KERNEL_AVX2::PackedAType* D, + const uint8_t* A, + size_t lda, + size_t CountM, + size_t CountK, + int32_t* RowSumBuffer, + bool AIsSigned +) +{ + MLAS_UNREFERENCED_PARAMETER(AIsSigned); + MlasGemmS8CopyPackAAvx2Vnni(D, A, lda, CountM, CountK, RowSumBuffer); +} + +template <> +MLAS_FORCEINLINE void +MlasGemmQuantCopyPackB( + MLAS_GEMM_S8U8_KERNEL_AVX2::PackedBType* D, + const uint8_t* B, + size_t ldb, + size_t CountN, + size_t CountK, + int32_t* ColumnSumBuffer, + bool BIsSigned +) +{ + MLAS_UNREFERENCED_PARAMETER(BIsSigned); + MlasGemmU8CopyPackBAvx2Vnni(D, B, ldb, CountN, CountK, ColumnSumBuffer); +} + +template <> +MLAS_FORCEINLINE + size_t + MlasGemmQuantKernel( + const MLAS_GEMM_S8U8_KERNEL_AVX2::PackedAType* A, + const MLAS_GEMM_S8U8_KERNEL_AVX2::PackedBType* B, + int32_t* C, + size_t PackedCountK, + size_t CountM, + size_t CountN, + size_t ldc, + const int32_t* RowSumBuffer, + const int32_t* ColumnSumBuffer, + const int32_t* ZeroPointB, + bool ZeroMode + ) +{ + return GetMlasPlatform().GemmS8U8Kernel(A, B, C, PackedCountK, CountM, CountN, ldc, RowSumBuffer, ColumnSumBuffer, ZeroPointB, ZeroMode); +} + +const MLAS_GEMM_QUANT_DISPATCH MlasGemmS8U8DispatchAvx2Vnni = { + MlasGemmQuantOperation, + MlasGemmQuantPackedOperation, + MlasGemmQuantCopyPackB, + MLAS_GEMM_S8U8_KERNEL_AVX2::PackedK, + MLAS_GEMM_S8U8_KERNEL_AVX2::PackedStrides.K, + 6 // assembly kernel M stride +}; diff --git a/onnxruntime/core/mlas/lib/x86_64/AssembleAvxVnni.h b/onnxruntime/core/mlas/lib/x86_64/AssembleAvxVnni.h index 0d025ce483..676102391d 100644 --- a/onnxruntime/core/mlas/lib/x86_64/AssembleAvxVnni.h +++ b/onnxruntime/core/mlas/lib/x86_64/AssembleAvxVnni.h @@ -176,3 +176,153 @@ Arguments: VnniXmmXmmXmm 0x53, \DestReg\(), \Src1Reg\(), \Src2Reg\() .endm + +/*++ + +Macro Description: + + This macro builds a VNNI instruction of the form: + + instr ymm1,ymm2,ymm3 + +Arguments: + + Opcode - Specifies the opcode for the VNNI instruction. + + Prefix - Specifies the opcode prefix for payload 1 + + DestReg - Specifies the destination register. + + Src1Reg - Specifies the first source register. + + Src2Reg - Specifies the second source register. + +--*/ + .macro Avx2VnniYmmYmmYmm Opcode, Prefix, DestReg, Src1Reg, Src2Reg + + .set Payload0, 0x02 # "0F 38" prefix + .set Payload0, Payload0 + ((((.LYmmIndex_\DestReg\() >> 3) & 1) ^ 1) << 7) + .set Payload0, Payload0 + (1 << 6) + .set Payload0, Payload0 + ((((.LYmmIndex_\Src2Reg\() >> 3) & 1) ^ 1) << 5) + + .set Payload1, 0x04 + \Prefix\() # 256-bit length and opcode prefix + .set Payload1, Payload1 + (((.LYmmIndex_\Src1Reg\() & 15) ^ 15) << 3) + + .set ModRMByte, 0xC0 # register form + .set ModRMByte, ModRMByte + ((.LYmmIndex_\DestReg\() & 7) << 3) + .set ModRMByte, ModRMByte + (.LYmmIndex_\Src2Reg\() & 7) + + .byte 0xC4, Payload0, Payload1, \Opcode\(), ModRMByte + + .endm + + .macro VpdpbssdYmmYmmYmm DestReg, Src1Reg, Src2Reg + + Avx2VnniYmmYmmYmm 0x50, 0x03, \DestReg\(), \Src1Reg\(), \Src2Reg\() + + .endm + + .macro VpdpbssdsYmmYmmYmm DestReg, Src1Reg, Src2Reg + + Avx2VnniYmmYmmYmm 0x51, 0x03, \DestReg\(), \Src1Reg\(), \Src2Reg\() + + .endm + + .macro VpdpbsudYmmYmmYmm DestReg, Src1Reg, Src2Reg + + Avx2VnniYmmYmmYmm 0x50, 0x02, \DestReg\(), \Src1Reg\(), \Src2Reg\() + + .endm + + .macro VpdpbsudsYmmYmmYmm DestReg, Src1Reg, Src2Reg + + Avx2VnniYmmYmmYmm 0x51, 0x02, \DestReg\(), \Src1Reg\(), \Src2Reg\() + + .endm + + .macro VpdpbuudYmmYmmYmm DestReg, Src1Reg, Src2Reg + + Avx2VnniYmmYmmYmm 0x50, 0x00, \DestReg\(), \Src1Reg\(), \Src2Reg\() + + .endm + + .macro VpdpbuudsYmmYmmYmm DestReg, Src1Reg, Src2Reg + + Avx2VnniYmmYmmYmm 0x51, 0x00, \DestReg\(), \Src1Reg\(), \Src2Reg\() + + .endm + +/*++ + +Macro Description: + + This macro builds a VNNI instruction of the form: + + instr xmm1,xmm2,xmm3 + +Arguments: + + Opcode - Specifies the opcode for the VNNI instruction. + + Prefix - Specifies the opcode prefix for payload 1 + + DestReg - Specifies the destination register. + + Src1Reg - Specifies the first source register. + + Src2Reg - Specifies the second source register. + +--*/ + .macro Avx2VnniXmmXmmXmm Opcode, Prefix, DestReg, Src1Reg, Src2Reg + + .set Payload0, 0x02 # "0F 38" prefix + .set Payload0, Payload0 + ((((.LYmmIndex_\DestReg\() >> 3) & 1) ^ 1) << 7) + .set Payload0, Payload0 + (1 << 6) + .set Payload0, Payload0 + ((((.LYmmIndex_\Src2Reg\() >> 3) & 1) ^ 1) << 5) + + .set Payload1, 0x00 + \Prefix\() # 128-bit length and opcode prefix + .set Payload1, Payload1 + (((.LYmmIndex_\Src1Reg\() & 15) ^ 15) << 3) + + .set ModRMByte, 0xC0 # register form + .set ModRMByte, ModRMByte + ((.LYmmIndex_\DestReg\() & 7) << 3) + .set ModRMByte, ModRMByte + (.LYmmIndex_\Src2Reg\() & 7) + + .byte 0xC4, Payload0, Payload1, \Opcode\(), ModRMByte + + .endm + + .macro VpdpbssdXmmXmmXmm DestReg, Src1Reg, Src2Reg + + Avx2VnniXmmXmmXmm 0x50, 0x03, \DestReg\(), \Src1Reg\(), \Src2Reg\() + + .endm + + .macro VpdpbssdsXmmXmmXmm DestReg, Src1Reg, Src2Reg + + Avx2VnniXmmXmmXmm 0x51, 0x03, \DestReg\(), \Src1Reg\(), \Src2Reg\() + + .endm + + .macro VpdpbsudXmmXmmXmm DestReg, Src1Reg, Src2Reg + + Avx2VnniXmmXmmXmm 0x50, 0x02, \DestReg\(), \Src1Reg\(), \Src2Reg\() + + .endm + + .macro VpdpbsudsXmmXmmXmm DestReg, Src1Reg, Src2Reg + + Avx2VnniXmmXmmXmm 0x51, 0x02, \DestReg\(), \Src1Reg\(), \Src2Reg\() + + .endm + + .macro VpdpbuudXmmXmmXmm DestReg, Src1Reg, Src2Reg + + Avx2VnniXmmXmmXmm 0x50, 0x00, \DestReg\(), \Src1Reg\(), \Src2Reg\() + + .endm + + .macro VpdpbuudsXmmXmmXmm DestReg, Src1Reg, Src2Reg + + Avx2VnniXmmXmmXmm 0x51, 0x00, \DestReg\(), \Src1Reg\(), \Src2Reg\() + + .endm diff --git a/onnxruntime/core/mlas/lib/x86_64/QgemmU8S8KernelAvx2.S b/onnxruntime/core/mlas/lib/x86_64/QgemmU8S8KernelAvx2.S index 5068d41cea..3066eadaf1 100644 --- a/onnxruntime/core/mlas/lib/x86_64/QgemmU8S8KernelAvx2.S +++ b/onnxruntime/core/mlas/lib/x86_64/QgemmU8S8KernelAvx2.S @@ -14,35 +14,37 @@ Abstract: multiply operation (QGEMM). This implementation uses AVX2 instructions. + Support for AVX-VNNI-INT8 for certain code paths. --*/ #include "asmmacro.h" +#include "AssembleAvxVnni.h" .intel_syntax noprefix // -// Stack frame layout for the U8S8 CopyPackA routine. +// Stack frame layout for the Int8 CopyPackA routine. // - .equ .LGemmU8S8CopyPackAFrame_PaddedMatrixAData, -72 - .equ .LGemmU8S8CopyPackAFrame_Padding, -8 - .equ .LGemmU8S8CopyPackAFrame_SavedR13, 0 - .equ .LGemmU8S8CopyPackAFrame_SavedR12, 8 - .equ .LGemmU8S8CopyPackAFrame_SavedRbx, 16 - .equ .LGemmU8S8CopyPackAFrame_SavedRbp, 24 - .equ .LGemmU8S8CopyPackAFrame_ReturnAddress, 32 + .equ .LGemmInt8CopyPackAFrame_PaddedMatrixAData, -72 + .equ .LGemmInt8CopyPackAFrame_Padding, -8 + .equ .LGemmInt8CopyPackAFrame_SavedR13, 0 + .equ .LGemmInt8CopyPackAFrame_SavedR12, 8 + .equ .LGemmInt8CopyPackAFrame_SavedRbx, 16 + .equ .LGemmInt8CopyPackAFrame_SavedRbp, 24 + .equ .LGemmInt8CopyPackAFrame_ReturnAddress, 32 // -// Stack frame layout for the U8S8 CopyPackB routine. +// Stack frame layout for the Int8 CopyPackB routine. // - .equ .LGemmU8S8CopyPackBFrame_PaddedMatrixBData, -72 - .equ .LGemmU8S8CopyPackBFrame_Padding, -8 - .equ .LGemmU8S8CopyPackBFrame_SavedRbx, 0 - .equ .LGemmU8S8CopyPackBFrame_SavedRbp, 8 - .equ .LGemmU8S8CopyPackBFrame_ReturnAddress, 16 - .equ .LGemmU8S8CopyPackBFrame_BIsSigned, 24 + .equ .LGemmInt8CopyPackBFrame_PaddedMatrixBData, -72 + .equ .LGemmInt8CopyPackBFrame_Padding, -8 + .equ .LGemmInt8CopyPackBFrame_SavedRbx, 0 + .equ .LGemmInt8CopyPackBFrame_SavedRbp, 8 + .equ .LGemmInt8CopyPackBFrame_ReturnAddress, 16 + .equ .LGemmInt8CopyPackBFrame_BIsSigned, 24 .text @@ -75,7 +77,7 @@ Return Value: --*/ - FUNCTION_ENTRY MlasGemmU8S8CopyPackAAvx2 +.macro MlasGemmCopyPackAAvx2 ASigned push rbp push rbx @@ -107,18 +109,18 @@ Return Value: // Zero initialize the padded stack buffers. // - vpxor xmm0,xmm0,xmm0 - vmovdqu YMMWORD PTR .LGemmU8S8CopyPackAFrame_PaddedMatrixAData[rsp],ymm0 - vmovdqu YMMWORD PTR .LGemmU8S8CopyPackAFrame_PaddedMatrixAData[rsp+32],ymm0 + vpxor xmm0,xmm0,xmm0 + vmovdqu YMMWORD PTR .LGemmInt8CopyPackAFrame_PaddedMatrixAData[rsp],ymm0 + vmovdqu YMMWORD PTR .LGemmInt8CopyPackAFrame_PaddedMatrixAData[rsp+32],ymm0 // // Process 4 rows of matrix A in a loop. // sub r11,4 - jb .LCopyPackA.ProcessRemainingRows + jb .LCopyPackA.ProcessRemainingRows\@ -.LCopyPackA.ProcessNextRowM4: +.LCopyPackA.ProcessNextRowM4\@: vpxor xmm0,xmm0,xmm0 # clear row accumulators vpxor xmm1,xmm1,xmm1 vpxor xmm2,xmm2,xmm2 @@ -131,9 +133,9 @@ Return Value: lea rdi,[rdi+r12*4] # advance next matrix D by 4 rows mov rbx,r8 # reload columns remaining sub rbx,32 - jb .LCopyPackA.ProcessRemainingColumnsM4 + jb .LCopyPackA.ProcessRemainingColumnsM4\@ -.LCopyPackA.ProcessNextColumnLoopM4: +.LCopyPackA.ProcessNextColumnLoopM4\@: vmovdqu ymm4,YMMWORD PTR [rdx] vmovdqu ymm5,YMMWORD PTR [rdx+r10] vmovdqu ymm6,YMMWORD PTR [rdx+r10*2] @@ -142,6 +144,12 @@ Return Value: vmovdqu YMMWORD PTR [rcx+r12],ymm5 vmovdqu YMMWORD PTR [rcx+r12*2],ymm6 vmovdqu YMMWORD PTR [rcx+rax],ymm7 +.if \ASigned\() == 1 + VpdpbssdYmmYmmYmm ymm0,ymm4,ymm9 + VpdpbssdYmmYmmYmm ymm1,ymm5,ymm9 + VpdpbssdYmmYmmYmm ymm2,ymm6,ymm9 + VpdpbssdYmmYmmYmm ymm3,ymm7,ymm9 +.else vpmaddubsw ymm4,ymm4,ymm9 # horizontal byte+byte=word per row vpaddw ymm0,ymm0,ymm4 # add words to row accumulators vpmaddubsw ymm5,ymm5,ymm9 @@ -150,16 +158,17 @@ Return Value: vpaddw ymm2,ymm2,ymm6 vpmaddubsw ymm7,ymm7,ymm9 vpaddw ymm3,ymm3,ymm7 +.endif add rdx,32 # advance matrix A by 32 bytes add rcx,32 # advance matrix D by 32 bytes sub rbx,32 # subtract columns remaining - jae .LCopyPackA.ProcessNextColumnLoopM4 + jae .LCopyPackA.ProcessNextColumnLoopM4\@ -.LCopyPackA.ProcessRemainingColumnsM4: +.LCopyPackA.ProcessRemainingColumnsM4\@: add rbx,32 # correct for over-subtract above - jz .LCopyPackA.ReduceRowSumBufferM4 + jz .LCopyPackA.ReduceRowSumBufferM4\@ test bl,16 # (CountK & 16) != 0? - jz .LCopyPackA.CopyRemainingCountKLessThan16M4 + jz .LCopyPackA.CopyRemainingCountKLessThan16M4\@ vmovdqu xmm4,XMMWORD PTR [rdx] vmovdqu xmm5,XMMWORD PTR [rdx+r10] vmovdqu xmm6,XMMWORD PTR [rdx+r10*2] @@ -168,6 +177,12 @@ Return Value: vmovdqu XMMWORD PTR [rcx+r12],xmm5 vmovdqu XMMWORD PTR [rcx+r12*2],xmm6 vmovdqu XMMWORD PTR [rcx+rax],xmm7 +.if \ASigned\() == 1 + VpdpbssdYmmYmmYmm ymm0,ymm4,ymm9 + VpdpbssdYmmYmmYmm ymm1,ymm5,ymm9 + VpdpbssdYmmYmmYmm ymm2,ymm6,ymm9 + VpdpbssdYmmYmmYmm ymm3,ymm7,ymm9 +.else vpmaddubsw xmm4,xmm4,xmm9 # horizontal byte+byte=word per row vpaddw ymm0,ymm0,ymm4 # add words to row accumulators vpmaddubsw xmm5,xmm5,xmm9 @@ -176,19 +191,20 @@ Return Value: vpaddw ymm2,ymm2,ymm6 vpmaddubsw xmm7,xmm7,xmm9 vpaddw ymm3,ymm3,ymm7 +.endif add rdx,16 # advance matrix A by 16 bytes add rcx,16 # advance matrix D by 16 bytes test bl,15 # test for unaligned columns - jz .LCopyPackA.ReduceRowSumBufferM4 + jz .LCopyPackA.ReduceRowSumBufferM4\@ // // Copy the unaligned CountK columns to a zero padded stack buffer. // -.LCopyPackA.CopyRemainingCountKLessThan16M4: - lea rbp,.LGemmU8S8CopyPackAFrame_PaddedMatrixAData[rsp] +.LCopyPackA.CopyRemainingCountKLessThan16M4\@: + lea rbp,.LGemmInt8CopyPackAFrame_PaddedMatrixAData[rsp] test bl,8 # (CountK & 8) != 0? - jz .LCopyPackA.CopyRemainingCountKLessThan8M4 + jz .LCopyPackA.CopyRemainingCountKLessThan8M4\@ mov rax,QWORD PTR [rdx] mov QWORD PTR [rbp],rax mov rax,QWORD PTR [rdx+r10] @@ -200,9 +216,9 @@ Return Value: add rdx,8 add rbp,8 # advance padded buffer destination -.LCopyPackA.CopyRemainingCountKLessThan8M4: +.LCopyPackA.CopyRemainingCountKLessThan8M4\@: test bl,4 # (CountK & 4) != 0? - jz .LCopyPackA.CopyRemainingCountKLessThan4M4 + jz .LCopyPackA.CopyRemainingCountKLessThan4M4\@ mov eax,DWORD PTR [rdx] mov DWORD PTR [rbp],eax mov eax,DWORD PTR [rdx+r10] @@ -214,9 +230,9 @@ Return Value: add rdx,4 add rbp,4 # advance padded buffer destination -.LCopyPackA.CopyRemainingCountKLessThan4M4: +.LCopyPackA.CopyRemainingCountKLessThan4M4\@: test bl,2 # (CountK & 2) != 0? - jz .LCopyPackA.CopyRemainingCountKLessThan2M4 + jz .LCopyPackA.CopyRemainingCountKLessThan2M4\@ movzx eax,WORD PTR [rdx] mov WORD PTR [rbp],ax movzx eax,WORD PTR [rdx+r10] @@ -228,9 +244,9 @@ Return Value: add rdx,2 add rbp,2 # advance padded buffer destination -.LCopyPackA.CopyRemainingCountKLessThan2M4: +.LCopyPackA.CopyRemainingCountKLessThan2M4\@: test bl,1 # (CountK & 1) != 0? - jz .LCopyPackA.ProcessPaddedMatrixADataM4 + jz .LCopyPackA.ProcessPaddedMatrixADataM4\@ movzx eax,BYTE PTR [rdx] mov BYTE PTR [rbp],al movzx eax,BYTE PTR [rdx+r10] @@ -244,16 +260,22 @@ Return Value: // Process the remaining CountK columns using the zero padded stack buffer. // -.LCopyPackA.ProcessPaddedMatrixADataM4: - vmovdqu xmm4,XMMWORD PTR .LGemmU8S8CopyPackAFrame_PaddedMatrixAData[rsp] - vmovdqu xmm5,XMMWORD PTR .LGemmU8S8CopyPackAFrame_PaddedMatrixAData[rsp+16] - vmovdqu xmm6,XMMWORD PTR .LGemmU8S8CopyPackAFrame_PaddedMatrixAData[rsp+32] - vmovdqu xmm7,XMMWORD PTR .LGemmU8S8CopyPackAFrame_PaddedMatrixAData[rsp+48] +.LCopyPackA.ProcessPaddedMatrixADataM4\@: + vmovdqu xmm4,XMMWORD PTR .LGemmInt8CopyPackAFrame_PaddedMatrixAData[rsp] + vmovdqu xmm5,XMMWORD PTR .LGemmInt8CopyPackAFrame_PaddedMatrixAData[rsp+16] + vmovdqu xmm6,XMMWORD PTR .LGemmInt8CopyPackAFrame_PaddedMatrixAData[rsp+32] + vmovdqu xmm7,XMMWORD PTR .LGemmInt8CopyPackAFrame_PaddedMatrixAData[rsp+48] lea rax,[rcx+r12*2] # compute matrix D plus 2 rows vpmaskmovd XMMWORD PTR [rcx],xmm10,xmm4 vpmaskmovd XMMWORD PTR [rcx+r12],xmm10,xmm5 vpmaskmovd XMMWORD PTR [rax],xmm10,xmm6 vpmaskmovd XMMWORD PTR [rax+r12],xmm10,xmm7 +.if \ASigned\() == 1 + VpdpbssdYmmYmmYmm ymm0,ymm4,ymm9 + VpdpbssdYmmYmmYmm ymm1,ymm5,ymm9 + VpdpbssdYmmYmmYmm ymm2,ymm6,ymm9 + VpdpbssdYmmYmmYmm ymm3,ymm7,ymm9 +.else vpmaddubsw xmm4,xmm4,xmm9 # horizontal byte+byte=word per row vpaddw ymm0,ymm0,ymm4 # add words to row accumulators vpmaddubsw xmm5,xmm5,xmm9 @@ -262,17 +284,22 @@ Return Value: vpaddw ymm2,ymm2,ymm6 vpmaddubsw xmm7,xmm7,xmm9 vpaddw ymm3,ymm3,ymm7 +.endif // // Reduce the sums for the four rows of output. // -.LCopyPackA.ReduceRowSumBufferM4: +.LCopyPackA.ReduceRowSumBufferM4\@: +.if \ASigned\() == 1 + vphaddd ymm0,ymm0,ymm1 +.else vpmaddwd ymm0,ymm0,ymm8 # horizontal word+word=dword per row vpmaddwd ymm1,ymm1,ymm8 vphaddd ymm0,ymm0,ymm1 # reduce and interleave Sum1/Sum0 vpmaddwd ymm2,ymm2,ymm8 vpmaddwd ymm3,ymm3,ymm8 +.endif vphaddd ymm1,ymm2,ymm3 # reduce and interleave Sum3/Sum2 vphaddd ymm0,ymm0,ymm1 # reduce and interleave Sum3/Sum2/Sum1/Sum0 vextracti128 xmm1,ymm0,1 # extract high dwords @@ -280,17 +307,17 @@ Return Value: vmovdqu XMMWORD PTR [r9],xmm0 add r9,4*4 # advance row sum buffer by 4 dwords sub r11,4 # subtract rows remaining - jae .LCopyPackA.ProcessNextRowM4 + jae .LCopyPackA.ProcessNextRowM4\@ -.LCopyPackA.ProcessRemainingRows: +.LCopyPackA.ProcessRemainingRows\@: add r11,4 # correct for over-subtract above - jz .LCopyPackA.ExitRoutine + jz .LCopyPackA.ExitRoutine\@ // // Process a single row of matrix A in a loop. // -.LCopyPackA.ProcessNextRowM1: +.LCopyPackA.ProcessNextRowM1\@: vpxor xmm0,xmm0,xmm0 # clear row accumulator mov rdx,rsi mov rcx,rdi @@ -298,64 +325,72 @@ Return Value: add rdi,r12 mov rbx,r8 # reload columns remaining sub rbx,32 - jb .LCopyPackA.ProcessRemainingColumnsM1 + jb .LCopyPackA.ProcessRemainingColumnsM1\@ -.LCopyPackA.ProcessNextColumnLoopM1: +.LCopyPackA.ProcessNextColumnLoopM1\@: vmovdqu ymm4,YMMWORD PTR [rdx] vmovdqu YMMWORD PTR [rcx],ymm4 +.if \ASigned\() == 1 + VpdpbssdYmmYmmYmm ymm0,ymm4,ymm9 +.else vpmaddubsw ymm4,ymm4,ymm9 # horizontal byte+byte=word per row vpaddw ymm0,ymm0,ymm4 # add words to row accumulators +.endif add rdx,32 # advance matrix A by 32 bytes add rcx,32 # advance matrix D by 32 bytes sub rbx,32 # subtract columns remaining - jae .LCopyPackA.ProcessNextColumnLoopM1 + jae .LCopyPackA.ProcessNextColumnLoopM1\@ -.LCopyPackA.ProcessRemainingColumnsM1: +.LCopyPackA.ProcessRemainingColumnsM1\@: add rbx,32 # correct for over-subtract above - jz .LCopyPackA.ReduceRowSumBufferM1 + jz .LCopyPackA.ReduceRowSumBufferM1\@ test bl,16 # (CountK & 16) != 0? - jz .LCopyPackA.CopyRemainingCountKLessThan16M1 + jz .LCopyPackA.CopyRemainingCountKLessThan16M1\@ vmovdqu xmm4,XMMWORD PTR [rdx] vmovdqu XMMWORD PTR [rcx],xmm4 +.if \ASigned\() == 1 + VpdpbssdYmmYmmYmm ymm0,ymm4,ymm9 +.else vpmaddubsw xmm4,xmm4,xmm9 # horizontal byte+byte=word per row vpaddw ymm0,ymm0,ymm4 # add words to row accumulators +.endif add rdx,16 # advance matrix A by 16 bytes add rcx,16 # advance matrix D by 16 bytes test bl,15 # test for unaligned columns - jz .LCopyPackA.ReduceRowSumBufferM1 + jz .LCopyPackA.ReduceRowSumBufferM1\@ // // Copy the unaligned CountK columns to a zero padded stack buffer. // -.LCopyPackA.CopyRemainingCountKLessThan16M1: - lea rbp,.LGemmU8S8CopyPackAFrame_PaddedMatrixAData[rsp] +.LCopyPackA.CopyRemainingCountKLessThan16M1\@: + lea rbp,.LGemmInt8CopyPackAFrame_PaddedMatrixAData[rsp] test bl,8 # (CountK & 8) != 0? - jz .LCopyPackA.CopyRemainingCountKLessThan8M1 + jz .LCopyPackA.CopyRemainingCountKLessThan8M1\@ mov rax,QWORD PTR [rdx] mov QWORD PTR [rbp],rax add rdx,8 add rbp,8 # advance padded buffer destination -.LCopyPackA.CopyRemainingCountKLessThan8M1: +.LCopyPackA.CopyRemainingCountKLessThan8M1\@: test bl,4 # (CountK & 4) != 0? - jz .LCopyPackA.CopyRemainingCountKLessThan4M1 + jz .LCopyPackA.CopyRemainingCountKLessThan4M1\@ mov eax,DWORD PTR [rdx] mov DWORD PTR [rbp],eax add rdx,4 add rbp,4 # advance padded buffer destination -.LCopyPackA.CopyRemainingCountKLessThan4M1: +.LCopyPackA.CopyRemainingCountKLessThan4M1\@: test bl,2 # (CountK & 2) != 0? - jz .LCopyPackA.CopyRemainingCountKLessThan2M1 + jz .LCopyPackA.CopyRemainingCountKLessThan2M1\@ movzx eax,WORD PTR [rdx] mov WORD PTR [rbp],ax add rdx,2 add rbp,2 # advance padded buffer destination -.LCopyPackA.CopyRemainingCountKLessThan2M1: +.LCopyPackA.CopyRemainingCountKLessThan2M1\@: test bl,1 # (CountK & 1) != 0? - jz .LCopyPackA.ProcessPaddedMatrixADataM1 + jz .LCopyPackA.ProcessPaddedMatrixADataM1\@ movzx eax,BYTE PTR [rdx] mov BYTE PTR [rbp],al @@ -363,18 +398,24 @@ Return Value: // Process the remaining CountK columns using the zero padded stack buffer. // -.LCopyPackA.ProcessPaddedMatrixADataM1: - vmovdqu xmm4,XMMWORD PTR .LGemmU8S8CopyPackAFrame_PaddedMatrixAData[rsp] +.LCopyPackA.ProcessPaddedMatrixADataM1\@: + vmovdqu xmm4,XMMWORD PTR .LGemmInt8CopyPackAFrame_PaddedMatrixAData[rsp] vpmaskmovd XMMWORD PTR [rcx],xmm10,xmm4 +.if \ASigned\() == 1 + VpdpbssdYmmYmmYmm ymm0,ymm4,ymm9 +.else vpmaddubsw ymm4,ymm4,ymm9 # horizontal byte+byte=word per row vpaddw ymm0,ymm0,ymm4 # accumulate per row along columns +.endif // // Reduce the sum for the single row of output. // -.LCopyPackA.ReduceRowSumBufferM1: +.LCopyPackA.ReduceRowSumBufferM1\@: +.if \ASigned\() == 0 vpmaddwd ymm0,ymm0,ymm8 # horizontal word+word=dword per row +.endif vextracti128 xmm1,ymm0,1 # extract high dwords vpaddd xmm0,xmm0,xmm1 # reduction vphaddd xmm0,xmm0,xmm0 @@ -382,13 +423,13 @@ Return Value: vmovd DWORD PTR [r9],xmm0 add r9,4 # advance row sum buffer by 1 dword dec r11 # decrement rows remaining - jnz .LCopyPackA.ProcessNextRowM1 + jnz .LCopyPackA.ProcessNextRowM1\@ // // Restore non-volatile registers and return. // -.LCopyPackA.ExitRoutine: +.LCopyPackA.ExitRoutine\@: vzeroupper pop r13 @@ -396,6 +437,13 @@ Return Value: pop rbx pop rbp ret +.endm + + FUNCTION_ENTRY MlasGemmU8S8CopyPackAAvx2 + MlasGemmCopyPackAAvx2 0 + + FUNCTION_ENTRY MlasGemmS8CopyPackAAvx2Vnni + MlasGemmCopyPackAAvx2 1 /*++ @@ -428,7 +476,7 @@ Return Value: --*/ - FUNCTION_ENTRY MlasGemmU8S8CopyPackBAvx2 +.macro MlasGemmCopyPackBAvx2 IsVnni, BSigned push rbp push rbx @@ -445,36 +493,37 @@ Return Value: // vpxor xmm9,xmm9,xmm9 # generate word vector [0x0000] - cmp BYTE PTR .LGemmU8S8CopyPackBFrame_BIsSigned[rsp],0 - jnz .LCopyPackB.SkipUnsignedBitFlipVector +.if \IsVnni\() == 0 + cmp BYTE PTR .LGemmInt8CopyPackBFrame_BIsSigned[rsp],0 + jnz .LCopyPackB.SkipUnsignedBitFlipVector\@ vpsllw ymm9,ymm8,7 # generate word vector [0x8080] - -.LCopyPackB.SkipUnsignedBitFlipVector: +.endif +.LCopyPackB.SkipUnsignedBitFlipVector\@: // // Process 16 columns of matrix B in a loop. // sub rcx,16 - jb .LCopyPackB.ProcessRemainingColumns + jb .LCopyPackB.ProcessRemainingColumns\@ -.LCopyPackB.ProcessNextColumnN16: +.LCopyPackB.ProcessNextColumnN16\@: vpxor xmm0,xmm0,xmm0 # clear column accumulators vpxor xmm1,xmm1,xmm1 mov rdx,rsi add rsi,16 # advance next matrix B by 16 columns mov rbx,r8 # reload rows remaining sub rbx,4 - jb .LCopyPackB.ProcessRemainingRowsN16 + jb .LCopyPackB.ProcessRemainingRowsN16\@ -.LCopyPackB.ProcessNextRowLoopN16: +.LCopyPackB.ProcessNextRowLoopN16\@: vmovdqu xmm2,XMMWORD PTR [rdx] # load 4 rows vmovdqu xmm3,XMMWORD PTR [rdx+r10] vmovdqu xmm4,XMMWORD PTR [rdx+r10*2] vmovdqu xmm5,XMMWORD PTR [rdx+r11] lea rdx,[rdx+r10*4] # advance matrix B by 4 rows -.LCopyPackB.InterleaveRowDataN16: +.LCopyPackB.InterleaveRowDataN16\@: vpunpcklbw xmm6,xmm2,xmm3 # interleave row data vpunpckhbw xmm3,xmm2,xmm3 vpunpcklbw xmm2,xmm4,xmm5 @@ -485,56 +534,68 @@ Return Value: vpunpckhwd xmm3,xmm3,xmm5 vinserti128 ymm4,ymm4,xmm6,1 vinserti128 ymm2,ymm2,xmm3,1 +.if \IsVnni\() == 0 vpxor ymm4,ymm4,ymm9 # optionally adjust unsigned data vpxor ymm2,ymm2,ymm9 +.endif vmovdqu YMMWORD PTR [rdi],ymm4 # store interleaved rows vmovdqu YMMWORD PTR [rdi+32],ymm2 +.if \IsVnni\() == 1 + .if \BSigned\() == 1 + VpdpbssdYmmYmmYmm ymm0,ymm4,ymm8 + VpdpbssdYmmYmmYmm ymm1,ymm2,ymm8 + .else + VpdpbuudYmmYmmYmm ymm0,ymm4,ymm8 + VpdpbuudYmmYmmYmm ymm1,ymm2,ymm8 + .endif +.else vpmaddubsw ymm4,ymm8,ymm4 # horizontal byte+byte=word per row vpmaddwd ymm4,ymm4,ymm7 # horizontal word+word=dword per row vpaddd ymm0,ymm0,ymm4 # accumulate per column vpmaddubsw ymm2,ymm8,ymm2 vpmaddwd ymm2,ymm2,ymm7 vpaddd ymm1,ymm1,ymm2 +.endif add rdi,64 # advance matrix D by 64 bytes sub rbx,4 # subtract rows remaining - jae .LCopyPackB.ProcessNextRowLoopN16 + jae .LCopyPackB.ProcessNextRowLoopN16\@ // // Process the less than 4 remaining rows where the row has 16 columns. // -.LCopyPackB.ProcessRemainingRowsN16: +.LCopyPackB.ProcessRemainingRowsN16\@: add rbx,4 # correct for over-subtract above - jz .LCopyPackB.StoreColumnSumBufferN16 + jz .LCopyPackB.StoreColumnSumBufferN16\@ vmovdqu xmm2,XMMWORD PTR [rdx] 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 + jz .LCopyPackB.InterleaveRowDataN16\@ vmovdqu xmm3,XMMWORD PTR [rdx+r10] test r8b,1 # (CountK & 1) != 0? - jz .LCopyPackB.InterleaveRowDataN16 + jz .LCopyPackB.InterleaveRowDataN16\@ vmovdqu xmm4,XMMWORD PTR [rdx+r10*2] - jmp .LCopyPackB.InterleaveRowDataN16 + jmp .LCopyPackB.InterleaveRowDataN16\@ -.LCopyPackB.StoreColumnSumBufferN16: +.LCopyPackB.StoreColumnSumBufferN16\@: vmovdqu YMMWORD PTR [r9],ymm0 vmovdqu YMMWORD PTR [r9+32],ymm1 add r9,16*4 # advance column sum buffer by 16 dwords sub rcx,16 # subtract columns remaining - jae .LCopyPackB.ProcessNextColumnN16 + jae .LCopyPackB.ProcessNextColumnN16\@ -.LCopyPackB.ProcessRemainingColumns: +.LCopyPackB.ProcessRemainingColumns\@: add rcx,16 # correct for over-subtract above - jnz .LCopyPackB.ProcessColumnNUnaligned + jnz .LCopyPackB.ProcessColumnNUnaligned\@ // // Restore non-volatile registers and return. // -.LCopyPackB.ExitRoutine: +.LCopyPackB.ExitRoutine\@: vzeroupper pop rbx @@ -545,19 +606,19 @@ Return Value: // Process the remaining columns of matrix B. // -.LCopyPackB.ProcessColumnNUnaligned: +.LCopyPackB.ProcessColumnNUnaligned\@: vpxor xmm0,xmm0,xmm0 # clear column accumulators vpxor xmm1,xmm1,xmm1 - vmovdqu YMMWORD PTR .LGemmU8S8CopyPackBFrame_PaddedMatrixBData[rsp],ymm9 - vmovdqu YMMWORD PTR .LGemmU8S8CopyPackBFrame_PaddedMatrixBData[rsp+32],ymm9 + vmovdqu YMMWORD PTR .LGemmInt8CopyPackBFrame_PaddedMatrixBData[rsp],ymm9 + vmovdqu YMMWORD PTR .LGemmInt8CopyPackBFrame_PaddedMatrixBData[rsp+32],ymm9 sub r8,4 - jb .LCopyPackB.ProcessRemainingRowsNUnaligned + jb .LCopyPackB.ProcessRemainingRowsNUnaligned\@ -.LCopyPackB.ProcessNextRowLoopNUnaligned: +.LCopyPackB.ProcessNextRowLoopNUnaligned\@: mov rdx,rsi - lea rbp,.LGemmU8S8CopyPackBFrame_PaddedMatrixBData[rsp] + lea rbp,.LGemmInt8CopyPackBFrame_PaddedMatrixBData[rsp] test cl,8 # (CountN & 8) != 0? - jz .LCopyPackB.CopyRemainingCountNLessThan8K4 + jz .LCopyPackB.CopyRemainingCountNLessThan8K4\@ mov rax,QWORD PTR [rdx] mov QWORD PTR [rbp],rax mov rax,QWORD PTR [rdx+r10] @@ -569,9 +630,9 @@ Return Value: add rdx,8 # advance matrix B add rbp,8 # advance padded buffer destination -.LCopyPackB.CopyRemainingCountNLessThan8K4: +.LCopyPackB.CopyRemainingCountNLessThan8K4\@: test cl,4 # (CountN & 4) != 0? - jz .LCopyPackB.CopyRemainingCountNLessThan4K4 + jz .LCopyPackB.CopyRemainingCountNLessThan4K4\@ mov eax,DWORD PTR [rdx] mov DWORD PTR [rbp],eax mov eax,DWORD PTR [rdx+r10] @@ -583,9 +644,9 @@ Return Value: add rdx,4 # advance matrix B add rbp,4 # advance padded buffer destination -.LCopyPackB.CopyRemainingCountNLessThan4K4: +.LCopyPackB.CopyRemainingCountNLessThan4K4\@: test cl,2 # (CountN & 2) != 0? - jz .LCopyPackB.CopyRemainingCountNLessThan2K4 + jz .LCopyPackB.CopyRemainingCountNLessThan2K4\@ movzx eax,WORD PTR [rdx] mov WORD PTR [rbp],ax movzx eax,WORD PTR [rdx+r10] @@ -597,9 +658,9 @@ Return Value: add rdx,2 # advance matrix B add rbp,2 # advance padded buffer destination -.LCopyPackB.CopyRemainingCountNLessThan2K4: +.LCopyPackB.CopyRemainingCountNLessThan2K4\@: test cl,1 # (CountN & 1) != 0? - jz .LCopyPackB.ProcessPaddedMatrixBData + jz .LCopyPackB.ProcessPaddedMatrixBData\@ movzx eax,BYTE PTR [rdx] mov BYTE PTR [rbp],al movzx eax,BYTE PTR [rdx+r10] @@ -609,11 +670,11 @@ Return Value: movzx eax,BYTE PTR [rdx+r11] mov BYTE PTR [rbp+48],al -.LCopyPackB.ProcessPaddedMatrixBData: - vmovdqu xmm2,XMMWORD PTR .LGemmU8S8CopyPackBFrame_PaddedMatrixBData[rsp] - vmovdqu xmm3,XMMWORD PTR .LGemmU8S8CopyPackBFrame_PaddedMatrixBData[rsp+16] - vmovdqu xmm4,XMMWORD PTR .LGemmU8S8CopyPackBFrame_PaddedMatrixBData[rsp+32] - vmovdqu xmm5,XMMWORD PTR .LGemmU8S8CopyPackBFrame_PaddedMatrixBData[rsp+48] +.LCopyPackB.ProcessPaddedMatrixBData\@: + vmovdqu xmm2,XMMWORD PTR .LGemmInt8CopyPackBFrame_PaddedMatrixBData[rsp] + vmovdqu xmm3,XMMWORD PTR .LGemmInt8CopyPackBFrame_PaddedMatrixBData[rsp+16] + vmovdqu xmm4,XMMWORD PTR .LGemmInt8CopyPackBFrame_PaddedMatrixBData[rsp+32] + vmovdqu xmm5,XMMWORD PTR .LGemmInt8CopyPackBFrame_PaddedMatrixBData[rsp+48] vpunpcklbw xmm6,xmm2,xmm3 # interleave row data vpunpckhbw xmm3,xmm2,xmm3 vpunpcklbw xmm2,xmm4,xmm5 @@ -624,75 +685,98 @@ Return Value: vpunpckhwd xmm3,xmm3,xmm5 vinserti128 ymm4,ymm4,xmm6,1 vinserti128 ymm2,ymm2,xmm3,1 +.if \IsVnni\() == 0 vpxor ymm4,ymm4,ymm9 # optionally adjust unsigned data vpxor ymm2,ymm2,ymm9 +.endif vmovdqu YMMWORD PTR [rdi],ymm4 # store interleaved rows vmovdqu YMMWORD PTR [rdi+32],ymm2 +.if \IsVnni\() == 1 + .if \BSigned\() == 1 + VpdpbssdYmmYmmYmm ymm0,ymm4,ymm8 + VpdpbssdYmmYmmYmm ymm1,ymm2,ymm8 + .else + VpdpbuudYmmYmmYmm ymm0,ymm4,ymm8 + VpdpbuudYmmYmmYmm ymm1,ymm2,ymm8 + .endif +.else vpmaddubsw ymm4,ymm8,ymm4 # horizontal byte+byte=word per row vpmaddwd ymm4,ymm4,ymm7 # horizontal word+word=dword per row vpaddd ymm0,ymm0,ymm4 # accumulate per column vpmaddubsw ymm2,ymm8,ymm2 vpmaddwd ymm2,ymm2,ymm7 vpaddd ymm1,ymm1,ymm2 +.endif lea rsi,[rsi+r10*4] # advance next matrix B by 4 rows add rdi,64 # advance matrix D by 64 bytes sub r8,4 # subtract rows remaining - jae .LCopyPackB.ProcessNextRowLoopNUnaligned + jae .LCopyPackB.ProcessNextRowLoopNUnaligned\@ -.LCopyPackB.ProcessRemainingRowsNUnaligned: +.LCopyPackB.ProcessRemainingRowsNUnaligned\@: add r8,4 - jz .LCopyPackB.StoreColumnSumBufferNUnaligned + jz .LCopyPackB.StoreColumnSumBufferNUnaligned\@ // // Process the less than 4 remaining rows where the row has less than 16 columns. // - lea rbp,.LGemmU8S8CopyPackBFrame_PaddedMatrixBData[rsp] + lea rbp,.LGemmInt8CopyPackBFrame_PaddedMatrixBData[rsp] vmovdqu YMMWORD PTR [rbp],ymm9 vmovdqu YMMWORD PTR [rbp+32],ymm9 -.LCopyPackB.CopyUnalignedRowLoop: +.LCopyPackB.CopyUnalignedRowLoop\@: lea r11,[rbp+16] # advance next padded buffer by 16 bytes mov rdx,rsi test cl,8 # (CountN & 8) != 0? - jz .LCopyPackB.CopyRemainingCountNLessThan8KSmall + jz .LCopyPackB.CopyRemainingCountNLessThan8KSmall\@ mov rax,QWORD PTR [rdx] mov QWORD PTR [rbp],rax add rdx,8 # advance matrix B add rbp,8 # advance padded buffer destination -.LCopyPackB.CopyRemainingCountNLessThan8KSmall: +.LCopyPackB.CopyRemainingCountNLessThan8KSmall\@: test cl,4 # (CountN & 4) != 0? - jz .LCopyPackB.CopyRemainingCountNLessThan4KSmall + jz .LCopyPackB.CopyRemainingCountNLessThan4KSmall\@ mov eax,DWORD PTR [rdx] mov DWORD PTR [rbp],eax add rdx,4 # advance matrix B add rbp,4 # advance padded buffer destination -.LCopyPackB.CopyRemainingCountNLessThan4KSmall: +.LCopyPackB.CopyRemainingCountNLessThan4KSmall\@: test cl,2 # (CountN & 2) != 0? - jz .LCopyPackB.CopyRemainingCountNLessThan2KSmall + jz .LCopyPackB.CopyRemainingCountNLessThan2KSmall\@ movzx eax,WORD PTR [rdx] mov WORD PTR [rbp],ax add rdx,2 # advance matrix B add rbp,2 # advance padded buffer destination -.LCopyPackB.CopyRemainingCountNLessThan2KSmall: +.LCopyPackB.CopyRemainingCountNLessThan2KSmall\@: test cl,1 # (CountN & 1) != 0? - jz .LCopyPackB.DoneCopyRemainingCountNKSmall + jz .LCopyPackB.DoneCopyRemainingCountNKSmall\@ movzx eax,BYTE PTR [rdx] mov BYTE PTR [rbp],al -.LCopyPackB.DoneCopyRemainingCountNKSmall: +.LCopyPackB.DoneCopyRemainingCountNKSmall\@: dec r8 - jz .LCopyPackB.ProcessPaddedMatrixBData + jz .LCopyPackB.ProcessPaddedMatrixBData\@ add rsi,r10 # advance next matrix B by 1 row mov rbp,r11 - jmp .LCopyPackB.CopyUnalignedRowLoop + jmp .LCopyPackB.CopyUnalignedRowLoop\@ -.LCopyPackB.StoreColumnSumBufferNUnaligned: +.LCopyPackB.StoreColumnSumBufferNUnaligned\@: vmovdqu YMMWORD PTR [r9],ymm0 vmovdqu YMMWORD PTR [r9+32],ymm1 - jmp .LCopyPackB.ExitRoutine + jmp .LCopyPackB.ExitRoutine\@ + +.endm + + FUNCTION_ENTRY MlasGemmU8S8CopyPackBAvx2 + MlasGemmCopyPackBAvx2 0, 0 # sign variable not checked if IsVnni = 0 + + FUNCTION_ENTRY MlasGemmU8CopyPackBAvx2Vnni + MlasGemmCopyPackBAvx2 1, 0 + + FUNCTION_ENTRY MlasGemmS8CopyPackBAvx2Vnni + MlasGemmCopyPackBAvx2 1, 1 .end diff --git a/onnxruntime/core/mlas/lib/x86_64/QgemmU8X8KernelAvx2.S b/onnxruntime/core/mlas/lib/x86_64/QgemmU8X8KernelAvx2.S index b0f7be63c4..af2a475ea0 100644 --- a/onnxruntime/core/mlas/lib/x86_64/QgemmU8X8KernelAvx2.S +++ b/onnxruntime/core/mlas/lib/x86_64/QgemmU8X8KernelAvx2.S @@ -14,6 +14,7 @@ Abstract: multiply operation (QGEMM). This implementation uses AVX2 and AVX VNNI instructions. + AVX-VNNI-INT8 support also included. --*/ @@ -23,20 +24,20 @@ Abstract: .intel_syntax noprefix // -// Stack frame layout for the U8X8 kernel. +// Stack frame layout for the Int8 kernel. // - .equ .LGemmU8X8KernelFrame_type, -8 - .equ .LGemmU8X8KernelFrame_SavedR13, 0 - .equ .LGemmU8X8KernelFrame_SavedR12, 8 - .equ .LGemmU8X8KernelFrame_SavedRbx, 16 - .equ .LGemmU8X8KernelFrame_SavedRbp, 24 - .equ .LGemmU8X8KernelFrame_ReturnAddress, 32 - .equ .LGemmU8X8KernelFrame_ldc, 40 - .equ .LGemmU8X8KernelFrame_RowSumBuffer, 48 - .equ .LGemmU8X8KernelFrame_ColumnSumBuffer, 56 - .equ .LGemmU8X8KernelFrame_ZeroPointB, 64 - .equ .LGemmU8X8KernelFrame_ZeroMode, 72 + .equ .LGemmInt8KernelFrame_type, -8 + .equ .LGemmInt8KernelFrame_SavedR13, 0 + .equ .LGemmInt8KernelFrame_SavedR12, 8 + .equ .LGemmInt8KernelFrame_SavedRbx, 16 + .equ .LGemmInt8KernelFrame_SavedRbp, 24 + .equ .LGemmInt8KernelFrame_ReturnAddress, 32 + .equ .LGemmInt8KernelFrame_ldc, 40 + .equ .LGemmInt8KernelFrame_RowSumBuffer, 48 + .equ .LGemmInt8KernelFrame_ColumnSumBuffer, 56 + .equ .LGemmInt8KernelFrame_ZeroPointB, 64 + .equ .LGemmInt8KernelFrame_ZeroMode, 72 /*++ @@ -115,7 +116,7 @@ Implicit Arguments: --*/ - .macro ComputeBlockU8S8Avx2 ColumnCount, RowCount, VectorOffset, BroadcastOffset + .macro ComputeBlockAvx2 ColumnCount, RowCount, VectorOffset, BroadcastOffset, ASigned, BSigned .if \RowCount\() == 1 vpbroadcastd ymm2,DWORD PTR [rdi+\BroadcastOffset\()] @@ -170,13 +171,40 @@ Implicit Arguments: --*/ - .macro MultiplyAccumulateRowU8S8AvxVnni ColumnCount, Vec1Reg, Vec2Reg + .macro MultiplyAccumulateRowAvxVnni ColumnCount, Vec1Reg, Vec2Reg, ASigned, BSigned -.if \ColumnCount\() == 16 - VpdpbusdsYmmYmmYmm \Vec1Reg\(),ymm2,ymm0 - VpdpbusdsYmmYmmYmm \Vec2Reg\(),ymm2,ymm1 +.if \ASigned\() == 1 + .if \BSigned\() == 1 + .if \ColumnCount\() == 16 + VpdpbssdYmmYmmYmm \Vec1Reg\(),ymm2,ymm0 + VpdpbssdYmmYmmYmm \Vec2Reg\(),ymm2,ymm1 + .else + VpdpbssdYmmYmmYmm \Vec2Reg\(),ymm2,ymm0 + .endif + .else + .if \ColumnCount\() == 16 + VpdpbsudYmmYmmYmm \Vec1Reg\(),ymm2,ymm0 + VpdpbsudYmmYmmYmm \Vec2Reg\(),ymm2,ymm1 + .else + VpdpbsudYmmYmmYmm \Vec2Reg\(),ymm2,ymm0 + .endif + .endif .else - VpdpbusdsYmmYmmYmm \Vec2Reg\(),ymm2,ymm0 + .if \BSigned\() == 1 + .if \ColumnCount\() == 16 + VpdpbusdYmmYmmYmm \Vec1Reg\(),ymm2,ymm0 + VpdpbusdYmmYmmYmm \Vec2Reg\(),ymm2,ymm1 + .else + VpdpbusdYmmYmmYmm \Vec2Reg\(),ymm2,ymm0 + .endif + .else + .if \ColumnCount\() == 16 + VpdpbuudYmmYmmYmm \Vec1Reg\(),ymm2,ymm0 + VpdpbuudYmmYmmYmm \Vec2Reg\(),ymm2,ymm1 + .else + VpdpbuudYmmYmmYmm \Vec2Reg\(),ymm2,ymm0 + .endif + .endif .endif .endm @@ -212,22 +240,22 @@ Implicit Arguments: --*/ - .macro ComputeBlockU8S8AvxVnni ColumnCount, RowCount, VectorOffset, BroadcastOffset + .macro ComputeBlockAvxVnni ColumnCount, RowCount, VectorOffset, BroadcastOffset, ASigned, BSigned vmovdqu ymm0,YMMWORD PTR [rsi+\VectorOffset\()] EmitIfCountGE \ColumnCount\(), 16, "vmovdqu ymm1,YMMWORD PTR [rsi+\VectorOffset\()+32]" EmitIfCountGE \RowCount\(), 1, "vpbroadcastd ymm2,DWORD PTR [rdi+\BroadcastOffset\()]" - EmitIfCountGE \RowCount\(), 1, "MultiplyAccumulateRowU8S8AvxVnni \ColumnCount\(), ymm4, ymm5" + EmitIfCountGE \RowCount\(), 1, "MultiplyAccumulateRowAvxVnni \ColumnCount\(), ymm4, ymm5, \ASigned\(), \BSigned\()" EmitIfCountGE \RowCount\(), 2, "vpbroadcastd ymm2,DWORD PTR [rdi+rcx+\BroadcastOffset\()]" - EmitIfCountGE \RowCount\(), 2, "MultiplyAccumulateRowU8S8AvxVnni \ColumnCount\(), ymm6, ymm7" + EmitIfCountGE \RowCount\(), 2, "MultiplyAccumulateRowAvxVnni \ColumnCount\(), ymm6, ymm7, \ASigned\(), \BSigned\()" EmitIfCountGE \RowCount\(), 3, "vpbroadcastd ymm2,DWORD PTR [rdi+rcx*2+\BroadcastOffset\()]" - EmitIfCountGE \RowCount\(), 3, "MultiplyAccumulateRowU8S8AvxVnni \ColumnCount\(), ymm8, ymm9" + EmitIfCountGE \RowCount\(), 3, "MultiplyAccumulateRowAvxVnni \ColumnCount\(), ymm8, ymm9, \ASigned\(), \BSigned\()" EmitIfCountGE \RowCount\(), 4, "vpbroadcastd ymm2,DWORD PTR [r8+\BroadcastOffset\()]" - EmitIfCountGE \RowCount\(), 4, "MultiplyAccumulateRowU8S8AvxVnni \ColumnCount\(), ymm10, ymm11" + EmitIfCountGE \RowCount\(), 4, "MultiplyAccumulateRowAvxVnni \ColumnCount\(), ymm10, ymm11, \ASigned\(), \BSigned\()" EmitIfCountGE \RowCount\(), 5, "vpbroadcastd ymm2,DWORD PTR [r8+rcx+\BroadcastOffset\()]" - EmitIfCountGE \RowCount\(), 5, "MultiplyAccumulateRowU8S8AvxVnni \ColumnCount\(), ymm12, ymm13" + EmitIfCountGE \RowCount\(), 5, "MultiplyAccumulateRowAvxVnni \ColumnCount\(), ymm12, ymm13, \ASigned\(), \BSigned\()" EmitIfCountGE \RowCount\(), 6, "vpbroadcastd ymm2,DWORD PTR [r8+rcx*2+\BroadcastOffset\()]" - EmitIfCountGE \RowCount\(), 6, "MultiplyAccumulateRowU8S8AvxVnni \ColumnCount\(), ymm14, ymm15" + EmitIfCountGE \RowCount\(), 6, "MultiplyAccumulateRowAvxVnni \ColumnCount\(), ymm14, ymm15, \ASigned\(), \BSigned\()" .endm @@ -260,7 +288,7 @@ Implicit Arguments: --*/ - .macro ComputeBlockLoopU8S8 Isa, ColumnCount, RowCount + .macro ComputeBlockLoop Isa, ColumnCount, RowCount, ASigned, BSigned mov rbp,rcx # reload row length remaining @@ -269,10 +297,10 @@ Implicit Arguments: jb .LProcessRemainingBlocks\@ .LComputeBlockBy4Loop\@: - ComputeBlockU8S8\Isa\() \ColumnCount\(), \RowCount\(), 0*64, 0 - ComputeBlockU8S8\Isa\() \ColumnCount\(), \RowCount\(), 1*64, 4 - ComputeBlockU8S8\Isa\() \ColumnCount\(), \RowCount\(), 2*64, 8 - ComputeBlockU8S8\Isa\() \ColumnCount\(), \RowCount\(), 3*64, 12 + ComputeBlock\Isa\() \ColumnCount\(), \RowCount\(), 0*64, 0, \ASigned\(), \BSigned\() + ComputeBlock\Isa\() \ColumnCount\(), \RowCount\(), 1*64, 4, \ASigned\(), \BSigned\() + ComputeBlock\Isa\() \ColumnCount\(), \RowCount\(), 2*64, 8, \ASigned\(), \BSigned\() + ComputeBlock\Isa\() \ColumnCount\(), \RowCount\(), 3*64, 12, \ASigned\(), \BSigned\() add rdi,4*4 # advance matrix A by 4 quads add rsi,4*64 # advance matrix B sub rbp,4*4 @@ -284,7 +312,7 @@ Implicit Arguments: .endif .LComputeBlockBy1Loop\@: - ComputeBlockU8S8\Isa\() \ColumnCount\(), \RowCount\(), 0, 0 + ComputeBlock\Isa\() \ColumnCount\(), \RowCount\(), 0, 0, \ASigned\(), \BSigned\() add rdi,4 # advance matrix A by 1 quad .if \RowCount\() > 3 add r8,4 # advance matrix A plus 3 rows by 1 quad @@ -487,7 +515,7 @@ Implicit Arguments: --*/ - .macro ProduceOutputBlock ColumnCount, RowCount + .macro ProduceOutputBlock ColumnCount, RowCount, ASigned, BSigned // // Initialize the accumulators with the row and column sums. @@ -565,16 +593,16 @@ Implicit Arguments: lea r8,[rcx*2+rcx] add r8,rdi # compute matrix A plus 3 rows .endif - cmp DWORD PTR .LGemmU8X8KernelFrame_type[rsp],0 + cmp DWORD PTR .LGemmInt8KernelFrame_type[rsp],0 jg .LProduceWithU8U8Avx2\@ .if \RowCount\() <= 4 - jl .LProduceWithU8S8AvxVnni\@ - ComputeBlockLoopU8S8 Avx2, \ColumnCount\(), \RowCount\() + jl .LProduceWithInt8AvxVnni\@ + ComputeBlockLoop Avx2, \ColumnCount\(), \RowCount\(), \ASigned\(), \BSigned\() jmp .LExitProduceOutputBlock\@ .endif -.LProduceWithU8S8AvxVnni\@: - ComputeBlockLoopU8S8 AvxVnni, \ColumnCount\(), \RowCount\() +.LProduceWithInt8AvxVnni\@: + ComputeBlockLoop AvxVnni, \ColumnCount\(), \RowCount\(), \ASigned\(), \BSigned\() jmp .LExitProduceOutputBlock\@ .LProduceWithU8U8Avx2\@: @@ -624,13 +652,13 @@ Implicit Arguments: --*/ - .macro ProcessCountM RowCount + .macro ProcessCountM RowCount, ASigned, BSigned cmp r9,8 jbe .LProcessRemainingCountN\@ .LProcessNextColumnLoop16xN\@: - ProduceOutputBlock 16, \RowCount\() + ProduceOutputBlock 16, \RowCount\(), \ASigned\(), \BSigned\() sub r9,16 jb .LOutputMasked16xNBlock\@ test r10b,r10b # ZeroMode? @@ -673,7 +701,7 @@ Implicit Arguments: jmp .LExitKernel .LProcessRemainingCountN\@: - ProduceOutputBlock 8, \RowCount\() + ProduceOutputBlock 8, \RowCount\(), \ASigned\(), \BSigned\() cmp r9,8 jb .LOutputMasked8xNBlock\@ test r10b,r10b # ZeroMode? @@ -747,26 +775,6 @@ Implicit Arguments: .endm -// -// Reduce code size for the various types of kernels by sharing the outer logic -// and switching on the selector codes (using sign bit to discriminate). -// - - FUNCTION_ENTRY MlasGemmU8S8KernelAvxVnni - - mov eax,-1 - jmp C_UNDERSCORE(MlasGemmU8X8KernelAvx2) - - FUNCTION_ENTRY MlasGemmU8U8KernelAvx2 - - mov eax,1 - jmp C_UNDERSCORE(MlasGemmU8X8KernelAvx2) - - FUNCTION_ENTRY MlasGemmU8S8KernelAvx2 - - xor eax,eax - jmp C_UNDERSCORE(MlasGemmU8X8KernelAvx2) - /*++ Routine Description: @@ -777,10 +785,10 @@ Routine Description: Arguments: A (rdi) - Supplies the address of matrix A. The matrix data has been packed - using MlasGemmU8X8CopyPackAAvx2. + using MlasGemmCopyPackAAvx2. B (rsi) - Supplies the address of matrix B. The matrix data has been packed - using MlasGemmU8X8CopyPackBAvx2. + using MlasGemmCopyPackBAvx2. C (rdx) - Supplies the address of matrix C. @@ -818,51 +826,62 @@ Return Value: --*/ - FUNCTION_ENTRY MlasGemmU8X8KernelAvx2 +.macro MlasGemmInt8KernelAvx2 ASigned, BSigned push rbp push rbx push r12 push r13 - mov DWORD PTR .LGemmU8X8KernelFrame_type[rsp],eax + mov DWORD PTR .LGemmInt8KernelFrame_type[rsp],eax mov rbx,rdi - mov rax,.LGemmU8X8KernelFrame_ldc[rsp] + mov rax,.LGemmInt8KernelFrame_ldc[rsp] shl rax,2 # convert ldc to bytes shl rcx,2 # convert to row length - movzx r10,BYTE PTR .LGemmU8X8KernelFrame_ZeroMode[rsp] - mov r11,.LGemmU8X8KernelFrame_RowSumBuffer[rsp] - mov r12,.LGemmU8X8KernelFrame_ColumnSumBuffer[rsp] - mov r13,.LGemmU8X8KernelFrame_ZeroPointB[rsp] + movzx r10,BYTE PTR .LGemmInt8KernelFrame_ZeroMode[rsp] + mov r11,.LGemmInt8KernelFrame_RowSumBuffer[rsp] + mov r12,.LGemmInt8KernelFrame_ColumnSumBuffer[rsp] + mov r13,.LGemmInt8KernelFrame_ZeroPointB[rsp] vpcmpeqw ymm12,ymm12,ymm12 # generate 256-bit word vector [0xFFFF] vpsrlw ymm12,ymm12,15 # generate 256-bit word vector [0x0001] - cmp DWORD PTR .LGemmU8X8KernelFrame_type[rsp],0 - je .LCheckCountM4OrMore # U8S8 AVX2 kernel requires extra registers + cmp DWORD PTR .LGemmInt8KernelFrame_type[rsp],0 + je .LCheckCountM4OrMore\@ # U8S8 AVX2 kernel requires extra registers // // Process CountM rows of the matrices. // -.LCheckCountM6OrMore: +.LCheckCountM6OrMore\@: cmp r8,5 - ja .LProcessCountM6 - je .LProcessCountM5 + ja .LProcessCountM6\@ + je .LProcessCountM5\@ -.LCheckCountM4OrMore: +.LCheckCountM4OrMore\@: cmp r8,3 - ja .LProcessCountM4 - je .LProcessCountM3 + ja .LProcessCountM4\@ + je .LProcessCountM3\@ cmp r8,1 - je .LProcessCountM1 + je .LProcessCountM1\@ -.LProcessCountM2: - ProcessCountM 2 +.LProcessCountM2\@: + ProcessCountM 2, \ASigned\(), \BSigned\() -.LProcessCountM4: - ProcessCountM 4 +.LProcessCountM4\@: + ProcessCountM 4, \ASigned\(), \BSigned\() -.LProcessCountM6: - ProcessCountM 6 +.LProcessCountM6\@: + ProcessCountM 6, \ASigned\(), \BSigned\() + +.LProcessCountM1\@: + ProcessCountM 1, \ASigned\(), \BSigned\() + +.LProcessCountM3\@: + ProcessCountM 3, \ASigned\(), \BSigned\() + +.LProcessCountM5\@: + ProcessCountM 5, \ASigned\(), \BSigned\() + +.endm // // Restore non-volatile registers and return. @@ -877,13 +896,39 @@ Return Value: pop rbp ret -.LProcessCountM1: - ProcessCountM 1 +// +// Reduce code size for the various types of kernels by sharing the outer logic +// and switching on the selector codes (using sign bit to discriminate). +// -.LProcessCountM3: - ProcessCountM 3 + FUNCTION_ENTRY MlasGemmU8S8KernelAvxVnni -.LProcessCountM5: - ProcessCountM 5 + mov eax,-1 + MlasGemmInt8KernelAvx2 0, 1 + + FUNCTION_ENTRY MlasGemmU8U8KernelAvx2Vnni + + mov eax,-1 + MlasGemmInt8KernelAvx2 0, 0 + + FUNCTION_ENTRY MlasGemmU8U8KernelAvx2 + + mov eax,1 + MlasGemmInt8KernelAvx2 0, 0 + + FUNCTION_ENTRY MlasGemmU8S8KernelAvx2 + + xor eax,eax + MlasGemmInt8KernelAvx2 0, 1 + + FUNCTION_ENTRY MlasGemmS8S8KernelAvx2Vnni + + mov eax,-1 + MlasGemmInt8KernelAvx2 1, 1 + + FUNCTION_ENTRY MlasGemmS8U8KernelAvx2Vnni + + mov eax,-1 + MlasGemmInt8KernelAvx2 1, 0 .end diff --git a/onnxruntime/test/mlas/unittest/test_qgemm.cpp b/onnxruntime/test/mlas/unittest/test_qgemm.cpp index 6bb93d3535..12955e6f04 100644 --- a/onnxruntime/test/mlas/unittest/test_qgemm.cpp +++ b/onnxruntime/test/mlas/unittest/test_qgemm.cpp @@ -10,6 +10,8 @@ static size_t QGemmRegistLongExecute() { count += MlasLongExecuteTests>::RegisterLongExecute(); count += MlasLongExecuteTests>::RegisterLongExecute(); count += MlasLongExecuteTests>::RegisterLongExecute(); + count += MlasLongExecuteTests>::RegisterLongExecute(); + count += MlasLongExecuteTests>::RegisterLongExecute(); if (GetMlasThreadPool() != nullptr) { count += MlasLongExecuteTests>::RegisterLongExecute(); @@ -18,6 +20,8 @@ static size_t QGemmRegistLongExecute() { count += MlasLongExecuteTests>::RegisterLongExecute(); count += MlasLongExecuteTests>::RegisterLongExecute(); count += MlasLongExecuteTests>::RegisterLongExecute(); + count += MlasLongExecuteTests>::RegisterLongExecute(); + count += MlasLongExecuteTests>::RegisterLongExecute(); } return count; @@ -32,6 +36,8 @@ static size_t QGemmRegistShortExecute() { count += QgemmShortExecuteTest::RegisterShortExecuteTests(); count += QgemmShortExecuteTest::RegisterShortExecuteTests(); count += QgemmShortExecuteTest::RegisterShortExecuteTests(); + count += QgemmShortExecuteTest::RegisterShortExecuteTests(); + count += QgemmShortExecuteTest::RegisterShortExecuteTests(); if (MlasGemmPackBSize(128, 128, false /*AIsSigned*/, false /*BIsSigned*/) > 0) { // QGEMM U8U8=float packed tests count += QgemmShortExecuteTest::RegisterShortExecuteTests(); @@ -45,11 +51,17 @@ static size_t QGemmRegistShortExecute() { count += QgemmShortExecuteTest::RegisterShortExecuteTests(); } if (MlasGemmPackBSize(128, 128, true /*AIsSigned*/, true /*BIsSigned*/) > 0) { - // QGEMM U8S8=float packed tests + // QGEMM S8S8=float packed tests count += QgemmShortExecuteTest::RegisterShortExecuteTests(); - // QGEMM U8S8=int32_t packed tests + // QGEMM S8S8=int32_t packed tests count += QgemmShortExecuteTest::RegisterShortExecuteTests(); } + if (MlasGemmPackBSize(128, 128, true /*AIsSigned*/, false /*BIsSigned*/) > 0) { + // QGEMM S8U8=float packed tests + count += QgemmShortExecuteTest::RegisterShortExecuteTests(); + // QGEMM S8U8=int32_t packed tests + count += QgemmShortExecuteTest::RegisterShortExecuteTests(); + } if (GetMlasThreadPool() != nullptr) { count += QgemmShortExecuteTest::RegisterShortExecuteTests(); @@ -58,6 +70,8 @@ static size_t QGemmRegistShortExecute() { count += QgemmShortExecuteTest::RegisterShortExecuteTests(); count += QgemmShortExecuteTest::RegisterShortExecuteTests(); count += QgemmShortExecuteTest::RegisterShortExecuteTests(); + count += QgemmShortExecuteTest::RegisterShortExecuteTests(); + count += QgemmShortExecuteTest::RegisterShortExecuteTests(); if (MlasGemmPackBSize(128, 128, false /*AIsSigned*/, false /*BIsSigned*/) > 0) { count += QgemmShortExecuteTest::RegisterShortExecuteTests(); count += QgemmShortExecuteTest::RegisterShortExecuteTests(); @@ -70,6 +84,10 @@ static size_t QGemmRegistShortExecute() { count += QgemmShortExecuteTest::RegisterShortExecuteTests(); count += QgemmShortExecuteTest::RegisterShortExecuteTests(); } + if (MlasGemmPackBSize(128, 128, true /*AIsSigned*/, false /*BIsSigned*/) > 0) { + count += QgemmShortExecuteTest::RegisterShortExecuteTests(); + count += QgemmShortExecuteTest::RegisterShortExecuteTests(); + } } return count;