Contents

namespace

SushiAI::Graph

Declared in
include/SushiAI/graph/arena.hpp

Contains

Enumerations

enum class FusedForm

The catalog: every fused kernel this build actually contains.

Not a set of flags to combine - each entry is a separate template instantiation that exists in the binary, and a combination not listed here does not exist. See the file comment.

NONE

Not a fused node.

GEMM_BIAS

MATMUL then BIAS_ADD, as one GEMM with a bias epilogue.

GEMM_BIAS_RELU

MATMUL, BIAS_ADD then RELU, as one GEMM.

GEMM_BIAS_SIGMOID

MATMUL, BIAS_ADD then SIGMOID, as one GEMM.

GEMM_BIAS_TANH

MATMUL, BIAS_ADD then TANH, as one GEMM.

GEMM_BIAS_GELU

MATMUL, BIAS_ADD then GELU, spilling the pre-activation.

EXP_SUB_SCALE

EXP, SUB then SCALE, as one elementwise pass.

enum class FusionDecline

Why a structurally matching chain was not fused.

Only reasons a reader would want to act on. A chain that simply is not in the graph is not a decline and is not recorded - the report would be nothing but noise.

NONE

Placeholder; never recorded.

VALUE_HAS_OTHER_CONSUMERS

An interior value feeds a consumer outside the chain.

VALUE_IS_OBSERVED

An interior value is read from outside the graph.

GRADIENT_NEEDS_PRE_ACTIVATION

Indicates the activation gradient differentiates at its input.

SPILL_READ_TOO_EARLY

Indicates a spilled value is read before the fused node writes it.

NO_CATALOG_ENTRY

The activation could be folded and the catalog has no entry.

DISABLED

The caller switched this form off through FusionOptions.

enum class Lifetime

Defines the lifetime class and owning arena category for a value.

PARAMETER

A learnable weight.

GRADIENT

A gradient of a parameter.

OPTIMIZER_STATE

Optimizer state such as a momentum or variance buffer.

INPUT

A minibatch input or target.

ACTIVATION

An intermediate.

SCRATCH

A transient the planner may overlap freely with any other scratch.

enum class TensorLayout

Specifies the dimension ordering for a rank-4 tensor value.

NHWC

Channels innermost - the single native layout (CV-3).

NCHW

Channels before the spatial axes; a boundary only.

Typedefs

using SushiAI::Graph::ValueId = std::uint32_t

Identifies a value in a Graph; an index, not a pointer.

using SushiAI::Graph::NodeId = std::uint32_t

Identifies a node in a Graph; an index, not a pointer.

using SushiAI::Graph::ArenaId = std::uint32_t

Identifies an arena in a MemoryPlan; an index, not a pointer.

Variables

constexpr ValueId NO_VALUE = ~ValueId{0}

The absent-value sentinel, used for "no producer" and "no gradient".

constexpr NodeId NO_NODE = ~NodeId{0}

The absent-node sentinel, marking a value that no node produces.

constexpr ArenaId NO_ARENA = ~ArenaId{0}

The "this value is not packed" sentinel; it gets its own allocation.

Functions

constexpr int64_t output_extent(int64_t in, int64_t pad_before, int64_t pad_after, int64_t dilation, int64_t kernel, int64_t stride) noexcept

Computes the output extent along one convolution axis.

Parameters

stride

Step between windows; must be positive.

kernel

Number of taps; must be positive.

dilation

Spacing between taps; must be positive.

Returns

The output extent, or 0 when the window does not fit.

SpatialExtents window_extents(int64_t height, int64_t width, const ConvGeometry &g)

Computes the extents a window yields over a height by width input.

Parameters

height

Input height, padded above and below by the geometry's top pad.

width

Input width, padded left and right by the geometry's left pad.

g

Needs a positive kernel, stride and dilation, a non-negative padding and a window the padded input holds; anything else throws.

Returns

The number of window rows and of window columns, both positive.

SpatialExtents upsampled_extents(int64_t height, int64_t width, int64_t scale_h, int64_t scale_w)

Computes the extents UPSAMPLE2D yields from a height by width input.

Parameters

scale_h

Vertical replication factor; below one throws.

scale_w

Horizontal replication factor; below one throws.

Returns

The scaled height and width; a product beyond 64 bits throws.

bool is_legal_leaky_relu_slope(double slope) noexcept

Tells whether LEAKY_RELU and its gradient accept slope.

Returns

True when slope lies strictly inside (0, 1), compared in double.

CastReport insert_casts(Graph &graph, const CastOptions &options={})

Rewrites graph so every node's operands match its computed dtype.

Parameters

graph

The graph to rewrite in place.

options

The policy and preserved values.

Returns

A summary of inserted conversions.

Exceptions

Error

If topologically unordered or conversion fails.

std::vector< DTypeMismatch > find_dtype_mismatches(const Graph &graph, const IPrecisionPolicy &policy)

Finds every node whose operands or results disagree with policy.

Parameters

graph

The graph to check.

policy

The precision policy to conform to.

Returns

Disagreements found, or empty if invariant holds.

std::string describe(const DTypeMismatch &mismatch)

A one-line description of mismatch, for a test failure or a report.

constexpr std::string_view fused_form_name(FusedForm form) noexcept

A human-readable name for form, for diagnostics.

FusedForm fused_form_from_name(std::string_view name) noexcept

Returns the catalog entry name refers to.

Parameters

name

Name as fused_form_name spells it.

Returns

Matching form, or FusedForm::NONE when unmatched.

OpID fused_form_op(FusedForm form) noexcept

The operation identity a node of form carries.

Parameters

form

The catalog entry.

Returns

Its Ops identity, or Ops::LEAF.id for FusedForm::NONE.

FusedForm fused_form_of(OpID op) noexcept

The catalog entry behind an operation identity.

Parameters

op

The identity to look up.

Returns

The form, or FusedForm::NONE when op is not a fused one.

constexpr bool is_gemm_epilogue(FusedForm form) noexcept

Returns whether form is a GEMM-epilogue entry.

Parameters

form

The catalog entry.

Returns

True for the GEMM_BIAS family.

OpID fused_form_activation(FusedForm form) noexcept

The activation a GEMM-epilogue form folds in, if any.

Parameters

form

The catalog entry.

Returns

The forward activation's identity, or Ops::LEAF.id when the form folds no activation.

bool fused_form_spills(FusedForm form) noexcept

Returns whether a node of form writes an auxiliary second value.

Parameters

form

The catalog entry.

Returns

True when a node of form declares two outputs.

bool catalog_folds_activation(OpID activation) noexcept

Returns whether any catalog entry folds activation into a GEMM.

Parameters

activation

A forward activation identity.

Returns

True when some entry's chain ends in activation.

constexpr std::string_view fusion_decline_reason(FusionDecline reason) noexcept

A sentence explaining reason, for the report.

FusionReport fuse(Graph &graph, const FusionOptions &options={})

Partitions graph into maximal fusable subtrees and rewrites them.

Parameters

graph

The graph to rewrite in place.

options

What the pass is allowed to do.

Returns

What it did and what it refused.

Exceptions

Error

If the graph is not in topological insertion order.

constexpr const char * lifetime_name(Lifetime lifetime) noexcept

A human-readable name for lifetime, for diagnostics.

constexpr bool is_persistent(Lifetime lifetime) noexcept

True when a value of this class outlives a single step.

Persistent values are allocated once and never overlapped; transient ones are what the liveness pass actually packs.

bool narrows_under_mixed_precision(OpID op) noexcept

Checks whether an operation is narrowed under mixed precision.

Parameters

op

Operation identifier to check.

Returns

True if op is narrowed; false if preserved at declared precision.

void rematerialise_im2col(Graph &graph)

Recomputes every stored IM2COL result at its later consumers.

Inserts gathers before later readers to shorten column matrix live ranges.

Parameters

graph

The graph to rewrite in place.