Numpy.tensordot报错tuple index out of range,如何实现PyTorch conv2d向量化版本
错误原因
IndexError: tuple index out of range错误核心是轴索引和张量实际维度不匹配,优先做两个排查:- 打印
blocks.ndim和filters.ndim确认实际维度是否和预期的6维、4维一致,大概率是生成blocks时维度少了一维,比如实际只有5维,访问索引5就会超出范围 - 确认
blocks的第3、4、5维尺寸和filters的第0、1、2维尺寸完全匹配
- 打印
修复方案
你写的tensordot语法本身没有问题,确认张量维度和尺寸符合预期后即可正常运行,验证示例如下:
import numpy as np # 按你给出的维度随机生成测试数据 a,b,c,d,e,f,g = 2,3,4,3,3,3,16 blocks = np.random.randn(a,b,c,d,e,f) filters = np.random.randn(d,e,f,g) output = np.tensordot(blocks, filters, axes=([3,4,5], [0,1,2])) print(output.shape) # 输出(2, 3, 4, 16),符合你期望的(a,b,c,g)尺寸
np.einsum替代实现
完全可以用np.einsum实现,写法更直观不易出错,等价实现代码如下:
output_einsum = np.einsum('abcdef,defg->abcg', blocks, filters)
可以用np.allclose(output, output_einsum)验证两种写法的计算结果完全一致。
内容的提问来源于stack exchange,提问作者wsun88
相关产品推荐
相关产品推荐

