/src/xnnpack/src/f16-dwconv/gen/f16-dwconv-4p16c-minmax-fma3.c
Line | Count | Source |
1 | | // clang-format off |
2 | | // Auto-generated file. Do not edit! |
3 | | // Template: src/f16-dwconv/unipass-fma3.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 <immintrin.h> |
16 | | |
17 | | #include "src/xnnpack/common.h" |
18 | | #include "src/xnnpack/dwconv.h" |
19 | | #include "src/xnnpack/intrinsics-polyfill.h" |
20 | | #include "src/xnnpack/math.h" |
21 | | #include "src/xnnpack/microparams.h" |
22 | | |
23 | | |
24 | | void xnn_f16_dwconv_minmax_ukernel_4p16c__fma3( |
25 | | size_t channels, |
26 | | size_t output_width, |
27 | | const xnn_float16** input, |
28 | | const xnn_float16* weights, |
29 | | xnn_float16* output, |
30 | | intptr_t input_stride, |
31 | | size_t output_increment, |
32 | | size_t input_offset, |
33 | | size_t input_pixel_stride, |
34 | | const xnn_float16* zero, |
35 | | const struct xnn_f16_minmax_params* restrict params) XNN_OOB_READS |
36 | 0 | { |
37 | 0 | assert(channels != 0); |
38 | 0 | assert(output_width != 0); |
39 | | |
40 | 0 | const __m256 vmin = _mm256_cvtph_ps(_mm_set1_epi16(*(const uint16_t*) ¶ms->scalar.min)); |
41 | 0 | const __m256 vmax = _mm256_cvtph_ps(_mm_set1_epi16(*(const uint16_t*) ¶ms->scalar.max)); |
42 | 0 | XNN_FORCE_REALIZATION(vmin); |
43 | 0 | XNN_FORCE_REALIZATION(vmax); |
44 | |
|
45 | 0 | uint16_t* o = (uint16_t*) output; |
46 | 0 | do { |
47 | 0 | const uint16_t* i0 = (const uint16_t*) input[0]; |
48 | 0 | assert(i0 != NULL); |
49 | 0 | if XNN_UNPREDICTABLE(i0 != (const uint16_t*) zero) { |
50 | 0 | i0 = (const uint16_t*) ((uintptr_t) i0 + input_offset); |
51 | 0 | } |
52 | 0 | const uint16_t* i1 = (const uint16_t*) input[1]; |
53 | 0 | assert(i1 != NULL); |
54 | 0 | if XNN_UNPREDICTABLE(i1 != (const uint16_t*) zero) { |
55 | 0 | i1 = (const uint16_t*) ((uintptr_t) i1 + input_offset); |
56 | 0 | } |
57 | 0 | const uint16_t* i2 = (const uint16_t*) input[2]; |
58 | 0 | assert(i2 != NULL); |
59 | 0 | if XNN_UNPREDICTABLE(i2 != (const uint16_t*) zero) { |
60 | 0 | i2 = (const uint16_t*) ((uintptr_t) i2 + input_offset); |
61 | 0 | } |
62 | 0 | const uint16_t* i3 = (const uint16_t*) input[3]; |
63 | 0 | assert(i3 != NULL); |
64 | 0 | if XNN_UNPREDICTABLE(i3 != (const uint16_t*) zero) { |
65 | 0 | i3 = (const uint16_t*) ((uintptr_t) i3 + input_offset); |
66 | 0 | } |
67 | 0 | input = (const xnn_float16**) ((uintptr_t) input + input_stride); |
68 | |
|
69 | 0 | size_t c = channels; |
70 | 0 | const uint16_t* w = (const uint16_t*)weights; |
71 | 0 | for (; c >= 16; c -= 16) { |
72 | 0 | __m256 vacc01234567p0 = _mm256_cvtph_ps(_mm_load_si128((const __m128i*) w)); |
73 | 0 | __m256 vacc89ABCDEFp0 = _mm256_cvtph_ps(_mm_load_si128((const __m128i*) (w + 8))); |
74 | | |
75 | |
|
76 | 0 | const __m256 vi0x01234567 = _mm256_cvtph_ps(_mm_loadu_si128((const __m128i*) i0)); |
77 | 0 | const __m256 vi0x89ABCDEF = _mm256_cvtph_ps(_mm_loadu_si128((const __m128i*) (i0 + 8))); |
78 | 0 | i0 += 16; |
79 | |
|
80 | 0 | const __m256 vk0x01234567 = _mm256_cvtph_ps(_mm_loadu_si128((const __m128i*) (w + 16))); |
81 | 0 | const __m256 vk0x89ABCDEF = _mm256_cvtph_ps(_mm_loadu_si128((const __m128i*) (w + 24))); |
82 | 0 | vacc01234567p0 = _mm256_cvtph_ps(_mm256_cvtps_ph(_mm256_fmadd_ps(vi0x01234567, vk0x01234567, vacc01234567p0), _MM_FROUND_TO_NEAREST_INT)); |
83 | 0 | vacc89ABCDEFp0 = _mm256_cvtph_ps(_mm256_cvtps_ph(_mm256_fmadd_ps(vi0x89ABCDEF, vk0x89ABCDEF, vacc89ABCDEFp0), _MM_FROUND_TO_NEAREST_INT)); |
84 | |
|
85 | 0 | const __m256 vi1x01234567 = _mm256_cvtph_ps(_mm_loadu_si128((const __m128i*) i1)); |
86 | 0 | const __m256 vi1x89ABCDEF = _mm256_cvtph_ps(_mm_loadu_si128((const __m128i*) (i1 + 8))); |
87 | 0 | i1 += 16; |
88 | |
|
89 | 0 | const __m256 vk1x01234567 = _mm256_cvtph_ps(_mm_loadu_si128((const __m128i*) (w + 32))); |
90 | 0 | const __m256 vk1x89ABCDEF = _mm256_cvtph_ps(_mm_loadu_si128((const __m128i*) (w + 40))); |
91 | 0 | vacc01234567p0 = _mm256_cvtph_ps(_mm256_cvtps_ph(_mm256_fmadd_ps(vi1x01234567, vk1x01234567, vacc01234567p0), _MM_FROUND_TO_NEAREST_INT)); |
92 | 0 | vacc89ABCDEFp0 = _mm256_cvtph_ps(_mm256_cvtps_ph(_mm256_fmadd_ps(vi1x89ABCDEF, vk1x89ABCDEF, vacc89ABCDEFp0), _MM_FROUND_TO_NEAREST_INT)); |
93 | |
|
94 | 0 | const __m256 vi2x01234567 = _mm256_cvtph_ps(_mm_loadu_si128((const __m128i*) i2)); |
95 | 0 | const __m256 vi2x89ABCDEF = _mm256_cvtph_ps(_mm_loadu_si128((const __m128i*) (i2 + 8))); |
96 | 0 | i2 += 16; |
97 | |
|
98 | 0 | const __m256 vk2x01234567 = _mm256_cvtph_ps(_mm_loadu_si128((const __m128i*) (w + 48))); |
99 | 0 | const __m256 vk2x89ABCDEF = _mm256_cvtph_ps(_mm_loadu_si128((const __m128i*) (w + 56))); |
100 | 0 | vacc01234567p0 = _mm256_cvtph_ps(_mm256_cvtps_ph(_mm256_fmadd_ps(vi2x01234567, vk2x01234567, vacc01234567p0), _MM_FROUND_TO_NEAREST_INT)); |
101 | 0 | vacc89ABCDEFp0 = _mm256_cvtph_ps(_mm256_cvtps_ph(_mm256_fmadd_ps(vi2x89ABCDEF, vk2x89ABCDEF, vacc89ABCDEFp0), _MM_FROUND_TO_NEAREST_INT)); |
102 | |
|
103 | 0 | const __m256 vi3x01234567 = _mm256_cvtph_ps(_mm_loadu_si128((const __m128i*) i3)); |
104 | 0 | const __m256 vi3x89ABCDEF = _mm256_cvtph_ps(_mm_loadu_si128((const __m128i*) (i3 + 8))); |
105 | 0 | i3 += 16; |
106 | |
|
107 | 0 | const __m256 vk3x01234567 = _mm256_cvtph_ps(_mm_loadu_si128((const __m128i*) (w + 64))); |
108 | 0 | const __m256 vk3x89ABCDEF = _mm256_cvtph_ps(_mm_loadu_si128((const __m128i*) (w + 72))); |
109 | 0 | vacc01234567p0 = _mm256_cvtph_ps(_mm256_cvtps_ph(_mm256_fmadd_ps(vi3x01234567, vk3x01234567, vacc01234567p0), _MM_FROUND_TO_NEAREST_INT)); |
110 | 0 | vacc89ABCDEFp0 = _mm256_cvtph_ps(_mm256_cvtps_ph(_mm256_fmadd_ps(vi3x89ABCDEF, vk3x89ABCDEF, vacc89ABCDEFp0), _MM_FROUND_TO_NEAREST_INT)); |
111 | |
|
112 | 0 | w += 80; |
113 | | |
114 | |
|
115 | 0 | __m256 vacc01234567 = _mm256_max_ps(vacc01234567p0, vmin); |
116 | 0 | __m256 vacc89ABCDEF = _mm256_max_ps(vacc89ABCDEFp0, vmin); |
117 | 0 | vacc01234567 = _mm256_min_ps(vacc01234567, vmax); |
118 | 0 | vacc89ABCDEF = _mm256_min_ps(vacc89ABCDEF, vmax); |
119 | |
|
120 | 0 | _mm_storeu_si128((__m128i*) o, _mm256_cvtps_ph(vacc01234567, _MM_FROUND_TO_NEAREST_INT)); |
121 | 0 | _mm_storeu_si128((__m128i*) (o + 8), _mm256_cvtps_ph(vacc89ABCDEF, _MM_FROUND_TO_NEAREST_INT)); |
122 | 0 | o += 16; |
123 | 0 | } |
124 | 0 | for (; c >= 8; c -= 8) { |
125 | 0 | __m256 vacc01234567p0 = _mm256_cvtph_ps(_mm_load_si128((const __m128i*) w)); |
126 | |
|
127 | 0 | const __m256 vi0x01234567 = _mm256_cvtph_ps(_mm_loadu_si128((const __m128i*) i0)); |
128 | 0 | i0 += 8; |
129 | |
|
130 | 0 | const __m256 vk0x01234567 = _mm256_cvtph_ps(_mm_load_si128((const __m128i*) (w + 16))); |
131 | 0 | vacc01234567p0 = _mm256_cvtph_ps(_mm256_cvtps_ph(_mm256_fmadd_ps(vi0x01234567, vk0x01234567, vacc01234567p0), _MM_FROUND_TO_NEAREST_INT)); |
132 | |
|
133 | 0 | const __m256 vi1x01234567 = _mm256_cvtph_ps(_mm_loadu_si128((const __m128i*) i1)); |
134 | 0 | i1 += 8; |
135 | |
|
136 | 0 | const __m256 vk1x01234567 = _mm256_cvtph_ps(_mm_load_si128((const __m128i*) (w + 32))); |
137 | 0 | vacc01234567p0 = _mm256_cvtph_ps(_mm256_cvtps_ph(_mm256_fmadd_ps(vi1x01234567, vk1x01234567, vacc01234567p0), _MM_FROUND_TO_NEAREST_INT)); |
138 | |
|
139 | 0 | const __m256 vi2x01234567 = _mm256_cvtph_ps(_mm_loadu_si128((const __m128i*) i2)); |
140 | 0 | i2 += 8; |
141 | |
|
142 | 0 | const __m256 vk2x01234567 = _mm256_cvtph_ps(_mm_load_si128((const __m128i*) (w + 48))); |
143 | 0 | vacc01234567p0 = _mm256_cvtph_ps(_mm256_cvtps_ph(_mm256_fmadd_ps(vi2x01234567, vk2x01234567, vacc01234567p0), _MM_FROUND_TO_NEAREST_INT)); |
144 | |
|
145 | 0 | const __m256 vi3x01234567 = _mm256_cvtph_ps(_mm_loadu_si128((const __m128i*) i3)); |
146 | 0 | i3 += 8; |
147 | |
|
148 | 0 | const __m256 vk3x01234567 = _mm256_cvtph_ps(_mm_load_si128((const __m128i*) (w + 64))); |
149 | 0 | vacc01234567p0 = _mm256_cvtph_ps(_mm256_cvtps_ph(_mm256_fmadd_ps(vi3x01234567, vk3x01234567, vacc01234567p0), _MM_FROUND_TO_NEAREST_INT)); |
150 | |
|
151 | 0 | w += 8; |
152 | | |
153 | |
|
154 | 0 | __m256 vacc01234567 = _mm256_max_ps(vacc01234567p0, vmin); |
155 | 0 | vacc01234567 = _mm256_min_ps(vacc01234567, vmax); |
156 | |
|
157 | 0 | _mm_storeu_si128((__m128i*) o, _mm256_cvtps_ph(vacc01234567, _MM_FROUND_TO_NEAREST_INT)); |
158 | 0 | o += 8; |
159 | 0 | } |
160 | 0 | if XNN_UNLIKELY(c != 0) { |
161 | 0 | assert(c >= 1); |
162 | 0 | assert(c <= 7); |
163 | | |
164 | 0 | __m256 vacc01234567p0 = _mm256_cvtph_ps(_mm_load_si128((const __m128i*) w)); |
165 | |
|
166 | 0 | const __m256 vi0x01234567 = _mm256_cvtph_ps(_mm_loadu_si128((const __m128i*) i0)); |
167 | |
|
168 | 0 | const __m256 vk0x01234567 = _mm256_cvtph_ps(_mm_load_si128((const __m128i*) (w + 16))); |
169 | 0 | vacc01234567p0 = _mm256_cvtph_ps(_mm256_cvtps_ph(_mm256_fmadd_ps(vi0x01234567, vk0x01234567, vacc01234567p0), _MM_FROUND_TO_NEAREST_INT)); |
170 | |
|
171 | 0 | const __m256 vi1x01234567 = _mm256_cvtph_ps(_mm_loadu_si128((const __m128i*) i1)); |
172 | |
|
173 | 0 | const __m256 vk1x01234567 = _mm256_cvtph_ps(_mm_load_si128((const __m128i*) (w + 32))); |
174 | 0 | vacc01234567p0 = _mm256_cvtph_ps(_mm256_cvtps_ph(_mm256_fmadd_ps(vi1x01234567, vk1x01234567, vacc01234567p0), _MM_FROUND_TO_NEAREST_INT)); |
175 | |
|
176 | 0 | const __m256 vi2x01234567 = _mm256_cvtph_ps(_mm_loadu_si128((const __m128i*) i2)); |
177 | |
|
178 | 0 | const __m256 vk2x01234567 = _mm256_cvtph_ps(_mm_load_si128((const __m128i*) (w + 48))); |
179 | 0 | vacc01234567p0 = _mm256_cvtph_ps(_mm256_cvtps_ph(_mm256_fmadd_ps(vi2x01234567, vk2x01234567, vacc01234567p0), _MM_FROUND_TO_NEAREST_INT)); |
180 | |
|
181 | 0 | const __m256 vi3x01234567 = _mm256_cvtph_ps(_mm_loadu_si128((const __m128i*) i3)); |
182 | |
|
183 | 0 | const __m256 vk3x01234567 = _mm256_cvtph_ps(_mm_load_si128((const __m128i*) (w + 64))); |
184 | 0 | vacc01234567p0 = _mm256_cvtph_ps(_mm256_cvtps_ph(_mm256_fmadd_ps(vi3x01234567, vk3x01234567, vacc01234567p0), _MM_FROUND_TO_NEAREST_INT)); |
185 | | |
186 | |
|
187 | 0 | __m256 vacc01234567 = _mm256_max_ps(vacc01234567p0, vmin); |
188 | 0 | vacc01234567 = _mm256_min_ps(vacc01234567, vmax); |
189 | |
|
190 | 0 | __m128i vh01234567 = _mm256_cvtps_ph(vacc01234567, _MM_FROUND_TO_NEAREST_INT); |
191 | 0 | if (c & 4) { |
192 | 0 | _mm_storel_epi64((__m128i*) o, vh01234567); |
193 | 0 | vh01234567 = _mm_unpackhi_epi64(vh01234567, vh01234567); |
194 | 0 | o += 4; |
195 | 0 | } |
196 | 0 | if (c & 2) { |
197 | 0 | _mm_storeu_si32(o, vh01234567); |
198 | 0 | vh01234567 = _mm_srli_epi64(vh01234567, 32); |
199 | 0 | o += 2; |
200 | 0 | } |
201 | 0 | if (c & 1) { |
202 | 0 | *o = (uint16_t) _mm_extract_epi16(vh01234567, 0); |
203 | 0 | o += 1; |
204 | 0 | } |
205 | 0 | } |
206 | | |
207 | 0 | input_offset += input_pixel_stride; |
208 | 0 | o = (uint16_t*) ((uintptr_t) o + output_increment); |
209 | 0 | } while (--output_width != 0); |
210 | 0 | } |