076 使用MLIR实现Batch Normalization的融合
从一次推理延迟异常说起
去年在给某款边缘AI芯片做模型部署时,遇到一个诡异现象:同一个ResNet50模型,在PyTorch上跑FP32推理时延迟是12ms,转成我们自研的MLIR-based编译器后,延迟反而飙到了18ms。当时第一反应是“编译器优化出bug了”,但逐层打印IR后发现,问题出在Batch Normalization(BN)层没有被正确融合进卷积层。
BN融合这个操作,理论上能减少一次全局内存读写和一次element-wise计算,在边缘设备上效果尤其明显。但MLIR的pass pipeline里,如果只是简单地把BN拆成乘加操作,反而会因为引入额外的tensor reshape和broadcast导致性能倒退。今天这篇笔记,就聊聊如何在MLIR中正确实现BN融合,以及那些容易踩坑的细节。
BN融合的本质:把“归一化”塞进卷积的权重里
先回忆一下BN的数学形式:
y = gamma * (x - mean) / sqrt(var + eps) + beta
在推理阶段,mean和var是训练好的统计量,所以这个公式可以重写为:
y = (gamma / sqrt(var + eps)) * x + (beta - gamma * mean / sqrt(var + eps))
看到没?这就是一个线性变换:y = a * x + b
订阅专栏 解锁全文

1764

被折叠的 条评论
为什么被折叠?



