Contents

namespace

SushiAI::Autograd

Declared in
include/SushiAI/autograd/backward.hpp

Contains

Typedefs

using SushiAI::Autograd::ScalarFunction = std::function<double(SushiRuntime::span<const double> inputs)>

Evaluates the scalar under test at a point.

Parameters

inputs

The point, in FLOAT64.

Returns

The scalar value at that point.

using SushiAI::Autograd::GradientFunction = std::function<void(SushiRuntime::span<const double> inputs, SushiRuntime::span<double> gradient)>

Fills gradient with the analytic gradient at inputs.

Parameters

inputs

The point, in FLOAT64.

gradient

The output, exactly as long as inputs.

using SushiAI::Autograd::BackwardRule = void (*)(const GradientTape&)

What every backward rule looks like.

Variables

constexpr std::array< RuleEntry, 6 > ACTIVATION_RULES = {{ {Ops::RELU.id, &activation_backward}, {Ops::GELU.id, &activation_backward}, {Ops::SIGMOID.id, &activation_backward}, {Ops::TANH.id, &activation_backward}, {Ops::LEAKY_RELU.id, &activation_backward}, {Ops::SILU.id, &activation_backward}, }}

Activation operations' backward rules.

constexpr auto RULES = concat_rules( concat_rules( concat_rules(ELEMENTWISE_RULES, LINALG_RULES), concat_rules(ACTIVATION_RULES, LOSS_RULES)), concat_rules( concat_rules(SHAPE_RULES, SPATIAL_RULES), NORM_RULES))

Every backward rule, as one flat compile-time table.

Adding an operation to a domain changes only that domain's file; adding a domain is one line here.

constexpr std::array< RuleEntry, 14 > ELEMENTWISE_RULES = {{ {Ops::ADD.id, &add_backward}, {Ops::SUB.id, &sub_backward}, {Ops::MUL.id, &mul_backward}, {Ops::SCALE.id, &scale_backward}, {Ops::EXP.id, &exp_backward}, {Ops::ACCUMULATE.id, &accumulate_backward}, {Ops::DIV.id, &div_backward}, {Ops::MINIMUM.id, &minimum_backward}, {Ops::MAXIMUM.id, &maximum_backward}, {Ops::ABS.id, &abs_backward}, {Ops::ATAN.id, &atan_backward}, {Ops::SOFTPLUS.id, &softplus_backward}, {Ops::SELECT.id, &select_backward}, {Ops::LESS_EQUAL.id, &less_equal_backward}, }}

Holds backward rule entries for elementwise operations.

constexpr std::array< RuleEntry, 3 > LINALG_RULES = {{ {Ops::MATMUL.id, &matmul_backward}, {Ops::BIAS_ADD.id, &bias_add_backward}, {Ops::SUM_ROWS.id, &sum_rows_backward}, }}

Holds backward rule entries for linear algebra operations.

constexpr std::array< RuleEntry, 1 > LOSS_RULES = {{ {Ops::NLL_LOSS.id, &nll_loss_backward}, }}

Holds backward rule entries for loss operations.

constexpr std::array< RuleEntry, 2 > NORM_RULES = {{ {Ops::BATCHNORM2D.id, &batchnorm2d_backward}, {Ops::CHANNEL_AFFINE.id, &channel_affine_backward}, }}

Holds backward rule entries for normalization operations.

constexpr std::array< RuleEntry, 2 > SHAPE_RULES = {{ {Ops::RESHAPE.id, &reshape_backward}, {Ops::LAYOUT_CONVERT.id, &layout_convert_backward}, }}

Holds backward rule entries for structural shape and layout operations.

constexpr std::array< RuleEntry, 8 > SPATIAL_RULES = {{ {Ops::IM2COL.id, &im2col_backward}, {Ops::COL2IM.id, &col2im_backward}, {Ops::MAXPOOL2D.id, &maxpool2d_backward}, {Ops::AVGPOOL2D.id, &avgpool2d_backward}, {Ops::UPSAMPLE2D.id, &upsample2d_backward}, {Ops::SPATIAL_MEAN.id, &spatial_mean_backward}, {Ops::CONCAT_CHANNELS.id, &concat_channels_backward}, {Ops::SLICE_CHANNELS.id, &slice_channels_backward}, }}

Spatial operations' backward rules.

Functions

bool has_backward_rule(OpID op) noexcept

Returns whether an operation can be differentiated.

Parameters

op

The operation to look up.

Returns

True when a backward rule exists for it.

std::vector< Graph::ValueId > differentiate(Graph::Graph &graph, Graph::ValueId loss, SushiRuntime::span< const Graph::ValueId > wrt, Graph::ValueId seed_factor=Graph::NO_VALUE, Provenance *provenance=nullptr)

Extends graph with reverse-mode gradients of loss.

Parameters

loss

Scalar value being differentiated; must have one element.

seed_factor

Optional scaling scalar, or Graph::NO_VALUE.

provenance

Filled with the origin of every appended node when not null.

Returns

Gradient ValueId per wrt entry, or Graph::NO_VALUE if unreachable.

void push_off_kinks(SushiRuntime::span< double > inputs, const GradCheckOptions &options={})

Moves coordinates off a kink at zero, in place.

Parameters

inputs

The point to adjust.

options

Supplies the step size and the margin.

GradCheckResult check_full_jacobian(const ScalarFunction &value, const GradientFunction &gradient, SushiRuntime::span< double > inputs, const GradCheckOptions &options={})

Compares every coordinate of the gradient against a central difference.

Parameters

inputs

Coordinates to evaluate at; restored on return.

options

Numerical step sizes and comparison tolerance.

Returns

Comparison result across all input coordinates.

GradCheckResult check_directional(const ScalarFunction &value, const GradientFunction &gradient, SushiRuntime::span< double > inputs, std::size_t directions=8, const GradCheckOptions &options={})

Compares the gradient against central differences along random directions.

Parameters

inputs

Coordinates to evaluate at; restored on return.

directions

Number of random unit directions to probe.

options

Numerical step sizes and comparison tolerance.

Returns

Comparison result across all sampled directions.

template <std::size_t N, std::size_t M>
constexpr std::array< RuleEntry, N+M > concat_rules(const std::array< RuleEntry, N > &a, const std::array< RuleEntry, M > &b)

Concatenates two rule tables at compile time.

Parameters

a

The left rule table.

b

The right rule table.

Returns

Combined array of rule entries.

void activation_backward(const GradientTape &tape)

Differentiates activation operations using metadata in Ops::ACTIVATIONS.

Parameters

tape

Traversal context providing forward nodes and gradient emitters.

constexpr bool has_rule(OpID op) noexcept

Whether op has a rule in RULES.

constexpr bool every_differentiable_op_has_a_rule(SushiRuntime::span< const Ops::OpName > table) noexcept

Verifies that every differentiable operation possesses a backward rule.

Parameters

table

Operation metadata table to inspect.

Returns

True when all differentiable operations have registered rules.

void add_backward(const GradientTape &tape)

Differentiates binary addition: d(a+b) = [dy, dy].

void sub_backward(const GradientTape &tape)

Differentiates binary subtraction: d(a-b) = [dy, -dy].

void mul_backward(const GradientTape &tape)

Differentiates binary multiplication: d(a*b) = [dy*b, dy*a].

void scale_backward(const GradientTape &tape)

Differentiates scalar scaling: d(alpha*x) = alpha*dy.

void exp_backward(const GradientTape &tape)

Differentiates elementwise exponential: d(exp x) = dy * exp(x).

void accumulate_backward(const GradientTape &tape)

Propagates accumulation gradients to both operands.

void div_backward(const GradientTape &tape)

Differentiates division: d(a/b) = [dy/b, -(dy/b)*(a/b)].

void minimum_backward(const GradientTape &tape)

Differentiates the minimum: dy goes to a where a <= b, to b elsewhere.

void maximum_backward(const GradientTape &tape)

Differentiates the maximum: dy goes to a where a >= b, to b elsewhere.

void abs_backward(const GradientTape &tape)

Differentiates the absolute value: d|x| = dy * sign(x), zero at x = 0.

void atan_backward(const GradientTape &tape)

Differentiates the arctangent: d(atan x) = dy / (1 + x*x).

void softplus_backward(const GradientTape &tape)

Differentiates softplus at its input: d(softplus x) = dy * sigmoid(x).

void select_backward(const GradientTape &tape)

Differentiates selection: dy goes to a where the condition holds, else to b.

void less_equal_backward(const GradientTape &tape)

Differentiates the comparison a <= b, which contributes to neither operand.

void matmul_backward(const GradientTape &tape)

Differentiates 2D matrix multiplication across transposition configurations.

void bias_add_backward(const GradientTape &tape)

Differentiates row-vector bias addition: identity in X, row-summed in bias.

void sum_rows_backward(const GradientTape &tape)

Differentiates matrix row summation by broadcasting incoming gradients.

void log_softmax_backward(const GradientTape &tape)

Differentiates standalone log-softmax operations.

Parameters

tape

Traversal context providing forward nodes and gradient emitters.

void nll_loss_backward(const GradientTape &tape)

Differentiates negative log-likelihood loss, fusing adjacent log-softmax if present.

Parameters

tape

Traversal context providing forward nodes and gradient emitters.

void batchnorm2d_backward(const GradientTape &tape)

Differentiates 2D batch normalization across input, scale, and shift.

Parameters

tape

Traversal context providing forward nodes and gradient emitters.

void channel_affine_backward(const GradientTape &tape)

Differentiates channel affine transformation with per-channel parameter sums.

Parameters

tape

Traversal context providing forward nodes and gradient emitters.

void reshape_backward(const GradientTape &tape)

Differentiates reshape operations by restoring operand shape on incoming gradients.

void layout_convert_backward(const GradientTape &tape)

Differentiates layout conversions by restoring operand layout on incoming gradients.

void im2col_backward(const GradientTape &tape)

IM2COL's gradient: col2im, scattering each row back onto the pixels its window read, summing where windows overlapped.

Parameters

tape

The differentiation state for this node.

void col2im_backward(const GradientTape &tape)

COL2IM's gradient: im2col with the forward's own geometry, the adjoint of the adjoint.

Parameters

tape

The differentiation state for this node.

void maxpool2d_backward(const GradientTape &tape)

MAXPOOL2D's gradient: routes each gradient through the argmax indices its own forward wrote.

Parameters

tape

The differentiation state for this node.

void avgpool2d_backward(const GradientTape &tape)

AVGPOOL2D's gradient: spreads each gradient evenly over the window it came from, dividing by the constant window area.

Parameters

tape

The differentiation state for this node.

void upsample2d_backward(const GradientTape &tape)

UPSAMPLE2D's gradient: sums each source pixel's replicas.

Parameters

tape

The differentiation state for this node.

void spatial_mean_backward(const GradientTape &tape)

SPATIAL_MEAN's gradient: broadcasts the [N,C] gradient back over [N,H,W,C], scaled by 1/(H*W).

Parameters

tape

The differentiation state for this node.

void concat_channels_backward(const GradientTape &tape)

CONCAT_CHANNELS's gradient: each operand receives the slice of the incoming gradient at that operand's own running channel offset.

Parameters

tape

The differentiation state for this node.

void slice_channels_backward(const GradientTape &tape)

SLICE_CHANNELS's gradient: a scatter into zeros, expressed as a concat of the zero ranges before and after the slice with the incoming gradient itself.

Parameters

tape

The differentiation state for this node.