Unroll Arm Neon DotProd variance kernels for w >= 32 Unroll the Arm Neon DotProd variance kernels and split the accumulators for blocks having width >= 32. This change reduces the time required to execute these kernels by up to 40% on a Neoverse V3 when compiled with Clang 22. Change-Id: Ibc8541bde727cbba5df460704877dc5e61455d1f
diff --git a/aom_dsp/arm/variance_neon_dotprod.c b/aom_dsp/arm/variance_neon_dotprod.c index 68a6445..68bb93f 100644 --- a/aom_dsp/arm/variance_neon_dotprod.c +++ b/aom_dsp/arm/variance_neon_dotprod.c
@@ -103,47 +103,110 @@ *sse = horizontal_add_u32x4(sse_u32); } +static inline void variance_32xh_neon_dotprod(const uint8_t *src, + int src_stride, + const uint8_t *ref, + int ref_stride, int h, + uint32_t *sse, int *sum) { + uint32x4_t src_sum[2] = { vdupq_n_u32(0), vdupq_n_u32(0) }; + uint32x4_t ref_sum[2] = { vdupq_n_u32(0), vdupq_n_u32(0) }; + uint32x4_t sse_u32[2] = { vdupq_n_u32(0), vdupq_n_u32(0) }; + + do { + uint8x16_t s0 = vld1q_u8(src); + uint8x16_t r0 = vld1q_u8(ref); + + src_sum[0] = vdotq_u32(src_sum[0], s0, vdupq_n_u8(1)); + ref_sum[0] = vdotq_u32(ref_sum[0], r0, vdupq_n_u8(1)); + + uint8x16_t abs_diff0 = vabdq_u8(s0, r0); + sse_u32[0] = vdotq_u32(sse_u32[0], abs_diff0, abs_diff0); + + uint8x16_t s1 = vld1q_u8(src + 16); + uint8x16_t r1 = vld1q_u8(ref + 16); + + src_sum[1] = vdotq_u32(src_sum[1], s1, vdupq_n_u8(1)); + ref_sum[1] = vdotq_u32(ref_sum[1], r1, vdupq_n_u8(1)); + + uint8x16_t abs_diff1 = vabdq_u8(s1, r1); + sse_u32[1] = vdotq_u32(sse_u32[1], abs_diff1, abs_diff1); + + src += src_stride; + ref += ref_stride; + } while (--h != 0); + + src_sum[0] = vaddq_u32(src_sum[0], src_sum[1]); + ref_sum[0] = vaddq_u32(ref_sum[0], ref_sum[1]); + int32x4_t sum_diff = vsubq_s32(vreinterpretq_s32_u32(src_sum[0]), + vreinterpretq_s32_u32(ref_sum[0])); + *sum = horizontal_add_s32x4(sum_diff); + *sse = horizontal_add_u32x4(vaddq_u32(sse_u32[0], sse_u32[1])); +} + static inline void variance_large_neon_dotprod(const uint8_t *src, int src_stride, const uint8_t *ref, int ref_stride, int w, int h, uint32_t *sse, int *sum) { - uint32x4_t src_sum = vdupq_n_u32(0); - uint32x4_t ref_sum = vdupq_n_u32(0); - uint32x4_t sse_u32 = vdupq_n_u32(0); + uint32x4_t src_sum[4] = { vdupq_n_u32(0), vdupq_n_u32(0), vdupq_n_u32(0), + vdupq_n_u32(0) }; + uint32x4_t ref_sum[4] = { vdupq_n_u32(0), vdupq_n_u32(0), vdupq_n_u32(0), + vdupq_n_u32(0) }; + uint32x4_t sse_u32[4] = { vdupq_n_u32(0), vdupq_n_u32(0), vdupq_n_u32(0), + vdupq_n_u32(0) }; do { int i = 0; do { - uint8x16_t s = vld1q_u8(src + i); - uint8x16_t r = vld1q_u8(ref + i); + uint8x16_t s0 = vld1q_u8(src + i + 0); + uint8x16_t r0 = vld1q_u8(ref + i + 0); - src_sum = vdotq_u32(src_sum, s, vdupq_n_u8(1)); - ref_sum = vdotq_u32(ref_sum, r, vdupq_n_u8(1)); + src_sum[0] = vdotq_u32(src_sum[0], s0, vdupq_n_u8(1)); + ref_sum[0] = vdotq_u32(ref_sum[0], r0, vdupq_n_u8(1)); - uint8x16_t abs_diff = vabdq_u8(s, r); - sse_u32 = vdotq_u32(sse_u32, abs_diff, abs_diff); + uint8x16_t abs_diff0 = vabdq_u8(s0, r0); + sse_u32[0] = vdotq_u32(sse_u32[0], abs_diff0, abs_diff0); - i += 16; + uint8x16_t s1 = vld1q_u8(src + i + 16); + uint8x16_t r1 = vld1q_u8(ref + i + 16); + + src_sum[1] = vdotq_u32(src_sum[1], s1, vdupq_n_u8(1)); + ref_sum[1] = vdotq_u32(ref_sum[1], r1, vdupq_n_u8(1)); + + uint8x16_t abs_diff1 = vabdq_u8(s1, r1); + sse_u32[1] = vdotq_u32(sse_u32[1], abs_diff1, abs_diff1); + + uint8x16_t s2 = vld1q_u8(src + i + 32); + uint8x16_t r2 = vld1q_u8(ref + i + 32); + + src_sum[2] = vdotq_u32(src_sum[2], s2, vdupq_n_u8(1)); + ref_sum[2] = vdotq_u32(ref_sum[2], r2, vdupq_n_u8(1)); + + uint8x16_t abs_diff2 = vabdq_u8(s2, r2); + sse_u32[2] = vdotq_u32(sse_u32[2], abs_diff2, abs_diff2); + + uint8x16_t s3 = vld1q_u8(src + i + 48); + uint8x16_t r3 = vld1q_u8(ref + i + 48); + + src_sum[3] = vdotq_u32(src_sum[3], s3, vdupq_n_u8(1)); + ref_sum[3] = vdotq_u32(ref_sum[3], r3, vdupq_n_u8(1)); + + uint8x16_t abs_diff3 = vabdq_u8(s3, r3); + sse_u32[3] = vdotq_u32(sse_u32[3], abs_diff3, abs_diff3); + + i += 64; } while (i < w); src += src_stride; ref += ref_stride; } while (--h != 0); - int32x4_t sum_diff = - vsubq_s32(vreinterpretq_s32_u32(src_sum), vreinterpretq_s32_u32(ref_sum)); + src_sum[0] = horizontal_add_4d_u32x4(src_sum); + ref_sum[0] = horizontal_add_4d_u32x4(ref_sum); + int32x4_t sum_diff = vsubq_s32(vreinterpretq_s32_u32(src_sum[0]), + vreinterpretq_s32_u32(ref_sum[0])); *sum = horizontal_add_s32x4(sum_diff); - *sse = horizontal_add_u32x4(sse_u32); -} - -static inline void variance_32xh_neon_dotprod(const uint8_t *src, - int src_stride, - const uint8_t *ref, - int ref_stride, int h, - uint32_t *sse, int *sum) { - variance_large_neon_dotprod(src, src_stride, ref, ref_stride, 32, h, sse, - sum); + *sse = horizontal_add_u32x4(horizontal_add_4d_u32x4(sse_u32)); } static inline void variance_64xh_neon_dotprod(const uint8_t *src,