pulsatrix
Loading...
Searching...
No Matches
hip_backend.hpp
Go to the documentation of this file.
1
5#pragma once
6
7#include <hip/hip_runtime.h>
8#include <hipblas/hipblas.h>
9
11
12namespace pulsatrix {
13
35class HIPBackend : public DeviceBackend {
36public:
38 ~HIPBackend() override;
39
40 HIPBackend(const HIPBackend&) = delete;
41 HIPBackend& operator=(const HIPBackend&) = delete;
42
43 [[nodiscard]] DeviceType device() const noexcept override { return DeviceType::Hip; }
44
45 [[nodiscard]] void* allocate(size_t bytes) override;
46 void free(void* ptr) noexcept override;
47 void copy(void* dst, const void* src, size_t bytes, CopyDirection dir) override;
48 void fill(void* ptr, float value, size_t n) override;
49 void gemm(const float* a, const float* b, float* out, size_t m, size_t k, size_t n) override;
50 void elementwise(ElementwiseOp op, const float* in, float* out, size_t n) override;
51 void add(const float* a, const float* b, float* out, size_t n) override;
52 void mul(const float* a, const float* b, float* out, size_t n) override;
53 void gemm_ex(const float* a, bool transpose_a, const float* b, bool transpose_b, float* out, size_t m, size_t k,
54 size_t n, float beta) override;
55 void column_sums(const float* in, float* out, size_t rows, size_t cols, float beta) override;
56 void add_row_vector(const float* in, const float* row, float* out, size_t rows, size_t cols) override;
57 void elementwise_backward(ElementwiseOp op, const float* x, const float* grad_out, float* grad_in,
58 size_t n) override;
59 void axpby(float alpha, const float* x, float beta, const float* y, float* out, size_t n) override;
60 [[nodiscard]] float dot(const float* a, const float* b, size_t n) override;
61 void softmax_rows(const float* in, float* out, size_t rows, size_t cols) override;
62 void softmax_rows_backward(const float* y, const float* dy, float* dx, size_t rows, size_t cols) override;
63 void logsumexp_rows(const float* in, float* out, size_t rows, size_t cols) override;
64 void adam_step(float* param, const float* grad, float* m, float* v, size_t n, float lr, float beta1, float beta2,
65 float eps, float bias_correction1, float bias_correction2) override;
66 [[nodiscard]] float sum(const float* in, size_t n) override;
67 void dropout_forward(const float* in, float* out, float* mask, size_t n, float p, float scale, uint64_t seed,
68 uint64_t offset) override;
69 void bce_with_logits(const float* logits, const float* target, float* out, size_t n) override;
70 void bce_with_logits_grad(const float* logits, const float* target, float* grad, size_t n, float scale) override;
71 void layer_norm_forward(const float* in, const float* gamma, const float* beta, float* xhat, float* out,
72 float* row_std, size_t rows, size_t cols, float eps) override;
73 void layer_norm_backward(const float* grad_out, const float* gamma, const float* xhat, const float* row_std,
74 float* grad_in, size_t rows, size_t cols) override;
75 void rms_norm_forward(const float* in, const float* gamma, float* out, float* row_rms, size_t rows, size_t cols,
76 float eps) override;
77 void rms_norm_backward(const float* grad_out, const float* gamma, const float* in, const float* row_rms,
78 float* grad_in, float* gamma_terms, size_t rows, size_t cols) override;
79 void rope_rotate(const float* in, const float* cos_table, const float* sin_table, float* out, size_t num_slices,
80 size_t seq_len, size_t head_dim, bool inverse) override;
81 void permute_0213(const float* in, float* out, size_t d0, size_t d1, size_t d2, size_t d3) override;
82 void gather_rows(const float* table, const float* indices, float* out, size_t count, size_t dim) override;
83 void scatter_add_rows(const float* src, const float* indices, float* table, size_t count, size_t dim) override;
84 void tanh_gaussian_forward(const float* mean, const float* log_std, const float* eps, float* action,
85 float* std_cache, float* log_prob, size_t rows, size_t cols, float stabilizer,
86 double half_log_two_pi) override;
87 void tanh_gaussian_backward(const float* action, const float* std_cache, const float* eps,
88 const float* grad_action, const float* grad_log_prob, float* grad_mean,
89 float* grad_log_std, size_t n, float stabilizer) override;
90 void lrp_linear(const float* x, const float* w, const float* z, const float* r, float* r_in, size_t rows,
91 size_t in_features, size_t out_features, float eps) override;
92 void lrp_residual_split(const float* a, const float* b, const float* r, float* r_a, float* r_b, size_t n,
93 float eps) override;
94 void lrp_bilinear_elementwise(const float* a, const float* b, const float* r, float* r_out, size_t n,
95 float eps) override;
96 void lrp_bilinear_matmul(const float* a, const float* b, const float* o, const float* r_o, float* r_a, float* r_b,
97 size_t slices, size_t m, size_t p, size_t q, float eps, bool b_transposed) override;
98 void lrp_softmax_rows(const float* x, const float* y, const float* r, float* r_in, size_t rows,
99 size_t cols) override;
100 void lrp_rope(const float* x, const float* y, const float* r, const float* cos_table, const float* sin_table,
101 float* r_in, size_t slices, size_t seq_len, size_t head_dim, float eps) override;
102 void logic_pointwise(LogicOp op, int norm, const float* a, const float* b, const float* g_or_r, const float* y,
103 float* out_a, float* out_b, size_t n, float eps) override;
104 void aggregator_forward(const float* x, float* mean_pow, float* out, size_t n, size_t cols, float p) override;
105 void aggregator_backward(const float* x, const float* mean_pow, const float* grad_out, float* grad_in, size_t n,
106 size_t cols, float p) override;
107 void aggregator_lrp(const float* x, const float* mean_pow, const float* r_out, float* r_in, size_t n, size_t cols,
108 float p, float eps) override;
109 void im2col(const float* in, float* col, size_t n, size_t c, size_t h, size_t w,
110 const ConvGeometry& geometry) override;
111 void col2im_add(const float* col, float* out, size_t n, size_t c, size_t h, size_t w,
112 const ConvGeometry& geometry) override;
113 void add_channel_vector(const float* in, const float* vec, float* out, size_t n, size_t c, size_t inner) override;
114 void lrp_conv(const float* col, const float* kernel, const float* pre_bias, const float* r, float* r_col, size_t n,
115 size_t out_channels, size_t p, size_t q, float eps) override;
116 void lrp_stabilized_divide(const float* r, const float* denom, const float* gate, float* out, size_t n, float eps,
117 LrpGate gate_mode) override;
118 void max_pool_forward(const float* in, float* out, float* argmax, size_t planes, size_t h, size_t w, size_t kh,
119 size_t kw) override;
120 void max_unpool(const float* src, const float* argmax, float* dst, size_t planes, size_t h, size_t w, size_t kh,
121 size_t kw) override;
122 void avg_pool_forward(const float* in, float* out, size_t planes, size_t h, size_t w, size_t kh, size_t kw)
123 override;
124 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
125 kw) override;
126 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
127 kw, float eps) override;
128 void batch_norm_forward(const float* in, const float* gamma, const float* beta, float* xhat, float* out, float*
129 channel_std, size_t n, size_t c, size_t spatial, float eps) override;
130 void batch_norm_backward(const float* grad_out, const float* gamma, const float* xhat, const float* channel_std,
131 float* grad_in, float* gamma_grad, float* beta_grad, size_t n, size_t c, size_t spatial)
132 override;
133 void batch_norm_update_running(const float* in, float* running_mean, float* running_var, size_t n, size_t c,
134 size_t spatial, float momentum) override;
135 void batch_norm_eval_forward(const float* in, const float* gamma, const float* beta, const float* running_mean,
136 const float* running_var, float* xhat, float* out, float* channel_std, size_t n,
137 size_t c, size_t spatial, float eps) override;
138 void batch_norm_eval_backward(const float* grad_out, const float* gamma, const float* xhat,
139 const float* channel_std, float* grad_in, float* gamma_grad, float* beta_grad,
140 size_t n, size_t c, size_t spatial) override;
141 void group_norm_forward(const float* in, const float* gamma, const float* beta, float* xhat, float* out, float*
142 group_std, size_t n, size_t c, size_t spatial, size_t num_groups, float eps) override;
143 void group_norm_backward(const float* grad_out, const float* gamma, const float* xhat, const float* group_std,
144 float* grad_in, float* gamma_grad, float* beta_grad, size_t n, size_t c, size_t spatial,
145 size_t num_groups) override;
146 void copy_2d(float* dst, size_t dst_stride, const float* src, size_t src_stride, size_t rows,
147 size_t cols) override;
148 void accumulate_rows(const float* in, float* out, size_t rows, size_t cols) override;
149 void recurrent_cell(RecurrentCellOp op, const RecurrentCellArgs& args, size_t n) override;
150 void gru_lrp_hprev(const float* h_prev, const float* w_hn, const float* hn, const float* r_term_b,
151 const float* direct, float* r_hprev, size_t rows, size_t hidden, float eps) override;
152 void ssm_pass(SsmPassOp op, const SsmPassArgs& args) override;
153 void rl_rows(RlRowOp op, const RlRowArgs& args) override;
154 void top_k_rows(const float* in, float* values, float* indices, size_t rows, size_t cols, size_t k,
155 bool largest) override;
156
157private:
158 hipStream_t stream_;
159 hipblasHandle_t hipblas_handle_;
160 // One device float that dot() reduces into before copying it to the host; allocated once
161 // so dot() costs no per-call device allocation.
162 float* dot_result_ = nullptr;
163};
164
165} // namespace pulsatrix
Vendor-agnostic compute/memory backend. CPUBackend, CUDABackend (Phase 1.5), and HIPBackend (Phase 1....
Definition device_backend.hpp:219
HIP-resident DeviceBackend implementation, targeting AMD GPUs via ROCm.
Definition hip_backend.hpp:35
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_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 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.
float sum(const float *in, size_t n) override
sum_i in[i], returned to the host. Same reduction order as dot(). Synchronizes.
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_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_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.
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 rl_rows(RlRowOp op, const RlRowArgs &args) override
One fused RL loss / target / Polyak pass (see RlRowOp for lanes and slots).
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 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 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 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 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 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 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).
HIPBackend(const HIPBackend &)=delete
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 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 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 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 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.
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 copy(void *dst, const void *src, size_t bytes, CopyDirection dir) override
Copies bytes between buffers.
void fill(void *ptr, float value, size_t n) override
Fills every element of a float buffer with a constant value.
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 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 free(void *ptr) noexcept override
Frees a buffer previously returned by allocate(). Safe to call with nullptr.
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 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 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 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 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).
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 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 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 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 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 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 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 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 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.
HIPBackend & operator=(const HIPBackend &)=delete
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 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 ssm_pass(SsmPassOp op, const SsmPassArgs &args) override
One fused Mamba / RWKV / RetNet pass (see SsmPassOp for lanes and slots).
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 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.
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.
DeviceType device() const noexcept override
Which device this backend's buffers reside on.
Definition hip_backend.hpp:43
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)).
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_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 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 recurrent_cell(RecurrentCellOp op, const RecurrentCellArgs &args, size_t n) override
One fused recurrent-cell pass over n elements (see RecurrentCellOp for slots).
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 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...
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 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 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 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 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 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 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 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 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 elementwise(ElementwiseOp op, const float *in, float *out, size_t n) override
Applies a unary elementwise operation to every element of a buffer.
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 * allocate(size_t bytes) override
Allocates a buffer of the given size.
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