MLAS: implement u8x8 GEMM for aarch32 (#5580)

This commit is contained in:
Tracy Sharpe 2020-10-25 23:05:12 -07:00 committed by GitHub
parent b2da700e4d
commit 502f67ba58
No known key found for this signature in database
GPG key ID: 4AEE18F83AFDEB23
7 changed files with 761 additions and 17 deletions

View file

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

View file

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

View 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

View 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

View file

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

View file

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

View file

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