/src/xnnpack/src/x32-packw/gen/x32-packw-x8-gemm-goi-avx-u4.c
Line | Count | Source |
1 | | // clang-format off |
2 | | // Auto-generated file. Do not edit! |
3 | | // Template: src/x32-packw/avx.c.in |
4 | | // Generator: tools/xngen |
5 | | // |
6 | | // Copyright 2023 Google LLC |
7 | | // |
8 | | // This source code is licensed under the BSD-style license found in the |
9 | | // LICENSE file in the root directory of this source tree. |
10 | | |
11 | | |
12 | | #include <assert.h> |
13 | | #include <stddef.h> |
14 | | #include <stdint.h> |
15 | | |
16 | | #include <immintrin.h> |
17 | | |
18 | | #include "src/xnnpack/common.h" |
19 | | #include "src/xnnpack/intrinsics-polyfill.h" |
20 | | #include "src/xnnpack/packw.h" |
21 | | |
22 | | |
23 | | void xnn_x32_packw_gemm_goi_ukernel_x8__avx_u4( |
24 | | size_t g, |
25 | | size_t nc, |
26 | | size_t kc, |
27 | | size_t nr, |
28 | | size_t kr, |
29 | | size_t sr, |
30 | | size_t n_stride, |
31 | | const uint32_t* weights, |
32 | | const uint32_t* bias, |
33 | | const void* scale, |
34 | | uint32_t* packed_weights, |
35 | | size_t extra_bytes, |
36 | | const void* params) |
37 | 0 | { |
38 | 0 | assert(g != 0); |
39 | 0 | assert(nc != 0); |
40 | 0 | assert(kc != 0); |
41 | 0 | assert(nr == 8); // This kernel is for NR=8 |
42 | 0 | assert(kr == 1); |
43 | 0 | assert(sr == 1); |
44 | 0 | assert(weights != NULL); |
45 | 0 | assert(packed_weights != NULL); |
46 | | |
47 | 0 | const float* b = (const float*) bias; |
48 | 0 | float* packed_w = (float*) packed_weights; |
49 | 0 | do { |
50 | | // NC main loop multiple of 8 |
51 | 0 | const float* w0 = (const float*) weights; |
52 | 0 | size_t n = nc; |
53 | |
|
54 | 0 | for (; n >= 8; n -= 8) { |
55 | 0 | if XNN_LIKELY(b != NULL) { |
56 | 0 | const __m256 vb0 = _mm256_loadu_ps(b); |
57 | 0 | _mm256_store_ps(packed_w, vb0); |
58 | 0 | b += 8; |
59 | 0 | } else { |
60 | 0 | const __m256 vzero = _mm256_setzero_ps(); |
61 | 0 | _mm256_store_ps(packed_w, vzero); |
62 | 0 | } |
63 | 0 | packed_w += 8; |
64 | |
|
65 | 0 | const float* w1 = w0 + n_stride; |
66 | 0 | const float* w2 = w1 + n_stride; |
67 | 0 | const float* w3 = w2 + n_stride; |
68 | 0 | const float* w4 = w3 + n_stride; |
69 | 0 | const float* w5 = w4 + n_stride; |
70 | 0 | const float* w6 = w5 + n_stride; |
71 | 0 | const float* w7 = w6 + n_stride; |
72 | | |
73 | | // KC main loop multiple of 4 |
74 | 0 | size_t k = kc; |
75 | 0 | for (; k >= 4; k -= 4) { |
76 | | // Read blocks of 4x4 |
77 | | // a b c d |
78 | | // e f g h |
79 | | // i j k l |
80 | | // m n o p |
81 | | // Load first 4 rows of N into low part of each register |
82 | 0 | __m256 v0x0123 = _mm256_castps128_ps256(_mm_loadu_ps(w0)); |
83 | 0 | w0 += 4; |
84 | 0 | __m256 v1x0123 = _mm256_castps128_ps256(_mm_loadu_ps(w1)); |
85 | 0 | w1 += 4; |
86 | 0 | __m256 v2x0123 = _mm256_castps128_ps256(_mm_loadu_ps(w2)); |
87 | 0 | w2 += 4; |
88 | 0 | __m256 v3x0123 = _mm256_castps128_ps256(_mm_loadu_ps(w3)); |
89 | 0 | w3 += 4; |
90 | | // Load next 4 rows of N into the high part of each register |
91 | 0 | v0x0123 = _mm256_insertf128_ps(v0x0123, _mm_loadu_ps(w4), 1); |
92 | 0 | w4 += 4; |
93 | 0 | v1x0123 = _mm256_insertf128_ps(v1x0123, _mm_loadu_ps(w5), 1); |
94 | 0 | w5 += 4; |
95 | 0 | v2x0123 = _mm256_insertf128_ps(v2x0123, _mm_loadu_ps(w6), 1); |
96 | 0 | w6 += 4; |
97 | 0 | v3x0123 = _mm256_insertf128_ps(v3x0123, _mm_loadu_ps(w7), 1); |
98 | 0 | w7 += 4; |
99 | | |
100 | | // Transpose 2x2 |
101 | 0 | const __m256 vtmp0x0123 = _mm256_unpacklo_ps(v0x0123, v1x0123); // a e b f from row 0, 1 |
102 | 0 | const __m256 vtmp1x0123 = _mm256_unpacklo_ps(v2x0123, v3x0123); // i m j n from row 2, 3 |
103 | 0 | const __m256 vtmp2x0123 = _mm256_unpackhi_ps(v0x0123, v1x0123); // c g d h from row 0, 1 |
104 | 0 | const __m256 vtmp3x0123 = _mm256_unpackhi_ps(v2x0123, v3x0123); // k o l p from row 2, 3 |
105 | | // Transpose 4x4 |
106 | 0 | v0x0123 = _mm256_castpd_ps(_mm256_unpacklo_pd(_mm256_castps_pd(vtmp0x0123), _mm256_castps_pd(vtmp1x0123))); // a e i m from row 0, 1 |
107 | 0 | v1x0123 = _mm256_castpd_ps(_mm256_unpackhi_pd(_mm256_castps_pd(vtmp0x0123), _mm256_castps_pd(vtmp1x0123))); // b f j n from row 0, 1 |
108 | 0 | v2x0123 = _mm256_castpd_ps(_mm256_unpacklo_pd(_mm256_castps_pd(vtmp2x0123), _mm256_castps_pd(vtmp3x0123))); // c g k o from row 2, 3 |
109 | 0 | v3x0123 = _mm256_castpd_ps(_mm256_unpackhi_pd(_mm256_castps_pd(vtmp2x0123), _mm256_castps_pd(vtmp3x0123))); // d h l p from row 2, 3 |
110 | |
|
111 | 0 | _mm256_store_ps(packed_w, v0x0123); |
112 | 0 | _mm256_store_ps(packed_w + 8, v1x0123); |
113 | 0 | _mm256_store_ps(packed_w + 16, v2x0123); |
114 | 0 | _mm256_store_ps(packed_w + 24, v3x0123); |
115 | 0 | packed_w += 32; |
116 | 0 | } |
117 | | |
118 | | // KC remainder (1..3) |
119 | 0 | if XNN_UNLIKELY(k != 0) { |
120 | 0 | assert(k >= 1); |
121 | 0 | assert(k <= 3); |
122 | 0 | if (k & 2) { |
123 | | // Read blocks of 4x2 |
124 | | // a b |
125 | | // c d |
126 | | // e f |
127 | | // g h |
128 | 0 | __m128 v0 = _mm_castsi128_ps(_mm_loadl_epi64((const __m128i*) w0)); |
129 | 0 | w0 += 2; |
130 | 0 | __m128 v1 = _mm_castsi128_ps(_mm_loadl_epi64((const __m128i*) w1)); |
131 | 0 | w1 += 2; |
132 | 0 | __m128 v2 = _mm_castsi128_ps(_mm_loadl_epi64((const __m128i*) w2)); |
133 | 0 | w2 += 2; |
134 | 0 | __m128 v3 = _mm_castsi128_ps(_mm_loadl_epi64((const __m128i*) w3)); |
135 | 0 | w3 += 2; |
136 | 0 | __m128 v4 = _mm_castsi128_ps(_mm_loadl_epi64((const __m128i*) w4)); |
137 | 0 | w4 += 2; |
138 | 0 | __m128 v5 = _mm_castsi128_ps(_mm_loadl_epi64((const __m128i*) w5)); |
139 | 0 | w5 += 2; |
140 | 0 | __m128 v6 = _mm_castsi128_ps(_mm_loadl_epi64((const __m128i*) w6)); |
141 | 0 | w6 += 2; |
142 | 0 | __m128 v7 = _mm_castsi128_ps(_mm_loadl_epi64((const __m128i*) w7)); |
143 | 0 | w7 += 2; |
144 | | |
145 | | // Transpose 2x2 |
146 | 0 | const __m128 vtmp0 = _mm_unpacklo_ps(v0, v1); // a c b d from row 0, 1 |
147 | 0 | const __m128 vtmp1 = _mm_unpacklo_ps(v2, v3); // e g f h from row 2, 3 |
148 | | // Transpose 4x4 |
149 | 0 | v0 = _mm_castpd_ps(_mm_unpacklo_pd(_mm_castps_pd(vtmp0), _mm_castps_pd(vtmp1))); // a c e g from row 0, 1 |
150 | 0 | v1 = _mm_castpd_ps(_mm_unpackhi_pd(_mm_castps_pd(vtmp0), _mm_castps_pd(vtmp1))); // b d f h from row 0, 1 |
151 | | // Transpose 2x2 |
152 | 0 | const __m128 vtmp4 = _mm_unpacklo_ps(v4, v5); // a c b d from row 0, 1 |
153 | 0 | const __m128 vtmp5 = _mm_unpacklo_ps(v6, v7); // e g f h from row 2, 3 |
154 | | // Transpose 4x4 |
155 | 0 | v4 = _mm_castpd_ps(_mm_unpacklo_pd(_mm_castps_pd(vtmp4), _mm_castps_pd(vtmp5))); // a c e g from row 0, 1 |
156 | 0 | v5 = _mm_castpd_ps(_mm_unpackhi_pd(_mm_castps_pd(vtmp4), _mm_castps_pd(vtmp5))); // b d f h from row 0, 1 |
157 | |
|
158 | 0 | _mm_store_ps(packed_w, v0); |
159 | 0 | _mm_store_ps(packed_w + 4, v4); |
160 | 0 | _mm_store_ps(packed_w + 8, v1); |
161 | 0 | _mm_store_ps(packed_w + 12, v5); |
162 | 0 | packed_w += 16; |
163 | 0 | } |
164 | 0 | if (k & 1) { |
165 | | // Read blocks of 4x1 |
166 | | // a |
167 | | // b |
168 | | // c |
169 | | // d |
170 | 0 | __m128 v0 = _mm_load_ss(w0); w0 += 1; |
171 | 0 | __m128 v1 = _mm_load_ss(w1); w1 += 1; |
172 | 0 | __m128 v2 = _mm_load_ss(w2); w2 += 1; |
173 | 0 | __m128 v3 = _mm_load_ss(w3); w3 += 1; |
174 | 0 | __m128 v4 = _mm_load_ss(w4); w4 += 1; |
175 | 0 | __m128 v5 = _mm_load_ss(w5); w5 += 1; |
176 | 0 | __m128 v6 = _mm_load_ss(w6); w6 += 1; |
177 | 0 | __m128 v7 = _mm_load_ss(w7); w7 += 1; |
178 | | |
179 | | // Transpose 2x2 |
180 | 0 | const __m128 vtmp0 = _mm_unpacklo_ps(v0, v1); // a b from row 0, 1 |
181 | 0 | const __m128 vtmp1 = _mm_unpacklo_ps(v2, v3); // c d from row 2, 3 |
182 | | // Transpose 4x4 |
183 | 0 | v0 = _mm_castpd_ps(_mm_unpacklo_pd(_mm_castps_pd(vtmp0), _mm_castps_pd(vtmp1))); // a b c d from row 0, 1 |
184 | | // Transpose 2x2 |
185 | 0 | const __m128 vtmp4 = _mm_unpacklo_ps(v4, v5); // a b from row 0, 1 |
186 | 0 | const __m128 vtmp5 = _mm_unpacklo_ps(v6, v7); // c d from row 2, 3 |
187 | | // Transpose 4x4 |
188 | 0 | v4 = _mm_castpd_ps(_mm_unpacklo_pd(_mm_castps_pd(vtmp4), _mm_castps_pd(vtmp5))); // a b c d from row 0, 1 |
189 | |
|
190 | 0 | _mm_store_ps(packed_w, v0); |
191 | 0 | _mm_store_ps(packed_w + 4, v4); |
192 | 0 | packed_w += 8; |
193 | 0 | } |
194 | 0 | } |
195 | 0 | packed_w = (float*) ((uintptr_t) packed_w + extra_bytes); |
196 | 0 | w0 = w7 + n_stride - kc; |
197 | 0 | } |
198 | | |
199 | | // NC remainder (1..7) |
200 | 0 | if XNN_UNLIKELY(n != 0) { |
201 | 0 | assert(n >= 1); |
202 | 0 | assert(n <= 7); |
203 | 0 | if XNN_LIKELY(b != NULL) { |
204 | 0 | size_t nb = n; |
205 | 0 | do { |
206 | 0 | *packed_w++ = *b++; |
207 | 0 | } while (--nb != 0); |
208 | 0 | packed_w += (8 - n); |
209 | 0 | } else { |
210 | 0 | const __m256 vzero = _mm256_setzero_ps(); |
211 | 0 | _mm256_store_ps(packed_w, vzero); |
212 | 0 | packed_w += 8; |
213 | 0 | } |
214 | | |
215 | | // NR remainder has less than 8 rows so last row is not loaded |
216 | | // For SR=4 the |
217 | 0 | const float* w1 = w0 + n_stride; |
218 | 0 | if XNN_UNPREDICTABLE(n < 2) { |
219 | 0 | w1 = w0; |
220 | 0 | } |
221 | 0 | const float* w2 = w1 + n_stride; |
222 | 0 | if XNN_UNPREDICTABLE(n <= 2) { |
223 | 0 | w2 = w1; |
224 | 0 | } |
225 | 0 | const float* w3 = w2 + n_stride; |
226 | 0 | if XNN_UNPREDICTABLE(n < 4) { |
227 | 0 | w3 = w2; |
228 | 0 | } |
229 | 0 | const float* w4 = w3 + n_stride; |
230 | 0 | if XNN_UNPREDICTABLE(n <= 4) { |
231 | 0 | w4 = w3; |
232 | 0 | } |
233 | 0 | const float* w5 = w4 + n_stride; |
234 | 0 | if XNN_UNPREDICTABLE(n < 6) { |
235 | 0 | w5 = w4; |
236 | 0 | } |
237 | 0 | const float* w6 = w5 + n_stride; |
238 | 0 | if XNN_UNPREDICTABLE(n <= 6) { |
239 | 0 | w6 = w5; |
240 | 0 | } |
241 | | |
242 | | // KC main loop multiple of 4 |
243 | 0 | size_t k = kc; |
244 | 0 | for (; k >= 4; k -= 4) { |
245 | | // Read blocks of 4x4 |
246 | | // a b c d |
247 | | // e f g h |
248 | | // i j k l |
249 | | // m n o p |
250 | | // Load first 4 rows of N into low part of each register |
251 | 0 | __m256 v0x0123 = _mm256_castps128_ps256(_mm_loadu_ps(w0)); |
252 | 0 | w0 += 4; |
253 | 0 | __m256 v1x0123 = _mm256_castps128_ps256(_mm_loadu_ps(w1)); |
254 | 0 | w1 += 4; |
255 | 0 | __m256 v2x0123 = _mm256_castps128_ps256(_mm_loadu_ps(w2)); |
256 | 0 | w2 += 4; |
257 | | // castps leaves upper 128 bits undefined, so zero them. |
258 | 0 | __m256 v3x0123 = _mm256_zextps128_ps256(_mm_loadu_ps(w3)); |
259 | 0 | w3 += 4; |
260 | | // Load next 4 rows of N into the high part of each register |
261 | 0 | v0x0123 = _mm256_insertf128_ps(v0x0123, _mm_loadu_ps(w4), 1); |
262 | 0 | w4 += 4; |
263 | 0 | v1x0123 = _mm256_insertf128_ps(v1x0123, _mm_loadu_ps(w5), 1); |
264 | 0 | w5 += 4; |
265 | 0 | v2x0123 = _mm256_insertf128_ps(v2x0123, _mm_loadu_ps(w6), 1); |
266 | 0 | w6 += 4; |
267 | | |
268 | | // Transpose 2x2 |
269 | 0 | const __m256 vtmp0x0123 = _mm256_unpacklo_ps(v0x0123, v1x0123); // a e b f from row 0, 1 |
270 | 0 | const __m256 vtmp1x0123 = _mm256_unpacklo_ps(v2x0123, v3x0123); // i m j n from row 2, 3 |
271 | 0 | const __m256 vtmp2x0123 = _mm256_unpackhi_ps(v0x0123, v1x0123); // c g d h from row 0, 1 |
272 | 0 | const __m256 vtmp3x0123 = _mm256_unpackhi_ps(v2x0123, v3x0123); // k o l p from row 2, 3 |
273 | | // Transpose 4x4 |
274 | 0 | v0x0123 = _mm256_castpd_ps(_mm256_unpacklo_pd(_mm256_castps_pd(vtmp0x0123), _mm256_castps_pd(vtmp1x0123))); // a e i m from row 0, 1 |
275 | 0 | v1x0123 = _mm256_castpd_ps(_mm256_unpackhi_pd(_mm256_castps_pd(vtmp0x0123), _mm256_castps_pd(vtmp1x0123))); // b f j n from row 0, 1 |
276 | 0 | v2x0123 = _mm256_castpd_ps(_mm256_unpacklo_pd(_mm256_castps_pd(vtmp2x0123), _mm256_castps_pd(vtmp3x0123))); // c g k o from row 2, 3 |
277 | 0 | v3x0123 = _mm256_castpd_ps(_mm256_unpackhi_pd(_mm256_castps_pd(vtmp2x0123), _mm256_castps_pd(vtmp3x0123))); // d h l p from row 2, 3 |
278 | |
|
279 | 0 | _mm256_store_ps(packed_w, v0x0123); |
280 | 0 | _mm256_store_ps(packed_w + 8, v1x0123); |
281 | 0 | _mm256_store_ps(packed_w + 16, v2x0123); |
282 | 0 | _mm256_store_ps(packed_w + 24, v3x0123); |
283 | 0 | packed_w += 32; |
284 | 0 | } |
285 | | |
286 | | // KC remainder (1..3) |
287 | 0 | if XNN_UNLIKELY(k != 0) { |
288 | 0 | assert(k >= 1); |
289 | 0 | assert(k <= 3); |
290 | 0 | if (k & 2) { |
291 | | // Read blocks of 4x2 |
292 | | // a b |
293 | | // c d |
294 | | // e f |
295 | | // g h |
296 | 0 | __m128 v0 = _mm_castsi128_ps(_mm_loadl_epi64((const __m128i*) w0)); |
297 | 0 | w0 += 2; |
298 | 0 | __m128 v1 = _mm_castsi128_ps(_mm_loadl_epi64((const __m128i*) w1)); |
299 | 0 | w1 += 2; |
300 | 0 | __m128 v2 = _mm_castsi128_ps(_mm_loadl_epi64((const __m128i*) w2)); |
301 | 0 | w2 += 2; |
302 | 0 | __m128 v3 = _mm_castsi128_ps(_mm_loadl_epi64((const __m128i*) w3)); |
303 | 0 | w3 += 2; |
304 | 0 | __m128 v4 = _mm_castsi128_ps(_mm_loadl_epi64((const __m128i*) w4)); |
305 | 0 | w4 += 2; |
306 | 0 | __m128 v5 = _mm_castsi128_ps(_mm_loadl_epi64((const __m128i*) w5)); |
307 | 0 | w5 += 2; |
308 | 0 | __m128 v6 = _mm_castsi128_ps(_mm_loadl_epi64((const __m128i*) w6)); |
309 | 0 | w6 += 2; |
310 | | |
311 | | // Transpose 2x2 |
312 | 0 | const __m128 vtmp0 = _mm_unpacklo_ps(v0, v1); // a c b d from row 0, 1 |
313 | 0 | const __m128 vtmp1 = _mm_unpacklo_ps(v2, v3); // e g f h from row 2, 3 |
314 | | // Transpose 4x4 |
315 | 0 | v0 = _mm_castpd_ps(_mm_unpacklo_pd(_mm_castps_pd(vtmp0), _mm_castps_pd(vtmp1))); // a c e g from row 0, 1 |
316 | 0 | v1 = _mm_castpd_ps(_mm_unpackhi_pd(_mm_castps_pd(vtmp0), _mm_castps_pd(vtmp1))); // b d f h from row 0, 1 |
317 | | // Transpose 2x2 |
318 | 0 | const __m128 vtmp4 = _mm_unpacklo_ps(v4, v5); // a c b d from row 0, 1 |
319 | 0 | const __m128 vtmp5 = _mm_unpacklo_ps(v6, v6); // e g f h from row 2, 3 |
320 | | // Transpose 4x4 |
321 | 0 | v4 = _mm_castpd_ps(_mm_unpacklo_pd(_mm_castps_pd(vtmp4), _mm_castps_pd(vtmp5))); // a c e g from row 0, 1 |
322 | 0 | v5 = _mm_castpd_ps(_mm_unpackhi_pd(_mm_castps_pd(vtmp4), _mm_castps_pd(vtmp5))); // b d f h from row 0, 1 |
323 | |
|
324 | 0 | _mm_store_ps(packed_w, v0); |
325 | 0 | _mm_store_ps(packed_w + 4, v4); |
326 | 0 | _mm_store_ps(packed_w + 8, v1); |
327 | 0 | _mm_store_ps(packed_w + 12, v5); |
328 | 0 | packed_w += 16; |
329 | 0 | } |
330 | 0 | if (k & 1) { |
331 | | // Read blocks of 4x1 |
332 | | // a |
333 | | // b |
334 | | // c |
335 | | // d |
336 | 0 | __m128 v0 = _mm_load_ss(w0); w0 += 1; |
337 | 0 | __m128 v1 = _mm_load_ss(w1); w1 += 1; |
338 | 0 | __m128 v2 = _mm_load_ss(w2); w2 += 1; |
339 | 0 | __m128 v3 = _mm_load_ss(w3); w3 += 1; |
340 | 0 | __m128 v4 = _mm_load_ss(w4); w4 += 1; |
341 | 0 | __m128 v5 = _mm_load_ss(w5); w5 += 1; |
342 | 0 | __m128 v6 = _mm_load_ss(w6); w6 += 1; |
343 | | |
344 | | // Transpose 2x2 |
345 | 0 | const __m128 vtmp0 = _mm_unpacklo_ps(v0, v1); // a b from row 0, 1 |
346 | 0 | const __m128 vtmp1 = _mm_unpacklo_ps(v2, v3); // c d from row 2, 3 |
347 | | // Transpose 4x4 |
348 | 0 | v0 = _mm_castpd_ps(_mm_unpacklo_pd(_mm_castps_pd(vtmp0), _mm_castps_pd(vtmp1))); // a b c d from row 0, 1 |
349 | | // Transpose 2x2 |
350 | 0 | const __m128 vtmp4 = _mm_unpacklo_ps(v4, v5); // a b from row 0, 1 |
351 | 0 | const __m128 vtmp5 = _mm_unpacklo_ps(v6, v6); // c d from row 2, 3 |
352 | | // Transpose 4x4 |
353 | 0 | v4 = _mm_castpd_ps(_mm_unpacklo_pd(_mm_castps_pd(vtmp4), _mm_castps_pd(vtmp5))); // a b c d from row 0, 1 |
354 | |
|
355 | 0 | _mm_store_ps(packed_w, v0); |
356 | 0 | _mm_store_ps(packed_w + 4, v4); |
357 | 0 | packed_w += 8; |
358 | 0 | } |
359 | 0 | } |
360 | 0 | } |
361 | 0 | weights += nc * n_stride; |
362 | 0 | } while (--g != 0); |
363 | 0 | } |