如何优化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
相关产品推荐
相关产品推荐

