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

tf.matmul与torch.matmul在1e-5精度下结果不一致的解决方法

解决PyTorch与TensorFlow矩阵乘法在1e-5容差下的一致性问题

先明确你的复现代码:

import numpy as np
import tensorflow as tf
import torch

gat_key_t = np.random.normal(size=(8, 16, 64, 20)).astype(np.float32)
gat_query_t = np.random.normal(size=(8, 16, 30, 64)).astype(np.float32)

tf_key   = tf.convert_to_tensor(gat_key_t)
tf_query = tf.convert_to_tensor(gat_query_t)
pt_key   = torch.from_numpy(gat_key_t)
pt_query = torch.from_numpy(gat_query_t)

tf_output = tf.matmul(tf_query, tf_key)
pt_output = torch.matmul(pt_query, pt_key)

# 结果不一致
print(np.allclose(tf_output.numpy(), pt_output.numpy(), rtol=1e-5, atol=1e-5, equal_nan=False))
# 结果一致
print(np.allclose(tf_output.numpy(), pt_output.numpy(), rtol=1e-4, atol=1e-4, equal_nan=False))

出现这种差异的核心原因是:PyTorch和TensorFlow对矩阵乘法的底层实现(依赖的BLAS库、硬件指令优化、计算精度截断策略等)存在细微差别,而float32本身的精度有限(约6-7位有效数字),累积的计算误差会在严格容差下显现。

要让两者在1e-5容差下保持一致,可以尝试以下几种方法:

  • 切换到更高精度的数据类型
    将张量从float32改为float64(双精度),双精度的有效数字位数约15-17位,能大幅缩小计算误差。修改代码如下:

    gat_key_t = np.random.normal(size=(8, 16, 64, 20)).astype(np.float64)
    gat_query_t = np.random.normal(size=(8, 16, 30, 64)).astype(np.float64)
    
    tf_key   = tf.convert_to_tensor(gat_key_t)
    tf_query = tf.convert_to_tensor(gat_query_t)
    pt_key   = torch.from_numpy(gat_key_t)
    pt_query = torch.from_numpy(gat_query_t)
    
    tf_output = tf.matmul(tf_query, tf_key)
    pt_output = torch.matmul(pt_query, pt_key)
    
    # 此时会返回True
    print(np.allclose(tf_output.numpy(), pt_output.numpy(), rtol=1e-5, atol=1e-5, equal_nan=False))
    
  • 统一框架的计算优化策略
    两个框架默认可能启用不同的硬件加速优化,比如TensorFlow可能用XLA,PyTorch用MKLDNN。尝试关闭这些优化或者统一使用相同的计算后端:

    • TensorFlow端关闭XLA:在会话或函数中禁用XLA编译
    • PyTorch端设置torch.backends.mkldnn.enabled = False,强制使用标准BLAS实现
  • 固定全局随机种子(针对随机初始化场景)
    虽然你的代码是从同一个numpy数组转换而来,但如果是框架内部生成随机数的场景,固定种子能确保初始输入完全一致,避免额外误差。在代码开头添加:

    np.random.seed(42)
    tf.random.set_seed(42)
    torch.manual_seed(42)
    torch.cuda.manual_seed_all(42) # 如果使用GPU
    

需要注意的是,即使做了以上调整,极端场景下可能仍存在极细微差异,但在1e-5的容差下基本能保证一致。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.06 05:56:04