惯性聚合 高效追踪和阅读你感兴趣的博客、新闻、科技资讯
阅读原文 在惯性聚合中打开

推荐订阅源

A
About on SuperTechFans
WordPress大学
WordPress大学
雷峰网
雷峰网
Threat Intelligence Blog | Flashpoint
Threat Intelligence Blog | Flashpoint
Latest news
Latest news
Spread Privacy
Spread Privacy
T
Threat Research - Cisco Blogs
T
Tor Project blog
博客园 - Franky
U
Unit 42
K
Kaspersky official blog
博客园_首页
G
GRAHAM CLULEY
美团技术团队
I
Intezer
T
The Exploit Database - CXSecurity.com
P
Proofpoint News Feed
Engineering at Meta
Engineering at Meta
The Hacker News
The Hacker News
B
Blog
云风的 BLOG
云风的 BLOG
cs.CL updates on arXiv.org
cs.CL updates on arXiv.org
S
Securelist
Last Week in AI
Last Week in AI
F
Fortinet All Blogs
N
Netflix TechBlog - Medium
M
MIT News - Artificial intelligence
Martin Fowler
Martin Fowler
Schneier on Security
Schneier on Security
cs.AI updates on arXiv.org
cs.AI updates on arXiv.org
S
SegmentFault 最新的问题
W
WeLiveSecurity
Cyber Security Advisories - MS-ISAC
Cyber Security Advisories - MS-ISAC
J
Java Code Geeks
Cyberwarzone
Cyberwarzone
爱范儿
爱范儿
K
KPMG report finds enterprise disconnect between AI and its ROI | CIO
The GitHub Blog
The GitHub Blog
S
Security @ Cisco Blogs
GbyAI
GbyAI
P
Proofpoint News Feed
D
Docker
D
DataBreaches.Net
aimingoo的专栏
aimingoo的专栏
博客园 - 司徒正美
T
Tenable Blog
V2EX - 技术
V2EX - 技术
The Register - Security
The Register - Security
V
Vulnerabilities – Threatpost
B
Blog RSS Feed

Пусть этот камень будет более крепким, чем человек

【琐记】烟火与尘埃 【Triton】Triton实现矩阵乘 【LLM推理加速】FlashAttention 【LLM推理加速】PagedAttention 【LLM推理加速】Online Softmax LLM基础知识【1】 Transformer模型 【AI编译】LayerGroup Tiling Tile的疑惑和思考 【AI编译】深度优先的Tile调度,万事大吉? 【AI编译】多级流水线Tile调度策略 【CUDA C++】GPU内存使用【3】 【AI编译】Cache缓存地址映射 【CUDA C++】GPU存储【2】 【CUDA C++】GPU基本介绍【1】 【00】0序章-不受欢迎的来客 【转载】我来了——持续低熵 【Halide】调度优化【2】 【感想】写作进度报告5 【Halide】调度优化【1】 【转载】北大中文男足战报2 【BYOC】TVM切分子图 【转载】北大中文男足战报1 【AI编译】张量生命周期管理 SystemC 用寄存器同步建模方法 【脉动阵列】脉动阵列类型 【im2col】AScend conv accelerate 【感想】写作进度报告4 【BYOC】TVM添加自定义编译器 ccompiler 【感想】写作进度报告3 【Tengine】推理流程脑图【2】 【Tengine】推理流程脑图【1】 【NCNN】学习ncnn模型转换 【编译器】使用llvm编译自定义语言【3】编译 object 【编译器】使用llvm编译自定义语言【2】转llvm IR 【编译器】使用llvm编译自定义语言【1】构建AST 【AI编译】如何进行内存分配 【感想】写作进度报告2 【AI编译】layer-group之后如何tiling 【AI编译】如何进行layer-group 【量化】连续卷积层首尾量化的可行性 【Gemm】内存对齐 【gemm】Gemm计算加速 【TVM】通过代码学习编译流程【5】FuseOps 【TVM】通过代码学习编译流程【6】CodeGen 【TVM】通过代码学习编译流程【4】BuildRelay 【AI编译】Tiling操作能优化什么时间 【TVM】通过代码学习编译流程【3】模型编译 【TVM】通过代码学习编译流程【2】模型转换 【TVM】通过代码学习编译流程【1】必要知识 【感想】写作进度报告1 【Winograd】卷积加速算法原理及实现 SystemC 等待异步事件解决方案 【TVM】Python脚本实现模型编译和保存 【推理引擎】常见AI推理框架 【3D建模】T110E3卡迪夫蓝调皮肤模型 【TVM】C++部署运行TVM 【推理引擎】NCNN和Tengine量化推理逻辑对比 【3D建模】IS-7攻城锤流纹岩皮肤展示 【TVM】根据例子走通代码库 博客汇总目录 【Im2Col】卷积加速算法【2】NHWC 【Im2Col】卷积加速算法【1】 NCHW openBlas库的安装与简单使用 C语言工程调用Cpp库解决方案 foo Hello World
【TVM】通过代码学习类【3.5】Pass
Post author: XianMu@Пусть этот камень будет более крепким, чем ч · 2024-10-22 · via Пусть этот камень будет более крепким, чем человек

# 前言

文章 《【TVM】通过代码学习编译流程》系列 主要介绍 TVM 在模型编译过程的流程,有时候感觉缺少了对类及其属性和方法的介绍。所以决定在系列文章的中间插入一些 “类的结构及其属性方法” 的介绍。

本篇文章主要介绍 Pass 及其相关类。

作为初学者,错误在所难免,还望不吝赐教。

可以再回顾一下在《【TVM】通过代码学习编译流程【4】》中讲到的本体、桥梁、指针的关系。

先看一看 Pass 的基类, 位于 include/tvm/ir/transform.h 。 Pass 本体 PassNode 。内容很少,主要就是 Pass 的执行函数: IRModule operator()(IRModule mod) 函数重载了 “()” 运算符。里面调用自身含有两个参数的 "()" 重载函数。

含有两个参数的 "()" 重载函数 virtual IRModule operator()(IRModule mod, const PassContext& pass_ctx) const = 0; 是个虚函数,这意味着 PassNode 的派生类需要重写该函数,实现 Pass 的实际功能。

class PassNode : public Object {
 public:
  virtual ~PassNode() {}
  
   * \brief Get the pass information/meta data. */
  virtual PassInfo Info() const = 0;
  IRModule operator()(IRModule mod) const {  
    return this->operator()(std::move(mod), PassContext::Current());  
  }
  virtual IRModule operator()(IRModule mod, const PassContext& pass_ctx) const = 0; 
  void VisitAttrs(AttrVisitor* v) {}
  static constexpr const char* _type_key = "transform.Pass";
  TVM_DECLARE_BASE_OBJECT_INFO(PassNode, Object);
};

Pass 指针 Pass ,指向 PassNode 本体。相当于给本体套了个壳子。
壳子中的 IRModule operator()(IRModule mod) const; 函数同样是调用自身含有两个参数的 "()" 重载函数。
含有两个参数的 "()" 重载函数 IRModule Pass::operator()(IRModule mod, const PassContext& pass_ctx) 调用的是本体 PassNode 的功能。

class Pass : public ObjectRef {
 public:
  
  IRModule operator()(IRModule mod) const;
  IRModule operator()(IRModule mod, const PassContext& pass_ctx) const;
  TVM_DEFINE_OBJECT_REF_METHODS(Pass, ObjectRef, PassNode);
 private:
  IRModule static AssertImmutableModule(const IRModule& mod, const PassNode* node,
                                        const PassContext& pass_ctx);
};
 
IRModule Pass::operator()(IRModule mod) const {  
  return this->operator()(std::move(mod), PassContext::Current());
}
IRModule Pass::operator()(IRModule mod, const PassContext& pass_ctx) const {  
  const PassNode* node = operator->();
  ICHECK(node != nullptr);
  const PassInfo& pass_info = node->Info();
  if (!pass_ctx.InstrumentBeforePass(mod, pass_info)) {
    DLOG(INFO) << "Skipping pass : " << pass_info->name
               << " with opt level: " << pass_info->opt_level;
    return mod;
  }
  IRModule ret;
  if (pass_ctx->GetConfig<Bool>("testing.immutable_module", Bool(false)).value()) {
    ret = Pass::AssertImmutableModule(mod, node, pass_ctx);
  } else {
    ret = node->operator()(std::move(mod), pass_ctx);
  }
  pass_ctx.InstrumentAfterPass(ret, pass_info);
  return std::move(ret);
}

所以总结来说,Pass 修改模型的功能由 Pass 的派生类重载的 Pass::operator()(IRModule mod, const PassContext& pass_ctx) 函数实现。那么它有哪些派生类呢?后文提供了三个派生类,分别是 FunctionPassSequentialModulePass 。他们有不同的功能作用。

# FunctionPass

FunctionPassNode :: PassNode

Function-level Pass 的实现类,该类是 Pass 的派生类。接收 Module 中函数表达式列表中的一个 function 进行优化。
pass_func 具体实现 function 优化的函数:由外部提供,以 function 为输入,如 Pass DefuseOps()FoldConstant() 等函数提供他们各自的 pass_func ,以实现不同的功能。

Pass::operator()(IRModule mod, const PassContext& pass_ctx) 函数, FunctionPass 对该函数的实现也在下方。

  • 先遍历模型中的 function
  • AsOptimizableFunctionNode() 函数 :过滤掉不能被优化的 function,如 kCompiler (指定编译器的),kExtern (外部编译器的),kSkipOptimization (指明跳过的)
  • 调用 pass_func 优化 function
class FunctionPassNode : public PassNode {
 public:
  PassInfo pass_info;
  runtime::TypedPackedFunc<Function(Function, IRModule, PassContext)> pass_func;  
  FunctionPassNode() = default;
  void VisitAttrs(tvm::AttrVisitor* v) { v->Visit("pass_info", &pass_info); }
  IRModule operator()(IRModule mod, const PassContext& pass_ctx) const final;  
  PassInfo Info() const override { return pass_info; }
  static constexpr const char* _type_key = "relay.FunctionPass";
  TVM_DECLARE_FINAL_OBJECT_INFO(FunctionPassNode, PassNode);
};
IRModule FunctionPassNode::operator()(IRModule mod, const PassContext& pass_ctx) const {  
  DiagnosticContext previous = DiagnosticContext::Default(mod);
  IRModule updated_mod = mod->ShallowCopy();
  std::vector<std::pair<GlobalVar, Function>> updates;
  for (const auto& kv : mod->functions) {  
    
    if (const auto* function_node = AsOptimizableFunctionNode(kv.second)) {  
      Function updated_func = pass_func(GetRef<Function>(function_node), updated_mod, pass_ctx);  
      updates.push_back({kv.first, std::move(updated_func)});
    }
  }
  return transform::InferType()(updated_mod);
}

# ModulePass

ModulePassNode :: PassNode

Module-level Pass 的实现类, FunctionPass 优化的是 Relay Module 包含的多个 function ,作用于 function 内部,不能实现 function 增删; ModulePassNode 优化的是整个 Module,能够实现 function 增删等 Module 范围的优化。
pass_func 具体实现 Module 优化的函数:由外部提供,以 Module 为输入
Pass::operator()(IRModule mod, const PassContext& pass_ctx) 函数,实现在下方。

  • 调用 pass_func 优化 Module
 * \brief Module-level passes are designed to implement global
 * analysis/optimizations, i.e. interprocedural optimizations (IPO), etc. Passes
 * at this level have the full control of a given Relay program including
 * addition and deletion of functions.
 */
class ModulePassNode : public PassNode {
 public:
  
  PassInfo pass_info;
  runtime::TypedPackedFunc<IRModule(IRModule, PassContext)> pass_func;  
  ModulePassNode() = default;
  void VisitAttrs(tvm::AttrVisitor* v) { v->Visit("pass_info", &pass_info); }
  IRModule operator()(IRModule mod, const PassContext& pass_ctx) const final;  
  
   * \brief Get the pass information/meta data.
   */
  PassInfo Info() const override { return pass_info; }
  static constexpr const char* _type_key = "transform.ModulePass";
  TVM_DECLARE_FINAL_OBJECT_INFO(ModulePassNode, PassNode);
};
IRModule ModulePassNode::operator()(IRModule mod, const PassContext& pass_ctx) const {
  DiagnosticContext previous = DiagnosticContext::Default(mod);
  const PassInfo& pass_info = Info();
  mod = pass_func(std::move(mod), pass_ctx);  
  pass_ctx->diag_ctx.value().Render();
  pass_ctx->diag_ctx = previous;
  return mod;
}

# Sequential

Sequential :Sequential 类包含多个按照顺序执行的 Pass,类似于 pytorch 里面的 nn.Sequential

  • tvm::Array<Pass> passes :数组,包含多个 Pass,如前面提到的 FunctionPassModulePass
  • Pass::operator()(IRModule mod, const PassContext& pass_ctx) 函数,实现在下方。
    • 遍历所有包含的 pass
    • 调用 Pass 执行对模型的优化
 * \brief The SequentialNode contains a set of passes that transform Relay/Relax
 * programs from one AST to another semantically equivalent one.
 *
 * One example of this level of pass is that the pass manager needs to correctly
 * perform a host of optimizations with a given optimization level and disabled
 * passes.
 */
class SequentialNode : public PassNode {
 public:
  
  PassInfo pass_info;
  
  tvm::Array<Pass> passes;  
  PassInfo Info() const override { return pass_info; }
  void ResolveDependency(const IRModule& mod);
  IRModule operator()(IRModule mod, const PassContext& pass_ctx) const final;
  static constexpr const char* _type_key = "transform.Sequential";
  TVM_DECLARE_FINAL_OBJECT_INFO(SequentialNode, PassNode);
};
IRModule SequentialNode::operator()(IRModule mod, const PassContext& pass_ctx) const {
  for (const Pass& pass : passes) {  
    
    
    for (const auto& it : pass_info->required) {
      mod = GetPass(it)(std::move(mod), pass_ctx);
    }
  
    if (pass_ctx->trace_stack.size() && !pass_info->traceable &&
        (!pass_ctx->make_traceable.defined() ||
         pass_ctx->make_traceable.value().count(pass_info->name))) {
      
      String transform_func_key = "relax.tuning_api.Choice.default_transform_func";
      String constr_func_key = "relax.tuning_api.Choice.default_constr_func";
      relax::Knob knob = relax::Knob(
          pass_info->name, <!--swig0-->);
      
      auto trace = Downcast<relax::Trace>(pass_ctx->trace_stack.back());
      trace->Add(knob, "Applied");
      mod = pass(std::move(mod), pass_ctx);  
      trace->SetOutMod(mod);
    } else {
      mod = pass(std::move(mod), pass_ctx);  
    }
  }
  return mod;
}

# 后记

本博客目前以及可预期的将来都不会支持评论功能。各位大侠如若有指教和问题,可以在我的 github 项目 或随便一个项目下提出 issue,并指明哪一篇博客,我看到一定及时回复。