class
SushiAI::Autograd::GradientTape
Holds traversal state and emitter interfaces for backward rules.
- Declared in
include/SushiAI/autograd/tape.hpp
Public member functions
const Graph::GraphBuilder & builder() const noexceptThe emitter every rule uses; it shape-checks what a rule writes.
const Graph::Node & node() constThe node currently being differentiated.
Graph::NodeId node_id() const noexceptThe id of the node currently being differentiated.
Graph::ValueId output_grad(std::size_t index) constReturns the gradient ValueId of the node's result.
Parameters
indexOutput index on the current node.
Returns
Gradient ValueId, or Graph::NO_VALUE if unseeded or inactive.
bool output_is_root(std::size_t index) constWhether the node's index-th result is the differentiation root.
bool input_requires_grad(std::size_t index) constWhether a gradient should be produced for the node's index-th operand.
Parameters
indexWhich operand.
Returns
True when that operand is on a path from something that requires one.
void contribute(Graph::ValueId primal, Graph::ValueId gradient) constRegisters a gradient contribution to a forward primal value.
Parameters
primalForward value receiving the gradient.
gradientContribution value to accumulate.
void contribute_input(std::size_t index, Graph::ValueId gradient) constConvenience: contribute to the current node's index-th operand.
void skip(Graph::NodeId node) constMarks an upstream node as already handled by the current rule.
Parameters
nodeNode identifier to skip during reverse traversal.
Graph::ValueId seed_factor() const noexceptReturns 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) constMultiplies gradient by seed_factor if present, returning the result.
Parameters
gradientRoot 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) noexceptInternal: constructs a tape over a differentiation run.
void retarget(Graph::NodeId node) noexceptInternal: points the tape at the next node.

