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.