如何在PyTorch中用make_functional仅计算部分参数的梯度?
优化PyTorch自动求导:仅计算指定参数的雅可比矩阵
问题背景
原本的代码通过make_functional结合vmap和jacrev计算模型全参数的雅可比矩阵(形状为数据数量×参数数量),但实际需求仅需计算随机选中的部分目标参数的梯度,全量计算会造成算力浪费,需要优化自动求导步骤。
原全参数雅可比计算代码:
func_model, func_params = make_functional(self.model) def fm(x, func_params): fx = func_model(func_params, x) return fx.squeeze(0).squeeze(0) def floss(func_params, input): fx = fm(input, func_params) return fx per_sample_grads = vmap(jacrev(floss), (None, 0))(func_params, input) cnt=0 for g in per_sample_grads: g = g.detach() J_d = g.reshape(len(g),-1) if cnt == 0 else torch.hstack([J_d,g.reshape(len(g),-1)]) cnt = 1 result = J_d.detach()
目标参数选择逻辑:
params = torch.cat([p.view(-1) for p in self.model.parameters()], dim=0) selected_columns = torch.random.choice(p_number, opt_num, replace=False) target_params = params[selected_columns] # 仅需计算这些参数的梯度
优化方案
核心思路是只对选中的参数求导,将未选中的参数视为固定常量,避免不必要的梯度计算。以下是两种可行的实现方式:
方式一:基于展平参数的简化实现
直接在展平后的参数张量上操作,代码简洁易维护:
# 先记录原始参数的形状信息,用于后续恢复结构 param_shapes = [p.shape for p in self.model.parameters()] total_params = sum(p.numel() for p in self.model.parameters()) # 将func_params展平为单一张量 full_params_flat = torch.cat([p.view(-1) for p in func_params]) def floss_selected(selected_params, fixed_flat_params, selected_idxs, input): # 将选中参数放回原始展平张量的对应位置 updated_flat = fixed_flat_params.clone() updated_flat[selected_idxs] = selected_params # 恢复为原始参数结构 full_params = [] ptr = 0 for shape in param_shapes: num_elem = shape.numel() full_params.append(updated_flat[ptr:ptr+num_elem].reshape(shape)) ptr += num_elem # 计算模型输出 fx = func_model(full_params, input).squeeze(0).squeeze(0) return fx # 仅对selected_params求导,用vmap批量处理每个样本 per_sample_grads = vmap(jacrev(floss_selected, argnums=0), (None, None, None, 0))(target_params, full_params_flat, selected_columns, input) # 最终结果形状:[样本数, 选中参数数] result = per_sample_grads.detach()
方式二:保留原始参数结构的高效实现
如果需要保留原始参数的分片结构(避免大张量操作),可以用以下方式:
# 先记录每个参数张量的长度和累计偏移量,建立索引映射 param_offsets = [0] total_params = 0 param_shapes = [] for p in self.model.parameters(): num_elem = p.numel() param_shapes.append(p.shape) param_offsets.append(total_params + num_elem) total_params += num_elem # 拆分固定参数和选中参数 fixed_params = [] selected_params_list = [] for param_idx, p in enumerate(func_params): # 找到当前参数张量中被选中的局部索引 local_selected = [idx - param_offsets[param_idx] for idx in selected_columns if param_offsets[param_idx] <= idx < param_offsets[param_idx+1]] if not local_selected: # 无选中元素,全部作为固定参数 fixed_params.append(p) continue # 提取选中元素作为可导参数 local_selected_tensor = torch.tensor(local_selected, device=p.device) selected_p = p.take(local_selected_tensor) selected_params_list.append(selected_p) # 提取未选中元素作为固定参数 mask = torch.ones(p.numel(), dtype=torch.bool, device=p.device) mask[local_selected_tensor] = False fixed_p = p.take(mask).reshape(p.shape) if mask.any() else torch.empty_like(p) fixed_params.append(fixed_p) # 合并选中参数为单一张量(也可保持列表形式) selected_params = torch.cat(selected_params_list) def reconstruct_full_params(selected_params, fixed_params): """将选中参数和固定参数合并回原始参数结构""" full_params = [] selected_ptr = 0 for param_idx, (fixed_p, shape) in enumerate(zip(fixed_params, param_shapes)): num_selected = len([idx for idx in selected_columns if param_offsets[param_idx] <= idx < param_offsets[param_idx+1]]) if num_selected == 0: full_params.append(fixed_p) continue # 取出当前参数对应的选中片段 selected_segment = selected_params[selected_ptr:selected_ptr+num_selected] selected_ptr += num_selected # 重建完整参数张量 full_p = torch.zeros(shape, device=selected_params.device) local_selected = [idx - param_offsets[param_idx] for idx in selected_columns if param_offsets[param_idx] <= idx < param_offsets[param_idx+1]] # 填充固定参数 mask = torch.ones(shape.numel(), dtype=torch.bool, device=selected_params.device) mask[local_selected] = False full_p.view(-1)[mask] = fixed_p.view(-1) # 填充选中参数 full_p.view(-1)[local_selected] = selected_segment full_params.append(full_p) return full_params def fm(x, full_params): fx = func_model(full_params, x) return fx.squeeze(0).squeeze(0) def floss(selected_params, fixed_params, input): full_params = reconstruct_full_params(selected_params, fixed_params) return fm(input, full_params) # 仅对选中参数求导,批量处理样本 per_sample_grads = vmap(jacrev(floss, argnums=0), (None, None, 0))(selected_params, fixed_params, input) # 最终结果形状:[样本数, 选中参数数] result = per_sample_grads.detach()
关键优化点说明
- 缩小求导范围:通过将未选中参数设为常量,
jacrev仅计算目标参数的梯度,避免全量参数的无效计算,大幅节省算力。 - 参数映射正确性:通过索引映射确保选中参数能准确放回原始结构,保证模型输出计算的正确性。
- 保持批量效率:保留
vmap的批量处理逻辑,避免单样本循环的低效问题。
内容的提问来源于stack exchange,提问作者Klae zhou
相关产品推荐
相关产品推荐

