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;