使用PyTorch-Lightning训练TransformerEncoder遇MisconfigurationException错误
问题详情
使用PyTorch-Lightning的Trainer类训练torch.nn.TransformerEncoderLayer时,第一个epoch启动前触发如下错误:
MisconfigurationException: The closure hasn't been executed. HINT: did you call
optimizer_closure()in youroptimizer_stephook? It could also happen because theoptimizer.step(optimizer_closure)call did not execute it internally.
已正确定义configure_optimizers()方法,该方法在LSTM、GRU、MultiHeadAttention等模型上均可正常工作,但替换为TransformerEncoder后出现上述错误。
模型代码
PositionalEncoder类
import torch import torch.nn as nn from torch.nn.functional import sin, cos, sqrt class PositionalEncoder(nn.Module): def __init__(self, d_model=512, max_seq_len=512): super().__init__() self.d_model = d_model pe = torch.zeros(max_seq_len, d_model) for pos in range(max_seq_len): for i in range(0, d_model, 2): pe[pos, i] = sin(pos / (10000 ** ((2 * i)/d_model))) pe[pos, i+1] = cos(pos / (10000 ** ((2 * (i + 1))/d_model))) pe = pe.unsqueeze(0) self.register_buffer('pe', pe) def forward(self, x): x *= sqrt(self.d_model) x += self.pe[:,:x.size(1)] return x
TRANSFORMER类(PyTorch-Lightning模块)
import pytorch_lightning as pl from torch.utils.data import DataLoader from transformers import get_cosine_schedule_with_warmup # 假设CRF和f1score已导入 class TRANSFORMER(pl.LightningModule): def __init__(self, input_dim, d_model=512, nhead=8, num_layers=6, dropout=0.5, use_scheduler=True, num_tags=len(TAG2IDX), total_steps=1024, train_dataset=None, val_dataset=None, test_dataset=None): super().__init__() self.crf = CRF(num_tags=num_tags, batch_first=True) self.fc = nn.Linear(d_model, num_tags) self.use_scheduler = use_scheduler self.embedding = nn.Embedding(num_embeddings=input_dim, embedding_dim=d_model, padding_idx=0) self.pos_encoder = PositionalEncoder(d_model=d_model) self.encoder_layer = nn.TransformerEncoderLayer(d_model=d_model, nhead=nhead, dropout=dropout, activation="gelu", batch_first=True) self.encoder = nn.TransformerEncoder(encoder_layer=self.encoder_layer, num_layers=num_layers) ## Hyperparameters ## self.learning_rate = LEARNING_RATE self.weight_decay = WEIGHT_DECAY self.total_steps = total_steps self.batch_size = BATCH_SIZE ## Datasets ## self.train_dataset = train_dataset self.val_dataset = val_dataset self.test_dataset = test_dataset ## steps ## if self.use_scheduler: self.total_steps = len(train_dataset) // self.batch_size # 创建数据加载器 def train_dataloader(self): return DataLoader(self.train_dataset, batch_size=self.batch_size, num_workers=N_JOBS, shuffle=True, drop_last=False) def val_dataloader(self): return DataLoader(self.val_dataset, batch_size=self.batch_size, num_workers=N_JOBS, shuffle=False, drop_last=False) def test_dataloader(self): return DataLoader(self.test_dataset, batch_size=self.batch_size, num_workers=N_JOBS, shuffle=False, drop_last=False) def forward(self, input_ids, masks): out = self.embedding(input_ids) out = self.pos_encoder(out) out = self.encoder(out, src_key_padding_mask=~masks) out = self.fc(out) return out def _shared_evaluation_step(self, batch, batch_idx): ids, masks, lbls = batch emissions = self(ids, masks) loss = -self.crf(emissions, lbls, mask=masks) pred = self.crf.decode(emissions, mask=masks) r, p, f1 = f1score(lbls, pred) return loss, r, p, f1 def training_step(self, batch, batch_idx): loss, r, p, f1 = self._shared_evaluation_step(batch, batch_idx) self.log("train_loss", loss, on_step=False, on_epoch=True, prog_bar=True) self.log("train_recall", r, on_step=False, on_epoch=True, prog_bar=True) self.log("train_precision", p, on_step=False, on_epoch=True, prog_bar=True) self.log("train_f1score", f1, on_step=False, on_epoch=True, prog_bar=True) return loss def validation_step(self, batch, batch_idx): loss, r, p, f1 = self._shared_evaluation_step(batch, batch_idx) self.log("val_loss", loss, on_step=False, on_epoch=True, prog_bar=True) self.log("val_recall", r, on_step=False, on_epoch=True, prog_bar=True) self.log("val_precision", p, on_step=False, on_epoch=True, prog_bar=True) self.log("val_f1score", f1, on_step=False, on_epoch=True, prog_bar=True) def test_step(self, batch, batch_idx): loss, r, p, f1 = self._shared_evaluation_step(batch, batch_idx) self.log("test_loss", loss, on_step=False, on_epoch=True, prog_bar=True) self.log("test_recall", r, on_step=False, on_epoch=True, prog_bar=True) self.log("test_precision", p, on_step=False, on_epoch=True, prog_bar=True) self.log("test_f1score", f1, on_step=False, on_epoch=True, prog_bar=True) def predict_step(self, batch, batch_idx, dataloader_idx=0): ids, masks, _ = batch return self.crf.decode(self(ids, masks), mask=masks) def configure_optimizers(self): optimizer = Ranger(self.parameters(), lr=self.learning_rate, weight_decay=self.weight_decay) if self.use_scheduler: scheduler = get_cosine_schedule_with_warmup(optimizer=optimizer, num_warmup_steps=1, num_training_steps=self.total_steps) lr_scheduler = { 'scheduler': scheduler, 'interval': 'epoch', 'frequency': 1 } return [optimizer], [lr_scheduler] else: return [optimizer]
Trainer使用代码
trainer = pl.Trainer(accelerator="gpu", max_epochs=EPOCHS, precision=32, log_every_n_steps=1, callbacks=[earlystopping_callback, checkpoint_callback])
问题原因与解决方案
这个错误核心是PyTorch-Lightning的优化器闭包执行逻辑与自定义优化器/Transformer模型的兼容性问题,以下是针对性解决办法:
1. 手动重写optimizer_step方法
在TRANSFORMER类中添加该方法,强制执行闭包:
def optimizer_step(self, epoch, batch_idx, optimizer, optimizer_idx, optimizer_closure, on_tpu=False, using_native_amp=False, using_lbfgs=False): # 先执行闭包计算损失与梯度 optimizer_closure() # 再执行优化器更新 optimizer.step()
2. 替换为PyTorch原生优化器
自定义优化器(如Ranger)可能未完全适配PyTorch-Lightning的闭包机制,暂时换成AdamW验证:
def configure_optimizers(self): optimizer = torch.optim.AdamW(self.parameters(), lr=self.learning_rate, weight_decay=self.weight_decay) if self.use_scheduler: scheduler = get_cosine_schedule_with_warmup(optimizer=optimizer, num_warmup_steps=1, num_training_steps=self.total_steps) lr_scheduler = { 'scheduler': scheduler, 'interval': 'epoch', 'frequency': 1 } return [optimizer], [lr_scheduler] else: return [optimizer]
3. 验证Transformer的mask逻辑
确认src_key_padding_mask的格式正确:PyTorch中该mask的True表示对应位置为padding需要被屏蔽。你的代码中用~masks,需确保masks的True代表有效输入位置,取反后逻辑正确。可在forward方法中打印~masks的形状和部分值验证。
4. 版本与参数检查
- 确保PyTorch和PyTorch-Lightning版本匹配,建议升级到稳定版(如PyTorch 2.x + PyTorch-Lightning 2.x)。
- 修正
total_steps计算:当数据集长度无法被batch_size整除时,需加1避免调度器步数不足:if self.use_scheduler: self.total_steps = len(train_dataset) // self.batch_size + (1 if len(train_dataset) % self.batch_size !=0 else 0)
内容的提问来源于stack exchange,提问作者LazyCoder11011

