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

Albumentations中自定义掩码依赖增强的参数传递问题

自定义Albumentations背景替换增强的问题与解决

问题场景

需要实现随机替换图像背景但保留掩码区域像素的增强效果,基于Albumentations的DualTransform自定义类时,无法将mask传递给内部的apply函数,运行时出现两种错误:

  • 保留get_params时:AttributeError: 'RandomBackground' object has no attribute 'mask'
  • 注释get_params时:TypeError: RandomBackground.apply() missing 1 required positional argument: 'mask'

错误原因

对Albumentations的DualTransform工作逻辑理解偏差:

  • DualTransform的默认apply方法仅接收图像和通用参数,不会自动传入mask
  • 自定义get_params中试图访问self.mask是错误的,类实例不会自动存储外部传入的mask数据

修正后的完整代码

from numpy.random import default_rng
import matplotlib.pyplot as plt
import numpy as np
from albumentations.core.transforms_interface import DualTransform

def show(arr, title = None):
    '''Quick helper function to reliably display arrays as images'''
    if len(arr.shape) == 3:
        if arr.shape[0] == 3:
            plt.imshow(np.moveaxis(arr, 0, -1))
        elif arr.shape[2] == 3:
            plt.imshow(arr)
        elif arr.shape[0] == 1:
            plt.imshow(arr.squeeze())
    else:
        plt.imshow(arr)
    plt.title(title)
    plt.show()
    
def apply_random_background(rgb: np.ndarray, mask: np.ndarray, **kwargs) -> np.ndarray:
    '''Function that does the conversion of the background to a uniform random colour'''
    rng = default_rng()
    rand_colour = rng.integers(low = 0, high = 255, size = [3,1,1]).astype(np.uint8)
    random_bkgnd = np.ones_like(rgb, dtype = np.uint8) * rand_colour

    mask3d = np.stack([mask]*3) 
    nrgb = rgb.copy()
    nrgb[~mask3d] = random_bkgnd[~mask3d]

    return nrgb


class RandomBackground(DualTransform):
    '''Class that extends Albumentations.DualTransform with the hope of applying the
    apply_random_background function'''
    def __init__(self, always_apply=False, p=1.0):
        super().__init__(always_apply, p)

    def apply_with_params(self, params, force_apply=False, **kwargs):
        # 从kwargs中获取mask,传递给apply函数
        img = kwargs['image']
        mask = kwargs['mask']
        transformed_img = self.apply(img, mask=mask, **params)
        transformed_mask = self.apply_to_mask(mask, **params)
        return {'image': transformed_img, 'mask': transformed_mask}

    def apply(self, img: np.ndarray, mask: np.ndarray, **params) -> np.ndarray:
        return apply_random_background(img, mask)

    def apply_to_mask(self, img: np.ndarray, **params) -> np.ndarray:
        return img
    
    
def create_random_image_and_mask(height = 100, width = 100, target_box_coords = [40,60]):
    '''Function to create demo "image" and mask, creating a grey square in the middle
    as the area we're interested in amongst a sea of random colours'''
    rng = default_rng()
    height, width = 100, 100
    img = rng.integers(low = 0, high = 255, size = [3,height,width]).astype(np.uint8)
    img[:,target_box_coords[0]: target_box_coords[1], target_box_coords[0]:target_box_coords[1]] = 125
    mask = np.zeros(shape = [height, width], dtype = np.uint8)
    mask[target_box_coords[0]: target_box_coords[1], target_box_coords[0]:target_box_coords[1]] = 1
    mask= mask.astype(bool)
    
    return img, mask

img, orig_mask = create_random_image_and_mask()
show(img, 'Original Image')
show(orig_mask, 'Original Mask')

newrgb = apply_random_background(img, orig_mask)
show(newrgb, 'New Image from Function')

randbg = RandomBackground(p=1)
rand = randbg(image = img, mask = orig_mask)
show(rand['image'], 'New Image from Class')

关键修正点

  1. 重写apply_with_params方法:这个方法是DualTransform中处理多数据(图像、掩码等)的入口,可以直接从kwargs中获取传入的mask,再传递给apply函数
  2. 移除错误的get_params方法:不需要手动构建mask参数,避免访问不存在的实例属性
  3. 保持apply_to_mask不变:因为我们只需要修改图像背景,掩码不需要做任何变换

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.23 16:33:16