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

类与函数方式加载PyTorch模型的速度差异原因问询

类实现与函数实现的特征提取速度差异原因分析

你遇到的速度差异核心原因非常明确:函数实现中每次调用特征提取都会重新加载整个模型,而类实现仅在初始化时加载一次模型,后续复用已加载的实例。

具体差异拆解

  • 类实现逻辑:模型在__init__方法中完成网络初始化、权重加载、结构裁剪(去掉顶层全连接层)、设置eval模式这些操作,之后每次调用extract_image_features只执行图像预处理和模型前向传播,这部分是真正的特征提取耗时。
  • 函数实现逻辑:在extract_image_features内部,每次都调用load_feature_extractor,这意味着每次都要重复执行:
    • 初始化ResNet50网络结构
    • 读取并解析.pth.tar权重文件
    • 处理权重字典(移除module.前缀)并加载到模型
    • 构建特征提取子网络并设置eval模式
      这些步骤的开销远大于单次前向传播的时间,直接导致整体耗时比类实现高3-4倍。

修正后的函数实现

将模型加载与特征提取分离,确保模型仅加载一次,修正后耗时会和类实现基本一致:

from typing import List, Union
import time
import numpy as np
import torch
import torchvision.models as models
from PIL import Image
from torch.autograd import Variable
from torchvision import transforms as trn
import cv2

def load_feature_extractor(model_path: str):
    """加载预训练ResNet50特征提取模型,仅需调用一次"""
    pretrained_model = models.__dict__['resnet50'](num_classes=365)
    checkpoint = torch.load(model_path, map_location=lambda storage, loc: storage)
    state_dict = {str.replace(k, 'module.', ''): v for k, v in checkpoint['state_dict'].items()}
    pretrained_model.load_state_dict(state_dict)
    res50_conv = torch.nn.Sequential(*list(pretrained_model.children())[:-1])
    res50_conv.eval()
    return res50_conv

def extract_image_features(img: np.ndarray, model) -> np.ndarray:
    """复用已加载的模型提取图像特征"""
    t_start = time.time()
    centre_crop = trn.Compose([
        trn.Resize((256, 256)),
        trn.CenterCrop(224),
        trn.ToTensor(),
        trn.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225])
    ])
    img = Image.fromarray(np.uint8(img)).convert('RGB')
    preprocessed_img = Variable(centre_crop(img).unsqueeze(0))

    places_features = model.forward(preprocessed_img).detach().numpy().flatten()
    t_end = time.time()
    print(t_end-t_start)
    return places_features

# 全局仅加载一次模型
feature_extractor_model_path = '/home/michael/Gitlab/backend/packages/izirecord-api-cv/izirecord/cv/models/resnet50_places365.pth.tar'
model = load_feature_extractor(feature_extractor_model_path)

# 测试特征提取
image = cv2.imread('/home/michael/Documents/thumbs/.png00001.png')
extract_image_features(image, model)

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.11 16:33:27