为什么会有TorchInductor?
众所周知,Python是一种解释型语言,在运行期间才去解析代码,然后执行。这种方式相比于C++这一类的需要编译后执行的语言来说效率很低的。随着AI的发展,基于Python这种对编程水平要求没那么高的编程语言进行算法研究的科研人员来说,自己实现模型结构也变得更加简单了,Pytorch和TF也就逐渐成为AI领域的重要工具。但是随着时间的推理,AI算法工程师对于模型推理性能的越发看重,解释型语言已经不满足要求了,尤其是对于计算密集型的AI算子来说,纯Python的实现更加不可能成为AI发展的最终形态,所以各个对性能要求比较高的工具组合使用Python/C++来提升工具的性能,也成为了现在的通用行为,Python层面负责用户便利性,C++层面进行各种各样的性能优化,这部分C++代码逐渐发展成为一个个推理框架的核心。
再随着时间推移,模型越来越大,模型结构也愈发复杂,单个算子调用C++的方式也不能满足实际的推理需求,而且这种没有Graph结构的方式对于AI编译器来说十分不利于优化,工程师无法拿到完整的模型结构,也就没有办法进行更多的优化(例如各种的算子融合、分支融合、buffer生命周期管理相关的buffer复用等等优化)。这个阶段torch的做法是先自研了TorchScript做图捕获,同时也接入了ONNX这个编译推理部署的推理框架,让工程师可以获取到完整的模型结构,十分利于AI编译器工程师进行编译期的优化(算子的等效融合、替换,Layout优化适配,分支融合,buffer依赖分析等)。但是外部工具导致torch只要版本升级总会产生一定的开发周期Gap,需要下游频繁适配上游的feature变化。逐渐就有了Torch自己的AI编译器–inductor,通过torch.compile()这个工具允许用户直接获取torch的模型结构,不再需要借助ONNX。inductor 拿到模型Graph后会进行大量优化,来生成高性能的内核。
Inductor具体做了什么?
主要工作分为了三部分:Lowering; Fuse and Optimize; CodeGen
Lowering
前面我们说了,torch模型经过torch.compile()之后获取到了模型结构,而模型结构的表示还是用FX Graph(High Level IR)来表示的,而Inductor的Lowering做的事情就是将FX Graph转换为自己的IR表示:Loop-based IR,这一层IR已经展开了FX Graph这一层的shape以及循环。Shape展开可以基本上确定了Buffer大小,所以已经可以提前分配buffer;如果是动态Shape,则使用符号化表示动态Axis,在运行时也可以获取到真实的shape来确定真实buffer大小。循环展开以后就带来了更多的计算层面融合的机会,减少数据的搬移次数,降低中间buffer的开销。
动态shape因为需要在推理期进行实际的buffer分配(即便使用内存池也无法完全消除runtime buffer分配的开销),以及因为循环边界的不确定性,无法完全消除循环,优化kernel中的block size也是不能完全确定的,同时也必然会存在一定的strip mining问题,会增加一些相应的分支和判断,所以kernel的实现会更加通用。因此性能相比于纯静态的模型来说必然会有所降低。
Fuse and Optimize
这部分工作是Inductor的重点,其中Fuse是最为核心的功能,避免中间结果从计算单元频繁进行内存读写,减少带宽消耗。同时也会进行buffer复用等分析,以及其他优化。
CodeGen
这个阶段通常是硬件相关的代码,inductor输出的高性能kernel目前支持GPU(生成Triton字符串代码,然后由Triton编译器编译为可执行文件)和CPU(C++代码,同时使用OMP并行,再编译为.so 供Pytorch调用)
WHY TorchInductor
优势
- 无需手写kernel:Inductor中集成了针对硬件特定优化pass,可以生成调优后的kernel
- E2E加速:相比逐个执行算子,Inductor中大量的FusedOp可以消除大量的memory I/O,减少Op数量也就意味着可以减少大量的KernelLaunch开销
劣势
- 重复编译:当shape、stride、device出现变化时,可能会触发重新编译
- 调试门槛高:配合
TORCH_LOGS=inductor、torch._inductor.config等调试选项来查看LOG,大量LOG,(丢给AI分析,可能已经不是多高的门槛了) - graph breaks:遇到不支持的一些控制流、动态对象、非tensor操作会导致Graph被切分,优化范围减小。
- 自定义算子不支持,混合Numpy、三方库等操作不支持,会fallback到eager
- 编译缓存会占用大量内存,特化的kernel代码会导致编译产物增长
- 数值精度:Fused Op会导致计算顺序和eager不同,从而出现精度误差。