【问题标题】:early stopping in PyTorch在 PyTorch 中提前停止
【发布时间】:2022-12-18 14:02:35
【问题描述】:

我尝试实现提前停止功能以避免我的神经网络模型过度拟合。我很确定逻辑没问题,但由于某种原因,它不起作用。 我希望当验证损失大于某些时期的训练损失时,提前停止函数返回 True。但它始终返回 False,即使验证损失变得比训练损失大得多。请问你能看出问题出在哪里吗?

提前停止功能

def early_stopping(train_loss, validation_loss, min_delta, tolerance):

    counter = 0
    if (validation_loss - train_loss) > min_delta:
        counter +=1
        if counter >= tolerance:
          return True

在训练期间调用函数

for i in range(epochs):
    
    print(f"Epoch {i+1}")
    epoch_train_loss, pred = train_one_epoch(model, train_dataloader, loss_func, optimiser, device)
    train_loss.append(epoch_train_loss)

    # validation 

    with torch.no_grad(): 
       epoch_validate_loss = validate_one_epoch(model, validate_dataloader, loss_func, device)
       validation_loss.append(epoch_validate_loss)
    
    # early stopping
    if early_stopping(epoch_train_loss, epoch_validate_loss, min_delta=10, tolerance = 20):
      print("We are at epoch:", i)
      break

编辑: 训练和验证损失:

编辑2:

def train_validate (model, train_dataloader, validate_dataloader, loss_func, optimiser, device, epochs):
    preds = []
    train_loss =  []
    validation_loss = []
    min_delta = 5
    

    for e in range(epochs):
        
        print(f"Epoch {e+1}")
        epoch_train_loss, pred = train_one_epoch(model, train_dataloader, loss_func, optimiser, device)
        train_loss.append(epoch_train_loss)

        # validation 
        with torch.no_grad(): 
           epoch_validate_loss = validate_one_epoch(model, validate_dataloader, loss_func, device)
           validation_loss.append(epoch_validate_loss)
        
        # early stopping
        early_stopping = EarlyStopping(tolerance=2, min_delta=5)
        early_stopping(epoch_train_loss, epoch_validate_loss)
        if early_stopping.early_stop:
            print("We are at epoch:", e)
            break

    return train_loss, validation_loss

【问题讨论】:

    标签: python deep-learning neural-network pytorch early-stopping


    【解决方案1】:

    虽然 @KarelZe's response 充分而优雅地解决了您的问题,但我想提供一个可以说更好的替代早期停止标准。

    您的提前停止标准基于验证损失与训练损失的差异程度(以及持续时间)。当验证损失确实在减少但通常不够接近训练损失时,这将会中断。训练模型的目标是鼓励减少验证损失,而不是减少训练损失和验证损失之间的差距。

    因此,我认为更好的早期停止标准是单独观察验证损失的趋势,即,如果训练没有导致验证损失的降低,则终止它。这是一个示例实现:

    class EarlyStopper:
        def __init__(self, patience=1, min_delta=0):
            self.patience = patience
            self.min_delta = min_delta
            self.counter = 0
            self.min_validation_loss = np.inf
    
        def early_stop(self, validation_loss):
            if validation_loss < self.min_validation_loss:
                self.min_validation_loss = validation_loss
                self.counter = 0
            elif validation_loss > (self.min_validation_loss + self.min_delta):
                self.counter += 1
                if self.counter >= self.patience:
                    return True
            return False
    

    以下是您将如何使用它:

    early_stopper = EarlyStopper(patience=3, min_delta=10)
    for epoch in np.arange(n_epochs):
        train_loss = train_one_epoch(model, train_loader)
        validation_loss = validate_one_epoch(model, validation_loader)
        if early_stopper.early_stop(validation_loss):             
            break
    

    【讨论】:

    • 非常感谢您的回答。这是一个新想法,太神奇了。你很热心!
    • 感谢您提供此解决方案!我只是想知道为什么早期的解决方案要检查 train 和 val 之间的差距?那不应该是标准,不是吗?还是我错过了什么?
    【解决方案2】:

    您的实现的问题在于,每当您调用early_stopping() 时,计数器都会使用0 重新初始化。

    这是使用面向 oo 的方法与 __call__()__init__() 代替的工作解决方案:

    class EarlyStopping:
        def __init__(self, tolerance=5, min_delta=0):
    
            self.tolerance = tolerance
            self.min_delta = min_delta
            self.counter = 0
            self.early_stop = False
    
        def __call__(self, train_loss, validation_loss):
            if (validation_loss - train_loss) > self.min_delta:
                self.counter +=1
                if self.counter >= self.tolerance:  
                    self.early_stop = True
    

    像这样称呼它:

    early_stopping = EarlyStopping(tolerance=5, min_delta=10)
    
    for i in range(epochs):
        
        print(f"Epoch {i+1}")
        epoch_train_loss, pred = train_one_epoch(model, train_dataloader, loss_func, optimiser, device)
        train_loss.append(epoch_train_loss)
    
        # validation 
        with torch.no_grad(): 
           epoch_validate_loss = validate_one_epoch(model, validate_dataloader, loss_func, device)
           validation_loss.append(epoch_validate_loss)
        
        # early stopping
        early_stopping(epoch_train_loss, epoch_validate_loss)
        if early_stopping.early_stop:
          print("We are at epoch:", i)
          break
    

    例子:

    early_stopping = EarlyStopping(tolerance=2, min_delta=5)
    
    train_loss = [
        642.14990234,
        601.29278564,
        561.98400879,
        530.01501465,
        497.1098938,
        466.92709351,
        438.2364502,
        413.76028442,
        391.5090332,
        370.79074097,
    ]
    validate_loss = [
        509.13619995,
        497.3125,
        506.17315674,
        497.68960571,
        505.69918823,
        459.78610229,
        480.25592041,
        418.08630371,
        446.42675781,
        372.09902954,
    ]
    
    for i in range(len(train_loss)):
    
        early_stopping(train_loss[i], validate_loss[i])
        print(f"loss: {train_loss[i]} : {validate_loss[i]}")
        if early_stopping.early_stop:
            print("We are at epoch:", i)
            break
    
    

    输出:

    loss: 642.14990234 : 509.13619995
    loss: 601.29278564 : 497.3125
    loss: 561.98400879 : 506.17315674
    loss: 530.01501465 : 497.68960571
    loss: 497.1098938 : 505.69918823
    loss: 466.92709351 : 459.78610229
    loss: 438.2364502 : 480.25592041
    We are at epoch: 6
    

    【讨论】:

    • 非常感谢您的回答。这样写比较优雅。但它也不起作用! :( P.S. 我对你的代码做了一个小修改:self.counter +=1 和 self.counter >= self.tolerance
    • 是的当然。
    • @龙猫。谢谢。很高兴调查它。
    • @Totoro 下次请提供打印输出作为文本。我添加了一个例子。鉴于您提供的样本损失,训练提前停止。不确定您是如何或在何处添加它的。
    • 太感谢了。下次我会按照你说的提供数据。我不知道。
    猜你喜欢
    • 1970-01-01
    • 2021-10-25
    • 2018-02-27
    • 1970-01-01
    • 2020-05-19
    • 2017-08-17
    • 2018-03-24
    • 2013-06-11
    • 2021-11-04
    相关资源
    最近更新 更多