【问题标题】:Multiple unmatched matrices in backpropagation through time随时间反向传播的多个不匹配矩阵
【发布时间】:2020-04-19 04:06:30
【问题描述】:

我将通过循环神经网络 (RNN) 实现二进制加法作为示例。我已经解决了一个通过 Python 实现它的问题,所以我决定在那里分享我的问题以提出解决它的想法。

在我的notebook code(反向传播时间 (BPTT) 部分)中可以看到, 有一个像下面这样的链式规则来更新输入权重矩阵,如下所示:

我的问题是这部分:

我尝试在我的Python code 或notebook code(class input_layer,backward 方法)中实现这部分,但不匹配的尺寸会引发错误。

在我的示例代码中,W_hidden 是 16*16,而 delta pre_hidden 的结果是 1*2。这会导致错误。如果你运行代码,你会看到错误。

我花了很多时间来检查我的链式规则以及我的代码。我想我的连锁规则是正确的。出现此错误的唯一原因是我的代码。

据我所知,就维度而言,多个不匹配的矩阵是不可能的。如果我的链式规则是正确的,那么它如何由 Python 实现? 有什么想法吗?

提前致谢。

【问题讨论】:

    标签: python deep-learning matrix-multiplication recurrent-neural-network back-propagation-through-time


    【解决方案1】:

    您需要在渐变上应用维度平衡。取自 Stanford's cs231n 课程,归结为两个简单的修改:

    鉴于 和,我们将有:

    ,

    这是我用来确保梯度计算正确的代码。您应该能够相应地更新您的代码。

    import torch
    
    torch.random.manual_seed(0)
    
    x_1, x_2 = torch.zeros(size=(1, 8)).normal_(0, 0.01), torch.zeros(size=(1, 8)).normal_(0, 0.01)
    y = torch.zeros(size=(1, 8)).normal_(0, 0.01)
    
    h_0 = torch.zeros(size=(1, 16)).normal_(0, 0.01)
    weight_ih = torch.zeros(size=(8, 16)).normal_(mean=0, std=0.01).requires_grad_(True)
    weight_hh = torch.zeros(size=(16, 16)).normal_(mean=0, std=0.01).requires_grad_(True)
    weight_ho = torch.zeros(size=(16, 8)).normal_(mean=0, std=0.01).requires_grad_(True)
    
    h_1 = x_1.mm(weight_ih) + h_0.mm(weight_hh)
    h_2 = x_2.mm(weight_ih) + h_1.mm(weight_hh)
    g_2 = h_2.sigmoid()
    j_2 = g_2.mm(weight_ho)
    y_predicted = j_2.sigmoid()
    
    loss = 0.5 * (y - y_predicted).pow(2).sum()
    
    loss.backward()
    
    
    delta_1 = -1 * (y - y_predicted) * y_predicted * (1 - y_predicted)
    delta_2 = delta_1.mm(weight_ho.t()) * (g_2 * (1 - g_2))
    delta_3 = delta_2.mm(weight_hh.t())
    
    # 16 x 8
    weight_ho_grad = g_2.t() * delta_1
    
    # 16 x 16
    weight_hh_grad = h_1.t() * delta_2 + (h_0.t() * delta_3)
    
    # 8 x 16
    weight_ih_grad = x_2.t() * delta_2 + x_1.t() * delta_3
    
    atol = 1e-10
    assert torch.allclose(weight_ho.grad, weight_ho_grad, atol=atol)
    assert torch.allclose(weight_hh.grad, weight_hh_grad, atol=atol)
    assert torch.allclose(weight_ih.grad, weight_ih_grad, atol=atol)
    

    【讨论】:

      猜你喜欢
      • 2018-05-30
      • 1970-01-01
      • 2011-08-25
      • 2019-05-23
      • 2017-01-20
      • 2019-12-10
      • 2021-07-27
      • 1970-01-01
      • 2020-11-23
      相关资源
      最近更新 更多