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

使用PRIME数据增强遇NameError:pil_mask未定义的技术咨询

问题分析与解决:PRIME数据增强中pil_mask未定义错误

问题背景

用户定义了以下图像增强操作函数:

def identity(pil_img, pil_mask, _):
    return pil_img, pil_mask

def autocontrast(pil_img, pil_mask, _):
    return ImageOps.autocontrast(pil_img), pil_mask

def equalize(pil_img, pil_mask, _):
    return ImageOps.equalize(pil_img), pil_mask

def rotate(pil_img, pil_mask, level):
    degrees = int_parameter(level, min_max_vals.rotate.max)
    if np.random.uniform() > 0.5:
        degrees = -degrees
    return pil_img.rotate(degrees, resample=Image.BILINEAR), pil_mask.rotate(degrees, resample=Image.BILINEAR)

在使用PRIME数据增强时,运行触发以下错误:

NameError: name 'pil_mask' is not defined

对应的PRIME实现代码:

augmentations = [
    (identity, 1.0)
    ]
class PRIMEAugModule(torch.nn.Module):
    def __init__(self):
        super().__init__()
        self.augmentations = augmentations
        self.num_transforms = len(augmentations)

    def forward(self, x, mask_t):
        x_tensor = torch.from_numpy(x)
        aug_x = torch.zeros_like(x_tensor)
        for i in range(self.num_transforms):
            fn, weight = self.augmentations[i]
            if fn.__name__ == 'identity':
                aug_x += fn(x_tensor, pil_mask, _) * mask_t[:, i] * weight
            else:
                aug_x += fn(x_tensor, pil_mask) * mask_t[:, i] * weight
        return aug_x

用户疑问:应该在何处以及如何定义pil_mask?


解决方案

1. 明确pil_mask的本质与来源

你的增强函数设计为接收PIL格式的图像和掩码,但当前forward方法存在两个核心问题:

  • 未传入掩码数据:forward只接收了x(图像)和mask_t(PRIME的权重掩码),没有传入任务对应的掩码输入
  • 数据类型不匹配:x被转成了tensor,但增强函数预期的是PIL图像

2. 调整forward方法的输入与数据转换

如果你的任务需要处理掩码,修改forward方法,新增掩码参数并完成格式转换:

class PRIMEAugModule(torch.nn.Module):
    def __init__(self):
        super().__init__()
        self.augmentations = augmentations
        self.num_transforms = len(augmentations)

    def forward(self, x, task_mask, mask_t):
        # 将numpy格式的图像转成PIL格式
        pil_img = Image.fromarray(x)
        # 将任务掩码转成PIL格式(根据掩码实际数据类型调整,比如单通道灰度图)
        pil_mask = Image.fromarray(task_mask.astype('uint8'))
        
        x_tensor = torch.from_numpy(x)
        aug_x = torch.zeros_like(x_tensor)
        
        for i in range(self.num_transforms):
            fn, weight = self.augmentations[i]
            # 统一按函数定义传参,第三个参数按需传入(比如rotate需要level,可随机生成)
            augmented_img, _ = fn(pil_img, pil_mask, None)
            # 将增强后的PIL图像转回tensor
            augmented_tensor = torch.from_numpy(np.array(augmented_img))
            aug_x += augmented_tensor * mask_t[:, i] * weight
        
        return aug_x

3. 修正增强函数的调用参数

原代码中存在函数参数不匹配问题(比如部分调用只传2个参数,但函数需要3个),必须统一按函数定义传参,第三个参数如果不需要可以传None;如果是rotate这类需要level的函数,需要额外生成对应level值(比如从预设范围随机采样)。

4. 无掩码场景的简化方案

如果你的任务不需要处理掩码,直接修改增强函数去掉pil_mask参数,减少冗余:

def identity(pil_img, _):
    return pil_img

def autocontrast(pil_img, _):
    return ImageOps.autocontrast(pil_img)

# 对应调整PRIME模块内的调用逻辑,不再传入掩码参数

内容的提问来源于stack exchange,提问作者Green_Cat

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.27 23:03:23