【问题标题】:PyTorch - efficiently apply TopK gradient coordinates in NNPyTorch - 在 NN 中有效地应用 TopK 梯度坐标
【发布时间】:2021-04-03 15:24:48
【问题描述】:

我有一个顺序神经网络(标准ResNet 模型),对于一个常数k(大约1000,但将来可能会增加)我想做以下事情:

  1. 求 NN 的梯度。
  2. 识别梯度的Top-k坐标(k绝对值最大的坐标)
  3. 仅使用这些坐标应用梯度下降步骤(即其他梯度坐标为 0)

我能做的如下:

  • 扁平化向量:
param_grads = [param.grad.to(cpu).flatten() for param in model.parameters()]
grad = torch.cat(param_grads) 
  • 识别排序向量中Top-k坐标的索引(我也可以使用topk函数):
sorted_grad = grad.abs().sort()[1]

现在,问题是如何仅应用这些坐标。我可以编写一个函数,手动将扁平矢量坐标转换为原始坐标(包括相应的参数),为每个参数制作一个切片,并为每个参数在该切片之外设置零梯度。但是,我怀疑它会非常低效。实现这一目标的最佳方法是什么?

【问题讨论】:

    标签: python neural-network pytorch gradient


    【解决方案1】:

    我最终得到了以下解决方案。它由以下几部分组成:

    • 将每个展平坐标映射为 1) 其原始参数 2) 其在参数中的原始坐标)
    • 收集展平向量的 top-k 坐标
    • 高效执行更新

    以下函数从第一项构建地图

    def build_index_map(model):
      ind_to_ind = []
      ind_to_param = []
      for param in model.parameters():
        if torch.numel(param.data) == 0:
          continue
        shape = param.data.shape
        if len(shape) == 1:
          for i in range(shape[0]):
            ind_to_ind.append((i,))
            ind_to_param.append(param)
        elif len(shape) == 2:
          ...
      return ind_to_ind, ind_to_param
    

    我使用展平向量中的坐标的顺序与我遍历所有参数的所有索引的顺序相同。 我刚刚对模型中遇到的一些形状进行了硬编码。

    下一部分将向量展平:

      param_grads = []
      for param in model.parameters():
        vec = param.grad.flatten()
        if len(vec) == 0:
          continue
        param_grads.append(vec)
    

    然后我找到Top-k坐标

      grad = torch.cat(param_grads)
      topk_abs, topk_coords = torch.topk(grad.abs(), k)
    

    并将当前参数的梯度归零:

      for param in model.parameters():
        param.grad.zero_()
    

    下一部分是使用扁平向量中的坐标。虽然可以只迭代 top-k 坐标并一个一个地分配它们的值,但这非常慢(我想一次与 GPU 通信一个坐标是低效的)。以下解决方案对我来说快了 10 倍。 这个想法是累积所有参数的所有更新并同时应用它们。我创建了以下地图,对于每个参数,存储执行更新的坐标和相应的更新值:

      param_to_ind, param_to_vals = {}, {}
    

    最后,我将执行这些更新:

      for p in param_to_ind:
        p.grad[param_to_ind[p]] = torch.tensor(param_to_vals[p], device=device)
    

    还有待填满这些地图。代码如下:

    def add_update(param_to_ind, param_to_vals, param, index, val):
      coords = param_to_ind.get(param)
      if coords is None:
        param_to_ind[param] = tuple([i] for i in index)
        param_to_vals[param] = [val]
      else:
        for i, ind in enumerate(index):
          coords[i].append(ind)
        param_to_vals[param].append(val)
    

    其中param_to_indparam_to_vals 是之前的映射,param 是模型参数,index 是更新坐标,val 是更新值。 它看起来有点讨厌,因为 PyTorch 期望切片的格式如下(即a[ind] = x):如果张量(即a)的形状为s,那么切片(即ind)应该是s-元组,其中每一项都是一个列表:第一个列表包含更新坐标的第一个坐标(即c1_1, c2_1, ..., ck_1),第二个列表包含第二个坐标(即c1_2, c2_2, ..., ck_2)等

    【讨论】:

      猜你喜欢
      • 2019-04-29
      • 2020-08-28
      • 2021-06-12
      • 1970-01-01
      • 2016-04-13
      • 1970-01-01
      • 1970-01-01
      • 2021-12-31
      • 2018-10-23
      相关资源
      最近更新 更多