如何在PyTorch中实现与Numpy原版结果一致的PWCCA?
PWCCA实现与谷歌原版结果对齐方案
误差来源修正
- 修正小值维度过滤逻辑
原版svcca的维度过滤是基于协方差矩阵的特征值,而非协方差对角元,将代码中判断x_idxs、y_idxs的逻辑替换为对sigma_xx做特征分解,保留特征值大于epsilon的维度,和原版逻辑完全对齐。 - 统一全链路数值精度
所有输入张量强制转换为torch.float64类型后再进入计算流程,底层CCA实现的SVD、矩阵求逆等操作也统一用float64精度,避免不同精度带来的数值差异。 - 对齐CCA输出顺序
确保底层cca函数输出的奇异值diag是按从大到小排列,对应投影矩阵a、b的列顺序和原版svcca输出一致,部分PyTorch的SVD实现会返回从小到大排列的奇异值,会直接导致权重计算错误。 - 修复原版接口调用报错
调用原版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
相关产品推荐
相关产品推荐

