【问题标题】:Does the Pytorch Lightning Trainer use the validation data to optimize the models weights?Pytorch Lightning Trainer 是否使用验证数据来优化模型权重?
【发布时间】:2022-04-13 21:32:27
【问题描述】:

我目前正在使用大量使用 Pytorch Lightning 的 Pytorch Forecasting。在这里,我应用 Pytorch Lightning Trainer 来训练一个 Temporal Fusion Transformer Model,大致遵循这个 example 的轮廓。我粗略的训练代码和模型定义如下:

training = TimeSeriesDataSet(
    df_train[lambda x: x.time_idx <= training_cutoff],
    time_idx="time_idx",
    target="target",
    group_ids=["group"],
    max_prediction_length=90,
    min_encoder_length=365 // 2,
    max_encoder_length=365, 
    time_varying_unknown_reals=["target"], 
    time_varying_known_reals=["time_idx"]
)

validation = TimeSeriesDataSet.from_dataset(training, df_train, predict=True, stop_randomization=True)

# create dataloaders for model
batch_size = 4  
train_dataloader = training.to_dataloader(train=True, batch_size=batch_size, num_workers=0)
val_dataloader = validation.to_dataloader(train=False, batch_size=batch_size, num_workers=0)

tft = TemporalFusionTransformer.from_dataset(
    training,
    learning_rate=res.suggestion(),
    hidden_size=16,
    attention_head_size=1,
    dropout=0.1,
    hidden_continuous_size=8,
    output_size=7,  
    loss=QuantileLoss(),
    log_interval=10,  
    reduce_on_plateau_patience=4,
    time_varying_reals_encoder=["target"],
    time_varying_reals_decoder=["target"]
)

trainer = pl.Trainer(
    max_epochs=15,
    gpus=0,
    weights_summary="top",
    gradient_clip_val=0.1,
    limit_train_batches=30,
    callbacks=[lr_logger, early_stop_callback],
    logger=logger,
)

trainer.fit(
    tft,
    train_dataloader,
    val_dataloader
)

现在我的问题是,验证数据是否对模型的优化有影响?我一直在玩max_prediction_length 参数,当我将验证时间窗口设置为更大的时间范围时,模型的性能似乎更好。 Pytorch Lightning Trainer 是否使用验证数据来优化模型,还是我遗漏了什么?

提前非常感谢!

【问题讨论】:

  • 我看到您正在使用 Early Stopping。您还没有指定如何实例化 early_stop_callback ?它可能会使用验证指标来停止训练——这就是早期停止的工作原理。
  • 谢谢,我真的应该更彻底地研究我复制的代码!
  • 不,他们不应该因为它会泄露数据,在这种情况下,你的验证数据将成为训练数据......

标签: python pytorch forecasting pytorch-lightning


【解决方案1】:

由于PyTorch-forecasting是建立在PyTorch-lightning之上的抽象,我们可以参考后者的trainer抽象的文档框架 (https://pytorch-lightning.readthedocs.io/en/stable/common/trainer.html)。

# put model in train mode
model.train()
torch.set_grad_enabled(True)

losses = []
for batch in train_dataloader:
    # calls hooks like this one
    on_train_batch_start()

    # train step
    loss = training_step(batch)

    # clear gradients
    optimizer.zero_grad()

    # backward
    loss.backward()

    # update parameters
    optimizer.step()

    losses.append(loss)

在上面的示例中,我们可以看到 trainer 仅计算 train_dataloader 中的批次损失并将损失传播回去。这意味着验证集不用于更新模型的权重。

【讨论】:

    猜你喜欢
    • 2021-09-05
    • 2017-02-16
    • 1970-01-01
    • 2021-02-11
    • 2020-09-12
    • 1970-01-01
    • 2021-08-27
    • 2017-10-30
    • 2021-02-06
    相关资源
    最近更新 更多