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

推荐订阅源

Blog — PlanetScale
Blog — PlanetScale
博客园_首页
WordPress大学
WordPress大学
博客园 - 聂微东
P
Privacy International News Feed
Forbes - Security
Forbes - Security
Threat Intelligence Blog | Flashpoint
Threat Intelligence Blog | Flashpoint
Last Week in AI
Last Week in AI
C
CERT Recently Published Vulnerability Notes
月光博客
月光博客
NISL@THU
NISL@THU
美团技术团队
T
Tailwind CSS Blog
Jina AI
Jina AI
Cyber Security Advisories - MS-ISAC
Cyber Security Advisories - MS-ISAC
Apple Machine Learning Research
Apple Machine Learning Research
C
Cisco Blogs
freeCodeCamp Programming Tutorials: Python, JavaScript, Git & More
The Hacker News
The Hacker News
B
Blog
P
Palo Alto Networks Blog
L
Lohrmann on Cybersecurity
有赞技术团队
有赞技术团队
The Register - Security
The Register - Security
S
Securelist
A
Arctic Wolf
MyScale Blog
MyScale Blog
H
Help Net Security
N
Netflix TechBlog - Medium
CTFtime.org: upcoming CTF events
CTFtime.org: upcoming CTF events
T
Threatpost
Recent Commits to openclaw:main
Recent Commits to openclaw:main
Security Latest
Security Latest
T
Tor Project blog
V
Vulnerabilities – Threatpost
V
V2EX
AI
AI
Hugging Face - Blog
Hugging Face - Blog
大猫的无限游戏
大猫的无限游戏
博客园 - Franky
Simon Willison's Weblog
Simon Willison's Weblog
小众软件
小众软件
OSCHINA 社区最新新闻
OSCHINA 社区最新新闻
H
Hackread – Cybersecurity News, Data Breaches, AI and More
T
Troy Hunt's Blog
Schneier on Security
Schneier on Security
cs.AI updates on arXiv.org
cs.AI updates on arXiv.org
H
Heimdal Security Blog
Google Online Security Blog
Google Online Security Blog
Know Your Adversary
Know Your Adversary

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

【琐记】烟火与尘埃 【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,并指明哪一篇博客,我看到一定及时回复。