同一函数不同输入尺寸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
相关产品推荐
相关产品推荐

