pulsatrix
Loading...
Searching...
No Matches
device_backend.hpp
Go to the documentation of this file.
1
5#pragma once
6
7#include <cstddef>
8#include <cstdint>
9
10namespace pulsatrix {
11
17enum class DeviceType {
18 Cpu,
19 Cuda,
20 Hip
21};
22
30
43enum class ElementwiseOp {
44 Relu,
45 Neg,
46 Tanh,
47 Sigmoid,
48 Silu,
49 Exp
50};
51
53enum class LrpGate {
54 None,
55 Positive,
57};
58
72
87
90 const float* in[10] = {};
91 float* out[6] = {};
92 float eps = 0.0f;
93};
94
154
157 const float* in[12] = {};
158 float* out[10] = {};
159 int64_t n = 0;
160 int64_t l = 0;
161 int64_t d = 0;
162 int64_t s = 0;
163 float eps = 0.0f;
164 float gamma = 0.0f;
165};
166
172 size_t kh;
173 size_t kw;
174 size_t stride_h = 1;
175 size_t stride_w = 1;
176 size_t pad_h = 0;
177 size_t pad_w = 0;
178};
179
197
199struct RlRowArgs {
200 const float* in[5] = {};
201 float* out[4] = {};
202 int64_t rows = 0;
203 int64_t cols = 0;
204 float scale = 0.0f;
205 float lower = 0.0f;
206 float upper = 0.0f;
207 float gamma = 0.0f;
208 float tau = 0.0f;
209};
210
220public:
221 virtual ~DeviceBackend() = default;
222
231 [[nodiscard]] virtual DeviceType device() const noexcept = 0;
232
241 [[nodiscard]] virtual void* allocate(size_t bytes) = 0;
242
244 virtual void free(void* ptr) noexcept = 0;
245
254 virtual void copy(void* dst, const void* src, size_t bytes, CopyDirection dir) = 0;
255
262 virtual void fill(void* ptr, float value, size_t n) = 0;
263
274 virtual void gemm(const float* a, const float* b, float* out, size_t m, size_t k, size_t n) = 0;
275
284 virtual void elementwise(ElementwiseOp op, const float* in, float* out, size_t n) = 0;
285
299 virtual void add(const float* a, const float* b, float* out, size_t n) = 0;
300
315 virtual void mul(const float* a, const float* b, float* out, size_t n) = 0;
316
317 // ---- GPU-native-kernels Mission 1: primitives for device-resident training ----------
318 // Every reduction below has a fixed summation order on every backend (no atomics), so
319 // results are reproducible run to run.
320
332 virtual void gemm_ex(const float* a, bool transpose_a, const float* b, bool transpose_b, float* out, size_t m,
333 size_t k, size_t n, float beta) = 0;
334
339 virtual void column_sums(const float* in, float* out, size_t rows, size_t cols, float beta) = 0;
340
345 virtual void add_row_vector(const float* in, const float* row, float* out, size_t rows, size_t cols) = 0;
346
356 virtual void elementwise_backward(ElementwiseOp op, const float* x, const float* grad_out, float* grad_in,
357 size_t n) = 0;
358
364 virtual void axpby(float alpha, const float* x, float beta, const float* y, float* out, size_t n) = 0;
365
371 [[nodiscard]] virtual float dot(const float* a, const float* b, size_t n) = 0;
372
374 virtual void softmax_rows(const float* in, float* out, size_t rows, size_t cols) = 0;
375
380 virtual void softmax_rows_backward(const float* y, const float* dy, float* dx, size_t rows, size_t cols) = 0;
381
383 virtual void logsumexp_rows(const float* in, float* out, size_t rows, size_t cols) = 0;
384
392 virtual void adam_step(float* param, const float* grad, float* m, float* v, size_t n, float lr, float beta1,
393 float beta2, float eps, float bias_correction1, float bias_correction2) = 0;
394
395 // ---- GPU-native-kernels Mission 1b ---------------------------------------------------
396
398 [[nodiscard]] virtual float sum(const float* in, size_t n) = 0;
399
410 virtual void dropout_forward(const float* in, float* out, float* mask, size_t n, float p, float scale,
411 uint64_t seed, uint64_t offset) = 0;
412
420 virtual void bce_with_logits(const float* logits, const float* target, float* out, size_t n) = 0;
421
426 virtual void bce_with_logits_grad(const float* logits, const float* target, float* grad, size_t n,
427 float scale) = 0;
428
429 // ---- GPU-native-kernels Mission 2: transformer building blocks -------------------------
430 // The fused row operations below run the same per-row source on every backend (src/
431 // row_math.hpp): CPU loops over rows, GPUs run one thread per row. row_std / row_rms /
432 // log_prob are per-row device buffers (rows entries). Index buffers hold whole numbers as
433 // floats, exact below 2^24.
434
436 virtual void layer_norm_forward(const float* in, const float* gamma, const float* beta, float* xhat, float* out,
437 float* row_std, size_t rows, size_t cols, float eps) = 0;
438
440 virtual void layer_norm_backward(const float* grad_out, const float* gamma, const float* xhat,
441 const float* row_std, float* grad_in, size_t rows, size_t cols) = 0;
442
444 virtual void rms_norm_forward(const float* in, const float* gamma, float* out, float* row_rms, size_t rows,
445 size_t cols, float eps) = 0;
446
451 virtual void rms_norm_backward(const float* grad_out, const float* gamma, const float* in, const float* row_rms,
452 float* grad_in, float* gamma_terms, size_t rows, size_t cols) = 0;
453
460 virtual void rope_rotate(const float* in, const float* cos_table, const float* sin_table, float* out,
461 size_t num_slices, size_t seq_len, size_t head_dim, bool inverse) = 0;
462
468 virtual void permute_0213(const float* in, float* out, size_t d0, size_t d1, size_t d2, size_t d3) = 0;
469
471 virtual void gather_rows(const float* table, const float* indices, float* out, size_t count, size_t dim) = 0;
472
478 virtual void scatter_add_rows(const float* src, const float* indices, float* table, size_t count, size_t dim) = 0;
479
484 virtual void tanh_gaussian_forward(const float* mean, const float* log_std, const float* eps, float* action,
485 float* std_cache, float* log_prob, size_t rows, size_t cols,
486 float stabilizer, double half_log_two_pi) = 0;
487
489 virtual void tanh_gaussian_backward(const float* action, const float* std_cache, const float* eps,
490 const float* grad_action, const float* grad_log_prob, float* grad_mean,
491 float* grad_log_std, size_t n, float stabilizer) = 0;
492
493 // ---- GPU-native-kernels Mission 3: LRP rules and logic modules --------------------------
494 // Shared per-output source in src/lrp_math.hpp (CPU loops, GPU one thread per output).
495 // Every reduction is owned by one thread and runs in the original loop order:
496 // deterministic, no atomics.
497
502 virtual void lrp_linear(const float* x, const float* w, const float* z, const float* r, float* r_in, size_t rows,
503 size_t in_features, size_t out_features, float eps) = 0;
504
506 virtual void lrp_residual_split(const float* a, const float* b, const float* r, float* r_a, float* r_b, size_t n,
507 float eps) = 0;
508
510 virtual void lrp_bilinear_elementwise(const float* a, const float* b, const float* r, float* r_out, size_t n,
511 float eps) = 0;
512
518 virtual void lrp_bilinear_matmul(const float* a, const float* b, const float* o, const float* r_o, float* r_a,
519 float* r_b, size_t slices, size_t m, size_t p, size_t q, float eps,
520 bool b_transposed) = 0;
521
523 virtual void lrp_softmax_rows(const float* x, const float* y, const float* r, float* r_in, size_t rows,
524 size_t cols) = 0;
525
527 virtual void lrp_rope(const float* x, const float* y, const float* r, const float* cos_table,
528 const float* sin_table, float* r_in, size_t slices, size_t seq_len, size_t head_dim,
529 float eps) = 0;
530
538 virtual void logic_pointwise(LogicOp op, int norm, const float* a, const float* b, const float* g_or_r,
539 const float* y, float* out_a, float* out_b, size_t n, float eps) = 0;
540
542 virtual void aggregator_forward(const float* x, float* mean_pow, float* out, size_t n, size_t cols, float p) = 0;
544 virtual void aggregator_backward(const float* x, const float* mean_pow, const float* grad_out, float* grad_in,
545 size_t n, size_t cols, float p) = 0;
547 virtual void aggregator_lrp(const float* x, const float* mean_pow, const float* r_out, float* r_in, size_t n,
548 size_t cols, float p, float eps) = 0;
549
550 // ---- GPU-native-kernels Mission 4: convolution, pooling, spatial norms -----------------
551 // Shared per-output source in src/cnn_math.hpp; deterministic, no atomics.
552
558 virtual void im2col(const float* in, float* col, size_t n, size_t c, size_t h, size_t w,
559 const ConvGeometry& geometry) = 0;
560
566 virtual void col2im_add(const float* col, float* out, size_t n, size_t c, size_t h, size_t w,
567 const ConvGeometry& geometry) = 0;
568
570 virtual void add_channel_vector(const float* in, const float* vec, float* out, size_t n, size_t c, size_t inner) =
571 0;
572
577 virtual void lrp_conv(const float* col, const float* kernel, const float* pre_bias, const float* r, float* r_col,
578 size_t n, size_t out_channels, size_t p, size_t q, float eps) = 0;
579
580 // ---- LRP-rules campaign Mission 2: Zennit-compatible Gamma / AlphaBeta / ZBox ----------
587 virtual void lrp_stabilized_divide(const float* r, const float* denom, const float* gate, float* out, size_t n,
588 float eps, LrpGate gate_mode) = 0;
589
591 virtual void max_pool_forward(const float* in, float* out, float* argmax, size_t planes, size_t h, size_t w, size_t
592 kh, size_t kw) = 0;
593
595 virtual void max_unpool(const float* src, const float* argmax, float* dst, size_t planes, size_t h, size_t w, size_t
596 kh, size_t kw) = 0;
597
599 virtual void avg_pool_forward(const float* in, float* out, size_t planes, size_t h, size_t w, size_t kh, size_t kw)
600 = 0;
601
603 virtual void avg_pool_backward(const float* grad_out, float* grad_in, size_t planes, size_t h, size_t w, size_t kh,
604 size_t kw) = 0;
605
607 virtual void lrp_avg_pool(const float* x, const float* r, float* r_in, size_t planes, size_t h, size_t w, size_t kh,
608 size_t kw, float eps) = 0;
609
611 virtual void batch_norm_forward(const float* in, const float* gamma, const float* beta, float* xhat, float* out,
612 float* channel_std, size_t n, size_t c, size_t spatial, float eps) = 0;
613
615 virtual void batch_norm_backward(const float* grad_out, const float* gamma, const float* xhat, const float*
616 channel_std, float* grad_in, float* gamma_grad, float* beta_grad, size_t n, size_t
617 c, size_t spatial) = 0;
618
621 virtual void batch_norm_update_running(const float* in, float* running_mean, float* running_var, size_t n,
622 size_t c, size_t spatial, float momentum) = 0;
623
625 virtual void batch_norm_eval_forward(const float* in, const float* gamma, const float* beta,
626 const float* running_mean, const float* running_var, float* xhat, float* out,
627 float* channel_std, size_t n, size_t c, size_t spatial, float eps) = 0;
628
631 virtual void batch_norm_eval_backward(const float* grad_out, const float* gamma, const float* xhat,
632 const float* channel_std, float* grad_in, float* gamma_grad,
633 float* beta_grad, size_t n, size_t c, size_t spatial) = 0;
634
636 virtual void group_norm_forward(const float* in, const float* gamma, const float* beta, float* xhat, float* out,
637 float* group_std, size_t n, size_t c, size_t spatial, size_t num_groups, float eps)
638 = 0;
639
641 virtual void group_norm_backward(const float* grad_out, const float* gamma, const float* xhat, const float*
642 group_std, float* grad_in, float* gamma_grad, float* beta_grad, size_t n, size_t c,
643 size_t spatial, size_t num_groups) = 0;
644
645 // ---- GPU-native-kernels Mission 5: recurrent networks ------------------------------------
646
651 virtual void copy_2d(float* dst, size_t dst_stride, const float* src, size_t src_stride, size_t rows,
652 size_t cols) = 0;
653
660 virtual void accumulate_rows(const float* in, float* out, size_t rows, size_t cols) = 0;
661
663 virtual void recurrent_cell(RecurrentCellOp op, const RecurrentCellArgs& args, size_t n) = 0;
664
669 virtual void gru_lrp_hprev(const float* h_prev, const float* w_hn, const float* hn, const float* r_term_b,
670 const float* direct, float* r_hprev, size_t rows, size_t hidden, float eps) = 0;
671
672 // ---- GPU-native-kernels Mission 6: state-space models -------------------------------------
673 // Shared per-lane source in src/ssm_math.hpp; deterministic, no atomics.
674
681 virtual void ssm_pass(SsmPassOp op, const SsmPassArgs& args) = 0;
682
683 // ---- GPU-native-kernels Mission 7: reinforcement learning -----------------------------------
684 // Shared per-row source in src/rl_math.hpp; deterministic, no atomics.
685
692 virtual void rl_rows(RlRowOp op, const RlRowArgs& args) = 0;
693
694 // ---- FND-3: selection ------------------------------------------------------------------
695
704 virtual void top_k_rows(const float* in, float* values, float* indices, size_t rows, size_t cols, size_t k,
705 bool largest) = 0;
706};
707
708} // namespace pulsatrix
Vendor-agnostic compute/memory backend. CPUBackend, CUDABackend (Phase 1.5), and HIPBackend (Phase 1....
Definition device_backend.hpp:219
virtual void copy_2d(float *dst, size_t dst_stride, const float *src, size_t src_stride, size_t rows, size_t cols)=0
Strided 2-D copy: rows of cols floats from src (row stride src_stride) to dst (row stride dst_stride)...
virtual void scatter_add_rows(const float *src, const float *indices, float *table, size_t count, size_t dim)=0
table[indices[i]][:] += src[i][:] for i = 0..count-1, in increasing i.
virtual void aggregator_lrp(const float *x, const float *mean_pow, const float *r_out, float *r_in, size_t n, size_t cols, float p, float eps)=0
AggregatorModule epsilon rule, per column.
virtual void bce_with_logits_grad(const float *logits, const float *target, float *grad, size_t n, float scale)=0
BCE-with-logits gradient: grad[i] = (sigmoid(x) - y) * scale, using the overflow-free sigmoid (exp(x)...
virtual void layer_norm_backward(const float *grad_out, const float *gamma, const float *xhat, const float *row_std, float *grad_in, size_t rows, size_t cols)=0
LayerNorm input gradient per row, from the cached xhat and per-row std.
virtual void gemm(const float *a, const float *b, float *out, size_t m, size_t k, size_t n)=0
Row-major matrix multiply: out = a * b.
virtual void lrp_rope(const float *x, const float *y, const float *r, const float *cos_table, const float *sin_table, float *r_in, size_t slices, size_t seq_len, size_t head_dim, float eps)=0
RoPEModule epsilon rule over (slices, seq_len, head_dim), tables as for rope_rotate.
virtual void * allocate(size_t bytes)=0
Allocates a buffer of the given size.
virtual void column_sums(const float *in, float *out, size_t rows, size_t cols, float beta)=0
Per-column sum of a (rows x cols) row-major matrix: out[j] = beta*out[j] + sum_i in[i][j].
virtual void rms_norm_backward(const float *grad_out, const float *gamma, const float *in, const float *row_rms, float *grad_in, float *gamma_terms, size_t rows, size_t cols)=0
RMSNorm input gradient per row, plus gamma_terms[r][i] = grad_out * x / rms – the per-row contributio...
virtual ~DeviceBackend()=default
virtual void dropout_forward(const float *in, float *out, float *mask, size_t n, float p, float scale, uint64_t seed, uint64_t offset)=0
Inverted dropout with a counter-based RNG: element i is dropped iff uniform(seed, offset + i) < p; ke...
virtual void add_channel_vector(const float *in, const float *vec, float *out, size_t n, size_t c, size_t inner)=0
out[i][ch][k] = in[i][ch][k] + vec[ch] over (n, c, inner). out may alias in.
virtual void layer_norm_forward(const float *in, const float *gamma, const float *beta, float *xhat, float *out, float *row_std, size_t rows, size_t cols, float eps)=0
LayerNorm forward per row: xhat = (x - mean)/sqrt(var + eps), out = gamma*xhat + beta.
virtual DeviceType device() const noexcept=0
Which device this backend's buffers reside on.
virtual void recurrent_cell(RecurrentCellOp op, const RecurrentCellArgs &args, size_t n)=0
One fused recurrent-cell pass over n elements (see RecurrentCellOp for slots).
virtual void permute_0213(const float *in, float *out, size_t d0, size_t d1, size_t d2, size_t d3)=0
Swaps the middle two axes: in (d0, d1, d2, d3) -> out (d0, d2, d1, d3).
virtual void fill(void *ptr, float value, size_t n)=0
Fills every element of a float buffer with a constant value.
virtual void softmax_rows(const float *in, float *out, size_t rows, size_t cols)=0
Row-wise softmax of a (rows x cols) matrix, max-subtracted. out may alias in.
virtual void elementwise_backward(ElementwiseOp op, const float *x, const float *grad_out, float *grad_in, size_t n)=0
Activation backward: grad_in[i] = grad_out[i] * f'(x[i]), f = op, x = the forward input.
virtual void batch_norm_backward(const float *grad_out, const float *gamma, const float *xhat, const float *channel_std, float *grad_in, float *gamma_grad, float *beta_grad, size_t n, size_t c, size_t spatial)=0
BatchNorm input gradient plus this call's gamma/beta gradients (overwritten, per channel).
virtual void max_unpool(const float *src, const float *argmax, float *dst, size_t planes, size_t h, size_t w, size_t kh, size_t kw)=0
dst[plane][argmax] = src for every pooled element; dst must be zeroed by the caller.
virtual void axpby(float alpha, const float *x, float beta, const float *y, float *out, size_t n)=0
out[i] = alpha * x[i] + beta * y[i]. out may alias x or y.
virtual void accumulate_rows(const float *in, float *out, size_t rows, size_t cols)=0
out[j] += in[i][j] for i = 0..rows-1 in order, accumulating straight into out.
virtual void rope_rotate(const float *in, const float *cos_table, const float *sin_table, float *out, size_t num_slices, size_t seq_len, size_t head_dim, bool inverse)=0
Rotary position embedding over (num_slices, seq_len, head_dim) data.
virtual void lrp_conv(const float *col, const float *kernel, const float *pre_bias, const float *r, float *r_col, size_t n, size_t out_channels, size_t p, size_t q, float eps)=0
Conv2D epsilon rule in patch space: r_col (n, p, q) from the cached patches, the kernel (out_channels...
virtual float sum(const float *in, size_t n)=0
sum_i in[i], returned to the host. Same reduction order as dot(). Synchronizes.
virtual void top_k_rows(const float *in, float *values, float *indices, size_t rows, size_t cols, size_t k, bool largest)=0
Per row of a row-major (rows, cols) matrix: the k largest (or smallest) values in rank order into val...
virtual float dot(const float *a, const float *b, size_t n)=0
Dot product sum_i a[i]*b[i], returned to the host.
virtual void lrp_stabilized_divide(const float *r, const float *denom, const float *gate, float *out, size_t n, float eps, LrpGate gate_mode)=0
Gated stabilized division, the one non-gemm step of the affine LRP rules: out[i] = passes(gate[i]) ?...
virtual void add(const float *a, const float *b, float *out, size_t n)=0
Elementwise binary addition: out[i] = a[i] + b[i] for i in [0, n).
virtual void lrp_linear(const float *x, const float *w, const float *z, const float *r, float *r_in, size_t rows, size_t in_features, size_t out_features, float eps)=0
LinearModule epsilon rule: r_in[n][i] = sum_j (x[n][i] w[i][j] / stab(z[n][j])) r[n][j].
virtual void tanh_gaussian_forward(const float *mean, const float *log_std, const float *eps, float *action, float *std_cache, float *log_prob, size_t rows, size_t cols, float stabilizer, double half_log_two_pi)=0
TanhGaussianPolicy sampling per (rows, cols) row: action = tanh(mean + exp(log_std)*eps),...
virtual void copy(void *dst, const void *src, size_t bytes, CopyDirection dir)=0
Copies bytes between buffers.
virtual void rms_norm_forward(const float *in, const float *gamma, float *out, float *row_rms, size_t rows, size_t cols, float eps)=0
RMSNorm forward per row: out = gamma * x / sqrt(mean(x^2) + eps).
virtual void free(void *ptr) noexcept=0
Frees a buffer previously returned by allocate(). Safe to call with nullptr.
virtual void lrp_residual_split(const float *a, const float *b, const float *r, float *r_a, float *r_b, size_t n, float eps)=0
Epsilon split of a residual sum y = a + b: r_a = (a / stab(y)) r, r_b = (b / stab(y)) r.
virtual void softmax_rows_backward(const float *y, const float *dy, float *dx, size_t rows, size_t cols)=0
Softmax backward from its output y: dx[i][j] = y[i][j] * (dy[i][j] - sum_k y[i][k]*dy[i][k]).
virtual void gather_rows(const float *table, const float *indices, float *out, size_t count, size_t dim)=0
out[i][:] = table[indices[i]][:] for count rows of width dim.
virtual void aggregator_forward(const float *x, float *mean_pow, float *out, size_t n, size_t cols, float p)=0
AggregatorModule power mean over the leading axis of an (n, cols) input, per column.
virtual void logic_pointwise(LogicOp op, int norm, const float *a, const float *b, const float *g_or_r, const float *y, float *out_a, float *out_b, size_t n, float eps)=0
One elementwise pass of Conjunction/Disjunction for operands a, b.
virtual void aggregator_backward(const float *x, const float *mean_pow, const float *grad_out, float *grad_in, size_t n, size_t cols, float p)=0
AggregatorModule input gradient, per column.
virtual void lrp_avg_pool(const float *x, const float *r, float *r_in, size_t planes, size_t h, size_t w, size_t kh, size_t kw, float eps)=0
AvgPool2D epsilon rule; r_in must be zeroed by the caller.
virtual void max_pool_forward(const float *in, float *out, float *argmax, size_t planes, size_t h, size_t w, size_t kh, size_t kw)=0
Non-overlapping max pool over planes of (h, w); argmax = flat in-plane index (first max wins).
virtual void batch_norm_eval_forward(const float *in, const float *gamma, const float *beta, const float *running_mean, const float *running_var, float *xhat, float *out, float *channel_std, size_t n, size_t c, size_t spatial, float eps)=0
Eval-mode BatchNorm from the running statistics: a per-channel affine map. FND-5.
virtual void batch_norm_forward(const float *in, const float *gamma, const float *beta, float *xhat, float *out, float *channel_std, size_t n, size_t c, size_t spatial, float eps)=0
BatchNorm over (n, spatial) per channel of (n, c, spatial) data.
virtual void elementwise(ElementwiseOp op, const float *in, float *out, size_t n)=0
Applies a unary elementwise operation to every element of a buffer.
virtual void mul(const float *a, const float *b, float *out, size_t n)=0
Elementwise binary (Hadamard) multiplication: out[i] = a[i] * b[i] for i in [0, n).
virtual void lrp_bilinear_elementwise(const float *a, const float *b, const float *r, float *r_out, size_t n, float eps)=0
Eq. 15 for an elementwise product c = a*b: r_out = (a b / (2c + eps sign c)) r (same for both).
virtual void lrp_bilinear_matmul(const float *a, const float *b, const float *o, const float *r_o, float *r_a, float *r_b, size_t slices, size_t m, size_t p, size_t q, float eps, bool b_transposed)=0
Eq. 15 for slices independent matmuls O = A @ B (A (M x P), B (P x Q), O and r_o (M x Q)).
virtual void ssm_pass(SsmPassOp op, const SsmPassArgs &args)=0
One fused Mamba / RWKV / RetNet pass (see SsmPassOp for lanes and slots).
virtual void logsumexp_rows(const float *in, float *out, size_t rows, size_t cols)=0
Per-row log-sum-exp, max-subtracted: out[i] = log sum_j exp(in[i][j]).
virtual void tanh_gaussian_backward(const float *action, const float *std_cache, const float *eps, const float *grad_action, const float *grad_log_prob, float *grad_mean, float *grad_log_std, size_t n, float stabilizer)=0
TanhGaussianPolicy gradients w.r.t. mean and log_std, per element.
virtual void avg_pool_forward(const float *in, float *out, size_t planes, size_t h, size_t w, size_t kh, size_t kw)=0
Non-overlapping average pool over planes of (h, w).
virtual void group_norm_forward(const float *in, const float *gamma, const float *beta, float *xhat, float *out, float *group_std, size_t n, size_t c, size_t spatial, size_t num_groups, float eps)=0
GroupNorm per (example, group) of (n, c, spatial) data; group_std is (n, num_groups).
virtual void rl_rows(RlRowOp op, const RlRowArgs &args)=0
One fused RL loss / target / Polyak pass (see RlRowOp for lanes and slots).
virtual void lrp_softmax_rows(const float *x, const float *y, const float *r, float *r_in, size_t rows, size_t cols)=0
SoftmaxModule rule per row: r_in = x * (r - y * sum(r)).
virtual void im2col(const float *in, float *col, size_t n, size_t c, size_t h, size_t w, const ConvGeometry &geometry)=0
Unfolds (n, c, h, w) into (n, c*kh*kw, out_h*out_w) patches, with out_h = (h + 2*pad_h - kh) / stride...
virtual void avg_pool_backward(const float *grad_out, float *grad_in, size_t planes, size_t h, size_t w, size_t kh, size_t kw)=0
Spreads grad_out / (kh*kw) over each window; grad_in must be zeroed by the caller.
virtual void add_row_vector(const float *in, const float *row, float *out, size_t rows, size_t cols)=0
Broadcast row add: out[i][j] = in[i][j] + row[j] for a (rows x cols) matrix.
virtual void col2im_add(const float *col, float *out, size_t n, size_t c, size_t h, size_t w, const ConvGeometry &geometry)=0
Folds (n, c*kh*kw, out_h*out_w) patches back, adding into out (n, c, h, w).
virtual void batch_norm_update_running(const float *in, float *running_mean, float *running_var, size_t n, size_t c, size_t spatial, float momentum)=0
Folds this batch's per-channel mean and unbiased variance into the running ones (PyTorch's momentum r...
virtual void gemm_ex(const float *a, bool transpose_a, const float *b, bool transpose_b, float *out, size_t m, size_t k, size_t n, float beta)=0
General row-major matrix multiply: out = op(A) * op(B) + beta * out.
virtual void gru_lrp_hprev(const float *h_prev, const float *w_hn, const float *hn, const float *r_term_b, const float *direct, float *r_hprev, size_t rows, size_t hidden, float eps)=0
GRU's R_hprev for (rows, hidden): the recurrent epsilon-rule sum through W_hn, with each element's di...
virtual void group_norm_backward(const float *grad_out, const float *gamma, const float *xhat, const float *group_std, float *grad_in, float *gamma_grad, float *beta_grad, size_t n, size_t c, size_t spatial, size_t num_groups)=0
GroupNorm input gradient plus this call's gamma/beta gradients (overwritten, per channel).
virtual void batch_norm_eval_backward(const float *grad_out, const float *gamma, const float *xhat, const float *channel_std, float *grad_in, float *gamma_grad, float *beta_grad, size_t n, size_t c, size_t spatial)=0
Eval-mode BatchNorm gradient: grad_out * gamma / std, plus gamma/beta gradients (overwritten,...
virtual void adam_step(float *param, const float *grad, float *m, float *v, size_t n, float lr, float beta1, float beta2, float eps, float bias_correction1, float bias_correction2)=0
One fused Adam update over n parameters.
virtual void bce_with_logits(const float *logits, const float *target, float *out, size_t n)=0
Per-element binary cross-entropy with logits: out[i] = max(x, 0) - x*y + log1p(exp(-|x|)),...
Definition acquisition_functions.hpp:16
LogicOp
Elementwise passes of the fuzzy-logic modules, for DeviceBackend::logic_pointwise.
Definition device_backend.hpp:64
DeviceType
Which physical device a Tensor's buffer resides on.
Definition device_backend.hpp:17
CopyDirection
Direction of a DeviceBackend::copy() call.
Definition device_backend.hpp:24
SsmPassOp
Fused passes of the state-space / linear-recurrence modules (MambaModule, RWKVModule,...
Definition device_backend.hpp:133
RecurrentCellOp
Fused per-element recurrent-cell passes, for DeviceBackend::recurrent_cell. Slots (in[] / out[]),...
Definition device_backend.hpp:86
LrpGate
Elementwise boolean gate for DeviceBackend::lrp_stabilized_divide().
Definition device_backend.hpp:53
@ Positive
passes where gate > 0
@ None
every element passes
@ Negative
passes where gate < 0
RlRowOp
Fused per-row reinforcement-learning passes, for DeviceBackend::rl_rows. One lane per batch row (per ...
Definition device_backend.hpp:196
ElementwiseOp
Unary elementwise operations supported by DeviceBackend::elementwise().
Definition device_backend.hpp:43
@ Silu
x * sigmoid(x) – a.k.a. swish; the gate half of SwiGLU
@ Sigmoid
1 / (1 + exp(-x))
@ Exp
exp(x) – GPU-native-kernels Mission 1b (Reparameterize, KL divergence)
Window geometry for DeviceBackend::im2col / col2im_add: kernel size, stride and zero padding per axis...
Definition device_backend.hpp:171
size_t kh
Definition device_backend.hpp:172
size_t pad_w
Definition device_backend.hpp:177
size_t stride_w
Definition device_backend.hpp:175
size_t kw
Definition device_backend.hpp:173
size_t stride_h
Definition device_backend.hpp:174
size_t pad_h
Definition device_backend.hpp:176
Operand pointers for DeviceBackend::recurrent_cell (passed to kernels by value).
Definition device_backend.hpp:89
float eps
LRP stabilizer (LstmLrp, GruLrp)
Definition device_backend.hpp:92
float * out[6]
Definition device_backend.hpp:91
const float * in[10]
Definition device_backend.hpp:90
Operand pointers and dims for DeviceBackend::rl_rows (passed to kernels by value).
Definition device_backend.hpp:199
float lower
PPO 1 - clip_epsilon.
Definition device_backend.hpp:205
int64_t cols
Definition device_backend.hpp:203
float tau
Polyak blend factor.
Definition device_backend.hpp:208
int64_t rows
Definition device_backend.hpp:202
float * out[4]
Definition device_backend.hpp:201
float gamma
discount factor
Definition device_backend.hpp:207
float upper
PPO 1 + clip_epsilon.
Definition device_backend.hpp:206
const float * in[5]
Definition device_backend.hpp:200
float scale
gradient batch-mean scale (1/N or 2/N)
Definition device_backend.hpp:204
Operand pointers and dims for DeviceBackend::ssm_pass (passed to kernels by value).
Definition device_backend.hpp:156
int64_t d
Definition device_backend.hpp:161
float gamma
RetNet decay.
Definition device_backend.hpp:164
float eps
LRP stabilizer.
Definition device_backend.hpp:163
int64_t l
Definition device_backend.hpp:160
int64_t s
Definition device_backend.hpp:162
int64_t n
Definition device_backend.hpp:159
const float * in[12]
Definition device_backend.hpp:157
float * out[10]
Definition device_backend.hpp:158