Skip to content

Commit 398ac42

Browse files
authored
Clean up and optimize SIMD kernels (#241)
* Remove unused SAD and diff kernels Signed-off-by: KP Choi <kp5.choi@samsung.com> * Saturate the quantized level in the NEON quant kernel Signed-off-by: KP Choi <kp5.choi@samsung.com> * Complete NEON kernel coverage for dquant, itx_part, itx_adj, and had8x8 Signed-off-by: KP Choi <kp5.choi@samsung.com> * Add an AVX2 version of the DC removed Hadamard kernel Signed-off-by: KP Choi <kp5.choi@samsung.com> * Speed up the AVX quant kernel and simplify dquant Signed-off-by: KP Choi <kp5.choi@samsung.com> --------- Signed-off-by: KP Choi <kp5.choi@samsung.com>
1 parent cd96e35 commit 398ac42

11 files changed

Lines changed: 388 additions & 536 deletions

File tree

src/avx/oapv_sad_avx.c

Lines changed: 85 additions & 115 deletions
Original file line numberDiff line numberDiff line change
@@ -34,72 +34,6 @@
3434
#if X86_SSE
3535

3636
/* SAD ***********************************************************************/
37-
static int sad_16b_avx_8x8(int w, int h, void* src1, void* src2, int s_src1, int s_src2)
38-
{
39-
s16* s1 = (s16*)src1;
40-
s16* s2 = (s16*)src2;
41-
__m256i zero_vector = _mm256_setzero_si256();
42-
__m256i s1_vector, s2_vector, diff_vector, diff_abs1, diff_abs2;
43-
// Because we are working with 16 elements at a time, stride is multiplied by 2.
44-
s16 s1_stride = 2 * s_src1;
45-
s16 s2_stride = 2 * s_src2;
46-
{ // Row 0 and Row 1
47-
// Load Row 0 and Row 1 data into registers.
48-
s1_vector = _mm256_loadu_si256((const __m256i*)(s1));
49-
s1 += s1_stride;
50-
s2_vector = _mm256_loadu_si256((const __m256i*)(s2));
51-
s2 += s2_stride;
52-
// Calculate absolute difference between two rows.
53-
diff_vector = _mm256_sub_epi16(s1_vector, s2_vector);
54-
diff_abs1 = _mm256_abs_epi16(diff_vector);
55-
}
56-
{ // Row 2 and Row 3
57-
s1_vector = _mm256_loadu_si256((const __m256i*)(s1));
58-
s1 += s1_stride;
59-
s2_vector = _mm256_loadu_si256((const __m256i*)(s2));
60-
s2 += s2_stride;
61-
diff_vector = _mm256_sub_epi16(s1_vector, s2_vector);
62-
diff_abs2 = _mm256_abs_epi16(diff_vector);
63-
}
64-
// Add absolute differences to running total.
65-
__m256i sum = _mm256_add_epi16(diff_abs1, diff_abs2);
66-
{ // Row 4 and Row 5
67-
s1_vector = _mm256_loadu_si256((const __m256i*)(s1));
68-
s1 += s1_stride;
69-
s2_vector = _mm256_loadu_si256((const __m256i*)(s2));
70-
s2 += s2_stride;
71-
diff_vector = _mm256_sub_epi16(s1_vector, s2_vector);
72-
diff_abs2 = _mm256_abs_epi16(diff_vector);
73-
sum = _mm256_add_epi16(sum, diff_abs2);
74-
}
75-
{ // Row 6 and Row 7
76-
s1_vector = _mm256_loadu_si256((const __m256i*)(s1));
77-
s2_vector = _mm256_loadu_si256((const __m256i*)(s2));
78-
diff_vector = _mm256_sub_epi16(s1_vector, s2_vector);
79-
diff_abs2 = _mm256_abs_epi16(diff_vector);
80-
sum = _mm256_add_epi16(sum, diff_abs2);
81-
}
82-
// Convert 16-bit integers to 32-bit integers for summation.
83-
__m128i sum_low = _mm256_extracti128_si256(sum, 0);
84-
__m128i sum_high = _mm256_extracti128_si256(sum, 1);
85-
__m256i sum_low_32 = _mm256_cvtepi16_epi32(sum_low);
86-
__m256i sum_high_32 = _mm256_cvtepi16_epi32(sum_high);
87-
// Sum up all the values in the array to get final SAD value.
88-
sum = _mm256_add_epi32(sum_low_32, sum_high_32);
89-
__m256i sum_hadd = _mm256_hadd_epi32(sum, zero_vector); // Horizontal add with zeros
90-
sum = _mm256_hadd_epi32(sum_hadd, zero_vector); // Horizontal add with zeros
91-
int sum1 = _mm256_extract_epi32(sum, 0);
92-
int sum2 = _mm256_extract_epi32(sum, 4);
93-
int sad = sum1 + sum2;
94-
return sad;
95-
}
96-
97-
const oapv_fn_sad_t oapv_tbl_fn_sad_16b_avx[2] =
98-
{
99-
sad_16b_avx_8x8,
100-
NULL
101-
};
102-
10337
/* SSD ***********************************************************************/
10438
static s64 ssd_16b_avx_8x8(int w, int h, void* src1, void* src2, int s_src1, int s_src2)
10539
{
@@ -165,56 +99,92 @@ const oapv_fn_ssd_t oapv_tbl_fn_ssd_16b_avx[2] =
16599
NULL
166100
};
167101

168-
/* DIFF ***********************************************************************/
169-
static void diff_16b_avx_8x8(int w, int h, void* src1, void* src2, int s_src1, int s_src2, int s_diff, s16 *diff)
102+
int oapv_dc_removed_had8x8_avx(pel* org, int s_org)
170103
{
171-
s16* s1 = (s16*)src1;
172-
s16* s2 = (s16*)src2;
173-
__m256i s1_vector, s2_vector, diff_vector;
174-
// Because we are working with 16 elements at a time, stride is multiplied by 2.
175-
s16 s1_stride = 2 * s_src1;
176-
s16 s2_stride = 2 * s_src2;
177-
s16 diff_stride = 2 * s_diff;
178-
{ // Row 0 and Row 1
179-
// Load Row 0 and Row 1 data into registers.
180-
s1_vector = _mm256_loadu_si256((const __m256i*)(s1));
181-
s1 += s1_stride;
182-
s2_vector = _mm256_loadu_si256((const __m256i*)(s2));
183-
s2 += s2_stride;
184-
// Calculate difference between two rows and store it in diff buffer.
185-
diff_vector = _mm256_sub_epi16(s1_vector, s2_vector);
186-
_mm256_storeu_si256((__m256i*)diff, diff_vector);
187-
diff += diff_stride;
188-
}
189-
{ // Row 2 and Row 3
190-
s1_vector = _mm256_loadu_si256((const __m256i*)(s1));
191-
s1 += s1_stride;
192-
s2_vector = _mm256_loadu_si256((const __m256i*)(s2));
193-
s2 += s2_stride;
194-
diff_vector = _mm256_sub_epi16(s1_vector, s2_vector);
195-
_mm256_storeu_si256((__m256i*)diff, diff_vector);
196-
diff += diff_stride;
197-
}
198-
{ // Row 4 and Row 5
199-
s1_vector = _mm256_loadu_si256((const __m256i*)(s1));
200-
s1 += s1_stride;
201-
s2_vector = _mm256_loadu_si256((const __m256i*)(s2));
202-
s2 += s2_stride;
203-
diff_vector = _mm256_sub_epi16(s1_vector, s2_vector);
204-
_mm256_storeu_si256((__m256i*)diff, diff_vector);
205-
diff += diff_stride;
206-
}
207-
{ // Row 6 and Row 7
208-
s1_vector = _mm256_loadu_si256((const __m256i*)(s1));
209-
s2_vector = _mm256_loadu_si256((const __m256i*)(s2));
210-
diff_vector = _mm256_sub_epi16(s1_vector, s2_vector);
211-
_mm256_storeu_si256((__m256i*)diff, diff_vector);
212-
}
104+
/* first pass is register-wise on 128-bit row vectors; after a transpose
105+
the second pass runs register-wise on 256-bit s32 vectors, since its
106+
values can reach 64 * 4095 and do not fit in s16 */
107+
__m128i r0 = _mm_loadu_si128((__m128i*)(org)); org += s_org;
108+
__m128i r1 = _mm_loadu_si128((__m128i*)(org)); org += s_org;
109+
__m128i r2 = _mm_loadu_si128((__m128i*)(org)); org += s_org;
110+
__m128i r3 = _mm_loadu_si128((__m128i*)(org)); org += s_org;
111+
__m128i r4 = _mm_loadu_si128((__m128i*)(org)); org += s_org;
112+
__m128i r5 = _mm_loadu_si128((__m128i*)(org)); org += s_org;
113+
__m128i r6 = _mm_loadu_si128((__m128i*)(org)); org += s_org;
114+
__m128i r7 = _mm_loadu_si128((__m128i*)(org));
115+
116+
/* pass 1: vertical butterflies */
117+
__m128i a0 = _mm_add_epi16(r0, r4), a4 = _mm_sub_epi16(r0, r4);
118+
__m128i a1 = _mm_add_epi16(r1, r5), a5 = _mm_sub_epi16(r1, r5);
119+
__m128i a2 = _mm_add_epi16(r2, r6), a6 = _mm_sub_epi16(r2, r6);
120+
__m128i a3 = _mm_add_epi16(r3, r7), a7 = _mm_sub_epi16(r3, r7);
121+
122+
__m128i b0 = _mm_add_epi16(a0, a2), b2 = _mm_sub_epi16(a0, a2);
123+
__m128i b1 = _mm_add_epi16(a1, a3), b3 = _mm_sub_epi16(a1, a3);
124+
__m128i b4 = _mm_add_epi16(a4, a6), b6 = _mm_sub_epi16(a4, a6);
125+
__m128i b5 = _mm_add_epi16(a5, a7), b7 = _mm_sub_epi16(a5, a7);
126+
127+
r0 = _mm_add_epi16(b0, b1); r1 = _mm_sub_epi16(b0, b1);
128+
r2 = _mm_add_epi16(b2, b3); r3 = _mm_sub_epi16(b2, b3);
129+
r4 = _mm_add_epi16(b4, b5); r5 = _mm_sub_epi16(b4, b5);
130+
r6 = _mm_add_epi16(b6, b7); r7 = _mm_sub_epi16(b6, b7);
131+
132+
/* 8x8 s16 transpose */
133+
__m128i t0 = _mm_unpacklo_epi16(r0, r1), t1 = _mm_unpackhi_epi16(r0, r1);
134+
__m128i t2 = _mm_unpacklo_epi16(r2, r3), t3 = _mm_unpackhi_epi16(r2, r3);
135+
__m128i t4 = _mm_unpacklo_epi16(r4, r5), t5 = _mm_unpackhi_epi16(r4, r5);
136+
__m128i t6 = _mm_unpacklo_epi16(r6, r7), t7 = _mm_unpackhi_epi16(r6, r7);
137+
__m128i u0 = _mm_unpacklo_epi32(t0, t2), u1 = _mm_unpackhi_epi32(t0, t2);
138+
__m128i u2 = _mm_unpacklo_epi32(t1, t3), u3 = _mm_unpackhi_epi32(t1, t3);
139+
__m128i u4 = _mm_unpacklo_epi32(t4, t6), u5 = _mm_unpackhi_epi32(t4, t6);
140+
__m128i u6 = _mm_unpacklo_epi32(t5, t7), u7 = _mm_unpackhi_epi32(t5, t7);
141+
r0 = _mm_unpacklo_epi64(u0, u4); r1 = _mm_unpackhi_epi64(u0, u4);
142+
r2 = _mm_unpacklo_epi64(u1, u5); r3 = _mm_unpackhi_epi64(u1, u5);
143+
r4 = _mm_unpacklo_epi64(u2, u6); r5 = _mm_unpackhi_epi64(u2, u6);
144+
r6 = _mm_unpacklo_epi64(u3, u7); r7 = _mm_unpackhi_epi64(u3, u7);
145+
146+
/* widen to s32 and run pass 2 register-wise */
147+
__m256i w0 = _mm256_cvtepi16_epi32(r0);
148+
__m256i w1 = _mm256_cvtepi16_epi32(r1);
149+
__m256i w2 = _mm256_cvtepi16_epi32(r2);
150+
__m256i w3 = _mm256_cvtepi16_epi32(r3);
151+
__m256i w4 = _mm256_cvtepi16_epi32(r4);
152+
__m256i w5 = _mm256_cvtepi16_epi32(r5);
153+
__m256i w6 = _mm256_cvtepi16_epi32(r6);
154+
__m256i w7 = _mm256_cvtepi16_epi32(r7);
155+
156+
__m256i c0 = _mm256_add_epi32(w0, w4), c4 = _mm256_sub_epi32(w0, w4);
157+
__m256i c1 = _mm256_add_epi32(w1, w5), c5 = _mm256_sub_epi32(w1, w5);
158+
__m256i c2 = _mm256_add_epi32(w2, w6), c6 = _mm256_sub_epi32(w2, w6);
159+
__m256i c3 = _mm256_add_epi32(w3, w7), c7 = _mm256_sub_epi32(w3, w7);
160+
161+
__m256i d0 = _mm256_add_epi32(c0, c2), d2 = _mm256_sub_epi32(c0, c2);
162+
__m256i d1 = _mm256_add_epi32(c1, c3), d3 = _mm256_sub_epi32(c1, c3);
163+
__m256i d4 = _mm256_add_epi32(c4, c6), d6 = _mm256_sub_epi32(c4, c6);
164+
__m256i d5 = _mm256_add_epi32(c5, c7), d7 = _mm256_sub_epi32(c5, c7);
165+
166+
w0 = _mm256_abs_epi32(_mm256_add_epi32(d0, d1));
167+
w1 = _mm256_abs_epi32(_mm256_sub_epi32(d0, d1));
168+
w2 = _mm256_abs_epi32(_mm256_add_epi32(d2, d3));
169+
w3 = _mm256_abs_epi32(_mm256_sub_epi32(d2, d3));
170+
w4 = _mm256_abs_epi32(_mm256_add_epi32(d4, d5));
171+
w5 = _mm256_abs_epi32(_mm256_sub_epi32(d4, d5));
172+
w6 = _mm256_abs_epi32(_mm256_add_epi32(d6, d7));
173+
w7 = _mm256_abs_epi32(_mm256_sub_epi32(d6, d7));
174+
175+
__m256i sum = _mm256_add_epi32(w0, w1);
176+
sum = _mm256_add_epi32(sum, _mm256_add_epi32(w2, w3));
177+
sum = _mm256_add_epi32(sum, _mm256_add_epi32(w4, w5));
178+
sum = _mm256_add_epi32(sum, _mm256_add_epi32(w6, w7));
179+
180+
__m128i s128 = _mm_add_epi32(_mm256_castsi256_si128(sum), _mm256_extracti128_si256(sum, 1));
181+
s128 = _mm_add_epi32(s128, _mm_srli_si128(s128, 8));
182+
s128 = _mm_add_epi32(s128, _mm_srli_si128(s128, 4));
183+
184+
int satd = _mm_cvtsi128_si32(s128) - _mm_cvtsi128_si32(_mm256_castsi256_si128(w0)); // remove DC
185+
return (satd + 2) >> 2;
213186
}
214187

215-
const oapv_fn_diff_t oapv_tbl_fn_diff_16b_avx[2] =
216-
{
217-
diff_16b_avx_8x8,
218-
NULL
219-
};
188+
189+
/* DIFF ***********************************************************************/
220190
#endif

src/avx/oapv_sad_avx.h

Lines changed: 1 addition & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -36,9 +36,8 @@
3636
#if X86_SSE
3737
#include <immintrin.h>
3838

39-
extern const oapv_fn_sad_t oapv_tbl_fn_sad_16b_avx[2];
4039
extern const oapv_fn_ssd_t oapv_tbl_fn_ssd_16b_avx[2];
41-
extern const oapv_fn_diff_t oapv_tbl_fn_diff_16b_avx[2];
40+
int oapv_dc_removed_had8x8_avx(pel* org, int s_org);
4241
#endif /* X86_SSE */
4342

4443
#endif /* _OAPV_SAD_AVX_H_ */

0 commit comments

Comments
 (0)