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

使用ResNet-50提取特征时遇RuntimeError:输入与权重类型不匹配

解决ResNet-50特征提取时的设备不匹配RuntimeError问题

错误原因

你碰到的RuntimeError: Input type (torch.FloatTensor) and weight type (torch.cuda.FloatTensor),核心问题是输入张量仍在CPU上,但模型已经部署到GPU运行。

看你的代码,虽然把转换后的图像张量移到了GPU,但在get_features_image函数里,image_processor处理生成的inputs张量并没有同步到GPU,导致输入数据和模型权重的设备不一致,触发了这个错误。

解决方案

修改get_features_image函数,将处理后的输入张量强制移到模型所在的设备(GPU/CPU),同时可以删掉多余的transforms.ToTensor()步骤(因为AutoImageProcessor已经包含了图像转张量的逻辑)。

修改后的完整代码

import torch
from transformers import AutoImageProcessor, ResNetModel
from tqdm import tqdm

# 检查GPU可用性,自动选择设备
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")

# 加载预训练的图像处理器和ResNet-50模型
image_processor = AutoImageProcessor.from_pretrained("microsoft/resnet-50")
model = ResNetModel.from_pretrained("microsoft/resnet-50").to(device)

# 定义图像特征提取函数
def get_features_image(image):
    inputs = image_processor(image, return_tensors="pt")
    # 将输入张量统一移到模型所在设备
    inputs = {key: tensor.to(device) for key, tensor in inputs.items()}
    with torch.no_grad():
        outputs = model(**inputs)
    # 返回最后一层隐藏状态作为特征
    return outputs.last_hidden_state

# 存储训练集图像特征的列表
chart_features_train = []

# 遍历训练数据集提取特征
for image in tqdm(train_chart_dataset):
    # 确保图像为RGB格式
    image_rgb = image.convert("RGB")
    # 直接传入RGB格式的PIL图像,无需额外转张量
    features = get_features_image(image_rgb)
    chart_features_train.append(features)

关键修改点说明

  1. 同步输入设备:在get_features_image中,将image_processor生成的所有输入张量移到device,保证和模型运行设备一致;
  2. 移除冗余预处理:删掉transforms.ToTensor(),因为AutoImageProcessor已经完成了图像转张量、归一化等模型要求的预处理步骤,重复处理可能导致输入格式异常。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.14 19:42:02