如何加载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
相关产品推荐
相关产品推荐

