【问题标题】:Nested optimization in pytorchpytorch中的嵌套优化
【发布时间】:2021-12-03 22:32:05
【问题描述】:

我写了一个简短的 sn-p 来训练一个分类模型,并学习其优化算法的学习率。在我的示例中,我尝试在内部优化循环中更新网络的权重,并使用外部优化循环(元优化)来学习权重更新的学习率。我收到了错误:

RuntimeError:梯度计算所需的变量之一已被就地操作修改:[torch.FloatTensor [3, 10]],即 AsStridedBackward0 的输出 0,版本为 12;而是预期的版本 2。提示:使用 torch.autograd.set_detect_anomaly(True) 启用异常检测以查找未能计算其梯度的操作。

我的代码 sn-p 如下(注意:我使用的是_statelessnn 的实验性功能 API。您需要使用每晚构建的 pytorch 运行。)

import torch
from torch import nn, optim
from torch.utils.data import Dataset, DataLoader

from torch.nn.utils import _stateless


class MyDataset(Dataset):
    def __init__(self, N):
        self.N = N
        self.x = torch.rand(self.N, 10)
        self.y = torch.randint(0, 3, (self.N,))

    def __len__(self):
        return self.N

    def __getitem__(self, idx):
        return self.x[idx], self.y[idx]


class MyModel(nn.Module):
    def __init__(self):
        super(MyModel, self).__init__()
        self.fc1 = nn.Linear(10, 10)
        self.fc2 = nn.Linear(10, 3)

        self.relu = nn.ReLU()

        self.alpha = nn.Parameter(torch.randn(1))
        self.beta = nn.Parameter(torch.randn(1))

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

epochs = 20
N = 100
dataset = DataLoader(dataset=MyDataset(N), batch_size=10)
model = MyModel()
loss_func = nn.CrossEntropyLoss()

optim = optim.Adam([model.alpha], lr=1e-3)

params = dict(model.named_parameters())
for i in range(epochs):
    model.train()
    train_loss = 0
    for batch_idx, (x, y) in enumerate(dataset):
        logits = _stateless.functional_call(model, params, x)             # predict
        loss_inner = loss_func(logits, y)                                 # loss
        optim.zero_grad()                                                 # reset grad
        loss_inner.backward(create_graph=True, inputs=params.values())    # compute grad
        train_loss += loss_inner.item()                                   # store loss
        for k, p in params.items():
            if k is not 'alpha' and k is not 'beta':
                p.update = - model.alpha * p.grad
                params[k] = p + p.update                      # update weight

    print('Train Epoch: {}\tLoss: {:.6f}'.format(i, train_loss / N))
    logits = _stateless.functional_call(model, params, x)                 # predict
    loss_meta = loss_func(logits, y)
    loss_meta.backward()
    loss_meta.step()

从错误消息中,我了解到问题来自网络第二层权重的权重更新,这表明我的内部循环优化存在错误。任何建议将不胜感激。

【问题讨论】:

标签: pytorch


【解决方案1】:

检查此链接并在每个时期保存 PARAM 并使用相同的内部批次: https://discuss.pytorch.org/t/issue-using-parameters-internal-method/134549/11

for i in range(epochs):
    model.train()
    train_loss = 0
    params = dict(model.named_parameters()) # add this
    for batch_idx, (x, y) in enumerate(dataset):
        params = {k: v.clone() for k,v in params.items()} # add this
        logits = _stateless.functional_call(model, params, x)             # predict
        loss_inner = loss_func(logits, y)
        ..................                     

【讨论】:

    【解决方案2】:

    您应该更新 params[k].data 而不是 params[k]

    (为了避免分心,删掉了例子)

    让我进入一种基本的讨论(不是你的问题的答案)。

    如果我理解正确,您想计算 loss(f(w[i], x)) 并计算 w[i+1,j] = w[i,j] + g(v[j], w[i,j].grad(w.r.t loss)) 。那么最后你要计算v[j+1] = v[j] + v[j].grad(w.r.t loss)

    v[j] 的梯度是使用反向传播计算的,作为 grad w[i,j] 的函数。所以你要做的是选择v[j],这会产生一个好的w[i,j]。我会问:如果你可以直接控制w[i,j],你为什么还要关心v[j]?这就是标准方法。

    【讨论】:

    • 这是不正确的,因为它根本没有更新内部优化的学习率。您可以在第 66 行之后使用make_dot(logits, params=dict(list(model.named_parameters()))).render('comp_graph_beta', format='png') 绘制计算图,并查看图中不存在model.alpha。 (使用from torchviz import make_dot
    • 是的,您定义了 alpha,但您没有在模型或损失函数中使用。你期望什么,alpha 应该乘以学习率?
    • 在这种情况下,您必须在每次迭代中创建一组新参数,并且外部优化将看到非常深的神经网络。容易出现梯度爆炸/消失,我不确定梯度下降法是优化一维参数空间的好方法。
    • 我确实在我的模型中使用了它,它是内部优化算法的学习率。请注意,我在内循环中的模型是MyModel,在外循环中它是MyModel + 内循环更新。因此,在内循环中引入的任何参数都将在外循环中进行优化。如果我按照您的建议更新params[k].data,则永远不会将model.alpha 引入模型。这样想:当你做loss_meta.backward()时,外部优化模型应该看到model.alpha作为参数,所以它实际上应该更新params[k]
    • 确实在这种情况下它是一维参数空间,但这只是一个玩具示例。但是,为什么你认为它容易出现梯度爆炸/消失?
    猜你喜欢
    • 2016-04-20
    • 1970-01-01
    • 2017-03-03
    • 1970-01-01
    • 1970-01-01
    • 1970-01-01
    • 1970-01-01
    • 2020-08-31
    相关资源
    最近更新 更多