diff --git a/cmake/onnxruntime_mlas.cmake b/cmake/onnxruntime_mlas.cmake index 12d1e271a9..263f580977 100644 --- a/cmake/onnxruntime_mlas.cmake +++ b/cmake/onnxruntime_mlas.cmake @@ -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 diff --git a/cmake/onnxruntime_webassembly.cmake b/cmake/onnxruntime_webassembly.cmake index 588c1e9da9..99db7cb7c6 100644 --- a/cmake/onnxruntime_webassembly.cmake +++ b/cmake/onnxruntime_webassembly.cmake @@ -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 diff --git a/onnxruntime/core/common/cpuid_info.cc b/onnxruntime/core/common/cpuid_info.cc index 226f1a08c0..fbc8fb03a5 100644 --- a/onnxruntime/core/common/cpuid_info.cc +++ b/onnxruntime/core/common/cpuid_info.cc @@ -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 } diff --git a/onnxruntime/core/mlas/lib/aarch64/AssembleDotProduct.h b/onnxruntime/core/mlas/lib/aarch64/AssembleDotProduct.h index 59ffe0fc52..3af76acddb 100644 --- a/onnxruntime/core/mlas/lib/aarch64/AssembleDotProduct.h +++ b/onnxruntime/core/mlas/lib/aarch64/AssembleDotProduct.h @@ -83,3 +83,4 @@ Arguments: .inst Instruction .endm + diff --git a/onnxruntime/core/mlas/lib/aarch64/ConvSymU8KernelDot.S b/onnxruntime/core/mlas/lib/aarch64/ConvSymU8KernelDot.S new file mode 100644 index 0000000000..a040427c95 --- /dev/null +++ b/onnxruntime/core/mlas/lib/aarch64/ConvSymU8KernelDot.S @@ -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 diff --git a/onnxruntime/core/mlas/lib/aarch64/ConvSymU8KernelNeon.S b/onnxruntime/core/mlas/lib/aarch64/ConvSymU8KernelNeon.S new file mode 100644 index 0000000000..4cb68827f0 --- /dev/null +++ b/onnxruntime/core/mlas/lib/aarch64/ConvSymU8KernelNeon.S @@ -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 diff --git a/onnxruntime/core/mlas/lib/arm64/ConvSymU8KernelDot.asm b/onnxruntime/core/mlas/lib/arm64/ConvSymU8KernelDot.asm new file mode 100644 index 0000000000..ecb6c6578b --- /dev/null +++ b/onnxruntime/core/mlas/lib/arm64/ConvSymU8KernelDot.asm @@ -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 diff --git a/onnxruntime/core/mlas/lib/arm64/ConvSymU8KernelNeon.asm b/onnxruntime/core/mlas/lib/arm64/ConvSymU8KernelNeon.asm new file mode 100644 index 0000000000..5b3c4f5d9e --- /dev/null +++ b/onnxruntime/core/mlas/lib/arm64/ConvSymU8KernelNeon.asm @@ -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 diff --git a/onnxruntime/core/mlas/lib/convsym.cpp b/onnxruntime/core/mlas/lib/convsym.cpp index 1033b1d0c0..4deecc3553 100644 --- a/onnxruntime/core/mlas/lib/convsym.cpp +++ b/onnxruntime/core/mlas/lib/convsym.cpp @@ -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::max() + : ConvSymDispatch->KernelChannelCount; const size_t KernelOutputCount = ConvSymDispatch->KernelOutputCount; const size_t KernelSize = Params.KernelSize; diff --git a/onnxruntime/core/mlas/lib/mlasi.h b/onnxruntime/core/mlas/lib/mlasi.h index 427c0559c1..6e0bd775c9 100644 --- a/onnxruntime/core/mlas/lib/mlasi.h +++ b/onnxruntime/core/mlas/lib/mlasi.h @@ -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 @@ -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) diff --git a/onnxruntime/core/mlas/lib/platform.cpp b/onnxruntime/core/mlas/lib/platform.cpp index f87a539e2e..9e747ce07a 100644 --- a/onnxruntime/core/mlas/lib/platform.cpp +++ b/onnxruntime/core/mlas/lib/platform.cpp @@ -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 diff --git a/onnxruntime/core/mlas/lib/pooling.cpp b/onnxruntime/core/mlas/lib/pooling.cpp index 649e137182..de56a3a42e 100644 --- a/onnxruntime/core/mlas/lib/pooling.cpp +++ b/onnxruntime/core/mlas/lib/pooling.cpp @@ -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. diff --git a/onnxruntime/core/mlas/lib/threading.cpp b/onnxruntime/core/mlas/lib/threading.cpp index 8769abdf08..317101d8c0 100644 --- a/onnxruntime/core/mlas/lib/threading.cpp +++ b/onnxruntime/core/mlas/lib/threading.cpp @@ -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); // diff --git a/onnxruntime/test/mlas/unittest/test_main.cpp b/onnxruntime/test/mlas/unittest/test_main.cpp index 4b7419cc53..66b5a6a15d 100644 --- a/onnxruntime/test/mlas/unittest/test_main.cpp +++ b/onnxruntime/test/mlas/unittest/test_main.cpp @@ -6,7 +6,7 @@ #include #include -#if !defined(MLAS_NO_ONNXRUNTIME_THREADPOOL) +#if !defined(BUILD_MLAS_NO_ONNXRUNTIME) MLAS_THREADPOOL* GetMlasThreadPool(void) { static MLAS_THREADPOOL* threadpool = new onnxruntime::concurrency::ThreadPool( diff --git a/onnxruntime/test/mlas/unittest/test_util.h b/onnxruntime/test/mlas/unittest/test_util.h index c14f1f57ce..c3d97eb3bb 100644 --- a/onnxruntime/test/mlas/unittest/test_util.h +++ b/onnxruntime/test/mlas/unittest/test_util.h @@ -18,7 +18,7 @@ #else #include #endif -#if !defined(MLAS_NO_ONNXRUNTIME_THREADPOOL) +#if !defined(BUILD_MLAS_NO_ONNXRUNTIME) #include "core/platform/threadpool.h" #endif diff --git a/onnxruntime/test/providers/cpu/nn/qlinearconv_op_test.cc b/onnxruntime/test/providers/cpu/nn/qlinearconv_op_test.cc index 178e13aff3..e68168efea 100644 --- a/onnxruntime/test/providers/cpu/nn/qlinearconv_op_test.cc +++ b/onnxruntime/test/providers/cpu/nn/qlinearconv_op_test.cc @@ -595,16 +595,6 @@ TEST(QLinearConvTest, Conv2D_U8S8_Sym_M64_C64) { test.Run(); } -TEST(QLinearConvTest, Conv2D_U8S8_Sym_M16_C4) { - QLinearConvOpTester 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 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 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 test; test.GenerateRandomInput({3, 24, 15, 11}, .05f, 4);