Contents

class

SushiAI::Autograd::GradientTape

Holds traversal state and emitter interfaces for backward rules.

Declared in
include/SushiAI/autograd/tape.hpp

Public member functions

Graph::Graph & graph() const noexcept

The graph being extended.

Rules emit through builder().

const Graph::GraphBuilder & builder() const noexcept

The emitter every rule uses; it shape-checks what a rule writes.

const Graph::Node & node() const

The node currently being differentiated.

Graph::NodeId node_id() const noexcept

The id of the node currently being differentiated.

Graph::ValueId output_grad(std::size_t index) const

Returns the gradient ValueId of the node's result.

Parameters

index

Output index on the current node.

Returns

Gradient ValueId, or Graph::NO_VALUE if unseeded or inactive.

bool output_is_root(std::size_t index) const

Whether the node's index-th result is the differentiation root.

bool input_requires_grad(std::size_t index) const

Whether a gradient should be produced for the node's index-th operand.

Parameters

index

Which operand.

Returns

True when that operand is on a path from something that requires one.

void contribute(Graph::ValueId primal, Graph::ValueId gradient) const

Registers a gradient contribution to a forward primal value.

Parameters

primal

Forward value receiving the gradient.

gradient

Contribution value to accumulate.

void contribute_input(std::size_t index, Graph::ValueId gradient) const

Convenience: contribute to the current node's index-th operand.

void skip(Graph::NodeId node) const

Marks an upstream node as already handled by the current rule.

Parameters

node

Node identifier to skip during reverse traversal.

Graph::ValueId seed_factor() const noexcept

Returns the optional device-resident scaling factor for the root gradient.

Returns

Single-element ValueId, or Graph::NO_VALUE if unscaled.

Graph::ValueId apply_seed_factor(Graph::ValueId gradient) const

Multiplies gradient by seed_factor if present, returning the result.

Parameters

gradient

Root gradient to scale.

Returns

Scaled gradient ValueId, or gradient unchanged if unscaled.

GradientTape(Graph::GraphBuilder builder, Differentiator &owner, Graph::ValueId seed_factor=Graph::NO_VALUE) noexcept

Internal: constructs a tape over a differentiation run.

void retarget(Graph::NodeId node) noexcept

Internal: points the tape at the next node.