/src/xnnpack/src/f32-gemm/gen/f32-gemm-4x2c4-minmax-sse.c
Line | Count | Source |
1 | | // clang-format off |
2 | | // Auto-generated file. Do not edit! |
3 | | // Template: src/f32-gemm/MRx2c4-sse.c.in |
4 | | // Generator: tools/xngen |
5 | | // |
6 | | // Copyright 2019 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 | | #include <assert.h> |
12 | | #include <stddef.h> |
13 | | #include <stdint.h> |
14 | | |
15 | | #include <xmmintrin.h> |
16 | | |
17 | | #include "src/xnnpack/common.h" |
18 | | #include "src/xnnpack/microparams.h" |
19 | | #include "src/xnnpack/gemm.h" |
20 | | |
21 | | |
22 | | void xnn_f32_gemm_minmax_ukernel_4x2c4__sse( |
23 | | size_t mr, |
24 | | size_t nc, |
25 | | size_t kc, |
26 | | const float* restrict a, |
27 | | size_t a_stride, |
28 | | const float* restrict w, |
29 | | float* restrict c, |
30 | | size_t cm_stride, |
31 | | size_t cn_stride, |
32 | | const struct xnn_f32_minmax_params* restrict params) XNN_OOB_READS |
33 | 0 | { |
34 | 0 | assert(mr != 0); |
35 | 0 | assert(mr <= 4); |
36 | 0 | assert(nc != 0); |
37 | 0 | assert(kc != 0); |
38 | 0 | assert(kc % sizeof(float) == 0); |
39 | 0 | assert(a != NULL); |
40 | 0 | assert(w != NULL); |
41 | 0 | assert(c != NULL); |
42 | | |
43 | 0 | const float* a0 = a; |
44 | 0 | float* c0 = c; |
45 | 0 | const float* a1 = (const float*) ((uintptr_t) a0 + a_stride); |
46 | 0 | float* c1 = (float*) ((uintptr_t) c0 + cm_stride); |
47 | 0 | if XNN_UNPREDICTABLE(mr < 2) { |
48 | 0 | a1 = a0; |
49 | 0 | c1 = c0; |
50 | 0 | } |
51 | 0 | const float* a2 = (const float*) ((uintptr_t) a1 + a_stride); |
52 | 0 | float* c2 = (float*) ((uintptr_t) c1 + cm_stride); |
53 | 0 | if XNN_UNPREDICTABLE(mr <= 2) { |
54 | 0 | a2 = a1; |
55 | 0 | c2 = c1; |
56 | 0 | } |
57 | 0 | const float* a3 = (const float*) ((uintptr_t) a2 + a_stride); |
58 | 0 | float* c3 = (float*) ((uintptr_t) c2 + cm_stride); |
59 | 0 | if XNN_UNPREDICTABLE(mr != 4) { |
60 | 0 | a3 = a2; |
61 | 0 | c3 = c2; |
62 | 0 | } |
63 | |
|
64 | 0 | const __m128 vmin = _mm_set1_ps(params->scalar.min); |
65 | 0 | const __m128 vmax = _mm_set1_ps(params->scalar.max); |
66 | 0 | XNN_FORCE_REALIZATION(vmin); |
67 | 0 | XNN_FORCE_REALIZATION(vmax); |
68 | |
|
69 | 0 | do { |
70 | 0 | __m128 vacc0x0c4 = _mm_load_ss(w); |
71 | 0 | __m128 vacc0x1c4 = _mm_load_ss(w + 1); |
72 | 0 | __m128 vacc1x0c4 = vacc0x0c4; |
73 | 0 | __m128 vacc1x1c4 = vacc0x1c4; |
74 | 0 | __m128 vacc2x0c4 = vacc0x0c4; |
75 | 0 | __m128 vacc2x1c4 = vacc0x1c4; |
76 | 0 | __m128 vacc3x0c4 = vacc0x0c4; |
77 | 0 | __m128 vacc3x1c4 = vacc0x1c4; |
78 | 0 | w += 2; |
79 | |
|
80 | 0 | size_t k = kc; |
81 | 0 | for (; k >= 4 * sizeof(float); k -= 4 * sizeof(float)) { |
82 | 0 | const __m128 va0 = _mm_loadu_ps(a0); |
83 | 0 | a0 += 4; |
84 | 0 | const __m128 va1 = _mm_loadu_ps(a1); |
85 | 0 | a1 += 4; |
86 | 0 | const __m128 va2 = _mm_loadu_ps(a2); |
87 | 0 | a2 += 4; |
88 | 0 | const __m128 va3 = _mm_loadu_ps(a3); |
89 | 0 | a3 += 4; |
90 | |
|
91 | 0 | const __m128 vb0 = _mm_loadu_ps(w); |
92 | 0 | const __m128 vb1 = _mm_loadu_ps(w + 4); |
93 | 0 | w += 8; |
94 | |
|
95 | 0 | vacc0x0c4 = _mm_add_ps(vacc0x0c4, _mm_mul_ps(va0, vb0)); |
96 | 0 | vacc0x1c4 = _mm_add_ps(vacc0x1c4, _mm_mul_ps(va0, vb1)); |
97 | 0 | vacc1x0c4 = _mm_add_ps(vacc1x0c4, _mm_mul_ps(va1, vb0)); |
98 | 0 | vacc1x1c4 = _mm_add_ps(vacc1x1c4, _mm_mul_ps(va1, vb1)); |
99 | 0 | vacc2x0c4 = _mm_add_ps(vacc2x0c4, _mm_mul_ps(va2, vb0)); |
100 | 0 | vacc2x1c4 = _mm_add_ps(vacc2x1c4, _mm_mul_ps(va2, vb1)); |
101 | 0 | vacc3x0c4 = _mm_add_ps(vacc3x0c4, _mm_mul_ps(va3, vb0)); |
102 | 0 | vacc3x1c4 = _mm_add_ps(vacc3x1c4, _mm_mul_ps(va3, vb1)); |
103 | 0 | } |
104 | 0 | if XNN_UNLIKELY(k != 0) { |
105 | 0 | const __m128 va0 = _mm_loadu_ps(a0); |
106 | 0 | a0 = (const float*) ((uintptr_t) a0 + k); |
107 | 0 | const __m128 va1 = _mm_loadu_ps(a1); |
108 | 0 | a1 = (const float*) ((uintptr_t) a1 + k); |
109 | 0 | const __m128 va2 = _mm_loadu_ps(a2); |
110 | 0 | a2 = (const float*) ((uintptr_t) a2 + k); |
111 | 0 | const __m128 va3 = _mm_loadu_ps(a3); |
112 | 0 | a3 = (const float*) ((uintptr_t) a3 + k); |
113 | |
|
114 | 0 | const __m128 vb0 = _mm_loadu_ps(w); |
115 | 0 | const __m128 vb1 = _mm_loadu_ps(w + 4); |
116 | 0 | w += 8; |
117 | |
|
118 | 0 | const __m128 vmask0 = _mm_cmpeq_ps(_mm_setzero_ps(), vb0); |
119 | 0 | const __m128 vmask1 = _mm_cmpeq_ps(_mm_setzero_ps(), vb1); |
120 | |
|
121 | 0 | vacc0x0c4 = _mm_add_ps(vacc0x0c4, _mm_mul_ps(_mm_andnot_ps(vmask0, va0), vb0)); |
122 | 0 | vacc0x1c4 = _mm_add_ps(vacc0x1c4, _mm_mul_ps(_mm_andnot_ps(vmask1, va0), vb1)); |
123 | 0 | vacc1x0c4 = _mm_add_ps(vacc1x0c4, _mm_mul_ps(_mm_andnot_ps(vmask0, va1), vb0)); |
124 | 0 | vacc1x1c4 = _mm_add_ps(vacc1x1c4, _mm_mul_ps(_mm_andnot_ps(vmask1, va1), vb1)); |
125 | 0 | vacc2x0c4 = _mm_add_ps(vacc2x0c4, _mm_mul_ps(_mm_andnot_ps(vmask0, va2), vb0)); |
126 | 0 | vacc2x1c4 = _mm_add_ps(vacc2x1c4, _mm_mul_ps(_mm_andnot_ps(vmask1, va2), vb1)); |
127 | 0 | vacc3x0c4 = _mm_add_ps(vacc3x0c4, _mm_mul_ps(_mm_andnot_ps(vmask0, va3), vb0)); |
128 | 0 | vacc3x1c4 = _mm_add_ps(vacc3x1c4, _mm_mul_ps(_mm_andnot_ps(vmask1, va3), vb1)); |
129 | 0 | } |
130 | |
|
131 | 0 | const __m128 vacc0x01c2 = _mm_add_ps(_mm_unpacklo_ps(vacc0x0c4, vacc0x1c4), _mm_unpackhi_ps(vacc0x0c4, vacc0x1c4)); |
132 | 0 | const __m128 vacc1x01c2 = _mm_add_ps(_mm_unpacklo_ps(vacc1x0c4, vacc1x1c4), _mm_unpackhi_ps(vacc1x0c4, vacc1x1c4)); |
133 | 0 | const __m128 vacc2x01c2 = _mm_add_ps(_mm_unpacklo_ps(vacc2x0c4, vacc2x1c4), _mm_unpackhi_ps(vacc2x0c4, vacc2x1c4)); |
134 | 0 | const __m128 vacc3x01c2 = _mm_add_ps(_mm_unpacklo_ps(vacc3x0c4, vacc3x1c4), _mm_unpackhi_ps(vacc3x0c4, vacc3x1c4)); |
135 | |
|
136 | 0 | __m128 vacc01x01 = _mm_add_ps(_mm_movelh_ps(vacc0x01c2, vacc1x01c2), _mm_movehl_ps(vacc1x01c2, vacc0x01c2)); |
137 | 0 | __m128 vacc23x01 = _mm_add_ps(_mm_movelh_ps(vacc2x01c2, vacc3x01c2), _mm_movehl_ps(vacc3x01c2, vacc2x01c2)); |
138 | |
|
139 | 0 | vacc01x01 = _mm_min_ps(vacc01x01, vmax); |
140 | 0 | vacc23x01 = _mm_min_ps(vacc23x01, vmax); |
141 | |
|
142 | 0 | vacc01x01 = _mm_max_ps(vacc01x01, vmin); |
143 | 0 | vacc23x01 = _mm_max_ps(vacc23x01, vmin); |
144 | |
|
145 | 0 | if XNN_LIKELY(nc >= 2) { |
146 | 0 | _mm_storel_pi((__m64*) c0, vacc01x01); |
147 | 0 | c0 = (float*) ((uintptr_t) c0 + cn_stride); |
148 | 0 | a0 = (const float*) ((uintptr_t) a0 - kc); |
149 | 0 | _mm_storeh_pi((__m64*) c1, vacc01x01); |
150 | 0 | c1 = (float*) ((uintptr_t) c1 + cn_stride); |
151 | 0 | a1 = (const float*) ((uintptr_t) a1 - kc); |
152 | 0 | _mm_storel_pi((__m64*) c2, vacc23x01); |
153 | 0 | c2 = (float*) ((uintptr_t) c2 + cn_stride); |
154 | 0 | a2 = (const float*) ((uintptr_t) a2 - kc); |
155 | 0 | _mm_storeh_pi((__m64*) c3, vacc23x01); |
156 | 0 | c3 = (float*) ((uintptr_t) c3 + cn_stride); |
157 | 0 | a3 = (const float*) ((uintptr_t) a3 - kc); |
158 | |
|
159 | 0 | nc -= 2; |
160 | 0 | } else { |
161 | 0 | assert(nc == 1); |
162 | 0 | _mm_store_ss(c0, vacc01x01); |
163 | 0 | _mm_store_ss(c1, _mm_movehl_ps(vacc01x01, vacc01x01)); |
164 | 0 | _mm_store_ss(c2, vacc23x01); |
165 | 0 | _mm_store_ss(c3, _mm_movehl_ps(vacc23x01, vacc23x01)); |
166 | |
|
167 | 0 | nc = 0; |
168 | 0 | } |
169 | 0 | } while (nc != 0); |
170 | 0 | } |