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

如何在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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.08 21:24:58