namespace
SushiAI::Graph
- Declared in
include/SushiAI/graph/arena.hpp
Contains
SushiAI::Graph::ArenaSushiAI::Graph::ArenaSetSushiAI::Graph::AttributesSushiAI::Graph::BatchNormGradsSushiAI::Graph::BatchNormResultSushiAI::Graph::CastOptionsSushiAI::Graph::CastReportSushiAI::Graph::CompiledStepSushiAI::Graph::ConvGeometrySushiAI::Graph::DeclinedFusionSushiAI::Graph::DeviceFusionSelectorSushiAI::Graph::DTypeMismatchSushiAI::Graph::FusedSubtreeSushiAI::Graph::FusionChoiceSushiAI::Graph::FusionOptionsSushiAI::Graph::FusionQuerySushiAI::Graph::FusionReportSushiAI::Graph::FusionSelectionSushiAI::Graph::GraphSushiAI::Graph::GraphBuilderSushiAI::Graph::IFusionSelectorSushiAI::Graph::InsertedCastSushiAI::Graph::IPrecisionPolicySushiAI::Graph::IStepRecorderSushiAI::Graph::LoweringSushiAI::Graph::MemoryPlanSushiAI::Graph::MixedPrecisionPolicySushiAI::Graph::NodeSushiAI::Graph::PlacementSushiAI::Graph::PlanOptionsSushiAI::Graph::PoolResultSushiAI::Graph::PrecisionChoiceSushiAI::Graph::PrecisionQuerySushiAI::Graph::SpatialExtentsSushiAI::Graph::TensorPoolSushiAI::Graph::UniformPrecisionPolicySushiAI::Graph::Value
Enumerations
enum class FusedFormThe 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.
NONENot a fused node.
GEMM_BIASMATMULthenBIAS_ADD, as one GEMM with a bias epilogue.GEMM_BIAS_RELUMATMUL,BIAS_ADDthenRELU, as one GEMM.GEMM_BIAS_SIGMOIDMATMUL,BIAS_ADDthenSIGMOID, as one GEMM.GEMM_BIAS_TANHMATMUL,BIAS_ADDthenTANH, as one GEMM.GEMM_BIAS_GELUMATMUL,BIAS_ADDthenGELU, spilling the pre-activation.EXP_SUB_SCALEEXP,SUBthenSCALE, as one elementwise pass.
enum class FusionDeclineWhy 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.
NONEPlaceholder; never recorded.
VALUE_HAS_OTHER_CONSUMERSAn interior value feeds a consumer outside the chain.
VALUE_IS_OBSERVEDAn interior value is read from outside the graph.
GRADIENT_NEEDS_PRE_ACTIVATIONIndicates the activation gradient differentiates at its input.
SPILL_READ_TOO_EARLYIndicates a spilled value is read before the fused node writes it.
NO_CATALOG_ENTRYThe activation could be folded and the catalog has no entry.
DISABLEDThe caller switched this form off through
FusionOptions.
enum class LifetimeDefines the lifetime class and owning arena category for a value.
PARAMETERA learnable weight.
GRADIENTA gradient of a parameter.
OPTIMIZER_STATEOptimizer state such as a momentum or variance buffer.
INPUTA minibatch input or target.
ACTIVATIONAn intermediate.
SCRATCHA transient the planner may overlap freely with any other scratch.
enum class TensorLayoutSpecifies the dimension ordering for a rank-4 tensor value.
NHWCChannels innermost - the single native layout (CV-3).
NCHWChannels before the spatial axes; a boundary only.
Typedefs
using SushiAI::Graph::ValueId = std::uint32_tIdentifies a value in a Graph; an index, not a pointer.
using SushiAI::Graph::NodeId = std::uint32_tIdentifies a node in a Graph; an index, not a pointer.
using SushiAI::Graph::ArenaId = std::uint32_tIdentifies 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) noexceptComputes the output extent along one convolution axis.
Parameters
strideStep between windows; must be positive.
kernelNumber of taps; must be positive.
dilationSpacing 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
heightInput height, padded above and below by the geometry's top pad.
widthInput width, padded left and right by the geometry's left pad.
gNeeds 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_hVertical replication factor; below one throws.
scale_wHorizontal 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) noexceptTells 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
graphThe graph to rewrite in place.
optionsThe policy and preserved values.
Returns
A summary of inserted conversions.
Exceptions
ErrorIf 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
graphThe graph to check.
policyThe 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) noexceptA human-readable name for form, for diagnostics.
FusedForm fused_form_from_name(std::string_view name) noexceptReturns the catalog entry name refers to.
Parameters
nameName as fused_form_name spells it.
Returns
Matching form, or FusedForm::NONE when unmatched.
OpID fused_form_op(FusedForm form) noexceptThe operation identity a node of form carries.
Parameters
formThe catalog entry.
Returns
Its Ops identity, or Ops::LEAF.id for FusedForm::NONE.
FusedForm fused_form_of(OpID op) noexceptThe catalog entry behind an operation identity.
Parameters
opThe identity to look up.
Returns
The form, or FusedForm::NONE when op is not a fused one.
constexpr bool is_gemm_epilogue(FusedForm form) noexceptReturns whether form is a GEMM-epilogue entry.
Parameters
formThe catalog entry.
Returns
True for the GEMM_BIAS family.
OpID fused_form_activation(FusedForm form) noexceptThe activation a GEMM-epilogue form folds in, if any.
Parameters
formThe catalog entry.
Returns
The forward activation's identity, or Ops::LEAF.id when the form folds no activation.
bool fused_form_spills(FusedForm form) noexceptReturns whether a node of form writes an auxiliary second value.
Parameters
formThe catalog entry.
Returns
True when a node of form declares two outputs.
bool catalog_folds_activation(OpID activation) noexceptReturns whether any catalog entry folds activation into a GEMM.
Parameters
activationA forward activation identity.
Returns
True when some entry's chain ends in activation.
constexpr std::string_view fusion_decline_reason(FusionDecline reason) noexceptA sentence explaining reason, for the report.
FusionReport fuse(Graph &graph, const FusionOptions &options={})Partitions graph into maximal fusable subtrees and rewrites them.
Parameters
graphThe graph to rewrite in place.
optionsWhat the pass is allowed to do.
Returns
What it did and what it refused.
Exceptions
ErrorIf the graph is not in topological insertion order.
constexpr const char * lifetime_name(Lifetime lifetime) noexceptA human-readable name for lifetime, for diagnostics.
constexpr bool is_persistent(Lifetime lifetime) noexceptTrue 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) noexceptChecks whether an operation is narrowed under mixed precision.
Parameters
opOperation 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
graphThe graph to rewrite in place.

