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

如何基于PyTorch调用FastAI训练的图像分类模型执行推理?

嘿,我来帮你捋清楚怎么用PyTorch加载fastai训练的模型做推理,以及对应的自定义Dataset该怎么写~

用PyTorch加载fastai训练的模型做推理

第一步:提取fastai训练的PyTorch模型

fastai训练的Learner虽然带有框架封装,但底层就是标准的PyTorch模型,我们可以直接提取出来:

from fastai.vision.all import *

# 加载训练好的fastai Learner
learn = load_learner('path/to/your/trained_model.pkl')
# 拿到底层的PyTorch模型
pytorch_model = learn.model
# 切换到评估模式(必须!否则BatchNorm、Dropout等层会出问题)
pytorch_model.eval()

第二步:编写自定义PyTorch Dataset

推理时的核心是复现训练时的图像预处理逻辑,这样模型才能输出正确结果。下面是对应实现:

from torch.utils.data import Dataset, DataLoader
from PIL import Image
import torchvision.transforms as transforms

class ImageInferenceDataset(Dataset):
    def __init__(self, image_paths, img_size=(224,224), mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]):
        self.image_paths = image_paths
        # 完全对齐fastai训练时的预处理流程(仅保留确定性操作)
        self.transform = transforms.Compose([
            transforms.Resize(img_size),
            transforms.ToTensor(),
            transforms.Normalize(mean=mean, std=std)
        ])
        
    def __len__(self):
        return len(self.image_paths)
    
    def __getitem__(self, idx):
        img_path = self.image_paths[idx]
        # 读取图像并转为RGB(和训练时保持一致)
        img = Image.open(img_path).convert('RGB')
        # 应用预处理得到张量
        img_tensor = self.transform(img)
        return img_tensor, img_path  # 返回张量和路径,方便后续对应结果

第三步:执行推理流程

有了模型和Dataset,就可以用PyTorch原生的方式做推理了:

import torch

# 你的测试图像路径列表
test_image_paths = ['test_cat.jpg', 'test_dog.jpg', 'test_bird.jpg']

# 创建Dataset和DataLoader
dataset = ImageInferenceDataset(test_image_paths)
dataloader = DataLoader(dataset, batch_size=4, shuffle=False)

# 关闭梯度计算,节省内存和加速推理
with torch.no_grad():
    for batch in dataloader:
        imgs, paths = batch
        # 前向传播得到输出
        outputs = pytorch_model(imgs)
        # 获取概率最大的类别索引
        preds = torch.argmax(outputs, dim=1)
        # 用fastai的类别词汇表把索引转成实际类别名称
        class_names = learn.dls.vocab
        for path, pred in zip(paths, preds):
            print(f"图像 {path} 的预测类别是: {class_names[pred]}")

关键注意事项

  • 预处理必须完全对齐:如果训练时用了自定义的图像尺寸、均值/标准差,一定要在Dataset里同步修改,否则推理结果会完全错误
  • 不要忘记eval模式:模型在训练时是train()模式,推理前必须切换到eval(),否则BatchNorm和Dropout层的行为会不符合预期
  • 类别映射要准确:fastai的learn.dls.vocab保存了训练时的类别顺序,直接用它就能把预测索引转成可读的类别名称

内容的提问来源于stack exchange,提问作者user1274878

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.20 08:58:24