#4503·mlx

mx.compile 会用 7 位有效数字对浮点标量常量进行内联,因此编译后的结果与即时计算结果相差 1 个 ulp

作者: freddyhaddad创建于 2026年9月14日更新于 2026年9月14日

□ 描述错误

将 Python 浮点数(或 0-d 常数阵列) 打印到生成的内核源的已编译函数,其中“ sted: set精确度(std:: numeric limits : digiters10 + 1)” (“ mlx/后端/common/compited.h”,“print float constant ”)。 这是7个重要数字,但 " float32 " 需要 " max digiters10 " (9),因此,金属编译器将 " 1/3 " 、 " 128 **-0.5 " 或 " 1 sqrt(2) " 等常数分析为相邻浮点。 然后编译的函数与几乎每个元素上 1 的急切执行不同。 来回旅行(0.3'、0.1'、`1e-6')的七位数小数形式的常数不受影响,因此在试验中很容易错过。

同样的二分法适用于 " 二分法 " ( " 数字10+1 " =16; " 最大-数字10 " =17)。

□ 重现

[Python] 导入 mlx.core 为 mx

mx.random. seed( 0) (中文(简体) ). x = mx.random.ormal ((65536)) (中文(简体) ). r = mx.random. ormal ((65536)). mx.eval(x, r) (中文(简体) ).

(0.3, 1/3, 128 **-0.5, 0.7071067811865476): 热=(x * r) * s 已编译 = mx.compile(lambda x, r, s=s:(x * r) * s (x, r) mx.eval( 已编辑) print(f"s={s!r}): {int (eager!=编译). sum ()}65536个元素不同").

{\fn方正粗倩简体\fs12\an8\1cHFFFF00\b0}同一个常数通过一个输入而不是捕获: 完全相同 sa = mx.array (128 ** - 0.5, mx.float32) (中文(简体) ). 热=(x *r) * (128 **-0.5) 已编译 = mx.compile(lambda x, r, sa:(x * r) * sa (x, r, sa) mx.eval( 已编辑) 打印(“ 作为输入 ” ), int( (eager!=编译)). sum())


mlx0.32.2的产出:

s=0.3:65536个元素中的0个不同 s=0.33333333333333333333:65536元素中的65536有差异 s=0.08838834764831845:65536个元素中的60474个元素不同 s=0.7071067811865476: 65536个元素中的60474有差异 作为输入: 0


□ 预期行为

所编译的函数应该计算出它急切的形态;所捕获的常数应该完全嵌入. 使用'std:::numeric limits<T>:::max digiters10'在'print float constant'(或发出'std::hexfloat')中使嵌入式文字可以进行回转.

□ 桌面

- 操作系统:macOS 26.4
- 芯片:苹果M3 Ultra
- Python 3.13,mlx 0.32.2 (自2026-09-15起,也载于“主要”的“编译后h”中)

□ 附加上下文

在编译线性意向解码步骤(`x * rsqrt( sum(x*x) + eps) * head dim**-0.5')并对照急切进行比特-比特检查时发现:在计算了JIT下的'mx.sigmoid''快"exp"后,这是最后剩下的差. 作为数组输入在它周围工作,通过比例表.