PyTorch Lightning 2.0.1验证时机Bug:单训练批次后触发全量验证
PyTorch Lightning 2.0.1 训练-验证循环逻辑修正
当前问题
现有伪代码中,if should_check_val: val_loop()被嵌套在for batch in train_dataloader():的训练批次循环内部,导致每完成一个训练批次就触发一次全量验证流程,不符合常规的"整轮训练后验证"逻辑。
预期行为
需将验证触发逻辑移至训练批次循环外部,确保整个训练轮次的所有批次完成训练后,再执行全量验证流程。
修正后的伪代码
def train_on_device(model): setup("fit") configure_optimizers() on_fit_start() # the sanity check runs here on_train_start() for epoch in epochs: fit_loop() #! CALLS Method Below on_train_end() on_fit_end() teardown("fit") def fit_loop(): model.train() torch.set_grad_enabled(True) on_train_epoch_start() for batch in train_dataloader(): # TRAINING Loop on_train_batch_start() on_before_batch_transfer() transfer_batch_to_device() on_after_batch_transfer() out = training_step() on_before_zero_grad() optimizer_zero_grad() on_before_backward() backward() on_after_backward() on_before_optimizer_step() configure_gradient_clipping() optimizer_step() on_train_batch_end(out, batch, batch_idx) # --- 修正:将验证逻辑移至训练批次循环外部 --- if should_check_val: val_loop() #! VALIDATION Loop on_train_epoch_end()
内容的提问来源于stack exchange,提问作者oolveea
相关产品推荐
相关产品推荐

