tvm
Loading...
Searching...
No Matches
transform.h
Go to the documentation of this file.
1/*
2 * Licensed to the Apache Software Foundation (ASF) under one
3 * or more contributor license agreements. See the NOTICE file
4 * distributed with this work for additional information
5 * regarding copyright ownership. The ASF licenses this file
6 * to you under the Apache License, Version 2.0 (the
7 * "License"); you may not use this file except in compliance
8 * with the License. You may obtain a copy of the License at
9 *
10 * http://www.apache.org/licenses/LICENSE-2.0
11 *
12 * Unless required by applicable law or agreed to in writing,
13 * software distributed under the License is distributed on an
14 * "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY
15 * KIND, either express or implied. See the License for the
16 * specific language governing permissions and limitations
17 * under the License.
18 */
19
24#ifndef TVM_TIR_TRANSFORM_H_
25#define TVM_TIR_TRANSFORM_H_
26
27#include <tvm/ir/prim/expr.h>
28#include <tvm/ir/transform.h>
29#include <tvm/target/target.h>
30#include <tvm/tirx/function.h>
31
32#include <string>
33#include <vector>
34
35namespace tvm {
36namespace tirx {
37namespace transform {
38
47
48/*
49 * \brief Create a function pass that optimizes PrimFuncs.
50 *
51 * \param pass_func The packed function that contains the optimization.
52 * \param opt_level The optimization level of the function pass.
53 * \param name The name of the function pass.
54 * \param required The list of the passes that the function pass is dependent on.
55 *
56 * \return The created function pass.
57 */
59 int opt_level, ffi::String name,
60 tvm::ffi::Array<ffi::String> required, bool traceable = false);
61
70
80
88
95
102
115
143
154TVM_DLL Pass RemapThreadAxis(ffi::Map<ffi::String, IterVar> axis_map);
155
175
182
186static constexpr const char* kDisableLowerTVMBuiltin = "disable_lower_builtin";
187
193
200
206
216
224
231
240
246
253
260
270
277
284
292
298
304
309TVM_DLL Pass Filter(ffi::TypedFunction<bool(PrimFunc)> fcond);
310
320
326
334
340
341} // namespace transform
342} // namespace tirx
343} // namespace tvm
344
345#endif // TVM_TIR_TRANSFORM_H_
Managed reference class to IRModuleNode.
Definition module.h:255
Managed reference class to TargetNode.
Definition target.h:134
RAII wrapper function to enter and exit a context object similar to python's with syntax.
Definition with_context.h:59
Managed reference to PrimFuncNode.
Definition function.h:105
PassContextNode contains the information that a pass can rely on, such as analysis results.
Definition transform.h:80
PassContext that is used to configure the pass behavior.
Definition transform.h:151
Meta data that will be used to help optimization and analysis.
Definition transform.h:336
Managed reference class for PassInfoNode.
Definition transform.h:367
PassNode is the base type of differnt types of optimization passes. It is designed as a pure class an...
Definition transform.h:387
Definition transform.h:417
Definition transform.h:508
TIR expressions.
Pass SplitHostDevice()
Annotate, split, and lower host/device functions.
Pass PointerValueTypeRewrite()
Rewrite the pointer content type of arguments, as well as Alloc internal to the function to use the m...
Pass LowerTIRx()
Lower the TIR to a lower level IR for the given target.
Pass SkipAssert()
skip assert stmt.
Pass RemapThreadAxis(ffi::Map< ffi::String, IterVar > axis_map)
Remap the thread axis.
Pass BindTarget(Target target)
Annotate a PrimFunc with a given target.
Pass AnnotateEntryFunc()
Set a PrimFunc as the entry point if it is only function in IRModule.
Pass CommonSubexprElim()
Implements Common Subexpression Elimination (CSE) for TIR which introduces Bind statements for duplic...
Pass MakePackedAPI()
Transform the high-level PrimFunc to a low-level version that can be used as an API function.
Pass LowerTIRxOpaque()
Lower opaque constructs in TIRX programs: AllocBuffer, For(thread_binding), unit loop elimination....
Pass FlattenBuffer()
Flatten the multi-dimensional TensorLoad and BufferStore to single dimensional BufferLoad/BufferStore...
Pass StorageRewrite()
Rewrite storage allocation pattern. Moves the allocation to outer most possible scope....
Pass ConvertSSA()
Convert an IRModule to be SSA form.
Pass LowerTIRxCleanup()
Finalize TIRx lowering by applying layout rewriters and cleanup passes.
Pass VectorizeLoop(bool enable_vectorize=true)
Lower vectorization loops.
Pass UnifiedStaticMemoryPlanner()
This is the unified static memory planner pass that will plan for memory intra- and inter- PrimFuncs ...
Pass LowerWarpMemory()
Lower warp memory access to low-level device related function calls.
Pass StmtSimplify()
Run statement-level arithmetic simplifications on the TIR PrimFunc.
Pass BF16ComputeLegalize()
Legalize bf16 compute Ops. Add a cast to fp32 before Ops, then add a cast back to bf16.
Pass CreatePrimFuncPass(std::function< PrimFunc(PrimFunc, IRModule, PassContext)> pass_func, int opt_level, ffi::String name, tvm::ffi::Array< ffi::String > required, bool traceable=false)
Pass NarrowDataType(int target_bits)
Narrow down PrimExpr datatype in stmt to target_bits.
Pass RemoveNoOp()
Remove No Op from the Stmt.
Pass TilePrimitiveDispatch()
Lower TIRx op calls using registered op dispatchers for the given target.
Pass InlinePrivateFunctions()
Inline calls to private functions.
Pass LowerIntrin()
Lower the target specific function intrinsics in each of the function.
Pass Filter(ffi::TypedFunction< bool(PrimFunc)> fcond)
Filter PrimFuncs with a given condition.
Pass UnrollLoop()
unroll the constant loop marked by unroll. This pass also automatically attach pragma unroll tag to l...
Pass FP8ComputeLegalize(ffi::String promote_dtype="float16")
Legalize fp8 compute Ops. Add a cast to fp16/fp32 before Ops, then add a cast back to fp8.
Pass LowerTVMBuiltin()
Lower builtin intrinsics.
Pass BF16StorageLegalize()
Legalize bf16 storage types to u16.
Pass FP8StorageLegalize()
Legalize fp8 storage types to u8.
Pass ForceNarrowIndexToInt32()
Force to narrow down indexing expressions and integer buffers to int32 dtype.
Pass CreateModulePass(std::function< IRModule(IRModule, PassContext)> pass_func, int opt_level, ffi::String name, ffi::Array< ffi::String > required, bool traceable=false)
An object that builds and maintains block scope and StmtSref mapping for Dependence analysis.
Definition analyzer.h:40
Compilation target object.
TIR Function.