Contents

namespace

SushiAI::Ops

Declared in
include/SushiAI/core/ops.hpp

Contains

Enumerations

enum class GradientSource

Which tensor an activation's gradient rule differentiates at.

INPUT

The forward input; the pre-activation must survive the step.

OUTPUT

The forward output; the pre-activation is dead on arrival.

Variables

constexpr OpName LEAF = op("sushiai.leaf")

A graph leaf: a parameter or an input, produced by nothing.

constexpr OpName ACCUMULATE = differentiable_op("sushiai.accumulate")

Accumulates a gradient contribution into an existing buffer.

constexpr OpName CAST = op("sushiai.cast")

Copies a value into another element type, out of place.

constexpr OpName MATMUL = differentiable_op("sushiai.matmul")

2-D matrix product, transposition expressed by flags not views.

constexpr OpName BIAS_ADD = differentiable_op("sushiai.bias_add")

Adds a row vector to every row of a matrix.

constexpr OpName ADD = differentiable_op("sushiai.add")

Elementwise sum of two equally shaped tensors.

constexpr OpName SUB = differentiable_op("sushiai.sub")

Elementwise difference of two equally shaped tensors.

constexpr OpName MUL = differentiable_op("sushiai.mul")

Elementwise product of two equally shaped tensors.

constexpr OpName SCALE = differentiable_op("sushiai.scale")

Multiplies every element by a host-side constant.

constexpr OpName DYNAMIC_SCALE = op("sushiai.dynamic_scale")

Multiplies every element by a device-resident scalar.

constexpr OpName EXP = differentiable_op("sushiai.exp")

Elementwise natural exponential, out of place.

constexpr OpName ZERO = op("sushiai.zero")

An all-zero tensor, read by nothing upstream.

constexpr OpName DIV = differentiable_op("sushiai.div")

Divides two equally shaped tensors elementwise: a / b.

constexpr OpName MINIMUM = differentiable_op("sushiai.minimum")

Takes the elementwise smaller of two equally shaped tensors.

constexpr OpName MAXIMUM = differentiable_op("sushiai.maximum")

Takes the elementwise larger of two equally shaped tensors.

constexpr OpName ABS = differentiable_op("sushiai.abs")

Takes the elementwise absolute value, out of place.

constexpr OpName ATAN = differentiable_op("sushiai.atan")

Takes the elementwise arctangent, out of place.

constexpr OpName SOFTPLUS = differentiable_op("sushiai.softplus")

Computes log(1 + exp(x)) elementwise, out of place.

constexpr OpName LESS_EQUAL = op("sushiai.less_equal")

Compares elementwise: 1 where a <= b and 0 elsewhere, in the operands' dtype.

constexpr OpName SELECT = differentiable_op("sushiai.select")

Selects elementwise: a where the condition is non-zero, b elsewhere.

constexpr OpName RELU = differentiable_op("sushiai.relu")

Rectified linear unit.

constexpr OpName GELU = differentiable_op("sushiai.gelu")

Gaussian error linear unit, tanh approximation.

constexpr OpName SIGMOID = differentiable_op("sushiai.sigmoid")

Logistic sigmoid.

constexpr OpName TANH = differentiable_op("sushiai.tanh")

Hyperbolic tangent.

constexpr OpName SILU = differentiable_op("sushiai.silu")

Computes the sigmoid linear unit, x * sigmoid(x).

constexpr OpName RELU_BACKWARD = op("sushiai.relu_backward")

Gradient of RELU; second operand is the forward output.

constexpr OpName GELU_BACKWARD = op("sushiai.gelu_backward")

Gradient of GELU; second operand is the forward input.

constexpr OpName SIGMOID_BACKWARD = op("sushiai.sigmoid_backward")

Gradient of SIGMOID; second operand is the forward output.

constexpr OpName TANH_BACKWARD = op("sushiai.tanh_backward")

Gradient of TANH; second operand is the forward output.

constexpr OpName SILU_BACKWARD = op("sushiai.silu_backward")

Computes the gradient of SILU; second operand is the forward input.

constexpr OpName RESHAPE = differentiable_op("sushiai.reshape")

Reinterprets extents without moving data.

constexpr OpName LAYOUT_CONVERT = differentiable_op("sushiai.layout_convert")

The only node permitted to change a value's TensorLayout (CV-4).

constexpr OpName IM2COL = differentiable_op("sushiai.im2col")

Gathers each window into a row of a column matrix.

constexpr OpName COL2IM = differentiable_op("sushiai.col2im")

The adjoint of IM2COL; a gather, never an atomic scatter (CV-6).

constexpr OpName MAXPOOL2D = differentiable_op("sushiai.maxpool2d")

Windowed maximum; also writes the argmax indices.

constexpr OpName MAXPOOL2D_BACKWARD = op("sushiai.maxpool2d_backward")

Routes each gradient back through the recorded argmax.

constexpr OpName AVGPOOL2D = differentiable_op("sushiai.avgpool2d")

Windowed mean, dividing by the window area (count_include_pad).

constexpr OpName AVGPOOL2D_BACKWARD = op("sushiai.avgpool2d_backward")

Spreads each gradient evenly over the window it came from.

constexpr OpName SPATIAL_MEAN = differentiable_op("sushiai.spatial_mean")

Global average over H and W: [N,H,W,C] -> [N,C].

constexpr OpName SPATIAL_MEAN_BACKWARD = op("sushiai.spatial_mean_backward")

Broadcasts an [N,C] gradient over [N,H,W,C], scaled by 1/(H*W).

constexpr OpName UPSAMPLE2D = differentiable_op("sushiai.upsample2d")

Nearest-neighbour upsample by an integer scale.

constexpr OpName UPSAMPLE2D_BACKWARD = op("sushiai.upsample2d_backward")

Sums each source pixel's replicas back together.

constexpr OpName CONCAT_CHANNELS = differentiable_op("sushiai.concat_channels")

Joins several values along the channel axis.

constexpr OpName SLICE_CHANNELS = differentiable_op("sushiai.slice_channels")

Extracts a channel range.

constexpr OpName BATCHNORM2D = differentiable_op("sushiai.batchnorm2d")

Batch normalisation's forward pass over an NHWC value.

constexpr OpName BATCHNORM2D_BACKWARD = op("sushiai.batchnorm2d_backward")

Batch normalisation's gradient, for input, gamma and beta.

constexpr OpName CHANNEL_AFFINE = differentiable_op("sushiai.channel_affine")

Per-channel scale and shift: y = gamma * x + beta.

constexpr OpName BATCHNORM_UPDATE_STATS = op("sushiai.batchnorm_update_stats")

Updates running statistics for batch normalisation.

constexpr OpName LEAKY_RELU = differentiable_op("sushiai.leaky_relu")

Leaky ReLU; its slope is attribute slot 0.

constexpr OpName LEAKY_RELU_BACKWARD = op("sushiai.leaky_relu_backward")

Gradient of LEAKY_RELU; second operand is the forward output.

constexpr OpName LOG_SOFTMAX = op("sushiai.log_softmax")

Row-wise log-softmax, computed max-shifted for stability.

constexpr OpName NLL_LOSS = differentiable_op("sushiai.nll_loss")

Mean negative log-likelihood against integer-valued targets.

constexpr OpName SUM_ROWS = differentiable_op("sushiai.sum_rows")

Sums a matrix along its rows, yielding one value per column.

constexpr OpName ADAMW_UPDATE = op("sushiai.optim.adamw")

One parameter's whole AdamW update, fused into one kernel.

constexpr OpName SGD_UPDATE = op("sushiai.optim.sgd")

One parameter's whole SGD update, fused into one kernel.

constexpr OpName GEMM_BIAS = op("sushiai.fused.gemm_bias")

GEMM with a bias broadcast folded into its store.

constexpr OpName GEMM_BIAS_RELU = op("sushiai.fused.gemm_bias_relu")

GEMM with a bias broadcast and a ReLU folded into its store.

constexpr OpName GEMM_BIAS_SIGMOID = op("sushiai.fused.gemm_bias_sigmoid")

GEMM with a bias broadcast and a sigmoid folded into its store.

constexpr OpName GEMM_BIAS_TANH = op("sushiai.fused.gemm_bias_tanh")

GEMM with a bias broadcast and a tanh folded into its store.

constexpr OpName GEMM_BIAS_GELU = op("sushiai.fused.gemm_bias_gelu")

Fused GEMM with bias broadcast and GELU, spilling pre-activation.

constexpr OpName EXP_SUB_SCALE = op("sushiai.fused.exp_sub_scale")

One pass computing out = (exp(a) - b) * alpha.

The cross-entropy gradient's whole elementwise tail: five recorded operations - a copy, an exp, a subtraction, a copy and a scal - become one, with both temporaries living in registers.

constexpr std::array ALL = { LEAF, ACCUMULATE, CAST, MATMUL, BIAS_ADD, ADD, SUB, MUL, SCALE, DYNAMIC_SCALE, EXP, ZERO, DIV, MINIMUM, MAXIMUM, ABS, ATAN, SOFTPLUS, LESS_EQUAL, SELECT, RELU, GELU, SIGMOID, TANH, LEAKY_RELU, SILU, RELU_BACKWARD, GELU_BACKWARD, SIGMOID_BACKWARD, TANH_BACKWARD, LEAKY_RELU_BACKWARD, SILU_BACKWARD, LOG_SOFTMAX, NLL_LOSS, SUM_ROWS, ADAMW_UPDATE, SGD_UPDATE, GEMM_BIAS, GEMM_BIAS_RELU, GEMM_BIAS_SIGMOID, GEMM_BIAS_TANH, GEMM_BIAS_GELU, EXP_SUB_SCALE, RESHAPE, LAYOUT_CONVERT, IM2COL, COL2IM, MAXPOOL2D, MAXPOOL2D_BACKWARD, AVGPOOL2D, AVGPOOL2D_BACKWARD, SPATIAL_MEAN, SPATIAL_MEAN_BACKWARD, UPSAMPLE2D, UPSAMPLE2D_BACKWARD, CONCAT_CHANNELS, SLICE_CHANNELS, BATCHNORM2D, BATCHNORM2D_BACKWARD, CHANNEL_AFFINE, BATCHNORM_UPDATE_STATS, }

Every identity declared above, in one enumerable table.

The uniqueness proof reads this, and so does the diagnostic that turns an OpID back into a printable name.

constexpr std::array< Activation, 6 > ACTIVATIONS = {{ {RELU.id, RELU_BACKWARD.id, GradientSource::OUTPUT}, {GELU.id, GELU_BACKWARD.id, GradientSource::INPUT}, {SIGMOID.id, SIGMOID_BACKWARD.id, GradientSource::OUTPUT}, {TANH.id, TANH_BACKWARD.id, GradientSource::OUTPUT}, {LEAKY_RELU.id, LEAKY_RELU_BACKWARD.id, GradientSource::OUTPUT}, {SILU.id, SILU_BACKWARD.id, GradientSource::INPUT}, }}

Every activation this library differentiates, in one table.

Functions

constexpr OpName op(std::string_view n)

Builds an OpName, hashing the name at compile time.

constexpr rather than consteval because consteval is C++20 and this library is C++17. Nothing is lost: every caller below initialises an inline constexpr variable, which forces constant evaluation anyway, so no name is hashed at run time.

constexpr OpName differentiable_op(std::string_view n)

Builds an OpName that declares itself differentiable.

The claim is checked, not trusted: autograd/rules/all.hpp's static_assert walks ALL and fails the build if any entry built this way has no matching backward rule.

constexpr bool all_ids_unique(SushiRuntime::span< const OpName > table)

Reports whether every identity in table hashes distinctly.

Parameters

table

The identities to check.

Returns

True when no two entries share an OpID.

constexpr std::string_view name_of(OpID id) noexcept

The declared name behind an OpID, for diagnostics and reports.

Parameters

id

The identity to look up.

Returns

The matching name, or "sushiai.unknown" when id is not ours.

constexpr const Activation * activation_of(OpID forward) noexcept

The activation entry for a forward operation.

Parameters

forward

The forward operation's identity.

Returns

A pointer into ACTIVATIONS, or null when forward is not an activation.

constexpr bool is_activation_backward(OpID op) noexcept

Reports whether op is the gradient operation of a row in ACTIVATIONS.

constexpr bool activation_survives_fusion(OpID forward) noexcept

Reports whether folding forward needs no auxiliary spill.

Parameters

forward

The forward activation identity.

Returns

True when folding forward requires no auxiliary spill.