TORCH.FX: PRACTICAL PROGRAM CAPTURE AND TRANSFORMATION FOR DEEP LEARNING IN PYTHON

发表时间: 2021-12 · arXiv:2112.08429

原文: https://arxiv.org/abs/2112.08429

文章标题:TORCH.FX:Python中深度学习的实用程序捕获与转换
作者/机构:James K Reed, Zachary DeVito, Horace He, Ansley Ussery, Jason Ansel (均为 Facebook AI)

速读

一句话结论 本文提出了 `torch.fx`,一个专为 PyTorch 设计的纯 Python 程序捕获与转换库,通过舍弃对长尾复杂语法的支持换取极简的中间表示,在量化、算子融合和 TensorRT 导出等任务中实现了最高 3.7 倍的推理加速,并大幅降低了模型转换的开发门槛。

要解决什么问题 现代深度学习框架的即时执行(eager execution)模式虽然带来了极佳的开发体验,但牺牲了获取全局程序结构的能力,导致难以进行性能优化、程序分析和硬件适配等高级转换。为了在即时模式下找回这种能力,现有的程序捕获系统(例如 TorchScript)试图完全忠实地模拟 Python 的复杂语义,包括可变状态、复杂的控制流和多样的数据类型。这种大而全的设计路线导致捕获技术和生成的中间表示(IR)变得异常臃肿复杂。事实上,绝大多数典型的神经网络(如 CNN、Transformer 甚至封装好的 RNN)在本质上只是没有控制流的扁平张量操作序列,即基本块程序。为了支持极少数的长尾用例而在 IR 中引入控制流,会给编写转换逻辑带来巨大的灾难。例如,当 IR 中存在控制流时,原本简单的张量形状正向传播就会退化为复杂且容易出错的不动点数据流分析,跨循环迭代传递的张量可能会呈现无限多种形状,最终只能得到一个无用的动态值,从而彻底阻碍后续需要具体形状信息的编译优化。此外,现有系统为了处理张量的别名和突变语义,需要引入代价高昂的别名分析,保守的假设往往会进一步阻碍优化空间的探索。

怎么做的 核心思路是放弃对 Python 长尾复杂语法的全面支持,转而针对深度学习中最常见的有向无环图(DAG)结构,设计一个极简的、纯 Python 的程序捕获与转换框架。该方法主要由三个核心部件构成。首先是基于代理的符号追踪(Symbolic Tracing)机制。它通过一个名为 `Proxy` 的鸭子类型对象来充当具体程序值的抽象替代品,利用 `__torch_function__` 协议来拦截并记录 PyTorch 算子的派发。追踪过程是预先(AoT)且不进行特化的,它会直接展开与输入无关的控制流。用户可以通过重写 `Tracer` 类的 `is_leaf_module` 方法来高度定制追踪行为,指定哪些模块需要被展开,哪些保留为不透明的黑盒调用。其次是仅包含 6 种指令的极简中间表示(IR)。程序被表示为一个包含线性 `Node` 序列的 `Graph` 对象。每个节点可以通过如下元组定义其核心结构: $$ Node = (opcode, target, args, kwargs) $$ 其中 $opcode$ 仅限于占位符、函数调用、方法调用、模块调用、属性获取和输出这 6 种操作。为了保持极简,IR 内部刻意剔除了控制流和状态突变原语,节点之间的数据依赖直接通过 $args$ 和 $kwargs$ 中的引用表达,且支持直接将 Python 内置类型(如整数、列表)作为节点参数嵌入,无需单独的对象构造节点。最后是代码生成系统与状态管理。`torch.fx` 将转换后的图与模型参数统一封装在 `GraphModule` 中。它不依赖定制的跨语言运行时,而是直接从 IR 重新生成有效的 Python 源代码,并将其安装为 `GraphModule` 的 `forward` 方法。这使得转换后的模型可以无缝接入原生的 PyTorch 生态,甚至可以再次被追踪和进行多轮转换。

效果如何 实验在多种模型架构和硬件(Intel Xeon CPU、NVIDIA V100 GPU)上验证了该框架的有效性。在 IR 复杂度对比中,针对标准的 ResNet50 模型,`torch.fx` 生成的 IR 仅包含 445 个操作,而代表基于示例追踪路线的 `torch.jit.trace` 为 860 个,代表全面解析 Python 源码路线的 `torch.jit.script` 高达 2614 个,证明了其表示的极简性。在实际优化任务中,对 DeepRecommender 模型进行训练后量化,在 CPU 上比浮点模型实现了 3.3 倍的推理加速;对 ResNet50 应用卷积与批量归一化融合,在 CPU 多线程下延迟降低了 40%,在 GPU 上降低了 6%,且整个融合转换工具仅用不到 150 行 Python 代码即可实现。在分布式训练的程序调度优化中,通过替换阻塞式远程调用,将 QPS 提升了 9%。在设备导出场景下,通过将模型下沉至 NVIDIA TensorRT,ResNet50 和 LearningToPaint 模型分别获得了 3.7 倍和 1.54 倍的运行时加速。该方法的局限性在于,由于 IR 剔除了控制流和可变性语义,它无法直接捕获那些真正包含依赖输入控制流的长尾模型,且将原位修改操作视为未定义行为,遇到此类情况需要用户手动配置追踪器来规避。

A1 主要贡献

本文旨在解决现代深度学习框架中的一个核心矛盾:即时执行(eager execution)模式虽然提升了开发效率和用户体验,但却牺牲了对程序结构的访问能力,而这种能力对于性能优化、可视化、分析和硬件集成等高级转换至关重要。为了在即时模式框架(如PyTorch)中重新获得这种能力,需要一种程序捕获机制。然而,现有的系统(如TorchScript)为了完全忠实地模拟Python的复杂语义(包括可变状态、控制流、复杂数据类型),其捕获技术和生成的中间表示(IR)都变得异常复杂,给转换(transform)的编写带来了巨大困难。

本文提出,可以通过专注于深度学习的典型用例(大多数神经网络模型的高层有向无环图(DAG)结构)而非长尾的复杂情况,来设计一个更简单、更高效的程序捕获与转换框架。基于这一理念,本文介绍了torch.fx,一个完全用Python编写的、为PyTorch设计的程序捕获与转换库,其核心目标是为机器学习从业者提供极高的开发生产力。

本文的主要贡献如下:

  1. 实用性分析:对深度学习程序中重要的程序捕获与转换特性进行了实用性分析。
  2. 纯Python程序捕获库:实现了一个纯Python的程序捕获库,该库可被定制以捕获不同层次的程序细节。
  3. 简单的6指令IR:提出了一种仅包含6个指令的简单中间表示(IR),其设计重点在于易于理解和进行静态分析。
  4. 代码生成系统:构建了一个代码生成系统,能够将转换后的代码无缝地返回到宿主语言(Python)的生态系统中。
  5. 案例研究:展示了torch.fx在实践中如何被用于性能优化、程序分析、设备适配(device lowering)等场景,实现了PyTorch生态系统中以前难以完成的工作流。

A3 背景知识与设计原则

背景知识

程序捕获、特化与IR设计的权衡。无论是即时模式还是图模式框架,在捕获和转换程序时都必须在程序结构的捕获、程序的特化(specialization)以及中间表示(IR)的设计之间做出选择。这些选择共同决定了框架能表示的程序范围、编写转换的难易程度以及转换后程序的性能。通常,为了支持更多程序并实现高性能,需要更复杂的捕获框架和IR,这反过来又使得转换的编写更加困难。

2.1 捕获程序结构

2.2 程序特化

2.3 中间表示(IR)设计

设计原则

现有框架的设计大多倾向于支持更广泛的深度学习程序,但牺牲了实现的简洁性。当捕获的程序是运行的唯一方式时,高保真度至关重要。但PyTorch主要作为即时执行框架使用,程序捕获仅用于特定转换,无需对整个程序都有效。此外,目标用户是机器学习从业者,他们更习惯使用Python而非编译器设计。

通过为典型的深度学习模型而非长尾用例进行设计,可以创建一个更易于使用和实现的框架。torch.fx的设计原则体现了这一理念:

A2 方法细节

TORCH.FX 概述

torch.fx采用符号追踪来捕获程序,使用一个简单的包含6个指令且基于Python的IR来表示它们,并从IR重新生成Python代码来执行。为避免JIT特化带来的重捕获复杂性,torch.fx本身不尝试特化程序,而是依赖于转换过程来决定需要执行何种特化。符号追踪过程是可配置的,用户可以定制以处理更特殊的用例。

from torch.fx import Graph
def replace_activation(g: Graph, old, new):
    for n in g.nodes:
        if n.op == 'call_function' and n.target == old:
            # create IR to call new activate
            with g.inserting_after(n):
                new_n = g.call_function(new, n.args)
            n.replace_all_uses_with(new_n)
            g.erase_node(n)
            # or for this simplified case: 'n.target = new'

replace_activation(traced.graph, torch.relu, torch.nn.functional.gelu)
traced.recompile()
图1. torch.fx使用符号追踪将程序捕获到一个简单的IR中,并从该IR生成Python代码。图2. 变换,如此处替换激活函数的变换,是直接用Python编写的。

4.1 程序捕获

4.2 中间表示

4.3 源码到源码的转换

设计决策

torch.fx融合并扩展了先前工作中的方法,提供了一个易于使用、实现简单且可配置的库。

5.1 符号追踪

5.2 可配置的程序捕获

5.3 预先(AoT)捕获而不进行特化

5.4 基于Python的IR和变换

5.5 IR内部无控制流

def loop_shapes(x, itr):
  # x is an input tensor of size [1, N]
  for _ in range(itr):
    x = torch.cat((x, x), dim=0)

  # Depending on the number of loop iterations, x may have an
  # arbitrary leading dimension i.e. x \in [*dynamic*, N]
  return x

IR本身不包含控制流,并不妨碍变换在更大模型中的基本块子图上工作;如何组合这些子图的细节留给变换的编写者或用户来决定。

5.6 函数式图与有状态模块

A4 实验环境

A5 实验结果

6.1 IR 复杂度

6.2 性能优化

6.3 程序分析

torch.fx已被应用于多种程序分析场景:

6.4 设备和运行时导出/编译

A6 结论

本文介绍了torch.fx,一个纯Python的系统,用于捕获和转换PyTorch程序。通过分析相关系统(如控制流、可变性、数据模型)的复杂性来源,本文展示了torch.fx如何通过专注于常见用例和提供可定制性来避免这些复杂性。通过对优化、分析和设备下沉等多个用例的研究,本文证明了torch.fx的API设计如何成功地实现了这些功能。

A7 附录

A. TORCH.FX 节点语义

A.1 操作码(Opcode)含义

下表描述了torch.fx中每个Nodeopcode的含义。

A.2 args/kwargs 行为

下表描述了不同opcodeargskwargs字段的预期行为。

B 量化评估数值数据

下表为第6.2.1节量化实验的详细运行时间数据(单位:秒)。

C 融合评估数值数据

下表为第6.2.2节融合实验的详细运行时间数据(单位:秒)。

D TensorRT评估数值数据

下表为第6.4节TensorRT实验的详细运行时间数据(单位:秒)。