我最终得到了以下解决方案。它由以下几部分组成:
- 将每个展平坐标映射为 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_ind 和param_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)等