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

如何加载Fastai训练的模型并实现单张图像的预测功能

模型无predict接口问题

你当前构造的是原生PyTorch的nn.Sequential模型实例,predict是fastai框架中Learner类的专属方法,原生PyTorch模型本身不提供该接口。要使用predict方法,你需要构造和训练时结构一致的fastai Learner实例加载权重。

修正后的加载代码如下:

from fastai import * 
from fastai.vision import *
import torch

# 定义和训练时完全对齐的图像变换规则
tfms = get_transforms(do_flip=True, flip_vert=True, max_rotate=10.0, max_zoom=1.1, max_lighting=0.2, max_warp=0.2)
# 构造匹配训练配置的dummy databunch,分类数、图像尺寸、归一化规则都要和训练时一致
dummy_path = Path('./')
data = ImageDataBunch.single_from_classes(dummy_path, classes=['NORMAL', 'CNV', 'DME', 'DRUSEN'], ds_tfms=tfms, size=224).normalize(imagenet_stats)
# 构造和训练时结构一致的Learner实例
learn = cnn_learner(data, models.resnet18, metrics=accuracy)
# 加载已保存的模型权重,不需要加.pth后缀
learn.load('/content/gdrive/MyDrive/Data Exports/35k data/stage-1')

加载完成后即可直接调用learn.predict()方法做预测。


单张图像预测失败问题

失败有两个核心原因:

  • 训练时fastai会自动对输入做resize、标准化、通道调整等预处理,直接传入numpy数组没有对齐预处理逻辑,会导致模型输入分布和训练时不匹配
  • PyTorch模型要求输入维度为[batch_size, channels, height, width],单张图直接传入会缺少batch维度,触发维度不匹配错误

如果使用上述fastai Learner方案,单张图预测代码如下:

# 用fastai自带的open_image读取图像,会自动对齐预处理逻辑
img = open_image('你的单张图像本地路径')
pred_class, pred_idx, outputs = learn.predict(img)
print(f'预测类别:{pred_class},置信度:{outputs.max().item():.4f}')

如果你坚持使用原生PyTorch模型做预测,参考以下代码:

from fastai import * 
from fastai.vision import *
import torch
from PIL import Image
import torchvision.transforms as transforms

# 初始化模型结构
body = create_body(models.resnet18, True, None)
data_classes = 4
nf = callbacks.hooks.num_features_model(body) * 2
head = create_head(nf, data_classes, None, ps=0.5, bn_final=False)
model = nn.Sequential(body, head)

# 加载权重
model.load_state_dict(torch.load('/content/gdrive/MyDrive/Data Exports/35k data/stage-1.pth'))
model.eval()

# 定义和训练时完全对齐的预处理规则
preprocess = transforms.Compose([
    transforms.Resize((224,224)),
    transforms.ToTensor(),
    transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225])
])

# 单张图预处理
img = Image.open('你的单张图像本地路径').convert('RGB')
img_tensor = preprocess(img)
img_tensor = img_tensor.unsqueeze(0) # 新增batch维度,适配PyTorch输入要求

# 预测
with torch.no_grad():
    outputs = model(img_tensor)
pred_idx = outputs.argmax(dim=1).item()
class_map = ['NORMAL', 'CNV', 'DME', 'DRUSEN']
print(f'预测类别:{class_map[pred_idx]},置信度:{torch.softmax(outputs, dim=1).max().item():.4f}')

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.10.05 14:15:02