You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何将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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.07.25 04:23:09