pulsatrix
Loading...
Searching...
No Matches
computation_graph.hpp
Go to the documentation of this file.
1
6#pragma once
7
8#include <memory>
9#include <optional>
10#include <string>
11#include <vector>
12
13#include "pulsatrix/node.hpp"
14#include "pulsatrix/op_type.hpp"
15#include "pulsatrix/shape.hpp"
16
17namespace pulsatrix {
18
28public:
39 NodeId add_node(OpType op_type, Shape shape, std::optional<std::string> label = std::nullopt,
40 std::vector<NodeId> parent_ids = {});
41
47 [[nodiscard]] const Node& node(NodeId id) const;
48
50 [[nodiscard]] size_t node_count() const { return nodes_.size(); }
51
59 [[nodiscard]] std::vector<NodeId> nodes_by_op_type(OpType op_type) const;
60
70 [[nodiscard]] std::vector<NodeId> topological_order() const;
71
72private:
73 std::vector<std::unique_ptr<Node>> nodes_;
74};
75
76} // namespace pulsatrix
Owns every Node in a computation graph and exposes read access for graph-walking code (autograd's bac...
Definition computation_graph.hpp:27
size_t node_count() const
Number of nodes currently in the graph.
Definition computation_graph.hpp:50
NodeId add_node(OpType op_type, Shape shape, std::optional< std::string > label=std::nullopt, std::vector< NodeId > parent_ids={})
Adds a node to the graph and wires it to its parents.
std::vector< NodeId > nodes_by_op_type(OpType op_type) const
Finds every node with the given op type.
std::vector< NodeId > topological_order() const
Returns every node id in a valid topological order (every node after all its parents).
const Node & node(NodeId id) const
Looks up a node by id.
A single computation graph node. Owned exclusively by its ComputationGraph (see computation_graph....
Definition node.hpp:29
An N-dimensional shape. A plain aggregate of dimensions with no invariant beyond "non-negative dimens...
Definition shape.hpp:24
Definition acquisition_functions.hpp:16
size_t NodeId
Stable identifier for a Node within its owning ComputationGraph.
Definition node.hpp:18
OpType
The op-type tag a Node carries. Charter Part 2 §3: nodes are tagged by a small closed set of op types...
Definition op_type.hpp:19
Computation graph node – op type, shape, optional label, parent/child edges.
Closed set of operation categories every graph Node is tagged with.
Tensor dimension arithmetic – rank, element count, per-dimension access.