使用torch.matmul计算两个Tensor时出现NAN值的原因求助
解决思路分析
首先明确两种矩阵乘法的维度差异:
torch.matmul(recon_1.t(), x):recon_1转置后为1433×2708,与2708×1433的x相乘,得到1433×1433的矩阵,每个元素是**recon_1的某列与x的某列的点积**。torch.matmul(recon_1, x.t()):x转置后为1433×2708,与2708×1433的recon_1相乘,得到2708×2708的矩阵,每个元素是**recon_1的某行与x的某行的点积**。
NaN出现的核心原因是两种运算的数值分布、精度累积差异,具体排查方向如下:
1. 排查数值溢出/下溢
第一种运算的点积是对2708个元素求和(列向量长度),如果recon_1或x的列向量存在极大/极小值,累积求和后可能超出浮点类型的范围(比如float32的最大值约为3.4e38),产生inf,后续运算会转化为NaN。而第二种运算的行向量点积数值范围更稳定。
- 操作:
- 检查列/行向量的极值:
# 列向量最大绝对值 recon_col_max = recon_1.abs().max(dim=0)[0] x_col_max = x.abs().max(dim=0)[0] # 行向量最大绝对值 recon_row_max = recon_1.abs().max(dim=1)[0] x_row_max = x.abs().max(dim=1)[0] - 对比列与行的极值差异,若列极值远大于行极值,尝试对列做归一化(如
torch.nn.functional.normalize(recon_1, dim=0)),或切换为更高精度的float64dtype计算。
- 检查列/行向量的极值:
2. 检查浮点精度丢失
float32的有效位数有限,当1433个元素的乘积累积求和时,舍入误差可能导致数值异常(比如极小值被放大、正负值抵消后出现NaN)。
- 操作:
- 转换为
float64重试:result = torch.matmul(recon_1.t().double(), x.double()) print(torch.any(torch.isnan(result))) - 若转换后NaN消失,说明是精度问题,可在计算阶段临时切换 dtype,或对张量做缩放处理(如整体除以最大极值)。
- 转换为
3. 定位异常元素对
既然第二种运算无NaN,说明原始张量中不存在NaN/inf,可定位第一种运算中产生NaN的具体位置,分析对应列向量的数值:
- 操作:
- 分步计算点积,找到异常列对:
recon_t = recon_1.t() for i in range(recon_t.shape[0]): for j in range(x.shape[1]): dot = torch.dot(recon_t[i], x[:, j]) if torch.isnan(dot): print(f"NaN at column pair ({i}, {j})") print("recon_1 column:", recon_1[:, i]) print("x column:", x[:, j]) exit() - 针对异常列向量,查看是否存在极端值、大量零值或正负值完全抵消的情况,针对性处理(如过滤异常值、调整预处理逻辑)。
- 分步计算点积,找到异常列对:
4. 训练场景下的梯度问题(若涉及反向传播)
如果是在模型训练中出现NaN,可能是第一种运算的梯度传播导致爆炸:1433×1433的矩阵梯度在反向传播时,累积效应可能比2708×2708更显著,触发梯度溢出。
- 操作:
- 使用梯度检测工具追踪:
with torch.autograd.detect_anomaly(): # 执行前向计算和反向传播 loss = ... # 基于torch.matmul(recon_1.t(), x)的损失 loss.backward() - 添加梯度裁剪限制梯度范围:
torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0)
- 使用梯度检测工具追踪:
内容的提问来源于stack exchange,提问作者Aaron
相关产品推荐
相关产品推荐

