tvm
Loading...
Searching...
No Matches
analysis.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_S_TIR_ANALYSIS_H_
25#define TVM_S_TIR_ANALYSIS_H_
26
27#include <tvm/ir/module.h>
28#include <tvm/ir/prim/expr.h>
29#include <tvm/ir/transform.h>
30#include <tvm/target/target.h>
31#include <tvm/tirx/function.h>
32#include <tvm/tirx/stmt.h>
33
34#include <optional>
35
36namespace tvm {
37namespace tirx {
38
51TVM_DLL ffi::Array<ffi::Array<BufferRegion>> GetSBlockAccessRegion(
52 const SBlock& block, const ffi::Map<Var, BufferVar>& buffer_var_map);
53
62TVM_DLL ffi::Array<ffi::Array<BufferRegion>> GetSBlockReadWriteRegion(
63 const SBlock& block, const ffi::Map<Var, BufferVar>& buffer_var_map);
64
73TVM_DLL ffi::Map<BufferVar, ffi::Optional<Stmt>> DetectBufferAccessLCA(const PrimFunc& func);
74
90
91} // namespace tirx
92
93namespace arith {
94class AnalyzerObj;
95class Analyzer;
96} // namespace arith
97
98namespace s_tir {
99using namespace tvm::prim;
100
101using namespace tvm::tirx;
102
108TVM_DLL double EstimateTIRFlops(const Stmt& stmt);
109
116
123TVM_DLL bool IsPureFunction(const PrimFunc& func, bool assert_on_error = false);
124
131TVM_DLL bool VerifyGPUCode(const PrimFunc& func, ffi::Map<ffi::String, PrimExpr> constraints);
132
138
144TVM_DLL std::optional<MemCpyDetails> IdentifyMemCpy(const For& loop,
146
152TVM_DLL ffi::Map<ffi::String, ffi::Map<ffi::String, int64_t>> CalculateAllocatedBytes(
153 const PrimFunc& func);
154
160TVM_DLL ffi::Map<ffi::String, ffi::Map<ffi::String, int64_t>> CalculateAllocatedBytes(
161 const IRModule& mod);
162
167TVM_DLL ffi::Array<tvm::transform::Pass> GetVTCMCompactionPasses();
168
176
184
185namespace transform {
186
189
195TVM_DLL Pass VerifyGPUCode(ffi::Map<ffi::String, PrimExpr> constraints);
196
202TVM_DLL Pass VerifyVTCMLimit(ffi::Optional<Target> default_target = std::nullopt);
203
209
210} // namespace transform
211} // namespace s_tir
212} // namespace tvm
213#endif // TVM_S_TIR_ANALYSIS_H_
Managed reference class to IRModuleNode.
Definition module.h:255
RAII wrapper function to enter and exit a context object similar to python's with syntax.
Definition with_context.h:59
Managed reference to AnalyzerObj.
Definition analyzer.h:931
Managed reference to BufferRegionNode.
Definition buffer_region.h:66
Managed reference to ForNode.
Definition stmt.h:644
Managed reference to PrimFuncNode.
Definition function.h:105
A block is a basic schedule unit in TIR.
Definition stmt.h:833
Managed reference to SBlockNode.
Definition stmt.h:881
Container of all statements.
Definition stmt.h:67
PassContext that is used to configure the pass behavior.
Definition transform.h:151
Definition transform.h:417
TIR expressions.
IRModule that holds the functions and type definitions.
Definition builtin.h:25
Pass VerifyGPUCode(ffi::Map< ffi::String, PrimExpr > constraints)
Pass to verify GPU code constraints.
Pass VerifyVTCMLimit(ffi::Optional< Target > default_target=std::nullopt)
Pass to check if VTCM usage is within limit.
Pass OOBChecker()
Statically check TIR code for out of bounds array access.
std::optional< MemCpyDetails > IdentifyMemCpy(const For &loop, const arith::Analyzer &analyzer)
Identify whether a For loop is semantically equivalent to MemCpy.
bool VerifyGPUCode(const PrimFunc &func, ffi::Map< ffi::String, PrimExpr > constraints)
Verify the correctness of a GPU code.
bool IsPureFunction(const PrimFunc &func, bool assert_on_error=false)
Analyze the side effect of a function.
bool VerifyVTCMLimit(const IRModule &mod, int64_t limit)
Verifies that the VTCM usage for all prim_funcs in the given IRModule.
ffi::Array< tvm::transform::Pass > GetVTCMCompactionPasses()
Get the list of lowering passes to calculate the compacted VTCM allocation size.
double EstimateTIRFlops(const Stmt &stmt)
Estimate the FLOPs of a TIR fragment.
ffi::Map< ffi::String, ffi::Map< ffi::String, int64_t > > CalculateAllocatedBytes(const PrimFunc &func)
Calculate the allocated memory per scope in bytes needed inside the TIR PrimFunc.
Definition axis_group_graph.h:39
ffi::Array< ffi::Array< BufferRegion > > GetSBlockAccessRegion(const SBlock &block, const ffi::Map< Var, BufferVar > &buffer_var_map)
Auto detect the block access region according to its body stmt It will detect the access region as an...
ffi::Map< BufferVar, ffi::Optional< Stmt > > DetectBufferAccessLCA(const PrimFunc &func)
Detect the lowest common ancestor(LCA) of buffer access, including both high-level access(BufferLoad,...
const tirx::SBlockNode * FindAnchorBlock(const IRModule &mod)
Find the "anchor block" of the given module. We define the anchor block to be the block with (1) an i...
ffi::Array< ffi::Array< BufferRegion > > GetSBlockReadWriteRegion(const SBlock &block, const ffi::Map< Var, BufferVar > &buffer_var_map)
Auto detect the block read/write region according to its body stmt. An opaque access will be counted ...
An object that builds and maintains block scope and StmtSref mapping for Dependence analysis.
Definition analyzer.h:40
Helper struct for return value of IdentifyMemCpy.
Definition analysis.h:134
BufferRegion source
Definition analysis.h:135
BufferRegion dest
Definition analysis.h:136
Compilation target object.
TIR Function.
TIR statements.