ARTICLE DETAIL

资讯详情

深耕网站建设与运营推广的一线实战洞察。

【Bug已解决】PyTorch: is there a definitive training loop similar to Keras‘ fit()? 解决方案

【Bug已解决】PyTorch: is there a definitive training loop similar to Keras‘ fit()? 解决方案 【Bug已解决】PyTorch: is there a definitive training loop similar to Keras fit()? 解决方案问题描述从 Keras 迁移到 PyTorch 的开发者最常问的一个问题是PyTorch 有没有类似 Kerasmodel.fit()的一行式训练 APIKeras 的model.fit(x, y, epochs10)只需要一行代码就能完成完整的训练循环包括前向传播、反向传播、梯度更新、进度条、验证等。而 PyTorch 要求开发者手动编写训练循环涉及optimizer.zero_grad()、loss.backward()、optimizer.step()等细节。这种差异带来了几个问题初学者需要理解更多底层概念才能开始训练每个项目都要重复编写相似的训练循环代码容易在训练循环中引入 bug如忘记zero_grad、忘记eval()模式等缺少统一的训练日志和回调机制错误复现以下代码展示了 PyTorch 手动训练循环的繁琐性以及常见错误import torch import torch.nn as nn import torch.optim as optim # 繁琐的手动训练循环 model nn.Linear(100, 10) optimizer optim.Adam(model.parameters(), lr1e-3) criterion nn.CrossEntropyLoss() # 典型的 PyTorch 训练循环繁琐且容易出错 for epoch in range(10): # 容易忘记切换到训练模式 model.train() for batch_x, batch_y in train_loader: # 容易忘记清零梯度 optimizer.zero_grad() # 前向传播 output model(batch_x) loss criterion(output, batch_y) # 反向传播 loss.backward() # 参数更新 optimizer.step() # 验证容易忘记切换到 eval 模式和 no_grad model.eval() # 容易忘记 with torch.no_grad(): # 容易忘记 val_loss 0 for batch_x, batch_y in val_loader: output model(batch_x) val_loss criterion(output, batch_y) print(fEpoch {epoch}: val_loss {val_loss / len(val_loader)}) # 常见错误 # 1. 忘记 optimizer.zero_grad() - 梯度累积 # 2. 忘记 model.train() / model.eval() - BatchNorm/Dropout 行为错误 # 3. 忘记 torch.no_grad() - 验证时构建计算图内存泄漏 # 4. 忘记 optimizer.step() - 模型不更新 # 5. 没有学习率调度 - 训练效果不佳 # 6. 没有保存最佳模型 - 训练后无法恢复最佳状态根因分析1. PyTorch 的设计哲学PyTorch 采用显式优于隐式的设计哲学。与 Keras 的高层封装不同PyTorch 要求开发者明确写出每一步操作。这提供了最大的灵活性但增加了样板代码。2. Kerasfit()的封装内容Keras 的fit()内部封装了大量逻辑训练/验证模式切换梯度清零和更新损失计算和累积指标跟踪进度条显示回调系统早停、学习率调度、模型保存等验证集评估日志记录3. PyTorch 生态的解决方案虽然 PyTorch 核心库没有提供fit()等价物但 PyTorch 生态中有多个高级训练框架PyTorch Lightning最流行的 PyTorch 高级封装Hugging Face Trainer专为 NLP 模型设计fastai提供类似 Keras 的简洁 APIIgnitePyTorch 官方的高级训练工具解决方案方案一使用 PyTorch LightningPyTorch Lightning 是最接近 Kerasfit()体验的 PyTorch 高级框架# pip install pytorch-lightning import torch import torch.nn as nn import pytorch_lightning as pl from torch.utils.data import DataLoader, TensorDataset class LitModel(pl.LightningModule): Lightning 模型类似 Keras 的 Model 类 def __init__(self, input_dim100, hidden_dim128, num_classes10): super().__init__() self.save_hyperparameters() self.model nn.Sequential( nn.Linear(input_dim, hidden_dim), nn.ReLU(), nn.Dropout(0.3), nn.Linear(hidden_dim, hidden_dim), nn.ReLU(), nn.Linear(hidden_dim, num_classes), ) self.criterion nn.CrossEntropyLoss() def forward(self, x): return self.model(x) def training_step(self, batch, batch_idx): x, y batch logits self(x) loss self.criterion(logits, y) # 自动记录损失 self.log(train_loss, loss, prog_barTrue) # 计算准确率 acc (logits.argmax(dim1) y).float().mean() self.log(train_acc, acc, prog_barTrue) return loss def validation_step(self, batch, batch_idx): x, y batch logits self(x) loss self.criterion(logits, y) acc (logits.argmax(dim1) y).float().mean() self.log(val_loss, loss, prog_barTrue) self.log(val_acc, acc, prog_barTrue) def configure_optimizers(self): optimizer torch.optim.Adam(self.parameters(), lr1e-3) scheduler torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max10) return [optimizer], [scheduler] # 创建数据 x_train torch.randn(1000, 100) y_train torch.randint(0, 10, (1000,)) x_val torch.randn(200, 100) y_val torch.randint(0, 10, (200,)) train_loader DataLoader(TensorDataset(x_train, y_train), batch_size32, shuffleTrue) val_loader DataLoader(TensorDataset(x_val, y_val), batch_size32) # 创建模型和训练器 model LitModel() # 类似 Keras fit() 的一行式训练 trainer pl.Trainer( max_epochs10, callbacks[ pl.callbacks.EarlyStopping(monitorval_loss, patience3), pl.callbacks.ModelCheckpoint(monitorval_loss, save_top_k1), ], enable_progress_barTrue, ) # 一行训练 trainer.fit(model, train_loader, val_loader)方案二自定义 Keras 风格训练器如果不依赖外部库可以自己封装一个类似 Kerasfit()的训练器import torch import torch.nn as nn import torch.optim as optim from torch.utils.data import DataLoader from typing import Optional, Callable, List, Dict, Any import time import copy class KerasStyleTrainer: Keras 风格的 PyTorch 训练器。 提供 model.fit() 类似的简洁 API。 def __init__(self, model, optimizerNone, criterionNone, deviceauto): self.model model self.device self._get_device(device) self.model self.model.to(self.device) self.optimizer optimizer or optim.Adam(model.parameters(), lr1e-3) self.criterion criterion or nn.CrossEntropyLoss() self.history {train_loss: [], val_loss: [], train_acc: [], val_acc: []} self.callbacks [] self.best_model_state None self.best_val_loss float(inf) def _get_device(self, device): if device auto: return torch.device(cuda if torch.cuda.is_available() else cpu) return torch.device(device) def add_callback(self, callback): 添加回调 self.callbacks.append(callback) def fit(self, train_loader, val_loaderNone, epochs10, learning_rateNone, verboseTrue): 训练模型类似 Keras 的 model.fit()。 Args: train_loader: 训练数据 DataLoader val_loader: 验证数据 DataLoader可选 epochs: 训练轮数 learning_rate: 学习率可选覆盖 optimizer 的 lr verbose: 是否打印进度 if learning_rate: for param_group in self.optimizer.param_groups: param_group[lr] learning_rate # 触发训练开始回调 self._on_train_begin() for epoch in range(1, epochs 1): # 触发 epoch 开始回调 self._on_epoch_begin(epoch) start_time time.time() # 训练 train_metrics self._train_epoch(train_loader) # 验证 val_metrics {} if val_loader: val_metrics self._validate(val_loader) # 记录历史 self.history[train_loss].append(train_metrics[loss]) self.history[train_acc].append(train_metrics[accuracy]) if val_loader: self.history[val_loss].append(val_metrics[loss]) self.history[val_acc].append(val_metrics[accuracy]) elapsed time.time() - start_time if verbose: self._print_progress(epoch, epochs, train_metrics, val_metrics, elapsed) # 触发 epoch 结束回调 stop_training self._on_epoch_end(epoch, val_metrics) if stop_training: if verbose: print(fEarly stopping at epoch {epoch}) break # 恢复最佳模型 if self.best_model_state: self.model.load_state_dict(self.best_model_state) if verbose: print(fRestored best model (val_loss{self.best_val_loss:.4f})) return self.history def _train_epoch(self, dataloader): 训练一个 epoch self.model.train() total_loss 0 correct 0 total 0 for batch_x, batch_y in dataloader: batch_x batch_x.to(self.device) batch_y batch_y.to(self.device) self.optimizer.zero_grad() output self.model(batch_x) loss self.criterion(output, batch_y) loss.backward() self.optimizer.step() total_loss loss.item() pred output.argmax(dim1) correct pred.eq(batch_y).sum().item() total batch_y.size(0) return { loss: total_loss / len(dataloader), accuracy: 100. * correct / total, } def _validate(self, dataloader): 验证 self.model.eval() total_loss 0 correct 0 total 0 with torch.no_grad(): for batch_x, batch_y in dataloader: batch_x batch_x.to(self.device) batch_y batch_y.to(self.device) output self.model(batch_x) loss self.criterion(output, batch_y) total_loss loss.item() pred output.argmax(dim1) correct pred.eq(batch_y).sum().item() total batch_y.size(0) return { loss: total_loss / len(dataloader), accuracy: 100. * correct / total, } def _print_progress(self, epoch, epochs, train_m, val_m, elapsed): 打印进度 msg fEpoch {epoch}/{epochs} - {elapsed:.1f}s - msg floss: {train_m[loss]:.4f} - acc: {train_m[accuracy]:.2f}% if val_m: msg f - val_loss: {val_m[loss]:.4f} - val_acc: {val_m[accuracy]:.2f}% print(msg) def _on_train_begin(self): for cb in self.callbacks: if hasattr(cb, on_train_begin): cb.on_train_begin(self) def _on_epoch_begin(self, epoch): for cb in self.callbacks: if hasattr(cb, on_epoch_begin): cb.on_epoch_begin(self, epoch) def _on_epoch_end(self, epoch, val_metrics): 返回 True 表示应该停止训练 # 保存最佳模型 if val_metrics and val_metrics.get(loss, float(inf)) self.best_val_loss: self.best_val_loss val_metrics[loss] self.best_model_state copy.deepcopy(self.model.state_dict()) stop False for cb in self.callbacks: if hasattr(cb, on_epoch_end): if cb.on_epoch_end(self, epoch, val_metrics): stop True return stop def predict(self, dataloader): 预测 self.model.eval() predictions [] with torch.no_grad(): for batch_x in dataloader: if isinstance(batch_x, (list, tuple)): batch_x batch_x[0] batch_x batch_x.to(self.device) output self.model(batch_x) predictions.append(output.cpu()) return torch.cat(predictions) def save(self, path): 保存模型 torch.save({ model_state: self.model.state_dict(), ![配图](https://i-blog.csdnimg.cn/img_convert/43c7b3e7961b1828dc25f54036c7dc9a.png) optimizer_state: self.optimizer.state_dict(), history: self.history, }, path) def load(self, path): 加载模型 checkpoint torch.load(path) self.model.load_state_dict(checkpoint[model_state]) self.optimizer.load_state_dict(checkpoint[optimizer_state]) self.history checkpoint[history] # 回调系统 class EarlyStopping: 早停回调 def __init__(self, monitorval_loss, patience3, min_delta0.0): self.monitor monitor self.patience patience self.min_delta min_delta self.wait 0 self.best float(inf) def on_epoch_end(self, trainer, epoch, val_metrics): if not val_metrics or self.monitor not in val_metrics: return False current val_metrics[self.monitor] if current self.best - self.min_delta: self.best current self.wait 0 else: self.wait 1 if self.wait self.patience: return True # 停止训练 return False class ModelCheckpoint: 模型保存回调 def __init__(self, filepathbest_model.pth, monitorval_loss, save_best_onlyTrue): self.filepath filepath self.monitor monitor self.save_best_only save_best_only self.best float(inf) def on_epoch_end(self, trainer, epoch, val_metrics): if not val_metrics or self.monitor not in val_metrics: return False current val_metrics[self.monitor] if not self.save_best_only or current self.best: self.best current trainer.save(self.filepath) print(f Model saved to {self.filepath}) return False class LearningRateScheduler: 学习率调度回调 def __init__(self, scheduler): self.scheduler scheduler def on_epoch_end(self, trainer, epoch, val_metrics): self.scheduler.step() current_lr trainer.optimizer.param_groups[0][lr] print(f Learning rate: {current_lr:.6f}) return False class ReduceLROnPlateau: 验证损失停滞时降低学习率 def __init__(self, monitorval_loss, factor0.5, patience2, min_lr1e-6): self.monitor monitor self.factor factor self.patience patience self.min_lr min_lr self.wait 0 self.best float(inf) def on_epoch_end(self, trainer, epoch, val_metrics): if not val_metrics or self.monitor not in val_metrics: return False current val_metrics[self.monitor] if current self.best: self.best current self.wait 0 else: self.wait 1 if self.wait self.patience: for param_group in trainer.optimizer.param_groups: new_lr max(param_group[lr] * self.factor, self.min_lr) param_group[lr] new_lr self.wait 0 print(f Reduce LR to {new_lr:.6f}) return False # 使用示例 if __name__ __main__: from torch.utils.data import TensorDataset # 创建数据 x_train torch.randn(1000, 100) y_train torch.randint(0, 10, (1000,)) x_val torch.randn(200, 100) y_val torch.randint(0, 10, (200,)) train_loader DataLoader(TensorDataset(x_train, y_train), batch_size32, shuffleTrue) val_loader DataLoader(TensorDataset(x_val, y_val), batch_size32) # 创建模型 model nn.Sequential( nn.Linear(100, 128), nn.ReLU(), nn.Dropout(0.3), nn.Linear(128, 64), nn.ReLU(), nn.Linear(64, 10), ) # 创建训练器 trainer KerasStyleTrainer( modelmodel, optimizeroptim.Adam(model.parameters(), lr1e-3), criterionnn.CrossEntropyLoss(), deviceauto, ) # 添加回调 trainer.add_callback(EarlyStopping(monitorval_loss, patience3)) trainer.add_callback(ModelCheckpoint(best_model.pth, monitorval_loss)) trainer.add_callback(ReduceLROnPlateau(factor0.5, patience2)) # 一行训练类似 Keras fit() print( * 60) print(Keras 风格训练) print( * 60) history trainer.fit( train_loadertrain_loader, val_loaderval_loader, epochs20, verboseTrue, ) print(f\n训练历史: {history}) # 预测 predictions trainer.predict(val_loader) print(f预测形状: {predictions.shape})方案三使用 fastaifastai 提供了最接近 Keras 简洁性的 PyTorch API# pip install fastai from fastai.vision.all import * # 创建数据加载器 dls ImageDataLoaders.from_folder( pathdata/, valid_pct0.2, item_tfmsResize(224), batch_tfmsaug_transforms(), ) # 创建学习器 learn vision_learner(dls, resnet18, metricsaccuracy) # 一行训练 learn.fit(10) # 或使用 fine_tune迁移学习 learn.fine_tune(5) # 查看训练历史 learn.recorder.plot_loss()完整修复代码以下是一个完整的、生产级别的 Keras 风格训练框架import torch import torch.nn as nn import torch.optim as optim from torch.utils.data import DataLoader, Dataset from typing import Optional, List, Dict, Any, Callable, Union import time import copy import os import json from collections import defaultdict class Metric: 指标基类 def __init__(self, name): self.name name self.reset() def reset(self): self.total 0 self.count 0 def update(self, output, target): raise NotImplementedError def compute(self): raise NotImplementedError class Accuracy(Metric): 准确率指标 def __init__(self): super().__init__(accuracy) def update(self, output, target): pred output.argmax(dim1) self.total pred.eq(target).sum().item() self.count target.size(0) def compute(self): return 100. * self.total / self.count if self.count 0 else 0 class LossMetric(Metric): 损失指标 def __init__(self, criterion): super().__init__(loss) self.criterion criterion def update(self, output, target): self.total self.criterion(output, target).item() self.count 1 def compute(self): return self.total / self.count if self.count 0 else 0 class Callback: 回调基类 def on_train_begin(self, trainer): pass def on_train_end(self, trainer): pass def on_epoch_begin(self, trainer, epoch): pass def on_epoch_end(self, trainer, epoch, metrics) - bool: return False def on_batch_begin(self, trainer, batch_idx): pass def on_batch_end(self, trainer, batch_idx, loss): pass class EarlyStopping(Callback): def __init__(self, monitorval_loss, patience5, min_delta0.001): self.monitor monitor self.patience patience self.min_delta min_delta self.wait 0 self.best float(inf) def on_epoch_end(self, trainer, epoch, metrics): current metrics.get(self.monitor, float(inf)) if current self.best - self.min_delta: self.best current self.wait 0 else: self.wait 1 if self.wait self.patience: print(fEarly stopping: no improvement for {self.patience} epochs) return True return False class ModelCheckpoint(Callback): def __init__(self, filepath, monitorval_loss, save_best_onlyTrue): self.filepath filepath self.monitor monitor self.save_best_only save_best_only self.best float(inf) def on_epoch_end(self, trainer, epoch, metrics): if self.save_best_only: current metrics.get(self.monitor, float(inf)) if current self.best: self.best current trainer.save(self.filepath) print(f Checkpoint saved: {self.filepath}) else: trainer.save(self.filepath.format(epochepoch)) return False class ProgressLogger(Callback): def __init__(self, print_every1): self.print_every print_every def on_epoch_begin(self, trainer, epoch): self.epoch_start time.time() def on_epoch_end(self, trainer, epoch, metrics): if epoch % self.print_every 0: elapsed time.time() - self.epoch_start parts [fEpoch {epoch}/{trainer.epochs}] for k, v in metrics.items(): parts.append(f{k}: {v:.4f}) parts.append(f[{elapsed:.1f}s]) print( - .join(parts)) return False class GradientClipping(Callback): def __init__(self, max_norm1.0): self.max_norm max_norm def on_batch_end(self, trainer, batch_idx, loss): torch.nn.utils.clip_grad_norm_(trainer.model.parameters(), self.max_norm) class FitTrainer: 完整的 Keras 风格训练器。 支持 fit()、predict()、evaluate() 方法。 def __init__(self, model, optimizerNone, criterionNone, metricsNone, deviceauto): self.model model self.device self._get_device(device) self.model self.model.to(self.device) self.optimizer optimizer or optim.Adam(model.parameters(), lr1e-3) self.criterion criterion or nn.CrossEntropyLoss() self.metrics metrics or [] self.callbacks [] self.history defaultdict(list) self.epochs 0 def _get_device(self, device): if device auto: return torch.device(cuda if torch.cuda.is_available() else cpu) return torch.device(device) def compile(self, optimizerNone, criterionNone, metricsNone): 类似 Keras 的 compile() if optimizer: self.optimizer optimizer if criterion: self.criterion criterion if metrics: self.metrics metrics def fit(self, train_loader, val_loaderNone, epochs10, callbacksNone, verboseTrue): 类似 Keras 的 fit() self.epochs epochs all_callbacks list(callbacks or []) self.callbacks if verbose: all_callbacks.append(ProgressLogger()) # 触发训练开始 for cb in all_callbacks: cb.on_train_begin(self) for epoch in range(1, epochs 1): for cb in all_callbacks: cb.on_epoch_begin(self, epoch) # 训练 train_metrics self._run_epoch(train_loader, trainingTrue, prefixtrain) # 验证 val_metrics {} if val_loader: val_metrics self._run_epoch(val_loader, trainingFalse, prefixval) # 合并指标 all_metrics {**train_metrics, **val_metrics} # 记录历史 for k, v in all_metrics.items(): self.history[k].append(v) # 触发 epoch 结束 stop False for cb in all_callbacks: if cb.on_epoch_end(self, epoch, all_metrics): stop True if stop: break for cb in all_callbacks: cb.on_train_end(self) return dict(self.history) def _run_epoch(self, dataloader, training, prefix): 运行一个 epoch if training: self.model.train() else: self.model.eval() # 初始化指标 loss_metric LossMetric(self.criterion) metric_objects [m() if isinstance(m, type) else m for m in self.metrics] for m in metric_objects: m.reset() context torch.enable_grad() if training else torch.no_grad() with context: for batch_idx, (batch_x, batch_y) in enumerate(dataloader): batch_x batch_x.to(self.device) batch_y batch_y.to(self.device) if training: self.optimizer.zero_grad() output self.model(batch_x) loss self.criterion(output, batch_y) if training: loss.backward() self.optimizer.step() loss_metric.update(output, batch_y) for m in metric_objects: m.update(output, batch_y) # 构建指标字典 metrics {f{prefix}_loss: loss_metric.compute()} for m in metric_objects: metrics[f{prefix}_{m.name}] m.compute() return metrics torch.no_grad() def predict(self, dataloader): 预测 self.model.eval() results [] for batch in dataloader: if isinstance(batch, (list, tuple)): batch batch[0] batch batch.to(self.device) output self.model(batch) results.append(output.cpu()) return torch.cat(results) def evaluate(self, dataloader): 评估 return self._run_epoch(dataloader, trainingFalse, prefixeval) def save(self, path): 保存 torch.save({ model_state: self.model.state_dict(), optimizer_state: self.optimizer.state_dict(), history: dict(self.history), }, path) def load(self, path): 加载 ckpt torch.load(path, map_locationself.device) self.model.load_state_dict(ckpt[model_state]) self.optimizer.load_state_dict(ckpt[optimizer_state]) self.history defaultdict(list, ckpt.get(history, {})) # 使用示例 if __name__ __main__: # 创建数据 x_train torch.randn(1000, 50) y_train torch.randint(0, 5, (1000,)) x_val torch.randn(200, 50) y_val torch.randint(0, 5, (200,)) train_ds torch.utils.data.TensorDataset(x_train, y_train) val_ds torch.utils.data.TensorDataset(x_val, y_val) train_loader DataLoader(train_ds, batch_size32, shuffleTrue) val_loader DataLoader(val_ds, batch_size32) # 创建模型 model nn.Sequential( nn.Linear(50, 128), nn.ReLU(), nn.Dropout(0.3), nn.Linear(128, 64), nn.ReLU(), nn.Linear(64, 5), ) # 创建训练器 trainer FitTrainer( modelmodel, optimizeroptim.AdamW(model.parameters(), lr1e-3, weight_decay1e-4), criterionnn.CrossEntropyLoss(), metrics[Accuracy], deviceauto, ) # 添加回调 callbacks [ EarlyStopping(monitorval_loss, patience5), ModelCheckpoint(best_model.pth, monitorval_loss), GradientClipping(max_norm1.0), ] # 训练类似 Keras fit() print( * 60) print(FitTrainer 训练) print( * 60) history trainer.fit( train_loadertrain_loader, val_loaderval_loader, epochs20, callbackscallbacks, ) # 评估 print(\n--- 评估 ---) eval_metrics trainer.evaluate(val_loader) print(f评估指标: {eval_metrics}) # 预测 predictions trainer.predict(val_loader) print(f预测形状: {predictions.shape}) print(\n训练完成)常见陷阱与注意事项1.model.train()vsmodel.eval()训练时必须调用model.train()验证/推理时必须调用model.eval()。这影响Dropout和BatchNorm的行为。忘记切换是训练循环中最常见的 bug。2.torch.no_grad()的使用验证和推理时必须使用torch.no_grad()上下文管理器否则会构建计算图导致内存泄漏和速度下降。3. 梯度清零每次backward()前必须调用optimizer.zero_grad()。PyTorch 默认累积梯度不清零会导致梯度不断叠加。4. 学习率调度器的调用时机lr_scheduler.step()应该在每个 epoch 结束后调用对于StepLR、CosineAnnealingLR等而不是每个 batch 后。OneCycleLR例外它在每个 batch 后调用。5. 混合精度训练使用torch.cuda.amp进行混合精度训练时需要使用GradScaler和autocast不能简单地用loss.backward()和optimizer.step()。6. 多 GPU 训练使用nn.DataParallel或nn.DistributedDataParallel时训练循环需要相应调整。DistributedDataParallel需要设置DistributedSampler并在每个 epoch 调用sampler.set_epoch(epoch)。总结虽然 PyTorch 核心库没有提供 Kerasfit()等价的一行式 API但通过以下方式可以实现类似的简洁体验使用 PyTorch Lightning最成熟的 PyTorch 高级框架提供trainer.fit()一行式训练支持多 GPU、混合精度、日志记录等。自定义训练器封装KerasStyleTrainer或FitTrainer将训练循环的样板代码抽象为可复用的类。回调系统实现早停、模型保存、学习率调度等回调与 Keras 的回调系统类似。使用 fastai提供最简洁的 APIlearn.fit(10)一行完成训练。理解底层原理即使使用高级 API也要理解zero_grad、backward、step的原理以便调试和优化。通过合理使用高级训练框架或自定义训练器可以在保持 PyTorch 灵活性的同时获得 Keras 般的简洁开发体验。
返回列表