如何在PyTorch+Albumentations中向变换函数传递图像原始尺寸
问题解决:Albumentations动态获取图像原始尺寸参数
问题根源
你的核心问题在于变换初始化时机错误:
- 在
train.py中,transforms.medium_transforms()是在数据集初始化前就执行的,此时还没有加载任何图像,自然无法传入original_height和original_width参数,导致params为空。 - 而
A.Compose对象在调用时(即self.transforms(**result))才会接收图像相关参数,但此时你之前定义的变换已经提前构建好了,无法再动态获取尺寸。
另外注意:dataset.py中cv2.imread的调用有误——第二个参数是读取模式(如cv2.IMREAD_COLOR),不是颜色转换常量,会导致图像读取异常。
解决方案:用Lambda变换动态获取尺寸
利用Albumentations的Lambda变换,让尺寸参数在处理图像时才动态传入,而不是提前构建变换。
修改后的代码
1. dataset.py(修正图像读取错误)
from torch.utils.data import Dataset import cv2 class SegmentationDataset(Dataset): def __init__(self, imagePaths, maskPaths, transforms): self.imagePaths = imagePaths self.maskPaths = maskPaths self.transforms = transforms def __getitem__(self, idx): # 修正:先读取BGR格式,再转换为RGB image = cv2.imread(self.imagePaths[idx], cv2.IMREAD_COLOR) image = cv2.cvtColor(image, cv2.COLOR_BGR2RGB) mask = cv2.imread(self.maskPaths[idx], 0) result = { "image": image, "mask": mask, "original_height": image.shape[0], "original_width": image.shape[1] } if self.transforms is not None: result = self.transforms(**result) return result
2. transforms.py(改用Lambda动态构建变换)
import albumentations as A from . import config def medium_transforms(): # 定义动态变换逻辑,处理时自动获取原始尺寸 def apply_dynamic_transform(image, mask, original_height, original_width, **kwargs): # 动态构建需要原始尺寸的变换 size_transform = A.OneOf([ A.RandomSizedCrop( min_max_height=(50, 101), height=original_height, width=original_width, p=config.MEDIUM_TRANSFORMS_PROBABILITY ), A.PadIfNeeded( min_height=original_height, min_width=original_width, p=config.MEDIUM_TRANSFORMS_PROBABILITY ) ], p=1) # 执行变换并返回结果 return size_transform(image=image, mask=mask) # 返回Lambda变换,让它处理图像和mask return [ A.Lambda( image=apply_dynamic_transform, mask=apply_dynamic_transform, name="DynamicSizeTransform" ) ] def compose(transforms_to_compose): return A.Compose([ item for sublist in transforms_to_compose for item in sublist ])
3. train.py(调用方式不变)
import torch from . import transforms from .dataset import SegmentationDataset trainImages = ["./images/test.png"] trainMasks = ["./masks/test.png"] train_transforms = transforms.compose([ transforms.medium_transforms() ]) train_dataset = SegmentationDataset(imagePaths=trainImages, maskPaths=trainMasks, transforms=train_transforms)
原理说明
Lambda变换会在每次处理图像时,把__getitem__中传入的所有参数(包括original_height、original_width)传递给自定义的apply_dynamic_transform函数。- 这样就能在实际处理图像时,动态获取当前图像的原始尺寸,再构建并执行对应的变换,完美解决参数传递的时机问题。
内容的提问来源于stack exchange,提问作者Below the Radar
相关产品推荐
相关产品推荐

