Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
2 changes: 1 addition & 1 deletion src/relax/op/op.cc
Original file line number Diff line number Diff line change
Expand Up @@ -309,7 +309,7 @@ TVM_FFI_STATIC_INIT_BLOCK() {
* \return The `arg_ty`, if it can be inferred from the arguments.
* Otherwise, std::nullopt.
*/
static ffi::Optional<Type> InferCallTIROutputTypeFromArguments(
ffi::Optional<Type> InferCallTIROutputTypeFromArguments(
Type func_ty, Type arg_ty, ffi::Optional<ffi::Array<int64_t>> opt_inplace_indices) {
auto opt_callee_ty = func_ty.as<FuncType>();
TVM_FFI_CHECK(opt_callee_ty, TypeError)
Expand Down
7 changes: 7 additions & 0 deletions src/relax/op/op_common.h
Original file line number Diff line number Diff line change
Expand Up @@ -238,6 +238,13 @@ inline Type InferTypeUnary(const Call& call, const BlockBuilder&, FType f_comput
return InferTypeUnary<require_float_dtype>(call, f_compute_out_dtype);
}

/*!
* \brief Derive the output type of a call_tir from the callee signature and the
* argument types, or nullopt when it cannot be derived.
*/
ffi::Optional<Type> InferCallTIROutputTypeFromArguments(
Type func_ty, Type arg_ty, ffi::Optional<ffi::Array<int64_t>> opt_inplace_indices);

/*!
* \brief Infer the type by returning the type of the input argument.
* \param call The context Call to the operator.
Expand Down
37 changes: 37 additions & 0 deletions src/relax/transform/canonicalize_bindings.cc
Original file line number Diff line number Diff line change
Expand Up @@ -28,13 +28,16 @@
#include <tvm/ffi/extra/structural_mutate.h>
#include <tvm/ffi/reflection/registry.h>
#include <tvm/relax/analysis.h>
#include <tvm/relax/attrs/op.h>
#include <tvm/relax/expr.h>
#include <tvm/relax/expr_functor.h>
#include <tvm/relax/transform.h>
#include <tvm/relax/type.h>
#include <tvm/relax/utils.h>
#include <tvm/tirx/stmt_functor.h>

#include "../op/op_common.h"

namespace tvm {
namespace relax {
using namespace tvm::prim;
Expand All @@ -52,6 +55,40 @@ class SymbolicVarCanonicalizer : public ExprMutator {
return CanonicalizeShapeValue(expr);
}

Expr VisitExpr_(const CallNode* op) final {
static const Op& call_tir_op = Op::Get("relax.call_tir");
static const Op& call_tir_inplace_op = Op::Get("relax.call_tir_inplace");
static const Op& call_tir_with_grad_op = Op::Get("relax.call_tir_with_grad");
bool is_call_tir = op->op.same_as(call_tir_op) || op->op.same_as(call_tir_inplace_op) ||
op->op.same_as(call_tir_with_grad_op);
if (!is_call_tir || op->args.size() != 2 || op->ty_args.size() != 1) {
return ExprMutator::VisitExpr_(op);
}
// The out_ty of a call_tir is checked against the argument types. Canonicalizing the
// two independently can leave them mentioning different variables for the same value,
// so when the canonical out_ty no longer follows from the arguments, use the type the
// arguments imply.
Expr new_op = this->VisitExpr(op->op);
ffi::Array<Expr> new_args =
op->args.Map([this](const Expr& arg) { return this->VisitExpr(arg); });
Type out_ty = this->VisitExprDepTypeField(op->ty_args[0]);
ffi::Optional<ffi::Array<int64_t>> inplace_indices;
if (const auto* attrs = op->attrs.as<CallTIRInplaceAttrs>()) {
inplace_indices = attrs->inplace_indices;
}
auto implied = InferCallTIROutputTypeFromArguments(GetType(new_args[0]), GetType(new_args[1]),
inplace_indices);
if (implied.has_value() && !IsBaseOf(implied.value(), out_ty)) {
out_ty = implied.value();
}
bool unchanged =
new_op.same_as(op->op) && new_args.same_as(op->args) && out_ty.same_as(op->ty_args[0]);
if (unchanged) {
return ffi::GetRef<Expr>(op);
}
return Call(Type::Missing(), new_op, new_args, op->attrs, {out_ty}, op->span);
}

Expr VisitExpr_(const ShapeExprNode* op) final {
if (!canonicalize_shape_values_) return ffi::GetRef<Expr>(op);
ffi::Array<PrimExpr> values =
Expand Down
96 changes: 96 additions & 0 deletions tests/python/relax/test_transform_canonicalize_bindings.py
Original file line number Diff line number Diff line change
Expand Up @@ -26,6 +26,7 @@
from tvm.relax.transform.transform import CanonicalizeBindings
from tvm.script import ir as I
from tvm.script import relax as R
from tvm.script import s_tir as Ts
from tvm.script import tirx as T


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

if __name__ == "__main__":
tvm.testing.main()


def test_call_tir_out_ty_follows_arguments_when_match_cast_is_kept():
"""A call_tir's out_ty must match what its arguments imply.

The match_cast defining `n` stays because `n` is passed to a kernel at
runtime, so the argument of the later call_tir keeps `n` in its shape.
Canonicalizing the call's out_ty on its own would replace `n` with 12
and leave the two disagreeing.
"""
n = T.dynamic("n", "int64")
m = T.dynamic("m", "int64")

@I.ir_module
class Before:
@Ts.prim_func(private=True)
def transpose(
x: T.Buffer((1, n, 2, m, 2, 8), "float16"),
y: T.Buffer((1, n, m, 2, 2, 8), "float16"),
):
for i0, i1, i2, i3, i4, i5 in T.grid(1, n, m, 2, 2, 8):
with Ts.sblock("b"):
v0, v1, v2, v3, v4, v5 = Ts.axis.remap("SSSSSS", [i0, i1, i2, i3, i4, i5])
y[v0, v1, v2, v3, v4, v5] = x[v0, v1, v3, v2, v4, v5]

@Ts.prim_func(private=True)
def add_scalar(x: T.Buffer((1, 8), "float16"), s: T.int64, y: T.Buffer((1, 8), "float16")):
for i in T.serial(8):
with Ts.sblock("b"):
vi = Ts.axis.spatial(8, i)
y[0, vi] = x[0, vi] + T.Cast("float16", s)

@R.function
def main(x: R.Tensor((1, 12, 2, 12, 2, 8), "float16"), w: R.Tensor((1, 8), "float16")):
cls = Before
with R.dataflow():
lv = R.match_cast(x, R.Tensor((1, n, 2, m, 2, 8), "float16"))
p = R.call_tir(cls.transpose, (lv,), out_ty=R.Tensor((1, n, m, 2, 2, 8), "float16"))
y = R.call_tir(cls.add_scalar, (w, n), out_ty=R.Tensor((1, 8), "float16"))
gv = (p, y)
R.output(gv)
return gv

after = relax.transform.CanonicalizeBindings()(Before)
lv_binding, p_binding = after["main"].body.blocks[0].bindings[:2]
assert isinstance(lv_binding, relax.MatchCast)
assert p_binding.value.args[1].fields[0].same_as(lv_binding.var)
# The out_ty still names the variables the argument carries.
assert p_binding.value.ty_args[0].shape[1].same_as(n)
assert p_binding.value.ty_args[0].shape[2].same_as(m)
relax.analysis.well_formed(after)


def test_call_tir_out_ty_follows_arguments_through_chained_match_casts():
"""Two match_casts rename the same dimension before a call_tir."""
n = T.dynamic("n", "int64")
a = T.dynamic("a", "int64")
b = T.dynamic("b", "int64")

@I.ir_module
class Before:
@Ts.prim_func(private=True)
def copy(x: T.Buffer((n, 8), "float16"), y: T.Buffer((n, 8), "float16")):
for i, j in T.grid(n, 8):
with Ts.sblock("b"):
vi, vj = Ts.axis.remap("SS", [i, j])
y[vi, vj] = x[vi, vj]

@Ts.prim_func(private=True)
def add_scalar(x: T.Buffer((1, 8), "float16"), s: T.int64, y: T.Buffer((1, 8), "float16")):
for i in T.serial(8):
with Ts.sblock("b"):
vi = Ts.axis.spatial(8, i)
y[0, vi] = x[0, vi] + T.Cast("float16", s)

@R.function
def main(x: R.Tensor((n, 8), "float16"), w: R.Tensor((1, 8), "float16")):
cls = Before
with R.dataflow():
lv1 = R.match_cast(x, R.Tensor((a, 8), "float16"))
lv2 = R.match_cast(lv1, R.Tensor((b, 8), "float16"))
c = R.call_tir(cls.copy, (lv2,), out_ty=R.Tensor((b, 8), "float16"))
y = R.call_tir(cls.add_scalar, (w, n), out_ty=R.Tensor((1, 8), "float16"))
gv = (c, y)
R.output(gv)
return gv

after = relax.transform.CanonicalizeBindings()(Before)
for binding in after["main"].body.blocks[0].bindings:
value = binding.value
if isinstance(value, relax.Call) and value.op == tvm.ir.Op.get("relax.call_tir"):
# Rebuilding the call re-runs the validator, so this passes only when the
# out_ty agrees with the argument types.
relax.Call(value.op, value.args, value.attrs, value.ty_args)
relax.analysis.well_formed(after)
Loading