如何在指定维度拼接张量列表?PyTorch通道维度张量拼接实现(避免新增维度)
如何用torch.cat在通道维度拼接张量列表
没问题,这事儿用torch.cat()完全能轻松实现,核心就是找对要拼接的维度就行!
关键思路
你的每个张量形状是[batch, channel, width, height],对应索引分别是0、1、2、3。你要在通道维度合并,也就是指定dim=1作为torch.cat()的参数——因为torch.cat()是在已有的维度上进行拼接,不会像torch.stack()那样新增维度,刚好符合你的需求。
具体代码示例
直接看实操代码:
import torch # 假设你的my_list是包含多个[1,3,128,128]张量的列表 my_list = [torch.randn(1, 3, 128, 128) for _ in range(4)] # 这里用4个张量举例 # 在通道维度(dim=1)拼接 new_tensor = torch.cat(my_list, dim=1) # 验证形状 print(new_tensor.shape) # 输出: torch.Size([1, 12, 128, 128]),正好是3*4=12个通道
为什么不用torch.stack()?
顺便提一句,torch.stack()会给结果新增一个维度,比如用它处理你的列表会得到形状[4,1,3,128,128],这显然不是你想要的合并通道的效果,所以选torch.cat()就对了。
内容的提问来源于stack exchange,提问作者Ammar Ul Hassan
相关产品推荐
相关产品推荐

