使用PyTorch训练ResNet50时ImageNet归一化异常问题排查
问题描述
用ResNet50训练ImageNet数据集时,遇到归一化操作导致图像显示异常的问题:
无归一化的变换(显示正常)
train_transforms = transforms.Compose([ transforms.Resize((224, 224), antialias=True), transforms.RandomCrop(180), transforms.Resize((224, 224), antialias=True), transforms.RandomHorizontalFlip(), transforms.RandomVerticalFlip(), transforms.ColorJitter(brightness=0.2, contrast=0.2, saturation=0.2, hue=0.15), transforms.RandomApply([transforms.GaussianBlur(3, sigma=(0.1, 2.0))], p=0.5), transforms.ToTensor() ])
使用上述变换时,增强后的图像(比如华夫饼机)显示正常,属于轻度增强。
添加归一化后的变换(图像显示异常)
train_transforms = transforms.Compose([ transforms.Resize((224, 224), antialias=True), transforms.RandomCrop(180), transforms.Resize((224, 224), antialias=True), transforms.RandomHorizontalFlip(), transforms.RandomVerticalFlip(), transforms.ColorJitter(brightness=0.2, contrast=0.2, saturation=0.2, hue=0.15), transforms.RandomApply([transforms.GaussianBlur(3, sigma=(0.1, 2.0))], p=0.5), transforms.ToTensor(), transforms.Normalize(mean, std) ])
添加transforms.Normalize后,华夫饼机图像显示异常。
错误调整归一化顺序的尝试
尝试调整归一化在变换链中的位置,要么报错(归一化仅适用于张量),要么图像损坏:
train_transforms = transforms.Compose([ transforms.ToTensor(), transforms.Normalize(mean, std), transforms.Resize((224, 224), antialias=True), transforms.RandomCrop(180), transforms.Resize((224, 224), antialias=True), transforms.RandomHorizontalFlip(), transforms.RandomVerticalFlip(), transforms.ColorJitter(brightness=0.2, contrast=0.2, saturation=0.2, hue=0.15), transforms.RandomApply([transforms.GaussianBlur(3, sigma=(0.1, 2.0))], p=0.5), ])
补充信息
使用的均值和标准差为ImageNet标准值:
mean = [0.485, 0.456, 0.406] std = [0.229, 0.224, 0.225]
保存图像的代码
# Helper function to convert tensor to PIL Image to_pil = ToPILImage() idx = indices[i].item() label = labels[i] # Define the paths to save original and augmented images original_img_path = os.path.join(debug_dir, f"{global_id_counter}_{interpretable_label[i]}_original.png") augmented_img_path = os.path.join(debug_dir, f"{global_id_counter}_{interpretable_label[i]}_augmented.png") # Convert tensor to PIL Image and save original_pil_img = to_pil(original_img[i].cpu()) # Convert tensor to PIL Image original_pil_img.save(original_img_path) # Get the augmented image from inputs augmented_img_np = inputs_np[i].transpose(1, 2, 0) # Convert from float to uint8 augmented_img_np = (augmented_img_np * 255).astype(np.uint8) # Convert to PIL image and save augmented_img = Image.fromarray(augmented_img_np) augmented_img.save(augmented_img_path)
请问问题出在哪里,该如何解决?
问题原因与解决方法
核心原因
归一化操作会把张量的像素值从[0,1]范围转换为以均值为中心、标准差为缩放单位的分布(部分值会落在[-2,2]甚至更宽区间)。你保存图像时直接将归一化后的张量乘以255转成uint8,负数会被截断为0,超出1的部分被截断为255,最终导致图像颜色失真、对比度异常。
另外,你最初的归一化顺序是正确的:必须在transforms.ToTensor()之后、所有图像增强操作之后(大部分增强操作只支持PIL图像或[0,1]范围的张量),调整顺序的尝试违反了这个规则,自然会报错或损坏图像。
解决方法
保存归一化后的图像前,需要先反归一化,将像素值还原回[0,1]范围,再转成uint8:
修改保存增强图像的代码:
# Get the augmented image from inputs augmented_img_np = inputs_np[i].transpose(1, 2, 0) # 反归一化:先乘标准差再加均值,还原到[0,1]范围 mean_np = np.array(mean) std_np = np.array(std) augmented_img_np = augmented_img_np * std_np + mean_np # 裁剪到[0,1]范围,避免数值溢出导致的异常 augmented_img_np = np.clip(augmented_img_np, 0, 1) # Convert from float to uint8 augmented_img_np = (augmented_img_np * 255).astype(np.uint8) # Convert to PIL image and save augmented_img = Image.fromarray(augmented_img_np) augmented_img.save(augmented_img_path)
额外说明
- 训练时的归一化流程是正确的:
PIL图像 -> 增强操作 -> ToTensor() -> Normalize(),这个顺序不会影响模型训练,只是保存图像时需要反归一化才能正常显示。 - 不要调整归一化的顺序,
Resize、ColorJitter等操作需要处理像素值在合理范围(0-255或0-1)的图像,归一化后的张量不符合这些操作的输入要求。
内容的提问来源于stack exchange,提问作者Chuck
相关产品推荐
相关产品推荐

