PyTorch ResNet152模型推理报错:期望Byte类型却发现Float类型
解决ResNet152训练时
RuntimeError: expected scalar type Byte but found Float的问题 问题根源
输入的images张量数据类型是Byte(uint8,像素值0-255),但PyTorch卷积层及预训练ResNet模型要求输入为浮点类型(Float32/Float64),同时预训练ResNet必须使用匹配ImageNet的标准化输入格式才能正常运行。
修复方案
1. 规范数据预处理(推荐在数据加载阶段处理)
如果使用torchvision.datasets和DataLoader,直接在transforms中添加类型转换与标准化:
from torchvision import transforms # 构建符合预训练ResNet要求的预处理流程 transform = transforms.Compose([ transforms.ToTensor(), # 将Byte格式图像转为Float32,同时把像素值归一到[0,1]区间 transforms.Normalize( mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225] ) # 用ImageNet数据集的均值和标准差做标准化 ]) # 应用到数据集加载 train_dataset = YourCustomDataset(..., transform=transform) train_loader = torch.utils.data.DataLoader(train_dataset, batch_size=32, shuffle=True)
2. 临时手动处理输入(适用于无法修改数据加载流程的场景)
如果你的images已经是Byte张量,在传入模型前执行以下转换:
# 转换为Float类型并归一到[0,1] images = images.float() / 255.0 # 加载ImageNet的标准化参数,注意要和输入张量同设备 mean = torch.tensor([0.485, 0.456, 0.406]).view(1, 3, 1, 1).to(images.device) std = torch.tensor([0.229, 0.224, 0.225]).view(1, 3, 1, 1).to(images.device) # 执行标准化 images = (images - mean) / std # 再传入模型计算 output = model(images)
补充说明
你的ResNet152初始化代码没有问题,权重加载和冻结参数、替换全连接层的操作都是正确的,问题完全出在输入数据的类型和预处理不匹配上。
内容的提问来源于stack exchange,提问作者Codeman
相关产品推荐
相关产品推荐

