Applies causal mask (j > i => 0) and softmax to scores matrix. In-place on [num_heads, T, T] scores matrix.
148{
149 for (int h = 0; h < num_heads; ++h) {
150 for (int i = 0; i < num_tokens; ++i) {
151 int base = h * aligned_context_window * aligned_context_window
152 + i * aligned_context_window;
153 float *row = &scores[base];
154 int len = i + 1;
155
156#if defined(__AVX512F__)
157
158 __m512 max_vec = _mm512_set1_ps(-INFINITY);
159 int j = 0;
160 for (; j + 16 <= len; j += 16) {
161 __m512 v = _mm512_loadu_ps(&row[j]);
162 max_vec = _mm512_max_ps(max_vec, v);
163 }
164 float max_val = _mm512_reduce_max_ps(max_vec);
165 for (; j < len; ++j) {
166 if (row[j] > max_val) max_val = row[j];
167 }
168
169
170 __m512 max_broadcast = _mm512_set1_ps(max_val);
171 __m512 sum_vec = _mm512_setzero_ps();
172 j = 0;
173 for (; j + 16 <= len; j += 16) {
174 __m512 v = _mm512_loadu_ps(&row[j]);
175 __m512 e = exp512_approx(_mm512_sub_ps(v, max_broadcast));
176 _mm512_storeu_ps(&row[j], e);
177 sum_vec = _mm512_add_ps(sum_vec, e);
178 }
179 float sum = _mm512_reduce_add_ps(sum_vec);
180 for (; j < len; ++j) {
181 float e = expf(row[j] - max_val);
182 row[j] = e;
183 sum += e;
184 }
185
186
187 float inv_sum = 1.0f / sum;
188 __m512 inv_sum_vec = _mm512_set1_ps(inv_sum);
189 j = 0;
190 for (; j + 16 <= len; j += 16) {
191 __m512 v = _mm512_loadu_ps(&row[j]);
192 _mm512_storeu_ps(&row[j], _mm512_mul_ps(v, inv_sum_vec));
193 }
194 for (; j < len; ++j) {
195 row[j] *= inv_sum;
196 }
197
198
199 __m512 zero = _mm512_setzero_ps();
200 for (; j + 16 <= num_tokens; j += 16) {
201 _mm512_storeu_ps(&row[j], zero);
202 }
203 for (; j < num_tokens; ++j) {
204 row[j] = 0.0f;
205 }
206
207#elif defined(__AVX2__)
208
209 __m256 max_vec = _mm256_set1_ps(-INFINITY);
210 int j = 0;
211 for (; j + 8 <= len; j += 8) {
212 __m256 v = _mm256_loadu_ps(&row[j]);
213 max_vec = _mm256_max_ps(max_vec, v);
214 }
215 float max_val = hmax256_ps(max_vec);
216 for (; j < len; ++j) {
217 if (row[j] > max_val) max_val = row[j];
218 }
219
220
221 __m256 max_broadcast = _mm256_set1_ps(max_val);
222 __m256 sum_vec = _mm256_setzero_ps();
223 j = 0;
224 for (; j + 8 <= len; j += 8) {
225 __m256 v = _mm256_loadu_ps(&row[j]);
226 __m256 e = exp256_approx(_mm256_sub_ps(v, max_broadcast));
227 _mm256_storeu_ps(&row[j], e);
228 sum_vec = _mm256_add_ps(sum_vec, e);
229 }
230 float sum = hsum256_ps_softmax(sum_vec);
231 for (; j < len; ++j) {
232 float e = expf(row[j] - max_val);
233 row[j] = e;
234 sum += e;
235 }
236
237
238 float inv_sum = 1.0f / sum;
239 __m256 inv_sum_vec = _mm256_set1_ps(inv_sum);
240 j = 0;
241 for (; j + 8 <= len; j += 8) {
242 __m256 v = _mm256_loadu_ps(&row[j]);
243 _mm256_storeu_ps(&row[j], _mm256_mul_ps(v, inv_sum_vec));
244 }
245 for (; j < len; ++j) {
246 row[j] *= inv_sum;
247 }
248
249
250 __m256 zero = _mm256_setzero_ps();
251 for (; j + 8 <= num_tokens; j += 8) {
252 _mm256_storeu_ps(&row[j], zero);
253 }
254 for (; j < num_tokens; ++j) {
255 row[j] = 0.0f;
256 }
257
258#elif defined(__AVX__)
259
260 __m256 max_vec = _mm256_set1_ps(-INFINITY);
261 int j = 0;
262 for (; j + 8 <= len; j += 8) {
263 __m256 v = _mm256_loadu_ps(&row[j]);
264 max_vec = _mm256_max_ps(max_vec, v);
265 }
266 float max_val = hmax256_ps(max_vec);
267 for (; j < len; ++j) {
268 if (row[j] > max_val) max_val = row[j];
269 }
270
271
272 float sum = 0.0f;
273 for (j = 0; j < len; ++j) {
274 float e = expf(row[j] - max_val);
275 row[j] = e;
276 sum += e;
277 }
278
279
280 float inv_sum = 1.0f / sum;
281 __m256 inv_sum_vec = _mm256_set1_ps(inv_sum);
282 j = 0;
283 for (; j + 8 <= len; j += 8) {
284 __m256 v = _mm256_loadu_ps(&row[j]);
285 _mm256_storeu_ps(&row[j], _mm256_mul_ps(v, inv_sum_vec));
286 }
287 for (; j < len; ++j) {
288 row[j] *= inv_sum;
289 }
290
291
292 __m256 zero = _mm256_setzero_ps();
293 for (; j + 8 <= num_tokens; j += 8) {
294 _mm256_storeu_ps(&row[j], zero);
295 }
296 for (; j < num_tokens; ++j) {
297 row[j] = 0.0f;
298 }
299
300#else
301
302 float max_val = row[0];
303 for (int j = 1; j < len; ++j) {
304 if (row[j] > max_val) max_val = row[j];
305 }
306
307 float sum = 0.0f;
308 for (int j = 0; j < len; ++j) {
309 float e = expf(row[j] - max_val);
310 row[j] = e;
311 sum += e;
312 }
313
314 float inv_sum = 1.0f / sum;
315 for (int j = 0; j < len; ++j) {
316 row[j] *= inv_sum;
317 }
318
319 for (int j = len; j < num_tokens; ++j) {
320 row[j] = 0.0f;
321 }
322#endif
323 }
324 }
325}