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

如何优化CIFAR10子集参数梯度协方差计算中的Sherman-Morrison更新效率?

问题描述

需要计算CIFAR10数据集子集样本上参数梯度的协方差矩阵,当前实现代码如下,但效率极低,且因内存限制无法一次性计算所有样本梯度,采用了Sherman-Morrison方法仍未解决效率问题,寻求优化方案:

from torch.func import functional_call, vmap, grad

model1 = LogisticModel().to(device)

def loss_fn(predictions, targets):
  loss = nn.CrossEntropyLoss()
  return loss(predictions, targets)

def compute_loss(params, buffers, sample, target):
  batch = sample.unsqueeze(0)
  targets = target.unsqueeze(0)

  predictions = functional_call(model1, (params, buffers), (batch,))
  loss = loss_fn(predictions, targets)
  return loss

ft_compute_grad = grad(compute_loss)
ft_compute_sample_grad = vmap(ft_compute_grad, in_dims=(None, None, 0, 0))

def sherman_morrison_update(A, u, v):
  vT = v.T
  Au = A @ u

  alpha = 1/(1 + vT@Au)
  A = A - alpha*torch.outer(Au, vT@A)
  return A

testloader1 = DataLoader(test_dataset, batch_size = 512)
params = {k: v.detach() for k, v in model1.named_parameters()}
buffers = {k: v.detach() for k, v in model1.named_buffers()}
w = 0

p_covs = {p:torch.eye(q.flatten().shape[0]).to(device) for p,q in param_grads.items()}
param_grad_mean = {p:torch.zeros(q.flatten().shape[0]).to(device) for p,q in param_grads.items()}

for x,y in tqdm(testloader1):
  param_grads = ft_compute_sample_grad(params, buffers, x.to(device), y.to(device))
  for p,q, mean in zip(param_grads.values(), p_covs, param_grad_mean):
    for p_grad in p:
      w += 1
      diff = p_grad.flatten() - param_grad_mean[mean]
      param_grad_mean[mean] += diff / w
      p_covs[q] = sherman_morrison_update(A=p_covs[q], u=diff, v= diff)
优化方案

1. 优化Sherman-Morrison实现

当前实现存在冗余计算和张量复制开销,可针对你的对称场景(u=v)简化公式,并改为原地更新减少内存消耗:

def sherman_morrison_update_symmetric_inplace(A, u):
    uT = u.t()
    Au = A @ u
    alpha = 1. / (1. + uT @ Au)
    # 原地更新,避免创建新张量
    A.sub_(alpha * torch.outer(Au, uT @ A))
    return A

2. 批量处理梯度更新,取消逐样本循环

逐样本执行Sherman-Morrison是效率瓶颈,改用Sherman-Morrison-Woodbury公式批量处理整个batch的梯度,将循环从样本级提升到batch级:

def sherman_morrison_woodbury_batch(A, U):
    # U是形状为[batch_size, dim]的梯度差异矩阵
    UT = U.t()
    I_plus_UTAU = torch.eye(U.shape[0], device=U.device) + UT @ A @ U
    inv_term = torch.linalg.inv(I_plus_UTAU)
    A.sub_(A @ U @ inv_term @ UT @ A)
    return A

配套的batch级更新逻辑:

for x,y in tqdm(testloader1):
    x, y = x.to(device), y.to(device)
    param_grads = ft_compute_sample_grad(params, buffers, x, y)
    batch_size = x.shape[0]
    
    for param_name, grads in param_grads.items():
        # 扁平化当前batch的所有梯度
        grads_flat = grads.flatten(start_dim=1)
        # 批量更新均值
        old_mean = param_grad_mean[param_name].clone()
        param_grad_mean[param_name] = (w * old_mean + grads_flat.sum(dim=0)) / (w + batch_size)
        # 计算当前batch的梯度差异
        diffs = grads_flat - old_mean.unsqueeze(0)
        # 批量更新协方差矩阵
        p_covs[param_name] = sherman_morrison_woodbury_batch(p_covs[param_name], diffs)
    
    w += batch_size

3. 优化梯度计算流程

  • 提前初始化损失函数:避免在loss_fn中重复创建nn.CrossEntropyLoss实例:
    loss_fn = nn.CrossEntropyLoss()
    def compute_loss(params, buffers, sample, target):
        predictions = functional_call(model1, (params, buffers), (sample.unsqueeze(0),))
        return loss_fn(predictions, target.unsqueeze(0))
    
  • 简化vmap输入逻辑:如果模型支持单样本输入,可直接去掉unsqueeze操作,减少张量处理开销:
    def compute_loss(params, buffers, sample, target):
        predictions = functional_call(model1, (params, buffers), (sample,))
        return loss_fn(predictions, target)
    

4. 内存与计算的权衡优化

  • 低秩近似协方差矩阵:针对大维度参数,用低秩矩阵近似协方差(比如保留SVD分解的前k个奇异值),避免存储完整的方阵。
  • JIT编译加速:将核心计算函数用torch.jit.script编译,提升执行效率:
    @torch.jit.script
    def sherman_morrison_update_symmetric_inplace(A, u):
        uT = u.t()
        Au = A @ u
        alpha = 1. / (1. + uT @ Au)
        A.sub_(alpha * torch.outer(Au, uT @ A))
        return A
    

5. 其他细节优化

  • 给DataLoader设置pin_memory=True,提前将数据加载到固定内存,减少设备转移开销。
  • 预计算每个参数的扁平化维度,避免循环内重复调用flatten().shape[0]。

内容的提问来源于stack exchange,提问作者Pratik Karmakar

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.01 18:02:02