Skip to content
heliaAOT
HELIA HUB

shape_propagation

Machine-readable model

  • ShapeRuleResultclassBase type for shape propagation rule outcomes.
  • NotApplicableclassRule does not own this operator case and generic fallback may run.
  • UpdatedclassRule produced concrete output shape updates.
  • NoUpdateclassRule owns this operator case and no shape update is required.
  • BlockedclassRule owns this operator case and cannot safely infer an update.
  • ShapeInferenceResultclassShape propagation outcome.
  • propagate_shapesfunctionInfer tensor shapes and shape-expression values over the graph until fixed point.
class

Shape propagation outcome.

helia_aot/air/shape_propagation.py:92

ShapeInferenceResult(
changed_tensor_ids: set[str] = set(),
confirmed_tensor_ids: set[str] = set(),
evaluated_tensor_ids: set[str] = set(),
unresolved_tensor_ids: set[str] = set(),
deferred_conflict_tensor_ids: set[str] = set(),
diagnostics: list[str] = list(),
) -> None

dataclass

Shape propagation outcome.

attribute

helia_aot/air/shape_propagation.py:115

confirmed_tensor_ids: set[str] = field(default_factory=set)

Tensor IDs whose shapes a rule computed from proven operands, whether or not the stored shape changed, plus concrete graph inputs with a dynamic signature. A dynamic-signature tensor whose inferred shape already matched is resolved, not unresolved.

attribute

helia_aot/air/shape_propagation.py:118

deferred_conflict_tensor_ids: set[str] = field(default_factory=set)

Unresolved tensor IDs whose inferred shape conflicted with their declared contract while computed from an unresolved operand. Such a shape can come from a placeholder, so the update was dropped rather than raised.

attribute

helia_aot/air/shape_propagation.py:119

diagnostics: list[str] = field(default_factory=list)

Human-readable diagnostics describing inference updates and unresolved tensors. Skip diagnostics (rule/fallback skipped) are only included when DEBUG logging is enabled to avoid unbounded growth on large graphs.

function

Infer tensor shapes and shape-expression values over the graph until fixed point.

helia_aot/air/shape_propagation.py:1273

propagate_shapes(model: AirModel, *, strict: bool = False, warn_unresolved: bool = True) -> ShapeInferenceResult

Infer tensor shapes and shape-expression values over the graph until fixed point.

Shapes and values depend on each other: SHAPE turns a tensor’s shape into a value, and a RESHAPE turns its target value back into a shape. Both are facts in the same worklist, so a chain of any depth resolves without a fixed pass order. Values are computed only for SHAPE and for STRIDED_SLICE and PACK nodes over shape-derived tensors (see shape_expressions), and only from proven shapes, so a placeholder can never become a value. A shape the model declares fully static counts as proven, so a value can come from declared metadata as well as from the graph-input shapes. Each value is stored in AirTensor.inferred_value, never in data, and the graph structure is left unchanged; FoldStaticShapeExpressions turns the values into constants afterwards.

Parameters of propagate_shapes
NameTypeDefaultDescription
modelAirModelRequiredAIR model to update in place.
strictboolFalseIf True, raise when unresolved dynamic tensors remain.
warn_unresolvedboolTrueIf True, log unresolved tensors when not raising.
Returns of propagate_shapes
TypeDescription
ShapeInferenceResultShape inference result with updated tensors and diagnostics.
Errors raised by propagate_shapes
TypeDescription
UnsupportedModelErrorIf ``strict`` is True and unresolved dynamic tensors remain, or if a shape or value computed from proven inputs contradicts a tensor's declared contract.
TypeErrorIf an operator's options type does not match its shape rule's contract, or a rule returns an unknown result type.