PyTorch中如何创建形状为(k, x.shape)的新张量?
解决PyTorch中创建形状为(k, x.shape)张量的问题
你需要基于现有张量x创建一个开头新增k维度的张量y,直接使用torch.empty((k, x.shape))会报错,因为x.shape是torch.Size类型,无法和整数k直接组成合法的形状参数。这里提供几种通用的实现方法,无需预先知晓x的维度数量:
方法一:元组拼接
利用torch.Size是元组子类的特性,直接和包含k的元组拼接:y = torch.empty((k,) + x.shape)方法二:使用
torch.Size的insert方法
在原形状的最前面插入k维度,生成新的形状对象:y = torch.empty(torch.Size(x.shape).insert(0, k))方法三:列表转换插入
将原形状转为列表后插入k,再传入torch.empty:shape_list = list(x.shape) shape_list.insert(0, k) y = torch.empty(shape_list)
以上三种方法都能自动适配x的任意维度数,替代需要手动罗列x.shape[0]、x.shape[1]的繁琐写法。
内容的提问来源于stack exchange,提问作者0xbadf00d
相关产品推荐
相关产品推荐

