Skip to main content

ComputationGraph

Struct ComputationGraph 

Source
pub struct ComputationGraph { /* private fields */ }
Expand description

Computation graph that records operations for backward pass.

The graph uses a tape-based approach where operations are recorded in order during the forward pass, then gradients are computed in reverse order during the backward pass.

§Thread Safety

Each thread has its own computation graph (via thread_local storage in the parent module). This avoids synchronization overhead during single-threaded training.

Implementations§

Source§

impl ComputationGraph

Source

pub fn new() -> ComputationGraph

Create a new empty computation graph.

Source

pub fn clear(&mut self)

Clear all recorded operations.

Source

pub fn register_tensor(&mut self, tensor: Tensor)

Register a tensor that requires gradients.

Source

pub fn record( &mut self, output_id: TensorId, grad_fn: Arc<dyn GradFn>, input_ids: Vec<TensorId>, )

Record an operation to the tape.

Source

pub fn get_tensor(&self, id: TensorId) -> Option<&Tensor>

Get a tensor by ID.

Source

pub fn get_tensor_mut(&mut self, id: TensorId) -> Option<&mut Tensor>

Get a mutable tensor by ID.

Source

pub fn backward(&mut self, output_id: TensorId, grad_output: Tensor)

Compute gradients via backpropagation.

This implements the reverse-mode automatic differentiation algorithm:

  1. Start with grad_output for the output tensor
  2. Iterate through operations in reverse order
  3. For each operation, compute gradients w.r.t. inputs
  4. Accumulate gradients for tensors used multiple times
§Arguments
  • output_id - ID of the tensor to differentiate
  • grad_output - Initial gradient (typically ones for scalar loss)
Source

pub fn len(&self) -> usize

Get the number of recorded operations.

Source

pub fn is_empty(&self) -> bool

Check if the tape is empty.

Source

pub fn get_grad(&self, id: TensorId) -> Option<Tensor>

Get gradient for a tensor by ID (after backward).

Source

pub fn clear_grad(&mut self, id: TensorId)

Clear gradient for a specific tensor.

Trait Implementations§

Source§

impl Default for ComputationGraph

Source§

fn default() -> ComputationGraph

Returns the “default value” for a type. Read more

Auto Trait Implementations§

Blanket Implementations§

Source§

impl<T> Any for T
where T: 'static + ?Sized,

Source§

fn type_id(&self) -> TypeId

Gets the TypeId of self. Read more
Source§

impl<T> Borrow<T> for T
where T: ?Sized,

Source§

fn borrow(&self) -> &T

Immutably borrows from an owned value. Read more
Source§

impl<T> BorrowMut<T> for T
where T: ?Sized,

Source§

fn borrow_mut(&mut self) -> &mut T

Mutably borrows from an owned value. Read more
Source§

impl<ST, DT> CastableFrom<ST, Initialized, Initialized> for DT
where ST: ?Sized, DT: ?Sized,

Source§

impl<ST, DT> CastableFrom<ST, Uninit, Uninit> for DT
where ST: ?Sized, DT: ?Sized,

Source§

impl<T> Downcast<T> for T

Source§

fn downcast(&self) -> &T

Source§

impl<T> From<T> for T

Source§

fn from(t: T) -> T

Returns the argument unchanged.

Source§

impl<T, U> Into<U> for T
where U: From<T>,

Source§

fn into(self) -> U

Calls U::from(self).

That is, this conversion is whatever the implementation of From<T> for U chooses to do.

Source§

impl<T> IntoEither for T

Source§

fn into_either(self, into_left: bool) -> Either<Self, Self>

Converts self into a Left variant of Either<Self, Self> if into_left is true. Converts self into a Right variant of Either<Self, Self> otherwise. Read more
Source§

fn into_either_with<F>(self, into_left: F) -> Either<Self, Self>
where F: FnOnce(&Self) -> bool,

Converts self into a Left variant of Either<Self, Self> if into_left(&self) returns true. Converts self into a Right variant of Either<Self, Self> otherwise. Read more
Source§

impl<T> Pointable for T

Source§

const ALIGN: usize

The alignment of pointer.
Source§

type Init = T

The type for initializers.
Source§

unsafe fn init(init: <T as Pointable>::Init) -> usize

Initializes a with the given initializer. Read more
Source§

unsafe fn deref<'a>(ptr: usize) -> &'a T

Dereferences the given pointer. Read more
Source§

unsafe fn deref_mut<'a>(ptr: usize) -> &'a mut T

Mutably dereferences the given pointer. Read more
Source§

unsafe fn drop(ptr: usize)

Drops the object pointed to by the given pointer. Read more
Source§

impl<T> Read<Exclusive, BecauseExclusive> for T
where T: ?Sized,

Source§

impl<T> Same for T

Source§

type Output = T

Should always be Self
Source§

impl<T, U> TryFrom<U> for T
where U: Into<T>,

Source§

type Error = !

The type returned in the event of a conversion error.
Source§

fn try_from(value: U) -> Result<T, !>

Performs the conversion.
Source§

impl<T, U> TryInto<U> for T
where U: TryFrom<T>,

Source§

type Error = <U as TryFrom<T>>::Error

The type returned in the event of a conversion error.
Source§

fn try_into(self) -> Result<U, <U as TryFrom<T>>::Error>

Performs the conversion.
Source§

impl<T> Upcast<T> for T

Source§

fn upcast(&self) -> Option<&T>

Source§

impl<V, T> VZip<V> for T
where V: MultiLane<T>,

Source§

fn vzip(self) -> V

Source§

impl<T> WasmNotSend for T
where T: Send,

Source§

impl<T> WasmNotSendSync for T

Source§

impl<T> WasmNotSync for T
where T: Sync,