Skip to content

Commit a46ed51

Browse files
committed
[Relax] Keep a call_tir's out_ty consistent with its arguments when canonicalizing
CanonicalizeTIRVariables substitutes the known value of a symbolic variable into type annotations, but keeps a MatchCast whose variable is used at runtime. The argument of a later call_tir then keeps the variable while the call's out_ty had it replaced, and the call_tir validator added in #20480 rejects the pair. The same happens when two MatchCasts rename one dimension, since a value is substituted one level at a time. Canonicalize the arguments first and derive the type they imply. Use the canonical out_ty when it still follows from the arguments, otherwise the implied type. Expose InferCallTIROutputTypeFromArguments for that. MLC LLM hit this compiling Phi-3.5-vision after the pin moved past #20480.
1 parent 2393db6 commit a46ed51

4 files changed

Lines changed: 141 additions & 1 deletion

File tree

‎src/relax/op/op.cc‎

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -309,7 +309,7 @@ TVM_FFI_STATIC_INIT_BLOCK() {
309309
* \return The `arg_ty`, if it can be inferred from the arguments.
310310
* Otherwise, std::nullopt.
311311
*/
312-
static ffi::Optional<Type> InferCallTIROutputTypeFromArguments(
312+
ffi::Optional<Type> InferCallTIROutputTypeFromArguments(
313313
Type func_ty, Type arg_ty, ffi::Optional<ffi::Array<int64_t>> opt_inplace_indices) {
314314
auto opt_callee_ty = func_ty.as<FuncType>();
315315
TVM_FFI_CHECK(opt_callee_ty, TypeError)

‎src/relax/op/op_common.h‎

Lines changed: 7 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -238,6 +238,13 @@ inline Type InferTypeUnary(const Call& call, const BlockBuilder&, FType f_comput
238238
return InferTypeUnary<require_float_dtype>(call, f_compute_out_dtype);
239239
}
240240

241+
/*!
242+
* \brief Derive the output type of a call_tir from the callee signature and the
243+
* argument types, or nullopt when it cannot be derived.
244+
*/
245+
ffi::Optional<Type> InferCallTIROutputTypeFromArguments(
246+
Type func_ty, Type arg_ty, ffi::Optional<ffi::Array<int64_t>> opt_inplace_indices);
247+
241248
/*!
242249
* \brief Infer the type by returning the type of the input argument.
243250
* \param call The context Call to the operator.

‎src/relax/transform/canonicalize_bindings.cc‎

Lines changed: 37 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -28,13 +28,16 @@
2828
#include <tvm/ffi/extra/structural_mutate.h>
2929
#include <tvm/ffi/reflection/registry.h>
3030
#include <tvm/relax/analysis.h>
31+
#include <tvm/relax/attrs/op.h>
3132
#include <tvm/relax/expr.h>
3233
#include <tvm/relax/expr_functor.h>
3334
#include <tvm/relax/transform.h>
3435
#include <tvm/relax/type.h>
3536
#include <tvm/relax/utils.h>
3637
#include <tvm/tirx/stmt_functor.h>
3738

39+
#include "../op/op_common.h"
40+
3841
namespace tvm {
3942
namespace relax {
4043
using namespace tvm::prim;
@@ -52,6 +55,40 @@ class SymbolicVarCanonicalizer : public ExprMutator {
5255
return CanonicalizeShapeValue(expr);
5356
}
5457

58+
Expr VisitExpr_(const CallNode* op) final {
59+
static const Op& call_tir_op = Op::Get("relax.call_tir");
60+
static const Op& call_tir_inplace_op = Op::Get("relax.call_tir_inplace");
61+
static const Op& call_tir_with_grad_op = Op::Get("relax.call_tir_with_grad");
62+
bool is_call_tir = op->op.same_as(call_tir_op) || op->op.same_as(call_tir_inplace_op) ||
63+
op->op.same_as(call_tir_with_grad_op);
64+
if (!is_call_tir || op->args.size() != 2 || op->ty_args.size() != 1) {
65+
return ExprMutator::VisitExpr_(op);
66+
}
67+
// The out_ty of a call_tir is checked against the argument types. Canonicalizing the
68+
// two independently can leave them mentioning different variables for the same value,
69+
// so when the canonical out_ty no longer follows from the arguments, use the type the
70+
// arguments imply.
71+
Expr new_op = this->VisitExpr(op->op);
72+
ffi::Array<Expr> new_args =
73+
op->args.Map([this](const Expr& arg) { return this->VisitExpr(arg); });
74+
Type out_ty = this->VisitExprDepTypeField(op->ty_args[0]);
75+
ffi::Optional<ffi::Array<int64_t>> inplace_indices;
76+
if (const auto* attrs = op->attrs.as<CallTIRInplaceAttrs>()) {
77+
inplace_indices = attrs->inplace_indices;
78+
}
79+
auto implied = InferCallTIROutputTypeFromArguments(GetType(new_args[0]), GetType(new_args[1]),
80+
inplace_indices);
81+
if (implied.has_value() && !IsBaseOf(implied.value(), out_ty)) {
82+
out_ty = implied.value();
83+
}
84+
bool unchanged =
85+
new_op.same_as(op->op) && new_args.same_as(op->args) && out_ty.same_as(op->ty_args[0]);
86+
if (unchanged) {
87+
return ffi::GetRef<Expr>(op);
88+
}
89+
return Call(Type::Missing(), new_op, new_args, op->attrs, {out_ty}, op->span);
90+
}
91+
5592
Expr VisitExpr_(const ShapeExprNode* op) final {
5693
if (!canonicalize_shape_values_) return ffi::GetRef<Expr>(op);
5794
ffi::Array<PrimExpr> values =

‎tests/python/relax/test_transform_canonicalize_bindings.py‎

Lines changed: 96 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -26,6 +26,7 @@
2626
from tvm.relax.transform.transform import CanonicalizeBindings
2727
from tvm.script import ir as I
2828
from tvm.script import relax as R
29+
from tvm.script import s_tir as Ts
2930
from tvm.script import tirx as T
3031

3132

@@ -1394,3 +1395,98 @@ def main(param_tuple: R.Tuple([R.Tensor, R.Tensor, R.Tensor])):
13941395

13951396
if __name__ == "__main__":
13961397
tvm.testing.main()
1398+
1399+
1400+
def test_call_tir_out_ty_follows_arguments_when_match_cast_is_kept():
1401+
"""A call_tir's out_ty must match what its arguments imply.
1402+
1403+
The match_cast defining `n` stays because `n` is passed to a kernel at
1404+
runtime, so the argument of the later call_tir keeps `n` in its shape.
1405+
Canonicalizing the call's out_ty on its own would replace `n` with 12
1406+
and leave the two disagreeing.
1407+
"""
1408+
n = T.dynamic("n", "int64")
1409+
m = T.dynamic("m", "int64")
1410+
1411+
@I.ir_module
1412+
class Before:
1413+
@Ts.prim_func(private=True)
1414+
def transpose(
1415+
x: T.Buffer((1, n, 2, m, 2, 8), "float16"),
1416+
y: T.Buffer((1, n, m, 2, 2, 8), "float16"),
1417+
):
1418+
for i0, i1, i2, i3, i4, i5 in T.grid(1, n, m, 2, 2, 8):
1419+
with Ts.sblock("b"):
1420+
v0, v1, v2, v3, v4, v5 = Ts.axis.remap("SSSSSS", [i0, i1, i2, i3, i4, i5])
1421+
y[v0, v1, v2, v3, v4, v5] = x[v0, v1, v3, v2, v4, v5]
1422+
1423+
@Ts.prim_func(private=True)
1424+
def add_scalar(x: T.Buffer((1, 8), "float16"), s: T.int64, y: T.Buffer((1, 8), "float16")):
1425+
for i in T.serial(8):
1426+
with Ts.sblock("b"):
1427+
vi = Ts.axis.spatial(8, i)
1428+
y[0, vi] = x[0, vi] + T.Cast("float16", s)
1429+
1430+
@R.function
1431+
def main(x: R.Tensor((1, 12, 2, 12, 2, 8), "float16"), w: R.Tensor((1, 8), "float16")):
1432+
cls = Before
1433+
with R.dataflow():
1434+
lv = R.match_cast(x, R.Tensor((1, n, 2, m, 2, 8), "float16"))
1435+
p = R.call_tir(cls.transpose, (lv,), out_ty=R.Tensor((1, n, m, 2, 2, 8), "float16"))
1436+
y = R.call_tir(cls.add_scalar, (w, n), out_ty=R.Tensor((1, 8), "float16"))
1437+
gv = (p, y)
1438+
R.output(gv)
1439+
return gv
1440+
1441+
after = relax.transform.CanonicalizeBindings()(Before)
1442+
lv_binding, p_binding = after["main"].body.blocks[0].bindings[:2]
1443+
assert isinstance(lv_binding, relax.MatchCast)
1444+
assert p_binding.value.args[1].fields[0].same_as(lv_binding.var)
1445+
# The out_ty still names the variables the argument carries.
1446+
assert p_binding.value.ty_args[0].shape[1].same_as(n)
1447+
assert p_binding.value.ty_args[0].shape[2].same_as(m)
1448+
relax.analysis.well_formed(after)
1449+
1450+
1451+
def test_call_tir_out_ty_follows_arguments_through_chained_match_casts():
1452+
"""Two match_casts rename the same dimension before a call_tir."""
1453+
n = T.dynamic("n", "int64")
1454+
a = T.dynamic("a", "int64")
1455+
b = T.dynamic("b", "int64")
1456+
1457+
@I.ir_module
1458+
class Before:
1459+
@Ts.prim_func(private=True)
1460+
def copy(x: T.Buffer((n, 8), "float16"), y: T.Buffer((n, 8), "float16")):
1461+
for i, j in T.grid(n, 8):
1462+
with Ts.sblock("b"):
1463+
vi, vj = Ts.axis.remap("SS", [i, j])
1464+
y[vi, vj] = x[vi, vj]
1465+
1466+
@Ts.prim_func(private=True)
1467+
def add_scalar(x: T.Buffer((1, 8), "float16"), s: T.int64, y: T.Buffer((1, 8), "float16")):
1468+
for i in T.serial(8):
1469+
with Ts.sblock("b"):
1470+
vi = Ts.axis.spatial(8, i)
1471+
y[0, vi] = x[0, vi] + T.Cast("float16", s)
1472+
1473+
@R.function
1474+
def main(x: R.Tensor((n, 8), "float16"), w: R.Tensor((1, 8), "float16")):
1475+
cls = Before
1476+
with R.dataflow():
1477+
lv1 = R.match_cast(x, R.Tensor((a, 8), "float16"))
1478+
lv2 = R.match_cast(lv1, R.Tensor((b, 8), "float16"))
1479+
c = R.call_tir(cls.copy, (lv2,), out_ty=R.Tensor((b, 8), "float16"))
1480+
y = R.call_tir(cls.add_scalar, (w, n), out_ty=R.Tensor((1, 8), "float16"))
1481+
gv = (c, y)
1482+
R.output(gv)
1483+
return gv
1484+
1485+
after = relax.transform.CanonicalizeBindings()(Before)
1486+
for binding in after["main"].body.blocks[0].bindings:
1487+
value = binding.value
1488+
if isinstance(value, relax.Call) and value.op == tvm.ir.Op.get("relax.call_tir"):
1489+
# Rebuilding the call re-runs the validator, so this passes only when the
1490+
# out_ty agrees with the argument types.
1491+
relax.Call(value.op, value.args, value.attrs, value.ty_args)
1492+
relax.analysis.well_formed(after)

0 commit comments

Comments
 (0)