You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何在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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.06.18 09:25:05