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

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,提问作者祝望舒

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.01 10:44:55