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

不同操作系统与设备上JAX线性代数计算结果不一致问题

JAX跨CPU设备数值一致性问题

我有一个基于JAX的确定性程序,涉及大量线性代数运算。在三款不同CPU设备上运行该代码时,单台设备内的输出具有确定性,但不同设备间的结果存在差异:两台MacOS设备(分别搭载Sequoia系统的M1 Pro、Sonoma系统的M2)和一台Linux设备。

最小可复现示例

import jax
import optax
import flax.linen as nn
import jax.numpy as jnp

jax.config.update("jax_enable_x64", True)

variables = jnp.array([0.1, -3 * jnp.pi / 2])


class RNN(nn.Module):
    hidden_size: int
    output_size: int

    @nn.compact
    def __call__(self, input, hidden_state):
        gru_cell = nn.GRUCell(features=self.hidden_size)
        new_hidden_state, _ = gru_cell(hidden_state, input)
        output = nn.Dense(features=self.output_size)(new_hidden_state)
        return output, new_hidden_state

def _optimize(
    loss_fn,
    init_params,
    max_iter,
    learning_rate,
):
    optimizer = optax.adam(learning_rate)
    opt_state = optimizer.init(init_params)

    @jax.jit
    def step(params, state):
        grads = jax.grad(loss_fn)(params)
        updates, new_state = optimizer.update(grads, state, params)
        new_params = optax.apply_updates(params, updates)
        return new_params, new_state

    params = init_params
    for iter_idx in range(max_iter):
        params, opt_state = step(params, opt_state)
    return params, iter_idx + 1

def fun(gamma, delta):
    op = jnp.array([[0, -1j], [1j, 0]])
    angle = (gamma * op) + delta / 2
    return (jax.scipy.linalg.expm(1j * angle) + jax.scipy.linalg.expm(-1j * angle)) / 2

def loss(params):
    rnn = RNN(hidden_size=10, output_size=2)
    input = variables
    hidden_state = jnp.zeros((10,))
    output, _ = rnn.apply({'params': params}, input, hidden_state)
    params_out = output
    return jnp.real(jnp.trace(fun(params_out[0], params_out[1])))

if __name__ == "__main__":
    rng = jax.random.PRNGKey(0)
    rnn = RNN(hidden_size=10, output_size=2)
    input = variables
    hidden_state = jnp.zeros((10,))
    params = rnn.init(rng, input, hidden_state)['params']

    max_iter = 100
    learning_rate = 0.01
    convergence_threshold = 1e-6
    optimized_params, num_iterations = _optimize(
        loss,
        params,
        max_iter,
        learning_rate,
    )
    final_loss = loss(optimized_params)
    print("Final Loss:", final_loss)

设备输出差异

  • MacOS设备输出:-1.9979573829398634
  • Linux设备输出:-1.9979573808129485

二者差异出现在小数点后第8位,若程序规模更大、逻辑更复杂,数值差异可能显著扩大。在机器学习场景中,若模型收敛路径复杂且存在多个局部极小值,这类数值差异可能导致模型最终收敛至不同的极小值点。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.12 20:52:20