Contents

class

SushiAI::NN::Parameter

Holds a learnable tensor's metadata and graph value identifier.

Declared in
include/SushiAI/nn/module.hpp

Public member functions

Parameter()=default
Graph::ValueId declare(const Graph::GraphBuilder &builder, Shape shape, ParameterRole role, int64_t fan_in, int64_t fan_out, std::string_view debug_name, DType dtype=DType::FLOAT32)

Declares this parameter into the graph idempotently.

Parameters

builder

Graph builder where the parameter node is registered.

shape

Extents of the parameter tensor.

role

Role indicating whether parameter is weight or bias.

debug_name

Name that must outlive the graph builder.

Returns

ValueId assigned to this parameter.

bool declared() const noexcept

Whether forward() has declared this parameter yet.

Graph::ValueId id() const noexcept

The parameter's value id, or Graph::NO_VALUE before declaration.

const Shape & shape() const noexcept

The parameter's extents.

DType dtype() const noexcept

The parameter's element type.

ParameterRole role() const noexcept

Whether this is a weight or a bias.

int64_t fan_in() const noexcept

Inputs feeding one output unit.

int64_t fan_out() const noexcept

Outputs one input unit feeds.