Python张量reshape报错,如何将[800,1,64,32,32]转为[1,64,800,32,32]
报错原因
- 框架不兼容:你操作的是PyTorch张量,却调用了TensorFlow的
tf.reshape接口,两个框架的张量数据结构不互通,直接调用必然报错。 - 数据类型错误:此时
target['a']还只是存储PyTorch张量的Python列表,不是连续内存存储的张量对象,无法直接传入reshape类接口处理。 - 操作逻辑错误:即使你把列表转成形状为
(800, 1, 64, 32, 32)的张量,直接用reshape也无法得到你要的结果,reshape是按内存存储顺序重排张量形状,不会调整维度的先后顺序,你需要的是调整维度位置的换位操作,不是重排形状。
正确实现步骤
- 先将存储张量的Python列表堆叠为完整的PyTorch张量,
torch.stack默认会在新增的第0维拼接,得到形状为(800, 1, 64, 32, 32)的张量:
import torch # 处理a对应的张量 a_tensor = torch.stack(target['a']) # 如果要处理b、c,对应修改键名即可: # b_tensor = torch.stack(target['b']) # c_tensor = torch.stack(target['c'])
- 使用
permute接口调整维度顺序,将原本第0位的800维度移动到第2位,得到目标形状:
a_target = a_tensor.permute(1, 2, 0, 3, 4)
执行后a_target.shape就是你需要的torch.Size([1, 64, 800, 32, 32])。
内容的提问来源于stack exchange,提问作者Ghost
相关产品推荐
相关产品推荐

