使用自定义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
相关产品推荐
相关产品推荐

