diff --git a/src/relax/transform/adjust_matmul_order.cc b/src/relax/transform/adjust_matmul_order.cc index bdd28054d86d..64870b22b5ea 100644 --- a/src/relax/transform/adjust_matmul_order.cc +++ b/src/relax/transform/adjust_matmul_order.cc @@ -242,6 +242,10 @@ std::tuple)>> 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)) { diff --git a/src/relax/transform/fold_constant.cc b/src/relax/transform/fold_constant.cc index 142d4cb75b83..8b962369c52f 100644 --- a/src/relax/transform/fold_constant.cc +++ b/src/relax/transform/fold_constant.cc @@ -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(ndarray->data); int64_t num_elems = ndarray->shape[0]; ffi::Array 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(ndarray->data)); + break; + case 32: + append_values(static_cast(ndarray->data)); + break; + case 64: + append_values(static_cast(ndarray->data)); + break; + default: + return post_call; } return ShapeExpr(shape_values); } diff --git a/src/relax/transform/lambda_lift.cc b/src/relax/transform/lambda_lift.cc index d3391eb0aeb0..8a72f76db7c6 100644 --- a/src/relax/transform/lambda_lift.cc +++ b/src/relax/transform/lambda_lift.cc @@ -247,6 +247,10 @@ class LambdaLifter : public ExprMutator { current_lambda_var_ = binding->var; auto new_value = VisitExpr(binding->value); + const auto* call = new_value.as(); + if (call && call->op.same_as(make_closure_op_)) { + closures_.insert(binding->var); + } if (!rebind_map_.count(binding->var)) { ReEmitBinding(binding, new_value); } @@ -254,6 +258,14 @@ class LambdaLifter : public ExprMutator { current_lambda_var_ = cache; } + void VisitBinding_(const VarBindingNode* binding, const VarNode* var_node) final { + auto new_value = VisitExpr(ffi::GetRef(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 @@ -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()) { + if (IsClosure(var)) { // if the original op was pure, we should use invoke_pure_closure - Call orig_call = bound_value.value().as_or_throw(); - bool is_pure = [&]() -> bool { - if (auto op = orig_call->op.as()) { - static const auto& purity_map = Op::GetAttrMap("FPurity"); - return purity_map.get(op.value(), false); - } else if (const auto* func_ty = orig_call->op->ty.as()) { - 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()->purity; auto prev = call; call = diff --git a/src/relax/transform/remove_unused_outputs.cc b/src/relax/transform/remove_unused_outputs.cc index dee3149b17d9..59df39c326fd 100644 --- a/src/relax/transform/remove_unused_outputs.cc +++ b/src/relax/transform/remove_unused_outputs.cc @@ -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()) { + 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; diff --git a/tests/python/relax/test_transform_adjust_matmul_order.py b/tests/python/relax/test_transform_adjust_matmul_order.py index 170702a67629..5b616b8afa05 100644 --- a/tests/python/relax/test_transform_adjust_matmul_order.py +++ b/tests/python/relax/test_transform_adjust_matmul_order.py @@ -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() diff --git a/tests/python/relax/test_transform_fold_constant.py b/tests/python/relax/test_transform_fold_constant.py index 810873ce4af4..bad9a32834d9 100644 --- a/tests/python/relax/test_transform_fold_constant.py +++ b/tests/python/relax/test_transform_fold_constant.py @@ -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() diff --git a/tests/python/relax/test_transform_lambda_lift.py b/tests/python/relax/test_transform_lambda_lift.py index 22ed281a2038..1333ca50362d 100644 --- a/tests/python/relax/test_transform_lambda_lift.py +++ b/tests/python/relax/test_transform_lambda_lift.py @@ -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() diff --git a/tests/python/relax/test_transform_remove_unused_outputs.py b/tests/python/relax/test_transform_remove_unused_outputs.py index 1a2f305bfebf..0db447501a54 100644 --- a/tests/python/relax/test_transform_remove_unused_outputs.py +++ b/tests/python/relax/test_transform_remove_unused_outputs.py @@ -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()