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

PyTorch中nn.Parameter alpha梯度为None的L2C实现问题求助

L2C算法中alpha参数梯度为None的问题排查

问题描述

在实现非meta-L2C算法时,步骤18需要计算损失对nn.Parameter类型参数alpha的梯度,但访问alpha.grad始终返回None,尝试过retain_grad()方法仍无法解决。以下是模型定义与训练循环代码:

模型代码

class CNNCifar(nn.Module):
    def __init__(self):
        super(CNNCifar, self).__init__()
        self.conv1 = nn.Conv2d(3, 6, 5)
        self.pool = nn.MaxPool2d(2, 2)
        self.conv2 = nn.Conv2d(6, 16, 5)
        self.fc1 = nn.Linear(16 * 5 * 5, 120)
        self.fc2 = nn.Linear(120, 84)
        self.fc3 = nn.Linear(84, 10)
        self.alpha = nn.Parameter(torch.randn(100, 100), requires_grad=True)
        self.w = torch.randn((100, 100), requires_grad=True)

    def forward(self, x):
        x = self.pool(F.relu(self.conv1(x)))
        x = self.pool(F.relu(self.conv2(x)))
        x = x.view(-1, 16 * 5 * 5)
        x = F.relu(self.fc1(x))
        x = F.relu(self.fc2(x))
        x = self.fc3(x)
        return F.log_softmax(x, dim=1)

训练循环代码

k = len(neighbour_sets)
device = torch.device("cuda" if not torch.cuda.is_available() else "cpu")
model = CNNCifar().to(device)
criterion = nn.CrossEntropyLoss()
optimizer = optim.SGD(model.parameters(), lr=0.1)
l2c_optimizer = optim.Adam([model.alpha], lr=beta, weight_decay=0.01)

test_accuracies = [[] for _ in range(k)]

theta = [model.state_dict().copy() for _ in range(k)]
theta_half = [model.state_dict().copy() for _ in range(k)]

delta_theta = [model.state_dict().copy() for _ in range(k)]

with tqdm_output(tqdm(range(T))) as trange:
    for t in trange:
        for i in range(k):
            # Local SGD step
            log.info(f'Started training a Local SGD at node {i + 1}')

            model.load_state_dict(theta[i])
            for m in range(S):
                for _, data in enumerate(train_loaders[i]):
                    inputs, labels = data
                    inputs, labels = inputs.to(device), labels.to(device)
                    optimizer.zero_grad()
                    outputs = model(inputs)
                    loss = criterion(outputs, labels)
                    loss.backward()
                    optimizer.step()

            log.info(f'Finished training a Local SGD at node {i + 1}')

            # Change capturing
            log.info(f'Computing change capturing at node {i + 1}')
            for name, param in model.named_parameters():
                delta_theta[i][name] = theta[i][name] - theta_half[i][name]

            log.info(f'Computing mixing weights at node {i + 1}')
            # Mixing weights calculation
            model.w = model.w.clone()
            model.w[i] = compute_mixing_weights(model.alpha[i], neighbour_sets[i])

            # Aggregation
            log.info(f'Aggergating at node {i + 1}')
            theta_next = {}
            for name, param in model.named_parameters():
                theta_next[name] = theta[i][name].clone()

            for j in neighbour_sets[i]:
                for name, param in model.named_parameters():
                    theta_next[name] -= model.w[i][j].item() * delta_theta[i][name][j].clone()

            # Update L2C
            log.info(f'Updating L2C at node {i + 1}')
            model.load_state_dict(theta_next)
            model.train()
            # a training loop to find alpha that minimizes the validation loss
            for _, data in enumerate(val_loaders[i]):
                inputs, labels = data
                inputs, labels = inputs.to(device), labels.to(device)
                
                l2c_optimizer.zero_grad()
                model.alpha.requires_grad_(True)
                
                log.info(f'Forward pass check')
                outputs = model(inputs)
                loss = criterion(outputs, labels)
                model.alpha.retain_grad()
                loss.backward()
                print(f'gradient of alpha is {model.alpha.grad}')
                import pdb; pdb.set_trace()
                l2c_optimizer.step()

            # Remove edges for sparse topology
            if t == T_0:
                for _ in range(K_0):
                    j = min(neighbour_sets[i], key=lambda x: w[i][x])
                    neighbour_sets[i].delete(j)

            theta[i] = model.state_dict().copy()
            theta_half[i] = model.state_dict().copy()

            # Compute test accuracy for each local model
            test_accuracies = compute_test_acc(model, test_loaders[i], device, test_accuracies, i)
        
        log.info(f'Test accuracies atiteration at Comm_round {t} =  {sum(test_accuracies) / k}')

return theta, test_accuracies

问题排查与分析

1. Alpha未参与前向传播计算

PyTorch仅会计算参与前向传播链路的参数的梯度。当前模型的forward函数中,输入x的整个计算流程完全没有用到self.alpha,相当于alpha和最终的损失输出没有任何关联,因此反向传播时不会生成梯度,alpha.grad自然为None。这是核心问题。

L2C算法中,alpha是用来生成混合权重w,进而影响模型参数的聚合,最终影响模型的预测结果与损失。需要将alpha的影响传递到前向传播的损失计算中,比如确保聚合后的模型参数依赖于alpha,或者在forward中直接引入alpha的计算逻辑。

2. W参数的定义与赋值破坏梯度流

  • 模型中self.w被定义为普通torch.Tensor而非nn.Parameter,虽然初始化时设置了requires_grad=True,但不符合PyTorch Module的参数管理规范,容易导致梯度跟踪异常。
  • 训练循环中model.w = model.w.clone()会将w转换为普通Tensor,丢失参数属性;model.w[i] = compute_mixing_weights(...)的直接赋值操作会切断梯度流,导致alpha的梯度无法传递到后续计算。

3. 状态加载覆盖参数梯度属性

在更新L2C前执行model.load_state_dict(theta_next),如果theta_next中保存的alpha是没有梯度跟踪的张量(比如克隆时没有保留requires_grad),会覆盖当前模型的alpha参数,导致后续反向传播无法生成梯度。

4. 冗余的梯度设置操作

  • model.alpha.requires_grad_(True)属于冗余操作:alpha是nn.Parameter,初始化时已经设置requires_grad=True,除非被手动关闭,否则无需重复设置。
  • model.alpha.retain_grad()也无必要:参数类型的张量默认会保留梯度,只有非参数张量(比如中间计算结果)才需要调用该方法。

修复建议

  1. 打通alpha与前向传播的链路
    修改模型逻辑,让alpha参与到最终损失的计算中。比如,确保聚合后的模型参数依赖于alpha,或者在forward中根据alpha调整输出(具体逻辑需符合L2C算法的定义)。

  2. 修正w参数的定义与赋值方式
    将w改为nn.Parameter,并使用保留梯度的方式更新w:

    # 模型初始化中修改
    self.w = nn.Parameter(torch.randn(100, 100), requires_grad=True)
    
    # 训练循环中替换w的赋值逻辑
    # 替代model.w = model.w.clone()和model.w[i] = ...
    new_w_i = compute_mixing_weights(model.alpha[i], neighbour_sets[i])
    model.w.data[i] = new_w_i.data
    # 或使用torch.index_copy_确保梯度不中断
    model.w = nn.Parameter(torch.index_copy(model.w, 0, torch.tensor([i]), new_w_i))
    
  3. 确保状态加载时保留梯度属性
    在克隆state_dict时,保留参数的requires_grad状态,或者在加载后重新启用梯度跟踪:

    # 克隆state_dict时保留梯度属性
    theta[i] = {k: v.clone().requires_grad_(v.requires_grad) for k, v in model.state_dict().items()}
    
  4. 移除冗余的梯度设置代码
    删除model.alpha.requires_grad_(True)和model.alpha.retain_grad()这两行代码。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.15 11:08:10