pulsatrix
Loading...
Searching...
No Matches
pulsatrix::Autograd Class Reference

Computes gradients by walking a ComputationGraph in reverse topological order. More...

#include <autograd.hpp>

Public Types

using BackwardFn = std::function< Tensor(const Tensor &grad_output)>
 A function computing the gradient w.r.t. a node's single input, given the gradient w.r.t. its output.
 

Public Member Functions

void register_backward (NodeId id, BackwardFn fn)
 Registers how to compute this node's input gradient from its output gradient.
 
void backward (const ComputationGraph &graph, NodeId root, const Tensor &grad_output)
 Runs backward from a single root, seeding its gradient with grad_output.
 
void backward (const ComputationGraph &graph, std::vector< std::pair< NodeId, Tensor > > seeds)
 Runs backward from multiple seeded roots in one pass, so gradients that converge on a shared ancestor accumulate correctly within a single call.
 
const Tensor & gradient (NodeId id) const
 Retrieves the accumulated gradient for a node after backward() has run.
 
bool has_gradient (NodeId id) const
 Whether a gradient was accumulated for this node during the last backward() call.
 

Detailed Description

Computes gradients by walking a ComputationGraph in reverse topological order.

Note
Phase 0 has no Module/layer abstraction yet – backward functions are supplied directly by the caller via register_backward() rather than being derived from an op-type dispatch table. Phase 1's Module::forward() is expected to call register_backward() when it builds graph nodes; this class doesn't need to know anything about Linear/Conv/etc. to do its job.
Scoped to single-parent nodes: a registered backward function takes the gradient w.r.t. this node's output and returns the gradient w.r.t. its one input. Nodes that combine multiple parents (branching/residual composition) are explicitly out of scope here – charter Part 2 SS6 defers that composition to Module-level design (a TransformerBlock "knows its own composition"), which is Phase 1+/2 work.
backward() is not noexcept: Tensor/DeviceBackend allocation can throw (see cpp_style_guide's error-handling table), so this gives only the basic exception safety guarantee, not nothrow – campaign Decision Point 3, resolved here rather than assumed.

Member Typedef Documentation

◆ BackwardFn

using pulsatrix::Autograd::BackwardFn = std::function<Tensor(const Tensor& grad_output)>

A function computing the gradient w.r.t. a node's single input, given the gradient w.r.t. its output.

Member Function Documentation

◆ backward() [1/2]

void pulsatrix::Autograd::backward ( const ComputationGraph &  graph,
NodeId  root,
const Tensor &  grad_output 
)

Runs backward from a single root, seeding its gradient with grad_output.

Parameters
graphThe graph to walk. Must outlive this call.
rootNode to seed.
grad_outputGradient to seed at root.

◆ backward() [2/2]

void pulsatrix::Autograd::backward ( const ComputationGraph &  graph,
std::vector< std::pair< NodeId, Tensor > >  seeds 
)

Runs backward from multiple seeded roots in one pass, so gradients that converge on a shared ancestor accumulate correctly within a single call.

Parameters
graphThe graph to walk.
seeds(node id, seed gradient) pairs to start from.

◆ gradient()

const Tensor & pulsatrix::Autograd::gradient ( NodeId  id) const

Retrieves the accumulated gradient for a node after backward() has run.

Parameters
idNode id. Must satisfy has_gradient(id).
Returns
The accumulated gradient.

◆ has_gradient()

bool pulsatrix::Autograd::has_gradient ( NodeId  id) const
inline

Whether a gradient was accumulated for this node during the last backward() call.

◆ register_backward()

void pulsatrix::Autograd::register_backward ( NodeId  id,
BackwardFn  fn 
)

Registers how to compute this node's input gradient from its output gradient.

Parameters
idNode to register a backward function for. Must have at most one parent (see class-level
Note
on single-parent scope).
Parameters
fnThe backward function.

The documentation for this class was generated from the following file: