8-pytorch中常见的合并和分割
发布时间
阅读量:
阅读量
import torch
'''
1-数据合并:cat(沿着维度);stack(增加维度)
2-数据分割:split(按照长度);chunk(按照数量)
'''
a1 = torch.rand(2, 3, 28, 28)
a2 = torch.rand(4, 3, 28, 28)
'''
torch.cat([a1, a2], dim=0)
dim = 指定拼接的维度
注意,使用torch.cat进行拼接时除了拼接维度可以不同外,其他的维度必须相同
'''
def zqb_cat():
print(a1.shape, a2.shape) # torch.Size([2, 3, 28, 28]) torch.Size([4, 3, 28, 28])
a3 = torch.cat([a1, a2], dim=0)
print(a3.shape) # torch.Size([6, 3, 28, 28])
# a4 = torch.cat([a1, a2],dim=1) #RuntimeError: Sizes of tensors must match except in dimension
全部评论 (0)
还没有任何评论哟~
