← Back to C-Kernel-Engine Docs Doxygen Source Documentation
 
Loading...
Searching...
No Matches
ck_parallel_train.c File Reference

Thread-pool dispatch wrapper for FP32 training GEMM. More...

#include "ckernel_engine.h"
#include "ck_threadpool.h"
#include <stddef.h>

Go to the source code of this file.

Functions

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_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)
 
static void ck_train_gemm_backward_work (int ith, int nth, void *argp)
 
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_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_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_gemm_work (int ith, int nth, void *argp)
 
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_outer_t1_work (int ith, int nth, void *argp)
 
static int ck_train_pick_active_threads (int nth, size_t work_items, size_t min_chunk)
 
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)
 
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)
 
void gemm_blocked_serial_train_parallel_dispatch (const float *A, const float *B, const float *bias, float *C, int M, int N, int K)
 

Detailed Description

Thread-pool dispatch wrapper for FP32 training GEMM.

The training runtime currently emits many gemm_blocked_serial calls. This wrapper keeps kernel math unchanged and only parallelizes dispatch.

OpenMP removal note: Training backward used to fall back into legacy OpenMP GEMM paths for T>1 micro-batches, which mixed libiomp barriers with the CK threadpool and showed up as fork/barrier overhead in profiling. Keep the hot training path on the CK threadpool here, and leave older kernel-level OpenMP paths as compatibility fallbacks until they are fully retired.

Definition in file ck_parallel_train.c.

Function Documentation

◆ ck_train_bias_reduce_compute_range()

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

Definition at line 411 of file ck_parallel_train.c.

416 {
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}

Referenced by ck_train_gemm_backward_serial(), and ck_train_gemm_backward_work().

◆ ck_train_gemm_backward_serial()

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 
)
static

Definition at line 453 of file ck_parallel_train.c.

461 {
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}
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_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)

References ck_train_bias_reduce_compute_range(), ck_train_gemm_nn_compute_rows(), and ck_train_gemm_tn_compute_rows().

Referenced by gemm_backward_f32_train_parallel_dispatch(), and gemm_backward_f32_train_parallel_dispatch_v2().

◆ ck_train_gemm_backward_work()

static void ck_train_gemm_backward_work ( int  ith,
int  nth,
void *  argp 
)
static

Definition at line 556 of file ck_parallel_train.c.

556 {
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}

References ck_train_bias_reduce_compute_range(), ck_train_gemm_nn_compute_rows(), and ck_train_gemm_tn_compute_rows().

Referenced by gemm_backward_f32_train_parallel_dispatch(), and gemm_backward_f32_train_parallel_dispatch_v2().

◆ ck_train_gemm_nn_compute_rows()

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

Definition at line 274 of file ck_parallel_train.c.

281 {
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}
#define C(color)
Definition show_config.c:39

References C.

Referenced by ck_train_gemm_backward_serial(), ck_train_gemm_backward_work(), and ck_train_gemm_nn_work().

◆ ck_train_gemm_nn_work()

static void ck_train_gemm_nn_work ( int  ith,
int  nth,
void *  argp 
)
static

Definition at line 469 of file ck_parallel_train.c.

469 {
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}

References ck_train_gemm_nn_compute_rows().

Referenced by gemm_backward_f32_train_parallel_dispatch_v2().

◆ ck_train_gemm_nt_compute_rows()

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

Definition at line 48 of file ck_parallel_train.c.

55 {
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}

References C.

Referenced by ck_train_gemm_work().

◆ ck_train_gemm_tn_compute_rows()

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

Definition at line 341 of file ck_parallel_train.c.

349 {
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}

References C.

Referenced by ck_train_gemm_backward_serial(), and ck_train_gemm_backward_work().

◆ ck_train_gemm_work()

static void ck_train_gemm_work ( int  ith,
int  nth,
void *  argp 
)
static

Definition at line 119 of file ck_parallel_train.c.

119 {
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}
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)
void gemm_blocked_serial(const float *A, const float *B, const float *bias, float *C, int M, int N, int K)

References ck_train_gemm_nt_compute_rows(), and gemm_blocked_serial().

Referenced by gemm_blocked_serial_train_parallel_dispatch().

◆ ck_train_outer_t1_compute_range()

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

Definition at line 430 of file ck_parallel_train.c.

436 {
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}

Referenced by ck_train_outer_t1_work(), gemm_backward_f32_train_parallel_dispatch(), and gemm_backward_f32_train_parallel_dispatch_v2().

◆ ck_train_outer_t1_work()

static void ck_train_outer_t1_work ( int  ith,
int  nth,
void *  argp 
)
static

Definition at line 537 of file ck_parallel_train.c.

537 {
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}
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)

References ck_train_outer_t1_compute_range().

Referenced by gemm_backward_f32_train_parallel_dispatch(), and gemm_backward_f32_train_parallel_dispatch_v2().

◆ ck_train_pick_active_threads()

static int ck_train_pick_active_threads ( int  nth,
size_t  work_items,
size_t  min_chunk 
)
static

Definition at line 104 of file ck_parallel_train.c.

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}

Referenced by gemm_backward_f32_train_parallel_dispatch(), and gemm_backward_f32_train_parallel_dispatch_v2().

◆ gemm_backward_f32_train_parallel_dispatch()

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 
)

Definition at line 610 of file ck_parallel_train.c.

619 {
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}
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)
static int ck_train_pick_active_threads(int nth, size_t work_items, size_t min_chunk)
static void ck_train_gemm_backward_work(int ith, int nth, void *argp)
static void ck_train_outer_t1_work(int ith, int nth, void *argp)
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)

References ck_threadpool_dispatch_n(), ck_threadpool_global(), ck_threadpool_n_threads(), ck_train_gemm_backward_serial(), ck_train_gemm_backward_work(), ck_train_outer_t1_compute_range(), ck_train_outer_t1_work(), ck_train_pick_active_threads(), and gemm_nn_simd().

◆ gemm_backward_f32_train_parallel_dispatch_v2()

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 
)

Definition at line 699 of file ck_parallel_train.c.

708 {
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_gemm_nn_work(int ith, int nth, void *argp)

References ck_threadpool_dispatch_n(), ck_threadpool_global(), ck_threadpool_n_threads(), ck_train_gemm_backward_serial(), ck_train_gemm_backward_work(), ck_train_gemm_nn_work(), ck_train_outer_t1_compute_range(), ck_train_outer_t1_work(), ck_train_pick_active_threads(), and gemm_nn_simd().

◆ gemm_blocked_serial_train_parallel_dispatch()

void gemm_blocked_serial_train_parallel_dispatch ( const float *  A,
const float *  B,
const float *  bias,
float *  C,
int  M,
int  N,
int  K 
)

Definition at line 167 of file ck_parallel_train.c.

173 {
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}
static void ck_train_gemm_work(int ith, int nth, void *argp)

References C, ck_threadpool_dispatch_n(), ck_threadpool_global(), ck_threadpool_n_threads(), ck_train_gemm_work(), and gemm_blocked_serial().