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:
Chen Fu 2021-12-13 21:14:45 -08:00 committed by GitHub
parent 7e55a942cd
commit cd0af7ad44
No known key found for this signature in database
GPG key ID: 4AEE18F83AFDEB23
16 changed files with 2259 additions and 29 deletions

View file

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

View file

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

View file

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

View file

@ -83,3 +83,4 @@ Arguments:
.inst Instruction
.endm

View 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

View 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

View 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

View 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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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