Pass Infrastructure#
A pass transforms an IRModule into another IRModule. PassContext
provides scoped configuration and instrumentation, and Sequential composes
passes in an explicit order. The same utilities support transformations across
IR dialects and allow pass implementations in C++ or Python.
Passes do not declare dependencies or discover prerequisite passes by name. A pipeline supplies all of its passes in the order they should execute.
C++ Backend#
PassInfo contains a pass’s name and opt_level. The optimization
level controls whether a pass in a Sequential runs under the current
PassContext. See include/tvm/ir/transform.h for the core interfaces.
class PassInfoNode : public Object {
int opt_level;
ffi::String name;
};
PassContext#
PassContext configures transformations with an optimization level, explicit
required_pass and disabled_pass names, pass-specific configuration, and
instruments. These name lists control passes already present in a pipeline;
they do not create or insert passes. See Pass Instrument
for instrumentation behavior.
This class is designed for users to conveniently write the Python with
syntax to perform optimizations under a certain configuration. In addition, the
users can obtain the context that is available within a certain program scope in
a thread-safe way through PassContext::Current(), since a thread-local store
PassContextThreadLocalStore is used to hold the created pass context
objects. Examples will be provided later to show how we can use both the C++ and
Python APIs to create a compilation pipeline using pass context.
class PassContextNode : public Object {
public:
int opt_level{2};
ffi::Array<ffi::String> required_pass;
ffi::Array<ffi::String> disabled_pass;
ffi::Map<ffi::String, Any> config;
ffi::Array<PassInstrument> instruments;
};
class PassContext : public ObjectRef {
public:
TVM_DLL static PassContext Create();
TVM_DLL static PassContext Current();
TVM_DLL void InstrumentEnterPassContext();
TVM_DLL void InstrumentExitPassContext();
TVM_DLL bool InstrumentBeforePass(const IRModule& mod, const PassInfo& info) const;
TVM_DLL void InstrumentAfterPass(const IRModule& mod, const PassInfo& info) const;
/* Other fields are omitted. */
private:
// The entry of a pass context scope.
TVM_DLL void EnterWithScope();
// The exit of a pass context scope.
TVM_DLL void ExitWithScope();
// Classes to get the Python `with` like syntax.
friend class tvm::With<PassContext>;
};
struct PassContextThreadLocalEntry {
/*! \brief The default pass context. */
PassContext default_context;
/*! \brief The current pass context. */
std::stack<PassContext> context_stack;
PassContextThreadLocalEntry() {
default_context = PassContext(ffi::make_object<PassContextNode>());
}
};
Pass Constructs#
The pass infra is designed in a hierarchical manner, and it could work at
different granularities of Relax/TensorIR programs. A pure virtual class PassNode is
introduced to serve as the base of the different optimization passes. This class
contains several virtual methods that must be implemented by the
subclasses at the level of modules, functions, or sequences of passes.
class PassNode : Object {
virtual PassInfo Info() const = 0;
virtual Module operator()(const IRModule& mod
const PassContext& pass_ctx) const = 0;
};
The functor shows how a pass must be realized, i.e. it always works on a
IRModule under a certain context. All passes are designed in a Module to Module
manner. Therefore, optimizations governed by the pass infra will
always update the whole module.
Several subclasses have been created to implement different types of optimization passes, e.g., function-level passes, module-level passes, and sequential passes. Each adapts its transformation to the same module-to-module interface. Their definitions can be found in src/ir/transform.cc.
Module-Level Passes#
Module level passes are geared mainly for global and inter-procedural optimizations (IPO), which are similar to the module pass used in LLVM. Some typical passes in Relax that need the global picture of a module, such as A-normal form conversion and lambda lifting, etc., fall into this set. At this level, users can even add and/or delete functions in a module. Note that all passes
class ModulePassNode : PassNode {
PassInfo pass_info;
std::function<Module(Module, PassContext)> pass_func;
Module operator()(const Module& mod, const PassContext& pass_ctx) const final;
// Other members/methods are omitted
};
pass_info maintains the information needed by a module-level pass.
pass_func sketches the real optimization. For example, we may need to
perform dead code elimination on the module. We could implement the algorithm in
the pass_func and let it run on a module. It will then remove the dead code
including the unused functions in the module. Note that this field is designed
as a packed function, which enables the implementation of the optimization in
both C++ and Python.
Function-Level Passes#
Function-level passes are used to implement various intra-function level
optimizations for a given Relax/TensorIR module. It fetches one function at a time from
the function list of a module for optimization and yields a rewritten Relax
Function or TensorIR Function. Most of passes can be classified into this category, such as
common subexpression elimination and inference simplification in Relax as well as vectorization
and flattening storage in TensorIR, etc.
Note that the scope of passes at this level is either a Relax function or a TensorIR primitive function. Therefore, we cannot add or delete a function through these passes as they are not aware of the global information.
class FunctionPassNode : PassNode {
PassInfo pass_info;
std::function<Function(Function, Module, PassContext)> pass_func;
Module operator()(const Module& mod, const PassContext& pass_ctx) const final;
bool SkipFunction(const Function& func) const;
// Other members/methods are omitted...
};
pass_info is identical to what we just described in the module pass.
pass_func takes a function for optimization, it also needs a module as we
may use it for reporting errors. A function could be annotated with
“SkipOptimization” so that it will be ignored during optimization.
Sequential Passes#
Sequential applies its passes in the supplied order. A nested
Sequential preserves the same ordered composition.
class SequentialNode : PassNode {
PassInfo pass_info;
// Passes need to be executed.
ffi::Array<Pass> passes;
Module operator()(const Module& mod, const PassContext& pass_ctx) const final;
};
The following code shows how individual passes in a sequential pass are invoked. Essentially, we sequentially execute each pass in a sequential pass using the order that they were appended to the pass list.
Module SequentialNode::operator()(const Module& module,
const PassContext& pass_ctx) const {
Module mod = module;
for (const Pass& pass : passes) {
TVM_FFI_ICHECK(pass.defined()) << "Found undefined pass for optimization.";
const PassInfo& pass_info = pass->Info();
if (!pass_ctx.PassEnabled(pass_info)) continue;
mod = pass(mod, pass_ctx);
}
return mod;
}
Upon the invocation of a pass, we first check if this pass is enabled. This is
done by first checking if the pass is explicitly disabled by a user, followed by
inspecting if it is specified as a required pass by the user. If it is still
undetermined whether this pass is enabled, its opt_level will be checked.
This pass will be enabled and therefore executed only when its optimization
level is at most the configured optimization level in the pass context.
Callers construct the passes directly, including any prerequisites, before creating a sequence. For example:
pipeline = tvm.transform.Sequential([
relax.transform.Normalize(),
relax.transform.FoldConstant(),
])
Some helper functions are provided to create each type of these aforementioned passes. These helpers are also exposed to the Python frontend for users to favorably use Python APIs to create a specific pass object.
Pass CreateFunctionPass(
std::function<Function(Function, IRModule, PassContext)> pass_func,
int opt_level,
ffi::String name);
Pass CreateFunctionPass(
std::function<Function(Function, IRModule, PassContext)> pass_func,
int opt_level,
ffi::String name);
Pass CreateModulePass(
std::function<IRModule(IRModule, PassContext)> pass_func,
int opt_level,
ffi::String name);
Pass Sequential(tvm::ffi::Array<Pass> passes, PassInfo pass_info);
Pass Registration#
We’ve covered the concept of different level of passes and the context used for compilation. It would be interesting to see how easily users can register a pass. Let’s take const folding as an example. This pass has already been implemented to fold constants in a Relax function (found in src/relax/transform/fold_constant.cc).
An API was provided to perform the Expr to Expr transformation.
Expr FoldConstant(const Expr& expr);
In order to register this pass to the pass infra, we first need to decide at
which level this pass will be performed. As const folding happens on individual
functions, we should intuitively create a FunctionPass for it through
CreateFunctionPass. The pass_func is returned as a packed function that
invokes the Expr to Expr API on each function in an IRModule.
Meanwhile, a pass API endpoint is registered with the name
"relax.transform.FoldConstant". This pass, therefore, becomes an entry in the
registry that exposes the factory to Python. C++ callers use the declared
factory directly.
namespace transform {
Pass FoldConstant() {
auto pass_func =
[=](Function f, IRModule m, PassContext pc) { return ConstantFolder::Fold(f, m); };
return CreateFunctionPass(pass_func, 0, "FoldConstant");
}
TVM_FFI_STATIC_INIT_BLOCK() {
namespace refl = tvm::ffi::reflection;
refl::GlobalDef().def("relax.transform.FoldConstant", FoldConstant);
}
} // namespace transform
To allow other C++ modules to apply this pass, we declare a free function in include/tvm/relax/transform.h as the following:
TVM_DLL Pass FoldConstant();
Pass Instrument#
Pass Instrument is a mechanism to analyze the pass itself. For example, we can use the infrastructure to know how much time and memory a pass requires or how a pass can transform the IR module.
We introduce four instrument points in the life-cycle of PassContext.
TVM_DLL void InstrumentEnterPassContext();
TVM_DLL void InstrumentExitPassContext();
TVM_DLL bool InstrumentBeforePass(const IRModule& mod, const PassInfo& info) const;
TVM_DLL void InstrumentAfterPass(const IRModule& mod, const PassInfo& info) const;
InstrumentEnterPassContext is called immediately when entering the scope
of the PassContext instance.
InstrumentExitPassContext is called when leaving the scope of PassContext,
or exceptions occur during the execution of passes.
This method is also called when instruments is being overridden by override_instruments in tvm.transform.PassContext.
See Override Instruments in Current PassContext.
InstrumentBeforePass is called before execution.
InstrumentAfterPass is called after execution if the pass should be run. The behavior is like:
if (pass_ctx.InstrumentBeforePass(ir_module, pass_info)) {
new_ir_module = run_pass(ir_module, pass_ctx);
pass_ctx.InstrumentAfterPass(new_ir_module, pass_info);
return new_ir_module;
}
The PassInstrument interface allow you to run arbitrary code inside above four methods.
Multiple PassInstrument instances can be registed into a single
PassContext. PassInstrument instances are called sequentially in the order of
instruments argument passed to PassContext.
PassInstrument provides following interfaces:
namespace transform {
class PassInstrumentNode : public Object {
public:
ffi::String name;
virtual void EnterPassContext() const = 0;
virtual void ExitPassContext() const = 0;
virtual bool ShouldRun(const IRModule& mod, const PassInfo& info) const = 0;
virtual void RunBeforePass(const IRModule& mod, const PassInfo& info) const = 0;
virtual void RunAfterPass(const IRModule& mod, const PassInfo& info) const = 0;
/* Other fields are omitted. */
};
class PassInstrument : public ObjectRef {
public:
TVM_FFI_DEFINE_OBJECT_REF_METHODS_NULLABLE(PassInstrument, ObjectRef, PassInstrumentNode);
};
} // namespace transform
Python frontend are provided to implement PassInstrument quickly. See Pass Instrument.
Within a PassContext, the call sequence of a PassInstrument instance is like:
with PassContext(instruments=[pi]) # pi = a PassInstrument implementation.
pi.EnterPassContext()
if pi.ShouldRun(Pass1):
pi.RunBeforePass()
Pass1()
pi.RunAfterPass()
if pi.ShouldRun(Pass2):
pi.RunBeforePass()
Pass2()
pi.RunAfterPass()
pi.ExitPassContext()
Here is a brief introduction of relations between PassInstrument interfaces
and PassContext methods. See (src/ir/transform.cc) for more details.
InstrumentEnterPassContextEnterPassContext()is executed in the order ofinstrumentspassed to thePassContext.When an exception raises,
PassContextdisable the pass instrumentation by clearing all registeredPassInstrumentinstances.Then
PassContextexecuteExitPassContext()method of eachPassInstrumentinstances which successfully finishedEnterPassContext()For example, if
PassInstrumentA, B, and C are registered to aPassContextand A finishedEnterPassContext()while B throws an exception, then C is never executed;ExitPassContext()of A is executed.
InstrumentExitPassContextExitPassContext()of eachPassInstrumentinstances are executed in the order ofinstrumentspassed to thePassContext.While an exception occurs,
instrumentsis cleared.PassInstrumentInstances registered after the one throwing exceptions do not executeExitPassContext.
InstrumentBeforePassShouldRunis executed if the pass is not listed as a required pass.RunBeforePassis executed in the order ofinstrumentsif the pass is not blocked byShouldRun.Note that
InstrumentBeforePassreturns a boolean indicating whether or not the pass should be run.When an exception occur, it is thrown immediately. We rely on Python Context Manager to exit
PassContextsafely (meaningExitPassContextof each instruments will be run. For C++, please refer to include/tvm/support/with.h.)
InstrumentAfterPassRunAfterPassis executed in the order ofinstrumentspassed to thePassContext.When an exception occur, it is thrown immediately. We rely on Python Context Manager or
Withclass(include/tvm/support/with.h) to exitPassContextsafely
Built-in Instrument#
There are several built-in instruments.
PassTimingInstrument (see src/ir/instrument.cc)
Profile the execution time of passes.
PrintBeforeAll (see python/tvm/transform/instrument.py)
Print the IR module and pass info before each pass executes.
PrintAfterAll (see python/tvm/transform/instrument.py)
Print the IR module and pass info after each pass executes.
PassPrintingInstrument (see python/tvm/transform/instrument.py)
Selectively print the IR module before or after specific named passes.
DumpIR (see python/tvm/transform/instrument.py)
Dump the IR module to files after each pass executes.
Python Frontend#
Only some simple APIs are needed for the frontend side. For example, we can provide users the following APIs to create and execute a pass (full implementation is provided in python/tvm/relax/transform/transform.py and python/tvm/transform/core.py). The backend receives the information and decides which function it should use to create a Pass object.
PassContext#
Python frontend provides a wrapper for the PassContext to enable the
with syntax by overriding __enter__ and __exit__. A current
static method is offered for users to get the context that is in use under
a certain scope.
@tvm_ffi.register_object("transform.PassContext")
class PassContext(tvm.runtime.Object):
def __enter__(self):
_transform.EnterPassContext(self)
return self
def __exit__(self, ptype, value, trace, config):
_transform.ExitPassContext(self)
@staticmethod
def current():
"""Return the current pass context."""
return _transform.GetCurrentPassContext()
A PassContext is used to configure the compilation options, including the
optimization level and required/disabled passes. It can also take a dictionary
of configs so that different passes can conveniently fetch the passed data, such
as fallback device info and loop-unrolling limits. Register each configuration
key and value type directly in a static initialization block:
TVM_FFI_STATIC_INIT_BLOCK() {
tvm::transform::PassContext::RegisterConfigOption<UnrollLoopConfig>("tirx.UnrollLoop");
}
GetConfig<T>(key) returns an optional value. Scalar and target-dependent
defaults are explicit at the call site. GetConfigOrDefault<TConfig>(key)
returns a configured object or, only when absent, constructs a fresh object
using its reflection-defined defaults. Both APIs take only the key.
auto ctx = tvm::transform::PassContext::Current();
bool noalias = ctx->GetConfig<bool>("tirx.noalias").value_or(true);
auto unroll = ctx->GetConfigOrDefault<UnrollLoopConfig>("tirx.UnrollLoop");
Please refer to src/tirx/transform/unroll_loop.cc for more details.
Pass Instrument#
One can implement a PassInstrument by using the pass_instrument
decorator (python/tvm/transform/core.py) or by subclassing
tvm.transform.PassInstrument. Implement any of the following callbacks:
enter_pass_ctxThis method is run when entering
PassContext.
exit_pass_ctxThis method is run when exiting
PassContext.
should_runThis method is run before a pass is executed, returning a boolean indicating whether or not the pass should be run.
run_before_passIf a pass should be run, this method is run just before pass execution.
run_after_passThis method is run right after a pass has been executed.
PassInstrument instances can be registered through instruments argument in
tvm.transform.PassContext.
Core interfaces are exported from tvm.transform. Concrete tools, including
PassTimingInstrument and DumpIR, live in tvm.transform.instrument;
see python/tvm/transform/instrument.py for examples.
Override Instruments in Current PassContext#
override_instruments method is provided to override the instruments of current PassContext.
For example, if passes are run without explicitly creating a new PassContext,
one can still register PassInstrument into the global PassContext by:
cur_pass_ctx = tvm.transform.PassContext.current()
# override PassInstrument instances
cur_pass_ctx.override_instruments([pass_inst])
mod = pass_seq(mod)
result = pass_inst.get_result()
Note that when override_instruments is called, the exit_pass_ctx method of
old PassInstrument instances are called. Then the enter_pass_ctx method of
new PassInstrument are called.