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

PyTorch中求损失对学习率的梯度:自定义优化器问题

问题描述

我正在构建一个从Dirichlet分布采样学习率的自定义优化器,其中参数alpha需要在每次反向传播时更新。已经算出了∂η/∂α(η为学习率),但需要获取损失对学习率的梯度∂L/∂η,通过链式法则∂L/∂α=∂L/∂η*∂η/∂α来更新alpha,优化分布采样效果。

尝试用以下代码获取∂L/∂η时出现错误:

grad_learning_rate = torch.autograd.grad(loss, self.learning_rate, grad_outputs=torch.tensor(1.0, device=loss.device), retain_graph=True, allow_unused=True)[0]

错误信息:

One of the differentiated Tensors appears to not have been used in the graph. Set allow_unused=True if this is the desired behavior.

相关代码

模型

class MLP(nn.Module):
    def __init__(self, input_size, output_size, device: torch.device=None):
        super(MLP, self).__init__()
        self.fc1 = nn.Linear(input_size, 10, dtype=torch.float64)
        self.relu = nn.ReLU()

        # 根据输入设备或CPU部署模型
        self.device = device if device is not None else torch.device('cpu')
        self.to(self.device)

    def forward(self, x):
        x = self.fc1(x)
        return x

自定义优化器

class Dart(Optimizer):
''' 
训练过程中需要向优化器传入损失值
''' 
def __init__(self, params, betas=(0.9, 0.999),
             alpha_init=1.0, alpha_lr=0.0001, eps=1e-8, weight_decay=0): 
    defaults = dict(betas=betas, eps=eps, weight_decay=weight_decay)
    super(Dart, self).__init__(params, defaults)
    self.alpha_scaler = alpha_init
    self.alpha_lr = alpha_lr
    self.learning_rate = None
    self.alpha_grads = None
    
def sample_lr_candidates(self, mean=1e-3, std=1e-4, num_samples=(10, 1), min_lr=1e-6, max_lr=1e-1):
    # 从高斯分布采样学习率候选值
    lr_samples = torch.normal(mean=mean, std=std, size=(num_samples))
    
    # 裁剪到指定范围
    lr_samples = torch.clamp(lr_samples, min=min_lr, max=max_lr)
    
    return lr_samples.to(torch.float64)

def step(self, loss):
    for group in self.param_groups:  # 仅一个参数组
        for p in group['params']:
            if p.grad is None:
                continue

            dim = (10, 784) if p.shape == torch.Size([10, 784]) else (1, 10)
            
            state = self.state[p]  # 优化器为每个参数维护状态字典
            input = torch.empty(dim, device='cpu', dtype=torch.float64)

    
            if len(state) == 0:  # 初始化参数状态
                state['step'] = 0
                state['lr_candidates'] = self.sample_lr_candidates(num_samples=p.shape)
                state['alphas'] = torch.ones_like(input, memory_format=torch.preserve_format) * self.alpha_scaler
            
            state['step'] += 1  
            
            # 开启alpha的自动求导
            state['alphas'].requires_grad_(True)
            
            # 从Dirichlet分布采样(可导)
            samples = torch.distributions.Dirichlet(state['alphas']).rsample()
            total = state['alphas'].sum(-1, True).expand_as(state['alphas'])
            grad_samples = torch._dirichlet_grad(samples, state['alphas'], total)  # ∂samples/∂alphas
            
            # 计算学习率
            self.learning_rate = samples * state['lr_candidates']
            self.learning_rate.retain_grad()
            
            # 尝试计算损失对学习率的梯度(报错位置)
            grad_learning_rate = torch.autograd.grad(loss, self.learning_rate, grad_outputs=torch.tensor(1.0, device=loss.device), retain_graph=True)[0]
            
            # 更新alpha(待补充∂L/∂η项)
            state['alphas'] = state['alphas'] - self.alpha_lr * grad_samples * state['lr_candidates']
            self.alpha_grads = state['alphas']

            # 更新模型参数
            p.data.sub_(self.learning_rate.squeeze() * p.grad)

训练函数

def train(self, model: nn.Module, optim: optim.Optimizer, criterion: Callable[[Tensor, Tensor], Tensor]) -> dict:
    model.train()
    training_history = {}

    for epoch in range(self.epochs):
        losses = []
        learning_rates = []
        alpha_grads = []
        accuracies = []
        for images, labels in self.data:
            images = torch.squeeze(images)
            labels = torch.squeeze(labels)
            
            assert images.shape == torch.Size([128, 784]), f"输入图像维度不符:{images.shape}"
            assert labels.shape == torch.Size([128]), f"标签维度不符:{labels.shape}"
            
            predictions = model(images)
            loss = criterion(predictions, labels)

            optim.zero_grad()  # 重置梯度
            loss.backward(retain_graph=True)  # 计算模型参数梯度
            optim.step(loss)  # 执行参数更新并传入损失

            learning_rates.append(torch.mean(optim.learning_rate.to('cpu')))
            alpha_grads.append(torch.mean(optim.alpha_grads.to('cpu')))
            losses.append(loss.to('cpu'))
        self.store_training_history(history=training_history,
                               epoch_num=epoch,
                               loss=losses,
                               learning_rate=learning_rates,
                               alpha_grads = alpha_grads
                            )
        with torch.no_grad():
            print(f"完成第 {epoch+1}/{self.epochs} 轮训练,平均损失:{np.mean(losses):.4f}")
    return training_history

问题原因与解决方案

错误原因

报错的核心是**self.learning_rate和损失loss之间没有建立计算图依赖**:

  • 当前流程中,loss.backward()先计算了模型参数的梯度,而self.learning_rate是在optim.step()里才生成的,属于反向传播之后的计算,因此它并没有被包含在loss的计算图中。
  • 当调用torch.autograd.grad(loss, self.learning_rate)时,PyTorch找不到两者之间的关联路径,所以报错。

解决思路与修改方案

要让损失能追踪到学习率的依赖,需要让学习率的采样、参数更新过程都嵌入到计算图中,具体修改如下:

1. 调整训练流程,提前生成学习率

在反向传播前就完成学习率的采样,确保它参与到参数更新的计算图中。给优化器新增一个prepare_lr()方法:

class Dart(Optimizer):
    # ... 原有代码省略 ...
    
    def prepare_lr(self):
        for group in self.param_groups:
            for p in group['params']:
                if p.grad is None:
                    continue
                state = self.state[p]
                if len(state) == 0:
                    # 初始化参数状态
                    state['step'] = 0
                    dim = (10, 784) if p.shape == torch.Size([10, 784]) else (1, 10)
                    state['lr_candidates'] = self.sample_lr_candidates(num_samples=p.shape).to(p.device)
                    state['alphas'] = torch.ones(dim, device=p.device, dtype=torch.float64) * self.alpha_scaler
                    state['alphas'].requires_grad_(True)
                
                # 采样学习率并保留计算图
                samples = torch.distributions.Dirichlet(state['alphas']).rsample()
                state['samples'] = samples
                state['lr'] = samples * state['lr_candidates']
                state['lr'].retain_grad()

修改训练函数,提前调用prepare_lr()并重新计算关联损失:

def train(self, model: nn.Module, optim: optim.Optimizer, criterion: Callable[[Tensor, Tensor], Tensor]) -> dict:
    # ... 原有代码省略 ...
    for images, labels in self.data:
        # ... 原有代码省略 ...
        predictions = model(images)
        loss = criterion(predictions, labels)

        optim.zero_grad()
        optim.prepare_lr()  # 提前生成学习率,嵌入计算图
        
        # 用更新后的参数重新计算损失,建立学习率与损失的关联
        with torch.enable_grad():
            updated_params = []
            for p in model.parameters():
                updated_p = p - optim.state[p]['lr'].squeeze() * p.grad
                updated_params.append(updated_p)
            # 手动模拟更新后的前向传播
            new_predictions = F.linear(images, updated_params[0], updated_params[1])
            new_loss = criterion(new_predictions, labels)
        
        new_loss.backward(retain_graph=True)
        optim.step(new_loss)
        # ... 原有代码省略 ...

2. 修正优化器step方法,正确计算梯度

更新step()方法,现在可以正常获取∂L/∂η并完成alpha的链式更新:

def step(self, loss):
    for group in self.param_groups:
        for p in group['params']:
            if p.grad is None:
                continue
            state = self.state[p]
            
            # 获取损失对学习率的梯度
            grad_learning_rate = torch.autograd.grad(loss, state['lr'], grad_outputs=torch.tensor(1.0, device=loss.device), retain_graph=True)[0]
            
            # 计算∂samples/∂alphas
            total = state['alphas'].sum(-1, True).expand_as(state['alphas'])
            grad_samples = torch._dirichlet_grad(state['samples'], state['alphas'], total)
            
            # 链式法则计算∂L/∂alphas
            grad_alpha = grad_samples * state['lr_candidates'] * grad_learning_rate
            
            # 更新alpha(用in-place操作避免断开计算图)
            state['alphas'].data.sub_(self.alpha_lr * grad_alpha)
            
            # 更新模型参数
            p.data.sub_(state['lr'].squeeze() * p.grad)
            
            # 保存当前学习率和alpha梯度
            self.learning_rate = state['lr']
            self.alpha_grads = grad_alpha

3. 关键注意事项

  • 确保state['alphas']始终开启requires_grad=True,更新时使用data.sub_()这类in-place操作,避免断开计算图。
  • 必须让学习率的生成和参数更新过程都被包含在损失的计算图中,否则无法追踪梯度依赖。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.13 22:04:50