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
4 changes: 4 additions & 0 deletions src/relax/transform/adjust_matmul_order.cc
Original file line number Diff line number Diff line change
Expand Up @@ -242,6 +242,10 @@ std::tuple<DFPattern, ffi::TypedFunction<Expr(Expr, ffi::Map<DFPattern, Expr>)>>
transpose_shape_last_two_dims(shape_c);
}

// A rank-one middle operand is treated as a column vector in A @ B,
// but as a row vector in B @ C, so matrix associativity does not apply.
if (shape_b.size() < 2) return expr;

// If two of the three are compile-time, group those two values
// together, to allow them to be lifted out and pre-computed.
if (is_compile_time(expr_a) && is_compile_time(expr_b)) {
Expand Down
23 changes: 20 additions & 3 deletions src/relax/transform/fold_constant.cc
Original file line number Diff line number Diff line change
Expand Up @@ -385,11 +385,28 @@ class ConstantFolder : public ExprMutator {
TVM_FFI_ICHECK(ndarray.IsContiguous());
TVM_FFI_ICHECK_EQ(ndarray->byte_offset, 0);
TVM_FFI_ICHECK_EQ(ndarray->ndim, 1);
const int64_t* data = static_cast<const int64_t*>(ndarray->data);
int64_t num_elems = ndarray->shape[0];
ffi::Array<PrimExpr> shape_values;
for (int64_t i = 0; i < num_elems; i++) {
shape_values.push_back(IntImm::Int64(data[i]));
auto append_values = [&](const auto* data) {
for (int64_t i = 0; i < num_elems; i++) {
shape_values.push_back(IntImm::Int64(data[i]));
}
};
if (ndarray->dtype.code != kDLInt || ndarray->dtype.lanes != 1) {
return post_call;
}
switch (ndarray->dtype.bits) {
case 16:
append_values(static_cast<const int16_t*>(ndarray->data));
break;
case 32:
append_values(static_cast<const int32_t*>(ndarray->data));
break;
case 64:
append_values(static_cast<const int64_t*>(ndarray->data));
break;
default:
return post_call;
}
return ShapeExpr(shape_values);
}
Expand Down
31 changes: 14 additions & 17 deletions src/relax/transform/lambda_lift.cc
Original file line number Diff line number Diff line change
Expand Up @@ -247,13 +247,25 @@ class LambdaLifter : public ExprMutator {
current_lambda_var_ = binding->var;

auto new_value = VisitExpr(binding->value);
const auto* call = new_value.as<CallNode>();
if (call && call->op.same_as(make_closure_op_)) {
closures_.insert(binding->var);
}
if (!rebind_map_.count(binding->var)) {
ReEmitBinding(binding, new_value);
}

current_lambda_var_ = cache;
}

void VisitBinding_(const VarBindingNode* binding, const VarNode* var_node) final {
auto new_value = VisitExpr(ffi::GetRef<Var>(var_node));
if (IsClosure(new_value)) {
closures_.insert(binding->var);
}
ReEmitBinding(binding, new_value);
}

Expr VisitExpr_(const FunctionNode* func_node) final {
if (!current_lambda_var_) {
// Early bail-out for top-level functions
Expand Down Expand Up @@ -368,24 +380,9 @@ class LambdaLifter : public ExprMutator {

// Call "relax.invoke_closure" to invoke closure

auto bound_value = LookupBinding(var);
if (IsClosure(var) && bound_value.as<CallNode>()) {
if (IsClosure(var)) {
// if the original op was pure, we should use invoke_pure_closure
Call orig_call = bound_value.value().as_or_throw<Call>();
bool is_pure = [&]() -> bool {
if (auto op = orig_call->op.as<Op>()) {
static const auto& purity_map = Op::GetAttrMap<bool>("FPurity");
return purity_map.get(op.value(), false);
} else if (const auto* func_ty = orig_call->op->ty.as<FuncTypeNode>()) {
return func_ty->purity;
} else {
TVM_FFI_THROW(InternalError)
<< "Could not determine purity of call to " << orig_call->op
<< ", as it is neither a tvm::Op (type = \"" << orig_call->op->GetTypeKey()
<< "\"), "
<< "nor is is annotated with FuncType (ty = " << orig_call->op->ty << ")";
}
}();
bool is_pure = GetType(var).as_or_throw<FuncType>()->purity;

auto prev = call;
call =
Expand Down
13 changes: 13 additions & 0 deletions src/relax/transform/remove_unused_outputs.cc
Original file line number Diff line number Diff line change
Expand Up @@ -95,6 +95,19 @@ class PartialTupleUsageCollector : ExprVisitor {
known_bindings_.Set(binding->var, GetBoundValue(binding));
}

void VisitExpr_(const CallNode* op) override {
ExprVisitor::VisitExpr_(op);

static const Op make_closure_op = Op::Get("relax.make_closure");
if (op->op.same_as(make_closure_op) && !op->args.empty()) {
if (auto callee = op->args[0].as<GlobalVar>()) {
if (auto it = output_usage_mask_.find(callee.value()); it != output_usage_mask_.end()) {
std::fill(it->second.begin(), it->second.end(), true);
}
}
}
}

void VisitExpr_(const TupleGetItemNode* op) override {
if (auto* usage_mask_ptr = GetCalleeUsageMask(op->tuple)) {
auto& used_indices = *usage_mask_ptr;
Expand Down
17 changes: 17 additions & 0 deletions tests/python/relax/test_transform_adjust_matmul_order.py
Original file line number Diff line number Diff line change
Expand Up @@ -1030,5 +1030,22 @@ def test_identity_permute_dims_numerics(self):
np.testing.assert_array_equal(out_after, expected)


def test_rank_one_middle_operand_is_not_reassociated():
@I.ir_module
class Before:
@R.function
def main(
a: R.Tensor((2, 2), "float32"),
b: R.Tensor((2,), "float32"),
c: R.Tensor((2, 2), "float32"),
):
R.func_attr({"num_input": 1})
inner = R.matmul(a, b)
return R.matmul(inner, c)

after = relax.transform.AdjustMatmulOrder()(Before)
tvm.ir.assert_structural_equal(after, Before)


if __name__ == "__main__":
tvm.testing.main()
26 changes: 26 additions & 0 deletions tests/python/relax/test_transform_fold_constant.py
Original file line number Diff line number Diff line change
Expand Up @@ -611,5 +611,31 @@ def main(x: R.Tensor((m,), "float32")):
tvm.ir.assert_structural_equal(after, Module)


def test_fold_tensor_to_shape_int32():
@I.ir_module
class Module:
@R.function
def before(
data: R.Tensor((6,), "float32"), shape_data: R.Tensor((2,), "int32")
):
with R.dataflow():
shape: R.Shape(ndim=2) = R.tensor_to_shape(shape_data)
out: R.Tensor(ndim=2, dtype="float32") = R.reshape(data, shape)
R.output(out)
return out

@R.function
def expected(data: R.Tensor((6,), "float32")) -> R.Tensor((2, 3), "float32"):
with R.dataflow():
out = R.reshape(data, R.shape([2, 3]))
R.output(out)
return out

before = gen_mod(Module, "before", {"shape_data": np.array([2, 3], dtype="int32")})
expected = gen_mod(Module, "expected", {})
after = relax.transform.FoldConstant()(before)
tvm.ir.assert_structural_equal(after, expected)


if __name__ == "__main__":
tvm.testing.main()
17 changes: 17 additions & 0 deletions tests/python/relax/test_transform_lambda_lift.py
Original file line number Diff line number Diff line change
Expand Up @@ -553,5 +553,22 @@ def main_inner(
assert_structural_equal(Expected, After)


def test_closure_called_through_alias():
@I.ir_module
class Before:
@R.function
def main(x: R.Tensor((2,), "float32"), y: R.Tensor((2,), "float32")):
@R.function
def inner(z: R.Tensor((2,), "float32")):
return R.add(z, x)

alias = inner
return alias(y)

after = transform.LambdaLift()(Before)
assert relax.analysis.check_well_formed(after, check_ty=True)
assert "R.invoke_pure_closure" in after.script()


if __name__ == "__main__":
tvm.testing.main()
23 changes: 23 additions & 0 deletions tests/python/relax/test_transform_remove_unused_outputs.py
Original file line number Diff line number Diff line change
Expand Up @@ -150,5 +150,28 @@ def func(B: R.Tensor([16, 16], "int32")) -> R.Tuple(
tvm.ir.assert_structural_equal(After, Expected)


def test_closure_callee_outputs_are_preserved():
@I.ir_module
class Before:
@R.function(private=True)
def outputs(
value: R.Tensor((2,), "float32"), env: R.Tensor((2,), "float32")
) -> R.Tuple(R.Tensor((2,), "float32"), R.Tensor((2,), "float32")):
return value, R.add(value, env)

@R.function
def main(x: R.Tensor((2,), "float32"), y: R.Tensor((2,), "float32")):
closure = R.make_closure(Before.outputs, (x,))
result = R.invoke_pure_closure(
closure,
(y,),
ty_args=R.Tuple(R.Tensor((2,), "float32"), R.Tensor((2,), "float32")),
)
return result[0]

after = tvm.relax.transform.RemoveUnusedOutputs()(Before)
tvm.ir.assert_structural_equal(after, Before)


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