You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

同一函数不同输入尺寸VMAP后GPU float32结果偏差的疑问

关于JAX/Equinox VMAP下GPU float32数值差异的解释

问题1:该现象是否属于预期情况?

是,这属于GPU单精度浮点运算下的预期数值行为,并非框架或代码bug。

问题2:不一致性的来源是什么?

核心原因是GPU并行计算特性与float32精度限制的共同作用:

  • 浮点运算的非精确性:float32仅保留约7位有效数字,加法、乘法等运算的舍入误差会随计算顺序变化而累积。
  • GPU的SIMD并行架构:处理不同大小的batch时,JAX会根据输入规模调整线程块分配、内存布局及算子优化策略(比如调用cuBLAS/cuDNN的不同优化实现),这会直接改变运算执行顺序(例如大batch的并行累加顺序与小batch完全不同)。而浮点数的加法/乘法不满足严格的交换律和结合律,最终导致细微数值偏差。
  • VMAP的优化逻辑:JAX的vmap(包括Equinox的filter_vmap)会根据输入batch大小生成不同的优化计算图,大batch可能触发算子融合、更高程度的向量化并行计算,这些优化在提升效率的同时,也会引入与小batch不同的舍入误差路径。

问题3:为何仅在GPU上出现,CPU上无此问题?

主要源于CPU与GPU的架构差异及优化策略不同:

  • 计算架构:CPU的浮点运算单元(FPU)偏向串行或小批量并行,计算顺序相对固定;部分CPU还支持用扩展精度(如80位x87浮点数)进行中间累加,大幅减少舍入误差累积。而GPU为最大化并行吞吐量,采用大规模SIMD并行,强制改变运算顺序,导致误差差异被放大。
  • 算子实现:JAX在CPU上的算子优化更偏向保证计算顺序一致性,而GPU上的算子优先追求并行效率,会使用更多改变计算路径的优化策略(如张量核心、批量融合运算),这些策略会带来更明显的数值差异。

补充验证说明

  • 启用float64后差异消失:双精度浮点数有15-17位有效数字,运算顺序带来的舍入误差远小于单精度,因此差异会被掩盖到可忽略的程度。
  • 复现代码中部分输出为0:这是因为这些样本的计算路径恰好未因batch大小变化改变运算顺序,或误差被抵消,属于随机现象。

复现代码

import jax
import jax.numpy as jnp
import equinox as eqx

def equinox_vmap(x, mlp):
    out = eqx.filter_vmap(mlp.__call__)(x)
    return out

key = jax.random.PRNGKey(0)
key, network_key = jax.random.split(key, 2)
mlp = eqx.nn.MLP(2, 2, 10, 2, key=network_key)

key, key_x = jax.random.split(key, 2)
x = jax.random.normal(key_x, (10000, 2))

error_eqx = equinox_vmap(x[:10], mlp) - equinox_vmap(x, mlp)[:10]
print("eqx error:", error_eqx)

运行输出

eqx error: [[-1.4442205e-04  1.0999292e-04]
 [-5.9515238e-05 -9.1716647e-06]
 [ 1.4841557e-05  5.6132674e-05]
 [ 0.0000000e+00  0.0000000e+00]
 [-9.1642141e-06 -2.5466084e-05]
 [ 3.8832426e-05 -3.3110380e-05]
 [ 3.3825636e-05 -2.4946406e-05]
 [ 4.0918589e-05 -3.2216311e-05]
 [ 1.3601780e-04  8.7693334e-06]
 [ 0.0000000e+00  0.0000000e+00]]

环境配置

  • GPU:NVIDIA RTX 6000 Ada Generation 48 GB
  • Python 3.11.11
  • 依赖版本:
equinox                  0.13.0
jax                      0.7.0
jax-cuda12-pjrt          0.7.0
jax-cuda12-plugin        0.7.0
jaxlib                   0.7.0
jaxtyping                0.3.2
ml_dtypes                0.5.3
numpy                    2.3.2
nvidia-cublas-cu12       12.9.1.4
nvidia-cuda-cupti-cu12   12.9.79
nvidia-cuda-nvcc-cu12    12.9.86
nvidia-cuda-nvrtc-cu12   12.9.86
nvidia-cuda-runtime-cu12 12.9.79
nvidia-cudnn-cu12        9.11.0.98
nvidia-cufft-cu12        11.4.1.4
nvidia-cusolver-cu12     11.7.5.82
nvidia-cusparse-cu12     12.5.10.65
nvidia-nccl-cu12         2.27.6
nvidia-nvjitlink-cu12    12.9.86
nvidia-nvshmem-cu12      3.3.9
opt_einsum               3.4.0
pip                      24.0
scipy                    1.16.1
setuptools               65.5.0
typing_extensions        4.14.1
wadler_lindig            0.1.7

内容的提问来源于stack exchange,提问作者hvater

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.06.12 13:44:55