← Back to C-Kernel-Engine Docs Doxygen Source Documentation
 
Loading...
Searching...
No Matches
ck_parallel_train.c
Go to the documentation of this file.
1/**
2 * @file ck_parallel_train.c
3 * @brief Thread-pool dispatch wrapper for FP32 training GEMM.
4 *
5 * The training runtime currently emits many gemm_blocked_serial calls.
6 * This wrapper keeps kernel math unchanged and only parallelizes dispatch.
7 *
8 * OpenMP removal note:
9 * Training backward used to fall back into legacy OpenMP GEMM paths for T>1
10 * micro-batches, which mixed libiomp barriers with the CK threadpool and
11 * showed up as fork/barrier overhead in profiling. Keep the hot training
12 * path on the CK threadpool here, and leave older kernel-level OpenMP paths
13 * as compatibility fallbacks until they are fully retired.
14 */
15
16#include "ckernel_engine.h"
17#include "ck_threadpool.h"
18
19#include <stddef.h>
20#if defined(__AVX2__) || defined(__AVX__) || defined(__AVX512F__)
21#include <immintrin.h>
22#endif
23
24typedef struct {
25 const float *A;
26 const float *B;
27 const float *bias;
28 float *C;
29 int M;
30 int N;
31 int K;
32 int split_n;
33} ck_train_gemm_args_t;
34
35#if defined(__AVX__) && !defined(__AVX512F__)
36static inline float ck_train_hsum256_ps(__m256 v) {
37 __m128 lo = _mm256_castps256_ps128(v);
38 __m128 hi = _mm256_extractf128_ps(v, 1);
39 __m128 sum128 = _mm_add_ps(lo, hi);
40 __m128 shuf = _mm_movehdup_ps(sum128);
41 __m128 sums = _mm_add_ps(sum128, shuf);
42 shuf = _mm_movehl_ps(shuf, sums);
43 sums = _mm_add_ss(sums, shuf);
44 return _mm_cvtss_f32(sums);
45}
46#endif
47
48static void ck_train_gemm_nt_compute_rows(const float *A,
49 const float *B,
50 const float *bias,
51 float *C,
52 int row_start,
53 int row_end,
54 int N,
55 int K) {
56 if (!A || !B || !C || row_start >= row_end || N <= 0 || K <= 0) {
57 return;
58 }
59
60 for (int i = row_start; i < row_end; ++i) {
61 const float *a_row = A + (size_t)i * (size_t)K;
62 float *c_row = C + (size_t)i * (size_t)N;
63 for (int j = 0; j < N; ++j) {
64 const float *b_row = B + (size_t)j * (size_t)K;
65 float sum = bias ? bias[j] : 0.0f;
66#if defined(__AVX512F__)
67 __m512 acc = _mm512_setzero_ps();
68 int k = 0;
69 for (; k <= K - 16; k += 16) {
70 __m512 a_vec = _mm512_loadu_ps(a_row + k);
71 __m512 b_vec = _mm512_loadu_ps(b_row + k);
72 acc = _mm512_fmadd_ps(a_vec, b_vec, acc);
73 }
74 sum += _mm512_reduce_add_ps(acc);
75 for (; k < K; ++k) {
76 sum += a_row[k] * b_row[k];
77 }
78#elif defined(__AVX2__)
79 __m256 acc = _mm256_setzero_ps();
80 int k = 0;
81 for (; k <= K - 8; k += 8) {
82 __m256 a_vec = _mm256_loadu_ps(a_row + k);
83 __m256 b_vec = _mm256_loadu_ps(b_row + k);
84#if defined(__FMA__)
85 acc = _mm256_fmadd_ps(a_vec, b_vec, acc);
86#else
87 acc = _mm256_add_ps(acc, _mm256_mul_ps(a_vec, b_vec));
88#endif
89 }
90 sum += ck_train_hsum256_ps(acc);
91 for (; k < K; ++k) {
92 sum += a_row[k] * b_row[k];
93 }
94#else
95 for (int k = 0; k < K; ++k) {
96 sum += a_row[k] * b_row[k];
97 }
98#endif
99 c_row[j] = sum;
100 }
101 }
102}
103
104static int ck_train_pick_active_threads(int nth, size_t work_items, size_t min_chunk)
105{
106 if (nth <= 1 || work_items == 0 || min_chunk == 0) {
107 return 1;
108 }
109 size_t active = (work_items + min_chunk - 1u) / min_chunk;
110 if (active < 1u) {
111 active = 1u;
112 }
113 if (active > (size_t)nth) {
114 active = (size_t)nth;
115 }
116 return (int)active;
117}
118
119static void ck_train_gemm_work(int ith, int nth, void *argp) {
120 ck_train_gemm_args_t *a = (ck_train_gemm_args_t *)argp;
121 if (!a || a->M <= 0 || a->N <= 0 || a->K <= 0) {
122 return;
123 }
124
125 if (a->split_n) {
126 /* Decode/microbatch path (M=1): split output columns. */
127 int dn = (a->N + nth - 1) / nth;
128 int n0 = dn * ith;
129 int n1 = n0 + dn;
130 if (n0 >= a->N) {
131 return;
132 }
133 if (n1 > a->N) {
134 n1 = a->N;
135 }
136 const int n_chunk = n1 - n0;
137 if (n_chunk <= 0) {
138 return;
139 }
140
141 const float *B_chunk = a->B + (size_t)n0 * (size_t)a->K;
142 const float *bias_chunk = a->bias ? (a->bias + n0) : NULL;
143 float *C_chunk = a->C + n0; /* M=1 layout */
144
145 gemm_blocked_serial(a->A, B_chunk, bias_chunk, C_chunk, 1, n_chunk, a->K);
146 return;
147 }
148
149 /* Prefill/train path (M>1): split rows. */
150 int dm = (a->M + nth - 1) / nth;
151 int m0 = dm * ith;
152 int m1 = m0 + dm;
153 if (m0 >= a->M) {
154 return;
155 }
156 if (m1 > a->M) {
157 m1 = a->M;
158 }
159 const int m_chunk = m1 - m0;
160 if (m_chunk <= 0) {
161 return;
162 }
163
164 ck_train_gemm_nt_compute_rows(a->A, a->B, a->bias, a->C, m0, m1, a->N, a->K);
165}
166
168 const float *B,
169 const float *bias,
170 float *C,
171 int M,
172 int N,
173 int K) {
174 if (!A || !B || !C || M <= 0 || N <= 0 || K <= 0) {
175 return;
176 }
177
178 ck_threadpool_t *pool = ck_threadpool_global();
179 const int nth = pool ? ck_threadpool_n_threads(pool) : 1;
180
181 /* Keep small shapes on serial path to avoid dispatch overhead. */
182 const size_t work = (size_t)M * (size_t)N * (size_t)K;
183 if (!pool || nth <= 1 || work < (size_t)131072) {
184 gemm_blocked_serial(A, B, bias, C, M, N, K);
185 return;
186 }
187
188 int split_n = 0;
189 int active_nth = nth;
190 if (M == 1) {
191 /* Tensor-parallel decode-style split only when each worker has enough columns. */
192 const int cols_per_worker = (N + nth - 1) / nth;
193 if (cols_per_worker >= 256 && K >= 512 && work >= (size_t)2097152) {
194 split_n = 1;
195 } else {
196 gemm_blocked_serial(A, B, bias, C, M, N, K);
197 return;
198 }
199 } else {
200 /* Row split for prefill/train batches with coarser chunks. */
201 active_nth = M / 2;
202 if (active_nth > nth) {
203 active_nth = nth;
204 }
205 if (active_nth <= 1) {
206 gemm_blocked_serial(A, B, bias, C, M, N, K);
207 return;
208 }
209 }
210
211 ck_train_gemm_args_t args = {
212 .A = A,
213 .B = B,
214 .bias = bias,
215 .C = C,
216 .M = M,
217 .N = N,
218 .K = K,
219 .split_n = split_n,
220 };
221
222 ck_threadpool_dispatch_n(pool, active_nth, ck_train_gemm_work, &args);
223}
224
225typedef struct {
226 const float *A;
227 const float *B;
228 const float *bias;
229 float *C;
230 int M;
231 int N;
232 int K;
233 int split_n;
234} ck_train_gemm_nn_args_t;
235
236typedef struct {
237 const float *A;
238 const float *B;
239 const float *bias;
240 float *C;
241 int M;
242 int N;
243 int K;
244} ck_train_gemm_tn_args_t;
245
246typedef struct {
247 const float *d_output;
248 float *d_bias;
249 int T;
250 int aligned_out;
251} ck_train_bias_reduce_args_t;
252
253typedef struct {
254 const float *d_output;
255 const float *input;
256 float *d_W;
257 float *d_b;
258 int aligned_in;
259 int aligned_out;
260} ck_train_outer_t1_args_t;
261
262typedef struct {
263 const float *d_output;
264 const float *input;
265 const float *W;
266 float *d_input;
267 float *d_W;
268 float *d_b;
269 int T;
270 int aligned_in;
271 int aligned_out;
272} ck_train_gemm_backward_args_t;
273
274static void ck_train_gemm_nn_compute_rows(const float *A,
275 const float *B,
276 const float *bias,
277 float *C,
278 int row_start,
279 int row_end,
280 int N,
281 int K) {
282 if (!A || !B || !C || row_start >= row_end || N <= 0 || K <= 0) {
283 return;
284 }
285
286 for (int i = row_start; i < row_end; ++i) {
287 const float *a_row = A + (size_t)i * (size_t)K;
288 float *c_row = C + (size_t)i * (size_t)N;
289#if defined(__AVX512F__)
290 int j = 0;
291 for (; j <= N - 16; j += 16) {
292 __m512 sum = bias ? _mm512_loadu_ps(bias + j) : _mm512_setzero_ps();
293 for (int k = 0; k < K; ++k) {
294 __m512 av = _mm512_set1_ps(a_row[k]);
295 __m512 bv = _mm512_loadu_ps(B + (size_t)k * (size_t)N + (size_t)j);
296 sum = _mm512_fmadd_ps(av, bv, sum);
297 }
298 _mm512_storeu_ps(c_row + j, sum);
299 }
300 for (; j < N; ++j) {
301 float sum = bias ? bias[j] : 0.0f;
302 for (int k = 0; k < K; ++k) {
303 sum += a_row[k] * B[(size_t)k * (size_t)N + (size_t)j];
304 }
305 c_row[j] = sum;
306 }
307#elif defined(__AVX2__)
308 int j = 0;
309 for (; j <= N - 8; j += 8) {
310 __m256 sum = bias ? _mm256_loadu_ps(bias + j) : _mm256_setzero_ps();
311 for (int k = 0; k < K; ++k) {
312 __m256 av = _mm256_set1_ps(a_row[k]);
313 __m256 bv = _mm256_loadu_ps(B + (size_t)k * (size_t)N + (size_t)j);
314#if defined(__FMA__)
315 sum = _mm256_fmadd_ps(av, bv, sum);
316#else
317 sum = _mm256_add_ps(sum, _mm256_mul_ps(av, bv));
318#endif
319 }
320 _mm256_storeu_ps(c_row + j, sum);
321 }
322 for (; j < N; ++j) {
323 float sum = bias ? bias[j] : 0.0f;
324 for (int k = 0; k < K; ++k) {
325 sum += a_row[k] * B[(size_t)k * (size_t)N + (size_t)j];
326 }
327 c_row[j] = sum;
328 }
329#else
330 for (int j = 0; j < N; ++j) {
331 float sum = bias ? bias[j] : 0.0f;
332 for (int k = 0; k < K; ++k) {
333 sum += a_row[k] * B[(size_t)k * (size_t)N + (size_t)j];
334 }
335 c_row[j] = sum;
336 }
337#endif
338 }
339}
340
341static void ck_train_gemm_tn_compute_rows(const float *A,
342 const float *B,
343 const float *bias,
344 float *C,
345 int row_start,
346 int row_end,
347 int M,
348 int N,
349 int K) {
350 if (!A || !B || !C || row_start >= row_end || M <= 0 || N <= 0 || K <= 0) {
351 return;
352 }
353
354 for (int i = row_start; i < row_end; ++i) {
355 float *c_row = C + (size_t)i * (size_t)N;
356#if defined(__AVX512F__)
357 int j = 0;
358 for (; j <= N - 16; j += 16) {
359 __m512 sum = bias ? _mm512_loadu_ps(bias + j) : _mm512_setzero_ps();
360 for (int k = 0; k < K; ++k) {
361 __m512 av = _mm512_set1_ps(A[(size_t)k * (size_t)M + (size_t)i]);
362 __m512 bv = _mm512_loadu_ps(B + (size_t)k * (size_t)N + (size_t)j);
363 sum = _mm512_fmadd_ps(av, bv, sum);
364 }
365 _mm512_storeu_ps(c_row + j, sum);
366 }
367 for (; j < N; ++j) {
368 float sum = bias ? bias[j] : 0.0f;
369 for (int k = 0; k < K; ++k) {
370 sum += A[(size_t)k * (size_t)M + (size_t)i] *
371 B[(size_t)k * (size_t)N + (size_t)j];
372 }
373 c_row[j] = sum;
374 }
375#elif defined(__AVX2__)
376 int j = 0;
377 for (; j <= N - 8; j += 8) {
378 __m256 sum = bias ? _mm256_loadu_ps(bias + j) : _mm256_setzero_ps();
379 for (int k = 0; k < K; ++k) {
380 __m256 av = _mm256_set1_ps(A[(size_t)k * (size_t)M + (size_t)i]);
381 __m256 bv = _mm256_loadu_ps(B + (size_t)k * (size_t)N + (size_t)j);
382#if defined(__FMA__)
383 sum = _mm256_fmadd_ps(av, bv, sum);
384#else
385 sum = _mm256_add_ps(sum, _mm256_mul_ps(av, bv));
386#endif
387 }
388 _mm256_storeu_ps(c_row + j, sum);
389 }
390 for (; j < N; ++j) {
391 float sum = bias ? bias[j] : 0.0f;
392 for (int k = 0; k < K; ++k) {
393 sum += A[(size_t)k * (size_t)M + (size_t)i] *
394 B[(size_t)k * (size_t)N + (size_t)j];
395 }
396 c_row[j] = sum;
397 }
398#else
399 for (int j = 0; j < N; ++j) {
400 float sum = bias ? bias[j] : 0.0f;
401 for (int k = 0; k < K; ++k) {
402 sum += A[(size_t)k * (size_t)M + (size_t)i] *
403 B[(size_t)k * (size_t)N + (size_t)j];
404 }
405 c_row[j] = sum;
406 }
407#endif
408 }
409}
410
411static void ck_train_bias_reduce_compute_range(const float *d_output,
412 float *d_bias,
413 int T,
414 int out_start,
415 int out_end,
416 int aligned_out) {
417 if (!d_output || !d_bias || T <= 0 || out_start >= out_end || aligned_out <= 0) {
418 return;
419 }
420
421 for (int out_idx = out_start; out_idx < out_end; ++out_idx) {
422 float bias_grad = 0.0f;
423 for (int t = 0; t < T; ++t) {
424 bias_grad += d_output[(size_t)t * (size_t)aligned_out + (size_t)out_idx];
425 }
426 d_bias[out_idx] += bias_grad;
427 }
428}
429
430static void ck_train_outer_t1_compute_range(const float *d_output,
431 const float *input,
432 float *d_W,
433 float *d_b,
434 int out_start,
435 int out_end,
436 int aligned_in) {
437 if (!d_output || !input || !d_W || out_start >= out_end || aligned_in <= 0) {
438 return;
439 }
440
441 for (int out_idx = out_start; out_idx < out_end; ++out_idx) {
442 float g = d_output[out_idx];
443 if (d_b) {
444 d_b[out_idx] += g;
445 }
446 float *dw_row = d_W + (size_t)out_idx * (size_t)aligned_in;
447 for (int j = 0; j < aligned_in; ++j) {
448 dw_row[j] = g * input[j];
449 }
450 }
451}
452
453static void ck_train_gemm_backward_serial(const float *d_output,
454 const float *input,
455 const float *W,
456 float *d_input,
457 float *d_W,
458 float *d_b,
459 int T,
460 int aligned_in,
461 int aligned_out) {
462 ck_train_gemm_nn_compute_rows(d_output, W, NULL, d_input, 0, T, aligned_in, aligned_out);
463 ck_train_gemm_tn_compute_rows(d_output, input, NULL, d_W, 0, aligned_out, aligned_out, aligned_in, T);
464 if (d_b) {
465 ck_train_bias_reduce_compute_range(d_output, d_b, T, 0, aligned_out, aligned_out);
466 }
467}
468
469static void ck_train_gemm_nn_work(int ith, int nth, void *argp) {
470 ck_train_gemm_nn_args_t *a = (ck_train_gemm_nn_args_t *)argp;
471 if (!a || a->M <= 0 || a->N <= 0 || a->K <= 0) {
472 return;
473 }
474
475 if (a->split_n) {
476 int dn = (a->N + nth - 1) / nth;
477 int n0 = dn * ith;
478 int n1 = n0 + dn;
479 if (n0 >= a->N) {
480 return;
481 }
482 if (n1 > a->N) {
483 n1 = a->N;
484 }
485#if defined(__AVX2__)
486 int j = n0;
487 for (; j <= n1 - 8; j += 8) {
488 __m256 sum = a->bias ? _mm256_loadu_ps(a->bias + j) : _mm256_setzero_ps();
489 for (int k = 0; k < a->K; ++k) {
490 __m256 av = _mm256_set1_ps(a->A[k]);
491 __m256 bv = _mm256_loadu_ps(a->B + (size_t)k * (size_t)a->N + (size_t)j);
492#if defined(__FMA__)
493 sum = _mm256_fmadd_ps(av, bv, sum);
494#else
495 sum = _mm256_add_ps(sum, _mm256_mul_ps(av, bv));
496#endif
497 }
498 _mm256_storeu_ps(a->C + j, sum);
499 }
500 for (; j < n1; ++j) {
501 float sum = a->bias ? a->bias[j] : 0.0f;
502 for (int k = 0; k < a->K; ++k) {
503 sum += a->A[k] * a->B[(size_t)k * (size_t)a->N + (size_t)j];
504 }
505 a->C[j] = sum;
506 }
507#else
508 for (int j = n0; j < n1; ++j) {
509 float sum = a->bias ? a->bias[j] : 0.0f;
510 for (int k = 0; k < a->K; ++k) {
511 sum += a->A[k] * a->B[(size_t)k * (size_t)a->N + (size_t)j];
512 }
513 a->C[j] = sum;
514 }
515#endif
516 return;
517 }
518
519 int dm = (a->M + nth - 1) / nth;
520 int m0 = dm * ith;
521 int m1 = m0 + dm;
522 if (m0 >= a->M) {
523 return;
524 }
525 if (m1 > a->M) {
526 m1 = a->M;
527 }
528
529 const int m_chunk = m1 - m0;
530 if (m_chunk <= 0) {
531 return;
532 }
533
534 ck_train_gemm_nn_compute_rows(a->A, a->B, a->bias, a->C, m0, m1, a->N, a->K);
535}
536
537static void ck_train_outer_t1_work(int ith, int nth, void *argp) {
538 ck_train_outer_t1_args_t *a = (ck_train_outer_t1_args_t *)argp;
539 if (!a || a->aligned_in <= 0 || a->aligned_out <= 0) {
540 return;
541 }
542
543 int dn = (a->aligned_out + nth - 1) / nth;
544 int n0 = dn * ith;
545 int n1 = n0 + dn;
546 if (n0 >= a->aligned_out) {
547 return;
548 }
549 if (n1 > a->aligned_out) {
550 n1 = a->aligned_out;
551 }
552
553 ck_train_outer_t1_compute_range(a->d_output, a->input, a->d_W, a->d_b, n0, n1, a->aligned_in);
554}
555
556static void ck_train_gemm_backward_work(int ith, int nth, void *argp) {
557 ck_train_gemm_backward_args_t *a = (ck_train_gemm_backward_args_t *)argp;
558 if (!a || !a->d_output || !a->input || !a->W || !a->d_input || !a->d_W ||
559 a->T <= 0 || a->aligned_in <= 0 || a->aligned_out <= 0) {
560 return;
561 }
562
563 int dt = (a->T + nth - 1) / nth;
564 int t0 = dt * ith;
565 int t1 = t0 + dt;
566 if (t1 > a->T) {
567 t1 = a->T;
568 }
569 if (t0 < t1) {
571 a->d_output,
572 a->W,
573 NULL,
574 a->d_input,
575 t0,
576 t1,
577 a->aligned_in,
578 a->aligned_out);
579 }
580
581 int dn = (a->aligned_out + nth - 1) / nth;
582 int n0 = dn * ith;
583 int n1 = n0 + dn;
584 if (n1 > a->aligned_out) {
585 n1 = a->aligned_out;
586 }
587 if (n0 < n1) {
589 a->d_output,
590 a->input,
591 NULL,
592 a->d_W,
593 n0,
594 n1,
595 a->aligned_out,
596 a->aligned_in,
597 a->T);
598 if (a->d_b) {
600 a->d_output,
601 a->d_b,
602 a->T,
603 n0,
604 n1,
605 a->aligned_out);
606 }
607 }
608}
609
611 const float *input,
612 const float *W,
613 float *d_input,
614 float *d_W,
615 float *d_b,
616 int T,
617 int aligned_in,
618 int aligned_out,
619 int num_threads) {
620 if (!d_output || !input || !W || !d_input || !d_W) {
621 return;
622 }
623 if (T <= 0 || aligned_in <= 0 || aligned_out <= 0) {
624 return;
625 }
626
627 ck_threadpool_t *pool = ck_threadpool_global();
628 int nth = pool ? ck_threadpool_n_threads(pool) : 1;
629 if (num_threads > 0 && num_threads < nth) {
630 nth = num_threads;
631 }
632
633 if (!pool || nth <= 1) {
634 ck_train_gemm_backward_serial(d_output, input, W, d_input, d_W, d_b, T, aligned_in, aligned_out);
635 return;
636 }
637
638 /* T=1 is the dominant runtime shape in generated train microsteps.
639 * Use vectorized d_input and parallel outer-product for dW/db. */
640 if (T == 1) {
641 gemm_nn_simd(d_output, W, NULL, d_input, 1, aligned_in, aligned_out);
642
643 const size_t outer_work = (size_t)aligned_out * (size_t)aligned_in;
644 if (aligned_out < nth * 2 || aligned_in < 64 || outer_work < (size_t)524288) {
645 ck_train_outer_t1_compute_range(d_output, input, d_W, d_b, 0, aligned_out, aligned_in);
646 } else {
647 const int active_outer = ck_train_pick_active_threads(nth, (size_t)aligned_out, (size_t)64);
648 ck_train_outer_t1_args_t t1_args = {
649 .d_output = d_output,
650 .input = input,
651 .d_W = d_W,
652 .d_b = d_b,
653 .aligned_in = aligned_in,
654 .aligned_out = aligned_out,
655 };
656 if (active_outer <= 1) {
657 ck_train_outer_t1_compute_range(d_output, input, d_W, d_b, 0, aligned_out, aligned_in);
658 } else {
659 ck_threadpool_dispatch_n(pool, active_outer, ck_train_outer_t1_work, &t1_args);
660 }
661 }
662 return;
663 }
664
665 const size_t nn_work = (size_t)T * (size_t)aligned_in * (size_t)aligned_out;
666 const size_t tn_work = (size_t)aligned_out * (size_t)aligned_in * (size_t)T;
667 if ((T < 2 && aligned_out < nth * 2) || (nn_work < (size_t)131072 && tn_work < (size_t)131072)) {
668 ck_train_gemm_backward_serial(d_output, input, W, d_input, d_W, d_b, T, aligned_in, aligned_out);
669 return;
670 }
671
672 ck_train_gemm_backward_args_t bw_args = {
673 .d_output = d_output,
674 .input = input,
675 .W = W,
676 .d_input = d_input,
677 .d_W = d_W,
678 .d_b = d_b,
679 .T = T,
680 .aligned_in = aligned_in,
681 .aligned_out = aligned_out,
682 };
683 {
684 const size_t bw_rows = (size_t)aligned_out > (size_t)T ? (size_t)aligned_out : (size_t)T;
685 const int active_bw = ck_train_pick_active_threads(nth, bw_rows, (size_t)64);
686 if (active_bw <= 1) {
687 ck_train_gemm_backward_serial(d_output, input, W, d_input, d_W, d_b, T, aligned_in, aligned_out);
688 } else {
689 ck_threadpool_dispatch_n(pool, active_bw, ck_train_gemm_backward_work, &bw_args);
690 }
691 }
692}
693
694/*
695 * v2 training backward wrapper:
696 * - keeps kernel math contract untouched
697 * - uses more aggressive threadpool partitioning for T=1 microstep shapes
698 */
700 const float *input,
701 const float *W,
702 float *d_input,
703 float *d_W,
704 float *d_b,
705 int T,
706 int aligned_in,
707 int aligned_out,
708 int num_threads) {
709 if (!d_output || !input || !W || !d_input || !d_W) {
710 return;
711 }
712 if (T <= 0 || aligned_in <= 0 || aligned_out <= 0) {
713 return;
714 }
715
716 ck_threadpool_t *pool = ck_threadpool_global();
717 int nth = pool ? ck_threadpool_n_threads(pool) : 1;
718 if (num_threads > 0 && num_threads < nth) {
719 nth = num_threads;
720 }
721
722 if (!pool || nth <= 1) {
723 ck_train_gemm_backward_serial(d_output, input, W, d_input, d_W, d_b, T, aligned_in, aligned_out);
724 return;
725 }
726
727 if (T == 1) {
728 const size_t nn_work = (size_t)aligned_in * (size_t)aligned_out;
729 if (aligned_in >= nth * 32 && nn_work >= (size_t)131072) {
730 const int active_nn = ck_train_pick_active_threads(nth, (size_t)aligned_in, (size_t)64);
731 ck_train_gemm_nn_args_t nn_args = {
732 .A = d_output,
733 .B = W,
734 .bias = NULL,
735 .C = d_input,
736 .M = 1,
737 .N = aligned_in,
738 .K = aligned_out,
739 .split_n = 1,
740 };
741 if (active_nn <= 1) {
742 gemm_nn_simd(d_output, W, NULL, d_input, 1, aligned_in, aligned_out);
743 } else {
744 ck_threadpool_dispatch_n(pool, active_nn, ck_train_gemm_nn_work, &nn_args);
745 }
746 } else {
747 gemm_nn_simd(d_output, W, NULL, d_input, 1, aligned_in, aligned_out);
748 }
749
750 const size_t outer_work = (size_t)aligned_out * (size_t)aligned_in;
751 if (aligned_out >= nth && outer_work >= (size_t)131072) {
752 const int active_outer = ck_train_pick_active_threads(nth, (size_t)aligned_out, (size_t)64);
753 ck_train_outer_t1_args_t t1_args = {
754 .d_output = d_output,
755 .input = input,
756 .d_W = d_W,
757 .d_b = d_b,
758 .aligned_in = aligned_in,
759 .aligned_out = aligned_out,
760 };
761 if (active_outer <= 1) {
762 ck_train_outer_t1_compute_range(d_output, input, d_W, d_b, 0, aligned_out, aligned_in);
763 } else {
764 ck_threadpool_dispatch_n(pool, active_outer, ck_train_outer_t1_work, &t1_args);
765 }
766 } else {
767 ck_train_outer_t1_compute_range(d_output, input, d_W, d_b, 0, aligned_out, aligned_in);
768 }
769 return;
770 }
771
772 const size_t nn_work = (size_t)T * (size_t)aligned_in * (size_t)aligned_out;
773 const size_t tn_work = (size_t)aligned_out * (size_t)aligned_in * (size_t)T;
774 if ((T < nth && aligned_out < nth) || (nn_work < (size_t)131072 && tn_work < (size_t)131072)) {
775 ck_train_gemm_backward_serial(d_output, input, W, d_input, d_W, d_b, T, aligned_in, aligned_out);
776 return;
777 }
778
779 ck_train_gemm_backward_args_t bw_args = {
780 .d_output = d_output,
781 .input = input,
782 .W = W,
783 .d_input = d_input,
784 .d_W = d_W,
785 .d_b = d_b,
786 .T = T,
787 .aligned_in = aligned_in,
788 .aligned_out = aligned_out,
789 };
790 {
791 const size_t bw_rows = (size_t)aligned_out > (size_t)T ? (size_t)aligned_out : (size_t)T;
792 const int active_bw = ck_train_pick_active_threads(nth, bw_rows, (size_t)64);
793 if (active_bw <= 1) {
794 ck_train_gemm_backward_serial(d_output, input, W, d_input, d_W, d_b, T, aligned_in, aligned_out);
795 } else {
796 ck_threadpool_dispatch_n(pool, active_bw, ck_train_gemm_backward_work, &bw_args);
797 }
798 }
799}
static void ck_train_bias_reduce_compute_range(const float *d_output, float *d_bias, int T, int out_start, int out_end, int aligned_out)
static void ck_train_gemm_tn_compute_rows(const float *A, const float *B, const float *bias, float *C, int row_start, int row_end, int M, int N, int K)
static void ck_train_outer_t1_compute_range(const float *d_output, const float *input, float *d_W, float *d_b, int out_start, int out_end, int aligned_in)
static void ck_train_gemm_backward_serial(const float *d_output, const float *input, const float *W, float *d_input, float *d_W, float *d_b, int T, int aligned_in, int aligned_out)
void gemm_blocked_serial_train_parallel_dispatch(const float *A, const float *B, const float *bias, float *C, int M, int N, int K)
void gemm_backward_f32_train_parallel_dispatch_v2(const float *d_output, const float *input, const float *W, float *d_input, float *d_W, float *d_b, int T, int aligned_in, int aligned_out, int num_threads)
static int ck_train_pick_active_threads(int nth, size_t work_items, size_t min_chunk)
static void ck_train_gemm_nn_work(int ith, int nth, void *argp)
static void ck_train_gemm_nt_compute_rows(const float *A, const float *B, const float *bias, float *C, int row_start, int row_end, int N, int K)
static void ck_train_gemm_nn_compute_rows(const float *A, const float *B, const float *bias, float *C, int row_start, int row_end, int N, int K)
static void ck_train_gemm_backward_work(int ith, int nth, void *argp)
static void ck_train_gemm_work(int ith, int nth, void *argp)
static void ck_train_outer_t1_work(int ith, int nth, void *argp)
void gemm_backward_f32_train_parallel_dispatch(const float *d_output, const float *input, const float *W, float *d_input, float *d_W, float *d_b, int T, int aligned_in, int aligned_out, int num_threads)
Persistent pthread thread pool for CK-Engine inference.
void ck_threadpool_dispatch_n(ck_threadpool_t *pool, int active_threads, ck_work_fn_t fn, void *args)
ck_threadpool_t * ck_threadpool_global(void)
int ck_threadpool_n_threads(const ck_threadpool_t *pool)
void gemm_nn_simd(const float *A, const float *B, const float *bias, float *C, int M, int N, int K)
void gemm_blocked_serial(const float *A, const float *B, const float *bias, float *C, int M, int N, int K)
#define C(color)
Definition show_config.c:39