MLX后端与PyTorch在4个位置存在差异(stitching,CPM RevIN精炼,剩余活化,NaN处理)
作者: matdou创建于 2026年9月12日更新于 2026年9月12日
一直与TimesFM3的MLX后端合作,并找到四个地点,在相同的重量和输入上,它能产生与PyTorch后端不同的数字.
为了隔离这一点,我在两个后端都建了一个微小的模型,通过安全器回转将PyTorch的权重复制到MLX模型中,并证实了默认的-config decode ()调用匹配(到~1e-6). 然后我一次翻出一个设定:
- " 使用-stitching=虚假 " 已解析,但从未读取。 'decode ()' 总是运行基于缝合的公式:
# mlx/ model.py, 解码 ()
提取 len = min (2 * p, cfg.output patch len) 互联网档案馆的存檔,存档日期2013-12-02.
重叠=提取 len-p
num forecast patches = 最大 (math.ceil ((范围 - 重叠) / p), 1) ;
num hor patches = num forecast patches + capg.rolls - 1 (中文(简体) ).
添加 h = num hor patches * p
hor pad = 添加 h - 地平线
# capg.use stitching 从未检查过
Max abs diff对PyTorch的测试配置:3.68.
- 在
TimesFM3MlxConfig'中不存在使用'一词。 精细步骤在任何时间都有一个patch cpm mask'存在.它总是在每一个decode()'的呼号上:
[Python]
mlx/ model.py, forward logits()
如果补丁 cpm mask不是无: ref mu, ref sigma = cpm revin refine lib.cpm iterative revin refine (中文(简体) ). 原始, 运行 n, mu, sigma, 补丁 cpm msk,... (中文(简体) ).
最大 abs diff: 0.053 (中文(简体) ).
3. `ResidualBlock'硬码ReLU,所以`shish'实际上运行ReLU:
```2ZZ
# mlx/腾讯.py
def call (自, x: mx. array) - > mx. array:
返回自.output layer(nn.relu (self.hidden layer (x))) +自.residual layer(x)
以"活化=""活化"为主的"最大活化":0.997.
- 在RevIN统计之前没有NaN消毒。 PyTorch 的“ 前向( ) ” 这样做 :
# 火炬/ model.py, 前进( )
值 = 火炬.nan to num(值,nan=0.0)
值 = 火炬.clamp(值, -self.value clip,自.value clip)
mlx/model.py'的forward logits()'没有等同性。 它会直接从原始输入中计算数据。 1个未蒙面的NaN毒害了累积的正/快. Pytorch预测的情况并非如此。
我在“ResidualBlock”中为所有四个系统进行了修正,匹配了PyTorch逻辑、激活/前置/标识 skip支持,以及一个NaN sanitize+clip步骤镜像“torch/model.py”外加一个检查等效的回归测试文件。 我周末晚些时候再开公关.
内容来源: google-research/timesfm