Coverage Report

Created: 2026-09-28 07:02

next uncovered line (L), next uncovered region (R), next uncovered branch (B)
/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
}