mirror of
https://github.com/saymrwulf/onnxruntime.git
synced 2026-07-27 20:02:15 +00:00
MLAS: implement u8x8 GEMM for aarch32 (#5580)
This commit is contained in:
parent
b2da700e4d
commit
502f67ba58
7 changed files with 761 additions and 17 deletions
|
|
@ -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
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
610
onnxruntime/core/mlas/lib/aarch32/QgemmU8X8KernelNeon.S
Normal file
610
onnxruntime/core/mlas/lib/aarch32/QgemmU8X8KernelNeon.S
Normal file
|
|
@ -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
|
||||
95
onnxruntime/core/mlas/lib/aarch32/asmmacro.h
Normal file
95
onnxruntime/core/mlas/lib/aarch32/asmmacro.h
Normal file
|
|
@ -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
|
||||
|
|
@ -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.
|
||||
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
||||
|
|
|
|||
|
|
@ -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<MLAS_GEMM_U8U8_KERNEL_AVX2>(
|
|||
|
||||
#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<const uint32_t*>(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<uint32_t*>(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<int32_t*>(RowSumBuffer), RowSums, 0);
|
||||
vst1q_lane_u32(reinterpret_cast<uint32_t*>(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<MLAS_GEMM_U8X8_KERNEL_SSE>(&WorkBlock);
|
||||
#elif defined(MLAS_NEON64_INTRINSICS)
|
||||
#elif defined(MLAS_NEON_INTRINSICS)
|
||||
if (WorkBlock.BIsPacked) {
|
||||
MlasGemmU8X8PackedOperation<MLAS_GEMM_U8X8_KERNEL_NEON>(&WorkBlock);
|
||||
} else {
|
||||
MlasGemmU8X8Operation<MLAS_GEMM_U8X8_KERNEL_NEON>(&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
|
||||
|
|
|
|||
Loading…
Reference in a new issue