ShapeRuleResult
PythonBase type for shape propagation rule outcomes.
ShapeRuleResult()Base type for shape propagation rule outcomes.
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.Base type for shape propagation rule outcomes.
ShapeRuleResult()Base type for shape propagation rule outcomes.
Rule does not own this operator case and generic fallback may run.
NotApplicable() -> Nonedataclass
Rule does not own this operator case and generic fallback may run.
Rule produced concrete output shape updates.
Updated(shapes: dict[str, tuple[int, ...]]) -> Nonedataclass
Rule produced concrete output shape updates.
shapes: dict[str, tuple[int, ...]]Rule owns this operator case and no shape update is required.
NoUpdate(reason: str = '') -> Nonedataclass
Rule owns this operator case and no shape update is required.
reason: str = ''Rule owns this operator case and cannot safely infer an update.
Blocked(reason: str) -> Nonedataclass
Rule owns this operator case and cannot safely infer an update.
reason: strShape propagation outcome.
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(),) -> Nonedataclass
Shape propagation outcome.
changed_tensor_ids: set[str] = field(default_factory=set)Tensor IDs whose shapes changed during propagation.
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.
evaluated_tensor_ids: set[str] = field(default_factory=set)Shape-expression tensors whose compile-time value
propagation computed and stored in AirTensor.inferred_value.
unresolved_tensor_ids: set[str] = field(default_factory=set)Dynamic tensor IDs that remain unresolved.
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.
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.
Infer tensor shapes and shape-expression values over the graph until fixed point.
propagate_shapes(model: AirModel, *, strict: bool = False, warn_unresolved: bool = True) -> ShapeInferenceResultInfer 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
| Name | Type | Default | Description |
|---|---|---|---|
model | AirModel | Required | AIR model to update in place. |
strict | bool | False | If True, raise when unresolved dynamic tensors remain. |
warn_unresolved | bool | True | If True, log unresolved tensors when not raising. |
Returns
| Type | Description |
|---|---|
ShapeInferenceResult | Shape inference result with updated tensors and diagnostics. |
Raises
| Type | Description |
|---|---|
UnsupportedModelError | If ``strict`` is True and unresolved dynamic tensors remain, or if a shape or value computed from proven inputs contradicts a tensor's declared contract. |
TypeError | If an operator's options type does not match its shape rule's contract, or a rule returns an unknown result type. |