pulsatrix
Loading...
Searching...
No Matches
autograd.hpp
Go to the documentation of this file.
1
6#pragma once
7
8#include <functional>
9#include <unordered_map>
10#include <utility>
11#include <vector>
12
14#include "pulsatrix/node.hpp"
15#include "pulsatrix/tensor.hpp"
16
17namespace pulsatrix {
18
36class Autograd {
37public:
39 using BackwardFn = std::function<Tensor(const Tensor& grad_output)>;
40
48
55 void backward(const ComputationGraph& graph, NodeId root, const Tensor& grad_output);
56
63 void backward(const ComputationGraph& graph, std::vector<std::pair<NodeId, Tensor>> seeds);
64
70 [[nodiscard]] const Tensor& gradient(NodeId id) const;
71
73 [[nodiscard]] bool has_gradient(NodeId id) const { return gradients_.find(id) != gradients_.end(); }
74
75private:
76 std::unordered_map<NodeId, BackwardFn> backward_fns_;
77 std::unordered_map<NodeId, Tensor> gradients_;
78};
79
80} // namespace pulsatrix
Computes gradients by walking a ComputationGraph in reverse topological order.
Definition autograd.hpp:36
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...
std::function< Tensor(const Tensor &grad_output)> BackwardFn
A function computing the gradient w.r.t. a node's single input, given the gradient w....
Definition autograd.hpp:39
bool has_gradient(NodeId id) const
Whether a gradient was accumulated for this node during the last backward() call.
Definition autograd.hpp:73
const Tensor & gradient(NodeId id) const
Retrieves the accumulated gradient for a node after backward() has run.
void register_backward(NodeId id, BackwardFn fn)
Registers how to compute this node's input gradient from its output gradient.
Owns every Node in a computation graph and exposes read access for graph-walking code (autograd's bac...
Definition computation_graph.hpp:27
N-dimensional tensor. Owns its data buffer exclusively; a DeviceBackend* is injected (not owned) – th...
Definition tensor.hpp:29
Owns and exposes graph structure – the interpretability substrate every explainer (Phase 2+) walks.
Definition acquisition_functions.hpp:16
size_t NodeId
Stable identifier for a Node within its owning ComputationGraph.
Definition node.hpp:18
Computation graph node – op type, shape, optional label, parent/child edges.
N-dimensional tensor – owns a buffer via DeviceBackend*, RAII (Rule of Five).