PyTorch报错RuntimeError:零维张量无法拼接的解决方法
问题解决:PyTorch拼接0维张量报错
错误原因
losses_q是0维标量张量(shape为()),而torch.cat要求所有待拼接的张量必须维度一致。this_loss_q是1维张量(shape为[1]),两者维度不匹配,因此触发RuntimeError。
解决方法
只需把0维的losses_q转换成1维张量,再执行拼接即可,常用两种实现方式:
方式1:用unsqueeze(0)扩展维度
unsqueeze(0)会在第0维位置新增一个维度,将标量转为shape为[1]的1维张量:
# 先扩展losses_q的维度 losses_q = losses_q.unsqueeze(0) # 执行拼接操作 losses_q = torch.cat((losses_q, this_loss_q), dim=0)
方式2:用[None]索引扩展维度
这是PyTorch的语法糖,效果和unsqueeze(0)完全相同,写法更简洁:
losses_q = losses_q[None] losses_q = torch.cat((losses_q, this_loss_q), dim=0)
验证结果
拼接后得到shape为[2]的张量,示例输出:
tensor([0.0870, 0.0874], device='cuda:0', grad_fn=<CatBackward0>)
注:两种方式都不会破坏原张量的计算图,梯度可正常反向传播,适配训练场景。
内容的提问来源于stack exchange,提问作者S.EB
相关产品推荐
相关产品推荐

