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
7 changes: 6 additions & 1 deletion python/tvm/relax/transform/legalize_ops/manipulate.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
21 changes: 16 additions & 5 deletions src/relax/transform/fold_constant.cc
Original file line number Diff line number Diff line change
Expand Up @@ -30,6 +30,8 @@
#include <tvm/tirx/function.h>
#include <tvm/tirx/op.h>

#include "utils.h"

namespace tvm {
namespace relax {
using namespace tvm::prim;
Expand Down Expand Up @@ -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<PrimTypeNode>() && !infer_type_map.count(op) &&
!infer_type_with_builder_map.count(op)) {
Expand All @@ -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<Call>();
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<CallNode>();
if (call && call->op.same_as(call_tir_op)) {
return VisitCallTIR(ffi::GetRef<Call>(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<CallNode>();
if (call && call->op.same_as(call_tir_op)) {
return VisitCallTIR(ffi::GetRef<Call>(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"
Expand Down
114 changes: 27 additions & 87 deletions src/relax/transform/legalize_ops.cc
Original file line number Diff line number Diff line change
Expand Up @@ -36,6 +36,8 @@

#include <set>

#include "utils.h"

namespace tvm {
namespace relax {

Expand All @@ -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<TensorTypeNode>()) {
return tensor_ty->shape.has_value() && tensor_ty->shape.value()->IsInstance<ShapeExprNode>();
} else if (const auto* shape_ty = ty.as<ShapeTypeNode>()) {
return shape_ty->values.has_value();
} else if (const auto* tuple_ty = ty.as<TupleTypeNode>()) {
return std::all_of(tuple_ty->fields.begin(), tuple_ty->fields.end(),
[](Type field_ty) { return KnowAllShapeValues(field_ty); });
} else if (ty.as<PrimTypeNode>()) {
return true;
} else {
return false;
}
}

class LegalizeMutator : public ExprMutator {
public:
explicit LegalizeMutator(const IRModule& mod,
Expand Down Expand Up @@ -242,7 +223,6 @@ class LegalizeMutator : public ExprMutator {
Call visited_call = this->VisitExprPostOrder_(call).as_or_throw<Call>();
static const auto& legalize_map = Op::GetAttrMap<FLegalize>("FLegalize");
static const auto& call_packed_map = Op::GetAttrMap<FCallPacked>("FCallPacked");
static const auto& requires_arg_shapes_map = Op::GetAttrMap<bool>("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");
Expand All @@ -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<bool>("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;

Expand Down
82 changes: 82 additions & 0 deletions src/relax/transform/utils.h
Original file line number Diff line number Diff line change
Expand Up @@ -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<TensorTypeNode>()) {
return tensor_ty->shape.has_value() && tensor_ty->shape.value()->IsInstance<ShapeExprNode>();
} else if (const auto* shape_ty = ty.as<ShapeTypeNode>()) {
return shape_ty->values.has_value();
} else if (const auto* tuple_ty = ty.as<TupleTypeNode>()) {
return std::all_of(tuple_ty->fields.begin(), tuple_ty->fields.end(),
[](Type field_ty) { return KnowAllShapeValues(field_ty); });
} else if (ty.as<PrimTypeNode>()) {
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<bool>("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<bool>("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:
Expand Down
19 changes: 19 additions & 0 deletions tests/python/relax/test_transform_fold_constant.py
Original file line number Diff line number Diff line change
Expand Up @@ -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:
Expand Down
19 changes: 19 additions & 0 deletions tests/python/relax/test_transform_legalize_ops_manipulate.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
Loading