mirror of
https://github.com/saymrwulf/onnxruntime.git
synced 2026-07-29 20:14:01 +00:00
Symmetric quantized convolution kernel ARM64 (#9772)
Adding a symmetric quantized convolution kernel for ARM64 Note: Indirect conv performs worse for shallow convs (input channels are small). This is much more so for low end pre-dot CPUs, where only 128 or deeper conv is faster with indirect conv. With DOT-CPUs, 32 deep conv is already faster Co-authored-by: Chen Fu <fuchen@microsoft.com>
This commit is contained in:
parent
7e55a942cd
commit
cd0af7ad44
16 changed files with 2259 additions and 29 deletions
|
|
@ -47,6 +47,8 @@ function(setup_mlas_source_for_windows)
|
|||
)
|
||||
|
||||
set(mlas_platform_preprocess_srcs
|
||||
${MLAS_SRC_DIR}/arm64/ConvSymU8KernelDot.asm
|
||||
${MLAS_SRC_DIR}/arm64/ConvSymU8KernelNeon.asm
|
||||
${MLAS_SRC_DIR}/arm64/DepthwiseConvsymKernelNeon.asm
|
||||
${MLAS_SRC_DIR}/arm64/DepthwiseQConvKernelSize9Neon.asm
|
||||
${MLAS_SRC_DIR}/arm64/QgemmU8X8KernelNeon.asm
|
||||
|
|
@ -268,6 +270,8 @@ else()
|
|||
if(ARM64 AND MLAS_SOURCE_IS_NOT_SET )
|
||||
enable_language(ASM)
|
||||
set(mlas_platform_srcs
|
||||
${MLAS_SRC_DIR}/aarch64/ConvSymU8KernelDot.S
|
||||
${MLAS_SRC_DIR}/aarch64/ConvSymU8KernelNeon.S
|
||||
${MLAS_SRC_DIR}/aarch64/DepthwiseConvSymKernelNeon.S
|
||||
${MLAS_SRC_DIR}/aarch64/DepthwiseQConvKernelSize9Neon.S
|
||||
${MLAS_SRC_DIR}/aarch64/QgemmU8X8KernelNeon.S
|
||||
|
|
|
|||
|
|
@ -13,7 +13,7 @@ add_executable(onnxruntime_webassembly
|
|||
|
||||
if (NOT onnxruntime_ENABLE_WEBASSEMBLY_THREADS)
|
||||
add_compile_definitions(
|
||||
MLAS_NO_ONNXRUNTIME_THREADPOOL
|
||||
BUILD_MLAS_NO_ONNXRUNTIME
|
||||
)
|
||||
|
||||
# Override re2 compiler options to remove -pthread
|
||||
|
|
|
|||
|
|
@ -95,12 +95,16 @@ CPUIDInfo::CPUIDInfo() {
|
|||
}
|
||||
#endif
|
||||
|
||||
#if defined(CPUIDINFO_ARCH_ARM) && defined(CPUINFO_SUPPORTED)
|
||||
#if defined(CPUIDINFO_ARCH_ARM)
|
||||
#ifdef CPUINFO_SUPPORTED
|
||||
|
||||
// only works on ARM linux or android, does not work on Windows
|
||||
is_hybrid_ = cpuinfo_get_uarchs_count() > 1;
|
||||
has_arm_neon_dot_ = cpuinfo_has_arm_neon_dot();
|
||||
|
||||
#elif defined(_WIN32)
|
||||
// TODO implement hardware feature detection in windows.
|
||||
is_hybrid_ = true;
|
||||
#endif
|
||||
#endif
|
||||
|
||||
}
|
||||
|
|
|
|||
|
|
@ -83,3 +83,4 @@ Arguments:
|
|||
.inst Instruction
|
||||
|
||||
.endm
|
||||
|
||||
|
|
|
|||
628
onnxruntime/core/mlas/lib/aarch64/ConvSymU8KernelDot.S
Normal file
628
onnxruntime/core/mlas/lib/aarch64/ConvSymU8KernelDot.S
Normal file
|
|
@ -0,0 +1,628 @@
|
|||
/*++
|
||||
|
||||
Copyright (c) Microsoft Corporation. All rights reserved.
|
||||
|
||||
Licensed under the MIT License.
|
||||
|
||||
Module Name:
|
||||
|
||||
ConvSymKernelNeonDot.S
|
||||
|
||||
Abstract:
|
||||
|
||||
This module implements the kernels for the symmetric quantized integer
|
||||
convolution operation.
|
||||
|
||||
--*/
|
||||
|
||||
#include "asmmacro.h"
|
||||
#include "AssembleDotProduct.h"
|
||||
|
||||
.equ .LMLAS_CONV_SYM_FLAG_INPUT_DIRECT, 1
|
||||
.equ .LMLAS_CONV_SYM_FLAG_PER_CHANNEL_SCALE, 2
|
||||
|
||||
//
|
||||
// Stack frame layout for the symmetric convolution kernel.
|
||||
// d8-d15, x19-x30 need to be preserved if used
|
||||
//
|
||||
.equ .LConvSymFrame_SavedRegisters, (8 * 8)
|
||||
.equ .LConvSymFrame_PostProcessParams, 0 + .LConvSymFrame_SavedRegisters
|
||||
.equ .LConvSymFrame_KernelFlags, 8 + .LConvSymFrame_SavedRegisters
|
||||
|
||||
.equ .LConvSymPostProcessParams_Bias, 0
|
||||
.equ .LConvSymPostProcessParams_Scale, 8
|
||||
.equ .LConvSymPostProcessParams_Min, 16
|
||||
.equ .LConvSymPostProcessParams_Max, 20
|
||||
.equ .LConvSymPostProcessParams_ZeroPoint, 24
|
||||
|
||||
.text
|
||||
|
||||
/*++
|
||||
|
||||
Routine Description:
|
||||
|
||||
This routine is the inner kernel to compute a convolution for the elements
|
||||
of an output row for a set of filter rows.
|
||||
|
||||
Arguments:
|
||||
|
||||
Input (x0) - Points to the input buffer.
|
||||
|
||||
If MLAS_CONV_SYM_FLAG_INPUT_DIRECT is set, then the input buffer points
|
||||
directly at the input tensor.
|
||||
|
||||
If MLAS_CONV_SYM_FLAG_INPUT_DIRECT is clear, then the input buffer is an
|
||||
indirection buffer. Every pointer in the indirection buffer points at a
|
||||
InputChannels length vector (either from the input tensor or a vector of
|
||||
padding values). These are grouped in batches of length KernelSize.
|
||||
These batches are then repeated OutputCount times.
|
||||
|
||||
Filter (x1) - Points to the filter buffer.
|
||||
|
||||
Output (x2) - Points the output buffer.
|
||||
|
||||
KernelSize (x3/x9) - Size of the kernel (most commonly. 3x3=9, 5x5=25).
|
||||
|
||||
If MLAS_CONV_SYM_FLAG_INPUT_DIRECT is set, then kernel size should be 1.
|
||||
|
||||
InputChannels (x4/x7) - Number of input channels.
|
||||
|
||||
OutputChannels (x5) - Number of output channels.
|
||||
|
||||
ChannelCount (x6) - Number of output channels this iteration produces.
|
||||
|
||||
OutputCount (x7) - Number of output elements this iteration produces.
|
||||
|
||||
This implementation requires the count to be no larger than 4.
|
||||
|
||||
PostProcessParams (x8) - Points to the post process parameter block.
|
||||
|
||||
KernelFlags - (w10) Additional flags controlling the operation.
|
||||
|
||||
Return Value:
|
||||
|
||||
None.
|
||||
|
||||
--*/
|
||||
FUNCTION_ENTRY MlasConvSymKernelNeonDot
|
||||
|
||||
stp d8,d9,[sp,#-.LConvSymFrame_SavedRegisters]!
|
||||
ldr x8,[sp,#.LConvSymFrame_PostProcessParams]
|
||||
ldr w10,[sp,#.LConvSymFrame_KernelFlags]
|
||||
stp d10,d11,[sp,#16]
|
||||
stp d12,d13,[sp,#32]
|
||||
stp x19,x20,[sp,#48]
|
||||
|
||||
cmp x7,2 // OutputCount < 2 ?
|
||||
add x16,x2,x5 // x16 -> C1
|
||||
lsl x3,x3,#3 // KernelSize * sizeof(int8_t*)
|
||||
csel x16,x2,x16,lo // if OutputCount < 2 x16/C1 -> C0
|
||||
mov x20,x4
|
||||
add x4,x4,3 // InputChannels align to 4
|
||||
add x17,x16,x5 // x17 -> C2
|
||||
ldr x11,[x8,#.LConvSymPostProcessParams_Bias]
|
||||
csel x17,x16,x17,ls // if OutputCount <= 2 x17/C2 -> C1
|
||||
bic x4,x4,3
|
||||
cmp x7,4 // OutputCount < 4 ?
|
||||
add x5,x17,x5 // x5 -> C3
|
||||
ldr x19,[x8,#.LConvSymPostProcessParams_Scale]
|
||||
csel x5,x17,x5,lo // if OutputCount < 4 x5/C3 -> C2
|
||||
movi v12.16b,128 // for top bit flipping
|
||||
|
||||
OutputChannelLoop:
|
||||
ldp q16,q20,[x11],32 // Init accumulators with biases
|
||||
mov v17.16b,v16.16b
|
||||
mov v18.16b,v16.16b
|
||||
ldp q24,q28,[x11],32
|
||||
mov v19.16b,v16.16b
|
||||
mov v21.16b,v20.16b
|
||||
mov v22.16b,v20.16b
|
||||
mov v23.16b,v20.16b
|
||||
mov v25.16b,v24.16b
|
||||
mov v26.16b,v24.16b
|
||||
mov v27.16b,v24.16b
|
||||
mov v29.16b,v28.16b
|
||||
mov v30.16b,v28.16b
|
||||
mov v31.16b,v28.16b
|
||||
mov x9,x3 // restore KernelSize * sizeof(int8_t*)
|
||||
|
||||
KernelSizeLoop:
|
||||
tst w10,#.LMLAS_CONV_SYM_FLAG_INPUT_DIRECT
|
||||
beq InputIndirection
|
||||
|
||||
InputDirect:
|
||||
cmp x16,x2
|
||||
mov x12,x0 // x12 -> A0
|
||||
add x13,x0,x20 // x13 -> A1 = A0 + input channels
|
||||
csel x13,x0,x13,eq
|
||||
cmp x17,x16
|
||||
add x14,x0,x20,lsl#1 // x14 -> A2
|
||||
csel x14,x13,x14,eq
|
||||
cmp x5,x17
|
||||
add x15,x13,x20,lsl#1 // x15 -> A3
|
||||
csel x15,x14,x15,eq
|
||||
b FinishLoadAPtr
|
||||
|
||||
InputIndirection:
|
||||
ldr x12,[x0] // x12 -> A0
|
||||
cmp x16,x2
|
||||
b.eq SkipLoadA1 // C1==C0 -> A0=A1=A2=A3
|
||||
cmp x17,x16
|
||||
lsl x14,x3,#1
|
||||
ldr x13,[x0,x3] // x13 -> A1
|
||||
b.eq SkipLoadA2 // C2==C1 -> A1=A2=A3
|
||||
cmp x5,x17
|
||||
add x15,x3,x3,lsl#1
|
||||
ldr x14,[x0,x14] // x14 -> A2
|
||||
b.eq SkipLoadA3 // C3==C2 -> A2=A3
|
||||
ldr x15,[x0,x15] // x15 -> A3
|
||||
b FinishLoadAPtr
|
||||
SkipLoadA1:
|
||||
mov x13,x12
|
||||
SkipLoadA2:
|
||||
mov x14,x13
|
||||
SkipLoadA3:
|
||||
mov x15,x14
|
||||
|
||||
// Register Usage
|
||||
// B (x1) -> 4x16
|
||||
// ----------------------------------------------------------------------------
|
||||
// |v4.b[0]..v4.b[12] v5.b[0]..v5.b[12] v6.b[0]..v6.b[12] v7.b[0]..v7.b[12]|
|
||||
// | ... ... ... ... ... ... ... ... |
|
||||
// |v4.b[3]..v4.b[15] v5.b[3]..v5.b[15] v6.b[3]..v6.b[15] v7.b[3]..v7.b[15]|
|
||||
// A 4x4 ----------------------------------------------------------------------------
|
||||
// ------------------ ----------------------------------------------------------------------------
|
||||
// x12 |v0.b[0]..v0.b[3]| |v16.s[0]_v16.s[3] v20.s[0]_v20.s[3] v24.s[0]_v24.s[3] v28.s[0]_v28.s[3]| x2
|
||||
// x13 |v1.b[0]..v1.b[3]| |v17.s[0]_v17.s[3] v21.s[0]_v21.s[3] v25.s[0]_v25.s[3] v29.s[0]_v29.s[3]| x16
|
||||
// x14 |v2.b[0]..v2.b[3]| |v18.s[0]_v18.s[3] v22.s[0]_v23.s[3] v26.s[0]_v26.s[3] v30.s[0]_v31.s[3]| x17
|
||||
// x15 |v3.b[0]..v3.b[3]| |v19.s[0]_v19.s[3] v23.s[0]_v23.s[3] v27.s[0]_v27.s[3] v31.s[0]_v31.s[3]| x5
|
||||
// ------------------ ----------------------------------------------------------------------------
|
||||
|
||||
FinishLoadAPtr:
|
||||
subs x7,x4,16 // Need 16 input channels for loop
|
||||
add x0,x0,8 // indirect A advance to next pointer, prepare for kernel size loop
|
||||
b.lo InChannels8
|
||||
|
||||
ldr d0,[x12],8
|
||||
ldr q4,[x1],16
|
||||
ldr d1,[x13],8
|
||||
subs x7,x7,16
|
||||
ldr d2,[x14],8
|
||||
ldr d3,[x15],8
|
||||
ldr q5,[x1],16
|
||||
ldr q6,[x1],16
|
||||
ldr q7,[x1],16
|
||||
b.lo InChLoopEpilogue // Need 32 input channels for main loop
|
||||
|
||||
InputChannelLoop:
|
||||
eor v0.8b,v0.8b,v12.8b
|
||||
eor v1.8b,v1.8b,v12.8b
|
||||
SdotByElement 16, 4, 0,0
|
||||
eor v2.8b,v2.8b,v12.8b
|
||||
SdotByElement 17, 4, 1,0
|
||||
eor v3.8b,v3.8b,v12.8b
|
||||
ldr d8,[x12],8
|
||||
SdotByElement 18, 4, 2,0
|
||||
SdotByElement 19, 4, 3,0
|
||||
ldr q4,[x1],16
|
||||
SdotByElement 20, 5, 0,0
|
||||
SdotByElement 21, 5, 1,0
|
||||
ldr d9,[x13],8
|
||||
SdotByElement 22, 5, 2,0
|
||||
SdotByElement 23, 5, 3,0
|
||||
ldr q5,[x1],16
|
||||
SdotByElement 24, 6, 0,0
|
||||
SdotByElement 25, 6, 1,0
|
||||
ldr d10,[x14],8
|
||||
SdotByElement 26, 6, 2,0
|
||||
SdotByElement 27, 6, 3,0
|
||||
ldr q6,[x1],16
|
||||
SdotByElement 28, 7, 0,0
|
||||
SdotByElement 29, 7, 1,0
|
||||
ldr d11,[x15],8
|
||||
SdotByElement 30, 7, 2,0
|
||||
SdotByElement 31, 7, 3,0
|
||||
ldr q7,[x1],16
|
||||
SdotByElement 16, 4, 0,1
|
||||
SdotByElement 17, 4, 1,1
|
||||
SdotByElement 18, 4, 2,1
|
||||
SdotByElement 19, 4, 3,1
|
||||
ldr q4,[x1],16
|
||||
SdotByElement 20, 5, 0,1
|
||||
SdotByElement 21, 5, 1,1
|
||||
SdotByElement 22, 5, 2,1
|
||||
SdotByElement 23, 5, 3,1
|
||||
ldr q5,[x1],16
|
||||
SdotByElement 24, 6, 0,1
|
||||
SdotByElement 25, 6, 1,1
|
||||
SdotByElement 26, 6, 2,1
|
||||
SdotByElement 27, 6, 3,1
|
||||
ldr q6,[x1],16
|
||||
SdotByElement 28, 7, 0,1
|
||||
SdotByElement 29, 7, 1,1
|
||||
SdotByElement 30, 7, 2,1
|
||||
SdotByElement 31, 7, 3,1
|
||||
eor v8.8b,v8.8b,v12.8b
|
||||
ldr q7,[x1],16
|
||||
eor v9.8b,v9.8b,v12.8b
|
||||
SdotByElement 16, 4, 8,0
|
||||
eor v10.8b,v10.8b,v12.8b
|
||||
SdotByElement 17, 4, 9,0
|
||||
ldr d0,[x12],8
|
||||
eor v11.8b,v11.8b,v12.8b
|
||||
SdotByElement 18, 4,10,0
|
||||
SdotByElement 19, 4,11,0
|
||||
ldr q4,[x1],16
|
||||
SdotByElement 20, 5, 8,0
|
||||
SdotByElement 21, 5, 9,0
|
||||
ldr d1,[x13],8
|
||||
SdotByElement 22, 5,10,0
|
||||
SdotByElement 23, 5,11,0
|
||||
ldr q5,[x1],16
|
||||
SdotByElement 24, 6, 8,0
|
||||
SdotByElement 25, 6, 9,0
|
||||
ldr d2,[x14],8
|
||||
SdotByElement 26, 6,10,0
|
||||
SdotByElement 27, 6,11,0
|
||||
ldr q6,[x1],16
|
||||
SdotByElement 28, 7, 8,0
|
||||
SdotByElement 29, 7, 9,0
|
||||
ldr d3,[x15],8
|
||||
SdotByElement 30, 7,10,0
|
||||
SdotByElement 31, 7,11,0
|
||||
ldr q7,[x1],16
|
||||
SdotByElement 16, 4, 8,1
|
||||
SdotByElement 17, 4, 9,1
|
||||
SdotByElement 18, 4,10,1
|
||||
SdotByElement 19, 4,11,1
|
||||
ldr q4,[x1],16
|
||||
SdotByElement 20, 5, 8,1
|
||||
SdotByElement 21, 5, 9,1
|
||||
SdotByElement 22, 5,10,1
|
||||
SdotByElement 23, 5,11,1
|
||||
ldr q5,[x1],16
|
||||
SdotByElement 24, 6, 8,1
|
||||
SdotByElement 25, 6, 9,1
|
||||
SdotByElement 26, 6,10,1
|
||||
SdotByElement 27, 6,11,1
|
||||
ldr q6,[x1],16
|
||||
SdotByElement 28, 7, 8,1
|
||||
SdotByElement 29, 7, 9,1
|
||||
subs x7,x7,16 // InputChannels -= 16
|
||||
SdotByElement 30, 7,10,1
|
||||
SdotByElement 31, 7,11,1
|
||||
ldr q7,[x1],16
|
||||
b.hs InputChannelLoop
|
||||
|
||||
InChLoopEpilogue:
|
||||
eor v0.8b,v0.8b,v12.8b
|
||||
eor v1.8b,v1.8b,v12.8b
|
||||
SdotByElement 16, 4, 0,0
|
||||
eor v2.8b,v2.8b,v12.8b
|
||||
SdotByElement 17, 4, 1,0
|
||||
eor v3.8b,v3.8b,v12.8b
|
||||
ldr d8,[x12],8
|
||||
SdotByElement 18, 4, 2,0
|
||||
SdotByElement 19, 4, 3,0
|
||||
ldr q4,[x1],16
|
||||
SdotByElement 20, 5, 0,0
|
||||
SdotByElement 21, 5, 1,0
|
||||
ldr d9,[x13],8
|
||||
SdotByElement 22, 5, 2,0
|
||||
SdotByElement 23, 5, 3,0
|
||||
ldr q5,[x1],16
|
||||
SdotByElement 24, 6, 0,0
|
||||
SdotByElement 25, 6, 1,0
|
||||
ldr d10,[x14],8
|
||||
SdotByElement 26, 6, 2,0
|
||||
SdotByElement 27, 6, 3,0
|
||||
ldr q6,[x1],16
|
||||
SdotByElement 28, 7, 0,0
|
||||
SdotByElement 29, 7, 1,0
|
||||
ldr d11,[x15],8
|
||||
SdotByElement 30, 7, 2,0
|
||||
SdotByElement 31, 7, 3,0
|
||||
ldr q7,[x1],16
|
||||
SdotByElement 16, 4, 0,1
|
||||
SdotByElement 17, 4, 1,1
|
||||
SdotByElement 18, 4, 2,1
|
||||
SdotByElement 19, 4, 3,1
|
||||
ldr q4,[x1],16
|
||||
SdotByElement 20, 5, 0,1
|
||||
SdotByElement 21, 5, 1,1
|
||||
SdotByElement 22, 5, 2,1
|
||||
SdotByElement 23, 5, 3,1
|
||||
ldr q5,[x1],16
|
||||
SdotByElement 24, 6, 0,1
|
||||
SdotByElement 25, 6, 1,1
|
||||
SdotByElement 26, 6, 2,1
|
||||
SdotByElement 27, 6, 3,1
|
||||
ldr q6,[x1],16
|
||||
SdotByElement 28, 7, 0,1
|
||||
SdotByElement 29, 7, 1,1
|
||||
SdotByElement 30, 7, 2,1
|
||||
SdotByElement 31, 7, 3,1
|
||||
eor v8.8b,v8.8b,v12.8b
|
||||
ldr q7,[x1],16
|
||||
eor v9.8b,v9.8b,v12.8b
|
||||
SdotByElement 16, 4, 8,0
|
||||
eor v10.8b,v10.8b,v12.8b
|
||||
SdotByElement 17, 4, 9,0
|
||||
eor v11.8b,v11.8b,v12.8b
|
||||
SdotByElement 18, 4,10,0
|
||||
SdotByElement 19, 4,11,0
|
||||
ldr q4,[x1],16
|
||||
SdotByElement 20, 5, 8,0
|
||||
SdotByElement 21, 5, 9,0
|
||||
SdotByElement 22, 5,10,0
|
||||
SdotByElement 23, 5,11,0
|
||||
ldr q5,[x1],16
|
||||
SdotByElement 24, 6, 8,0
|
||||
SdotByElement 25, 6, 9,0
|
||||
SdotByElement 26, 6,10,0
|
||||
SdotByElement 27, 6,11,0
|
||||
ldr q6,[x1],16
|
||||
SdotByElement 28, 7, 8,0
|
||||
SdotByElement 29, 7, 9,0
|
||||
SdotByElement 30, 7,10,0
|
||||
SdotByElement 31, 7,11,0
|
||||
ldr q7,[x1],16
|
||||
SdotByElement 16, 4, 8,1
|
||||
SdotByElement 17, 4, 9,1
|
||||
SdotByElement 18, 4,10,1
|
||||
SdotByElement 19, 4,11,1
|
||||
SdotByElement 20, 5, 8,1
|
||||
SdotByElement 21, 5, 9,1
|
||||
SdotByElement 22, 5,10,1
|
||||
SdotByElement 23, 5,11,1
|
||||
SdotByElement 24, 6, 8,1
|
||||
SdotByElement 25, 6, 9,1
|
||||
SdotByElement 26, 6,10,1
|
||||
SdotByElement 27, 6,11,1
|
||||
SdotByElement 28, 7, 8,1
|
||||
SdotByElement 29, 7, 9,1
|
||||
SdotByElement 30, 7,10,1
|
||||
SdotByElement 31, 7,11,1
|
||||
|
||||
tst x7,15
|
||||
b.ne InChannels8 // 4 ~ 12 InputChannels
|
||||
|
||||
subs x9,x9,8 // KernelSize-=1
|
||||
b.hi KernelSizeLoop
|
||||
|
||||
Requantize:
|
||||
tst w10,#.LMLAS_CONV_SYM_FLAG_PER_CHANNEL_SCALE
|
||||
ldr w13,[x8,#.LConvSymPostProcessParams_ZeroPoint]
|
||||
beq BroadcastScaleValue
|
||||
ldp q0,q1,[x19],32 // load scale vector
|
||||
ldp q2,q3,[x19],32
|
||||
b AccumulatorsToFloat
|
||||
|
||||
BroadcastScaleValue:
|
||||
ld1r {v0.4s},[x19] // load scale Value
|
||||
mov v1.16b, v0.16b
|
||||
mov v2.16b, v0.16b
|
||||
mov v3.16b, v0.16b
|
||||
|
||||
AccumulatorsToFloat:
|
||||
scvtf v16.4s,v16.4s // convert to float
|
||||
scvtf v17.4s,v17.4s
|
||||
scvtf v18.4s,v18.4s
|
||||
scvtf v19.4s,v19.4s
|
||||
scvtf v20.4s,v20.4s
|
||||
scvtf v21.4s,v21.4s
|
||||
scvtf v22.4s,v22.4s
|
||||
scvtf v23.4s,v23.4s
|
||||
scvtf v24.4s,v24.4s
|
||||
scvtf v25.4s,v25.4s
|
||||
scvtf v26.4s,v26.4s
|
||||
scvtf v27.4s,v27.4s
|
||||
scvtf v28.4s,v28.4s
|
||||
scvtf v29.4s,v29.4s
|
||||
scvtf v30.4s,v30.4s
|
||||
scvtf v31.4s,v31.4s
|
||||
fmul v16.4s,v16.4s,v0.4s // multiply by scale
|
||||
fmul v17.4s,v17.4s,v0.4s
|
||||
fmul v18.4s,v18.4s,v0.4s
|
||||
fmul v19.4s,v19.4s,v0.4s
|
||||
fmul v20.4s,v20.4s,v1.4s
|
||||
fmul v21.4s,v21.4s,v1.4s
|
||||
fmul v22.4s,v22.4s,v1.4s
|
||||
fmul v23.4s,v23.4s,v1.4s
|
||||
fmul v24.4s,v24.4s,v2.4s
|
||||
fmul v25.4s,v25.4s,v2.4s
|
||||
fmul v26.4s,v26.4s,v2.4s
|
||||
fmul v27.4s,v27.4s,v2.4s
|
||||
fmul v28.4s,v28.4s,v3.4s
|
||||
fmul v29.4s,v29.4s,v3.4s
|
||||
fmul v30.4s,v30.4s,v3.4s
|
||||
fmul v31.4s,v31.4s,v3.4s
|
||||
fcvtns v16.4s,v16.4s // convert to int
|
||||
fcvtns v17.4s,v17.4s
|
||||
fcvtns v18.4s,v18.4s
|
||||
fcvtns v19.4s,v19.4s
|
||||
fcvtns v20.4s,v20.4s
|
||||
fcvtns v21.4s,v21.4s
|
||||
fcvtns v22.4s,v22.4s
|
||||
fcvtns v23.4s,v23.4s
|
||||
fcvtns v24.4s,v24.4s
|
||||
fcvtns v25.4s,v25.4s
|
||||
fcvtns v26.4s,v26.4s
|
||||
fcvtns v27.4s,v27.4s
|
||||
fcvtns v28.4s,v28.4s
|
||||
fcvtns v29.4s,v29.4s
|
||||
fcvtns v30.4s,v30.4s
|
||||
fcvtns v31.4s,v31.4s
|
||||
|
||||
sqxtn v16.4h,v16.4s
|
||||
sqxtn v17.4h,v17.4s
|
||||
sqxtn v18.4h,v18.4s
|
||||
sqxtn v19.4h,v19.4s
|
||||
sqxtn v24.4h,v24.4s
|
||||
sqxtn v25.4h,v25.4s
|
||||
sqxtn v26.4h,v26.4s
|
||||
sqxtn v27.4h,v27.4s
|
||||
dup v4.8h,w13 // zero point
|
||||
sqxtn2 v16.8h,v20.4s
|
||||
sqxtn2 v17.8h,v21.4s
|
||||
sqxtn2 v18.8h,v22.4s
|
||||
sqxtn2 v19.8h,v23.4s
|
||||
sqxtn2 v24.8h,v28.4s
|
||||
sqxtn2 v25.8h,v29.4s
|
||||
sqxtn2 v26.8h,v30.4s
|
||||
sqxtn2 v27.8h,v31.4s
|
||||
sqadd v16.8h,v16.8h,v4.8h
|
||||
sqadd v17.8h,v17.8h,v4.8h
|
||||
sqadd v18.8h,v18.8h,v4.8h
|
||||
sqadd v19.8h,v19.8h,v4.8h
|
||||
sqadd v24.8h,v24.8h,v4.8h
|
||||
sqadd v25.8h,v25.8h,v4.8h
|
||||
sqadd v26.8h,v26.8h,v4.8h
|
||||
sqadd v27.8h,v27.8h,v4.8h
|
||||
sqxtun v0.8b,v16.8h
|
||||
sqxtun v1.8b,v17.8h
|
||||
sqxtun v2.8b,v18.8h
|
||||
sqxtun v3.8b,v19.8h
|
||||
sqxtun2 v0.16b,v24.8h
|
||||
sqxtun2 v1.16b,v25.8h
|
||||
subs x6,x6,16 // processed 16 output channels
|
||||
sqxtun2 v2.16b,v26.8h
|
||||
sqxtun2 v3.16b,v27.8h
|
||||
b.lo PartialStore
|
||||
|
||||
st1 {v3.16b},[x5],16 // Store full 4 x 16
|
||||
st1 {v2.16b},[x17],16
|
||||
sub x0,x0,x3 // Restore pointer to A: a -= ks
|
||||
st1 {v1.16b},[x16],16
|
||||
st1 {v0.16b},[x2],16
|
||||
b.hi OutputChannelLoop
|
||||
|
||||
ExitKernel:
|
||||
ldp x19,x20,[sp,#48]
|
||||
ldp d12,d13,[sp,#32]
|
||||
ldp d10,d11,[sp,#16]
|
||||
ldp d8,d9,[sp],#.LConvSymFrame_SavedRegisters
|
||||
ret
|
||||
|
||||
InChannels8:
|
||||
tbz x7,3,InChannels4
|
||||
ldr d0,[x12],8
|
||||
ldr q4,[x1],16
|
||||
ldr d1,[x13],8
|
||||
ldr d2,[x14],8
|
||||
ldr d3,[x15],8
|
||||
eor v0.8b,v0.8b,v12.8b
|
||||
ldr q5,[x1],16
|
||||
eor v1.8b,v1.8b,v12.8b
|
||||
SdotByElement 16, 4, 0,0
|
||||
SdotByElement 17, 4, 1,0
|
||||
eor v2.8b,v2.8b,v12.8b
|
||||
ldp q6, q7, [x1], 32
|
||||
eor v3.8b,v3.8b,v12.8b
|
||||
SdotByElement 18, 4, 2,0
|
||||
SdotByElement 19, 4, 3,0
|
||||
SdotByElement 20, 5, 0,0
|
||||
SdotByElement 21, 5, 1,0
|
||||
SdotByElement 22, 5, 2,0
|
||||
SdotByElement 23, 5, 3,0
|
||||
SdotByElement 24, 6, 0,0
|
||||
SdotByElement 25, 6, 1,0
|
||||
ldp q4, q5, [x1], 32
|
||||
SdotByElement 26, 6, 2,0
|
||||
SdotByElement 27, 6, 3,0
|
||||
SdotByElement 28, 7, 0,0
|
||||
SdotByElement 29, 7, 1,0
|
||||
SdotByElement 30, 7, 2,0
|
||||
SdotByElement 31, 7, 3,0
|
||||
SdotByElement 16, 4, 0,1
|
||||
SdotByElement 17, 4, 1,1
|
||||
ldp q6, q7, [x1], 32
|
||||
SdotByElement 18, 4, 2,1
|
||||
SdotByElement 19, 4, 3,1
|
||||
SdotByElement 20, 5, 0,1
|
||||
SdotByElement 21, 5, 1,1
|
||||
SdotByElement 22, 5, 2,1
|
||||
SdotByElement 23, 5, 3,1
|
||||
SdotByElement 24, 6, 0,1
|
||||
SdotByElement 25, 6, 1,1
|
||||
SdotByElement 26, 6, 2,1
|
||||
SdotByElement 27, 6, 3,1
|
||||
SdotByElement 28, 7, 0,1
|
||||
SdotByElement 29, 7, 1,1
|
||||
SdotByElement 30, 7, 2,1
|
||||
SdotByElement 31, 7, 3,1
|
||||
tbz x7,2,SkipInCh4
|
||||
|
||||
InChannels4:
|
||||
ldr s0,[x12],4
|
||||
ldr q4,[x1],16
|
||||
ldr s1,[x13],4
|
||||
ldr s2,[x14],4
|
||||
ldr s3,[x15],4
|
||||
eor v0.8b,v0.8b,v12.8b
|
||||
ldr q5, [x1], 16
|
||||
eor v1.8b,v1.8b,v12.8b
|
||||
SdotByElement 16, 4, 0,0
|
||||
SdotByElement 17, 4, 1,0
|
||||
eor v2.8b,v2.8b,v12.8b
|
||||
ldp q6, q7, [x1], 32
|
||||
eor v3.8b,v3.8b,v12.8b
|
||||
SdotByElement 18, 4, 2,0
|
||||
SdotByElement 19, 4, 3,0
|
||||
SdotByElement 20, 5, 0,0
|
||||
SdotByElement 21, 5, 1,0
|
||||
SdotByElement 22, 5, 2,0
|
||||
SdotByElement 23, 5, 3,0
|
||||
SdotByElement 24, 6, 0,0
|
||||
SdotByElement 25, 6, 1,0
|
||||
SdotByElement 26, 6, 2,0
|
||||
SdotByElement 27, 6, 3,0
|
||||
SdotByElement 28, 7, 0,0
|
||||
SdotByElement 29, 7, 1,0
|
||||
SdotByElement 30, 7, 2,0
|
||||
SdotByElement 31, 7, 3,0
|
||||
|
||||
SkipInCh4:
|
||||
subs x9,x9,8 // ks -= 1
|
||||
b.hi KernelSizeLoop
|
||||
b Requantize
|
||||
|
||||
PartialStore:
|
||||
tbz x6,3,LT8Store
|
||||
str d3,[x5],8 // no less than 8 channels
|
||||
str d2,[x17],8
|
||||
dup d3,v3.d[1]
|
||||
dup d2,v2.d[1]
|
||||
str d1,[x16],8
|
||||
str d0,[x2],8
|
||||
dup d1,v1.d[1]
|
||||
dup d0,v0.d[1]
|
||||
LT8Store:
|
||||
tbz x6,2,LT4Store
|
||||
str s3,[x5],4
|
||||
str s2,[x17],4
|
||||
dup s3,v3.s[1]
|
||||
dup s2,v2.s[1]
|
||||
str s1,[x16],4
|
||||
str s0,[x2],4
|
||||
dup s1,v1.s[1]
|
||||
dup s0,v0.s[1]
|
||||
LT4Store:
|
||||
tbz x6,1, LT2Store
|
||||
str h3,[x5],2
|
||||
str h2,[x17],2
|
||||
dup h3,v3.h[1]
|
||||
dup h2,v2.h[1]
|
||||
str h1,[x16],2
|
||||
str h0,[x2],2
|
||||
dup h1,v1.h[1]
|
||||
dup h0,v0.h[1]
|
||||
LT2Store:
|
||||
tbz x6,0,ExitKernel
|
||||
str b3,[x5]
|
||||
str b2,[x17]
|
||||
str b1,[x16]
|
||||
str b0,[x2]
|
||||
b ExitKernel
|
||||
|
||||
.end
|
||||
454
onnxruntime/core/mlas/lib/aarch64/ConvSymU8KernelNeon.S
Normal file
454
onnxruntime/core/mlas/lib/aarch64/ConvSymU8KernelNeon.S
Normal file
|
|
@ -0,0 +1,454 @@
|
|||
/*++
|
||||
|
||||
Copyright (c) Microsoft Corporation. All rights reserved.
|
||||
|
||||
Licensed under the MIT License.
|
||||
|
||||
Module Name:
|
||||
|
||||
ConvSymKernelNeon.s
|
||||
|
||||
Abstract:
|
||||
|
||||
This module implements the kernels for the symmetric quantized integer
|
||||
convolution operation.
|
||||
|
||||
--*/
|
||||
|
||||
#include "asmmacro.h"
|
||||
|
||||
.equ .LMLAS_CONV_SYM_FLAG_INPUT_DIRECT, 1
|
||||
.equ .LMLAS_CONV_SYM_FLAG_PER_CHANNEL_SCALE, 2
|
||||
|
||||
//
|
||||
// Stack frame layout for the symmetric convolution kernel.
|
||||
// d8-d15, x19-x30 need to be preserved if used
|
||||
//
|
||||
.equ .LConvSymFrame_SavedNeonRegisters, (8 * 8)
|
||||
.equ .LConvSymFrame_SavedRegisters, .LConvSymFrame_SavedNeonRegisters
|
||||
.equ .LConvSymFrame_PostProcessParams, 0 + .LConvSymFrame_SavedRegisters
|
||||
.equ .LConvSymFrame_KernelFlags, 8 + .LConvSymFrame_SavedRegisters
|
||||
|
||||
.equ .LConvSymPostProcessParams_Bias, 0
|
||||
.equ .LConvSymPostProcessParams_Scale, 8
|
||||
.equ .LConvSymPostProcessParams_Min, 16
|
||||
.equ .LConvSymPostProcessParams_Max, 20
|
||||
.equ .LConvSymPostProcessParams_ZeroPoint, 24
|
||||
|
||||
.text
|
||||
|
||||
/*++
|
||||
|
||||
Routine Description:
|
||||
|
||||
This routine is the inner kernel to compute a convolution for the elements
|
||||
of an output row for a set of filter rows.
|
||||
|
||||
Arguments:
|
||||
|
||||
Input (x0) - Supplies the address of the input buffer.
|
||||
|
||||
If MLAS_CONV_SYM_FLAG_INPUT_DIRECT is set, then the input buffer points
|
||||
directly at the input tensor.
|
||||
|
||||
If MLAS_CONV_SYM_FLAG_INPUT_DIRECT is clear, then the input buffer is an
|
||||
indirection buffer. Every pointer in the indirection buffer points at a
|
||||
InputChannels length vector (either from the input tensor or a vector of
|
||||
padding values). These are grouped in batches of length KernelSize.
|
||||
These batches are then repeated OutputCount times.
|
||||
|
||||
Filter (x1) - Supplies the address of the filter buffer.
|
||||
|
||||
Output (x2) - Supplies the address of the output buffer.
|
||||
|
||||
KernelSize (x3) - Supplies the size of the kernel.
|
||||
|
||||
If MLAS_CONV_SYM_FLAG_INPUT_DIRECT is set, then kernel size should be 1.
|
||||
|
||||
InputChannels (x4) - Supplies the number of input channels.
|
||||
|
||||
This implementation requires the count to be a multiple of 8.
|
||||
|
||||
OutputChannels (x5) - Supplies the number of output channels.
|
||||
|
||||
ChannelCount (x6) - Supplies the number of channels this iteration produces.
|
||||
|
||||
This implementation requires the count to be 8.
|
||||
|
||||
OutputCount (x7) - Supplies the number of output elements this iteration produces.
|
||||
|
||||
This implementation requires the count to be 1 or 2.
|
||||
|
||||
PostProcessParams - Supplies the address of the post process parameter block.
|
||||
|
||||
KernelFlags - Supplies additional flags controlling the operation.
|
||||
|
||||
Return Value:
|
||||
|
||||
None.
|
||||
|
||||
--*/
|
||||
FUNCTION_ENTRY MlasConvSymKernelNeon
|
||||
|
||||
stp d8,d9,[sp,#-64]!
|
||||
ldr x8,[sp,#.LConvSymFrame_PostProcessParams]
|
||||
ldrb w10,[sp,#.LConvSymFrame_KernelFlags]
|
||||
stp d10,d11,[sp,#16]
|
||||
stp d12,d13,[sp,#32]
|
||||
stp d14,d15,[sp,#48]
|
||||
mov x9,x3 // save kernel size
|
||||
ldr x11,[x8,#.LConvSymPostProcessParams_Bias]
|
||||
mov x16,x4 // save input channels
|
||||
ldr x12,[x8,#.LConvSymPostProcessParams_Scale]
|
||||
cmp x7,2 // if OutputCount < 2
|
||||
add x5,x2,x5 // c1 = c0 + ldc
|
||||
add x4,x4,7 // kc = (kc + 7) & ~7
|
||||
csel x5,x2,x5,lo // if OutputCount < 2 c1 = c0
|
||||
bic x4,x4,7
|
||||
ldp s16,s18,[x11],8 // init accumulators with bias
|
||||
ldp s20,s22,[x11],8
|
||||
ldp s24,s26,[x11],8
|
||||
ldp s28,s30,[x11],8
|
||||
mov v17.16b,v16.16b
|
||||
mov v19.16b,v18.16b
|
||||
mov v21.16b,v20.16b
|
||||
mov v23.16b,v22.16b
|
||||
mov v25.16b,v24.16b
|
||||
mov v27.16b,v26.16b
|
||||
mov v29.16b,v28.16b
|
||||
mov v31.16b,v30.16b
|
||||
|
||||
// Nested loops, inner loop: input channel; outter loop: kernel size
|
||||
// Each inner iteration processes 8 input channels, 2 output pixels, 8 output channels.
|
||||
//
|
||||
// B 8x8
|
||||
// ------------------------------------------------------------------
|
||||
// |v4.b[0] v5.b[0] v4.b[0] v5.b[0] v4.b[0] v5.b[0] v4.b[0] v5.b[0] |
|
||||
// | ... ... ... ... ... ... ... ... |
|
||||
// |v4.b[7] v5.b[7] v4.b[7] v5.b[7] v4.b[7] v5.b[7] v4.b[7] v5.b[7] |
|
||||
// A 2x8 ------------------------------------------------------------------
|
||||
// ------------------ ------------------------------------------------------------------
|
||||
// x13-> |v0.b[0]..v0.b[7]| |v16.4s v18.4s v20.4s v22.4s v24.4s v26.4s v28.4s v30.4s |
|
||||
// x15-> |v1.b[0]..v1.b[7]| |v17.4s v19.4s v21.4s v23.4s v25.4s v27.4s v29.4s v31.4s |
|
||||
// ------------------ ------------------------------------------------------------------
|
||||
// When Input Channels greater than 16, unroll:
|
||||
// A registers v6 v7,
|
||||
// B registers v8 v9
|
||||
//
|
||||
|
||||
.LConvSym.KernelSizeLoop:
|
||||
|
||||
# Load next 2 A pointers
|
||||
tst w10,#.LMLAS_CONV_SYM_FLAG_INPUT_DIRECT
|
||||
ldr d4,[x1]
|
||||
ldr d5,[x1,8]
|
||||
beq .LConvSym.InputIndirection
|
||||
|
||||
.LConvSym.InputDirect:
|
||||
mov x13,x0 // x13 -> A0
|
||||
add x15,x0,x16 // x15 -> A1 = A0 + input channels
|
||||
b .LConvSym.BlockLoopPrologue
|
||||
|
||||
.LConvSym.InputIndirection:
|
||||
cmp x7,2 // test if OutputCount < 2
|
||||
ldr x13,[x0] // x13 -> A0
|
||||
blo .LConvSym.SkipLoadA1
|
||||
ldr x15,[x0,x3,lsl#3] // x15 -> A1
|
||||
.LConvSym.SkipLoadA1:
|
||||
|
||||
.LConvSym.BlockLoopPrologue:
|
||||
cmp x7,2 // test if OutputCount < 2
|
||||
add x0,x0,8 // indirect A advance to next pointer, prepare for kernel size loop
|
||||
csel x15,x13,x15,lo // if OutputCount < 2 x15 -> A0
|
||||
subs x14,x4,16 // input channel - 16
|
||||
movi v12.8b,128
|
||||
blo .LConvSym.8InputChannels // less than 16 deep, no unroll
|
||||
|
||||
ldr d0,[x13],8
|
||||
ldr d1,[x15],8
|
||||
ldr d8,[x1,64]
|
||||
ldr d9,[x1,72]
|
||||
ldr d6,[x13],8
|
||||
subs x14,x14,16 // input channel - 16
|
||||
ldr d7,[x15],8
|
||||
blo .LConvSym.BlockLoopEpilogue // need 32 input channel for full unrolled loop
|
||||
|
||||
.LConvSym.Blockloop:
|
||||
eor v0.8b,v0.8b,v12.8b
|
||||
eor v1.8b,v1.8b,v12.8b
|
||||
smull v2.8h,v4.8b,v0.8b
|
||||
smull v3.8h,v4.8b,v1.8b
|
||||
ldr d4,[x1,16]
|
||||
smull v10.8h,v5.8b,v0.8b
|
||||
smull v11.8h,v5.8b,v1.8b
|
||||
ldr d5,[x1,24]
|
||||
eor v6.8b,v6.8b,v12.8b
|
||||
eor v7.8b,v7.8b,v12.8b
|
||||
smlal v2.8h,v8.8b,v6.8b
|
||||
smlal v3.8h,v8.8b,v7.8b
|
||||
ldr d8,[x1,80]
|
||||
smlal v10.8h,v9.8b,v6.8b
|
||||
smlal v11.8h,v9.8b,v7.8b
|
||||
ldr d9,[x1,88]
|
||||
smull v12.8h,v4.8b,v0.8b
|
||||
sadalp v16.4s,v2.8h
|
||||
smull v13.8h,v4.8b,v1.8b
|
||||
ldr d4,[x1,32]
|
||||
sadalp v17.4s,v3.8h
|
||||
smull v14.8h,v5.8b,v0.8b
|
||||
sadalp v18.4s,v10.8h
|
||||
smull v15.8h,v5.8b,v1.8b
|
||||
ldr d5,[x1,40]
|
||||
sadalp v19.4s,v11.8h
|
||||
smlal v12.8h,v8.8b,v6.8b
|
||||
smlal v13.8h,v8.8b,v7.8b
|
||||
ldr d8,[x1,96]
|
||||
smlal v14.8h,v9.8b,v6.8b
|
||||
smlal v15.8h,v9.8b,v7.8b
|
||||
ldr d9,[x1,104]
|
||||
smull v2.8h,v4.8b,v0.8b
|
||||
sadalp v20.4s,v12.8h
|
||||
smull v3.8h,v4.8b,v1.8b
|
||||
ldr d4,[x1,48]
|
||||
sadalp v21.4s,v13.8h
|
||||
smull v10.8h,v5.8b,v0.8b
|
||||
sadalp v22.4s,v14.8h
|
||||
smull v11.8h,v5.8b,v1.8b
|
||||
ldr d5,[x1,56]
|
||||
sadalp v23.4s, v15.8h
|
||||
smlal v2.8h,v8.8b,v6.8b
|
||||
smlal v3.8h,v8.8b,v7.8b
|
||||
ldr d8,[x1,112]
|
||||
smlal v10.8h,v9.8b,v6.8b
|
||||
smlal v11.8h,v9.8b,v7.8b
|
||||
ldr d9,[x1,120]
|
||||
smull v12.8h,v4.8b,v0.8b
|
||||
add x1,x1,128
|
||||
sadalp v24.4s,v2.8h
|
||||
smull v13.8h,v4.8b,v1.8b
|
||||
ldr d4,[x1] // Read B
|
||||
sadalp v25.4s,v3.8h
|
||||
smull v14.8h,v5.8b,v0.8b
|
||||
ldr d0,[x13],8 // Read A0
|
||||
sadalp v26.4s,v10.8h
|
||||
smull v15.8h,v5.8b,v1.8b
|
||||
ldr d1,[x15],8 // Read A1
|
||||
sadalp v27.4s,v11.8h
|
||||
smlal v12.8h,v8.8b,v6.8b
|
||||
ldr d5,[x1,8] // Read B
|
||||
smlal v13.8h,v8.8b,v7.8b
|
||||
ldr d8,[x1,64] // Read B
|
||||
smlal v14.8h,v9.8b,v6.8b
|
||||
ldr d6,[x13],8 // Read A0
|
||||
smlal v15.8h,v9.8b,v7.8b
|
||||
ldr d7,[x15],8 // Read A1
|
||||
sadalp v28.4s,v12.8h
|
||||
ldr d9,[x1,72] // Read B
|
||||
sadalp v29.4s,v13.8h
|
||||
subs x14,x14,16
|
||||
sadalp v30.4s,v14.8h
|
||||
movi v12.8b,128
|
||||
sadalp v31.4s,v15.8h
|
||||
b.hs .LConvSym.Blockloop
|
||||
|
||||
.LConvSym.BlockLoopEpilogue: // remaining 16 input channels
|
||||
eor v0.8b,v0.8b,v12.8b
|
||||
eor v1.8b,v1.8b,v12.8b
|
||||
smull v2.8h,v4.8b,v0.8b
|
||||
smull v3.8h,v4.8b,v1.8b
|
||||
ldr d4,[x1,16]
|
||||
smull v10.8h,v5.8b,v0.8b
|
||||
smull v11.8h,v5.8b,v1.8b
|
||||
ldr d5,[x1,24]
|
||||
eor v6.8b,v6.8b,v12.8b
|
||||
eor v7.8b,v7.8b,v12.8b
|
||||
smlal v2.8h,v8.8b,v6.8b
|
||||
smlal v3.8h,v8.8b,v7.8b
|
||||
ldr d8,[x1,80]
|
||||
smlal v10.8h,v9.8b,v6.8b
|
||||
smlal v11.8h,v9.8b,v7.8b
|
||||
ldr d9,[x1,88]
|
||||
smull v12.8h,v4.8b,v0.8b
|
||||
sadalp v16.4s,v2.8h
|
||||
smull v13.8h,v4.8b,v1.8b
|
||||
ldr d4,[x1,32]
|
||||
sadalp v17.4s,v3.8h
|
||||
smull v14.8h,v5.8b,v0.8b
|
||||
sadalp v18.4s,v10.8h
|
||||
smull v15.8h,v5.8b,v1.8b
|
||||
sadalp v19.4s,v11.8h
|
||||
ldr d5,[x1,40]
|
||||
smlal v12.8h,v8.8b,v6.8b
|
||||
smlal v13.8h,v8.8b,v7.8b
|
||||
ldr d8,[x1,96]
|
||||
smlal v14.8h,v9.8b,v6.8b
|
||||
smlal v15.8h,v9.8b,v7.8b
|
||||
ldr d9,[x1,104]
|
||||
smull v2.8h,v4.8b,v0.8b
|
||||
sadalp v20.4s,v12.8h
|
||||
smull v3.8h,v4.8b,v1.8b
|
||||
ldr d4,[x1,48]
|
||||
sadalp v21.4s,v13.8h
|
||||
smull v10.8h,v5.8b,v0.8b
|
||||
sadalp v22.4s,v14.8h
|
||||
smull v11.8h,v5.8b,v1.8b
|
||||
sadalp v23.4s,v15.8h
|
||||
ldr d5,[x1,56]
|
||||
smlal v2.8h,v8.8b,v6.8b
|
||||
smlal v3.8h,v8.8b,v7.8b
|
||||
ldr d8,[x1,112]
|
||||
smlal v10.8h,v9.8b,v6.8b
|
||||
smlal v11.8h,v9.8b,v7.8b
|
||||
ldr d9,[x1,120]
|
||||
smull v12.8h,v4.8b,v0.8b
|
||||
sadalp v24.4s,v2.8h
|
||||
smull v13.8h,v4.8b,v1.8b
|
||||
sadalp v25.4s,v3.8h
|
||||
smull v14.8h,v5.8b,v0.8b
|
||||
sadalp v26.4s,v10.8h
|
||||
smull v15.8h,v5.8b,v1.8b
|
||||
sadalp v27.4s,v11.8h
|
||||
smlal v12.8h,v8.8b,v6.8b
|
||||
smlal v13.8h,v8.8b,v7.8b
|
||||
smlal v14.8h,v9.8b,v6.8b
|
||||
smlal v15.8h,v9.8b,v7.8b
|
||||
add x1,x1,128
|
||||
|
||||
sadalp v28.4s,v12.8h
|
||||
sadalp v29.4s,v13.8h
|
||||
sadalp v30.4s,v14.8h
|
||||
sadalp v31.4s,v15.8h
|
||||
movi v12.8b,128
|
||||
tbnz x14,3,.LConvSym.8InputChannels
|
||||
|
||||
subs x9,x9,1
|
||||
b.hi .LConvSym.KernelSizeLoop
|
||||
|
||||
.LConvSym.Requantize:
|
||||
ldr w11, [x8, #.LConvSymPostProcessParams_ZeroPoint]
|
||||
tst w10,#.LMLAS_CONV_SYM_FLAG_PER_CHANNEL_SCALE
|
||||
beq .LConvSym.BroadcastScaleValue
|
||||
ld1 {v4.4s,v5.4s},[x12] // load scale vector
|
||||
b .LConvSym.AccumulatorsToFloat
|
||||
|
||||
.LConvSym.BroadcastScaleValue:
|
||||
ld1r {v4.4s},[x12] // load scale Value
|
||||
mov v5.16b, v4.16b
|
||||
|
||||
.LConvSym.AccumulatorsToFloat:
|
||||
addp v16.4s,v16.4s,v18.4s
|
||||
addp v20.4s,v20.4s,v22.4s
|
||||
addp v24.4s,v24.4s,v26.4s
|
||||
addp v28.4s,v28.4s,v30.4s
|
||||
addp v17.4s,v17.4s,v19.4s
|
||||
addp v21.4s,v21.4s,v23.4s
|
||||
addp v25.4s,v25.4s,v27.4s
|
||||
addp v29.4s,v29.4s,v31.4s
|
||||
addp v0.4s,v16.4s,v20.4s
|
||||
addp v1.4s,v24.4s,v28.4s
|
||||
addp v2.4s,v17.4s,v21.4s
|
||||
addp v3.4s,v25.4s,v29.4s
|
||||
scvtf v0.4s,v0.4s // convert to float
|
||||
scvtf v1.4s,v1.4s
|
||||
scvtf v2.4s,v2.4s
|
||||
scvtf v3.4s,v3.4s
|
||||
fmul v0.4s,v0.4s,v4.4s // multiply by scale
|
||||
fmul v1.4s,v1.4s,v5.4s
|
||||
fmul v2.4s,v2.4s,v4.4s
|
||||
fmul v3.4s,v3.4s,v5.4s
|
||||
fcvtns v0.4s,v0.4s // convert to int
|
||||
fcvtns v1.4s,v1.4s
|
||||
dup v9.8h,w11
|
||||
fcvtns v2.4s,v2.4s
|
||||
fcvtns v3.4s,v3.4s
|
||||
sqxtn v0.4h,v0.4s
|
||||
sqxtn2 v0.8h,v1.4s
|
||||
sqxtn v2.4h,v2.4s
|
||||
sqxtn2 v2.8h,v3.4s
|
||||
subs x6, x6, 8
|
||||
sqadd v0.8h,v0.8h,v9.8h
|
||||
sqadd v2.8h,v2.8h,v9.8h
|
||||
sqxtun v0.8b,v0.8h // shorten to int8
|
||||
sqxtun2 v0.16b,v2.8h
|
||||
b.lo .LConvSym.PartialStore
|
||||
|
||||
st1 {v0.d}[1],[x5] // full 2x8 store to c
|
||||
st1 {v0.8b},[x2]
|
||||
|
||||
.LConvSym.ExitKernel:
|
||||
ldp d14,d15,[sp,#48]
|
||||
ldp d12,d13,[sp,#32]
|
||||
ldp d10,d11,[sp,#16]
|
||||
ldp d8,d9,[sp],#64
|
||||
ret
|
||||
|
||||
.LConvSym.8InputChannels:
|
||||
ldr d0,[x13]
|
||||
ldr d1,[x15]
|
||||
ldr d4,[x1]
|
||||
ldr d5,[x1,8]
|
||||
ldr d6,[x1,16]
|
||||
ldr d7,[x1,24]
|
||||
eor v0.8b,v0.8b,v12.8b
|
||||
eor v1.8b,v1.8b,v12.8b
|
||||
smull v2.8h,v4.8b,v0.8b
|
||||
smull v3.8h,v4.8b,v1.8b
|
||||
ldr d4,[x1,32]
|
||||
smull v10.8h,v5.8b,v0.8b
|
||||
smull v11.8h,v5.8b,v1.8b
|
||||
ldr d5,[x1,40]
|
||||
smull v12.8h,v6.8b,v0.8b
|
||||
sadalp v16.4s,v2.8h
|
||||
smull v13.8h,v6.8b,v1.8b
|
||||
ldr d6,[x1,48]
|
||||
sadalp v17.4s,v3.8h
|
||||
smull v14.8h,v7.8b,v0.8b
|
||||
sadalp v18.4s,v10.8h
|
||||
smull v15.8h,v7.8b,v1.8b
|
||||
ldr d7,[x1,56]
|
||||
sadalp v19.4s,v11.8h
|
||||
smull v2.8h,v4.8b,v0.8b
|
||||
sadalp v20.4s,v12.8h
|
||||
smull v3.8h,v4.8b,v1.8b
|
||||
sadalp v21.4s,v13.8h
|
||||
smull v10.8h,v5.8b,v0.8b
|
||||
sadalp v22.4s,v14.8h
|
||||
smull v11.8h,v5.8b,v1.8b
|
||||
sadalp v23.4s,v15.8h
|
||||
smull v12.8h,v6.8b,v0.8b
|
||||
sadalp v24.4s,v2.8h
|
||||
smull v13.8h,v6.8b,v1.8b
|
||||
sadalp v25.4s,v3.8h
|
||||
smull v14.8h,v7.8b,v0.8b
|
||||
sadalp v26.4s,v10.8h
|
||||
smull v15.8h,v7.8b,v1.8b
|
||||
sadalp v27.4s,v11.8h
|
||||
add x1,x1,64
|
||||
sadalp v28.4s,v12.8h
|
||||
sadalp v29.4s,v13.8h
|
||||
sadalp v30.4s,v14.8h
|
||||
sadalp v31.4s,v15.8h
|
||||
|
||||
# ks loop
|
||||
subs x9,x9,1
|
||||
b.hi .LConvSym.KernelSizeLoop
|
||||
b .LConvSym.Requantize
|
||||
|
||||
.LConvSym.PartialStore:
|
||||
tbz x6,2,.LConvSym.Store2
|
||||
st1 {v0.s}[2],[x5],4
|
||||
str s0,[x2],4
|
||||
EXT v0.16b,v0.16b,v0.16b,4
|
||||
|
||||
.LConvSym.Store2:
|
||||
tbz x6, 1, .LConvSym.Store1
|
||||
st1 {v0.h}[4], [x5], 2
|
||||
str h0, [x2], 2
|
||||
EXT v0.16b,v0.16b,v0.16b,2
|
||||
.LConvSym.Store1:
|
||||
tbz x6,0,.LConvSym.ExitKernel
|
||||
st1 {v0.b}[8],[x5]
|
||||
str b0,[x2]
|
||||
b .LConvSym.ExitKernel
|
||||
|
||||
.end
|
||||
631
onnxruntime/core/mlas/lib/arm64/ConvSymU8KernelDot.asm
Normal file
631
onnxruntime/core/mlas/lib/arm64/ConvSymU8KernelDot.asm
Normal file
|
|
@ -0,0 +1,631 @@
|
|||
/*++
|
||||
|
||||
Copyright (c) Microsoft Corporation. All rights reserved.
|
||||
|
||||
Licensed under the MIT License.
|
||||
|
||||
Module Name:
|
||||
|
||||
ConvSymKernelNeonDot.asm
|
||||
|
||||
Abstract:
|
||||
|
||||
This module implements the kernels for the symmetric quantized integer
|
||||
convolution operation.
|
||||
|
||||
--*/
|
||||
|
||||
#include "kxarm64.h"
|
||||
|
||||
#define MLAS_CONV_SYM_FLAG_INPUT_DIRECT 1
|
||||
#define MLAS_CONV_SYM_FLAG_PER_CHANNEL_SCALE 2
|
||||
|
||||
//
|
||||
// Stack frame layout for the symmetric convolution kernel.
|
||||
// d8-d15, x19-x30 need to be preserved if used
|
||||
//
|
||||
#define ConvSymFrame_SavedNeonRegisters (8 * 8)
|
||||
#define ConvSymFrame_SavedRegisters ConvSymFrame_SavedNeonRegisters
|
||||
#define ConvSymFrame_PostProcessParams 0 + ConvSymFrame_SavedRegisters
|
||||
#define ConvSymFrame_KernelFlags 8 + ConvSymFrame_SavedRegisters
|
||||
|
||||
#define ConvSymPostProcessParams_Bias 0
|
||||
#define ConvSymPostProcessParams_Scale 8
|
||||
#define ConvSymPostProcessParams_Min 16
|
||||
#define ConvSymPostProcessParams_Max 20
|
||||
#define ConvSymPostProcessParams_ZeroPoint 24
|
||||
|
||||
TEXTAREA
|
||||
|
||||
/*++
|
||||
|
||||
Routine Description:
|
||||
|
||||
This routine is the inner kernel to compute a convolution for the elements
|
||||
of an output row for a set of filter rows.
|
||||
|
||||
Arguments:
|
||||
|
||||
Input (x0) - Points to the input buffer.
|
||||
|
||||
If MLAS_CONV_SYM_FLAG_INPUT_DIRECT is set, then the input buffer points
|
||||
directly at the input tensor.
|
||||
|
||||
If MLAS_CONV_SYM_FLAG_INPUT_DIRECT is clear, then the input buffer is an
|
||||
indirection buffer. Every pointer in the indirection buffer points at a
|
||||
InputChannels length vector (either from the input tensor or a vector of
|
||||
padding values). These are grouped in batches of length KernelSize.
|
||||
These batches are then repeated OutputCount times.
|
||||
|
||||
Filter (x1) - Points to the filter buffer.
|
||||
|
||||
Output (x2) - Points the output buffer.
|
||||
|
||||
KernelSize (x3/x9) - Size of the kernel (most commonly. 3x3=9, 5x5=25).
|
||||
|
||||
If MLAS_CONV_SYM_FLAG_INPUT_DIRECT is set, then kernel size should be 1.
|
||||
|
||||
InputChannels (x4/x7) - Number of input channels.
|
||||
|
||||
OutputChannels (x5) - Number of output channels.
|
||||
|
||||
ChannelCount (x6) - Number of output channels this iteration produces.
|
||||
|
||||
OutputCount (x7) - Number of output elements this iteration produces.
|
||||
|
||||
This implementation requires the count to be no larger than 4.
|
||||
|
||||
PostProcessParams (x8) - Points to the post process parameter block.
|
||||
|
||||
KernelFlags - (w10) Additional flags controlling the operation.
|
||||
|
||||
Return Value:
|
||||
|
||||
None.
|
||||
|
||||
--*/
|
||||
NESTED_ENTRY MlasConvSymKernelNeonDot
|
||||
|
||||
PROLOG_SAVE_REG_PAIR d8,d9,#-64!
|
||||
PROLOG_NOP ldr x8,[sp,#ConvSymFrame_PostProcessParams]
|
||||
PROLOG_NOP ldr w10,[sp,#ConvSymFrame_KernelFlags]
|
||||
PROLOG_SAVE_REG_PAIR d10,d11,#16
|
||||
PROLOG_SAVE_REG_PAIR d12,d13,#32
|
||||
PROLOG_SAVE_REG_PAIR x19,x20,#48
|
||||
|
||||
// compute C pointers: x2, x16, x17, x5
|
||||
cmp x7,2 // OutputCount < 2 ?
|
||||
add x16,x2,x5 // x16 -> C1
|
||||
lsl x3,x3,#3 // KernelSize * sizeof(int8_t*)
|
||||
csel x16,x2,x16,lo // if OutputCount < 2 x16/C1 -> C0
|
||||
mov x20,x4
|
||||
add x4,x4,3 // InputChannels align to 4
|
||||
add x17,x16,x5 // x17 -> C2
|
||||
ldr x11,[x8,#ConvSymPostProcessParams_Bias]
|
||||
csel x17,x16,x17,ls // if OutputCount <= 2 x17/C2 -> C1
|
||||
bic x4,x4,3
|
||||
cmp x7,4 // OutputCount < 4 ?
|
||||
add x5,x17,x5 // x5 -> C3
|
||||
ldr x19,[x8,#ConvSymPostProcessParams_Scale]
|
||||
csel x5,x17,x5,lo // if OutputCount < 4 x5/C3 -> C2
|
||||
movi v12.16b,128 // for top bit flipping
|
||||
|
||||
OutputChannelLoop
|
||||
ldp q16,q20,[x11],32 // Init accumulators with biases
|
||||
mov v17.16b,v16.16b
|
||||
mov v18.16b,v16.16b
|
||||
ldp q24,q28,[x11],32
|
||||
mov v19.16b,v16.16b
|
||||
mov v21.16b,v20.16b
|
||||
mov v22.16b,v20.16b
|
||||
mov v23.16b,v20.16b
|
||||
mov v25.16b,v24.16b
|
||||
mov v26.16b,v24.16b
|
||||
mov v27.16b,v24.16b
|
||||
mov v29.16b,v28.16b
|
||||
mov v30.16b,v28.16b
|
||||
mov v31.16b,v28.16b
|
||||
mov x9,x3 // restore KernelSize * sizeof(int8_t*)
|
||||
|
||||
KernelSizeLoop
|
||||
tst w10,#MLAS_CONV_SYM_FLAG_INPUT_DIRECT
|
||||
beq InputIndirection
|
||||
|
||||
InputDirect
|
||||
cmp x16,x2
|
||||
mov x12,x0 // x12 -> A0
|
||||
add x13,x0,x20 // x13 -> A1 = A0 + input channels
|
||||
csel x13,x0,x13,eq
|
||||
cmp x17,x16
|
||||
add x14,x0,x20,lsl#1 // x14 -> A2
|
||||
csel x14,x13,x14,eq
|
||||
cmp x5,x17
|
||||
add x15,x13,x20,lsl#1 // x15 -> A3
|
||||
csel x15,x14,x15,eq
|
||||
b FinishLoadAPtr
|
||||
|
||||
InputIndirection
|
||||
ldr x12,[x0] // x12 -> A0
|
||||
cmp x16,x2
|
||||
b.eq SkipLoadA1 // C1==C0 -> A0=A1=A2=A3
|
||||
cmp x17,x16
|
||||
lsl x14,x3,#1
|
||||
ldr x13,[x0,x3] // x13 -> A1
|
||||
b.eq SkipLoadA2 // C2==C1 -> A1=A2=A3
|
||||
cmp x5,x17
|
||||
add x15,x3,x3,lsl#1
|
||||
ldr x14,[x0,x14] // x14 -> A2
|
||||
b.eq SkipLoadA3 // C3==C2 -> A2=A3
|
||||
ldr x15,[x0,x15] // x15 -> A3
|
||||
b FinishLoadAPtr
|
||||
SkipLoadA1
|
||||
mov x13,x12
|
||||
SkipLoadA2
|
||||
mov x14,x13
|
||||
SkipLoadA3
|
||||
mov x15,x14
|
||||
|
||||
// Register Usage
|
||||
// B (x1) -> 4x16
|
||||
// ----------------------------------------------------------------------------
|
||||
// |v4.b[0]..v4.b[12] v5.b[0]..v5.b[12] v6.b[0]..v6.b[12] v7.b[0]..v7.b[12]|
|
||||
// | ... ... ... ... ... ... ... ... |
|
||||
// |v4.b[3]..v4.b[15] v5.b[3]..v5.b[15] v6.b[3]..v6.b[15] v7.b[3]..v7.b[15]|
|
||||
// A 4x4 ----------------------------------------------------------------------------
|
||||
// ------------------ ----------------------------------------------------------------------------
|
||||
// x12 |v0.b[0]..v0.b[3]| |v16.s[0]_v16.s[3] v20.s[0]_v20.s[3] v24.s[0]_v24.s[3] v28.s[0]_v28.s[3]| x2
|
||||
// x13 |v1.b[0]..v1.b[3]| |v17.s[0]_v17.s[3] v21.s[0]_v21.s[3] v25.s[0]_v25.s[3] v29.s[0]_v29.s[3]| x16
|
||||
// x14 |v2.b[0]..v2.b[3]| |v18.s[0]_v18.s[3] v22.s[0]_v23.s[3] v26.s[0]_v26.s[3] v30.s[0]_v31.s[3]| x17
|
||||
// x15 |v3.b[0]..v3.b[3]| |v19.s[0]_v19.s[3] v23.s[0]_v23.s[3] v27.s[0]_v27.s[3] v31.s[0]_v31.s[3]| x5
|
||||
// ------------------ ----------------------------------------------------------------------------
|
||||
|
||||
FinishLoadAPtr
|
||||
subs x7,x4,16 // Need 16 input channels for loop
|
||||
add x0,x0,8 // indirect A advance to next pointer, prepare for kernel size loop
|
||||
b.lo InChannels8
|
||||
|
||||
ldr d0,[x12],8
|
||||
ldr q4,[x1],16
|
||||
ldr d1,[x13],8
|
||||
subs x7,x7,16
|
||||
ldr d2,[x14],8
|
||||
ldr d3,[x15],8
|
||||
ldr q5,[x1],16
|
||||
ldr q6,[x1],16
|
||||
ldr q7,[x1],16
|
||||
b.lo InChLoopEpilogue // Need 32 input channels for main loop
|
||||
|
||||
InputChannelLoop
|
||||
eor v0.8b,v0.8b,v12.8b
|
||||
eor v1.8b,v1.8b,v12.8b
|
||||
sdot v16.4s,v4.16b,v0.4b[0]
|
||||
eor v2.8b,v2.8b,v12.8b
|
||||
sdot v17.4s,v4.16b,v1.4b[0]
|
||||
eor v3.8b,v3.8b,v12.8b
|
||||
ldr d8,[x12],8
|
||||
sdot v18.4s,v4.16b,v2.4b[0]
|
||||
sdot v19.4s,v4.16b,v3.4b[0]
|
||||
ldr q4,[x1],16
|
||||
sdot v20.4s,v5.16b,v0.4b[0]
|
||||
sdot v21.4s,v5.16b,v1.4b[0]
|
||||
ldr d9,[x13],8
|
||||
sdot v22.4s,v5.16b,v2.4b[0]
|
||||
sdot v23.4s,v5.16b,v3.4b[0]
|
||||
ldr q5,[x1],16
|
||||
sdot v24.4s,v6.16b,v0.4b[0]
|
||||
sdot v25.4s,v6.16b,v1.4b[0]
|
||||
ldr d10,[x14],8
|
||||
sdot v26.4s,v6.16b,v2.4b[0]
|
||||
sdot v27.4s,v6.16b,v3.4b[0]
|
||||
ldr q6,[x1],16
|
||||
sdot v28.4s,v7.16b,v0.4b[0]
|
||||
sdot v29.4s,v7.16b,v1.4b[0]
|
||||
ldr d11,[x15],8
|
||||
sdot v30.4s,v7.16b,v2.4b[0]
|
||||
sdot v31.4s,v7.16b,v3.4b[0]
|
||||
ldr q7,[x1],16
|
||||
sdot v16.4s,v4.16b,v0.4b[1]
|
||||
sdot v17.4s,v4.16b,v1.4b[1]
|
||||
sdot v18.4s,v4.16b,v2.4b[1]
|
||||
sdot v19.4s,v4.16b,v3.4b[1]
|
||||
ldr q4,[x1],16
|
||||
sdot v20.4s,v5.16b,v0.4b[1]
|
||||
sdot v21.4s,v5.16b,v1.4b[1]
|
||||
sdot v22.4s,v5.16b,v2.4b[1]
|
||||
sdot v23.4s,v5.16b,v3.4b[1]
|
||||
ldr q5,[x1],16
|
||||
sdot v24.4s,v6.16b,v0.4b[1]
|
||||
sdot v25.4s,v6.16b,v1.4b[1]
|
||||
sdot v26.4s,v6.16b,v2.4b[1]
|
||||
sdot v27.4s,v6.16b,v3.4b[1]
|
||||
ldr q6,[x1],16
|
||||
sdot v28.4s,v7.16b,v0.4b[1]
|
||||
sdot v29.4s,v7.16b,v1.4b[1]
|
||||
sdot v30.4s,v7.16b,v2.4b[1]
|
||||
sdot v31.4s,v7.16b,v3.4b[1]
|
||||
eor v8.8b,v8.8b,v12.8b
|
||||
ldr q7,[x1],16
|
||||
eor v9.8b,v9.8b,v12.8b
|
||||
sdot v16.4s,v4.16b,v8.4b[0]
|
||||
eor v10.8b,v10.8b,v12.8b
|
||||
sdot v17.4s,v4.16b,v9.4b[0]
|
||||
ldr d0,[x12],8
|
||||
eor v11.8b,v11.8b,v12.8b
|
||||
sdot v18.4s,v4.16b,v10.4b[0]
|
||||
sdot v19.4s,v4.16b,v11.4b[0]
|
||||
ldr q4,[x1],16
|
||||
sdot v20.4s,v5.16b,v8.4b[0]
|
||||
sdot v21.4s,v5.16b,v9.4b[0]
|
||||
ldr d1,[x13],8
|
||||
sdot v22.4s,v5.16b,v10.4b[0]
|
||||
sdot v23.4s,v5.16b,v11.4b[0]
|
||||
ldr q5,[x1],16
|
||||
sdot v24.4s,v6.16b,v8.4b[0]
|
||||
sdot v25.4s,v6.16b,v9.4b[0]
|
||||
ldr d2,[x14],8
|
||||
sdot v26.4s,v6.16b,v10.4b[0]
|
||||
sdot v27.4s,v6.16b,v11.4b[0]
|
||||
ldr q6,[x1],16
|
||||
sdot v28.4s,v7.16b,v8.4b[0]
|
||||
sdot v29.4s,v7.16b,v9.4b[0]
|
||||
ldr d3,[x15],8
|
||||
sdot v30.4s,v7.16b,v10.4b[0]
|
||||
sdot v31.4s,v7.16b,v11.4b[0]
|
||||
ldr q7,[x1],16
|
||||
sdot v16.4s,v4.16b,v8.4b[1]
|
||||
sdot v17.4s,v4.16b,v9.4b[1]
|
||||
sdot v18.4s,v4.16b,v10.4b[1]
|
||||
sdot v19.4s,v4.16b,v11.4b[1]
|
||||
ldr q4,[x1],16
|
||||
sdot v20.4s,v5.16b,v8.4b[1]
|
||||
sdot v21.4s,v5.16b,v9.4b[1]
|
||||
sdot v22.4s,v5.16b,v10.4b[1]
|
||||
sdot v23.4s,v5.16b,v11.4b[1]
|
||||
ldr q5,[x1],16
|
||||
sdot v24.4s,v6.16b,v8.4b[1]
|
||||
sdot v25.4s,v6.16b,v9.4b[1]
|
||||
sdot v26.4s,v6.16b,v10.4b[1]
|
||||
sdot v27.4s,v6.16b,v11.4b[1]
|
||||
ldr q6,[x1],16
|
||||
sdot v28.4s,v7.16b,v8.4b[1]
|
||||
sdot v29.4s,v7.16b,v9.4b[1]
|
||||
subs x7,x7,16 // InputChannels -= 16
|
||||
sdot v30.4s,v7.16b,v10.4b[1]
|
||||
sdot v31.4s,v7.16b,v11.4b[1]
|
||||
ldr q7,[x1],16
|
||||
b.hs InputChannelLoop
|
||||
|
||||
InChLoopEpilogue
|
||||
eor v0.8b,v0.8b,v12.8b
|
||||
eor v1.8b,v1.8b,v12.8b
|
||||
sdot v16.4s,v4.16b,v0.4b[0]
|
||||
eor v2.8b,v2.8b,v12.8b
|
||||
sdot v17.4s,v4.16b,v1.4b[0]
|
||||
eor v3.8b,v3.8b,v12.8b
|
||||
ldr d8,[x12],8
|
||||
sdot v18.4s,v4.16b,v2.4b[0]
|
||||
sdot v19.4s,v4.16b,v3.4b[0]
|
||||
ldr q4,[x1],16
|
||||
sdot v20.4s,v5.16b,v0.4b[0]
|
||||
sdot v21.4s,v5.16b,v1.4b[0]
|
||||
ldr d9,[x13],8
|
||||
sdot v22.4s,v5.16b,v2.4b[0]
|
||||
sdot v23.4s,v5.16b,v3.4b[0]
|
||||
ldr q5,[x1],16
|
||||
sdot v24.4s,v6.16b,v0.4b[0]
|
||||
sdot v25.4s,v6.16b,v1.4b[0]
|
||||
ldr d10,[x14],8
|
||||
sdot v26.4s,v6.16b,v2.4b[0]
|
||||
sdot v27.4s,v6.16b,v3.4b[0]
|
||||
ldr q6,[x1],16
|
||||
sdot v28.4s,v7.16b,v0.4b[0]
|
||||
sdot v29.4s,v7.16b,v1.4b[0]
|
||||
ldr d11,[x15],8
|
||||
sdot v30.4s,v7.16b,v2.4b[0]
|
||||
sdot v31.4s,v7.16b,v3.4b[0]
|
||||
ldr q7,[x1],16
|
||||
sdot v16.4s,v4.16b,v0.4b[1]
|
||||
sdot v17.4s,v4.16b,v1.4b[1]
|
||||
sdot v18.4s,v4.16b,v2.4b[1]
|
||||
sdot v19.4s,v4.16b,v3.4b[1]
|
||||
ldr q4,[x1],16
|
||||
sdot v20.4s,v5.16b,v0.4b[1]
|
||||
sdot v21.4s,v5.16b,v1.4b[1]
|
||||
sdot v22.4s,v5.16b,v2.4b[1]
|
||||
sdot v23.4s,v5.16b,v3.4b[1]
|
||||
ldr q5,[x1],16
|
||||
sdot v24.4s,v6.16b,v0.4b[1]
|
||||
sdot v25.4s,v6.16b,v1.4b[1]
|
||||
sdot v26.4s,v6.16b,v2.4b[1]
|
||||
sdot v27.4s,v6.16b,v3.4b[1]
|
||||
ldr q6,[x1],16
|
||||
sdot v28.4s,v7.16b,v0.4b[1]
|
||||
sdot v29.4s,v7.16b,v1.4b[1]
|
||||
sdot v30.4s,v7.16b,v2.4b[1]
|
||||
sdot v31.4s,v7.16b,v3.4b[1]
|
||||
eor v8.8b,v8.8b,v12.8b
|
||||
ldr q7,[x1],16
|
||||
eor v9.8b,v9.8b,v12.8b
|
||||
sdot v16.4s,v4.16b,v8.4b[0]
|
||||
eor v10.8b,v10.8b,v12.8b
|
||||
sdot v17.4s,v4.16b,v9.4b[0]
|
||||
eor v11.8b,v11.8b,v12.8b
|
||||
sdot v18.4s,v4.16b,v10.4b[0]
|
||||
sdot v19.4s,v4.16b,v11.4b[0]
|
||||
ldr q4,[x1],16
|
||||
sdot v20.4s,v5.16b,v8.4b[0]
|
||||
sdot v21.4s,v5.16b,v9.4b[0]
|
||||
sdot v22.4s,v5.16b,v10.4b[0]
|
||||
sdot v23.4s,v5.16b,v11.4b[0]
|
||||
ldr q5,[x1],16
|
||||
sdot v24.4s,v6.16b,v8.4b[0]
|
||||
sdot v25.4s,v6.16b,v9.4b[0]
|
||||
sdot v26.4s,v6.16b,v10.4b[0]
|
||||
sdot v27.4s,v6.16b,v11.4b[0]
|
||||
ldr q6,[x1],16
|
||||
sdot v28.4s,v7.16b,v8.4b[0]
|
||||
sdot v29.4s,v7.16b,v9.4b[0]
|
||||
sdot v30.4s,v7.16b,v10.4b[0]
|
||||
sdot v31.4s,v7.16b,v11.4b[0]
|
||||
ldr q7,[x1],16
|
||||
sdot v16.4s,v4.16b,v8.4b[1]
|
||||
sdot v17.4s,v4.16b,v9.4b[1]
|
||||
sdot v18.4s,v4.16b,v10.4b[1]
|
||||
sdot v19.4s,v4.16b,v11.4b[1]
|
||||
sdot v20.4s,v5.16b,v8.4b[1]
|
||||
sdot v21.4s,v5.16b,v9.4b[1]
|
||||
sdot v22.4s,v5.16b,v10.4b[1]
|
||||
sdot v23.4s,v5.16b,v11.4b[1]
|
||||
sdot v24.4s,v6.16b,v8.4b[1]
|
||||
sdot v25.4s,v6.16b,v9.4b[1]
|
||||
sdot v26.4s,v6.16b,v10.4b[1]
|
||||
sdot v27.4s,v6.16b,v11.4b[1]
|
||||
sdot v28.4s,v7.16b,v8.4b[1]
|
||||
sdot v29.4s,v7.16b,v9.4b[1]
|
||||
sdot v30.4s,v7.16b,v10.4b[1]
|
||||
sdot v31.4s,v7.16b,v11.4b[1]
|
||||
|
||||
TST x7,15
|
||||
B.NE InChannels8 // 4 ~ 12 InputChannels
|
||||
|
||||
subs x9,x9,8 // KernelSize-=1
|
||||
b.hi KernelSizeLoop
|
||||
|
||||
Requantize
|
||||
tst w10,#MLAS_CONV_SYM_FLAG_PER_CHANNEL_SCALE
|
||||
ldr w13,[x8,#ConvSymPostProcessParams_ZeroPoint]
|
||||
beq BroadcastScaleValue
|
||||
ldp q0,q1,[x19],32 // load scale vector
|
||||
ldp q2,q3,[x19],32
|
||||
b AccumulatorsToFloat
|
||||
|
||||
BroadcastScaleValue
|
||||
ld1r {v0.4s},[x19] // load scale Value
|
||||
mov v1.16b, v0.16b
|
||||
mov v2.16b, v0.16b
|
||||
mov v3.16b, v0.16b
|
||||
|
||||
AccumulatorsToFloat
|
||||
scvtf v16.4s,v16.4s // convert to float
|
||||
scvtf v17.4s,v17.4s
|
||||
scvtf v18.4s,v18.4s
|
||||
scvtf v19.4s,v19.4s
|
||||
scvtf v20.4s,v20.4s
|
||||
scvtf v21.4s,v21.4s
|
||||
scvtf v22.4s,v22.4s
|
||||
scvtf v23.4s,v23.4s
|
||||
scvtf v24.4s,v24.4s
|
||||
scvtf v25.4s,v25.4s
|
||||
scvtf v26.4s,v26.4s
|
||||
scvtf v27.4s,v27.4s
|
||||
scvtf v28.4s,v28.4s
|
||||
scvtf v29.4s,v29.4s
|
||||
scvtf v30.4s,v30.4s
|
||||
scvtf v31.4s,v31.4s
|
||||
fmul v16.4s,v16.4s,v0.4s // multiply by scale
|
||||
fmul v17.4s,v17.4s,v0.4s
|
||||
fmul v18.4s,v18.4s,v0.4s
|
||||
fmul v19.4s,v19.4s,v0.4s
|
||||
fmul v20.4s,v20.4s,v1.4s
|
||||
fmul v21.4s,v21.4s,v1.4s
|
||||
fmul v22.4s,v22.4s,v1.4s
|
||||
fmul v23.4s,v23.4s,v1.4s
|
||||
fmul v24.4s,v24.4s,v2.4s
|
||||
fmul v25.4s,v25.4s,v2.4s
|
||||
fmul v26.4s,v26.4s,v2.4s
|
||||
fmul v27.4s,v27.4s,v2.4s
|
||||
fmul v28.4s,v28.4s,v3.4s
|
||||
fmul v29.4s,v29.4s,v3.4s
|
||||
fmul v30.4s,v30.4s,v3.4s
|
||||
fmul v31.4s,v31.4s,v3.4s
|
||||
fcvtns v16.4s,v16.4s // convert to int
|
||||
fcvtns v17.4s,v17.4s
|
||||
fcvtns v18.4s,v18.4s
|
||||
fcvtns v19.4s,v19.4s
|
||||
fcvtns v20.4s,v20.4s
|
||||
fcvtns v21.4s,v21.4s
|
||||
fcvtns v22.4s,v22.4s
|
||||
fcvtns v23.4s,v23.4s
|
||||
fcvtns v24.4s,v24.4s
|
||||
fcvtns v25.4s,v25.4s
|
||||
fcvtns v26.4s,v26.4s
|
||||
fcvtns v27.4s,v27.4s
|
||||
fcvtns v28.4s,v28.4s
|
||||
fcvtns v29.4s,v29.4s
|
||||
fcvtns v30.4s,v30.4s
|
||||
fcvtns v31.4s,v31.4s
|
||||
|
||||
sqxtn v16.4h,v16.4s
|
||||
sqxtn v17.4h,v17.4s
|
||||
sqxtn v18.4h,v18.4s
|
||||
sqxtn v19.4h,v19.4s
|
||||
sqxtn v24.4h,v24.4s
|
||||
sqxtn v25.4h,v25.4s
|
||||
sqxtn v26.4h,v26.4s
|
||||
sqxtn v27.4h,v27.4s
|
||||
dup v4.8h,w13 // zero point
|
||||
sqxtn2 v16.8h,v20.4s
|
||||
sqxtn2 v17.8h,v21.4s
|
||||
sqxtn2 v18.8h,v22.4s
|
||||
sqxtn2 v19.8h,v23.4s
|
||||
sqxtn2 v24.8h,v28.4s
|
||||
sqxtn2 v25.8h,v29.4s
|
||||
sqxtn2 v26.8h,v30.4s
|
||||
sqxtn2 v27.8h,v31.4s
|
||||
sqadd v16.8h,v16.8h,v4.8h
|
||||
sqadd v17.8h,v17.8h,v4.8h
|
||||
sqadd v18.8h,v18.8h,v4.8h
|
||||
sqadd v19.8h,v19.8h,v4.8h
|
||||
sqadd v24.8h,v24.8h,v4.8h
|
||||
sqadd v25.8h,v25.8h,v4.8h
|
||||
sqadd v26.8h,v26.8h,v4.8h
|
||||
sqadd v27.8h,v27.8h,v4.8h
|
||||
sqxtun v0.8b,v16.8h
|
||||
sqxtun v1.8b,v17.8h
|
||||
sqxtun v2.8b,v18.8h
|
||||
sqxtun v3.8b,v19.8h
|
||||
sqxtun2 v0.16b,v24.8h
|
||||
sqxtun2 v1.16b,v25.8h
|
||||
subs x6,x6,16 // processed 16 output channels
|
||||
sqxtun2 v2.16b,v26.8h
|
||||
sqxtun2 v3.16b,v27.8h
|
||||
b.lo PartialStore
|
||||
|
||||
st1 {v3.16b},[x5],16 // Store full 4 x 16
|
||||
st1 {v2.16b},[x17],16
|
||||
sub x0,x0,x3 // Restore pointer to A: a -= ks
|
||||
st1 {v1.16b},[x16],16
|
||||
st1 {v0.16b},[x2],16
|
||||
b.hi OutputChannelLoop
|
||||
|
||||
ExitKernel
|
||||
EPILOG_RESTORE_REG_PAIR x19,x20,#48
|
||||
EPILOG_RESTORE_REG_PAIR d12,d13,#32
|
||||
EPILOG_RESTORE_REG_PAIR d10,d11,#16
|
||||
EPILOG_RESTORE_REG_PAIR d8,d9,#64!
|
||||
EPILOG_RETURN
|
||||
|
||||
InChannels8
|
||||
tbz x7,3,InChannels4
|
||||
ldr d0,[x12],8
|
||||
ldr q4,[x1],16
|
||||
ldr d1,[x13],8
|
||||
ldr d2,[x14],8
|
||||
ldr d3,[x15],8
|
||||
eor v0.8b,v0.8b,v12.8b
|
||||
ldr q5,[x1],16
|
||||
eor v1.8b,v1.8b,v12.8b
|
||||
sdot v16.4s,v4.16b,v0.4b[0]
|
||||
sdot v17.4s,v4.16b,v1.4b[0]
|
||||
eor v2.8b,v2.8b,v12.8b
|
||||
ldp q6,q7,[x1],32
|
||||
eor v3.8b,v3.8b,v12.8b
|
||||
sdot v18.4s,v4.16b,v2.4b[0]
|
||||
sdot v19.4s,v4.16b,v3.4b[0]
|
||||
sdot v20.4s,v5.16b,v0.4b[0]
|
||||
sdot v21.4s,v5.16b,v1.4b[0]
|
||||
sdot v22.4s,v5.16b,v2.4b[0]
|
||||
sdot v23.4s,v5.16b,v3.4b[0]
|
||||
sdot v24.4s,v6.16b,v0.4b[0]
|
||||
sdot v25.4s,v6.16b,v1.4b[0]
|
||||
ldp q4,q5,[x1],32
|
||||
sdot v26.4s,v6.16b,v2.4b[0]
|
||||
sdot v27.4s,v6.16b,v3.4b[0]
|
||||
sdot v28.4s,v7.16b,v0.4b[0]
|
||||
sdot v29.4s,v7.16b,v1.4b[0]
|
||||
sdot v30.4s,v7.16b,v2.4b[0]
|
||||
sdot v31.4s,v7.16b,v3.4b[0]
|
||||
sdot v16.4s,v4.16b,v0.4b[1]
|
||||
sdot v17.4s,v4.16b,v1.4b[1]
|
||||
ldp q6,q7,[x1],32
|
||||
sdot v18.4s,v4.16b,v2.4b[1]
|
||||
sdot v19.4s,v4.16b,v3.4b[1]
|
||||
sdot v20.4s,v5.16b,v0.4b[1]
|
||||
sdot v21.4s,v5.16b,v1.4b[1]
|
||||
sdot v22.4s,v5.16b,v2.4b[1]
|
||||
sdot v23.4s,v5.16b,v3.4b[1]
|
||||
sdot v24.4s,v6.16b,v0.4b[1]
|
||||
sdot v25.4s,v6.16b,v1.4b[1]
|
||||
sdot v26.4s,v6.16b,v2.4b[1]
|
||||
sdot v27.4s,v6.16b,v3.4b[1]
|
||||
sdot v28.4s,v7.16b,v0.4b[1]
|
||||
sdot v29.4s,v7.16b,v1.4b[1]
|
||||
sdot v30.4s,v7.16b,v2.4b[1]
|
||||
sdot v31.4s,v7.16b,v3.4b[1]
|
||||
tbz x7,2,SkipInCh4
|
||||
|
||||
InChannels4
|
||||
ldr s0,[x12],4
|
||||
ldr q4,[x1],16
|
||||
ldr s1,[x13],4
|
||||
ldr s2,[x14],4
|
||||
ldr s3,[x15],4
|
||||
eor v0.8b,v0.8b,v12.8b
|
||||
ldr q5,[x1],16
|
||||
eor v1.8b,v1.8b,v12.8b
|
||||
sdot v16.4s,v4.16b,v0.4b[0]
|
||||
sdot v17.4s,v4.16b,v1.4b[0]
|
||||
eor v2.8b,v2.8b,v12.8b
|
||||
ldp q6,q7,[x1],32
|
||||
eor v3.8b,v3.8b,v12.8b
|
||||
sdot v18.4s,v4.16b,v2.4b[0]
|
||||
sdot v19.4s,v4.16b,v3.4b[0]
|
||||
sdot v20.4s,v5.16b,v0.4b[0]
|
||||
sdot v21.4s,v5.16b,v1.4b[0]
|
||||
sdot v22.4s,v5.16b,v2.4b[0]
|
||||
sdot v23.4s,v5.16b,v3.4b[0]
|
||||
sdot v24.4s,v6.16b,v0.4b[0]
|
||||
sdot v25.4s,v6.16b,v1.4b[0]
|
||||
sdot v26.4s,v6.16b,v2.4b[0]
|
||||
sdot v27.4s,v6.16b,v3.4b[0]
|
||||
sdot v28.4s,v7.16b,v0.4b[0]
|
||||
sdot v29.4s,v7.16b,v1.4b[0]
|
||||
sdot v30.4s,v7.16b,v2.4b[0]
|
||||
sdot v31.4s,v7.16b,v3.4b[0]
|
||||
|
||||
SkipInCh4
|
||||
subs x9,x9,8 // ks -= 1
|
||||
b.hi KernelSizeLoop
|
||||
b Requantize
|
||||
|
||||
PartialStore
|
||||
tbz x6,3,LT8Store
|
||||
str d3,[x5],8 // no less than 8 channels
|
||||
str d2,[x17],8
|
||||
dup d3,v3.d[1]
|
||||
dup d2,v2.d[1]
|
||||
str d1,[x16],8
|
||||
str d0,[x2],8
|
||||
dup d1,v1.d[1]
|
||||
dup d0,v0.d[1]
|
||||
LT8Store
|
||||
tbz x6,2,LT4Store
|
||||
str s3,[x5],4
|
||||
str s2,[x17],4
|
||||
dup s3,v3.s[1]
|
||||
dup s2,v2.s[1]
|
||||
str s1,[x16],4
|
||||
str s0,[x2],4
|
||||
dup s1,v1.s[1]
|
||||
dup s0,v0.s[1]
|
||||
LT4Store
|
||||
tbz x6,1, LT2Store
|
||||
str h3,[x5],2
|
||||
str h2,[x17],2
|
||||
dup h3,v3.h[1]
|
||||
dup h2,v2.h[1]
|
||||
str h1,[x16],2
|
||||
str h0,[x2],2
|
||||
dup h1,v1.h[1]
|
||||
dup h0,v0.h[1]
|
||||
LT2Store
|
||||
tbz x6,0,ExitKernel
|
||||
str b3,[x5]
|
||||
str b2,[x17]
|
||||
str b1,[x16]
|
||||
str b0,[x2]
|
||||
b ExitKernel
|
||||
|
||||
NESTED_END MlasConvSymKernelNeonDot
|
||||
|
||||
END
|
||||
436
onnxruntime/core/mlas/lib/arm64/ConvSymU8KernelNeon.asm
Normal file
436
onnxruntime/core/mlas/lib/arm64/ConvSymU8KernelNeon.asm
Normal file
|
|
@ -0,0 +1,436 @@
|
|||
/*++
|
||||
|
||||
Copyright (c) Microsoft Corporation. All rights reserved.
|
||||
|
||||
Licensed under the MIT License.
|
||||
|
||||
Module Name:
|
||||
|
||||
ConvSymKernelNeon.s
|
||||
|
||||
Abstract:
|
||||
|
||||
This module implements the kernels for the symmetric quantized integer
|
||||
convolution operation.
|
||||
|
||||
--*/
|
||||
|
||||
#include "kxarm64.h"
|
||||
|
||||
#define MLAS_CONV_SYM_FLAG_INPUT_DIRECT 1
|
||||
#define MLAS_CONV_SYM_FLAG_PER_CHANNEL_SCALE 2
|
||||
|
||||
//
|
||||
// Stack frame layout for the symmetric convolution kernel.
|
||||
// d8-d15, x19-x30 need to be preserved if used
|
||||
//
|
||||
#define ConvSymFrame_SavedNeonRegisters (8 * 8)
|
||||
#define ConvSymFrame_SavedRegisters ConvSymFrame_SavedNeonRegisters
|
||||
#define ConvSymFrame_PostProcessParams 0 + ConvSymFrame_SavedRegisters
|
||||
#define ConvSymFrame_KernelFlags 8 + ConvSymFrame_SavedRegisters
|
||||
|
||||
#define ConvSymPostProcessParams_Bias 0
|
||||
#define ConvSymPostProcessParams_Scale 8
|
||||
#define ConvSymPostProcessParams_Min 16
|
||||
#define ConvSymPostProcessParams_Max 20
|
||||
#define ConvSymPostProcessParams_ZeroPoint 24
|
||||
|
||||
TEXTAREA
|
||||
|
||||
/*++
|
||||
|
||||
Routine Description:
|
||||
|
||||
This routine is the inner kernel to compute a convolution for the elements
|
||||
of an output row for a set of filter rows.
|
||||
|
||||
Arguments:
|
||||
|
||||
Input (x0) - Supplies the address of the input buffer.
|
||||
|
||||
If MLAS_CONV_SYM_FLAG_INPUT_DIRECT is set, then the input buffer points
|
||||
directly at the input tensor.
|
||||
|
||||
If MLAS_CONV_SYM_FLAG_INPUT_DIRECT is clear, then the input buffer is an
|
||||
indirection buffer. Every pointer in the indirection buffer points at a
|
||||
InputChannels length vector (either from the input tensor or a vector of
|
||||
padding values). These are grouped in batches of length KernelSize.
|
||||
These batches are then repeated OutputCount times.
|
||||
|
||||
Filter (x1) - Supplies the address of the filter buffer.
|
||||
|
||||
Output (x2) - Supplies the address of the output buffer.
|
||||
|
||||
KernelSize (x3) - Supplies the size of the kernel.
|
||||
|
||||
If MLAS_CONV_SYM_FLAG_INPUT_DIRECT is set, then kernel size should be 1.
|
||||
|
||||
InputChannels (x4) - Supplies the number of input channels.
|
||||
|
||||
This implementation requires the count to be a multiple of 8.
|
||||
|
||||
OutputChannels (x5) - Supplies the number of output channels.
|
||||
|
||||
ChannelCount (x6) - Supplies the number of channels this iteration produces.
|
||||
|
||||
This implementation requires the count to be 8.
|
||||
|
||||
OutputCount (x7) - Supplies the number of output elements this iteration produces.
|
||||
|
||||
This implementation requires the count to be 1 or 2.
|
||||
|
||||
PostProcessParams - Supplies the address of the post process parameter block.
|
||||
|
||||
KernelFlags - Supplies additional flags controlling the operation.
|
||||
|
||||
Return Value:
|
||||
|
||||
None.
|
||||
|
||||
--*/
|
||||
NESTED_ENTRY MlasConvSymKernelNeon
|
||||
|
||||
PROLOG_SAVE_REG_PAIR d8,d9,#-64!
|
||||
PROLOG_NOP ldr x8,[sp,#ConvSymFrame_PostProcessParams]
|
||||
PROLOG_NOP ldrb w10,[sp,#ConvSymFrame_KernelFlags]
|
||||
PROLOG_SAVE_REG_PAIR d10,d11,#16
|
||||
PROLOG_SAVE_REG_PAIR d12,d13,#32
|
||||
PROLOG_SAVE_REG_PAIR d14,d15,#48
|
||||
mov x9,x3 // save kernel size
|
||||
ldr x11,[x8,#ConvSymPostProcessParams_Bias]
|
||||
mov x16,x4 // save input channels
|
||||
ldr x12,[x8,#ConvSymPostProcessParams_Scale]
|
||||
cmp x7,2 // if OutputCount < 2
|
||||
add x5,x2,x5 // c1 = c0 + ldc
|
||||
add x4,x4,7 // kc = (kc + 7) & ~7
|
||||
csel x5,x2,x5,lo // if OutputCount < 2 c1 = c0
|
||||
bic x4,x4,7
|
||||
ldp s16,s18,[x11],8 // init accumulators with bias
|
||||
ldp s20,s22,[x11],8
|
||||
ldp s24,s26,[x11],8
|
||||
ldp s28,s30,[x11],8
|
||||
mov v17.16b,v16.16b
|
||||
mov v19.16b,v18.16b
|
||||
mov v21.16b,v20.16b
|
||||
mov v23.16b,v22.16b
|
||||
mov v25.16b,v24.16b
|
||||
mov v27.16b,v26.16b
|
||||
mov v29.16b,v28.16b
|
||||
mov v31.16b,v30.16b
|
||||
|
||||
// Nested loops, inner loop: input channel; outter loop: kernel size
|
||||
// Each inner iteration processes 8 input channels, 2 output pixels, 8 output channels.
|
||||
//
|
||||
// B 8x8
|
||||
// ------------------------------------------------------------------
|
||||
// |v4.b[0] v5.b[0] v4.b[0] v5.b[0] v4.b[0] v5.b[0] v4.b[0] v5.b[0] |
|
||||
// | ... ... ... ... ... ... ... ... |
|
||||
// |v4.b[7] v5.b[7] v4.b[7] v5.b[7] v4.b[7] v5.b[7] v4.b[7] v5.b[7] |
|
||||
// A 2x8 ------------------------------------------------------------------
|
||||
// ------------------ ------------------------------------------------------------------
|
||||
// x13-> |v0.b[0]..v0.b[7]| |v16.4s v18.4s v20.4s v22.4s v24.4s v26.4s v28.4s v30.4s |
|
||||
// x15-> |v1.b[0]..v1.b[7]| |v17.4s v19.4s v21.4s v23.4s v25.4s v27.4s v29.4s v31.4s |
|
||||
// ------------------ ------------------------------------------------------------------
|
||||
// When Input Channels greater than 16, unroll:
|
||||
// A registers v6 v7,
|
||||
// B registers v8 v9
|
||||
//
|
||||
|
||||
KernelSizeLoop
|
||||
|
||||
// Load next 2 A pointers
|
||||
tst w10,#MLAS_CONV_SYM_FLAG_INPUT_DIRECT
|
||||
ldr d4,[x1]
|
||||
ldr d5,[x1,8]
|
||||
beq InputIndirection
|
||||
|
||||
InputDirect
|
||||
mov x13,x0 // x13 -> A0
|
||||
add x15,x0,x16 // x15 -> A1 = A0 + input channels
|
||||
b BlockLoopPrologue
|
||||
|
||||
InputIndirection
|
||||
cmp x7,2 // test if OutputCount < 2
|
||||
ldr x13,[x0] // x13 -> A0
|
||||
blo SkipLoadA1
|
||||
ldr x15,[x0,x3,lsl#3] // x15 -> A1
|
||||
SkipLoadA1
|
||||
|
||||
BlockLoopPrologue
|
||||
cmp x7,2 // test if OutputCount < 2
|
||||
add x0,x0,8 // indirect A advance to next pointer, prepare for kernel size loop
|
||||
csel x15,x13,x15,lo // if OutputCount < 2 x15 -> A0
|
||||
subs x14,x4,16 // input channel - 16
|
||||
movi v12.8b,128
|
||||
blo InputChannel8 // less than 16 deep, no unroll
|
||||
|
||||
ldr d0,[x13],8
|
||||
ldr d1,[x15],8
|
||||
ldr d8,[x1,64]
|
||||
ldr d9,[x1,72]
|
||||
ldr d6,[x13],8
|
||||
subs x14,x14,16 // input channel - 16
|
||||
ldr d7,[x15],8
|
||||
blo BlockLoopEpilogue // need 32 input channel for full unrolled loop
|
||||
|
||||
Blockloop
|
||||
eor v0.8b,v0.8b,v12.8b
|
||||
eor v1.8b,v1.8b,v12.8b
|
||||
smull v2.8h,v4.8b,v0.8b
|
||||
smull v3.8h,v4.8b,v1.8b
|
||||
ldr d4,[x1,16]
|
||||
smull v10.8h,v5.8b,v0.8b
|
||||
smull v11.8h,v5.8b,v1.8b
|
||||
ldr d5,[x1,24]
|
||||
eor v6.8b,v6.8b,v12.8b
|
||||
eor v7.8b,v7.8b,v12.8b
|
||||
smlal v2.8h,v8.8b,v6.8b
|
||||
smlal v3.8h,v8.8b,v7.8b
|
||||
ldr d8,[x1,80]
|
||||
smlal v10.8h,v9.8b,v6.8b
|
||||
smlal v11.8h,v9.8b,v7.8b
|
||||
ldr d9,[x1,88]
|
||||
smull v12.8h,v4.8b,v0.8b
|
||||
sadalp v16.4s,v2.8h
|
||||
smull v13.8h,v4.8b,v1.8b
|
||||
ldr d4,[x1,32]
|
||||
sadalp v17.4s,v3.8h
|
||||
smull v14.8h,v5.8b,v0.8b
|
||||
sadalp v18.4s,v10.8h
|
||||
smull v15.8h,v5.8b,v1.8b
|
||||
ldr d5,[x1,40]
|
||||
sadalp v19.4s,v11.8h
|
||||
smlal v12.8h,v8.8b,v6.8b
|
||||
smlal v13.8h,v8.8b,v7.8b
|
||||
ldr d8,[x1,96]
|
||||
smlal v14.8h,v9.8b,v6.8b
|
||||
smlal v15.8h,v9.8b,v7.8b
|
||||
ldr d9,[x1,104]
|
||||
smull v2.8h,v4.8b,v0.8b
|
||||
sadalp v20.4s,v12.8h
|
||||
smull v3.8h,v4.8b,v1.8b
|
||||
ldr d4,[x1,48]
|
||||
sadalp v21.4s,v13.8h
|
||||
smull v10.8h,v5.8b,v0.8b
|
||||
sadalp v22.4s,v14.8h
|
||||
smull v11.8h,v5.8b,v1.8b
|
||||
ldr d5,[x1,56]
|
||||
sadalp v23.4s, v15.8h
|
||||
smlal v2.8h,v8.8b,v6.8b
|
||||
smlal v3.8h,v8.8b,v7.8b
|
||||
ldr d8,[x1,112]
|
||||
smlal v10.8h,v9.8b,v6.8b
|
||||
smlal v11.8h,v9.8b,v7.8b
|
||||
ldr d9,[x1,120]
|
||||
smull v12.8h,v4.8b,v0.8b
|
||||
add x1,x1,128
|
||||
sadalp v24.4s,v2.8h
|
||||
smull v13.8h,v4.8b,v1.8b
|
||||
ldr d4,[x1] // Read B
|
||||
sadalp v25.4s,v3.8h
|
||||
smull v14.8h,v5.8b,v0.8b
|
||||
ldr d0,[x13],8 // Read A0
|
||||
sadalp v26.4s,v10.8h
|
||||
smull v15.8h,v5.8b,v1.8b
|
||||
ldr d1,[x15],8 // Read A1
|
||||
sadalp v27.4s,v11.8h
|
||||
smlal v12.8h,v8.8b,v6.8b
|
||||
ldr d5,[x1,8] // Read B
|
||||
smlal v13.8h,v8.8b,v7.8b
|
||||
ldr d8,[x1,64] // Read B
|
||||
smlal v14.8h,v9.8b,v6.8b
|
||||
ldr d6,[x13],8 // Read A0
|
||||
smlal v15.8h,v9.8b,v7.8b
|
||||
ldr d7,[x15],8 // Read A1
|
||||
sadalp v28.4s,v12.8h
|
||||
ldr d9,[x1,72] // Read B
|
||||
sadalp v29.4s,v13.8h
|
||||
subs x14,x14,16
|
||||
sadalp v30.4s,v14.8h
|
||||
movi v12.8b,128
|
||||
sadalp v31.4s,v15.8h
|
||||
b.hs Blockloop
|
||||
|
||||
BlockLoopEpilogue // remaining 16 input channels
|
||||
eor v0.8b,v0.8b,v12.8b
|
||||
eor v1.8b,v1.8b,v12.8b
|
||||
smull v2.8h,v4.8b,v0.8b
|
||||
smull v3.8h,v4.8b,v1.8b
|
||||
ldr d4,[x1,16]
|
||||
smull v10.8h,v5.8b,v0.8b
|
||||
smull v11.8h,v5.8b,v1.8b
|
||||
ldr d5,[x1,24]
|
||||
eor v6.8b,v6.8b,v12.8b
|
||||
eor v7.8b,v7.8b,v12.8b
|
||||
smlal v2.8h,v8.8b,v6.8b
|
||||
smlal v3.8h,v8.8b,v7.8b
|
||||
ldr d8,[x1,80]
|
||||
smlal v10.8h,v9.8b,v6.8b
|
||||
smlal v11.8h,v9.8b,v7.8b
|
||||
ldr d9,[x1,88]
|
||||
smull v12.8h,v4.8b,v0.8b
|
||||
sadalp v16.4s,v2.8h
|
||||
smull v13.8h,v4.8b,v1.8b
|
||||
ldr d4,[x1,32]
|
||||
sadalp v17.4s,v3.8h
|
||||
smull v14.8h,v5.8b,v0.8b
|
||||
sadalp v18.4s,v10.8h
|
||||
smull v15.8h,v5.8b,v1.8b
|
||||
sadalp v19.4s,v11.8h
|
||||
ldr d5,[x1,40]
|
||||
smlal v12.8h,v8.8b,v6.8b
|
||||
smlal v13.8h,v8.8b,v7.8b
|
||||
ldr d8,[x1,96]
|
||||
smlal v14.8h,v9.8b,v6.8b
|
||||
smlal v15.8h,v9.8b,v7.8b
|
||||
ldr d9,[x1,104]
|
||||
smull v2.8h,v4.8b,v0.8b
|
||||
sadalp v20.4s,v12.8h
|
||||
smull v3.8h,v4.8b,v1.8b
|
||||
ldr d4,[x1,48]
|
||||
sadalp v21.4s,v13.8h
|
||||
smull v10.8h,v5.8b,v0.8b
|
||||
sadalp v22.4s,v14.8h
|
||||
smull v11.8h,v5.8b,v1.8b
|
||||
sadalp v23.4s,v15.8h
|
||||
ldr d5,[x1,56]
|
||||
smlal v2.8h,v8.8b,v6.8b
|
||||
smlal v3.8h,v8.8b,v7.8b
|
||||
ldr d8,[x1,112]
|
||||
smlal v10.8h,v9.8b,v6.8b
|
||||
smlal v11.8h,v9.8b,v7.8b
|
||||
ldr d9,[x1,120]
|
||||
smull v12.8h,v4.8b,v0.8b
|
||||
sadalp v24.4s,v2.8h
|
||||
smull v13.8h,v4.8b,v1.8b
|
||||
sadalp v25.4s,v3.8h
|
||||
smull v14.8h,v5.8b,v0.8b
|
||||
sadalp v26.4s,v10.8h
|
||||
smull v15.8h,v5.8b,v1.8b
|
||||
sadalp v27.4s,v11.8h
|
||||
smlal v12.8h,v8.8b,v6.8b
|
||||
smlal v13.8h,v8.8b,v7.8b
|
||||
smlal v14.8h,v9.8b,v6.8b
|
||||
smlal v15.8h,v9.8b,v7.8b
|
||||
add x1,x1,128
|
||||
|
||||
sadalp v28.4s,v12.8h
|
||||
sadalp v29.4s,v13.8h
|
||||
sadalp v30.4s,v14.8h
|
||||
sadalp v31.4s,v15.8h
|
||||
movi v12.8b,128
|
||||
tbnz x14,3,InputChannel8
|
||||
|
||||
subs x9,x9,1
|
||||
b.hi KernelSizeLoop
|
||||
|
||||
Requantize
|
||||
ldr w11,[x8,#ConvSymPostProcessParams_ZeroPoint]
|
||||
tst w10,#MLAS_CONV_SYM_FLAG_PER_CHANNEL_SCALE
|
||||
beq BroadcastScaleValue
|
||||
ld1 {v4.4s,v5.4s},[x12] // load scale vector
|
||||
b AccumulatorsToFloat
|
||||
|
||||
BroadcastScaleValue
|
||||
ld1r {v4.4s},[x12] // load scale Value
|
||||
mov v5.16b, v4.16b
|
||||
|
||||
AccumulatorsToFloat
|
||||
addp v16.4s,v16.4s,v18.4s
|
||||
addp v20.4s,v20.4s,v22.4s
|
||||
addp v24.4s,v24.4s,v26.4s
|
||||
addp v28.4s,v28.4s,v30.4s
|
||||
addp v17.4s,v17.4s,v19.4s
|
||||
addp v21.4s,v21.4s,v23.4s
|
||||
addp v25.4s,v25.4s,v27.4s
|
||||
addp v29.4s,v29.4s,v31.4s
|
||||
addp v0.4s,v16.4s,v20.4s
|
||||
addp v1.4s,v24.4s,v28.4s
|
||||
addp v2.4s,v17.4s,v21.4s
|
||||
addp v3.4s,v25.4s,v29.4s
|
||||
scvtf v0.4s,v0.4s // convert to float
|
||||
scvtf v1.4s,v1.4s
|
||||
scvtf v2.4s,v2.4s
|
||||
scvtf v3.4s,v3.4s
|
||||
fmul v0.4s,v0.4s,v4.4s // multiply by scale
|
||||
fmul v1.4s,v1.4s,v5.4s
|
||||
fmul v2.4s,v2.4s,v4.4s
|
||||
fmul v3.4s,v3.4s,v5.4s
|
||||
fcvtns v0.4s,v0.4s // convert to int
|
||||
fcvtns v1.4s,v1.4s
|
||||
dup v9.8h,w11
|
||||
fcvtns v2.4s,v2.4s
|
||||
fcvtns v3.4s,v3.4s
|
||||
sqxtn v0.4h,v0.4s
|
||||
sqxtn2 v0.8h,v1.4s
|
||||
sqxtn v2.4h,v2.4s
|
||||
sqxtn2 v2.8h,v3.4s
|
||||
sqadd v0.8h,v0.8h,v9.8h
|
||||
sqadd v2.8h,v2.8h,v9.8h
|
||||
sqxtun v0.8b,v0.8h // shorten to int8
|
||||
sqxtun2 v0.16b,v2.8h
|
||||
st1 {v0.d}[1],[x5] // full 2x8 store to c
|
||||
st1 {v0.8b},[x2]
|
||||
|
||||
ExitKernel
|
||||
EPILOG_RESTORE_REG_PAIR d14,d15,#48
|
||||
EPILOG_RESTORE_REG_PAIR d12,d13,#32
|
||||
EPILOG_RESTORE_REG_PAIR d10,d11,#16
|
||||
EPILOG_RESTORE_REG_PAIR d8,d9,#64!
|
||||
EPILOG_RETURN
|
||||
|
||||
InputChannel8
|
||||
ldr d0,[x13]
|
||||
ldr d1,[x15]
|
||||
ldr d4,[x1]
|
||||
ldr d5,[x1,8]
|
||||
ldr d6,[x1,16]
|
||||
ldr d7,[x1,24]
|
||||
eor v0.8b,v0.8b,v12.8b
|
||||
eor v1.8b,v1.8b,v12.8b
|
||||
smull v2.8h,v4.8b,v0.8b
|
||||
smull v3.8h,v4.8b,v1.8b
|
||||
ldr d4,[x1,32]
|
||||
smull v10.8h,v5.8b,v0.8b
|
||||
smull v11.8h,v5.8b,v1.8b
|
||||
ldr d5,[x1,40]
|
||||
smull v12.8h,v6.8b,v0.8b
|
||||
sadalp v16.4s,v2.8h
|
||||
smull v13.8h,v6.8b,v1.8b
|
||||
ldr d6,[x1,48]
|
||||
sadalp v17.4s,v3.8h
|
||||
smull v14.8h,v7.8b,v0.8b
|
||||
sadalp v18.4s,v10.8h
|
||||
smull v15.8h,v7.8b,v1.8b
|
||||
ldr d7,[x1,56]
|
||||
sadalp v19.4s,v11.8h
|
||||
smull v2.8h,v4.8b,v0.8b
|
||||
sadalp v20.4s,v12.8h
|
||||
smull v3.8h,v4.8b,v1.8b
|
||||
sadalp v21.4s,v13.8h
|
||||
smull v10.8h,v5.8b,v0.8b
|
||||
sadalp v22.4s,v14.8h
|
||||
smull v11.8h,v5.8b,v1.8b
|
||||
sadalp v23.4s,v15.8h
|
||||
smull v12.8h,v6.8b,v0.8b
|
||||
sadalp v24.4s,v2.8h
|
||||
smull v13.8h,v6.8b,v1.8b
|
||||
sadalp v25.4s,v3.8h
|
||||
smull v14.8h,v7.8b,v0.8b
|
||||
sadalp v26.4s,v10.8h
|
||||
smull v15.8h,v7.8b,v1.8b
|
||||
sadalp v27.4s,v11.8h
|
||||
add x1,x1,64
|
||||
sadalp v28.4s,v12.8h
|
||||
sadalp v29.4s,v13.8h
|
||||
sadalp v30.4s,v14.8h
|
||||
sadalp v31.4s,v15.8h
|
||||
|
||||
// ks loop
|
||||
subs x9,x9,1
|
||||
b.hi KernelSizeLoop
|
||||
b Requantize
|
||||
|
||||
NESTED_END MlasConvSymKernelNeon
|
||||
|
||||
END
|
||||
|
|
@ -64,6 +64,8 @@ extern "C" {
|
|||
MLAS_CONV_SYM_KERNEL MlasConvSymKernelAvx512Vnni;
|
||||
MLAS_CONV_SYM_DEPTHWISE_KERNEL MlasConvSymDepthwiseKernelAvx512Vnni;
|
||||
#elif defined(MLAS_TARGET_ARM64)
|
||||
MLAS_CONV_SYM_KERNEL MlasConvSymKernelNeon;
|
||||
MLAS_CONV_SYM_KERNEL MlasConvSymKernelNeonDot;
|
||||
MLAS_CONV_SYM_DEPTHWISE_KERNEL MlasConvSymDepthwiseKernelNeon;
|
||||
MLAS_CONV_SYM_DEPTHWISE_ROUTINE_KERNELSIZE MlasConvSymDepthwiseKernelSize9Arm64;
|
||||
MLAS_CONV_SYM_DEPTHWISE_ROUTINE_KERNELSIZE MlasConvSymDepthwiseKernelSize25Arm;
|
||||
|
|
@ -153,18 +155,33 @@ const MLAS_CONV_SYM_DISPATCH MlasConvSymDispatchAvx512Vnni = {
|
|||
|
||||
#elif defined(MLAS_TARGET_ARM64)
|
||||
const MLAS_CONV_SYM_DISPATCH MlasConvSymDispatchNeon = {
|
||||
nullptr,
|
||||
MlasConvSymKernelNeon,
|
||||
MlasConvSymDepthwiseKernelNeon,
|
||||
4, // FilterInputChannelPackCount
|
||||
16, // FilterOutputChannelPackCount
|
||||
8, // FilterInputChannelPackCount
|
||||
8, // FilterOutputChannelPackCount
|
||||
8, // KernelChannelCount
|
||||
8, // KernelOutputCount
|
||||
4, // KernelInputChannelAlignment
|
||||
2, // KernelOutputCount
|
||||
8, // KernelInputChannelAlignment
|
||||
8, // KernelOutputChannelAlignment
|
||||
16, // KernelDepthwiseChannelCount
|
||||
4, // KernelDepthwiseOutputCount
|
||||
true
|
||||
};
|
||||
|
||||
const MLAS_CONV_SYM_DISPATCH MlasConvSymDispatchDot = {
|
||||
MlasConvSymKernelNeonDot,
|
||||
MlasConvSymDepthwiseKernelNeon,
|
||||
4, // FilterInputChannelPackCount
|
||||
16, // FilterOutputChannelPackCount
|
||||
0, // KernelChannelCount
|
||||
4, // KernelOutputCount
|
||||
4, // KernelInputChannelAlignment
|
||||
1, // KernelOutputChannelAlignment
|
||||
16, // KernelDepthwiseChannelCount
|
||||
4, // KernelDepthwiseOutputCount
|
||||
true
|
||||
};
|
||||
|
||||
#endif // MLAS_TARGET_AMD64
|
||||
|
||||
MLAS_FORCEINLINE
|
||||
|
|
@ -229,6 +246,16 @@ MlasConvSymPackWSize(
|
|||
|
||||
} else {
|
||||
|
||||
#ifdef MLAS_TARGET_ARM64
|
||||
// TODO!! remove this for functional testing!
|
||||
// TODO!! is there a way to know whether this is called by tests?
|
||||
if (InputChannels < 128) {
|
||||
// Shallow indirect conv runs slower.
|
||||
// TODO!! for DOT arch, threshold should be 32 for better perf
|
||||
return 0;
|
||||
}
|
||||
#endif
|
||||
|
||||
size_t OutputChannelPackCount = ConvSymDispatch->FilterOutputChannelPackCount;
|
||||
|
||||
if (ConvSymDispatch->Kernel == nullptr ||
|
||||
|
|
@ -344,7 +371,9 @@ MlasConvSym(
|
|||
|
||||
MlasConvSymSetOutputZeroPoint(PostProcessParams, Params.OutputZeroPoint, Params.InputIsSigned);
|
||||
|
||||
const size_t KernelChannelCount = ConvSymDispatch->KernelChannelCount;
|
||||
const size_t KernelChannelCount = (ConvSymDispatch->KernelChannelCount == 0)
|
||||
? std::numeric_limits<size_t>::max()
|
||||
: ConvSymDispatch->KernelChannelCount;
|
||||
const size_t KernelOutputCount = ConvSymDispatch->KernelOutputCount;
|
||||
|
||||
const size_t KernelSize = Params.KernelSize;
|
||||
|
|
|
|||
|
|
@ -96,14 +96,46 @@ Abstract:
|
|||
//
|
||||
// Select the threading model.
|
||||
//
|
||||
// N.B. MLAS_NO_ONNXRUNTIME_THREADPOOL is used to build MLAS test code outside
|
||||
// N.B. BUILD_MLAS_NO_ONNXRUNTIME is used to build MLAS test code outside
|
||||
// of the ONNX Runtime source tree. OpenMP may or may not be enabled in this
|
||||
// configuration.
|
||||
//
|
||||
|
||||
#if !defined(MLAS_NO_ONNXRUNTIME_THREADPOOL)
|
||||
#if !defined(BUILD_MLAS_NO_ONNXRUNTIME)
|
||||
#include "core/platform/threadpool.h"
|
||||
#endif
|
||||
|
||||
#if defined(MLAS_TARGET_ARM64) && defined(__linux__)
|
||||
|
||||
#include "core/common/cpuid_info.h"
|
||||
using MLAS_CPUIDINFO = onnxruntime::CPUIDInfo;
|
||||
|
||||
#endif // MLAS_TARGET_ARM64
|
||||
|
||||
#else // BUILD_MLAS_NO_ONNXRUNTIME
|
||||
|
||||
#if defined(MLAS_TARGET_ARM64) && defined(__linux__)
|
||||
class MLASCPUIDInfo
|
||||
{
|
||||
public:
|
||||
static const MLASCPUIDInfo& GetCPUIDInfo()
|
||||
{
|
||||
static MLASCPUIDInfo cpuid_info;
|
||||
return cpuid_info;
|
||||
}
|
||||
|
||||
// ARM
|
||||
bool HasArmNeonDot() const { return has_arm_neon_dot_; }
|
||||
|
||||
private:
|
||||
MLASCPUIDInfo();
|
||||
|
||||
bool has_arm_neon_dot_{false};
|
||||
};
|
||||
using MLAS_CPUIDINFO = MLASCPUIDInfo;
|
||||
|
||||
#endif // MLAS_TARGET_ARM64
|
||||
|
||||
#endif // BUILD_MLAS_NO_ONNXRUNTIME
|
||||
|
||||
#if defined(_OPENMP)
|
||||
#include <omp.h>
|
||||
|
|
@ -680,6 +712,7 @@ extern const MLAS_CONV_SYM_DISPATCH MlasConvSymDispatchAvxVnni;
|
|||
extern const MLAS_CONV_SYM_DISPATCH MlasConvSymDispatchAvx512Core;
|
||||
extern const MLAS_CONV_SYM_DISPATCH MlasConvSymDispatchAvx512Vnni;
|
||||
extern const MLAS_CONV_SYM_DISPATCH MlasConvSymDispatchNeon;
|
||||
extern const MLAS_CONV_SYM_DISPATCH MlasConvSymDispatchDot;
|
||||
|
||||
//
|
||||
// Quantized depthwise convolution kernels.
|
||||
|
|
@ -804,6 +837,8 @@ struct MLAS_PLATFORM {
|
|||
uint32_t NchwcBlockSize;
|
||||
uint32_t PreferredBufferAlignment;
|
||||
int32_t MaximumThreadCount;
|
||||
#elif defined(MLAS_TARGET_ARM64)
|
||||
static constexpr int32_t MaximumThreadCount = MLAS_MAXIMUM_THREAD_COUNT * 4;
|
||||
#else
|
||||
static constexpr int32_t MaximumThreadCount = MLAS_MAXIMUM_THREAD_COUNT;
|
||||
#endif
|
||||
|
|
@ -851,7 +886,7 @@ MlasGetMaximumThreadCount(
|
|||
MLAS_THREADPOOL* ThreadPool
|
||||
)
|
||||
{
|
||||
#if defined(MLAS_NO_ONNXRUNTIME_THREADPOOL)
|
||||
#if defined(BUILD_MLAS_NO_ONNXRUNTIME)
|
||||
MLAS_UNREFERENCED_PARAMETER(ThreadPool);
|
||||
|
||||
#if defined(_OPENMP)
|
||||
|
|
|
|||
|
|
@ -35,6 +35,11 @@ Abstract:
|
|||
#ifndef HWCAP_ASIMDDP
|
||||
#define HWCAP_ASIMDDP (1 << 20)
|
||||
#endif
|
||||
|
||||
#if defined(BUILD_MLAS_NO_ONNXRUNTIME)
|
||||
MLASCPUIDInfo::MLASCPUIDInfo() { has_arm_neon_dot_ = ((getauxval(AT_HWCAP) & HWCAP_ASIMDDP) != 0); }
|
||||
#endif
|
||||
|
||||
#endif
|
||||
#endif // MLAS_TARGET_ARM64
|
||||
|
||||
|
|
@ -364,13 +369,14 @@ Return Value:
|
|||
#if defined(_WIN32)
|
||||
HasDotProductInstructions = (IsProcessorFeaturePresent(PF_ARM_V82_DP_INSTRUCTIONS_AVAILABLE) != 0);
|
||||
#elif defined(__linux__)
|
||||
HasDotProductInstructions = ((getauxval(AT_HWCAP) & HWCAP_ASIMDDP) != 0);
|
||||
HasDotProductInstructions = MLAS_CPUIDINFO::GetCPUIDInfo().HasArmNeonDot();
|
||||
#else
|
||||
HasDotProductInstructions = false;
|
||||
#endif
|
||||
|
||||
if (HasDotProductInstructions) {
|
||||
this->GemmU8X8Dispatch = &MlasGemmU8X8DispatchUdot;
|
||||
this->ConvSymU8S8Dispatch = &MlasConvSymDispatchDot;
|
||||
}
|
||||
|
||||
#endif // MLAS_TARGET_ARM64
|
||||
|
|
|
|||
|
|
@ -1274,7 +1274,7 @@ Return Value:
|
|||
}
|
||||
}
|
||||
|
||||
#ifdef MLAS_NO_ONNXRUNTIME_THREADPOOL
|
||||
#ifdef BUILD_MLAS_NO_ONNXRUNTIME
|
||||
MLAS_UNREFERENCED_PARAMETER(ThreadPool);
|
||||
//
|
||||
// Execute the pooling kernel routine.
|
||||
|
|
|
|||
|
|
@ -33,7 +33,7 @@ MlasExecuteThreaded(
|
|||
return;
|
||||
}
|
||||
|
||||
#if defined(MLAS_NO_ONNXRUNTIME_THREADPOOL)
|
||||
#if defined(BUILD_MLAS_NO_ONNXRUNTIME)
|
||||
MLAS_UNREFERENCED_PARAMETER(ThreadPool);
|
||||
|
||||
//
|
||||
|
|
@ -75,7 +75,7 @@ MlasTrySimpleParallel(
|
|||
return;
|
||||
}
|
||||
|
||||
#if defined(MLAS_NO_ONNXRUNTIME_THREADPOOL)
|
||||
#if defined(BUILD_MLAS_NO_ONNXRUNTIME)
|
||||
MLAS_UNREFERENCED_PARAMETER(ThreadPool);
|
||||
|
||||
//
|
||||
|
|
|
|||
|
|
@ -6,7 +6,7 @@
|
|||
#include <list>
|
||||
#include <algorithm>
|
||||
|
||||
#if !defined(MLAS_NO_ONNXRUNTIME_THREADPOOL)
|
||||
#if !defined(BUILD_MLAS_NO_ONNXRUNTIME)
|
||||
|
||||
MLAS_THREADPOOL* GetMlasThreadPool(void) {
|
||||
static MLAS_THREADPOOL* threadpool = new onnxruntime::concurrency::ThreadPool(
|
||||
|
|
|
|||
|
|
@ -18,7 +18,7 @@
|
|||
#else
|
||||
#include <sys/mman.h>
|
||||
#endif
|
||||
#if !defined(MLAS_NO_ONNXRUNTIME_THREADPOOL)
|
||||
#if !defined(BUILD_MLAS_NO_ONNXRUNTIME)
|
||||
#include "core/platform/threadpool.h"
|
||||
#endif
|
||||
|
||||
|
|
|
|||
|
|
@ -595,16 +595,6 @@ TEST(QLinearConvTest, Conv2D_U8S8_Sym_M64_C64) {
|
|||
test.Run();
|
||||
}
|
||||
|
||||
TEST(QLinearConvTest, Conv2D_U8S8_Sym_M16_C4) {
|
||||
QLinearConvOpTester<uint8_t, int8_t> test;
|
||||
test.GenerateRandomInput({1, 4, 3, 3}, .05f, 4);
|
||||
test.GenerateRandomWeights({16, 4, 3, 3}, .125f, 0);
|
||||
test.GenerateRandomBias();
|
||||
test.SetPads({0, 0, 0, 0});
|
||||
test.SetOutputScaleAndZeroPoint(.55f, 54);
|
||||
test.Run();
|
||||
}
|
||||
|
||||
TEST(QLinearConvTest, Conv2D_U8S8_Sym_M16_C4_Bias) {
|
||||
QLinearConvOpTester<uint8_t, int8_t> test;
|
||||
test.GenerateRandomInput({1, 4, 3, 3}, .05f, 4);
|
||||
|
|
@ -645,6 +635,18 @@ TEST(QLinearConvTest, Conv2D_U8S8_Sym_M32_C32_Bias_Pads) {
|
|||
test.Run();
|
||||
}
|
||||
|
||||
TEST(QLinearConvTest, Conv2D_U8S8_Sym_M8_C8) {
|
||||
// Targeting code processing 8 channels, with odd number
|
||||
// of output pixels
|
||||
QLinearConvOpTester<uint8_t, int8_t> test;
|
||||
test.GenerateRandomInput({1, 8, 3, 5}, .85f, 4);
|
||||
test.GenerateRandomWeights({8, 8, 3, 3}, .125f, 0);
|
||||
test.GenerateRandomBias();
|
||||
test.SetPads({0, 0, 0, 0});
|
||||
test.SetOutputScaleAndZeroPoint(.55f, 54);
|
||||
test.Run();
|
||||
}
|
||||
|
||||
TEST(QLinearConvTest, Conv2D_U8S8) {
|
||||
QLinearConvOpTester<uint8_t, int8_t> test;
|
||||
test.GenerateRandomInput({3, 24, 15, 11}, .05f, 4);
|
||||
|
|
|
|||
Loading…
Reference in a new issue