8#include <cuda_runtime.h>
38 [[nodiscard]]
void*
allocate(
size_t bytes)
override;
39 void free(
void* ptr)
noexcept override;
41 void fill(
void* ptr,
float value,
size_t n)
override;
42 void gemm(
const float* a,
const float* b,
float* out,
size_t m,
size_t k,
size_t n)
override;
44 void add(
const float* a,
const float* b,
float* out,
size_t n)
override;
45 void mul(
const float* a,
const float* b,
float* out,
size_t n)
override;
46 void gemm_ex(
const float* a,
bool transpose_a,
const float* b,
bool transpose_b,
float* out,
size_t m,
size_t k,
47 size_t n,
float beta)
override;
48 void column_sums(
const float* in,
float* out,
size_t rows,
size_t cols,
float beta)
override;
49 void add_row_vector(
const float* in,
const float* row,
float* out,
size_t rows,
size_t cols)
override;
52 void axpby(
float alpha,
const float* x,
float beta,
const float* y,
float* out,
size_t n)
override;
53 [[nodiscard]]
float dot(
const float* a,
const float* b,
size_t n)
override;
54 void softmax_rows(
const float* in,
float* out,
size_t rows,
size_t cols)
override;
56 void logsumexp_rows(
const float* in,
float* out,
size_t rows,
size_t cols)
override;
57 void adam_step(
float* param,
const float* grad,
float* m,
float* v,
size_t n,
float lr,
float beta1,
float beta2,
58 float eps,
float bias_correction1,
float bias_correction2)
override;
59 [[nodiscard]]
float sum(
const float* in,
size_t n)
override;
60 void dropout_forward(
const float* in,
float* out,
float* mask,
size_t n,
float p,
float scale, uint64_t seed,
61 uint64_t offset)
override;
62 void bce_with_logits(
const float* logits,
const float* target,
float* out,
size_t n)
override;
63 void bce_with_logits_grad(
const float* logits,
const float* target,
float* grad,
size_t n,
float scale)
override;
64 void layer_norm_forward(
const float* in,
const float* gamma,
const float* beta,
float* xhat,
float* out,
65 float* row_std,
size_t rows,
size_t cols,
float eps)
override;
66 void layer_norm_backward(
const float* grad_out,
const float* gamma,
const float* xhat,
const float* row_std,
67 float* grad_in,
size_t rows,
size_t cols)
override;
68 void rms_norm_forward(
const float* in,
const float* gamma,
float* out,
float* row_rms,
size_t rows,
size_t cols,
70 void rms_norm_backward(
const float* grad_out,
const float* gamma,
const float* in,
const float* row_rms,
71 float* grad_in,
float* gamma_terms,
size_t rows,
size_t cols)
override;
72 void rope_rotate(
const float* in,
const float* cos_table,
const float* sin_table,
float* out,
size_t num_slices,
73 size_t seq_len,
size_t head_dim,
bool inverse)
override;
74 void permute_0213(
const float* in,
float* out,
size_t d0,
size_t d1,
size_t d2,
size_t d3)
override;
75 void gather_rows(
const float* table,
const float* indices,
float* out,
size_t count,
size_t dim)
override;
76 void scatter_add_rows(
const float* src,
const float* indices,
float* table,
size_t count,
size_t dim)
override;
78 float* std_cache,
float* log_prob,
size_t rows,
size_t cols,
float stabilizer,
79 double half_log_two_pi)
override;
81 const float* grad_action,
const float* grad_log_prob,
float* grad_mean,
82 float* grad_log_std,
size_t n,
float stabilizer)
override;
83 void lrp_linear(
const float* x,
const float* w,
const float* z,
const float* r,
float* r_in,
size_t rows,
84 size_t in_features,
size_t out_features,
float eps)
override;
85 void lrp_residual_split(
const float* a,
const float* b,
const float* r,
float* r_a,
float* r_b,
size_t n,
89 void lrp_bilinear_matmul(
const float* a,
const float* b,
const float* o,
const float* r_o,
float* r_a,
float* r_b,
90 size_t slices,
size_t m,
size_t p,
size_t q,
float eps,
bool b_transposed)
override;
91 void lrp_softmax_rows(
const float* x,
const float* y,
const float* r,
float* r_in,
size_t rows,
92 size_t cols)
override;
93 void lrp_rope(
const float* x,
const float* y,
const float* r,
const float* cos_table,
const float* sin_table,
94 float* r_in,
size_t slices,
size_t seq_len,
size_t head_dim,
float eps)
override;
96 float* out_a,
float* out_b,
size_t n,
float eps)
override;
97 void aggregator_forward(
const float* x,
float* mean_pow,
float* out,
size_t n,
size_t cols,
float p)
override;
98 void aggregator_backward(
const float* x,
const float* mean_pow,
const float* grad_out,
float* grad_in,
size_t n,
99 size_t cols,
float p)
override;
100 void aggregator_lrp(
const float* x,
const float* mean_pow,
const float* r_out,
float* r_in,
size_t n,
size_t cols,
101 float p,
float eps)
override;
102 void im2col(
const float* in,
float* col,
size_t n,
size_t c,
size_t h,
size_t w,
104 void col2im_add(
const float* col,
float* out,
size_t n,
size_t c,
size_t h,
size_t w,
106 void add_channel_vector(
const float* in,
const float* vec,
float* out,
size_t n,
size_t c,
size_t inner)
override;
107 void lrp_conv(
const float* col,
const float* kernel,
const float* pre_bias,
const float* r,
float* r_col,
size_t n,
108 size_t out_channels,
size_t p,
size_t q,
float eps)
override;
111 void max_pool_forward(
const float* in,
float* out,
float* argmax,
size_t planes,
size_t h,
size_t w,
size_t kh,
113 void max_unpool(
const float* src,
const float* argmax,
float* dst,
size_t planes,
size_t h,
size_t w,
size_t kh,
115 void avg_pool_forward(
const float* in,
float* out,
size_t planes,
size_t h,
size_t w,
size_t kh,
size_t kw)
117 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
119 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
120 kw,
float eps)
override;
121 void batch_norm_forward(
const float* in,
const float* gamma,
const float* beta,
float* xhat,
float* out,
float*
122 channel_std,
size_t n,
size_t c,
size_t spatial,
float eps)
override;
123 void batch_norm_backward(
const float* grad_out,
const float* gamma,
const float* xhat,
const float* channel_std,
124 float* grad_in,
float* gamma_grad,
float* beta_grad,
size_t n,
size_t c,
size_t spatial)
127 size_t spatial,
float momentum)
override;
129 const float* running_var,
float* xhat,
float* out,
float* channel_std,
size_t n,
130 size_t c,
size_t spatial,
float eps)
override;
132 const float* channel_std,
float* grad_in,
float* gamma_grad,
float* beta_grad,
133 size_t n,
size_t c,
size_t spatial)
override;
134 void group_norm_forward(
const float* in,
const float* gamma,
const float* beta,
float* xhat,
float* out,
float*
135 group_std,
size_t n,
size_t c,
size_t spatial,
size_t num_groups,
float eps)
override;
136 void group_norm_backward(
const float* grad_out,
const float* gamma,
const float* xhat,
const float* group_std,
137 float* grad_in,
float* gamma_grad,
float* beta_grad,
size_t n,
size_t c,
size_t spatial,
138 size_t num_groups)
override;
139 void copy_2d(
float* dst,
size_t dst_stride,
const float* src,
size_t src_stride,
size_t rows,
140 size_t cols)
override;
143 void gru_lrp_hprev(
const float* h_prev,
const float* w_hn,
const float* hn,
const float* r_term_b,
144 const float* direct,
float* r_hprev,
size_t rows,
size_t hidden,
float eps)
override;
147 void top_k_rows(
const float* in,
float* values,
float* indices,
size_t rows,
size_t cols,
size_t k,
148 bool largest)
override;
151 cudaStream_t stream_;
152 cublasHandle_t cublas_handle_;
155 float* dot_result_ =
nullptr;
CUDA-resident DeviceBackend implementation.
Definition cuda_backend.hpp:28
CUDABackend(const CUDABackend &)=delete
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) override
GroupNorm input gradient plus this call's gamma/beta gradients (overwritten, per channel).
void * allocate(size_t bytes) override
Allocates a buffer of the given size.
void column_sums(const float *in, float *out, size_t rows, size_t cols, float beta) override
Per-column sum of a (rows x cols) row-major matrix: out[j] = beta*out[j] + sum_i in[i][j].
void top_k_rows(const float *in, float *values, float *indices, size_t rows, size_t cols, size_t k, bool largest) override
Per row of a row-major (rows, cols) matrix: the k largest (or smallest) values in rank order into val...
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) override
TanhGaussianPolicy sampling per (rows, cols) row: action = tanh(mean + exp(log_std)*eps),...
void rl_rows(RlRowOp op, const RlRowArgs &args) override
One fused RL loss / target / Polyak pass (see RlRowOp for lanes and slots).
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) override
LayerNorm forward per row: xhat = (x - mean)/sqrt(var + eps), out = gamma*xhat + beta.
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) override
BatchNorm input gradient plus this call's gamma/beta gradients (overwritten, per channel).
void accumulate_rows(const float *in, float *out, size_t rows, size_t cols) override
out[j] += in[i][j] for i = 0..rows-1 in order, accumulating straight into out.
void fill(void *ptr, float value, size_t n) override
Fills every element of a float buffer with a constant value.
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) override
Eq. 15 for slices independent matmuls O = A @ B (A (M x P), B (P x Q), O and r_o (M x Q)).
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) override
AggregatorModule epsilon rule, per column.
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) override
BatchNorm over (n, spatial) per channel of (n, c, spatial) data.
CUDABackend & operator=(const CUDABackend &)=delete
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) override
RoPEModule epsilon rule over (slices, seq_len, head_dim), tables as for rope_rotate.
void lrp_residual_split(const float *a, const float *b, const float *r, float *r_a, float *r_b, size_t n, float eps) override
Epsilon split of a residual sum y = a + b: r_a = (a / stab(y)) r, r_b = (b / stab(y)) r.
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) override
Eval-mode BatchNorm gradient: grad_out * gamma / std, plus gamma/beta gradients (overwritten,...
void aggregator_forward(const float *x, float *mean_pow, float *out, size_t n, size_t cols, float p) override
AggregatorModule power mean over the leading axis of an (n, cols) input, per column.
void lrp_bilinear_elementwise(const float *a, const float *b, const float *r, float *r_out, size_t n, float eps) override
Eq. 15 for an elementwise product c = a*b: r_out = (a b / (2c + eps sign c)) r (same for both).
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) override
Non-overlapping max pool over planes of (h, w); argmax = flat in-plane index (first max wins).
void ssm_pass(SsmPassOp op, const SsmPassArgs &args) override
One fused Mamba / RWKV / RetNet pass (see SsmPassOp for lanes and slots).
void gather_rows(const float *table, const float *indices, float *out, size_t count, size_t dim) override
out[i][:] = table[indices[i]][:] for count rows of width dim.
void bce_with_logits_grad(const float *logits, const float *target, float *grad, size_t n, float scale) override
BCE-with-logits gradient: grad[i] = (sigmoid(x) - y) * scale, using the overflow-free sigmoid (exp(x)...
void elementwise(ElementwiseOp op, const float *in, float *out, size_t n) override
Applies a unary elementwise operation to every element of a buffer.
void softmax_rows_backward(const float *y, const float *dy, float *dx, size_t rows, size_t cols) override
Softmax backward from its output y: dx[i][j] = y[i][j] * (dy[i][j] - sum_k y[i][k]*dy[i][k]).
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) override
One fused Adam update over n parameters.
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) override
Spreads grad_out / (kh*kw) over each window; grad_in must be zeroed by the caller.
void copy(void *dst, const void *src, size_t bytes, CopyDirection dir) override
Copies bytes between buffers.
void bce_with_logits(const float *logits, const float *target, float *out, size_t n) override
Per-element binary cross-entropy with logits: out[i] = max(x, 0) - x*y + log1p(exp(-|x|)),...
void recurrent_cell(RecurrentCellOp op, const RecurrentCellArgs &args, size_t n) override
One fused recurrent-cell pass over n elements (see RecurrentCellOp for slots).
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) override
GroupNorm per (example, group) of (n, c, spatial) data; group_std is (n, num_groups).
void elementwise_backward(ElementwiseOp op, const float *x, const float *grad_out, float *grad_in, size_t n) override
Activation backward: grad_in[i] = grad_out[i] * f'(x[i]), f = op, x = the forward input.
void lrp_stabilized_divide(const float *r, const float *denom, const float *gate, float *out, size_t n, float eps, LrpGate gate_mode) override
Gated stabilized division, the one non-gemm step of the affine LRP rules: out[i] = passes(gate[i]) ?...
void dropout_forward(const float *in, float *out, float *mask, size_t n, float p, float scale, uint64_t seed, uint64_t offset) override
Inverted dropout with a counter-based RNG: element i is dropped iff uniform(seed, offset + i) < p; ke...
void add_channel_vector(const float *in, const float *vec, float *out, size_t n, size_t c, size_t inner) override
out[i][ch][k] = in[i][ch][k] + vec[ch] over (n, c, inner). out may alias in.
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) override
Folds this batch's per-channel mean and unbiased variance into the running ones (PyTorch's momentum r...
void avg_pool_forward(const float *in, float *out, size_t planes, size_t h, size_t w, size_t kh, size_t kw) override
Non-overlapping average pool over planes of (h, w).
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) override
LinearModule epsilon rule: r_in[n][i] = sum_j (x[n][i] w[i][j] / stab(z[n][j])) r[n][j].
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) override
General row-major matrix multiply: out = op(A) * op(B) + beta * out.
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) override
AvgPool2D epsilon rule; r_in must be zeroed by the caller.
void axpby(float alpha, const float *x, float beta, const float *y, float *out, size_t n) override
out[i] = alpha * x[i] + beta * y[i]. out may alias x or y.
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) override
GRU's R_hprev for (rows, hidden): the recurrent epsilon-rule sum through W_hn, with each element's di...
void scatter_add_rows(const float *src, const float *indices, float *table, size_t count, size_t dim) override
table[indices[i]][:] += src[i][:] for i = 0..count-1, in increasing i.
void logsumexp_rows(const float *in, float *out, size_t rows, size_t cols) override
Per-row log-sum-exp, max-subtracted: out[i] = log sum_j exp(in[i][j]).
void copy_2d(float *dst, size_t dst_stride, const float *src, size_t src_stride, size_t rows, size_t cols) override
Strided 2-D copy: rows of cols floats from src (row stride src_stride) to dst (row stride dst_stride)...
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) override
dst[plane][argmax] = src for every pooled element; dst must be zeroed by the caller.
void mul(const float *a, const float *b, float *out, size_t n) override
Elementwise binary (Hadamard) multiplication: out[i] = a[i] * b[i] for i in [0, n).
void softmax_rows(const float *in, float *out, size_t rows, size_t cols) override
Row-wise softmax of a (rows x cols) matrix, max-subtracted. out may alias in.
float dot(const float *a, const float *b, size_t n) override
Dot product sum_i a[i]*b[i], returned to the host.
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) override
RMSNorm input gradient per row, plus gamma_terms[r][i] = grad_out * x / rms – the per-row contributio...
void lrp_softmax_rows(const float *x, const float *y, const float *r, float *r_in, size_t rows, size_t cols) override
SoftmaxModule rule per row: r_in = x * (r - y * sum(r)).
void add(const float *a, const float *b, float *out, size_t n) override
Elementwise binary addition: out[i] = a[i] + b[i] for i in [0, n).
void add_row_vector(const float *in, const float *row, float *out, size_t rows, size_t cols) override
Broadcast row add: out[i][j] = in[i][j] + row[j] for a (rows x cols) matrix.
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) override
One elementwise pass of Conjunction/Disjunction for operands a, b.
void rms_norm_forward(const float *in, const float *gamma, float *out, float *row_rms, size_t rows, size_t cols, float eps) override
RMSNorm forward per row: out = gamma * x / sqrt(mean(x^2) + eps).
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) override
AggregatorModule input gradient, per column.
void col2im_add(const float *col, float *out, size_t n, size_t c, size_t h, size_t w, const ConvGeometry &geometry) override
Folds (n, c*kh*kw, out_h*out_w) patches back, adding into out (n, c, h, w).
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) override
TanhGaussianPolicy gradients w.r.t. mean and log_std, per element.
float sum(const float *in, size_t n) override
sum_i in[i], returned to the host. Same reduction order as dot(). Synchronizes.
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) override
Rotary position embedding over (num_slices, seq_len, head_dim) data.
void free(void *ptr) noexcept override
Frees a buffer previously returned by allocate(). Safe to call with nullptr.
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) override
Eval-mode BatchNorm from the running statistics: a per-channel affine map. FND-5.
void im2col(const float *in, float *col, size_t n, size_t c, size_t h, size_t w, const ConvGeometry &geometry) override
Unfolds (n, c, h, w) into (n, c*kh*kw, out_h*out_w) patches, with out_h = (h + 2*pad_h - kh) / stride...
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) override
LayerNorm input gradient per row, from the cached xhat and per-row std.
void gemm(const float *a, const float *b, float *out, size_t m, size_t k, size_t n) override
Row-major matrix multiply: out = a * b.
void permute_0213(const float *in, float *out, size_t d0, size_t d1, size_t d2, size_t d3) override
Swaps the middle two axes: in (d0, d1, d2, d3) -> out (d0, d2, d1, d3).
DeviceType device() const noexcept override
Which device this backend's buffers reside on.
Definition cuda_backend.hpp:36
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) override
Conv2D epsilon rule in patch space: r_col (n, p, q) from the cached patches, the kernel (out_channels...
Vendor-agnostic compute/memory backend. CPUBackend, CUDABackend (Phase 1.5), and HIPBackend (Phase 1....
Definition device_backend.hpp:219
Abstract interface isolating vendor-specific memory/compute operations from Tensor/ComputationGraph.
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
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
Window geometry for DeviceBackend::im2col / col2im_add: kernel size, stride and zero padding per axis...
Definition device_backend.hpp:171
Operand pointers for DeviceBackend::recurrent_cell (passed to kernels by value).
Definition device_backend.hpp:89
Operand pointers and dims for DeviceBackend::rl_rows (passed to kernels by value).
Definition device_backend.hpp:199
Operand pointers and dims for DeviceBackend::ssm_pass (passed to kernels by value).
Definition device_backend.hpp:156