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

使用自定义Albumentations转换器报错:TypeError遇意外参数'cols'

自定义Albumentations转换器报错解决

问题场景

自定义了如下Albumentations转换器:

import albumentations as A
from albumentations.pytorch import ToTensorV2
import numpy as np
import cv2

class RandomTranslateWithReflect:
    """Translate image randomly
    Translate vertically and horizontally by n pixels where
    n is integer drawn uniformly independently for each axis
    from [-max_translation, max_translation].
    Fill the uncovered blank area with reflect padding.
    """

    def __init__(self, max_translation):
        self.max_translation = max_translation

    def __call__(self, old_image):
        xtranslation, ytranslation = np.random.randint(-self.max_translation,
                                                       self.max_translation + 1,
                                                       size=2)
        # Apply the translation using the Albumentations library
        transform = A.ShiftScaleRotate(shift_limit=(xtranslation / old_image.shape[1], ytranslation / old_image.shape[0]),
                                       scale_limit=0,
                                       rotate_limit=0,
                                       border_mode=cv2.BORDER_REFLECT,
                                       p=1)
        new_image = transform(image=old_image)["image"]
        return new_image

并通过以下代码将其加入Compose:

train_transform = {
    'cifar10': A.Compose([
        A.Lambda(image=RandomTranslateWithReflect(4)),
        A.HorizontalFlip(p=0.5),
        A.Normalize(*meanstd['cifar10']),
        ToTensorV2()
    ])
}

运行时出现如下错误:

mg1 = transform(image=img)["image"]
  File "/home/student/anaconda3/envs/few-shot/lib/python3.6/site-packages/albumentations/core/composition.py", line 205, in __call__
    data = t(**data)
  File "/home/student/anaconda3/envs/few-shot/lib/python3.6/site-packages/albumentations/core/transforms_interface.py", line 118, in __call__
    return self.apply_with_params(params, **kwargs)
  File "/home/student/anaconda3/envs/few-shot/lib/python3.6/site-packages/albumentations/core/transforms_interface.py", line 131, in apply_with_params
    res[key] = target_function(arg, **dict(params, **target_dependencies))
  File "/home/student/anaconda3/envs/few-shot/lib/python3.6/site-packages/albumentations/augmentations/transforms.py", line 1648, in apply
    return fn(img, **params)
TypeError: __call__() got an unexpected keyword argument 'cols'

报错原因

Albumentations的A.Lambda组件在调用传入的函数时,会自动传递cols、rows等额外关键字参数,但自定义类的__call__方法仅定义了old_image一个参数,无法接收这些额外参数,导致参数不匹配报错。

解决方法

方法1:修改自定义类的__call__方法

在__call__方法中添加**kwargs来接收所有额外的关键字参数,修改后的代码如下:

import albumentations as A
from albumentations.pytorch import ToTensorV2
import numpy as np
import cv2

class RandomTranslateWithReflect:
    """Translate image randomly
    Translate vertically and horizontally by n pixels where
    n is integer drawn uniformly independently for each axis
    from [-max_translation, max_translation].
    Fill the uncovered blank area with reflect padding.
    """

    def __init__(self, max_translation):
        self.max_translation = max_translation

    def __call__(self, old_image, **kwargs):
        xtranslation, ytranslation = np.random.randint(-self.max_translation,
                                                       self.max_translation + 1,
                                                       size=2)
        # Apply the translation using the Albumentations library
        transform = A.ShiftScaleRotate(shift_limit=(xtranslation / old_image.shape[1], ytranslation / old_image.shape[0]),
                                       scale_limit=0,
                                       rotate_limit=0,
                                       border_mode=cv2.BORDER_REFLECT,
                                       p=1)
        new_image = transform(image=old_image)["image"]
        return new_image

方法2:直接使用Albumentations原生组件

实际上自定义类的功能可以直接通过原生的ShiftScaleRotate实现,无需自定义类,简化后的Compose代码如下:

train_transform = {
    'cifar10': A.Compose([
        A.ShiftScaleRotate(
            shift_limit=4/32,  # CIFAR10图像尺寸为32x32,4像素对应比例为4/32
            scale_limit=0,
            rotate_limit=0,
            border_mode=cv2.BORDER_REFLECT,
            p=1
        ),
        A.HorizontalFlip(p=0.5),
        A.Normalize(*meanstd['cifar10']),
        ToTensorV2()
    ])
}

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.26 14:59:56