Advertisement

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)

还没有任何评论哟~