Lunderberg commented on code in PR #16194: URL: https://github.com/apache/tvm/pull/16194#discussion_r1418204135
########## src/relax/transform/inline_functions.cc: ########## @@ -0,0 +1,228 @@ +/* + * Licensed to the Apache Software Foundation (ASF) under one + * or more contributor license agreements. See the NOTICE file + * distributed with this work for additional information + * regarding copyright ownership. The ASF licenses this file + * to you under the Apache License, Version 2.0 (the + * "License"); you may not use this file except in compliance + * with the License. You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, + * software distributed under the License is distributed on an + * "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY + * KIND, either express or implied. See the License for the + * specific language governing permissions and limitations + * under the License. + */ + +#include <tvm/relax/analysis.h> +#include <tvm/relax/expr.h> +#include <tvm/relax/expr_functor.h> +#include <tvm/relax/transform.h> + +#include <utility> + +#include "../../support/ordered_set.h" +#include "utils.h" + +namespace tvm { +namespace relax { + +namespace { + +class FunctionInliner : public ExprMutator { + public: + explicit FunctionInliner(const Map<Variant<String, GlobalVar>, Function>& replacements) + : replacements_(replacements) {} + + using ExprMutator::VisitExpr_; + + Expr VisitExpr_(const FunctionNode* op) override { + auto node = ExprMutator::VisitExpr_(op); + if (node.get() != op) { + node = CanonicalizeBindings(node); + node = RemoveAllUnused(node); + } + return node; + } + + Expr VisitExpr_(const CallNode* op) override { + auto node = Downcast<Call>(ExprMutator::VisitExpr_(op)); + + if (auto opt = node->op.as<GlobalVar>()) { + auto gvar = opt.value(); + if (auto opt = GetFunction(gvar)) { + auto callee = opt.value(); + CHECK_EQ(callee->params.size(), node->args.size()) + << "Attempted to inline call to " << gvar << ", which accepts " << callee->params.size() + << " parameters. " + << "However, it was called with " << node->args.size() << " arguments in expression " + << node; + + Expr inlined = InlinedCall(callee, node->args); + + CHECK(!inline_stack_.count(gvar)) + << "Relax function inlining does not support recursive functions. " + << "However, recursive function " << gvar << " was requested to be inlined."; + + inline_stack_.insert(gvar); + inlined = VisitExpr(std::move(inlined)); + inline_stack_.erase(gvar); + + return inlined; + } + } + + return std::move(node); + } + + private: + Optional<Function> GetFunction(const GlobalVar& gvar) const { + if (auto opt = replacements_.Get(gvar)) { + return opt; + } else if (auto opt = replacements_.Get(gvar->name_hint)) { + return opt; + } else { + return NullOpt; + } + } + + Expr InlinedCall(Function func, const Array<Expr>& args) const { + // Ensures that the inlined instance does not have duplicate usage + // with other inlined copies, or with the original callee. + func = CopyWithNewVars(std::move(func)); + + Array<Binding> param_bindings; + + Map<Var, Expr> param_map; + for (size_t i = 0; i < args.size(); i++) { + // Option 1: Use tvm::relax::Bind to substitute arguments into + // the body. If the arguments contain DataflowVar instances, + // but the subroutine does not use DataflowBlock, this would + // result in invalid AST. + // + // Option 2: Define a VarBinding `param[i] = args[i]` for each + // parameter, then rely on CanonicalizeBindings to replace with + // DataflowVar where possible. This would solve the invalid use + // of DataflowVar, but wouldn't handle symbolic variables. If + // the subroutine has symbolic variables defined by its + // arguments, the VarBinding would leave them undefined. + // + // Option 3: Define a MatchCast `param[i] = args[i]` for each + // parameter, followed by CanonicalizeBindings. This is the + // first option that would result in well-formed AST, but it + // wouldn't be optimal. Symbolic variables would have two + // copies, one from the initial definition, and one + // from the MatchCast inlined portion. + // + // Option 4: Define a VarBinding `param[i] = args[i]`, with + // CanonicalizeBindings to handle conversion of Var to + // DataflowVar, and tvm::relax::Bind to handle substitution of + // symbolic variables. This would result in a well-formed Relax + // function, with no duplicate definitions of symbolic + // variables. + // + // This implementation uses Option 4. Review Comment: Thank you. I may have gone a bit overboard with this comment, but I wanted to make sure that it was "clever robust" and not "clever unmaintainable". I'm also really liking the use of CanonicalizeBindings as a general utility, since the rules about `DataflowVar` have too many edge cases to duplicate the logic everywhere. -- This is an automated message from the Apache Git Service. To respond to the message, please log on to GitHub and use the URL above to go to the specific comment. To unsubscribe, e-mail: [email protected] For queries about this service, please contact Infrastructure at: [email protected]
