059、自动微分(AutoDiff)在MLIR中的表示

发布时间:2026/8/20 17:09:46
059、自动微分(AutoDiff)在MLIR中的表示 059、自动微分(AutoDiff)在MLIR中的表示上周调试一个自定义算子反向传播时,发现MLIR生成的梯度计算图里多了一个莫名其妙的linalg.generic操作,把输入张量原地改写了。查了两天,最后定位到是自动微分pass在lowering过程中把某个中间变量的梯度累加逻辑错误地折叠进了前向计算。这种问题在手工写反向传播时几乎不会犯,但一旦交给编译器自动推导,中间表示的细节就变得极其关键。今天这篇笔记,就围绕MLIR里自动微分的表示层展开。不聊高深的数学推导,只讲在IR层面,梯度是怎么被“翻译”成一组操作的。从“链式法则”到“操作图”自动微分在编译器层面的本质,是把数值计算图的正向传播和反向传播都显式地表示为IR中的操作序列。MLIR没有像PyTorch那样在运行时动态构建梯度图,而是在编译期通过pass对函数体进行变换。一个典型的正向计算:c = a + b,在MLIR中对应一个arith.addf操作。自动微分pass会为这个操作生成对应的反向操作——通常是arith.addf的梯度就是它自身(因为加法对两个输入的偏导都是1),但实际实现中,梯度会以“梯度累加”的形式注入到a和b的梯度变量中。这里有个关键点:MLIR的自动微分不是简单地为每个操作生成一个反向操作,而是把整个计算图看作一个有向无环图,然后为每个值(Value)维护一个梯度缓冲区。这些缓冲区在IR中表现为额外的memref/