PyTorch张量中PIL图像打补丁后转存异常:呈现负片效果求助
图像添加指定图案后保存出现负片效果的解决方案
我需要实现一个函数,接收PIL图像作为输入,为图像添加指定图案后保存到磁盘。但保存后的图像呈现出类似负片的异常效果。以下是我编写的代码:
import torch import torch.nn.functional as F import torchvision.transforms as transforms from PIL import Image import random DEBUG = True def default_trigger_pattern() -> torch.Tensor: tensor = torch.tensor([[0,255,0], [255,0,255], [0,255,0]],dtype=torch.uint8) resized_tensor = tensor.repeat_interleave(3, dim=0).repeat_interleave(3, dim=1) return resized_tensor def mask_for_trigger(trigger, image, x, y) -> torch.Tensor: ''' 创建触发器位置的掩码 参数: trigger: 形状为 (CxHxW) 的张量 image: 形状为 (CxHxW) 或 (NxCxHxW) 的图像张量 x: 触发器在图像中左上角的垂直坐标(高度方向) y: 触发器在图像中左上角的水平坐标(宽度方向) ''' mask = torch.zeros_like(image,dtype=torch.uint8) height, width = trigger.shape[-2], trigger.shape[-1] if image.ndim == 4 or image.ndim == 3: mask[...,x:x+height,y:y+width] = 1 else: raise ValueError(f"图像维度必须是3 (CxHxW) 或4 (NxCxHxW),当前为 {image.ndim}") return mask class ImagePatcher: """ 为输入图像添加指定触发器图案的工具类 参数: trigger_pattern(Tensor): 要添加的触发器图案 img_size(int): 图像尺寸 location(str): 触发器位置,可选值:default(右下角)、center(中心)、random(随机)、topleft(左上角) rand(bool): 是否每次生成随机触发器图案 """ def __init__( self, trigger_pattern:torch.Tensor=default_trigger_pattern(), img_size=224, location='default', rand=False ): self.trigger_pattern = trigger_pattern self.location = location self.input_size=img_size self.rand = rand def __call__(self, x): ''' 输入为PIL图像或张量,返回添加触发器后的PIL图像或张量 ''' # 生成随机触发器(如果开启) if self.rand: self.trigger_pattern = torch.randint_like(self.trigger_pattern, 2) * 255 trigger_channel = self.trigger_pattern.size(dim=-3) if self.trigger_pattern.ndim >=3 else 0 channel_count = 0 is_image = False # 处理输入类型转换 if isinstance(x, Image.Image): is_image = True x = transforms.PILToTensor()(x) channel_count = x.size(dim=0) elif isinstance(x, np.ndarray): x = torch.tensor(x) channel_count = x.size(dim=0) else: channel_count = x.size(dim=0) assert(channel_count > 0) width = x.shape[-1] height = x.shape[-2] trigger_height = self.trigger_pattern.size(dim=-2) trigger_width = self.trigger_pattern.size(dim=-1) # 设置触发器位置 if self.location == 'default': self.start_loc = (width - trigger_width, height - trigger_height) elif self.location == 'center': self.start_loc = (int((width - trigger_width) / 2), int((height - trigger_height) / 2)) elif self.location == 'random': x_random = random.randint(0, width-trigger_width) y_random = random.randint(0, height-trigger_height) self.start_loc = (x_random, y_random) elif self.location == 'topleft': self.start_loc = (0,0) pad_left = self.start_loc[0] pad_top = self.start_loc[1] pad_right_patch = width - self.start_loc[0] - trigger_width pad_bottom_patch = height - self.start_loc[1] - trigger_height # 扩展触发器到图像尺寸 patch_expanded = F.pad(self.trigger_pattern, (pad_left, pad_right_patch, pad_top, pad_bottom_patch), value=0) # 匹配图像通道数 if trigger_channel == 0: patch = patch_expanded.unsqueeze(dim=0).repeat([channel_count,1,1]) else : patch = patch_expanded mask = mask_for_trigger(self.trigger_pattern, x, self.start_loc[1],self.start_loc[0]) x = (1 - mask) * x + mask * patch if DEBUG: im = transforms.ToPILImage()(x) im.save("debug01.png") return transforms.ToPILImage()(x) if is_image else x
问题原因
负片效果的核心是uint8类型张量的计算溢出:
transforms.PILToTensor()将PIL图像转为uint8类型张量(像素值范围0-255)mask_for_trigger生成的是uint8类型的0/1掩码,计算1 - mask时,1会被强制转为uint8类型,此时1 - 0 = 255、1 - 1 = 0- 最终计算
(1 - mask) * x时,原图像区域的像素值被乘以255,直接导致负片效果
修复方案
1. 修改掩码生成函数,使用浮点类型计算
将掩码改为float32类型,避免uint8下的错误计算:
def mask_for_trigger(trigger, image, x, y) -> torch.Tensor: ''' 创建触发器位置的掩码 参数: trigger: 形状为 (CxHxW) 的张量 image: 形状为 (CxHxW) 或 (NxCxHxW) 的图像张量 x: 触发器在图像中左上角的垂直坐标(高度方向) y: 触发器在图像中左上角的水平坐标(宽度方向) ''' # 使用float32类型,避免uint8计算溢出 mask = torch.zeros_like(image, dtype=torch.float32) height, width = trigger.shape[-2], trigger.shape[-1] if image.ndim == 4 or image.ndim == 3: mask[..., x:x+height, y:y+width] = 1.0 else: raise ValueError(f"图像维度必须是3 (CxHxW) 或4 (NxCxHxW),当前为 {image.ndim}") return mask
2. 修改图像混合计算逻辑,确保数值范围正确
在混合图像与触发器前,将张量转为浮点类型,计算完成后再转回uint8并钳制数值范围:
# 替换原代码中的x = (1 - mask) * x + mask * patch x = x.to(torch.float32) patch = patch.to(torch.float32) x = (1 - mask) * x + mask * patch # 转回uint8并确保像素值在0-255之间 x = x.clamp(0, 255).to(torch.uint8)
3. 优化触发器生成(可选)
明确触发器的通道维度,避免后续处理歧义:
def default_trigger_pattern() -> torch.Tensor: tensor = torch.tensor([[0,255,0], [255,0,255], [0,255,0]], dtype=torch.uint8) resized_tensor = tensor.repeat_interleave(3, dim=0).repeat_interleave(3, dim=1) # 添加通道维度,变为(1, H, W),方便后续扩展为RGB通道 return resized_tensor.unsqueeze(0)
测试示例
# 加载测试图像 img = Image.open("test.jpg").convert("RGB") # 创建补丁器 patcher = ImagePatcher(location='default') # 应用补丁 patched_img = patcher(img) # 保存处理后的图像 patched_img.save("patched_test.jpg")
内容的提问来源于stack exchange,提问作者祝望舒
相关产品推荐
相关产品推荐

