PyTorch - efficiently apply TopK gradient coordinates in NN

Viewed 359

I have a sequential neural network (standard ResNet model), and for a constant k (about 1000, but may potentially increase in the future) I want to do the following:

  1. Find the gradient of NN.
  2. Identify Top-kcoordinates of the gradient (k coordinates with the largest absolute value)
  3. Apply the gradient descent step using only these coordinates (i.e. other gradient coordinates are 0)

What I can do is the following:

  • Flatten the vector:
param_grads = [param.grad.to(cpu).flatten() for param in model.parameters()]
grad = torch.cat(param_grads) 
  • Identify indices of Top-k coordinates in the sorted vector (I can also use topk function):
sorted_grad = grad.abs().sort()[1]

Now, the question is how to apply only these coordinates. I can write a function that manually translates flatten vector coordinates into the original coordinates (including the corresponding parameter), make a slice for each parameter, and zero gradient outside this slice for each parameter. However, I suspect that it'll be really inefficient. What's the best way to achieve this?

1 Answers

I ended up with the following solution. It consists of the following parts:

  • Map each flattened coordinate into 1) its original parameter 2) its original coordinate in the parameter)
  • Collect top-k coordinates of the flattened vector
  • Perform the update efficiently

The following function builds maps from the first item

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

I use that coordinates in the flattened vector will be in the same order as if I iterate over all indices of all parameters. I've just hard-coded some shapes I encounter in the model.

The next part flattens the vector:

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

Then I find Top-k coordinates

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

and zero the current parameters' gradients:

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

The next part is to use coordinates from the flattened vector. While it's possible to just iterate over top-k coordinates and assign their values one-by-one, this was pretty slow (I guess that it's inefficient to communicate to GPU one coordinate at a time). The following solution was 10 times faster for me. The idea is to accumulate all updates for all parameters and apply them simultaneously. I create the following maps which, for each parameter, store coordinates where updates are performed and the corresponding update values:

  param_to_ind, param_to_vals = {}, {}

In the end, I'll just perform these updates:

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

It remains to fill these maps. The code is the following:

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)

where param_to_ind and param_to_vals are maps from before, param is the model parameter, index is the update coordinate and val is the update value. It looks a bit nasty since PyTorch expects the following format for slice (i.e. a[ind] = x): if the tensor (i.e. a) has shape s, then slice (i.e. ind) should be an s-tuple, where each item is a list: the first list contains first coordinates of update coordinates (i.e. c1_1, c2_1, ..., ck_1), second list contains second coordinates (i.e. c1_2, c2_2, ..., ck_2), etc.

Related