PyTorch中索引切片得到的张量使用.t()转置报out of range错误怎么办?
PyTorch张量转置报错解决方案
报错触发原因
结合你使用的PyTorch版本(从报错路径的构建时间判断为0.4.x早期版本),该报错有三个常见触发原因:
.t()方法和torch.t()函数仅支持二维张量转置,如果你的tensordata实际维度为1维或者3维及以上,调用转置时就会触发边界检查报错。你可以先执行print(tensordata.shape)确认张量的实际维度是否符合预期。- 你通过索引生成的
tensordata属于非连续存储张量,早期版本PyTorch的.t()方法对非连续张量的兼容存在缺陷,会误触发越界报错。 - 额外隐藏原因:你的索引逻辑中
batch_item_index[i]包含的索引值,超出了self.linear1.weight第二维的最大合法范围(最大合法索引为第二维长度减1),索引时的越界问题会在后续张量操作中暴露出来。
修复方案
- 先确认张量维度,如果维度不符合二维要求,先调整索引逻辑保证输出为二维张量。
- 确认维度正常的前提下,转置前先将张量转为连续存储:
# 替换原有转置代码 transposed_data = tensordata.contiguous().t()
- 也可以使用兼容性更好的
permute方法完成二维转置,不受张量存储格式限制:
transposed_data = tensordata.permute(1, 0)
- 检查
batch_item_index[i]的数值范围,保证所有索引值小于self.linear1.weight.size(1),避免索引越界。
内容的提问来源于stack exchange,提问作者ddd
相关产品推荐
相关产品推荐

