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

如何在PyTorch中实现与Numpy原版结果一致的PWCCA?

PWCCA实现与谷歌原版结果对齐方案

误差来源修正

  1. 修正小值维度过滤逻辑
    原版svcca的维度过滤是基于协方差矩阵的特征值,而非协方差对角元,将代码中判断x_idxs、y_idxs的逻辑替换为对sigma_xx做特征分解,保留特征值大于epsilon的维度,和原版逻辑完全对齐。
  2. 统一全链路数值精度
    所有输入张量强制转换为torch.float64类型后再进入计算流程,底层CCA实现的SVD、矩阵求逆等操作也统一用float64精度,避免不同精度带来的数值差异。
  3. 对齐CCA输出顺序
    确保底层cca函数输出的奇异值diag是按从大到小排列,对应投影矩阵a、b的列顺序和原版svcca输出一致,部分PyTorch的SVD实现会返回从小到大排列的奇异值,会直接导致权重计算错误。
  4. 修复原版接口调用报错
    调用原版svcca接口时,给转numpy的操作加.copy(),确保传入的是连续数组:
acts1: np.ndarray = L1.T.detach().cpu().numpy().copy()
acts2: np.ndarray = L2.T.detach().cpu().numpy().copy()

性能优化

自定义PyTorch实现运行慢的问题,可以用(x.T @ x)/(x.size(0)-1)替代torch.cov计算协方差,计算效率会有数量级提升。

误差说明

numpy和PyTorch的底层线性代数库实现存在固有数值差异,不可能做到100%完全对齐,当前0.001量级的误差已经属于可接受范围,不会影响实际实验的结论判断。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.09.26 09:24:02