代码之家  ›  专栏  ›  技术社区  ›  Atul Balaji

使用pytorch对给定张量进行矢量化洗牌的方法

  •  0
  • Atul Balaji  · 技术社区  · 6 年前

    我有一个形状为(1,12,2,2)的张量a,如下所示:

      ([[[[1., 3.],
          [9., 11.],
    
         [[ 2.,  4.],
          [10., 12.]],
    
         [[ 5.,  7.],
          [13., 15.]],
    
         [[ 6.,  8.],
          [14., 16.]],
    
         [[17., 19.],
          [25., 27.]],
    
         [[18., 20.],
          [26., 28.]],
    
         [[21., 23.],
          [29., 31.]],
    
         [[22., 24.],
          [30., 32.]],
    
         [[33., 35.],
          [41., 43.]],
    
         [[34., 36.],
          [42., 44.]],
    
         [[37., 39.],
          [45., 47.]],
    
         [[38., 40.],
          [46., 48.]]]])
    

    我想用pytorch将其洗牌,以产生以下形状(1,3,4,4)的张量B:

    tensor([[[[ 1.,  6.,  3.,  8.],
              [21., 34., 23., 36.],
              [ 9., 14., 11., 16.],
              [29., 42., 31., 44.]],
    
             [[ 2., 17.,  4., 19.],
              [22., 37., 24., 39.],
              [10., 25., 12., 27.],
              [30., 45., 32., 47.]],
    
             [[ 5., 18.,  7., 20.],
              [33., 38., 35., 40.],
              [13., 26., 15., 28.],
              [41., 46., 43., 48.]]]])
    

    我使用两个for循环实现了这一点,如下所示:

    B = torch.zeros(1,3,4,4, dtype=torch.float)
    ctr = 0
    for i in range(2):
        for j in range(2):
            B[:,:,i:4:2,j:4:2] = A[:,ctr:ctr+3,:,:]
            ctr = ctr+3
    

    我正在寻找任何方法,在pytorch中以矢量化的方式实现这一点,而不需要这些for循环。也许使用像这样的函数 .permute()

    0 回复  |  直到 6 年前
        1
  •  4
  •   ddoGas    6 年前

    这样就行了

    B = A.reshape(2,2,3,2,2).permute(2,3,0,4,1).reshape(1,3,4,4)
    
        2
  •  1
  •   Mohit Lamba    6 年前

    只需将上述解决方案推广到任何上采样因子'r'中,如像素洗牌。

    B = A.reshape(-1,r,3,s,s).permute(2,3,0,4,1).reshape(1,3,rs,rs)
    

    对于手上的问题,s=2,r=2,解如下

    B = A.reshape(-1,2,3,2,2).permute(2,3,0,4,1).reshape(1,3,4,4)
    

    由@ddoGas发布

    类似地,如果“A”的大小为(1192356532),并且希望通过r=8do增加采样

    B = A.reshape(-1,8,3,356,532).permute(2,3,0,4,1).reshape(1,3,2848,4256)