本文将对torch compile技术进行总览性介绍,虽然不全也不精,但可以快速建立起一个对于torch compile的整体认知,在进行相关工程时不会变成完全黑盒。
故命名为二十一世纪火柴编译目录。
Background
PyTorch默认为eager模式,即“立即计算”,在计算过程中动态构建计算图。这使得研究人员可以方便地调试,可以在任何地方插入print,可以直接使用Python的if/for来控制计算流。
但是这种灵活性是有性能代价的:
- 增大了启动开销:每个PyTorch算子都要被独立调用,有很繁复的Kernel Launch开销。
- 增大优化难度:这也使得跨算子优化(算子融合)比较困难。
在PyTorch 1.x的时代,对于这个问题也有解决方案,就是TorchScript。这是一种DSL(类似Triton),使用torch.jit装饰器来对被编译的区域进行warp。TorchScript存在这样一些问题:
- 有自己独立于Python的类型系统和定义控制流的语法,例如不能用NumPy,不能用Python的字典推导式、很多标准库函数不能用等
- 图捕获能力有点残疾。TorchScript使用
torch.jit.trace进行图捕获,它从执行时调用的算子Op获取图,但是对于数据依赖的情况(比较典型的就是if-else依赖数据值时),torch.jit.trace不能感知到语法层面的控制流,只能感知到实际执行时调用的算子Op,因此会固化第一次运行时所走的分支。
所以,综上,TorchScript在Pythonic和无痛使用方面做得比较一般,大家还是更愿意在Triton里痛苦一会拿到更好的性能。
为了彻底解决这个问题,PyTorch 2.x大版本引入了一整套JIT编译系统。为了实现:
- 完整的图捕获能力
- 为了易用性牺牲一些强迫性:接受代码部分回退到eager模式
- 性能提升应当显著和可预测
Overview
torch的编译系统体系主要包含三大核心组件:
- TorchDynamo:用于捕获运行时生成的计算图,转化为一种叫FX Graph的计算图。
- AOTAutograd/AOTDispatcher:规范化FX Graph,将节点内容dispatch到aten ops上,生成反向计算图;前向后向中间激活值最小化(训练模式下)。
- TorchInductor:对FX Graph进行算子融合优化,然后生成Triton代码,调用Triton编译器进行编译并缓存kernel。
torch.compile的实际性能收益主要来源于:
- 算子融合,减小memory bound和kernel launch开销。
- 在训练时,AOTAutograd的激活值联合优化,减少显存开销。
- cuda graph(可选),降低CPU启动开销。
flowchart LR
A[代码] -->|Python代码| B[TorchDynamo]
B -->|FX Graph| C[AOTAutograd
AOTDispatcher]
C -->|Joint Graph| D[TorchInductor]
D -->|Triton代码| E[Triton编译器]
E --> F[可运行的kernel]
TorchDynamo
TorchDynamo 是整个编译系统的入口,负责从 Python 代码中提取计算图。
可以用自定义后端的方式来拦截TorchDynamo -> AOTAutograd/TorchInductor的过程,查看捕获的FX Graph:
import torch
from typing import Sequence, Callable
def custom_backend(
gm: torch.fx.GraphModule,
example_inputs: Sequence[torch.Tensor],
):
print("FX Graph:")
gm.graph.print_tabular()
print("\nGenerated Python code:")
print(gm.code)
return gm.forward
@torch.compile(backend=custom_backend)
def func(x: torch.Tensor, y: torch.Tensor) -> torch.Tensor:
z = x + y
return torch.relu(z) * 2
x = torch.randn((4,), dtype=torch.float32, device="cuda:1")
y = torch.randn((4,), dtype=torch.float32, device="cuda:1")
o = func(x, y)
x = torch.randn((4,), dtype=torch.float32, device="cuda:1")
y = torch.randn((4,), dtype=torch.float32, device="cuda:1")
o = func(x, y)
输出为:
FX Graph:
opcode name target args kwargs
------------- ------ ------------------------------------------------------- ------------ --------
placeholder l_x_ L_x_ () {}
placeholder l_y_ L_y_ () {}
call_function z <built-in function add> (l_x_, l_y_) {}
call_function relu <built-in method relu of type object at 0x7f8b4ed8b4a0> (z,) {}
call_function mul <built-in function mul> (relu, 2) {}
output output output ((mul,),) {}
Generated Python code:
def forward(self, L_x_ : torch.Tensor, L_y_ : torch.Tensor):
l_x_ = L_x_
l_y_ = L_y_
z = l_x_ + l_y_; l_x_ = l_y_ = None
relu = torch.relu(z); z = None
mul = relu * 2; relu = None
return (mul,)
FX Graph将自己以表格形式打印出来后,可以看到存在opcode、name、target、args字段。opcode代表的是操作的类型,比如占位符、调用函数、输出等。这一部分非常像TensorFlow编程时的概念,可以说静态计算图中的概念是一脉相承的。这是一种函数式的描述方法。 下面通过FX Graph生成的Python代码也基本可以和FX Graph一一对应。
处理控制流
对于简单的控制流,TorchDynamo 可以展开:
@torch.compile(backend=custom_backend)
def func(x: torch.Tensor, y: torch.Tensor) -> torch.Tensor:
z = x + y
for _ in range(3):
z += x + y
return torch.relu(z) * 2
输出将变为:
FX Graph:
opcode name target args kwargs
------------- ------ ------------------------------------------------------- ------------ --------
placeholder l_x_ L_x_ () {}
placeholder l_y_ L_y_ () {}
call_function z <built-in function add> (l_x_, l_y_) {}
call_function add_1 <built-in function add> (l_x_, l_y_) {}
call_function z_1 <built-in function iadd> (z, add_1) {}
call_function add_2 <built-in function add> (l_x_, l_y_) {}
call_function z_2 <built-in function iadd> (z_1, add_2) {}
call_function add_3 <built-in function add> (l_x_, l_y_) {}
call_function z_3 <built-in function iadd> (z_2, add_3) {}
call_function relu <built-in method relu of type object at 0x7f2d9658b4a0> (z_3,) {}
call_function mul <built-in function mul> (relu, 2) {}
output output output ((mul,),) {}
Generated Python code:
def forward(self, L_x_ : torch.Tensor, L_y_ : torch.Tensor):
l_x_ = L_x_
l_y_ = L_y_
z = l_x_ + l_y_
add_1 = l_x_ + l_y_
z += add_1; z_1 = z; z = add_1 = None
add_2 = l_x_ + l_y_
z_1 += add_2; z_2 = z_1; z_1 = add_2 = None
add_3 = l_x_ + l_y_; l_x_ = l_y_ = None
z_2 += add_3; z_3 = z_2; z_2 = add_3 = None
relu = torch.relu(z_3); z_3 = None
mul = relu * 2; relu = None
return (mul,)
循环被展开成了add_1、add_2、add_3。
断图
TorchDynamo遇到编译期无法确定的情况,就会回退到eager模式,触发graph break(断图)。常见的触发断图的情况有依赖具体数据的分支和print等。例如依赖输入的分支:
@torch.compile(backend=custom_backend)
def func(x: torch.Tensor, y: torch.Tensor) -> torch.Tensor:
z = x + y
if z.sum() > 0:
return z
return torch.relu(z) * 2
得到的输出为:
FX Graph:
opcode name target args kwargs
------------- ------ ----------------------- ------------ --------
placeholder l_x_ L_x_ () {}
placeholder l_y_ L_y_ () {}
call_function z <built-in function add> (l_x_, l_y_) {}
call_method sum_1 sum (z,) {}
call_function gt <built-in function gt> (sum_1, 0) {}
output output output ((gt, z),) {}
Generated Python code:
def forward(self, L_x_ : torch.Tensor, L_y_ : torch.Tensor):
l_x_ = L_x_
l_y_ = L_y_
z = l_x_ + l_y_; l_x_ = l_y_ = None
sum_1 = z.sum()
gt = sum_1 > 0; sum_1 = None
return (gt, z)
FX Graph:
opcode name target args kwargs
------------- ------ ------------------------------------------------------- --------- --------
placeholder l_z_ L_z_ () {}
call_function relu <built-in method relu of type object at 0x7f83bc38b4a0> (l_z_,) {}
call_function mul <built-in function mul> (relu, 2) {}
output output output ((mul,),) {}
Generated Python code:
def forward(self, L_z_ : torch.Tensor):
l_z_ = L_z_
relu = torch.relu(l_z_); l_z_ = None
mul = relu * 2; relu = None
return (mul,)
可以看到TorchDynamo最终调用了两次backend,这是因为被编译区域中间发生了断图,使其一分为二,变为两个子图。
断图会影响backend进行优化(例如这样就无法融合跨越边界的那些算子)。更进一步,断图对性能的影响有:
- 需要在GPU和CPU之间同步(如果触发断图的是
.item()和print(tensor)这类东西) - 编译区域被切碎,导致自动优化无法更好地顾及全局,做出更好的优化
- 可能有额外的Python函数调用开销,以及kernel launch开销。
预防与观测断图
可以使用fullgraph=True参数,强制生成完整的一张图,如果遇到断图就会直接报错,来在开发阶段预防断图:
@torch.compile(fullgraph=True)
def func(*args, **kwargs):
...
我们也可以在运行时使用TORCH_LOGS=graph_breaks环境变量来查看所有断图位置和原因:
TORCH_LOGS=graph_breaks python main.py
输出类似于:
[graph_breaks] Graph break at line 15:
Reason: call to print() [User code]
[graph_breaks] Graph break at line 23:
Reason: data-dependent control flow (if tensor.item() > 0)
守卫机制
TorchDynamo不会对一个已经送到后端处理的图,重新送到后端处理,而是复用编译好的缓存。但是如何判断不需要进行重编译呢?TorchDynamo使用Guard(守卫)机制来进行判断。
从直觉上来说,涉及到的tensor的shape不变,只变数值,那么肯定是不需要进行重编译的。类似于这种感觉,TorchDynamo设计了四种守卫:
- Tensor属性守卫:会检查每个Tensor的基本属性,
dtype、device、shape和stride。会被检查的tensor是在计算图中处于输入节点等一切无法被推导出完整基本属性的tensor。如果没有变化就不会进行重编译。 - 全局状态守卫:检查
torch.is_grad_enabled()等全局作用量是否变化。 - 类型守卫:检查tensor的类型(Python类型,与tensor属性守卫的dtype相区别),如检查是不是保持
nn.Parameter类型等。 - 动态形状守卫:如果启用了动态形状,即允许shape进行一些变化而使用同一个编译好的结果(一般是shape满足一个约束范围)。则会检查shape是否满足这个约束范围。
守卫是运行时检查的(具体来说,是调用时检查),如果通过Python包装的方式来去实现判断的话,开销是比较大的(考虑到编译出来的Op可能会被频繁调用)。因此守卫是会被代码生成并编译的,它的速度非常快。
AOTAutograd/AOTDispatcher
AOTAutograd和AOTDispatcher实际上是同一个东西(参考文献:这篇官方文档)。早期叫AOTAutograd,强调其在整个编译流程中处理反向传播(生成反向计算图)。但是其实它的工作不仅仅是处理反向传播,因此后续改名叫AOTDispatcher。现在依然经常被叫做AOTAutograd,可以说是一种历史遗留问题。
AOT是Ahead-Of-Time的缩写,含义是在运行之前编译。这就很奇怪了,因为torch.compile实际上对于一个Python程序来说是JIT编译,为什么这里torch.compile的其中一个流程会被叫做AOT?我的想法是,这里的AOT是针对被编译的这一个对象,或者说是对于这个编译系统来说的。这个系统仍然处于运行结果之前的状态,是先编译,再运行的模式,因此叫AOT;而对于调用torch.compile的Python程序来说,它处于运行状态,要用到这个东西,却没有编译,需要在运行时要用的时候临时编译这一段,因此是一种JIT的范式。
TorchDispatcher
在介绍AOTDispatcher之前,先介绍TorchDispatcher。这个机制是PyTorch为了管理多后端算子的一个重要的抽象层。试想,我们的每一个算子,需要管理:多后端(CPU、CUDA、MPS…),自动微分(多后端的自动微分,AutogradCPU、AutogradCUDA…)。如果在一个函数里实现这些内容,那会非常混乱。如果直接写在调用算子的函数里,可能会变成这样(非常伪的伪代码):
def op(*args, **kwargs):
if 需要自动微分:
总之这里要处理一堆有关自动微分的东西但是我学艺不精也写不出来个所以然所以就这样了当然这里肯定不是调用什么autograd kernel但是总之肯定和普通的kernel不一样
if 输入都是CPU:
op_cpu_autograd_kernel(*args, **kwargs)
elif 输入都是CUDA:
op_cuda_autograd_kernel(*args, **kwargs)
elif 输入都是MPS:
op_mps_autograd_kernel(*args, **kwargs)
...
else:
if 输入都是CPU:
op_cpu_kernel(*args, **kwargs)
elif 输入都是CUDA:
op_cuda_kernel(*args, **kwargs)
elif 输入都是MPS:
op_mps_kernel(*args, **kwargs)
...
每个op都要这么写就显得非常的丑陋,如果要扩展后端也非常的麻烦。于是我们觉得这样肯定不行。借鉴了一下C++虚函数的方法,PyTorch整了个叫Dispatcher的抽象层。它的作用有两个:
- 分析输入输出,将一个Op分配到一个合适的后端kernel实现。
- 处理自动微分Autograd,自动精度转换Autocast等Op调用与具体实现中间的功能层。
参考C++虚函数表,每个算子都有自己的一个表。这个表的key是不同的后端名称,value是可执行的kernel实现。从全局来看,是这样的:
dispatch_table[operator][dispatch_key] = executable_kernel
画出这张表,可以是这样的:
| key | value |
|---|---|
| CPU | CPU上执行的kernel |
| CUDA | CUDA上执行的kernel |
| MPS | MPS上执行的kernel |
当调用某个op,Dispatcher会通过输入Tensor的属性,获得当前应该使用哪个后端实现(即Dispatch Key)。具体来说,每个Tensor都有自己的Dispatch Key Set,Dispatcher会收集这些Dispatch Key Sets,合并成一个最终的Dispatch Key Set,此时第一个元素就是优先级最高的key(所有的Dispatch Key Set内部元素隐式有序)。
这样确实是处理了Op到具体kernel的映射,但是如何兼容合并中间的功能层呢?这个简单,直接把处理autograd等功能层的handler函数也作为一对key-value就行了。所以实际上会是:
| key | value |
|---|---|
| AutogradCUDA | 处理CUDA上自动微分的handler |
| AutogradCPU | 处理CPU上自动微分的handler |
| CPU | CPU上执行的kernel |
| CUDA | CUDA上执行的kernel |
| MPS | MPS上执行的kernel |
这个时候,value就不再是可执行的kernel(的函数指针),而是一类Dispatch Function(的函数指针)。它可以是可执行kernel,可以是实现某个具体功能的handler。
如果一个op被dispatch到了某个功能handler函数上,那不就无法具体执行了吗?因为真正做运算的只有kernel,而不是handler。如果每个handler既要处理自己的功能,又要帮op找到它们的具体执行kernel,就会带来功能层与执行层耦合的问题。导致功能的复合和模块化难以清晰地实现。所以,在功能handler函数中,有一步叫redispatch。
具体来说,功能函数的执行大概是这样的步骤:
- 实现对应的功能
- 把自己对应的dispatch key加入进一个本地屏蔽表(local exclude set)
- 重新dispatch(redispatch)
由于将自己加入了local exclude set,意味着自己不会被再次选中,redispatch时会顺延到自己之后的dispatch路径。这也可以实现功能的叠加:只需要把需要叠加的功能都放到具体kernel实现的前面,就会递归式地实现功能,最终redispatch多次后选择具体kernel实现的路径。
规范化Dynamo FX Graph
ATen化
AOTDispatcher得到的FX Graph实际上是比较高层次的(被称为Dynamo FX Graph)。它的节点所包含的算子往往是PyTorch的前端接口算子(例如torch.add、torch.nn.functional.linear等)。而AOTDispatcher所做的工作就是把它们转成ATen Ops。ATen算子实际上才是PyTorch最底层的算子。调用前端的计算实际上最终会被patch到一个或多个ATen Ops。而把Dynamo FX Graph转成只用ATen Ops的ATen Ops FX Graph,就是AOTDispatcher的核心工作。
实际上,这个规范化,是AOTDispatcher利用TorchDispatcher的能力进行实现的。具体来说,是当AOTDispatcher拿到Dynamo FX Graph后,会按照Dynamo FX Graph运行一遍。在运行过程中,必然会调用Torch Dispatcher进行算子实现的运行时分派,这会被AOTDispatcher截获。一方面,AOTDispatcher可以得知具体是调用了哪个ATen Op(这是显然的,因为到Torch Dispatcher这层就处理的是ATen Op了);另一方面,AOTDispatcher可以利用TorchDispatcher中的那些功能层内容,特别是Autograd。
截获方式是通过ProxyTensor和FakeTensor,对Dynamo FX Graph进行假运行:
- ProxyTensor:用于记录具体调用了哪个ATen Op,生成FX节点
- FakeTensor:没有数据,只有元信息(shape,dtype等)的Tensor,用于推导运算后的Tensor元信息
这其实相当于AOTDispatcher是一个使用类似hook的设计来去重用TorchDispatcher逻辑的模块。叫Dispatcher确实比叫Autograd更合理。
Functionalization
Dynamo FX Graph中的调用方式经常包含各种mutation式的调用,例如:
x += y # or x.add_(y)
这种inplace操作对编译器来说是不友好的(更难处理数据依赖关系)。因此需要将其转化为函数式的形式,逻辑上是这种:
new_x = torch.ops.aten.add.Tensor(x, y)
这属于一种编译前的语义规范化。
Decomposition
AOTDispatcher会将有些高级的Op展开成更基本的组成Ops,例如将aten.linear转换成aten.t与aten.addmm。这样后端更好优化,方便进行算子融合等的分析。
在规范化之后
AOTDispatcher可以利用TorchDispatcher的Autograd来实现生成反向图。这就可以与正向图合成一张正向+反向的joint graph,方便从全局进行考虑。当然了,如果在不需要梯度的情况下,AOTDispatcher就不会生成backward graph。
AOTDispatcher可以做到:
- 重计算策略:对于便宜的算子(如 relu),不保存结果,后向时重算
- 内存布局优化:统一规划前后向的内存布局,减少 transpose
重计算策略采用min-cut(最小割)方法实现。目标是最小化跨越cut的张量存储大小。展开讲又是一篇文章(当然了,肯定不是因为我懒得学),可以阅读这个discuss:戳我
总之我们现在将Dynamo FX Graph成功转换成了一个只包含ATen Ops,functional式调用,展开成低级Ops的FX Graph,并且如果启用梯度,那么还是一张Forward Graph加上一张Backward Graph的Joint Graph。在生成图时还考虑了最小化中间张量存储的策略。
TorchInductor
TorchInductor是PyTorch 2.x这套编译系统的后端。负责算子融合、生成Tirton代码并编译的工作。它接收ATen Ops FX Graph(也叫做Core ATen IR),然后会做如下工作:
- 继续Decomposition(是的有一部分工作依然会在TorchInductor这里去做),不过目的和AOTDispatcher阶段的Decomposition有一些不同。AOTDispatcher阶段的Decomposition是为了规范化和好分析;TorchInductor阶段是为了决定如何在硬件上运行。
- 融合分析:识别可以进行算子融合的模式,并进行算子融合。
- 调度决策(超参数生成):决定block size、tiling方式,vectorization等
- 代码生成:生成Triton/C++代码
- 编译:调用Triton Compiler/nvcc/g++
- 加载:Binding到Python层(生成一个Python warpper)
融合分析
TorchInductor将算子分为三类:
- Pointwise(逐点):在别的语境下也会被叫做element-wise。指输出的每个元素只依赖对应位置的输入的算子。例如add、mul、relu、exp、tanh等。
- Reduction(规约):输出元素依赖多个输入元素(元素依赖中内含Reduce型的拓扑)。例如:sum、softmax、layernorm等。
- Template(模板):复杂的结构化计算,有专门的实现。例如:matmul,conv2d等。
对于Pointwise类型的多个算子,融合策略是直接将它们的计算串联,写在一个kernel内;对于Reduction类型的算子,融合策略则是将它与后续的Pointwise类型算子进行融合(可以说是FlashAttention能实现算子融合的某个更本质原因之一了);对于Template类型的算子,融合策略是直接调库,一般不进行融合。
# 可以融合
x = input + bias # pointwise
y = x.relu() # pointwise
z = y * scale # pointwise
# → 融合成一个 kernel
# 可以融合(persistent reduction)
sum_val = x.sum(dim=-1, keepdim=True) # reduction
normalized = x / sum_val # pointwise
# → 融合成一个 kernel
# 不能融合
y = x @ weight # template (matmul)
z = y.relu() # pointwise
# → matmul 单独调用 cuBLAS,relu 是独立的 kernel
调度决策
GPU路径
需要决定:
- Block size:每个线程块有多少个线程。与CUDA、Triton中相同的概念,太小会导致并行度不足,太大会导致occupancy下降(寄存器资源压力太大)。
- Tiling:数据如何分块。目标是贴合memory hierarchy,最大化利用L1/L2 cache。
- 向量化:一次加载多少元素。增大合并访问程度,减少内存事务开销,提高带宽利用率。
生成的Triton代码示例:
@triton.jit
def fused_add_relu_mul(
in_ptr, out_ptr,
n_elements,
bias: tl.constexpr,
scale: tl.constexpr,
BLOCK_SIZE: tl.constexpr
):
pid = tl.program_id(0)
block_start = pid * BLOCK_SIZE
offsets = block_start + tl.arange(0, BLOCK_SIZE)
mask = offsets < n_elements
# Load
x = tl.load(in_ptr + offsets, mask=mask)
# Compute (融合的三个操作)
x = x + bias
x = tl.maximum(x, 0.0)
x = x * scale
# Store
tl.store(out_ptr + offsets, x, mask=mask)
CPU路径
使用OpenMP实现并行化。
需要决定:
- 线程数:通常与CPU核数相等。
- 向量化:使用SIMD指令(AVX),利用向量硬件进行加速。
- 循环分块:适应memory hierarchy,最大化利用L1/L2/L3 cache。
生成的代码示例:
#pragma omp parallel for
for (int64_t i = 0; i < n_elements; i++) {
float x = input[i];
x = x + bias;
x = std::max(x, 0.0f);
x = x * scale;
output[i] = x;
}
代码生成
对于Triton来说,TorchInductor会填充Triton模板(Jinja2管理),类似于这样:
triton_template = """
@triton.jit
def {{kernel_name}}({{params}}):
pid = tl.program_id(0)
offsets = pid * {{BLOCK_SIZE}} + tl.arange(0, {{BLOCK_SIZE}})
mask = offsets < {{n_elements}}
{% for load in loads %}
{{load.name}} = tl.load({{load.ptr}} + offsets, mask=mask)
{% endfor %}
{% for op in ops %}
{{op}}
{% endfor %}
tl.store({{output_ptr}} + offsets, {{output_var}}, mask=mask)
"""
C++会麻烦一些,这里不赘述了(感觉也不是很重要,我不想学了)。
编译与缓存
两条路径,分别对应Triton与C++:
- Triton代码 -> Triton Compiler -> ptx -> cuda driver compile -> cubin
- C++代码 -> g++/clang -> .so
编译结果将会被缓存(offload到磁盘中),使用代码的hash进行查找和防重复。
Autotuning
TorchInductor支持自动调优(Autotuning)。对于一个kernel,根据不同的block size和tiling生成多个不同的版本,随后实际运行测试真实性能,最终选择最快的版本。可以实现面对不同情景(不同硬件,输入的shape等)的一定程度的泛化能力,并且避免费时费力的人力试错调优。
Autotuning会在第一次编译时进行,随后最快的那个kernel会被缓存。
Summary
在2026年的今天,要编写高性能算子或融合算子,主要有三条路径(参考这篇很好的CUTLASS入门文章,虽然原文是2025年,但是按照我的知识边界来说在2026年依旧是这样的):
- 基于 Python DSL + PTX 编译器的路径。像 triton、CuTe DSL、TileLang、Mojo 等等都是走了这条路径。
- 基于 C++ 模板封装 PTX 的路径。例如 CUTLASS、Thunder Kittens 等。
- 基于自动编译的路径。典型例子是 torch.compile 。
Torch compile由于其易用性(直接搞个装饰器就能用),非常适合作为一个strong baseline。在动手去做人肉triton甚至cuda/cutlass之前,不妨先用一用torch compile确立一个基线。先将graph break解决了,发挥出torch compile的一个比较好的效能,再在torch compile的基础上,去做自动化规则所难以处理的优化(例如提到的Template类型算子的融合问题等)。
不过另外一方面,torch compile实际上压榨了很多简单场景下人力的发挥空间。假如我们的场景仅仅只是在研究算子,cuda graph上的性能问题(这些内容往往在单卡或单节点情境下讨论),通过为torch compile大人扫清障碍(特别是别写你那些个依赖cpu和io的dirty代码),我们就可以拿到一个很逼近算子和graph层面所能提供的优化极限的提升(说是极限,一般也就20%)。当然也并不是完全没机会,多多少少总会有自动化所无法照顾到的东西,运气好一些说不定自动化抽风了(或者确实需要人类智能)还会有比较大的优化点。
底层平台已经很成熟,也正是这种成熟使得人们能更加关注于上层建设。
私货
在写本文的时候发散性地走马观花了其他很多的概念。深深感受到:kernel哥该下岗了。无论是越发完善的机器学习编译体系,还是人力积淀越来越多的各种已经优化得非常极限的成熟算子库,亦或者是这段时间被研究和落地得非常多的kernel agent(孩子们,我真写不过ai)。ML编译大厦已经建成,后人要做的只是修修补补和follow新架构的特性(其实还有国产芯片百废待兴的工具链)。而且写kernel哪有搞scheduler好玩…
最近在研究Mooncake,集群平台的调度其实本质上和算子内的数据调度有同一性(这是某种递归性的东西),我能感受到这是同质的,整个infra的运行,甚至说整个计算机技术的运行建立在一层又一层本质相同的system上,这是一种类似分形几何的hierarchy。或许做system,本质上是一个对于这个在不同领域重复本质的把握,加上某个DSK(Domain Special Knowledge,刚自己造的词)。当然也有很强的人可以热插拔DSK导致做啥都行。
当我写到这里的时候,感觉到非常地疲惫,我或许应该找点东西吃。

欢迎友好讨论~