<?xml version="1.0" encoding="utf-8" standalone="yes"?><rss version="2.0" xmlns:atom="http://www.w3.org/2005/Atom"><channel><title>Ai Infra on GlowLED的后花园</title><link>https://blog.glowled.top/categories/ai-infra/</link><description>Recent content from GlowLED的后花园</description><generator>Hugo</generator><language>zh-CN</language><managingEditor>zpeiyu11@gamil.com (GlowLED)</managingEditor><webMaster>zpeiyu11@gamil.com (GlowLED)</webMaster><copyright>本博客所有文章除特别声明外，均采用 BY-NC-SA 许可协议。转载请注明出处！</copyright><lastBuildDate>Mon, 20 Jul 2026 16:38:14 +0800</lastBuildDate><atom:link href="https://blog.glowled.top/categories/ai-infra/index.xml" rel="self" type="application/rss+xml"/><item><title>二十一世纪火柴编译目录</title><link>https://blog.glowled.top/post/torch-compile-tech-overview/</link><pubDate>Mon, 20 Jul 2026 16:38:14 +0800</pubDate><author>zpeiyu11@gamil.com (GlowLED)</author><guid>https://blog.glowled.top/post/torch-compile-tech-overview/</guid><description>
<![CDATA[<h1>二十一世纪火柴编译目录</h1><p>作者：GlowLED（zpeiyu11@gamil.com）</p>
        
          <p>本文将对torch compile技术进行总览性介绍，虽然不全也不精，但可以快速建立起一个对于torch compile的整体认知，在进行相关工程时不会变成完全黑盒。</p>
<p>故命名为<strong>二十一世纪火柴编译目录</strong>。</p>
<h2 id="background">
<a class="header-anchor" href="#background"></a>
Background
</h2><p>PyTorch默认为<strong>eager模式</strong>，即“立即计算”，在计算过程中动态构建计算图。这使得研究人员可以方便地调试，可以在任何地方插入<code>print</code>，可以直接使用Python的<code>if/for</code>来控制计算流。
但是这种灵活性是有性能代价的：</p>
<ul>
<li>增大了<strong>启动开销</strong>：每个PyTorch算子都要被独立调用，有很繁复的Kernel Launch开销。</li>
<li>增大<strong>优化难度</strong>：这也使得跨算子优化（算子融合）比较困难。</li>
</ul>
<p>在PyTorch 1.x的时代，对于这个问题也有解决方案，就是<strong>TorchScript</strong>。这是一种DSL（类似Triton），使用<code>torch.jit</code>装饰器来对被编译的区域进行warp。TorchScript存在这样一些问题：</p>
<ul>
<li>有自己<strong>独立于Python</strong>的类型系统和定义控制流的语法，例如不能用NumPy，不能用Python的字典推导式、很多标准库函数不能用等</li>
<li>图<strong>捕获能力有点残疾</strong>。TorchScript使用<code>torch.jit.trace</code>进行图捕获，它从执行时调用的算子Op获取图，但是对于数据依赖的情况（比较典型的就是if-else依赖数据值时），<code>torch.jit.trace</code>不能感知到语法层面的控制流，只能感知到实际执行时调用的算子Op，因此会固化第一次运行时所走的分支。</li>
</ul>
<p>所以，综上，TorchScript在<strong>Pythonic</strong>和<strong>无痛使用</strong>方面做得比较一般，大家还是更愿意在Triton里痛苦一会拿到更好的性能。</p>
<p>为了彻底解决这个问题，PyTorch 2.x大版本引入了一整套JIT编译系统。为了实现：</p>
<ul>
<li>完整的图捕获能力</li>
<li>为了易用性牺牲一些强迫性：接受代码部分回退到eager模式</li>
<li>性能提升应当显著和可预测</li>
</ul>
<h2 id="overview">
<a class="header-anchor" href="#overview"></a>
Overview
</h2><p>torch的编译系统体系主要包含三大核心组件：</p>
<ul>
<li><strong>TorchDynamo</strong>：用于捕获运行时生成的计算图，转化为一种叫FX Graph的计算图。</li>
<li><strong>AOTAutograd/AOTDispatcher</strong>：规范化FX Graph，将节点内容dispatch到aten ops上，生成反向计算图；前向后向中间激活值最小化（训练模式下）。</li>
<li><strong>TorchInductor</strong>：对FX Graph进行算子融合优化，然后生成Triton代码，调用Triton编译器进行编译并缓存kernel。</li>
</ul>
<p>torch.compile的实际性能收益主要来源于：</p>
<ul>
<li><strong>算子融合</strong>，减小memory bound和kernel launch开销。</li>
<li>在训练时，AOTAutograd的<strong>激活值联合优化</strong>，减少显存开销。</li>
<li><strong>cuda graph</strong>（可选），降低CPU启动开销。</li>
</ul>
<pre class="mermaid">
  flowchart LR
    A[代码] -->|Python代码| B[TorchDynamo]
    B -->|FX Graph| C[AOTAutograd<br/>AOTDispatcher]

    C -->|Joint Graph| D[TorchInductor]

    D -->|Triton代码| E[Triton编译器]
    E --> F[可运行的kernel]
</pre><h2 id="torchdynamo">
<a class="header-anchor" href="#torchdynamo"></a>
TorchDynamo
</h2><p>TorchDynamo 是整个编译系统的入口，负责从 Python 代码中提取计算图。</p>
<p>可以用自定义后端的方式来拦截TorchDynamo -&gt; AOTAutograd/TorchInductor的过程，查看捕获的FX Graph：</p>
<div class="highlight"><pre tabindex="0" class="chroma"><code class="language-python" data-lang="python"><span class="line"><span class="cl"><span class="kn">import</span> <span class="nn">torch</span>
</span></span><span class="line"><span class="cl"><span class="kn">from</span> <span class="nn">typing</span> <span class="kn">import</span> <span class="n">Sequence</span><span class="p">,</span> <span class="n">Callable</span>
</span></span><span class="line"><span class="cl">
</span></span><span class="line"><span class="cl"><span class="k">def</span> <span class="nf">custom_backend</span><span class="p">(</span>
</span></span><span class="line"><span class="cl">    <span class="n">gm</span><span class="p">:</span> <span class="n">torch</span><span class="o">.</span><span class="n">fx</span><span class="o">.</span><span class="n">GraphModule</span><span class="p">,</span>
</span></span><span class="line"><span class="cl">    <span class="n">example_inputs</span><span class="p">:</span> <span class="n">Sequence</span><span class="p">[</span><span class="n">torch</span><span class="o">.</span><span class="n">Tensor</span><span class="p">],</span>
</span></span><span class="line"><span class="cl"><span class="p">):</span>
</span></span><span class="line"><span class="cl">    <span class="nb">print</span><span class="p">(</span><span class="s2">&#34;FX Graph:&#34;</span><span class="p">)</span>
</span></span><span class="line"><span class="cl">    
</span></span><span class="line"><span class="cl">    <span class="n">gm</span><span class="o">.</span><span class="n">graph</span><span class="o">.</span><span class="n">print_tabular</span><span class="p">()</span>
</span></span><span class="line"><span class="cl">    
</span></span><span class="line"><span class="cl">    <span class="nb">print</span><span class="p">(</span><span class="s2">&#34;</span><span class="se">\n</span><span class="s2">Generated Python code:&#34;</span><span class="p">)</span>
</span></span><span class="line"><span class="cl">    <span class="nb">print</span><span class="p">(</span><span class="n">gm</span><span class="o">.</span><span class="n">code</span><span class="p">)</span>
</span></span><span class="line"><span class="cl">    
</span></span><span class="line"><span class="cl">    <span class="k">return</span> <span class="n">gm</span><span class="o">.</span><span class="n">forward</span>
</span></span><span class="line"><span class="cl">    
</span></span><span class="line"><span class="cl"><span class="nd">@torch.compile</span><span class="p">(</span><span class="n">backend</span><span class="o">=</span><span class="n">custom_backend</span><span class="p">)</span>
</span></span><span class="line"><span class="cl"><span class="k">def</span> <span class="nf">func</span><span class="p">(</span><span class="n">x</span><span class="p">:</span> <span class="n">torch</span><span class="o">.</span><span class="n">Tensor</span><span class="p">,</span> <span class="n">y</span><span class="p">:</span> <span class="n">torch</span><span class="o">.</span><span class="n">Tensor</span><span class="p">)</span> <span class="o">-&gt;</span> <span class="n">torch</span><span class="o">.</span><span class="n">Tensor</span><span class="p">:</span>
</span></span><span class="line"><span class="cl">    <span class="n">z</span> <span class="o">=</span> <span class="n">x</span> <span class="o">+</span> <span class="n">y</span>
</span></span><span class="line"><span class="cl">    <span class="k">return</span> <span class="n">torch</span><span class="o">.</span><span class="n">relu</span><span class="p">(</span><span class="n">z</span><span class="p">)</span> <span class="o">*</span> <span class="mi">2</span>
</span></span><span class="line"><span class="cl">
</span></span><span class="line"><span class="cl"><span class="n">x</span> <span class="o">=</span> <span class="n">torch</span><span class="o">.</span><span class="n">randn</span><span class="p">((</span><span class="mi">4</span><span class="p">,),</span> <span class="n">dtype</span><span class="o">=</span><span class="n">torch</span><span class="o">.</span><span class="n">float32</span><span class="p">,</span> <span class="n">device</span><span class="o">=</span><span class="s2">&#34;cuda:1&#34;</span><span class="p">)</span>
</span></span><span class="line"><span class="cl"><span class="n">y</span> <span class="o">=</span> <span class="n">torch</span><span class="o">.</span><span class="n">randn</span><span class="p">((</span><span class="mi">4</span><span class="p">,),</span> <span class="n">dtype</span><span class="o">=</span><span class="n">torch</span><span class="o">.</span><span class="n">float32</span><span class="p">,</span> <span class="n">device</span><span class="o">=</span><span class="s2">&#34;cuda:1&#34;</span><span class="p">)</span>
</span></span><span class="line"><span class="cl">
</span></span><span class="line"><span class="cl"><span class="n">o</span> <span class="o">=</span> <span class="n">func</span><span class="p">(</span><span class="n">x</span><span class="p">,</span> <span class="n">y</span><span class="p">)</span>
</span></span><span class="line"><span class="cl">
</span></span><span class="line"><span class="cl"><span class="n">x</span> <span class="o">=</span> <span class="n">torch</span><span class="o">.</span><span class="n">randn</span><span class="p">((</span><span class="mi">4</span><span class="p">,),</span> <span class="n">dtype</span><span class="o">=</span><span class="n">torch</span><span class="o">.</span><span class="n">float32</span><span class="p">,</span> <span class="n">device</span><span class="o">=</span><span class="s2">&#34;cuda:1&#34;</span><span class="p">)</span>
</span></span><span class="line"><span class="cl"><span class="n">y</span> <span class="o">=</span> <span class="n">torch</span><span class="o">.</span><span class="n">randn</span><span class="p">((</span><span class="mi">4</span><span class="p">,),</span> <span class="n">dtype</span><span class="o">=</span><span class="n">torch</span><span class="o">.</span><span class="n">float32</span><span class="p">,</span> <span class="n">device</span><span class="o">=</span><span class="s2">&#34;cuda:1&#34;</span><span class="p">)</span>
</span></span><span class="line"><span class="cl">
</span></span><span class="line"><span class="cl"><span class="n">o</span> <span class="o">=</span> <span class="n">func</span><span class="p">(</span><span class="n">x</span><span class="p">,</span> <span class="n">y</span><span class="p">)</span>
</span></span></code></pre></div><p>输出为：</p>
        
        <hr><p>本文2026-07-20首发于<a href='https://blog.glowled.top/'>GlowLED的后花园</a>，最后修改于2026-07-20</p>]]></description><category>ai infra</category></item></channel></rss>