90 const float*
in[10] = {};
157 const float*
in[12] = {};
200 const float*
in[5] = {};
241 [[nodiscard]] virtual
void*
allocate(
size_t bytes) = 0;
244 virtual
void free(
void* ptr) noexcept = 0;
262 virtual
void fill(
void* ptr,
float value,
size_t n) = 0;
274 virtual
void gemm(const
float* a, const
float* b,
float* out,
size_t m,
size_t k,
size_t n) = 0;
299 virtual
void add(const
float* a, const
float* b,
float* out,
size_t n) = 0;
315 virtual
void mul(const
float* a, const
float* b,
float* out,
size_t n) = 0;
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;
339 virtual
void column_sums(const
float* in,
float* out,
size_t rows,
size_t cols,
float beta) = 0;
345 virtual
void add_row_vector(const
float* in, const
float* row,
float* out,
size_t rows,
size_t cols) = 0;
364 virtual
void axpby(
float alpha, const
float* x,
float beta, const
float* y,
float* out,
size_t n) = 0;
371 [[nodiscard]] virtual
float dot(const
float* a, const
float* b,
size_t n) = 0;
374 virtual
void softmax_rows(const
float* in,
float* out,
size_t rows,
size_t cols) = 0;
383 virtual
void logsumexp_rows(const
float* in,
float* out,
size_t rows,
size_t cols) = 0;
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;
398 [[nodiscard]] virtual
float sum(const
float* in,
size_t n) = 0;
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;
420 virtual
void bce_with_logits(const
float* logits, const
float* target,
float* out,
size_t n) = 0;
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;
441 const
float* row_std,
float* grad_in,
size_t rows,
size_t cols) = 0;
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;
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;
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;
468 virtual
void permute_0213(const
float* in,
float* out,
size_t d0,
size_t d1,
size_t d2,
size_t d3) = 0;
471 virtual
void gather_rows(const
float* table, const
float* indices,
float* out,
size_t count,
size_t dim) = 0;
478 virtual
void scatter_add_rows(const
float* src, const
float* indices,
float* table,
size_t count,
size_t dim) = 0;
485 float* std_cache,
float* log_prob,
size_t rows,
size_t cols,
486 float stabilizer,
double half_log_two_pi) = 0;
490 const
float* grad_action, const
float* grad_log_prob,
float* grad_mean,
491 float* grad_log_std,
size_t n,
float stabilizer) = 0;
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;
506 virtual
void lrp_residual_split(const
float* a, const
float* b, const
float* r,
float* r_a,
float* r_b,
size_t n,
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;
523 virtual
void lrp_softmax_rows(const
float* x, const
float* y, const
float* r,
float* r_in,
size_t rows,
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,
539 const
float* y,
float* out_a,
float* out_b,
size_t n,
float eps) = 0;
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;
558 virtual
void im2col(const
float* in,
float* col,
size_t n,
size_t c,
size_t h,
size_t w,
566 virtual
void col2im_add(const
float* col,
float* out,
size_t n,
size_t c,
size_t h,
size_t w,
570 virtual
void add_channel_vector(const
float* in, const
float* vec,
float* out,
size_t n,
size_t c,
size_t inner) =
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;
588 float eps,
LrpGate gate_mode) = 0;
591 virtual
void max_pool_forward(const
float* in,
float* out,
float* argmax,
size_t planes,
size_t h,
size_t w,
size_t
595 virtual
void max_unpool(const
float* src, const
float* argmax,
float* dst,
size_t planes,
size_t h,
size_t w,
size_t
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)
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,
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;
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;
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;
622 size_t c,
size_t spatial,
float momentum) = 0;
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;
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;
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)
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;
651 virtual
void copy_2d(
float* dst,
size_t dst_stride, const
float* src,
size_t src_stride,
size_t rows,
660 virtual
void accumulate_rows(const
float* in,
float* out,
size_t rows,
size_t cols) = 0;
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;
704 virtual
void top_k_rows(const
float* in,
float* values,
float* indices,
size_t rows,
size_t cols,
size_t k,
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