diff --git a/python/tvm/relax/transform/legalize_ops/manipulate.py b/python/tvm/relax/transform/legalize_ops/manipulate.py index ecc904b1755a..2647213bf09f 100644 --- a/python/tvm/relax/transform/legalize_ops/manipulate.py +++ b/python/tvm/relax/transform/legalize_ops/manipulate.py @@ -39,7 +39,12 @@ def reshape_call_te(bb: BlockBuilder, call: Call): # If target shape is Var, pass its bound expr only when it is ShapeExpr if isinstance(tgt_shape, Var): tgt_shape = bb.lookup_binding(tgt_shape) - assert isinstance(tgt_shape, ShapeExpr) + if not isinstance(tgt_shape, ShapeExpr): + raise ValueError( + f"Cannot legalize {call.op}: expected the bound target shape to be a " + f"ShapeExpr with statically known values, but got " + f"{type(tgt_shape).__name__} bound to {call.args[1]}." + ) return bb.call_te(te_func, call.args[0], tgt_shape, primfunc_name_hint=primfunc_name) return reshape_call_te diff --git a/src/relax/transform/fold_constant.cc b/src/relax/transform/fold_constant.cc index 142d4cb75b83..57de82ddde64 100644 --- a/src/relax/transform/fold_constant.cc +++ b/src/relax/transform/fold_constant.cc @@ -30,6 +30,8 @@ #include #include +#include "utils.h" + namespace tvm { namespace relax { using namespace tvm::prim; @@ -348,6 +350,7 @@ class ConstantFolder : public ExprMutator { } new_args.push_back(arg); } + Type original_ty = post_call->ty; Type ret_ty = Type::Missing(); if (post_call->ty.as() && !infer_type_map.count(op) && !infer_type_with_builder_map.count(op)) { @@ -362,11 +365,19 @@ class ConstantFolder : public ExprMutator { if (legalize_map.count(op)) { // Get the legalized expression Call post_call_normalized = builder_->Normalize(post_call).as_or_throw(); - Expr legalized_expr = builder_->Normalize(legalize_map[op](builder_, post_call_normalized)); - // If the legalized expression is call_tir, try to fold it. - const CallNode* call = legalized_expr.as(); - if (call && call->op.same_as(call_tir_op)) { - return VisitCallTIR(ffi::GetRef(call)).value_or(post_call); + // Only probe foldability once shapes are known, matching LegalizeOps's own gate. + if (CanLegalizeCall(op, post_call_normalized)) { + Expr legalized_expr = + builder_->Normalize(legalize_map[op](builder_, post_call_normalized)); + // If the legalized expression is call_tir, try to fold it. + const CallNode* call = legalized_expr.as(); + if (call && call->op.same_as(call_tir_op)) { + return VisitCallTIR(ffi::GetRef(call)).value_or(post_call); + } + } else { + // Restore the original (well-formed) type instead of the `Missing` type set above. + return Call(original_ty, post_call->op, post_call->args, post_call->attrs, + post_call->ty_args, post_call->span); } } else if (op->name == "relax.tensor_to_shape") { // Special handling for composite op "relax.tensor_to_shape" diff --git a/src/relax/transform/legalize_ops.cc b/src/relax/transform/legalize_ops.cc index b1b6343d7a39..3c599986a80d 100644 --- a/src/relax/transform/legalize_ops.cc +++ b/src/relax/transform/legalize_ops.cc @@ -36,6 +36,8 @@ #include +#include "utils.h" + namespace tvm { namespace relax { @@ -45,27 +47,6 @@ struct OpIdentityLess { TVM_REGISTER_PASS_CONFIG_OPTION("relax.transform.apply_legalize_ops", bool); -/*! - * \brief Check if a given Tensor/Shape/TupleType contains shapes whose - * values are all known. - * \param ty The Type to be checked. - * \return A boolean indicating the given type contains shape values that are all known. - */ -bool KnowAllShapeValues(const Type& ty) { - if (const auto* tensor_ty = ty.as()) { - return tensor_ty->shape.has_value() && tensor_ty->shape.value()->IsInstance(); - } else if (const auto* shape_ty = ty.as()) { - return shape_ty->values.has_value(); - } else if (const auto* tuple_ty = ty.as()) { - return std::all_of(tuple_ty->fields.begin(), tuple_ty->fields.end(), - [](Type field_ty) { return KnowAllShapeValues(field_ty); }); - } else if (ty.as()) { - return true; - } else { - return false; - } -} - class LegalizeMutator : public ExprMutator { public: explicit LegalizeMutator(const IRModule& mod, @@ -242,7 +223,6 @@ class LegalizeMutator : public ExprMutator { Call visited_call = this->VisitExprPostOrder_(call).as_or_throw(); static const auto& legalize_map = Op::GetAttrMap("FLegalize"); static const auto& call_packed_map = Op::GetAttrMap("FCallPacked"); - static const auto& requires_arg_shapes_map = Op::GetAttrMap("RequiresArgumentShapes"); static const Op call_pure_packed_op = Op::Get("relax.call_pure_packed"); static const Op call_tir_op = Op::Get("relax.call_tir"); static const Op call_dps_packed_op = Op::Get("relax.call_dps_packed"); @@ -258,71 +238,31 @@ class LegalizeMutator : public ExprMutator { return visited_call; } - bool shapes_are_known_if_required = [&]() -> bool { - bool requires_arg_shapes = requires_arg_shapes_map.get(op, true); - if (!requires_arg_shapes) { - // This operator does not require its arguments to have a - // known shape/dtype. For example, the "relax.tensor_ndim" - // operator can output the dimensionality of a tensor at - // runtime, and does not require the dimensionality to be - // known at compile-time. - return true; - } - - bool arg_shapes_defined = - std::all_of(visited_call->args.begin(), visited_call->args.end(), - [](Expr arg) { return KnowAllShapeValues(GetType(arg)); }); - if (!arg_shapes_defined) { - // This operator cannot be legalized, because legalization - // requires the argument shapes to be known. - // - // TODO(Lunderberg): - // - // Improve this fallback case, as failure to legalize can - // produce unexpected errors during CodeGenVM. This could - // be done by having `R.Tensor(ndim=2)` be syntactic sugar - // for `R.Tensor(shape=[m, n])`, where `m` and `n` are new - // shape variables. This would allow legalization into - // dynamic TIR PrimFuncs. - // - // This fallback would only be applicable for cases where - // both the dtype and the dimensionality are known. While - // Relax can express a tensor with unknown dtype and - // dimensionality as `TensorType(DLDataType{kDLOpaqueHandle, 0, 0}, - // kUnknownNDim)`, TIR cannot express unknown dtype or - // unknown dimensionality. - return false; - } - - bool is_data_dependent_op = [&]() -> bool { - if (Op::HasAttrMap("FDataDependent")) { - auto op_map = Op::GetAttrMap("FDataDependent"); - if (op_map.count(op)) { - return op_map[op]; - } - } - return false; - }(); - bool ret_shape_defined = KnowAllShapeValues(GetType(visited_call)); - if (!is_data_dependent_op && !ret_shape_defined) { - // This operator cannot be legalized, because legalization by - // default requires the output shape. The exception is - // data-dependent operators (e.g. `R.dynamic_strided_slice`), - // where the shape of the output depends on the runtime values - // stored in a tensor. - // - // For data-dependent ops, the output shape will be identified - // at runtime. The Legalizer will insert their shape - // functions, which are manually registered for each - // data-dependent op, and match cast to define symbolic output - // shapes. These symbolic output shapes at compile time can - // be by later operations to refer to the runtime shape. - return false; - } - - // All checks pass, this operator can be legalized. - return true; - }(); + // TODO(Lunderberg): + // + // Improve the "argument shape not known" fallback case in + // `CanLegalizeCall`, as failure to legalize can produce + // unexpected errors during CodeGenVM. This could be done by + // having `R.Tensor(ndim=2)` be syntactic sugar for + // `R.Tensor(shape=[m, n])`, where `m` and `n` are new shape + // variables. This would allow legalization into dynamic TIR + // PrimFuncs. + // + // This fallback would only be applicable for cases where both + // the dtype and the dimensionality are known. While Relax can + // express a tensor with unknown dtype and dimensionality as + // `TensorType(DLDataType{kDLOpaqueHandle, 0, 0}, kUnknownNDim)`, + // TIR cannot express unknown dtype or unknown dimensionality. + // + // For data-dependent ops (e.g. `R.dynamic_strided_slice`), the + // shape of the output depends on the runtime values stored in a + // tensor, and `CanLegalizeCall` allows the return shape to remain + // unknown. The Legalizer will insert their shape functions, which + // are manually registered for each data-dependent op, and match + // cast to define symbolic output shapes. These symbolic output + // shapes at compile time can be used by later operations to refer + // to the runtime shape. + bool shapes_are_known_if_required = CanLegalizeCall(op, visited_call); FLegalize legalization_func; diff --git a/src/relax/transform/utils.h b/src/relax/transform/utils.h index d2af2684fe48..811b1317282b 100644 --- a/src/relax/transform/utils.h +++ b/src/relax/transform/utils.h @@ -200,6 +200,88 @@ bool IsNestedTensor(const Type& ty); */ bool IsNestedTensor(const Expr& expr); +/*! + * \brief Check if a given Tensor/Shape/TupleType contains shapes whose + * values are all known. + * \param ty The Type to be checked. + * \return A boolean indicating the given type contains shape values that are all known. + */ +inline bool KnowAllShapeValues(const Type& ty) { + if (const auto* tensor_ty = ty.as()) { + return tensor_ty->shape.has_value() && tensor_ty->shape.value()->IsInstance(); + } else if (const auto* shape_ty = ty.as()) { + return shape_ty->values.has_value(); + } else if (const auto* tuple_ty = ty.as()) { + return std::all_of(tuple_ty->fields.begin(), tuple_ty->fields.end(), + [](Type field_ty) { return KnowAllShapeValues(field_ty); }); + } else if (ty.as()) { + return true; + } else { + return false; + } +} + +/*! + * \brief Check whether \p op can be legalized (via its registered `FLegalize`) for the + * given \p call right now, i.e. whether the argument/return shapes it needs are already + * statically known. + * + * This mirrors the gate `LegalizeOps` applies before invoking any `FLegalize` function, so + * that other callers of `FLegalize` (e.g. `FoldConstant`'s speculative "would this fold to a + * constant `call_tir`?" probe) do not violate the same invariant `FLegalize` implementations + * are allowed to assume. + * \param op The op to check. + * \param call The (post-order-visited) call to \p op. + * \return Whether \p op's registered `FLegalize` function can be safely invoked on \p call. + */ +inline bool CanLegalizeCall(const Op& op, const Call& call) { + static const auto& requires_arg_shapes_map = Op::GetAttrMap("RequiresArgumentShapes"); + + bool requires_arg_shapes = requires_arg_shapes_map.get(op, true); + if (!requires_arg_shapes) { + // This operator does not require its arguments to have a + // known shape/dtype. For example, the "relax.tensor_ndim" + // operator can output the dimensionality of a tensor at + // runtime, and does not require the dimensionality to be + // known at compile-time. + return true; + } + + bool arg_shapes_defined = std::all_of(call->args.begin(), call->args.end(), + [](Expr arg) { return KnowAllShapeValues(GetType(arg)); }); + if (!arg_shapes_defined) { + return false; + } + + bool is_data_dependent_op = [&]() -> bool { + if (Op::HasAttrMap("FDataDependent")) { + auto op_map = Op::GetAttrMap("FDataDependent"); + if (op_map.count(op)) { + return op_map[op]; + } + } + return false; + }(); + bool ret_shape_defined = KnowAllShapeValues(GetType(call)); + if (!is_data_dependent_op && !ret_shape_defined) { + // This operator cannot be legalized, because legalization by + // default requires the output shape. The exception is + // data-dependent operators (e.g. `R.dynamic_strided_slice`), + // where the shape of the output depends on the runtime values + // stored in a tensor. + // + // For data-dependent ops, the output shape will be identified + // at runtime. The Legalizer will insert their shape + // functions, which are manually registered for each + // data-dependent op, and match cast to define symbolic output + // shapes. These symbolic output shapes at compile time can + // be by later operations to refer to the runtime shape. + return false; + } + + return true; +} + // TODO(@bohan): implements some postorder function accepts a visitor closure class VarReplacer : public ExprMutator { public: diff --git a/tests/python/relax/test_transform_fold_constant.py b/tests/python/relax/test_transform_fold_constant.py index 810873ce4af4..01678c0621db 100644 --- a/tests/python/relax/test_transform_fold_constant.py +++ b/tests/python/relax/test_transform_fold_constant.py @@ -380,6 +380,25 @@ def expected(data: R.Tensor((256,), "float32")) -> R.Tensor((16, 16), dtype="flo tvm.ir.assert_structural_equal(after, expected) +def test_fold_constant_skips_data_dependent_reshape_with_unknown_shape(): + @tvm.script.ir_module + class Module: + @R.function + def before(data: R.Tensor((256,), "float32"), c0: R.Tensor((2,), "int64")): + with R.dataflow(): + lv2: R.Shape(ndim=2) = R.tensor_to_shape(c0) + gv: R.Tensor(ndim=2, dtype="float32") = R.reshape(data, lv2) + R.output(gv) + return gv + + # Deliberately leave `c0` unbound, so its value (and therefore the reshape's + # target shape) is never statically known. + before = gen_mod(Module, "before", {}) + + after = relax.transform.FoldConstant()(before) + tvm.ir.assert_structural_equal(after, before) + + def test_unsupported_fold_ops_legalized_to_multiple_calls(): @tvm.script.ir_module class Module: diff --git a/tests/python/relax/test_transform_legalize_ops_manipulate.py b/tests/python/relax/test_transform_legalize_ops_manipulate.py index 85ad10d49047..44744fc57584 100644 --- a/tests/python/relax/test_transform_legalize_ops_manipulate.py +++ b/tests/python/relax/test_transform_legalize_ops_manipulate.py @@ -759,6 +759,25 @@ def reshape( tvm.ir.assert_structural_equal(out_mod, Expected) +def test_data_dependent_reshape_without_decompose_ops_is_skipped(): + # fmt: off + @tvm.script.ir_module + class DDReshape: + @R.function + def main( + x: R.Tensor([2], dtype="int64"), + y: R.Tensor([16],dtype='float32'), + ): + lv: R.Shape(ndim=2) = R.tensor_to_shape(x) + gv = R.reshape(y, lv) + return gv + # fmt: on + + relax.analysis.well_formed(DDReshape) + out_mod = relax.transform.LegalizeOps()(DDReshape) + tvm.ir.assert_structural_equal(out_mod, DDReshape) + + def test_split_by_indices(): # fmt: off @tvm.script.ir_module