如何将Keras中的模型训练与评估代码转换为PyTorch实现?
将Keras训练代码转换为PyTorch实现
PyTorch没有Keras那样封装好的model.fit()高层API,需要手动编写训练循环。以下是对应你提供的Keras代码的PyTorch实现步骤:
1. 准备数据加载器(对应Keras的training_set/test_set)
假设你已经定义了PyTorch的Dataset,需要用DataLoader封装,设置batch_size=16(对应你Keras代码中的批量大小),训练集开启打乱:
from torch.utils.data import DataLoader # 假设train_dataset和test_dataset是你定义好的PyTorch Dataset train_loader = DataLoader(train_dataset, batch_size=16, shuffle=True) test_loader = DataLoader(test_dataset, batch_size=16, shuffle=False)
2. 定义核心组件
初始化模型、损失函数、优化器,以及对应Keras回调的早停和学习率调整器:
import torch import torch.nn as nn import torch.optim as optim from torch.optim.lr_scheduler import ReduceLROnPlateau # 假设你已经定义好了PyTorch模型 model = YourModel() device = torch.device("cuda" if torch.cuda.is_available() else "cpu") model.to(device) # 损失函数(根据任务选择,分类任务常用CrossEntropyLoss) criterion = nn.CrossEntropyLoss() # 优化器(示例用Adam) optimizer = optim.Adam(model.parameters(), lr=1e-3) # 对应Keras的ReduceLROnPlateau回调 lr_scheduler = ReduceLROnPlateau(optimizer, mode='max', factor=0.1, patience=5, verbose=1) # 自定义早停逻辑(PyTorch无内置早停,手动实现) class EarlyStopping: def __init__(self, patience=5, verbose=0, delta=0, path='checkpoint.pt'): self.patience = patience self.verbose = verbose self.counter = 0 self.best_score = None self.early_stop = False self.val_loss_min = float('inf') self.delta = delta self.path = path def __call__(self, val_loss, model): score = -val_loss if self.best_score is None: self.best_score = score self.save_checkpoint(val_loss, model) elif score < self.best_score + self.delta: self.counter += 1 if self.verbose > 0: print(f'EarlyStopping counter: {self.counter} out of {self.patience}') if self.counter >= self.patience: self.early_stop = True else: self.best_score = score self.save_checkpoint(val_loss, model) self.counter = 0 def save_checkpoint(self, val_loss, model): if self.verbose > 0: print(f'Validation loss decreased ({self.val_loss_min:.6f} --> {val_loss:.6f}). Saving model ...') torch.save(model.state_dict(), self.path) self.val_loss_min = val_loss # 初始化早停实例 early_stop = EarlyStopping(patience=5, verbose=1)
3. 编写训练循环(对应Keras的model.fit())
epochs = 100 # 用于记录训练过程的指标 history = {'train_loss': [], 'train_acc': [], 'val_loss': [], 'val_acc': []} for epoch in range(epochs): # 训练阶段 model.train() train_loss = 0.0 train_correct = 0 train_total = 0 for inputs, labels in train_loader: inputs, labels = inputs.to(device), labels.to(device) # 前向传播 outputs = model(inputs) loss = criterion(outputs, labels) # 反向传播与优化 optimizer.zero_grad() loss.backward() optimizer.step() # 统计训练损失与准确率 train_loss += loss.item() * inputs.size(0) _, predicted = torch.max(outputs.data, 1) train_total += labels.size(0) train_correct += (predicted == labels).sum().item() # 计算训练集平均指标 avg_train_loss = train_loss / train_total avg_train_acc = train_correct / train_total history['train_loss'].append(avg_train_loss) history['train_acc'].append(avg_train_acc) # 验证阶段 model.eval() val_loss = 0.0 val_correct = 0 val_total = 0 with torch.no_grad(): # 禁用梯度计算,节省资源 for inputs, labels in test_loader: inputs, labels = inputs.to(device), labels.to(device) outputs = model(inputs) loss = criterion(outputs, labels) val_loss += loss.item() * inputs.size(0) _, predicted = torch.max(outputs.data, 1) val_total += labels.size(0) val_correct += (predicted == labels).sum().item() # 计算验证集平均指标 avg_val_loss = val_loss / val_total avg_val_acc = val_correct / val_total history['val_loss'].append(avg_val_loss) history['val_acc'].append(avg_val_acc) # 更新学习率 lr_scheduler.step(avg_val_acc) # 若关注损失则传入avg_val_loss,需将scheduler的mode设为'min' # 检查早停条件 early_stop(avg_val_loss, model) if early_stop.early_stop: print("Early stopping") break # 打印当前epoch结果 print(f'Epoch {epoch+1}/{epochs}') print(f'Train Loss: {avg_train_loss:.4f} Acc: {avg_train_acc:.4f}') print(f'Val Loss: {avg_val_loss:.4f} Acc: {avg_val_acc:.4f}\n') # 加载训练过程中保存的最佳模型(可选) model.load_state_dict(torch.load('checkpoint.pt'))
4. 模型评估(对应Keras的model.evaluate())
model.eval() test_loss = 0.0 test_correct = 0 test_total = 0 with torch.no_grad(): for inputs, labels in test_loader: inputs, labels = inputs.to(device), labels.to(device) outputs = model(inputs) loss = criterion(outputs, labels) test_loss += loss.item() * inputs.size(0) _, predicted = torch.max(outputs.data, 1) test_total += labels.size(0) test_correct += (predicted == labels).sum().item() avg_test_loss = test_loss / test_total avg_test_acc = test_correct / test_total print(f'loss={avg_test_loss:.4f}, accuracy={avg_test_acc:.4f}')
关键说明
- PyTorch需手动切换模型模式:
model.train()开启训练相关层(如Dropout、BatchNorm),model.eval()关闭这些行为用于验证/测试。 - 梯度计算需手动管理:
optimizer.zero_grad()清零梯度、loss.backward()反向传播、optimizer.step()更新参数。 - 早停需自定义实现,学习率调整可直接使用PyTorch内置的
ReduceLROnPlateau。 - 数据加载依赖
DataLoader,批量大小需与Keras代码保持一致(16)。
内容的提问来源于stack exchange,提问作者seni
相关产品推荐
相关产品推荐

