From 5bcb5f5a3d25fa0c33b3909a5f9b1e8c57617bf4 Mon Sep 17 00:00:00 2001 From: Ramakrishnan Sivakumar <63471807+ramkrishna0116@users.noreply.github.com> Date: Mon, 26 Oct 2020 16:27:48 -0700 Subject: [PATCH] MLAS: Add support for AVXVNNI (#5592) Adds Gemm kernels with AVXVNNI support for Int8 acceleration --- cmake/onnxruntime_mlas.cmake | 4 + .../core/mlas/lib/amd64/AssembleAvxVnni.inc | 178 ++++++++ .../mlas/lib/amd64/QgemmU8S8KernelAvxVnni.asm | 298 ++++++++++++++ .../mlas/lib/amd64/QgemvU8S8KernelAvxVnni.asm | 385 ++++++++++++++++++ onnxruntime/core/mlas/lib/mlasi.h | 2 + onnxruntime/core/mlas/lib/platform.cpp | 19 + .../core/mlas/lib/x86_64/AssembleAvxVnni.h | 178 ++++++++ .../mlas/lib/x86_64/QgemmU8S8KernelAvxVnni.S | 268 ++++++++++++ .../mlas/lib/x86_64/QgemvU8S8KernelAvxVnni.S | 344 ++++++++++++++++ 9 files changed, 1676 insertions(+) create mode 100644 onnxruntime/core/mlas/lib/amd64/AssembleAvxVnni.inc create mode 100644 onnxruntime/core/mlas/lib/amd64/QgemmU8S8KernelAvxVnni.asm create mode 100644 onnxruntime/core/mlas/lib/amd64/QgemvU8S8KernelAvxVnni.asm create mode 100644 onnxruntime/core/mlas/lib/x86_64/AssembleAvxVnni.h create mode 100644 onnxruntime/core/mlas/lib/x86_64/QgemmU8S8KernelAvxVnni.S create mode 100644 onnxruntime/core/mlas/lib/x86_64/QgemvU8S8KernelAvxVnni.S diff --git a/cmake/onnxruntime_mlas.cmake b/cmake/onnxruntime_mlas.cmake index fee9648460..e3a5ce28cc 100644 --- a/cmake/onnxruntime_mlas.cmake +++ b/cmake/onnxruntime_mlas.cmake @@ -82,6 +82,8 @@ if(MSVC) ${ONNXRUNTIME_ROOT}/core/mlas/lib/amd64/QgemvU8S8KernelAvx512Core.asm ${ONNXRUNTIME_ROOT}/core/mlas/lib/amd64/QgemmU8S8KernelAvx512Vnni.asm ${ONNXRUNTIME_ROOT}/core/mlas/lib/amd64/QgemvU8S8KernelAvx512Vnni.asm + ${ONNXRUNTIME_ROOT}/core/mlas/lib/amd64/QgemmU8S8KernelAvxVnni.asm + ${ONNXRUNTIME_ROOT}/core/mlas/lib/amd64/QgemvU8S8KernelAvxVnni.asm ${ONNXRUNTIME_ROOT}/core/mlas/lib/amd64/QgemmU8U8KernelAvx2.asm ${ONNXRUNTIME_ROOT}/core/mlas/lib/amd64/QgemmU8U8KernelAvx512Core.asm ${ONNXRUNTIME_ROOT}/core/mlas/lib/amd64/DgemmKernelSse2.asm @@ -242,6 +244,8 @@ else() ${ONNXRUNTIME_ROOT}/core/mlas/lib/x86_64/QgemmU8S8KernelAvx2.S ${ONNXRUNTIME_ROOT}/core/mlas/lib/x86_64/QgemvU8S8KernelAvx2.S ${ONNXRUNTIME_ROOT}/core/mlas/lib/x86_64/QgemmU8U8KernelAvx2.S + ${ONNXRUNTIME_ROOT}/core/mlas/lib/x86_64/QgemmU8S8KernelAvxVnni.S + ${ONNXRUNTIME_ROOT}/core/mlas/lib/x86_64/QgemvU8S8KernelAvxVnni.S ${ONNXRUNTIME_ROOT}/core/mlas/lib/x86_64/DgemmKernelFma3.S ${ONNXRUNTIME_ROOT}/core/mlas/lib/x86_64/SgemmKernelFma3.S ${ONNXRUNTIME_ROOT}/core/mlas/lib/x86_64/SconvKernelFma3.S diff --git a/onnxruntime/core/mlas/lib/amd64/AssembleAvxVnni.inc b/onnxruntime/core/mlas/lib/amd64/AssembleAvxVnni.inc new file mode 100644 index 0000000000..8717528d0d --- /dev/null +++ b/onnxruntime/core/mlas/lib/amd64/AssembleAvxVnni.inc @@ -0,0 +1,178 @@ +;++ +; +; Copyright (c) 2020 Intel Corporation. All rights reserved. +; +; Licensed under the MIT License. +; +; Module Name: +; +; AssembleAvxVnni.inc +; +; Abstract: +; +; This module contains macros to build AVXVNNI instructions for toolchains that +; do not natively support this newer instruction set extension. +; +;-- + +; +; Map friendly register names to the encoded register index. +; + +YmmIndex_ymm0 EQU 0 +YmmIndex_ymm1 EQU 1 +YmmIndex_ymm2 EQU 2 +YmmIndex_ymm3 EQU 3 +YmmIndex_ymm4 EQU 4 +YmmIndex_ymm5 EQU 5 +YmmIndex_ymm6 EQU 6 +YmmIndex_ymm7 EQU 7 +YmmIndex_ymm8 EQU 8 +YmmIndex_ymm9 EQU 9 +YmmIndex_ymm10 EQU 10 +YmmIndex_ymm11 EQU 11 +YmmIndex_ymm12 EQU 12 +YmmIndex_ymm13 EQU 13 +YmmIndex_ymm14 EQU 14 +YmmIndex_ymm15 EQU 15 + +XmmIndex_xmm0 EQU 0 +XmmIndex_xmm1 EQU 1 +XmmIndex_xmm2 EQU 2 +XmmIndex_xmm3 EQU 3 +XmmIndex_xmm4 EQU 4 +XmmIndex_xmm5 EQU 5 +XmmIndex_xmm6 EQU 6 +XmmIndex_xmm7 EQU 7 +XmmIndex_xmm8 EQU 8 +XmmIndex_xmm9 EQU 9 +XmmIndex_xmm10 EQU 10 +XmmIndex_xmm11 EQU 11 +XmmIndex_xmm12 EQU 12 +XmmIndex_xmm13 EQU 13 +XmmIndex_xmm14 EQU 14 +XmmIndex_xmm15 EQU 15 + +; +; Macro Description: +; +; This macro builds a VNNI instruction of the form: +; +; instr ymm1,ymm2,ymm3 +; +; Arguments: +; +; Opcode - Specifies the opcode for the VNNI instruction. +; +; DestReg - Specifies the destination register. +; +; Src1Reg - Specifies the first source register. +; +; Src2Reg - Specifies the second source register. +; + +VnniYmmYmmYmm MACRO Opcode, 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 = 005h ; "66" 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 + +VpdpbusdYmmYmmYmm MACRO DestReg, Src1Reg, Src2Reg + + VnniYmmYmmYmm 050h, DestReg, Src1Reg, Src2Reg + + ENDM + +VpdpbusdsYmmYmmYmm MACRO DestReg, Src1Reg, Src2Reg + + VnniYmmYmmYmm 051h, DestReg, Src1Reg, Src2Reg + + ENDM + +VpdpwssdYmmYmmYmm MACRO DestReg, Src1Reg, Src2Reg + + VnniYmmYmmYmm 052h, DestReg, Src1Reg, Src2Reg + + ENDM + +VpdpwssdsYmmYmmYmm MACRO DestReg, Src1Reg, Src2Reg + + VnniYmmYmmYmm 053h, 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. +; +; DestReg - Specifies the destination register. +; +; Src1Reg - Specifies the first source register. +; +; Src2Reg - Specifies the second source register. +; + +VnniXmmXmmXmm MACRO Opcode, 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 = 001h ; "66" 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 + +VpdpbusdXmmXmmXmm MACRO DestReg, Src1Reg, Src2Reg + + VnniXmmXmmXmm 050h, DestReg, Src1Reg, Src2Reg + + ENDM + +VpdpbusdsXmmXmmXmm MACRO DestReg, Src1Reg, Src2Reg + + VnniXmmXmmXmm 051h, DestReg, Src1Reg, Src2Reg + + ENDM + +VpdpwssdXmmXmmXmm MACRO DestReg, Src1Reg, Src2Reg + + VnniXmmXmmXmm 052h, DestReg, Src1Reg, Src2Reg + + ENDM + +VpdpwssdsXmmXmmXmm MACRO DestReg, Src1Reg, Src2Reg + + VnniXmmXmmXmm 053h, DestReg, Src1Reg, Src2Reg + + ENDM \ No newline at end of file diff --git a/onnxruntime/core/mlas/lib/amd64/QgemmU8S8KernelAvxVnni.asm b/onnxruntime/core/mlas/lib/amd64/QgemmU8S8KernelAvxVnni.asm new file mode 100644 index 0000000000..613fd7b410 --- /dev/null +++ b/onnxruntime/core/mlas/lib/amd64/QgemmU8S8KernelAvxVnni.asm @@ -0,0 +1,298 @@ +;++ +; +; Copyright (c) Intel Corporation 2020. All rights reserved. +; +; Licensed under the MIT License. +; +; Module Name: +; +; QgemmU8S8KernelAvxVnni.asm +; +; Abstract: +; +; This module implements the kernels for the quantized integer matrix/matrix +; multiply operation (QGEMM). +; +; This implementation uses AVXVNNI instructions. +; +;-- + + .xlist +INCLUDE mlasi.inc +INCLUDE QgemmU8X8KernelAvx2Common.inc +INCLUDE AssembleAvxVnni.inc + .list + +; +; Macro Description: +; +; This macro generates code to multiply and accumulator a single row of the +; output block. +; +; Arguments: +; +; ColumnCount - Supplies the number of columns to produce. +; +; Vec1Reg - Supplies the high block accumulator register (when ColumnCount +; is 16). +; +; Vec2Reg - Supplies the low block accumulator register. +; +; Implicit Arguments: +; +; ymm0 - Supplies the first vector loaded from matrix B. +; +; ymm1 - Supplies the second vector loaded from matrix B (when ColumnCount +; is 16). +; +; ymm2 - Supplies the broadcast value loaded from matrix A. +; + +MultiplyAccumulateRow MACRO ColumnCount, Vec1Reg, Vec2Reg + +IF ColumnCount EQ 16 + VpdpbusdsYmmYmmYmm Vec1Reg,ymm2,ymm0 + VpdpbusdsYmmYmmYmm Vec2Reg,ymm2,ymm1 +ELSE + VpdpbusdsYmmYmmYmm Vec2Reg,ymm2,ymm0 +ENDIF + + ENDM + +; +; Macro Description: +; +; This macro generates code to multiply and accumulate each row of the output +; block. +; +; Arguments: +; +; ColumnCount - Supplies the number of columns to produce. +; +; RowCount - Supplies the number of rows to produce. +; +; VectorOffset - Supplies the byte offset from matrix B to fetch elements. +; +; BroadcastOffset - Supplies the byte offset from matrix A to fetch elements. +; +; Implicit Arguments: +; +; rbx - Supplies the address into the matrix A data plus 3 rows. +; +; rcx - Supplies the address into the matrix A data. +; +; rdx - Supplies the address into the matrix B data. +; +; r9 - Supplies the length in bytes of a row from matrix A. +; +; ymm4-ymm15 - Supplies the block accumulators. +; + +ComputeBlock MACRO ColumnCount, RowCount, VectorOffset, BroadcastOffset + + vmovdqu ymm0,YMMWORD PTR [rdx+VectorOffset] + EmitIfCountGE ColumnCount, 16, + EmitIfCountGE RowCount, 1, + EmitIfCountGE RowCount, 1, + EmitIfCountGE RowCount, 2, + EmitIfCountGE RowCount, 2, + EmitIfCountGE RowCount, 3, + EmitIfCountGE RowCount, 3, + EmitIfCountGE RowCount, 4, + EmitIfCountGE RowCount, 4, + EmitIfCountGE RowCount, 5, + EmitIfCountGE RowCount, 6, + EmitIfCountGE RowCount, 6, + + ENDM + +; +; Macro Description: +; +; This macro generates code to execute the block compute macro multiple +; times and advancing the matrix A and matrix B data pointers. +; +; Arguments: +; +; ColumnCount - Supplies the number of columns to produce. +; +; RowCount - Supplies the number of rows to produce. +; +; Implicit Arguments: +; +; rbx - Supplies the address into the matrix A data plus 3 rows. +; +; rcx - Supplies the address into the matrix A data. +; +; rdx - Supplies the address into the matrix B data. +; +; r9 - Supplies the length in bytes of a row from matrix A. +; +; ymm4-ymm15 - Supplies the block accumulators. +; + +ComputeBlockLoop MACRO ColumnCount, RowCount + + LOCAL ComputeBlockBy1Loop + + mov rsi,r9 ; reload row length remaining + +ComputeBlockBy1Loop: + ComputeBlock ColumnCount, RowCount, 0, 0 + add rcx,4 ; advance matrix A by 1 quad +IF RowCount GT 3 + add rbx,4 ; advance matrix A plus 3 rows by 1 quad +ENDIF + add rdx,64 ; advance matrix B + sub rsi,4 + jnz ComputeBlockBy1Loop + + ENDM + +;++ +; +; Routine Description: +; +; This routine is an inner kernel to compute matrix multiplication for a +; set of rows. +; +; Arguments: +; +; A (rcx) - Supplies the address of matrix A. The matrix data has been packed +; using MlasGemmU8S8CopyPackAAvx2. +; +; B (rdx) - Supplies the address of matrix B. The matrix data has been packed +; using MlasGemmU8S8CopyPackBAvx2. +; +; C (r8) - Supplies the address of matrix C. +; +; PackedCountK (r9) - Supplies the number of packed columns from matrix A and +; the number of packed rows from matrix B to iterate over. +; +; CountM - Supplies the maximum number of rows that can be processed for +; matrix A and matrix C. The actual number of rows handled for this +; invocation depends on the kernel implementation. +; +; CountN - Supplies the number of columns from matrix B and matrix C to iterate +; over. +; +; ldc - Supplies the first dimension of matrix C. +; +; RowSumBuffer - Supplies the sum of each row from matrix A multiplied by the +; zero point offset of matrix B. These values are accumulated into every +; row of matrix C. +; +; ColumnSumBuffer - Supplies the sum of each column from matrix B multiplied +; by the zero point offset of matrix A. These values are accumulated into +; every column of matrix C. +; +; DepthValue - Supplies the value CountK multiplied by the zero point offset +; of matrix A multplied by the zero point offset of matrix B. This value is +; accumulated into every element of matrix C. +; +; ZeroMode - Supplies true if the output matrix must be zero initialized, +; else false if the output matrix is accumulated into. +; +; Return Value: +; +; Returns the number of rows handled. +; +;-- + + NESTED_ENTRY MlasGemmU8S8KernelAvxVnni, _TEXT + + rex_push_reg rbp + push_reg rbx + push_reg rsi + 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 + + END_PROLOGUE + + mov rdi,rcx + mov rbp,GemmU8X8KernelFrame.CountN[rsp] + mov rax,GemmU8X8KernelFrame.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.CountM[rsp] + mov r12,GemmU8X8KernelFrame.RowSumBuffer[rsp] + mov r13,GemmU8X8KernelFrame.ColumnSumBuffer[rsp] + +; +; Process CountM rows of the matrices. +; + + cmp r11,5 + ja ProcessCountM6 + je ProcessCountM5 + cmp r11,3 + ja ProcessCountM4 + je ProcessCountM3 + cmp r11,1 + je ProcessCountM1 + +ProcessCountM2: + ProcessCountM 2 + +ProcessCountM4: + ProcessCountM 4 + +ProcessCountM6: + mov r11d,6 ; return 6 rows handled + ProcessCountM 6, Fallthrough + +; +; Restore non-volatile registers and return. +; + +ExitKernel: + mov eax,r11d + 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) + + BEGIN_EPILOGUE + + pop r13 + pop r12 + pop rdi + pop rsi + pop rbx + pop rbp + ret + +ProcessCountM1: + ProcessCountM 1 + +ProcessCountM3: + ProcessCountM 3 + +ProcessCountM5: + ProcessCountM 5 + + NESTED_END MlasGemmU8S8KernelAvxVnni, _TEXT + + END \ No newline at end of file diff --git a/onnxruntime/core/mlas/lib/amd64/QgemvU8S8KernelAvxVnni.asm b/onnxruntime/core/mlas/lib/amd64/QgemvU8S8KernelAvxVnni.asm new file mode 100644 index 0000000000..be13dcebbc --- /dev/null +++ b/onnxruntime/core/mlas/lib/amd64/QgemvU8S8KernelAvxVnni.asm @@ -0,0 +1,385 @@ +;++ +; +; Copyright (c) 2020 Intel Corporation. All rights reserved. +; +; Licensed under the MIT License. +; +; Module Name: +; +; QgemvU8S8KernelAvxVnni.asm +; +; Abstract: +; +; This module implements the kernels for the quantized integer matrix/vector +; multiply operation (QGEMV). +; +; This implementation uses AVXVNNI instructions. +; +;-- + + .xlist +INCLUDE mlasi.inc +INCLUDE AssembleAvxVnni.inc + .list + + EXTERN MlasMaskMoveAvx:NEAR + EXTERN MlasTranspose4x4BytesAvx:NEAR + +; +; Stack frame layout for the U8S8 kernel. +; + +GemvU8S8KernelFrame STRUCT + + SavedXmm6 OWORD ? + SavedXmm7 OWORD ? + SavedXmm8 OWORD ? + SavedXmm9 OWORD ? + SavedXmm10 OWORD ? + Padding QWORD ? + SavedRdi QWORD ? + SavedRsi QWORD ? + SavedRbx QWORD ? + SavedRbp QWORD ? + ReturnAddress QWORD ? + PreviousP1Home QWORD ? + PreviousP2Home QWORD ? + PreviousP3Home QWORD ? + PreviousP4Home QWORD ? + CountN QWORD ? + ldb QWORD ? + +GemvU8S8KernelFrame ENDS + +;++ +; +; Routine Description: +; +; This routine is an inner kernel to compute matrix/vector multiplication. +; +; Arguments: +; +; A (rcx) - Supplies the address of vector A. +; +; B (rdx) - Supplies the address of matrix B. +; +; C (r8) - Supplies the address of matrix C. +; +; CountK (r9) - Supplies the number of columns from vector A and the number +; of rows from matrix B to iterate over. +; +; CountN - Supplies the number of columns from matrix B and matrix C to iterate +; over. +; +; ldb - Supplies the first dimension of matrix B. +; +; Return Value: +; +; None. +; +;-- + + NESTED_ENTRY MlasGemvU8S8KernelAvxVnni, _TEXT + + rex_push_reg rbp + push_reg rbx + push_reg rsi + push_reg rdi + alloc_stack (GemvU8S8KernelFrame.SavedRdi) + save_xmm128 xmm6,GemvU8S8KernelFrame.SavedXmm6 + save_xmm128 xmm7,GemvU8S8KernelFrame.SavedXmm7 + save_xmm128 xmm8,GemvU8S8KernelFrame.SavedXmm8 + save_xmm128 xmm9,GemvU8S8KernelFrame.SavedXmm9 + save_xmm128 xmm10,GemvU8S8KernelFrame.SavedXmm10 + + END_PROLOGUE + + mov rsi,rdx + mov rdi,GemvU8S8KernelFrame.ldb[rsp] + mov r10,GemvU8S8KernelFrame.CountN[rsp] + mov r11,rsp ; set ZeroMode to any non-zero value + +; +; Process 4 rows of matrix B in a loop. +; + + sub r9,4 + jb ProcessRemainingRows + +ProcessRowLoop4: + mov rdx,rsi ; reload matrix B + lea rsi,[rsi+rdi*4] ; advance matrix B by 4 rows + mov rbx,r8 ; reload matrix C + mov rbp,r10 ; reload CountN + vpbroadcastd ymm0,DWORD PTR [rcx] + add rcx,4 ; advance matrix A by 4 bytes + +; +; Process sets of 32 columns from the 4 rows in a loop. +; +; Some permute operations are deferred until the final store of the 4x32 block +; as these permutes are expensive. +; + +ProcessColumnLoop4By32: + cmp rbp,32 + jb ProcessColumnLoop4By8 + lea rax,[rdx+rdi*2] ; compute matrix B plus 2 rows + vmovdqu ymm2,YMMWORD PTR [rdx] + vmovdqu ymm3,YMMWORD PTR [rdx+rdi] + vmovdqu ymm4,YMMWORD PTR [rax] + vmovdqu ymm5,YMMWORD PTR [rax+rdi] + vpunpcklbw ymm1,ymm2,ymm3 ; interleave row data bytes + vpunpckhbw ymm2,ymm2,ymm3 + vpxor ymm7,ymm7,ymm7 + vpunpcklbw ymm3,ymm4,ymm5 + vpunpckhbw ymm4,ymm4,ymm5 + vpxor ymm8,ymm8,ymm8 + vpunpcklwd ymm5,ymm1,ymm3 ; interleave row data words + vpunpckhwd ymm1,ymm1,ymm3 + vpxor ymm9,ymm9,ymm9 + vpunpcklwd ymm3,ymm2,ymm4 + vpunpckhwd ymm2,ymm2,ymm4 + vpxor ymm10,ymm10,ymm10 + VpdpbusdsYmmYmmYmm ymm7,ymm0,ymm5 + VpdpbusdsYmmYmmYmm ymm8,ymm0,ymm1 + VpdpbusdsYmmYmmYmm ymm9,ymm0,ymm3 + VpdpbusdsYmmYmmYmm ymm10,ymm0,ymm2 + test r11,r11 ; ZeroMode? + jnz SkipAccumulateOutput4By32 + vpaddd ymm7,ymm7,YMMWORD PTR [rbx] + vpaddd ymm8,ymm8,YMMWORD PTR [rbx+32] + vpaddd ymm9,ymm9,YMMWORD PTR [rbx+64] + vpaddd ymm10,ymm10,YMMWORD PTR [rbx+96] + +SkipAccumulateOutput4By32: + cmp r9,4 ; final 4x32 block? + jae StoreOutput4By32 + vperm2i128 ymm4,ymm7,ymm8,31h ; interleave vector results + vperm2i128 ymm7,ymm7,ymm8,20h + vperm2i128 ymm8,ymm9,ymm10,20h + vperm2i128 ymm10,ymm9,ymm10,31h + vmovaps ymm9,ymm4 + +StoreOutput4By32: + vmovdqu YMMWORD PTR [rbx],ymm7 + vmovdqu YMMWORD PTR [rbx+32],ymm8 + vmovdqu YMMWORD PTR [rbx+64],ymm9 + vmovdqu YMMWORD PTR [rbx+96],ymm10 + add rdx,32 ; advance matrix B by 32 bytes + add rbx,32*4 ; advance matrix C by 32 columns + sub rbp,32 ; decrement CountN + jnz ProcessColumnLoop4By32 + +AdvanceRowLoop4: + xor r11,r11 ; clear ZeroMode + sub r9,4 ; decrement CountK + jae ProcessRowLoop4 + +ProcessRemainingRows: + add r9,4 ; correct for over-subtract above + jnz ProcessRemainingSmallK + +; +; Restore non-volatile registers and return. +; + +ExitKernel: + vzeroupper + movaps xmm6,GemvU8S8KernelFrame.SavedXmm6[rsp] + movaps xmm7,GemvU8S8KernelFrame.SavedXmm7[rsp] + movaps xmm8,GemvU8S8KernelFrame.SavedXmm8[rsp] + movaps xmm9,GemvU8S8KernelFrame.SavedXmm9[rsp] + movaps xmm10,GemvU8S8KernelFrame.SavedXmm10[rsp] + add rsp,(GemvU8S8KernelFrame.SavedRdi) + + BEGIN_EPILOGUE + + pop rdi + pop rsi + pop rbx + pop rbp + ret + +; +; Process sets of 8 columns from the 4 rows in a loop. +; + +ProcessColumnLoop4By8: + cmp ebp,8 + jb ProcessColumn4By4 + lea rax,[rdx+rdi*2] ; compute matrix B plus 2 rows + vmovq xmm2,QWORD PTR [rdx] + vmovq xmm3,QWORD PTR [rdx+rdi] + vmovq xmm4,QWORD PTR [rax] + vmovq xmm5,QWORD PTR [rax+rdi] + vpunpcklbw xmm2,xmm2,xmm3 ; interleave row data bytes + vpunpcklbw xmm4,xmm4,xmm5 + vpunpcklwd xmm1,xmm2,xmm4 ; interleave row data words + vpunpckhwd xmm2,xmm2,xmm4 + vinserti128 ymm1,ymm1,xmm2,1 ; concatenate vector + vpxor ymm8,ymm8,ymm8 + VpdpbusdsYmmYmmYmm ymm8,ymm0,ymm1 + test r11,r11 ; ZeroMode? + jnz SkipAccumulateOutput4By8 + vpaddd ymm8,ymm8,YMMWORD PTR [rbx] + +SkipAccumulateOutput4By8: + vmovdqu YMMWORD PTR [rbx],ymm8 + add rdx,8 ; advance matrix B by 8 bytes + add rbx,8*4 ; advance matrix C by 8 columns + sub ebp,8 ; decrement CountN + jnz ProcessColumnLoop4By8 + jmp AdvanceRowLoop4 + +; +; Process a set of 4 columns from the 4 rows. +; + +ProcessColumn4By4: + test ebp,4 ; (CountN & 4) != 0? + jz ProcessColumn4BySmallN + lea rax,[rdx+rdi*2] ; compute matrix B plus 2 rows + vmovd xmm1,DWORD PTR [rdx] + vpinsrd xmm1,xmm1,DWORD PTR [rdx+rdi],1 + vpinsrd xmm1,xmm1,DWORD PTR [rax],2 + vpinsrd xmm1,xmm1,DWORD PTR [rax+rdi],3 + vpshufb xmm1,xmm1,XMMWORD PTR [MlasTranspose4x4BytesAvx] + vpxor xmm8,xmm8,xmm8 + VpdpbusdsXmmXmmXmm xmm8,xmm0,xmm1 + test r11,r11 ; ZeroMode? + jnz SkipAccumulateOutput4By4 + vpaddd xmm8,xmm8,XMMWORD PTR [rbx] + +SkipAccumulateOutput4By4: + vmovdqu XMMWORD PTR [rbx],xmm8 + and ebp,3 ; (CountN & 3) != 0? + jz AdvanceRowLoop4 + add rdx,4 ; advance matrix B by 4 bytes + add rbx,4*4 ; advance matrix C by 4 columns + +; +; Process the remaining 1 to 3 columns from the 4 rows. +; + +ProcessColumn4BySmallN: + mov DWORD PTR GemvU8S8KernelFrame.CountN[rsp],ebp + vbroadcastss xmm2,DWORD PTR GemvU8S8KernelFrame.CountN[rsp] + vpcmpgtd xmm2,xmm2,XMMWORD PTR [MlasMaskMoveAvx] + vpxor xmm1,xmm1,xmm1 + lea rax,[rdx+rdi*2] ; compute matrix B plus 2 rows + cmp ebp,2 ; (CountN & 2) != 0? + jb ProcessColumn4By1 + vpinsrw xmm1,xmm1,WORD PTR [rdx],0 + vpinsrw xmm1,xmm1,WORD PTR [rdx+rdi],2 + vpinsrw xmm1,xmm1,WORD PTR [rax],4 + vpinsrw xmm1,xmm1,WORD PTR [rax+rdi],6 + je ComputeOutput4BySmallN + vpinsrb xmm1,xmm1,BYTE PTR [rdx+2],2 + vpinsrb xmm1,xmm1,BYTE PTR [rdx+rdi+2],6 + vpinsrb xmm1,xmm1,BYTE PTR [rax+2],10 + vpinsrb xmm1,xmm1,BYTE PTR [rax+rdi+2],14 + jmp ComputeOutput4BySmallN + +ProcessColumn4By1: + vpinsrb xmm1,xmm1,BYTE PTR [rdx],0 + vpinsrb xmm1,xmm1,BYTE PTR [rdx+rdi],4 + vpinsrb xmm1,xmm1,BYTE PTR [rax],8 + vpinsrb xmm1,xmm1,BYTE PTR [rax+rdi],12 + +ComputeOutput4BySmallN: + vpshufb xmm1,xmm1,XMMWORD PTR [MlasTranspose4x4BytesAvx] + vpxor xmm8,xmm8,xmm8 + VpdpbusdsXmmXmmXmm xmm8,xmm0,xmm1 + test r11,r11 ; ZeroMode? + jnz StoreOutput4BySmallN + vpmaskmovd xmm3,xmm2,XMMWORD PTR [rbx] + vpaddd xmm8,xmm8,xmm3 + +StoreOutput4BySmallN: + vpmaskmovd XMMWORD PTR [rbx],xmm2,xmm8 + jmp AdvanceRowLoop4 + +; +; Broadcast the remaining 1 to 3 values from vector A. +; + +ProcessRemainingSmallK: + vpxor xmm5,xmm5,xmm5 ; keep zero vector for vpinsrb/vpinsrw + cmp r9d,2 + jb LoadVectorASingleRemainingByte + vpinsrw xmm0,xmm5,WORD PTR [rcx],0 + je BroadcastVectorARemainingBytes + vpinsrb xmm0,xmm0,BYTE PTR [rcx+2],2 + jmp BroadcastVectorARemainingBytes + +LoadVectorASingleRemainingByte: + vpinsrb xmm0,xmm5,BYTE PTR [rcx],0 + +BroadcastVectorARemainingBytes: + vpshufd xmm0,xmm0,0 ; broadcast values + +; +; Process a set of 4 columns from the remaining rows. +; + +ProcessColumnLoopSmallKBy4: + cmp r10,4 + jb ProcessColumnLoopSmallKBySmallN + vmovd xmm1,DWORD PTR [rsi] + cmp r9d,2 + jb ComputeOutputSmallKBy4 + vpinsrd xmm1,xmm1,DWORD PTR [rsi+rdi],1 + je ComputeOutputSmallKBy4 + vpinsrd xmm1,xmm1,DWORD PTR [rsi+rdi*2],2 + +ComputeOutputSmallKBy4: + vpshufb xmm1,xmm1,XMMWORD PTR [MlasTranspose4x4BytesAvx] + vpxor xmm8,xmm8,xmm8 + VpdpbusdsXmmXmmXmm xmm8,xmm0,xmm1 + test r11,r11 ; ZeroMode? + jnz SkipAccumulateOutputSmallKBy4 + vpaddd xmm8,xmm8,XMMWORD PTR [r8] + +SkipAccumulateOutputSmallKBy4: + vmovdqu XMMWORD PTR [r8],xmm8 + add rsi,4 ; advance matrix B by 4 bytes + add r8,4*4 ; advance matrix C by 4 columns + sub r10,4 ; decrement CountN + jnz ProcessColumnLoopSmallKBy4 + jmp ExitKernel + +; +; Process the remaining 1 to 3 columns from the remaining rows. +; +; Single step through each of the columns to keep code size small for the +; uncommon path (typically the row count is a multiple of 4). +; + +ProcessColumnLoopSmallKBySmallN: + vpinsrb xmm1,xmm5,BYTE PTR [rsi],0 + cmp r9d,2 + jb ComputeOutputSmallKBySmallN + vpinsrb xmm1,xmm1,BYTE PTR [rsi+rdi],1 + je ComputeOutputSmallKBySmallN + vpinsrb xmm1,xmm1,BYTE PTR [rsi+rdi*2],2 + +ComputeOutputSmallKBySmallN: + vpxor xmm8,xmm8,xmm8 + VpdpbusdsXmmXmmXmm xmm8,xmm0,xmm1 + test r11,r11 ; ZeroMode? + jnz SkipAccumulateOutputSmallKBySmallN + vmovd xmm3,DWORD PTR [r8] + vpaddd xmm8,xmm8,xmm3 + +SkipAccumulateOutputSmallKBySmallN: + vmovd DWORD PTR [r8],xmm8 + inc rsi ; advance matrix B by 1 byte + add r8,4 ; advance matrix C by 1 column + dec r10 + jnz ProcessColumnLoopSmallKBySmallN + jmp ExitKernel + + NESTED_END MlasGemvU8S8KernelAvxVnni, _TEXT + + END \ No newline at end of file diff --git a/onnxruntime/core/mlas/lib/mlasi.h b/onnxruntime/core/mlas/lib/mlasi.h index abe18515e2..d5e4aef0ca 100644 --- a/onnxruntime/core/mlas/lib/mlasi.h +++ b/onnxruntime/core/mlas/lib/mlasi.h @@ -530,6 +530,8 @@ extern "C" { MLAS_GEMV_U8S8_KERNEL MlasGemvU8S8KernelAvx512Core; MLAS_GEMM_U8S8_KERNEL MlasGemmU8S8KernelAvx512Vnni; MLAS_GEMV_U8S8_KERNEL MlasGemvU8S8KernelAvx512Vnni; + MLAS_GEMM_U8S8_KERNEL MlasGemmU8S8KernelAvxVnni; + MLAS_GEMV_U8S8_KERNEL MlasGemvU8S8KernelAvxVnni; MLAS_GEMM_U8U8_KERNEL MlasGemmU8U8KernelAvx2; MLAS_GEMM_U8U8_KERNEL MlasGemmU8U8KernelAvx512Core; #endif diff --git a/onnxruntime/core/mlas/lib/platform.cpp b/onnxruntime/core/mlas/lib/platform.cpp index 270010df43..42c1fe4d18 100644 --- a/onnxruntime/core/mlas/lib/platform.cpp +++ b/onnxruntime/core/mlas/lib/platform.cpp @@ -221,6 +221,25 @@ Return Value: this->QLinearAddS8Kernel = MlasQLinearAddS8KernelAvx2; this->QLinearAddU8Kernel = MlasQLinearAddU8KernelAvx2; this->ComputeSumExpF32Kernel = MlasComputeSumExpF32KernelFma3; + + // + // Check if the processor supports AVXVNNI features. + // + + unsigned Cpuid7_1[4]; +#if defined(_WIN32) + __cpuidex((int*)Cpuid7_1, 7, 1); +#else + __cpuid_count(7, 1, Cpuid7_1[0], Cpuid7_1[1], Cpuid7_1[2], Cpuid7_1[3]); +#endif + + if ((Cpuid7_1[0] & 0x10) != 0) { + + this->GemmU8U8Operation = MlasGemmU8X8Operation; + this->GemmU8U8PackedOperation = MlasGemmU8X8PackedOperation; + this->GemmU8S8Kernel = MlasGemmU8S8KernelAvxVnni; + this->GemvU8S8Kernel = MlasGemvU8S8KernelAvxVnni; + } #if !defined(MLAS_AVX512F_UNSUPPORTED) diff --git a/onnxruntime/core/mlas/lib/x86_64/AssembleAvxVnni.h b/onnxruntime/core/mlas/lib/x86_64/AssembleAvxVnni.h new file mode 100644 index 0000000000..0d025ce483 --- /dev/null +++ b/onnxruntime/core/mlas/lib/x86_64/AssembleAvxVnni.h @@ -0,0 +1,178 @@ +/*++ + +Copyright (c) 2020 Intel Corporation. All rights reserved. + +Licensed under the MIT License. + +Module Name: + + AssembleAvxVnni.h + +Abstract: + + This module contains macros to build VNNI instructions for toolchains that + do not natively support this newer instruction set extension. + +--*/ + +// +// Map friendly register names to the encoded register index. +// + + .equ .LYmmIndex_ymm0, 0 + .equ .LYmmIndex_ymm1, 1 + .equ .LYmmIndex_ymm2, 2 + .equ .LYmmIndex_ymm3, 3 + .equ .LYmmIndex_ymm4, 4 + .equ .LYmmIndex_ymm5, 5 + .equ .LYmmIndex_ymm6, 6 + .equ .LYmmIndex_ymm7, 7 + .equ .LYmmIndex_ymm8, 8 + .equ .LYmmIndex_ymm9, 9 + .equ .LYmmIndex_ymm10, 10 + .equ .LYmmIndex_ymm11, 11 + .equ .LYmmIndex_ymm12, 12 + .equ .LYmmIndex_ymm13, 13 + .equ .LYmmIndex_ymm14, 14 + .equ .LYmmIndex_ymm15, 15 + + .equ .LXmmIndex_xmm0, 0 + .equ .LXmmIndex_xmm1, 1 + .equ .LXmmIndex_xmm2, 2 + .equ .LXmmIndex_xmm3, 3 + .equ .LXmmIndex_xmm4, 4 + .equ .LXmmIndex_xmm5, 5 + .equ .LXmmIndex_xmm6, 6 + .equ .LXmmIndex_xmm7, 7 + .equ .LXmmIndex_xmm8, 8 + .equ .LXmmIndex_xmm9, 9 + .equ .LXmmIndex_xmm10, 10 + .equ .LXmmIndex_xmm11, 11 + .equ .LXmmIndex_xmm12, 12 + .equ .LXmmIndex_xmm13, 13 + .equ .LXmmIndex_xmm14, 14 + .equ .LXmmIndex_xmm15, 15 + +/*++ + +Macro Description: + + This macro builds a VNNI instruction of the form: + + instr ymm1,ymm2,ymm3 + +Arguments: + + Opcode - Specifies the opcode for the VNNI instruction. + + DestReg - Specifies the destination register. + + Src1Reg - Specifies the first source register. + + Src2Reg - Specifies the second source register. + +--*/ + + .macro VnniYmmYmmYmm Opcode, 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, 0x05 # "66" 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 VpdpbusdYmmYmmYmm DestReg, Src1Reg, Src2Reg + + VnniYmmYmmYmm 0x50, \DestReg\(), \Src1Reg\(), \Src2Reg\() + + .endm + + .macro VpdpbusdsYmmYmmYmm DestReg, Src1Reg, Src2Reg + + VnniYmmYmmYmm 0x51, \DestReg\(), \Src1Reg\(), \Src2Reg\() + + .endm + + .macro VpdpwssdYmmYmmYmm DestReg, Src1Reg, Src2Reg + + VnniYmmYmmYmm 0x52, \DestReg\(), \Src1Reg\(), \Src2Reg\() + + .endm + + .macro VpdpwssdsYmmYmmYmm DestReg, Src1Reg, Src2Reg + + VnniYmmYmmYmm 0x53, \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. + + DestReg - Specifies the destination register. + + Src1Reg - Specifies the first source register. + + Src2Reg - Specifies the second source register. + +--*/ + + .macro VnniXmmXmmXmm Opcode, DestReg, Src1Reg, Src2Reg + + .set Payload0, 0x02 # "0F 38" prefix + .set Payload0, Payload0 + ((((.LXmmIndex_\DestReg\() >> 3) & 1) ^ 1) << 7) + .set Payload0, Payload0 + (1 << 6) + .set Payload0, Payload0 + ((((.LXmmIndex_\Src2Reg\() >> 3) & 1) ^ 1) << 5) + + .set Payload1, 0x05 # "66" prefix + .set Payload1, Payload1 + (((.LXmmIndex_\Src1Reg\() & 15) ^ 15) << 3) + + .set ModRMByte, 0xC0 # register form + .set ModRMByte, ModRMByte + ((.LXmmIndex_\DestReg\() & 7) << 3) + .set ModRMByte, ModRMByte + (.LXmmIndex_\Src2Reg\() & 7) + + .byte 0xC4, Payload0, Payload1, \Opcode\(), ModRMByte + + .endm + + .macro VpdpbusdXmmXmmXmm DestReg, Src1Reg, Src2Reg + + VnniXmmXmmXmm 0x50, \DestReg\(), \Src1Reg\(), \Src2Reg\() + + .endm + + .macro VpdpbusdsXmmXmmXmm DestReg, Src1Reg, Src2Reg + + VnniXmmXmmXmm 0x51, \DestReg\(), \Src1Reg\(), \Src2Reg\() + + .endm + + .macro VpdpwssdXmmXmmXmm DestReg, Src1Reg, Src2Reg + + VnniXmmXmmXmm 0x52, \DestReg\(), \Src1Reg\(), \Src2Reg\() + + .endm + + .macro VpdpwssdsXmmXmmXmm DestReg, Src1Reg, Src2Reg + + VnniXmmXmmXmm 0x53, \DestReg\(), \Src1Reg\(), \Src2Reg\() + + .endm diff --git a/onnxruntime/core/mlas/lib/x86_64/QgemmU8S8KernelAvxVnni.S b/onnxruntime/core/mlas/lib/x86_64/QgemmU8S8KernelAvxVnni.S new file mode 100644 index 0000000000..29c4b006b9 --- /dev/null +++ b/onnxruntime/core/mlas/lib/x86_64/QgemmU8S8KernelAvxVnni.S @@ -0,0 +1,268 @@ +/*++ + +Copyright (c) 2020 Intel Corporation. All rights reserved. + +Licensed under the MIT License. + +Module Name: + + QgemmU8S8KernelAvxVnni.s + +Abstract: + + This module implements the kernels for the quantized integer matrix/matrix + multiply operation (QGEMM). + + This implementation uses AVXVNNI instructions. + +--*/ + +#include "asmmacro.h" +#include "QgemmU8X8KernelAvx2Common.h" +#include "AssembleAvxVnni.h" + + .intel_syntax noprefix + +/*++ +Macro Description: + + This macro generates code to multiply and accumulator a single row of the + output block. + +Arguments: + + ColumnCount - Supplies the number of columns to produce. + + Vec1Reg - Supplies the high block accumulator register (when ColumnCount + is 16). + + Vec2Reg - Supplies the low block accumulator register. + +Implicit Arguments: + + ymm0 - Supplies the first vector loaded from matrix B. + + ymm1 - Supplies the second vector loaded from matrix B (when ColumnCount + is 16). + + ymm2 - Supplies the broadcast value loaded from matrix A. + +--*/ + + .macro MultiplyAccumulateRow ColumnCount, Vec1Reg, Vec2Reg + +.if \ColumnCount\() == 16 + VpdpbusdsYmmYmmYmm \Vec1Reg\(),ymm2,ymm0 + VpdpbusdsYmmYmmYmm \Vec2Reg\(),ymm2,ymm1 +.else + VpdpbusdsYmmYmmYmm \Vec2Reg\(),ymm2,ymm0 +.endif + + .endm + +/*++ + +Macro Description: + + This macro generates code to multiply and accumulate each row of the output + block. + +Arguments: + + ColumnCount - Supplies the number of columns to produce. + + RowCount - Supplies the number of rows to produce. + + VectorOffset - Supplies the byte offset from matrix B to fetch elements. + + BroadcastOffset - Supplies the byte offset from matrix A to fetch elements. + +Implicit Arguments: + + rbx - Supplies the address into the matrix A data plus 3 rows. + + rcx - Supplies the address into the matrix A data. + + rsi - Supplies the address into the matrix B data. + + rcx - Supplies the length in bytes of a row from matrix A. + + ymm4-ymm15 - Supplies the block accumulators. + +--*/ + + .macro ComputeBlock ColumnCount, RowCount, VectorOffset, BroadcastOffset + + 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, "MultiplyAccumulateRow \ColumnCount\(), ymm4, ymm5" + EmitIfCountGE \RowCount\(), 2, "vpbroadcastd ymm2,DWORD PTR [rdi+rcx+\BroadcastOffset\()]" + EmitIfCountGE \RowCount\(), 2, "MultiplyAccumulateRow \ColumnCount\(), ymm6, ymm7" + EmitIfCountGE \RowCount\(), 3, "vpbroadcastd ymm2,DWORD PTR [rdi+rcx*2+\BroadcastOffset\()]" + EmitIfCountGE \RowCount\(), 3, "MultiplyAccumulateRow \ColumnCount\(), ymm8, ymm9" + EmitIfCountGE \RowCount\(), 4, "vpbroadcastd ymm2,DWORD PTR [rbx+\BroadcastOffset\()]" + EmitIfCountGE \RowCount\(), 4, "MultiplyAccumulateRow \ColumnCount\(), ymm10, ymm11" + EmitIfCountGE \RowCount\(), 5, "vpbroadcastd ymm2,DWORD PTR [rbx+r9+\BroadcastOffset\()]" + EmitIfCountGE \RowCount\(), 5, "MultiplyAccumulateRow \ColumnCount\(), ymm12, ymm13" + EmitIfCountGE \RowCount\(), 6, "vpbroadcastd ymm2,DWORD PTR [rbx+r9*2+\BroadcastOffset\()]" + EmitIfCountGE \RowCount\(), 6, "MultiplyAccumulateRow \ColumnCount\(), ymm14, ymm15" + + .endm + +/*++ + +Macro Description: + + This macro generates code to execute the block compute macro multiple + times and advancing the matrix A and matrix B data pointers. + +Arguments: + + ColumnCount - Supplies the number of columns to produce. + + RowCount - Supplies the number of rows to produce. + +Implicit Arguments: + + rbx - Supplies the address into the matrix A data plus 3 rows. + + rdi - Supplies the address into the matrix A data. + + rsi - Supplies the address into the matrix B data. + + rcx - Supplies the length in bytes of a row from matrix A. + + ymm4-ymm15 - Supplies the block accumulators. + +--*/ + + .macro ComputeBlockLoop ColumnCount, RowCount + + mov rbp,rcx # reload row length remaining + +.LComputeBlockBy1Loop\@: + ComputeBlock \ColumnCount\(), \RowCount\(), 0, 0 + add rdi,4 # advance matrix A by 1 quad +.if \RowCount\() > 3 + add rbx,4 # advance matrix A plus 3 rows by 1 quad +.endif + add rsi,64 # advance matrix B + sub rbp,4 + jnz .LComputeBlockBy1Loop\@ + + .endm + +/*++ + +Routine Description: + + This routine is an inner kernel to compute matrix multiplication for a + set of rows. + +Arguments: + + A (rdi) - Supplies the address of matrix A. The matrix data has been packed + using MlasGemmU8S8CopyPackAAvx2. + + B (rsi) - Supplies the address of matrix B. The matrix data has been packed + using MlasGemmU8S8CopyPackBAvx2. + + C (rdx) - Supplies the address of matrix C. + + PackedCountK (rcx) - Supplies the number of packed columns from matrix A + and the number of packed rows from matrix B to iterate over. + + CountM (r8) - Supplies the maximum number of rows that can be processed for + matrix A and matrix C. The actual number of rows handled for this + invocation depends on the kernel implementation. + + CountN (r9) - Supplies the number of columns from matrix B and matrix C to + iterate over. + + ldc - Supplies the first dimension of matrix C. + + RowSumBuffer - Supplies the sum of each row from matrix A multiplied by the + zero point offset of matrix B. These values are accumulated into every + row of matrix C. + + ColumnSumBuffer - Supplies the sum of each column from matrix B multiplied + by the zero point offset of matrix A. These values are accumulated into + every column of matrix C. + + DepthValue - Supplies the value CountK multiplied by the zero point offset + of matrix A multplied by the zero point offset of matrix B. This value + is accumulated into every element of matrix C. + + ZeroMode - Supplies true if the output matrix must be zero initialized, + else false if the output matrix is accumulated into. + +Return Value: + + Returns the number of rows handled. + +--*/ + + .globl C_UNDERSCORE(MlasGemmU8S8KernelAvxVnni) +C_UNDERSCORE(MlasGemmU8S8KernelAvxVnni): + + push rbp + push rbx + push r12 + push r13 + + mov rax,.LGemmU8X8KernelFrame_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,rdi + mov r12,.LGemmU8X8KernelFrame_RowSumBuffer[rsp] + mov r13,.LGemmU8X8KernelFrame_ColumnSumBuffer[rsp] + +// +// Process CountM rows of the matrices. +// + + cmp r8,5 + ja .LProcessCountM6 + je .LProcessCountM5 + cmp r8,3 + ja .LProcessCountM4 + je .LProcessCountM3 + cmp r8,1 + je .LProcessCountM1 + +.LProcessCountM2: + ProcessCountM 2 + +.LProcessCountM4: + ProcessCountM 4 + +.LProcessCountM6: + mov r8d,6 # return 6 rows handled + ProcessCountM 6, Fallthrough + +// +// Restore non-volatile registers and return. +// + +.LExitKernel: + mov eax,r8d + vzeroupper + + pop r13 + pop r12 + pop rbx + pop rbp + ret + +.LProcessCountM1: + ProcessCountM 1 + +.LProcessCountM3: + ProcessCountM 3 + +.LProcessCountM5: + ProcessCountM 5 + + .end diff --git a/onnxruntime/core/mlas/lib/x86_64/QgemvU8S8KernelAvxVnni.S b/onnxruntime/core/mlas/lib/x86_64/QgemvU8S8KernelAvxVnni.S new file mode 100644 index 0000000000..fb544ce34d --- /dev/null +++ b/onnxruntime/core/mlas/lib/x86_64/QgemvU8S8KernelAvxVnni.S @@ -0,0 +1,344 @@ +/*++ + +Copyright (c) 2020 Intel Corporation. All rights reserved. + +Licensed under the MIT License. + +Module Name: + + QgemvU8S8KernelAvxVnni.s + +Abstract: + + This module implements the kernels for the quantized integer matrix/vector + multiply operation (QGEMV). + + This implementation uses AVXVNNI instructions. + +--*/ + +#include "asmmacro.h" +#include "AssembleAvxVnni.h" + + .intel_syntax noprefix + +// +// Stack frame layout for the U8S8 kernel. +// + + .equ .LGemvU8S8KernelFrame_mask, -8 + .equ .LGemvU8S8KernelFrame_SavedRbx, 0 + .equ .LGemvU8S8KernelFrame_SavedRbp, 8 + .equ .LGemvU8S8KernelFrame_ReturnAddress, 16 + + .text + +/*++ + +Routine Description: + + This routine is an inner kernel to compute matrix/vector multiplication. + +Arguments: + + A (rdi) - Supplies the address of vector A. + + B (rsi) - Supplies the address of matrix B. + + C (rdx) - Supplies the address of matrix C. + + CountK (rcx) - Supplies the number of columns from vector A and the number + of rows from matrix B to iterate over. + + CountN (r8) - Supplies the number of columns from matrix B and matrix C to + iterate over. + + ldb (r9) - Supplies the first dimension of matrix B. + +Return Value: + + None. + +--*/ + + .globl C_UNDERSCORE(MlasGemvU8S8KernelAvxVnni) +C_UNDERSCORE(MlasGemvU8S8KernelAvxVnni): + + push rbp + push rbx + + mov r10,rdx + mov r11,rsp # set ZeroMode to any non-zero value + +// +// Process 4 rows of matrix B in a loop. +// + + sub rcx,4 + jb .LProcessRemainingRows + +.LProcessRowLoop4: + mov rdx,rsi # reload matrix B + lea rsi,[rsi+r9*4] # advance matrix B by 4 rows + mov rbx,r10 # reload matrix C + mov rbp,r8 # reload CountN + vpbroadcastd ymm0,DWORD PTR [rdi] + add rdi,4 # advance matrix A by 4 bytes + +// +// Process sets of 32 columns from the 4 rows in a loop. +// +// Some permute operations are deferred until the final store of the 4x32 block +// as these permutes are expensive. +// + +.LProcessColumnLoop4By32: + cmp rbp,32 + jb .LProcessColumnLoop4By8 + lea rax,[rdx+r9*2] # compute matrix B plus 2 rows + vmovdqu ymm2,YMMWORD PTR [rdx] + vmovdqu ymm3,YMMWORD PTR [rdx+r9] + vmovdqu ymm4,YMMWORD PTR [rax] + vmovdqu ymm5,YMMWORD PTR [rax+r9] + vpunpcklbw ymm1,ymm2,ymm3 # interleave row data bytes + vpunpckhbw ymm2,ymm2,ymm3 + vpxor ymm7,ymm7,ymm7 + vpunpcklbw ymm3,ymm4,ymm5 + vpunpckhbw ymm4,ymm4,ymm5 + vpxor ymm8,ymm8,ymm8 + vpunpcklwd ymm5,ymm1,ymm3 # interleave row data words + vpunpckhwd ymm1,ymm1,ymm3 + vpxor ymm9,ymm9,ymm9 + vpunpcklwd ymm3,ymm2,ymm4 + vpunpckhwd ymm2,ymm2,ymm4 + vpxor ymm10,ymm10,ymm10 + VpdpbusdsYmmYmmYmm ymm7,ymm0,ymm5 + VpdpbusdsYmmYmmYmm ymm8,ymm0,ymm1 + VpdpbusdsYmmYmmYmm ymm9,ymm0,ymm3 + VpdpbusdsYmmYmmYmm ymm10,ymm0,ymm2 + test r11,r11 # ZeroMode? + jnz .LSkipAccumulateOutput4By32 + vpaddd ymm7,ymm7,YMMWORD PTR [rbx] + vpaddd ymm8,ymm8,YMMWORD PTR [rbx+32] + vpaddd ymm9,ymm9,YMMWORD PTR [rbx+64] + vpaddd ymm10,ymm10,YMMWORD PTR [rbx+96] + +.LSkipAccumulateOutput4By32: + cmp rcx,4 # final 4x32 block? + jae .LStoreOutput4By32 + vperm2i128 ymm4,ymm7,ymm8,0x31 # interleave vector results + vperm2i128 ymm7,ymm7,ymm8,0x20 + vperm2i128 ymm8,ymm9,ymm10,0x20 + vperm2i128 ymm10,ymm9,ymm10,0x31 + vmovaps ymm3,ymm4 + +.LStoreOutput4By32: + vmovdqu YMMWORD PTR [rbx],ymm7 + vmovdqu YMMWORD PTR [rbx+32],ymm8 + vmovdqu YMMWORD PTR [rbx+64],ymm9 + vmovdqu YMMWORD PTR [rbx+96],ymm10 + add rdx,32 # advance matrix B by 32 bytes + add rbx,32*4 # advance matrix C by 32 columns + sub rbp,32 # decrement CountN + jnz .LProcessColumnLoop4By32 + +.LAdvanceRowLoop4: + xor r11,r11 # clear ZeroMode + sub rcx,4 # decrement CountK + jae .LProcessRowLoop4 + +.LProcessRemainingRows: + add rcx,4 # correct for over-subtract above + jnz .LProcessRemainingSmallK + +// +// Restore non-volatile registers and return. +// + +.LExitKernel: + vzeroupper + + pop rbx + pop rbp + ret + +// +// Process sets of 8 columns from the 4 rows in a loop. +// + +.LProcessColumnLoop4By8: + cmp ebp,8 + jb .LProcessColumn4By4 + lea rax,[rdx+r9*2] # compute matrix B plus 2 rows + vmovq xmm2,QWORD PTR [rdx] + vmovq xmm3,QWORD PTR [rdx+r9] + vmovq xmm4,QWORD PTR [rax] + vmovq xmm5,QWORD PTR [rax+r9] + vpunpcklbw xmm2,xmm2,xmm3 # interleave row data bytes + vpunpcklbw xmm4,xmm4,xmm5 + vpunpcklwd xmm1,xmm2,xmm4 # interleave row data words + vpunpckhwd xmm2,xmm2,xmm4 + vinserti128 ymm1,ymm1,xmm2,1 # concatenate vector + vpxor ymm8,ymm8,ymm8 + VpdpbusdsYmmYmmYmm ymm8,ymm0,ymm1 + test r11,r11 # ZeroMode? + jnz .LSkipAccumulateOutput4By8 + vpaddd ymm8,ymm8,YMMWORD PTR [rbx] + +.LSkipAccumulateOutput4By8: + vmovdqu YMMWORD PTR [rbx],ymm8 + add rdx,8 # advance matrix B by 8 bytes + add rbx,8*4 # advance matrix C by 8 columns + sub ebp,8 # decrement CountN + jnz .LProcessColumnLoop4By8 + jmp .LAdvanceRowLoop4 + +// +// Process a set of 4 columns from the 4 rows. +// + +.LProcessColumn4By4: + test ebp,4 # (CountN & 4) != 0? + jz .LProcessColumn4BySmallN + lea rax,[rdx+r9*2] # compute matrix B plus 2 rows + vmovd xmm1,DWORD PTR [rdx] + vpinsrd xmm1,xmm1,DWORD PTR [rdx+r9],1 + vpinsrd xmm1,xmm1,DWORD PTR [rax],2 + vpinsrd xmm1,xmm1,DWORD PTR [rax+r9],3 + vpshufb xmm1,xmm1,XMMWORD PTR C_UNDERSCORE(MlasTranspose4x4BytesAvx)[rip] + vpxor xmm8,xmm8,xmm8 + VpdpbusdsXmmXmmXmm xmm8,xmm0,xmm1 + test r11,r11 # ZeroMode? + jnz .LSkipAccumulateOutput4By4 + vpaddd xmm8,xmm8,XMMWORD PTR [rbx] + +.LSkipAccumulateOutput4By4: + vmovdqu XMMWORD PTR [rbx],xmm8 + and ebp,3 # (CountN & 3) != 0? + jz .LAdvanceRowLoop4 + add rdx,4 # advance matrix B by 4 bytes + add rbx,4*4 # advance matrix C by 4 columns + +// +// Process the remaining 1 to 3 columns from the 4 rows. +// + +.LProcessColumn4BySmallN: + mov DWORD PTR .LGemvU8S8KernelFrame_mask[rsp],ebp + vbroadcastss xmm2,DWORD PTR .LGemvU8S8KernelFrame_mask[rsp] + vpcmpgtd xmm2,xmm2,XMMWORD PTR C_UNDERSCORE(MlasMaskMoveAvx)[rip] + vpxor xmm1,xmm1,xmm1 + lea rax,[rdx+r9*2] # compute matrix B plus 2 rows + cmp ebp,2 # (CountN & 2) != 0? + jb .LProcessColumn4By1 + vpinsrw xmm1,xmm1,WORD PTR [rdx],0 + vpinsrw xmm1,xmm1,WORD PTR [rdx+r9],2 + vpinsrw xmm1,xmm1,WORD PTR [rax],4 + vpinsrw xmm1,xmm1,WORD PTR [rax+r9],6 + je .LComputeOutput4BySmallN + vpinsrb xmm1,xmm1,BYTE PTR [rdx+2],2 + vpinsrb xmm1,xmm1,BYTE PTR [rdx+r9+2],6 + vpinsrb xmm1,xmm1,BYTE PTR [rax+2],10 + vpinsrb xmm1,xmm1,BYTE PTR [rax+r9+2],14 + jmp .LComputeOutput4BySmallN + +.LProcessColumn4By1: + vpinsrb xmm1,xmm1,BYTE PTR [rdx],0 + vpinsrb xmm1,xmm1,BYTE PTR [rdx+r9],4 + vpinsrb xmm1,xmm1,BYTE PTR [rax],8 + vpinsrb xmm1,xmm1,BYTE PTR [rax+r9],12 + +.LComputeOutput4BySmallN: + vpshufb xmm1,xmm1,XMMWORD PTR C_UNDERSCORE(MlasTranspose4x4BytesAvx)[rip] + vpxor xmm8,xmm8,xmm8 + VpdpbusdsXmmXmmXmm xmm8,xmm0,xmm1 + test r11,r11 # ZeroMode? + jnz .LStoreOutput4BySmallN + vpmaskmovd xmm3,xmm2,XMMWORD PTR [rbx] + vpaddd xmm8,xmm8,xmm3 + +.LStoreOutput4BySmallN: + vpmaskmovd XMMWORD PTR [rbx],xmm2,xmm8 + jmp .LAdvanceRowLoop4 + +// +// Broadcast the remaining 1 to 3 values from vector A. +// + +.LProcessRemainingSmallK: + vpxor xmm5,xmm5,xmm5 # keep zero vector for vpinsrb/vpinsrw + cmp ecx,2 + jb .LLoadVectorASingleRemainingByte + vpinsrw xmm0,xmm5,WORD PTR [rdi],0 + je .LBroadcastVectorARemainingBytes + vpinsrb xmm0,xmm0,BYTE PTR [rdi+2],2 + jmp .LBroadcastVectorARemainingBytes + +.LLoadVectorASingleRemainingByte: + vpinsrb xmm0,xmm5,BYTE PTR [rdi],0 + +.LBroadcastVectorARemainingBytes: + vpshufd xmm0,xmm0,0 # broadcast values + +// +// Process a set of 4 columns from the remaining rows. +// + +.LProcessColumnLoopSmallKBy4: + cmp r8d,4 + jb .LProcessColumnLoopSmallKBySmallN + vmovd xmm1,DWORD PTR [rsi] + cmp ecx,2 + jb .LComputeOutputSmallKBy4 + vpinsrd xmm1,xmm1,DWORD PTR [rsi+r9],1 + je .LComputeOutputSmallKBy4 + vpinsrd xmm1,xmm1,DWORD PTR [rsi+r9*2],2 + +.LComputeOutputSmallKBy4: + vpshufb xmm1,xmm1,XMMWORD PTR C_UNDERSCORE(MlasTranspose4x4BytesAvx)[rip] + vpxor xmm8,xmm8,xmm8 + VpdpbusdsXmmXmmXmm xmm8,xmm0,xmm1 + test r11,r11 # ZeroMode? + jnz .LSkipAccumulateOutputSmallKBy4 + vpaddd xmm8,xmm8,XMMWORD PTR [r10] + +.LSkipAccumulateOutputSmallKBy4: + vmovdqu XMMWORD PTR [r10],xmm8 + add rsi,4 # advance matrix B by 4 bytes + add r10,4*4 # advance matrix C by 4 columns + sub r8d,4 # decrement CountN + jnz .LProcessColumnLoopSmallKBy4 + jmp .LExitKernel + +// +// Process the remaining 1 to 3 columns from the remaining rows. +// +// Single step through each of the columns to keep code size small for the +// uncommon path (typically the row count is a multiple of 4). +// + +.LProcessColumnLoopSmallKBySmallN: + vpinsrb xmm1,xmm5,BYTE PTR [rsi],0 + cmp ecx,2 + jb .LComputeOutputSmallKBySmallN + vpinsrb xmm1,xmm1,BYTE PTR [rsi+r9],1 + je .LComputeOutputSmallKBySmallN + vpinsrb xmm1,xmm1,BYTE PTR [rsi+r9*2],2 + +.LComputeOutputSmallKBySmallN: + vpxor xmm8,xmm8,xmm8 + VpdpbusdsXmmXmmXmm xmm8,xmm0,xmm1 + test r11,r11 # ZeroMode? + jnz .LSkipAccumulateOutputSmallKBySmallN + vmovd xmm3,DWORD PTR [r10] + vpaddd xmm8,xmm8,xmm3 + +.LSkipAccumulateOutputSmallKBySmallN: + vmovd DWORD PTR [r10],xmm8 + inc rsi # advance matrix B by 1 byte + add r10,4 # advance matrix C by 1 column + dec r8 + jnz .LProcessColumnLoopSmallKBySmallN + jmp .LExitKernel + + .end