如何仅用Torch将3×3张量转换为3×1×3张量?
3×3 PyTorch张量转3×1×3的实现方法
给定原始3×3张量:
tensor([[1., 1., 1.], [1., 1., 1.], [1., 1., 1.]])
需要转换为3×1×3的目标张量:
tensor([[[1., 1., 1.]], [[1., 1., 1.]], [[1., 1., 1.]]])
你提到不能用[:, None, :]索引方式,且unsqueeze(-1)会得到不符合需求的3×3×1张量,这里提供几种仅用PyTorch的可行方案:
方案1:指定维度的unsqueeze
unsqueeze支持传入具体维度索引,我们需要在第1个维度(从0开始计数)插入新维度,代码如下:import torch x = torch.ones(3, 3) x_reshaped = x.unsqueeze(1)执行后
x_reshaped的形状即为(3, 1, 3),完全匹配需求。方案2:使用reshape直接定义形状
利用reshape方法直接指定目标张量形状,只要元素总数和原始张量一致即可:x_reshaped = x.reshape(3, 1, 3)方案3:使用view方法
如果原始张量是连续存储的(比如torch.ones生成的张量),可以用view方法实现形状转换:x_reshaped = x.view(3, 1, 3)
注:unsqueeze(-1)不符合需求是因为-1代表最后一个维度,插入后会得到(3, 3, 1)的形状,而我们需要在中间维度插入,所以指定索引1即可。
内容的提问来源于stack exchange,提问作者Ant
相关产品推荐
相关产品推荐

