计算图基本知识
计算图的基本功能
统一计算过程:不同的计算后端(硬件)有不同的表示,深度学习框架需要把机器学习统一表达为一种IR。以进行一种共性优化。
自动化微分:对于任意的模型拓扑,计算梯度的方法必须要通用且自动运行。计算图可以记录辅助分析模型的梯度计算过程。
分析模型变量的生命周期:例如激活值和梯度。从而优化内存管理(ZeRO)
优化程序执行:根据网络拓扑计算图,来构建算子执行依赖关系。从而优化模型执行效率。类似编译器的优化。
计算图生成
计算图分为静态与动态。
动态图对每个操作都即时发送到计算内核然后得到计算结果。
静态图的计算图是获取了整个计算序列之后计算的。
因此动态图的有点是一步一计算,可以更好的调试,但是效率低。
静态图的优点是可以进行更多的优化,但是调试困难。
! 注意:pytorch默认计算动态图,但是可以通过装饰trace来转换为静态图。
静态图还可以通过避免重新编译来快捷部署。但是有些计算是延迟的,不能易于调试和断点恢复。
中间表示 (IR)
和传统编译器一样,深度学习框架也有中间表示。这个中间表示是对计算图的抽象,可以用来进行优化。但同时要保留源代码的完备性。
线性IR
线性IR类似汇编的顺序执行。例如堆栈机、三地址码等。
图IR
编译过程通过图保存。通过节点与边等结构表示。类似抽象语法树等。更多用在编译器前端。
混合IR
以LLVM为例,LLVM使用线性IR为基本块,用图表示控制流。每个基本块内部的变量只能赋值一次,并且要在使用前定义。
机器学习框架IR
机器学习框架相比于传统编译器,需要表达好张量数据。还需要考虑自动求导计算的数据流。针对优化方面,机器学习框架需要优化计算图本身以及硬件相关的优化工作。这些都依赖于中间表示的实现。IR不仅影响静态图生成的执行效率,也对动态图的JIT有影响。
Pytorch
Pytorch主要基于动态计算图。利用Torchscript来创建可序列化的模型。TorchScript通过JIT讲Python转换为模型文件。通过torch.jit.script装饰器来实现计算图。
JAX
JAX生成了JAXpr的中间表示。是一种函数式的中间表示,其具有强类型,只依赖局部变量。其中表达式只有原子表达式和其组成的符合表达式。
Tensorflow
Tensorflow同时支持静态与动态图。其编译过程会将程序逐层编译,最终得到最底层IR。
MLIR
MLIR提供了一个统一抽象的概念。使用MLIR的套件来定义自己需求的IR。使用Dialect概念来为特定名称空间下的抽象分组。抽象表示可以绑定op,从而生成MLIR的烈性。
自动微分
自动微分和数值微分与符号微分不同,为了避免导函数规模过大,讲运算拆解为基本运算集合,然后通过链式法则来计算导数。这样可以减少计算量。
自动微分的实现可以依靠反向计算的操作符重载,通过追踪每一个程序的控制流,来按照轨迹反向执行微分。但是也因此不能预编译,需要运行时信息来得到。
静态分析
静态类型检查可以避免出现一部分运行时的错误,也能够给编译器更多优化空间。类型推导也可以进行泛型特化。或者将抽象语义近似为实际语义。
前端编译优化
通常有一些简单的优化,与硬件无关。例如常量替换、死代码擦除等。
死代码消除
无用代码指输出结果没有被使用的代码。不可达代码指没有控制流使用的代码。因此可以进行删除
常量传播 折叠
常量传播通过把已知的值改成立即数替换,从而加速
常量折叠:通过计算常量表达式的值,从而减少计算量。
公共表达式消除
如果表达式的值,有重复计算,则将这个表达式的一个计算结果共享。
本博客所有文章除特别声明外,均采用 CC BY-SA 4.0 协议 ,转载请注明出处!