diff --git a/cmake/onnxruntime_mlas.cmake b/cmake/onnxruntime_mlas.cmake index 9aa7c8adef..fee9648460 100644 --- a/cmake/onnxruntime_mlas.cmake +++ b/cmake/onnxruntime_mlas.cmake @@ -161,13 +161,18 @@ else() endif() if(ARM) + enable_language(ASM) + + set(CMAKE_ASM_FLAGS "${CMAKE_ASM_FLAGS} -mfpu=neon") set(CMAKE_CXX_FLAGS "${CMAKE_CXX_FLAGS} -mfpu=neon") set(mlas_platform_srcs + ${ONNXRUNTIME_ROOT}/core/mlas/lib/aarch32/QgemmU8X8KernelNeon.S ${ONNXRUNTIME_ROOT}/core/mlas/lib/arm/sgemmc.cpp ) elseif(ARM64) enable_language(ASM) + set(mlas_platform_srcs ${ONNXRUNTIME_ROOT}/core/mlas/lib/aarch64/QgemmU8X8KernelNeon.S ${ONNXRUNTIME_ROOT}/core/mlas/lib/aarch64/SgemmKernelNeon.S @@ -179,6 +184,7 @@ else() ) elseif(X86) enable_language(ASM) + set(mlas_platform_srcs_sse2 ${ONNXRUNTIME_ROOT}/core/mlas/lib/x86/SgemmKernelSse2.S ) diff --git a/onnxruntime/core/mlas/inc/mlas.h b/onnxruntime/core/mlas/inc/mlas.h index 64abcd47a0..5e44308ded 100644 --- a/onnxruntime/core/mlas/inc/mlas.h +++ b/onnxruntime/core/mlas/inc/mlas.h @@ -61,7 +61,7 @@ Abstract: #define MLAS_SUPPORTS_GEMM_DOUBLE #endif -#if defined(MLAS_TARGET_AMD64_IX86) || defined(MLAS_TARGET_ARM64) +#if defined(MLAS_TARGET_AMD64_IX86) || defined(MLAS_TARGET_ARM64) || (defined(MLAS_TARGET_ARM) && !defined(_MSC_VER)) #define MLAS_SUPPORTS_GEMM_U8X8 #endif @@ -69,7 +69,7 @@ Abstract: #define MLAS_SUPPORTS_GEMM_U8X8_AND_REQUANTIZE_OUTPUT #endif -#if defined(MLAS_TARGET_AMD64) || defined(MLAS_TARGET_ARM64) +#if defined(MLAS_TARGET_AMD64) || defined(MLAS_TARGET_ARM64) || (defined(MLAS_TARGET_ARM) && !defined(_MSC_VER)) #define MLAS_SUPPORTS_PACKED_GEMM_U8X8 #endif diff --git a/onnxruntime/core/mlas/lib/aarch32/QgemmU8X8KernelNeon.S b/onnxruntime/core/mlas/lib/aarch32/QgemmU8X8KernelNeon.S new file mode 100644 index 0000000000..01799b196c --- /dev/null +++ b/onnxruntime/core/mlas/lib/aarch32/QgemmU8X8KernelNeon.S @@ -0,0 +1,610 @@ +/*++ + +Copyright (c) Microsoft Corporation. All rights reserved. + +Licensed under the MIT License. + +Module Name: + + QgemmU8X8KernelNeon.s + +Abstract: + + This module implements the kernels for the quantized integer matrix/matrix + multiply operation (QGEMM). + +--*/ + +#include "asmmacro.h" + + .syntax unified + .arch armv7-a + .thumb + +// +// Stack frame layout for the U8X8 kernel. +// + + .equ .LGemmU8X8KernelFrame_SavedGeneralRegisters, (6 * 4) + .equ .LGemmU8X8KernelFrame_SavedNeonRegisters, (8 * 8) + .equ .LGemmU8X8KernelFrame_SavedRegisters, .LGemmU8X8KernelFrame_SavedGeneralRegisters + .LGemmU8X8KernelFrame_SavedNeonRegisters + .equ .LGemmU8X8KernelFrame_CountM, 0 + .LGemmU8X8KernelFrame_SavedRegisters + .equ .LGemmU8X8KernelFrame_CountN, 4 + .LGemmU8X8KernelFrame_SavedRegisters + .equ .LGemmU8X8KernelFrame_ldc, 8 + .LGemmU8X8KernelFrame_SavedRegisters + .equ .LGemmU8X8KernelFrame_RowSumBuffer, 12 + .LGemmU8X8KernelFrame_SavedRegisters + .equ .LGemmU8X8KernelFrame_ColumnSumBuffer, 16 + .LGemmU8X8KernelFrame_SavedRegisters + .equ .LGemmU8X8KernelFrame_DepthValue, 20 + .LGemmU8X8KernelFrame_SavedRegisters + .equ .LGemmU8X8KernelFrame_ZeroMode, 24 + .LGemmU8X8KernelFrame_SavedRegisters + + .text + +/*++ + +Routine Description: + + This routine is an inner kernel to compute matrix multiplication for a + set of rows. + +Arguments: + + A (r0) - Supplies the address of matrix A. The matrix data has been packed + using MlasGemmU8X8CopyPackANeon. + + B (r1) - Supplies the address of matrix B. The matrix data has been packed + using MlasGemmU8X8CopyPackBNeon. + + C (r2) - Supplies the address of matrix C. + + PackedCountK (r3) - 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 + 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. + +--*/ + + FUNCTION_ENTRY MlasGemmU8X8KernelNeon + +// +// Register usage: +// +// q0-q1 (d0-d3) matrix B data +// q2-q3 (d4-d7) matrix A data +// q4 (d8-d9) packed matrix B data +// q5 (d10-d11) RowSumBufferData + DepthValue +// q6-q7 (d12-d15) ColumnSumBuffer data +// q8-q15 accumulators[4][2] +// + + push {r4,r5,r6,r7,r8,r10} + vpush {d8-d15} + ldr r4,[sp,#.LGemmU8X8KernelFrame_CountM] + ldr r5,[sp,#.LGemmU8X8KernelFrame_ZeroMode] + ldr r7,[sp,#.LGemmU8X8KernelFrame_RowSumBuffer] + ldr r8,[sp,#.LGemmU8X8KernelFrame_ColumnSumBuffer] + vldr s0,[sp,#.LGemmU8X8KernelFrame_DepthValue] + ldr r10,[sp,#.LGemmU8X8KernelFrame_ldc] + ldr r12,[sp,#.LGemmU8X8KernelFrame_CountN] + vdup.32 q5,d0[0] // broadcast DepthValue + vld1.32 {d12-d13},[r7] + mov r6,r0 + mov r7,r3 + vadd.u32 q5,q5,q6 // add row fixups and DepthValue + cmp r4,#1 // CountM == 1? + beq .LGemmU8X8.M1.ProcessNextColumnLoop + cmp r4,#4 // CountM < 4? + blo .LGemmU8X8.M2.ProcessNextColumnLoop + +// +// Process 4 rows of the matrices. +// + +.LGemmU8X8.M4.ProcessNextColumnLoop: + vldr d0,[r1] // load packed B0 + mov r0,r6 // reload matrix A + vld1.32 {d12-d15},[r8]! // load ColumnSumBuffer + mov r3,r7 // reload PackedCountK + vmovl.u8 q0,d0 + vdup.32 q9,d10[0] + vdup.32 q11,d10[1] + vdup.32 q13,d11[0] + vdup.32 q15,d11[1] + vldr d8,[r0] // load first packed A0 + vadd.u32 q8,q9,q6 + vadd.u32 q9,q9,q7 + vadd.u32 q10,q11,q6 + vadd.u32 q11,q11,q7 + vldr d9,[r0,#8] // load first packed A1 + vadd.u32 q12,q13,q6 + vadd.u32 q13,q13,q7 + vadd.u32 q14,q15,q6 + vadd.u32 q15,q15,q7 + +.LGemmU8X8.M4.ComputeBlockLoop: + vmovl.u8 q2,d8 + add r0,#16 + vmovl.u8 q3,d9 + vldr d2,[r1,#8] // load packed B1 + vmlal.u16 q8,d0,d4[0] + vmlal.u16 q9,d1,d4[0] + vmlal.u16 q10,d0,d5[0] + vmlal.u16 q11,d1,d5[0] + vmovl.u8 q1,d2 + vmlal.u16 q12,d0,d6[0] + vmlal.u16 q13,d1,d6[0] + vmlal.u16 q14,d0,d7[0] + vmlal.u16 q15,d1,d7[0] + vldr d0,[r1,#16] // load packed B2 + vmlal.u16 q8,d2,d4[1] + vmlal.u16 q9,d3,d4[1] + vmlal.u16 q10,d2,d5[1] + vmlal.u16 q11,d3,d5[1] + vmovl.u8 q0,d0 + vmlal.u16 q12,d2,d6[1] + vmlal.u16 q13,d3,d6[1] + vmlal.u16 q14,d2,d7[1] + vmlal.u16 q15,d3,d7[1] + vldr d2,[r1,#24] // load packed B3 + add r1,#32 + subs r3,#1 + beq .LGemmU8X8.M4.ComputeBlockLoopFinish + vmlal.u16 q8,d0,d4[2] + vmlal.u16 q9,d1,d4[2] + vmlal.u16 q10,d0,d5[2] + vmlal.u16 q11,d1,d5[2] + vmovl.u8 q1,d2 + vldr d8,[r0] // load next packed A0 + vmlal.u16 q12,d0,d6[2] + vmlal.u16 q13,d1,d6[2] + vmlal.u16 q14,d0,d7[2] + vmlal.u16 q15,d1,d7[2] + vldr d0,[r1] // load packed B0 + vmlal.u16 q8,d2,d4[3] + vmlal.u16 q9,d3,d4[3] + vmlal.u16 q10,d2,d5[3] + vmlal.u16 q11,d3,d5[3] + vmovl.u8 q0,d0 + vldr d9,[r0,#8] // load next packed A1 + vmlal.u16 q12,d2,d6[3] + vmlal.u16 q13,d3,d6[3] + vmlal.u16 q14,d2,d7[3] + vmlal.u16 q15,d3,d7[3] + b .LGemmU8X8.M4.ComputeBlockLoop + +.LGemmU8X8.M4.ComputeBlockLoopFinish: + vmlal.u16 q8,d0,d4[2] // finish computing tail vectors + vmlal.u16 q9,d1,d4[2] + add r0,r2,r10,lsl #2 // compute output row 2 + vmlal.u16 q10,d0,d5[2] + vmlal.u16 q11,d1,d5[2] + vmovl.u8 q1,d2 + vmlal.u16 q12,d0,d6[2] + vmlal.u16 q13,d1,d6[2] + vmlal.u16 q14,d0,d7[2] + vmlal.u16 q15,d1,d7[2] + add r3,r0,r10,lsl #2 // compute output row 3 + vmlal.u16 q8,d2,d4[3] + vmlal.u16 q9,d3,d4[3] + vmlal.u16 q10,d2,d5[3] + vmlal.u16 q11,d3,d5[3] + vmlal.u16 q12,d2,d6[3] + vmlal.u16 q13,d3,d6[3] + add r4,r3,r10,lsl #2 // compute output row 4 + vmlal.u16 q14,d2,d7[3] + vmlal.u16 q15,d3,d7[3] + subs r12,#8 // adjust CountN remaining + blo .LGemmU8X8.M4.StoreOutputPartial + cbnz r5,.LGemmU8X8.M4.SkipAccumulateOutput + vld1.32 {d0-d3},[r2] + vld1.32 {d4-d7},[r0] + vadd.u32 q8,q8,q0 + vadd.u32 q9,q9,q1 + vld1.32 {d0-d3},[r3] + vadd.u32 q10,q10,q2 + vadd.u32 q11,q11,q3 + vld1.32 {d4-d7},[r4] + vadd.u32 q12,q12,q0 + vadd.u32 q13,q13,q1 + vadd.u32 q14,q14,q2 + vadd.u32 q15,q15,q3 + +.LGemmU8X8.M4.SkipAccumulateOutput: + vst1.32 {d16-d19},[r2]! + vst1.32 {d20-d23},[r0] + vst1.32 {d24-d27},[r3] + vst1.32 {d28-d31},[r4] + cmp r12,#0 + bne .LGemmU8X8.M4.ProcessNextColumnLoop + +.LGemmU8X8.M4.ExitKernel: + mov r0,#4 // return number of rows handled + vpop {d8-d15} + pop {r4,r5,r6,r7,r8,r10} + bx lr + +// +// Store the partial 1 to 7 columns either overwriting the output matrix or +// accumulating into the existing contents of the output matrix. +// + +.LGemmU8X8.M4.StoreOutputPartial: + cbz r5,.LGemmU8X8.M4.StoreOutputPartial.AddMode + +.LGemmU8X8.M4.StoreOutputPartial.ZeroMode: + tst r12,#4 + beq .LGemmU8X8.M4.StoreOutputPartial2.ZeroMode + vst1.32 {d16-d17},[r2]! + vmov q8,q9 // shift remaining elements down + vst1.32 {d20-d21},[r0]! + vmov q10,q11 + vst1.32 {d24-d25},[r3]! + vmov q12,q13 + vst1.32 {d28-d29},[r4]! + vmov q14,q15 + +.LGemmU8X8.M4.StoreOutputPartial2.ZeroMode: + tst r12,#2 + beq .LGemmU8X8.M4.StoreOutputPartial1.ZeroMode + vst1.32 {d16},[r2]! + vmov d16,d17 // shift remaining elements down + vst1.32 {d20},[r0]! + vmov d20,d21 + vst1.32 {d24},[r3]! + vmov d24,d25 + vst1.32 {d28},[r4]! + vmov d28,d29 + +.LGemmU8X8.M4.StoreOutputPartial1.ZeroMode: + tst r12,#1 + beq .LGemmU8X8.M4.ExitKernel + vst1.32 d16[0],[r2] + vst1.32 d20[0],[r0] + vst1.32 d24[0],[r3] + vst1.32 d28[0],[r4] + b .LGemmU8X8.M4.ExitKernel + +.LGemmU8X8.M4.StoreOutputPartial.AddMode: + tst r12,#4 + beq .LGemmU8X8.M4.StoreOutputPartial2.AddMode + vld1.32 {d0-d1},[r2] + vld1.32 {d4-d5},[r0] + vadd.u32 q8,q8,q0 + vld1.32 {d0-d1},[r3] + vadd.u32 q10,q10,q2 + vld1.32 {d4-d5},[r4] + vadd.u32 q12,q12,q0 + vadd.u32 q14,q14,q2 + vst1.32 {d16-d17},[r2]! + vmov q8,q9 // shift remaining elements down + vst1.32 {d20-d21},[r0]! + vmov q10,q11 + vst1.32 {d24-d25},[r3]! + vmov q12,q13 + vst1.32 {d28-d29},[r4]! + vmov q14,q15 + +.LGemmU8X8.M4.StoreOutputPartial2.AddMode: + tst r12,#2 + beq .LGemmU8X8.M4.StoreOutputPartial1.AddMode + vld1.32 {d0},[r2] + vld1.32 {d4},[r0] + vadd.u32 d16,d16,d0 + vld1.32 {d0},[r3] + vadd.u32 d20,d20,d4 + vld1.32 {d4},[r4] + vadd.u32 d24,d24,d0 + vadd.u32 d28,d28,d4 + vst1.32 {d16},[r2]! + vmov d16,d17 // shift remaining elements down + vst1.32 {d20},[r0]! + vmov d20,d21 + vst1.32 {d24},[r3]! + vmov d24,d25 + vst1.32 {d28},[r4]! + vmov d28,d29 + +.LGemmU8X8.M4.StoreOutputPartial1.AddMode: + tst r12,#1 + beq .LGemmU8X8.M4.ExitKernel + vld1.32 d0[0],[r2] + vld1.32 d4[0],[r0] + vadd.u32 d16,d16,d0 + vld1.32 d0[0],[r3] + vadd.u32 d20,d20,d4 + vld1.32 d4[0],[r4] + vadd.u32 d24,d24,d0 + vadd.u32 d28,d28,d4 + vst1.32 d16[0],[r2] + vst1.32 d20[0],[r0] + vst1.32 d24[0],[r3] + vst1.32 d28[0],[r4] + b .LGemmU8X8.M4.ExitKernel + +// +// Process 2 rows of the matrices. +// + +.LGemmU8X8.M2.ProcessNextColumnLoop: + vldr d0,[r1] // load packed B0 + mov r0,r6 // reload matrix A + vld1.32 {d12-d15},[r8]! // load ColumnSumBuffer + mov r3,r7 // reload PackedCountK + vmovl.u8 q0,d0 + vdup.32 q9,d10[0] + vdup.32 q11,d10[1] + vld1.32 d8,[r0]! // load first packed A0 + vadd.u32 q8,q9,q6 + vadd.u32 q9,q9,q7 + vadd.u32 q10,q11,q6 + vadd.u32 q11,q11,q7 + +.LGemmU8X8.M2.ComputeBlockLoop: + vmovl.u8 q2,d8 + vldr d2,[r1,#8] // load packed B1 + vmlal.u16 q8,d0,d4[0] + vmlal.u16 q9,d1,d4[0] + vmlal.u16 q10,d0,d5[0] + vmlal.u16 q11,d1,d5[0] + vmovl.u8 q1,d2 + vldr d0,[r1,#16] // load packed B2 + vmlal.u16 q8,d2,d4[1] + vmlal.u16 q9,d3,d4[1] + vmlal.u16 q10,d2,d5[1] + vmlal.u16 q11,d3,d5[1] + vmovl.u8 q0,d0 + vldr d2,[r1,#24] // load packed B3 + add r1,#32 + subs r3,#1 + beq .LGemmU8X8.M2.ComputeBlockLoopFinish + vmlal.u16 q8,d0,d4[2] + vmlal.u16 q9,d1,d4[2] + vmlal.u16 q10,d0,d5[2] + vmlal.u16 q11,d1,d5[2] + vmovl.u8 q1,d2 + vld1.32 d8,[r0]! // load next packed A0 + vldr d0,[r1] // load packed B0 + vmlal.u16 q8,d2,d4[3] + vmlal.u16 q9,d3,d4[3] + vmlal.u16 q10,d2,d5[3] + vmlal.u16 q11,d3,d5[3] + vmovl.u8 q0,d0 + b .LGemmU8X8.M2.ComputeBlockLoop + +.LGemmU8X8.M2.ComputeBlockLoopFinish: + vmlal.u16 q8,d0,d4[2] // finish computing tail vectors + vmlal.u16 q9,d1,d4[2] + add r0,r2,r10,lsl #2 // compute output row 2 + vmlal.u16 q10,d0,d5[2] + vmlal.u16 q11,d1,d5[2] + vmovl.u8 q1,d2 + vmlal.u16 q8,d2,d4[3] + vmlal.u16 q9,d3,d4[3] + vmlal.u16 q10,d2,d5[3] + vmlal.u16 q11,d3,d5[3] + subs r12,#8 // adjust CountN remaining + blo .LGemmU8X8.M2.StoreOutputPartial + cbnz r5,.LGemmU8X8.M2.SkipAccumulateOutput + vld1.32 {d0-d3},[r2] + vld1.32 {d4-d7},[r0] + vadd.u32 q8,q8,q0 + vadd.u32 q9,q9,q1 + vadd.u32 q10,q10,q2 + vadd.u32 q11,q11,q3 + +.LGemmU8X8.M2.SkipAccumulateOutput: + vst1.32 {d16-d19},[r2]! + vst1.32 {d20-d23},[r0] + cmp r12,#0 + bne .LGemmU8X8.M2.ProcessNextColumnLoop + +.LGemmU8X8.M2.ExitKernel: + mov r0,#2 // return number of rows handled + vpop {d8-d15} + pop {r4,r5,r6,r7,r8,r10} + bx lr + +// +// Store the partial 1 to 7 columns either overwriting the output matrix or +// accumulating into the existing contents of the output matrix. +// + +.LGemmU8X8.M2.StoreOutputPartial: + cbz r5,.LGemmU8X8.M2.StoreOutputPartial.AddMode + +.LGemmU8X8.M2.StoreOutputPartial.ZeroMode: + tst r12,#4 + beq .LGemmU8X8.M2.StoreOutputPartial2.ZeroMode + vst1.32 {d16-d17},[r2]! + vmov q8,q9 // shift remaining elements down + vst1.32 {d20-d21},[r0]! + vmov q10,q11 + +.LGemmU8X8.M2.StoreOutputPartial2.ZeroMode: + tst r12,#2 + beq .LGemmU8X8.M2.StoreOutputPartial1.ZeroMode + vst1.32 {d16},[r2]! + vmov d16,d17 // shift remaining elements down + vst1.32 {d20},[r0]! + vmov d20,d21 + +.LGemmU8X8.M2.StoreOutputPartial1.ZeroMode: + tst r12,#1 + beq .LGemmU8X8.M2.ExitKernel + vst1.32 d16[0],[r2] + vst1.32 d20[0],[r0] + b .LGemmU8X8.M2.ExitKernel + +.LGemmU8X8.M2.StoreOutputPartial.AddMode: + tst r12,#4 + beq .LGemmU8X8.M2.StoreOutputPartial2.AddMode + vld1.32 {d0-d1},[r2] + vld1.32 {d4-d5},[r0] + vadd.u32 q8,q8,q0 + vadd.u32 q10,q10,q2 + vst1.32 {d16-d17},[r2]! + vmov q8,q9 // shift remaining elements down + vst1.32 {d20-d21},[r0]! + vmov q10,q11 + +.LGemmU8X8.M2.StoreOutputPartial2.AddMode: + tst r12,#2 + beq .LGemmU8X8.M2.StoreOutputPartial1.AddMode + vld1.32 {d0},[r2] + vld1.32 {d4},[r0] + vadd.u32 d16,d16,d0 + vadd.u32 d20,d20,d4 + vst1.32 {d16},[r2]! + vmov d16,d17 // shift remaining elements down + vst1.32 {d20},[r0]! + vmov d20,d21 + +.LGemmU8X8.M2.StoreOutputPartial1.AddMode: + tst r12,#1 + beq .LGemmU8X8.M2.ExitKernel + vld1.32 d0[0],[r2] + vld1.32 d4[0],[r0] + vadd.u32 d16,d16,d0 + vadd.u32 d20,d20,d4 + vst1.32 d16[0],[r2] + vst1.32 d20[0],[r0] + b .LGemmU8X8.M2.ExitKernel + +// +// Process 1 row of the matrices. +// + +.LGemmU8X8.M1.ProcessNextColumnLoop: + vldr d0,[r1] // load packed B0 + mov r0,r6 // reload matrix A + vld1.32 {d12-d15},[r8]! // load ColumnSumBuffer + mov r3,r7 // reload PackedCountK + vmovl.u8 q0,d0 + vdup.32 q9,d10[0] + vld1.32 d8[0],[r0]! // load first packed A0 + vadd.u32 q8,q9,q6 + vadd.u32 q9,q9,q7 + +.LGemmU8X8.M1.ComputeBlockLoop: + vmovl.u8 q2,d8 + vldr d2,[r1,#8] // load packed B1 + vmlal.u16 q8,d0,d4[0] + vmlal.u16 q9,d1,d4[0] + vmovl.u8 q1,d2 + vldr d0,[r1,#16] // load packed B2 + vmlal.u16 q8,d2,d4[1] + vmlal.u16 q9,d3,d4[1] + vmovl.u8 q0,d0 + vldr d2,[r1,#24] // load packed B3 + add r1,#32 + subs r3,#1 + beq .LGemmU8X8.M1.ComputeBlockLoopFinish + vmlal.u16 q8,d0,d4[2] + vmlal.u16 q9,d1,d4[2] + vmovl.u8 q1,d2 + vld1.32 d8[0],[r0]! // load next packed A0 + vldr d0,[r1] // load packed B0 + vmlal.u16 q8,d2,d4[3] + vmlal.u16 q9,d3,d4[3] + vmovl.u8 q0,d0 + b .LGemmU8X8.M1.ComputeBlockLoop + +.LGemmU8X8.M1.ComputeBlockLoopFinish: + vmlal.u16 q8,d0,d4[2] // finish computing tail vectors + vmlal.u16 q9,d1,d4[2] + vmovl.u8 q1,d2 + vmlal.u16 q8,d2,d4[3] + vmlal.u16 q9,d3,d4[3] + subs r12,#8 // adjust CountN remaining + blo .LGemmU8X8.M1.StoreOutputPartial + cbnz r5,.LGemmU8X8.M1.SkipAccumulateOutput + vld1.32 {d0-d3},[r2] + vadd.u32 q8,q8,q0 + vadd.u32 q9,q9,q1 + +.LGemmU8X8.M1.SkipAccumulateOutput: + vst1.32 {d16-d19},[r2]! + cmp r12,#0 + bne .LGemmU8X8.M1.ProcessNextColumnLoop + +.LGemmU8X8.M1.ExitKernel: + mov r0,#1 // return number of rows handled + vpop {d8-d15} + pop {r4,r5,r6,r7,r8,r10} + bx lr + +// +// Store the partial 1 to 7 columns either overwriting the output matrix or +// accumulating into the existing contents of the output matrix. +// + +.LGemmU8X8.M1.StoreOutputPartial: + cbz r5,.LGemmU8X8.M1.StoreOutputPartial.AddMode + +.LGemmU8X8.M1.StoreOutputPartial.ZeroMode: + tst r12,#4 + beq .LGemmU8X8.M1.StoreOutputPartial2.ZeroMode + vst1.32 {d16-d17},[r2]! + vmov q8,q9 // shift remaining elements down + +.LGemmU8X8.M1.StoreOutputPartial2.ZeroMode: + tst r12,#2 + beq .LGemmU8X8.M1.StoreOutputPartial1.ZeroMode + vst1.32 {d16},[r2]! + vmov d16,d17 // shift remaining elements down + +.LGemmU8X8.M1.StoreOutputPartial1.ZeroMode: + tst r12,#1 + beq .LGemmU8X8.M1.ExitKernel + vst1.32 d16[0],[r2] + b .LGemmU8X8.M1.ExitKernel + +.LGemmU8X8.M1.StoreOutputPartial.AddMode: + tst r12,#4 + beq .LGemmU8X8.M1.StoreOutputPartial2.AddMode + vld1.32 {d0-d1},[r2] + vadd.u32 q8,q8,q0 + vst1.32 {d16-d17},[r2]! + vmov q8,q9 // shift remaining elements down + +.LGemmU8X8.M1.StoreOutputPartial2.AddMode: + tst r12,#2 + beq .LGemmU8X8.M1.StoreOutputPartial1.AddMode + vld1.32 {d0},[r2] + vadd.u32 d16,d16,d0 + vst1.32 {d16},[r2]! + vmov d16,d17 // shift remaining elements down + +.LGemmU8X8.M1.StoreOutputPartial1.AddMode: + tst r12,#1 + beq .LGemmU8X8.M1.ExitKernel + vld1.32 d0[0],[r2] + vadd.u32 d16,d16,d0 + vst1.32 d16[0],[r2] + b .LGemmU8X8.M1.ExitKernel + + .end diff --git a/onnxruntime/core/mlas/lib/aarch32/asmmacro.h b/onnxruntime/core/mlas/lib/aarch32/asmmacro.h new file mode 100644 index 0000000000..72982db003 --- /dev/null +++ b/onnxruntime/core/mlas/lib/aarch32/asmmacro.h @@ -0,0 +1,95 @@ +/*++ + +Copyright (c) Microsoft Corporation. All rights reserved. + +Licensed under the MIT License. + +Module Name: + + asmmacro.h + +Abstract: + + This module implements common macros for the assembly modules. + +--*/ + +/*++ + +Macro Description: + + This macro emits the assembler directives to annotate a new function. + +Arguments: + + FunctionName - Supplies the name of the function. + +--*/ + + .macro FUNCTION_ENTRY FunctionName + + .p2align 2 +#if defined(__APPLE__) + .globl _\FunctionName\() +_\FunctionName\(): +#else + .globl \FunctionName\() + .type \FunctionName\(),%function +\FunctionName\(): +#endif + + .endm + +/*++ + +Macro Description: + + This macro conditionally emits the statement if Count is greater than or + equal to Value. + +Arguments: + + Count - Supplies the variable used in the comparison. + + Value - Supplies the static used in the comparison. + + Statement - Supplies the statement to conditionally emit. + +--*/ + + .macro EmitIfCountGE Count1, Value1, Statement + +.if (\Count1\() >= \Value1\()) + \Statement\() +.endif + + .endm + +/*++ + +Macro Description: + + This macro conditionally emits the statement if Count1 is greater than or + equal to Value1 and Count2 is greater than or equal to Value2. + +Arguments: + + Count1 - Supplies the variable used in the comparison. + + Value1 - Supplies the static used in the comparison. + + Count2 - Supplies the variable used in the comparison. + + Value2 - Supplies the static used in the comparison. + + Statement - Supplies the statement to conditionally emit. + +--*/ + + .macro EmitIfCount2GE Count1, Value1, Count2, Value2, Statement + +.if (\Count1\() >= \Value1\()) && (\Count2\() >= \Value2\()) + \Statement\() +.endif + + .endm diff --git a/onnxruntime/core/mlas/lib/aarch64/QgemmU8X8KernelNeon.S b/onnxruntime/core/mlas/lib/aarch64/QgemmU8X8KernelNeon.S index 6e93cb81cb..d1dad123ff 100644 --- a/onnxruntime/core/mlas/lib/aarch64/QgemmU8X8KernelNeon.S +++ b/onnxruntime/core/mlas/lib/aarch64/QgemmU8X8KernelNeon.S @@ -37,10 +37,10 @@ Routine Description: Arguments: A (x0) - Supplies the address of matrix A. The matrix data has been packed - using MlasGemmU8S8CopyPackANeon. + using MlasGemmU8X8CopyPackANeon. B (x1) - Supplies the address of matrix B. The matrix data has been packed - using MlasGemmU8S8CopyPackBNeon. + using MlasGemmU8X8CopyPackBNeon. C (x2) - Supplies the address of matrix C. diff --git a/onnxruntime/core/mlas/lib/arm64/QgemmU8X8KernelNeon.asm b/onnxruntime/core/mlas/lib/arm64/QgemmU8X8KernelNeon.asm index 06ccdb2578..d123bd67d4 100644 --- a/onnxruntime/core/mlas/lib/arm64/QgemmU8X8KernelNeon.asm +++ b/onnxruntime/core/mlas/lib/arm64/QgemmU8X8KernelNeon.asm @@ -48,10 +48,10 @@ Routine Description: Arguments: A (x0) - Supplies the address of matrix A. The matrix data has been packed - using MlasGemmU8S8CopyPackANeon. + using MlasGemmU8X8CopyPackANeon. B (x1) - Supplies the address of matrix B. The matrix data has been packed - using MlasGemmU8S8CopyPackBNeon. + using MlasGemmU8X8CopyPackBNeon. C (x2) - Supplies the address of matrix C. diff --git a/onnxruntime/core/mlas/lib/qgemm.cpp b/onnxruntime/core/mlas/lib/qgemm.cpp index 0e1a69964a..ad54fbc5c4 100644 --- a/onnxruntime/core/mlas/lib/qgemm.cpp +++ b/onnxruntime/core/mlas/lib/qgemm.cpp @@ -17,6 +17,8 @@ Abstract: #include "mlasi.h" +#ifdef MLAS_SUPPORTS_GEMM_U8X8 + // // Define the parameters to execute segments of a QGEMM operation on worker // threads. @@ -266,7 +268,7 @@ Return Value: offb = typename KernelType::OffsetBType(offb ^ 0x80); } } -#elif defined(MLAS_NEON64_INTRINSICS) +#elif defined(MLAS_NEON_INTRINSICS) if (WorkBlock->BIsSigned) { offb = typename KernelType::OffsetBType(offb ^ 0x80); } @@ -416,7 +418,7 @@ Return Value: offb = typename KernelType::OffsetBType(offb ^ 0x80); } } -#elif defined(MLAS_NEON64_INTRINSICS) +#elif defined(MLAS_NEON_INTRINSICS) if (WorkBlock->BIsSigned) { offb = typename KernelType::OffsetBType(offb ^ 0x80); } @@ -1416,7 +1418,7 @@ MlasGemmU8X8PackedOperation( #endif -#ifdef MLAS_NEON64_INTRINSICS +#ifdef MLAS_NEON_INTRINSICS // // Define the prototypes of the NEON routines written in assembly. @@ -1515,14 +1517,33 @@ Return Value: uint32x4_t v3 = vld1q_u32(reinterpret_cast(a3)); a3 += 16; +#if defined(MLAS_NEON32_INTRINSICS) + uint32x4x2_t z0 = vzipq_u32(v0, v2); + uint32x4x2_t z1 = vzipq_u32(v1, v3); + + v0 = z0.val[0]; + v1 = z0.val[1]; + v2 = z1.val[0]; + v3 = z1.val[1]; + + uint32x4x2_t z2 = vzipq_u32(v0, v2); + uint32x4x2_t z3 = vzipq_u32(v1, v3); + + v0 = z2.val[0]; + v1 = z2.val[1]; + v2 = z3.val[0]; + v3 = z3.val[1]; +#else uint32x4_t z0 = vzip1q_u32(v0, v2); uint32x4_t z1 = vzip2q_u32(v0, v2); uint32x4_t z2 = vzip1q_u32(v1, v3); uint32x4_t z3 = vzip2q_u32(v1, v3); + v0 = vzip1q_u32(z0, z2); v1 = vzip2q_u32(z0, z2); v2 = vzip1q_u32(z1, z3); v3 = vzip2q_u32(z1, z3); +#endif vst1q_u8(&D[0], vreinterpretq_u8_u32(v0)); vst1q_u8(&D[16], vreinterpretq_u8_u32(v1)); @@ -1721,14 +1742,18 @@ Return Value: RowSums = vpadalq_u16(RowSums, vpaddlq_u8(v)); } -#if defined(_M_ARM64) +#if defined(MLAS_NEON32_INTRINSICS) + uint32x2_t RowSumsLow = vpadd_u32(vget_high_u32(RowSums), vget_low_u32(RowSums)); + RowSumsLow = vpadd_u32(RowSumsLow, RowSumsLow); + vst1_lane_u32(reinterpret_cast(RowSumBuffer), RowSumsLow, 0); +#elif defined(_M_ARM64) // N.B. The workaround of defining a local vaddvq_u32 doesn't work here // as VS2019 added new intrinsics to make the operation work. Also, not // all build environments using VS2019 have the up-to-date arm64_neon.h, // so fallback to pairwise addition. RowSums = vpaddq_u32(RowSums, RowSums); RowSums = vpaddq_u32(RowSums, RowSums); - vst1q_lane_u32(reinterpret_cast(RowSumBuffer), RowSums, 0); + vst1q_lane_u32(reinterpret_cast(RowSumBuffer), RowSums, 0); #else *RowSumBuffer = int32_t(vaddvq_u32(RowSums)); #endif @@ -1749,7 +1774,11 @@ MlasGemmU8X8CopyPackBProcessNeon( uint16x8_t WordsRow = vmovl_u8(BytesRow); ColumnSums[0] = vaddq_u32(ColumnSums[0], vmovl_u16(vget_low_u16(WordsRow))); +#if defined(MLAS_NEON32_INTRINSICS) + ColumnSums[1] = vaddq_u32(ColumnSums[1], vmovl_u16(vget_high_u16(WordsRow))); +#else ColumnSums[1] = vaddq_u32(ColumnSums[1], vmovl_high_u16(WordsRow)); +#endif } void @@ -2050,12 +2079,14 @@ Return Value: GemmU8X8Operation(&WorkBlock); #elif defined(MLAS_SSE2_INTRINSICS) MlasGemmU8X8Operation(&WorkBlock); -#elif defined(MLAS_NEON64_INTRINSICS) +#elif defined(MLAS_NEON_INTRINSICS) if (WorkBlock.BIsPacked) { MlasGemmU8X8PackedOperation(&WorkBlock); } else { MlasGemmU8X8Operation(&WorkBlock); } +#else +#error Unsupported architecture. #endif } @@ -2331,7 +2362,9 @@ Return Value: MlasGemmU8X8Schedule(&WorkBlock, ThreadPool); } -#if defined(MLAS_TARGET_AMD64) || defined(MLAS_NEON64_INTRINSICS) +#endif // MLAS_SUPPORTS_GEMM_U8X8 + +#ifdef MLAS_SUPPORTS_PACKED_GEMM_U8X8 void MLASCALL @@ -2563,7 +2596,7 @@ Return Value: } else { return 0; } -#elif defined(MLAS_NEON64_INTRINSICS) +#elif defined(MLAS_NEON_INTRINSICS) MLAS_UNREFERENCED_PARAMETER(BIsSigned); PackedK = MLAS_GEMM_U8X8_KERNEL_NEON::PackedK; @@ -2650,7 +2683,7 @@ Return Value: throw std::runtime_error("packing unavailable"); #endif } -#elif defined(MLAS_NEON64_INTRINSICS) +#elif defined(MLAS_NEON_INTRINSICS) PackedK = MLAS_GEMM_U8X8_KERNEL_NEON::PackedK; StrideK = MLAS_GEMM_U8X8_KERNEL_NEON::PackedStrides.K; #else @@ -2700,7 +2733,7 @@ Return Value: } else { MLAS_GEMM_U8U8_KERNEL_AVX2::CopyPackB(pb, B + n, ldb, CountN, CountK, ColumnSumBuffer, BIsSigned); } -#elif defined(MLAS_NEON64_INTRINSICS) +#elif defined(MLAS_NEON_INTRINSICS) MLAS_GEMM_U8X8_KERNEL_NEON::CopyPackB(pb, B + n, ldb, CountN, CountK, ColumnSumBuffer, BIsSigned); #else #error Unknown architecture. @@ -2723,4 +2756,4 @@ Return Value: } } -#endif +#endif // MLAS_SUPPORTS_PACKED_GEMM_U8X8