Implementation of AVX-VNNI-INT8 dot product instructions into MLAS GEMM (#21984)

### Description
<!-- Describe your changes. -->
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
<!-- - Why is this change required? What problem does it solve?
- If it fixes an open issue, please link to the issue here. -->
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.
This commit is contained in:
mguynn-intc 2024-09-18 22:18:23 -07:00 committed by GitHub
parent 560778fd07
commit d5f6343a4a
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
11 changed files with 1238 additions and 385 deletions

View file

@ -175,4 +175,156 @@ VpdpwssdsXmmXmmXmm MACRO DestReg, Src1Reg, Src2Reg
VnniXmmXmmXmm 053h, DestReg, Src1Reg, Src2Reg
ENDM
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

View file

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

View file

@ -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, <vmovdqu ymm1,YMMWORD PTR [rdx+VectorOffset+32]>
EmitIfCountGE RowCount, 1, <vpbroadcastd ymm2,DWORD PTR [rcx+BroadcastOffset]>
EmitIfCountGE RowCount, 1, <MultiplyAccumulateRowU8S8AvxVnni ColumnCount, ymm4, ymm5>
EmitIfCountGE RowCount, 1, <MultiplyAccumulateRowAvxVnni ColumnCount, ymm4, ymm5, ASigned, BSigned>
EmitIfCountGE RowCount, 2, <vpbroadcastd ymm2,DWORD PTR [rcx+r9+BroadcastOffset]>
EmitIfCountGE RowCount, 2, <MultiplyAccumulateRowU8S8AvxVnni ColumnCount, ymm6, ymm7>
EmitIfCountGE RowCount, 2, <MultiplyAccumulateRowAvxVnni ColumnCount, ymm6, ymm7, ASigned, BSigned>
EmitIfCountGE RowCount, 3, <vpbroadcastd ymm2,DWORD PTR [rcx+r9*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 [rbx+BroadcastOffset]>
EmitIfCountGE RowCount, 4, <MultiplyAccumulateRowU8S8AvxVnni ColumnCount, ymm10, ymm11>
EmitIfCountGE RowCount, 4, <MultiplyAccumulateRowAvxVnni ColumnCount, ymm10, ymm11, ASigned, BSigned>
EmitIfCountGE RowCount, 5, <vpbroadcastd ymm2,DWORD PTR [rbx+r9+BroadcastOffset]>
EmitIfCountGE RowCount, 5, <MultiplyAccumulateRowU8S8AvxVnni ColumnCount, ymm12, ymm13>
EmitIfCountGE RowCount, 5, <MultiplyAccumulateRowAvxVnni ColumnCount, ymm12, ymm13, ASigned, BSigned>
EmitIfCountGE RowCount, 6, <vpbroadcastd ymm2,DWORD PTR [rbx+r9*2+BroadcastOffset]>
EmitIfCountGE RowCount, 6, <MultiplyAccumulateRowU8S8AvxVnni ColumnCount, ymm14, ymm15>
EmitIfCountGE RowCount, 6, <MultiplyAccumulateRowAvxVnni ColumnCount, ymm14, ymm15, ASigned, BSigned>
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

View file

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

View file

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

View file

@ -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<KernelType>(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) {

View file

@ -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>(
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>(
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<MLAS_GEMM_U8U8_KERNEL_AVX2VNNI>(
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<MLAS_GEMM_U8U8_KERNEL_AVX2VNNI>,
MlasGemmQuantPackedOperation<MLAS_GEMM_U8U8_KERNEL_AVX2VNNI>,
MlasGemmQuantCopyPackB<MLAS_GEMM_U8U8_KERNEL_AVX2VNNI>,
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>(
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>(
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<MLAS_GEMM_S8S8_KERNEL_AVX2>(
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<MLAS_GEMM_S8S8_KERNEL_AVX2>,
MlasGemmQuantPackedOperation<MLAS_GEMM_S8S8_KERNEL_AVX2>,
MlasGemmQuantCopyPackB<MLAS_GEMM_S8S8_KERNEL_AVX2>,
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>(
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>(
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<MLAS_GEMM_S8U8_KERNEL_AVX2>(
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<MLAS_GEMM_S8U8_KERNEL_AVX2>,
MlasGemmQuantPackedOperation<MLAS_GEMM_S8U8_KERNEL_AVX2>,
MlasGemmQuantCopyPackB<MLAS_GEMM_S8U8_KERNEL_AVX2>,
MLAS_GEMM_S8U8_KERNEL_AVX2::PackedK,
MLAS_GEMM_S8U8_KERNEL_AVX2::PackedStrides.K,
6 // assembly kernel M stride
};

View file

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

View file

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

View file

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

View file

@ -10,6 +10,8 @@ static size_t QGemmRegistLongExecute() {
count += MlasLongExecuteTests<MlasQgemmTest<uint8_t, uint8_t, int32_t, true, false>>::RegisterLongExecute();
count += MlasLongExecuteTests<MlasQgemmTest<int8_t, int8_t, int32_t, false, false>>::RegisterLongExecute();
count += MlasLongExecuteTests<MlasQgemmTest<int8_t, int8_t, int32_t, true, false>>::RegisterLongExecute();
count += MlasLongExecuteTests<MlasQgemmTest<int8_t, uint8_t, int32_t, false, false>>::RegisterLongExecute();
count += MlasLongExecuteTests<MlasQgemmTest<int8_t, uint8_t, int32_t, true, false>>::RegisterLongExecute();
if (GetMlasThreadPool() != nullptr) {
count += MlasLongExecuteTests<MlasQgemmTest<uint8_t, int8_t, int32_t, false, true>>::RegisterLongExecute();
@ -18,6 +20,8 @@ static size_t QGemmRegistLongExecute() {
count += MlasLongExecuteTests<MlasQgemmTest<uint8_t, uint8_t, int32_t, true, true>>::RegisterLongExecute();
count += MlasLongExecuteTests<MlasQgemmTest<int8_t, int8_t, int32_t, false, true>>::RegisterLongExecute();
count += MlasLongExecuteTests<MlasQgemmTest<int8_t, int8_t, int32_t, true, true>>::RegisterLongExecute();
count += MlasLongExecuteTests<MlasQgemmTest<int8_t, uint8_t, int32_t, false, true>>::RegisterLongExecute();
count += MlasLongExecuteTests<MlasQgemmTest<int8_t, uint8_t, int32_t, true, true>>::RegisterLongExecute();
}
return count;
@ -32,6 +36,8 @@ static size_t QGemmRegistShortExecute() {
count += QgemmShortExecuteTest<uint8_t, uint8_t, int32_t, false, false>::RegisterShortExecuteTests();
count += QgemmShortExecuteTest<int8_t, int8_t, float, false, false>::RegisterShortExecuteTests();
count += QgemmShortExecuteTest<int8_t, int8_t, int32_t, false, false>::RegisterShortExecuteTests();
count += QgemmShortExecuteTest<int8_t, uint8_t, float, false, false>::RegisterShortExecuteTests();
count += QgemmShortExecuteTest<int8_t, uint8_t, int32_t, false, false>::RegisterShortExecuteTests();
if (MlasGemmPackBSize(128, 128, false /*AIsSigned*/, false /*BIsSigned*/) > 0) {
// QGEMM U8U8=float packed tests
count += QgemmShortExecuteTest<uint8_t, uint8_t, float, true, false>::RegisterShortExecuteTests();
@ -45,11 +51,17 @@ static size_t QGemmRegistShortExecute() {
count += QgemmShortExecuteTest<uint8_t, int8_t, int32_t, true, false>::RegisterShortExecuteTests();
}
if (MlasGemmPackBSize(128, 128, true /*AIsSigned*/, true /*BIsSigned*/) > 0) {
// QGEMM U8S8=float packed tests
// QGEMM S8S8=float packed tests
count += QgemmShortExecuteTest<int8_t, int8_t, float, true, false>::RegisterShortExecuteTests();
// QGEMM U8S8=int32_t packed tests
// QGEMM S8S8=int32_t packed tests
count += QgemmShortExecuteTest<int8_t, int8_t, int32_t, true, false>::RegisterShortExecuteTests();
}
if (MlasGemmPackBSize(128, 128, true /*AIsSigned*/, false /*BIsSigned*/) > 0) {
// QGEMM S8U8=float packed tests
count += QgemmShortExecuteTest<int8_t, uint8_t, float, true, false>::RegisterShortExecuteTests();
// QGEMM S8U8=int32_t packed tests
count += QgemmShortExecuteTest<int8_t, uint8_t, int32_t, true, false>::RegisterShortExecuteTests();
}
if (GetMlasThreadPool() != nullptr) {
count += QgemmShortExecuteTest<uint8_t, int8_t, float, false, true>::RegisterShortExecuteTests();
@ -58,6 +70,8 @@ static size_t QGemmRegistShortExecute() {
count += QgemmShortExecuteTest<uint8_t, uint8_t, int32_t, false, true>::RegisterShortExecuteTests();
count += QgemmShortExecuteTest<int8_t, int8_t, float, false, true>::RegisterShortExecuteTests();
count += QgemmShortExecuteTest<int8_t, int8_t, int32_t, false, true>::RegisterShortExecuteTests();
count += QgemmShortExecuteTest<int8_t, uint8_t, float, false, true>::RegisterShortExecuteTests();
count += QgemmShortExecuteTest<int8_t, uint8_t, int32_t, false, true>::RegisterShortExecuteTests();
if (MlasGemmPackBSize(128, 128, false /*AIsSigned*/, false /*BIsSigned*/) > 0) {
count += QgemmShortExecuteTest<uint8_t, uint8_t, float, true, true>::RegisterShortExecuteTests();
count += QgemmShortExecuteTest<uint8_t, uint8_t, int32_t, true, true>::RegisterShortExecuteTests();
@ -70,6 +84,10 @@ static size_t QGemmRegistShortExecute() {
count += QgemmShortExecuteTest<int8_t, int8_t, float, true, true>::RegisterShortExecuteTests();
count += QgemmShortExecuteTest<int8_t, int8_t, int32_t, true, true>::RegisterShortExecuteTests();
}
if (MlasGemmPackBSize(128, 128, true /*AIsSigned*/, false /*BIsSigned*/) > 0) {
count += QgemmShortExecuteTest<int8_t, uint8_t, float, true, true>::RegisterShortExecuteTests();
count += QgemmShortExecuteTest<int8_t, uint8_t, int32_t, true, true>::RegisterShortExecuteTests();
}
}
return count;