#506·timesfm

MLX后端与PyTorch在4个位置存在差异(stitching,CPM RevIN精炼,剩余活化,NaN处理)

作者: matdou创建于 2026年9月12日更新于 2026年9月12日

一直与TimesFM3的MLX后端合作,并找到四个地点,在相同的重量和输入上,它能产生与PyTorch后端不同的数字.

为了隔离这一点,我在两个后端都建了一个微小的模型,通过安全器回转将PyTorch的权重复制到MLX模型中,并证实了默认的-config decode ()调用匹配(到~1e-6). 然后我一次翻出一个设定:

  1. " 使用-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.

  1. 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.

  1. 在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