使用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)
关键修改点说明
- 同步输入设备:在
get_features_image中,将image_processor生成的所有输入张量移到device,保证和模型运行设备一致; - 移除冗余预处理:删掉
transforms.ToTensor(),因为AutoImageProcessor已经完成了图像转张量、归一化等模型要求的预处理步骤,重复处理可能导致输入格式异常。
内容的提问来源于stack exchange,提问作者Mohamed Amine
相关产品推荐
相关产品推荐

