【问题标题】:PyTorch torch.no_grad() versus requires_grad=FalsePyTorch torch.no_grad() 与 requires_grad=False
【发布时间】:2020-12-26 07:58:55
【问题描述】:

我正在关注 PyTorch tutorial,它使用 Huggingface Transformers 库中的 BERT NLP 模型(特征提取器)。有两段相互关联的梯度更新代码我看不懂。

(1)torch.no_grad()

本教程有一个类,其中 forward() 函数围绕对 BERT 特征提取器的调用创建一个 torch.no_grad() 块,如下所示:

bert = BertModel.from_pretrained('bert-base-uncased')

class BERTGRUSentiment(nn.Module):
    
    def __init__(self, bert):
        super().__init__()
        self.bert = bert
        
    def forward(self, text):
        with torch.no_grad():
            embedded = self.bert(text)[0]

(2)param.requires_grad = False

在同一教程中的另一部分,BERT 参数被冻结。

for name, param in model.named_parameters():                
    if name.startswith('bert'):
        param.requires_grad = False

我什么时候需要 (1) 和/或 (2)?

  • 如果我想使用冻结的 BERT 进行训练,是否需要同时启用两者?
  • 如果我想训练以更新 BERT,是否需要同时禁用两者?

另外,我跑了所有四个组合,发现:

   with torch.no_grad   requires_grad = False  Parameters  Ran
   ------------------   ---------------------  ----------  ---
a. Yes                  Yes                      3M        Successfully
b. Yes                  No                     112M        Successfully
c. No                   Yes                      3M        Successfully
d. No                   No                     112M        CUDA out of memory

有人能解释一下发生了什么吗? 为什么我收到 CUDA out of memory 表示 (d) 而不是 (b)?两者都有 112M 的可学习参数。

【问题讨论】:

    标签: python machine-learning pytorch bert-language-model huggingface-transformers


    【解决方案1】:

    这是一个较早的讨论,多年来略有变化(主要是由于 with torch.no_grad() 作为模式的目的。可以在 on Stackoverflow already 找到一个很好的答案,也可以回答您的问题。
    但是,由于最初的问题有很大不同,我将避免将其标记为重复,尤其是因为第二部分是关于记忆的。

    no_grad的初步解释给出here

    with torch.no_grad() 是一个上下文管理器,用于防止计算梯度 [...]。

    另一方面使用requires_grad

    冻结部分模型并训练其余部分 [...]。

    再次来源the SO post

    基本上,使用requires_grad,您只是禁用网络的一部分,而no_grad 根本不会存储任何梯度,因为您可能将其用于推理而不是训练。
    为了分析您的参数组合的行为,让我们调查正在发生的事情:

    • a)b) 根本不存储任何渐变,这意味着无论参数数量如何,您都可以使用更多内存,因为您不会保留它们以用于潜在的向后传递。
    • c) 必须存储前向传递以供以后反向传播,但是,只存储了有限数量的参数(300 万),这使得这仍然可以管理。
    • 但是,d) 需要存储所有 1.12 亿个参数的正向传递,这会导致内存不足。

    【讨论】:

    • 谢谢。什么时候会使用(c)? (即收集梯度但冻结参数)
    • 只要不冻结网络的所有参数,就可以只训练特定的层。例如,如果您有一个非常大(但已经预训练)的嵌入层,您可以通过简单地冻结嵌入层来实现更快的训练时间,同时可能不会牺牲太多的准确性。
    猜你喜欢
    • 2019-01-15
    • 2021-12-01
    • 1970-01-01
    • 2020-07-03
    • 2019-09-01
    • 2021-05-01
    • 2019-04-26
    • 2021-07-04
    • 1970-01-01
    相关资源
    最近更新 更多