在PyTorch Lightning中使用Dropout时是否需调用model.train()?
PyTorch Lightning中含Dropout模型是否需手动调用model.train()?
在普通PyTorch教程中,由于模型包含Dropout层,作者会在每个epoch中调用model.train()。我基于PyTorch Lightning构建了含Dropout的模型,使用Trainer进行训练,想咨询:是否需要在Lightning Module中手动调用model.train()?若需要该如何实现,还是框架会自动处理?
以下是相关代码:
我的模型代码
class MulticlassClassificationLIGHT(pl.LightningModule): def __init__(self,class_weights): super(MulticlassClassificationLIGHT, self).__init__() self.num_feature=35 self.num_class=36 self.layer_1 = nn.Linear(self.num_feature, 512) self.layer_2 = nn.Linear(512, 128) self.layer_3 = nn.Linear(128, 64) self.layer_out = nn.Linear(64, self.num_class) self.relu = nn.ReLU() self.dropout = nn.Dropout(p=0.2) self.batchnorm1 = nn.BatchNorm1d(512) self.batchnorm2 = nn.BatchNorm1d(128) self.batchnorm3 = nn.BatchNorm1d(64) self.loss = nn.CrossEntropyLoss(weight=class_weights.to(device)) def forward(self, x): x = self.layer_1(x) x = self.batchnorm1(x) x = self.relu(x) x = self.layer_2(x) x = self.batchnorm2(x) x = self.relu(x) x = self.dropout(x) x = self.layer_3(x) x = self.batchnorm3(x) x = self.relu(x) x = self.dropout(x) x = self.layer_out(x) return x def training_step(self, batch, batch_idx): x, y = batch logits = self.forward(x) loss = self.loss(logits, y) self.log("train_loss", loss, prog_bar=True, logger=True) return loss
普通PyTorch教程示例代码
for e in tqdm(range(1, EPOCHS+1)): train_epoch_loss = 0 train_epoch_acc = 0 model.train() for X_train_batch, y_train_batch in train_loader: X_train_batch, y_train_batch = X_train_batch.to(device), y_train_batch.to(device) optimizer.zero_grad()
我的Trainer代码
trainer = pl.Trainer( devices="auto", accelerator="auto", auto_lr_find=False, auto_scale_batch_size=True, fast_dev_run=False, num_sanity_val_steps=3, logger=logger, min_epochs=EPOCHS, )
训练调用代码
trainer.fit(model, data_module_classifier.train_dataloader(),data_module_classifier.val_dataloader() )
回答
不需要手动调用model.train(),PyTorch Lightning框架会自动处理模型模式的切换:
- 当执行训练流程(如
trainer.fit()中的训练阶段)时,Lightning会自动将模型切换到train()模式,此时Dropout层会正常启用随机失活,BatchNorm层会更新运行时统计量。 - 当进入验证/测试阶段时,框架会自动切换到eval()模式,关闭Dropout的随机失活行为,同时固定BatchNorm的统计量,避免影响推理结果。
你当前的代码写法完全没问题,training_step中直接调用self.forward(x)即可,框架已经确保此时模型处于训练模式。
内容的提问来源于stack exchange,提问作者Master_Sniffer
相关产品推荐
相关产品推荐

