Skip to content

Commit a7fffaa

Browse files
authored
Add context-free Call reinference and binding type propagation (#20480)
- Separate native context-free Call type inference from builder-dependent Relax and distributed inference, with strict read-only reinference from current Call inputs. - Infer `tirx.buffer_data` from its BufferVar argument and reinfer after pointer storage-scope rewrites. - Migrate TVM structural visit and mutation hooks to native tvm-ffi callbacks and remove the structural compatibility header. - Propagate Bind and Let value type changes through their binder definitions and later uses while preserving unchanged node identity. - Reuse existing Call result Type objects when inference establishes an equivalent result cheaply.
1 parent b98f258 commit a7fffaa

87 files changed

Lines changed: 2661 additions & 1484 deletions

Some content is hidden

Large Commits have some content hidden by default. Use the searchbox below for content that may be hidden.

‎include/tvm/ir/base_expr.h‎

Lines changed: 2 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -26,10 +26,10 @@
2626

2727
#include <tvm/ffi/cast.h>
2828
#include <tvm/ffi/dtype.h>
29+
#include <tvm/ffi/extra/structural_mutate.h>
30+
#include <tvm/ffi/extra/structural_visit.h>
2931
#include <tvm/ffi/reflection/registry.h>
3032
#include <tvm/ffi/string.h>
31-
// Keep raw and typed structural hook return paths available to all TVM IR nodes.
32-
#include <tvm/ir/ffi_structural_compat.h>
3333
#include <tvm/ir/source_map.h>
3434

3535
#include <cstddef>

‎include/tvm/ir/expr.h‎

Lines changed: 10 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -385,6 +385,9 @@ class Var : public Expr {
385385
/*! \brief Return a fresh ordinary Var with a new primitive type. */
386386
TVM_DLL Var CopyWithDType(PrimType dtype) const;
387387

388+
/*! \brief Return a fresh ordinary Var with a new type, retaining its metadata. */
389+
TVM_DLL Var CopyWithType(Type type) const;
390+
388391
TVM_FFI_DEFINE_OBJECT_REF_METHODS_NULLABLE(Var, Expr, VarNode);
389392
};
390393

@@ -513,6 +516,13 @@ class Call : public Expr {
513516
/*! \brief Construct a provisional Call without invoking its Op validator. */
514517
TVM_DLL static Call Unchecked(Type ret_ty, Expr op, ffi::Array<Expr> args, Attrs attrs = Attrs(),
515518
ffi::Array<Type> ty_args = ffi::Array<Type>(), Span span = Span());
519+
/*! \brief Recompute a result type from the Call's current explicit inputs.
520+
*
521+
* This ignores the Call's stored result type and does not mutate the Call.
522+
* A context-free inference hook must be registered for the operator. Missing
523+
* or invalid inputs are reported by the hook rather than screened here.
524+
*/
525+
TVM_DLL static Type ReinferType(const CallNode* call);
516526

517527
TVM_FFI_DEFINE_OBJECT_REF_METHODS_NULLABLE(Call, Expr, CallNode);
518528
TVM_DEFINE_OBJECT_REF_COW_METHOD(CallNode);

‎include/tvm/ir/expr_functor.h‎

Lines changed: 13 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -510,6 +510,19 @@ class TVM_DLL ExprMutator : public ObjectMutator {
510510
virtual UnchangedOr<PrimExpr> Mutate_(const prim::ShuffleNode* node, InplaceMode inplace_mode);
511511

512512
protected:
513+
/*!
514+
* \brief Reinfer the result type of a Call after ordinary mutation.
515+
* \param mutated The ordinary mutation result, or Unchanged for the original Call.
516+
* \param original The borrowed Call before ordinary mutation.
517+
* \param inplace_mode Permission to modify the original Call in place.
518+
* \return The combined mutation result, preserving Unchanged when possible.
519+
* \note A unique replacement Call may be updated in place even when the original
520+
* could not be. Inference errors propagate to the caller.
521+
*/
522+
static UnchangedOr<Expr> ReinferMutatedCallType(UnchangedOr<Expr> mutated,
523+
const CallNode* original,
524+
InplaceMode inplace_mode);
525+
513526
/*!
514527
* \brief Construct a mutator with an extended native dispatch table.
515528
* \param vtable The finalized table, which must outlive this mutator.

‎include/tvm/ir/ffi_structural_compat.h‎

Lines changed: 0 additions & 118 deletions
This file was deleted.

‎include/tvm/ir/op.h‎

Lines changed: 3 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -43,6 +43,9 @@ namespace tvm {
4343
template <typename>
4444
class OpAttrMap;
4545

46+
/*! \brief Infer a Call's result type from its explicit inputs without builder state. */
47+
using FInferType = ffi::reflection::NativeFunctionView<ffi::Expected<Type>(const CallNode* call)>;
48+
4649
/*! \brief An operator argument's name and documentation. */
4750
class ArgumentInfoNode : public ffi::Object {
4851
public:

‎include/tvm/relax/op_attr_types.h‎

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -66,7 +66,7 @@ using FCallPacked = ffi::String;
6666
* \param call The call expression to be derived.
6767
* \param ctx The builder context.
6868
*/
69-
using FInferType = ffi::TypedFunction<Type(const Call& call, const BlockBuilder& ctx)>;
69+
using FInferTypeWithBuilder = ffi::TypedFunction<Type(const Call& call, const BlockBuilder& ctx)>;
7070

7171
/*!
7272
* \brief The function type of a normalization function.

‎python/tvm/ir/__init__.py‎

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -53,6 +53,7 @@
5353
Var,
5454
is_prim_expr,
5555
is_prim_var,
56+
reinfer_type,
5657
)
5758
from . import prim
5859
from .function import BaseFunc, CallingConv

‎python/tvm/ir/expr.py‎

Lines changed: 9 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -593,6 +593,15 @@ def unchecked(
593593
)
594594

595595

596+
def reinfer_type(call: Call) -> "tvm.ir.Type":
597+
"""Derive a Call's result type from its current inputs without changing the Call.
598+
599+
The operator must have a context-free inference rule. Invalid inputs and
600+
missing rules raise errors.
601+
"""
602+
return _ffi_api.reinfer_type(call)
603+
604+
596605
@tvm_ffi.register_object("ir.TensorRegion")
597606
class TensorRegion(Expr, Scriptable):
598607
"""A region of an arbitrary tensor expression.

‎python/tvm/relax/transform/legalize_ops/linear_algebra.py‎

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -52,7 +52,7 @@ def te_matmul(a: te.Tensor, b: te.Tensor) -> te.Tensor:
5252

5353
a_relax = relax.Var("a", relax.TensorType(a.shape))
5454
b_relax = relax.Var("b", relax.TensorType(b.shape))
55-
f_infer_ty = call.op.get_attr("FInferType")
55+
f_infer_ty = call.op.get_attr("relax.FInferTypeWithBuilder")
5656
output_shape = f_infer_ty(relax.op.matmul(a_relax, b_relax), bb).shape
5757
if isinstance(a_shape[-1], tirx.IntImm) and a_shape[-1] == 0:
5858
return te.compute(

0 commit comments

Comments
 (0)