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

解决PyTorch Lightning自定义数据集‘object is not callable’报错

问题:PyTorch Lightning自定义数据集替换时出现'CUSTOMDset' object is not callable错误

刚接触PyTorch Lightning开发模式,在时空Transformer项目中尝试将原有CIFAR10/MNIST数据集替换为自定义数据集时,数据模块触发错误。


相关代码

数据集创建函数create_dset

def create_dset(config):
    INV_SCALER = lambda x: x
    SCALER = lambda x: x
    NULL_VAL = None
    PLOT_VAR_IDXS = None
    PLOT_VAR_NAMES = None
    PAD_VAL = None


    if config.dset in ["mnist", "cifar"]:
        if config.dset == "mnist":
            config.target_points = 28 - config.context_points
            datasetCls = stf.data.image_completion.MNISTDset
            PLOT_VAR_IDXS = [18, 24]
            PLOT_VAR_NAMES = ["18th row", "24th row"]
        else:
            config.target_points = 32 * 32 - config.context_points
            datasetCls = stf.data.image_completion.CIFARDset
            PLOT_VAR_IDXS = [0]
            PLOT_VAR_NAMES = ["Reds"]
        DATA_MODULE = stf.data.DataModule(
            datasetCls=datasetCls,
            dataset_kwargs={"context_points": config.context_points},
            batch_size=config.batch_size,
            workers=config.workers,
            overfit=args.overfit,
        )

        return (
            DATA_MODULE,
            INV_SCALER,
            SCALER,
            NULL_VAL,
            PLOT_VAR_IDXS,
            PLOT_VAR_NAMES,
            PAD_VAL,
        )
    # 自定义数据集分支
    elif config.dset == "custom":
        data_dir = "./spacetimeformer/mydata"

        # 数据变换定义
        transform = transforms.Compose([
            transforms.Resize((256, 256)),  # 调整图像尺寸
            transforms.ToTensor(),          # 转换为张量
            transforms.Normalize((0.5, 0.5, 0.5), (0.5, 0.5, 0.5)),  # 归一化像素值
        ])
        
        # 创建自定义数据集实例
        dataset = CUSTOMDset(data_dir, context_points=config.context_points, transform=transform)
        
        DATA_MODULE = stf.data.DataModule(
            datasetCls=dataset,
            dataset_kwargs={"context_points": config.context_points},
            batch_size=config.batch_size,
            workers=config.workers,
            overfit=args.overfit,
        )

        PLOT_VAR_IDXS = [0]
        PLOT_VAR_NAMES = ["Reds"]
        PAD_VAL = None
    
        return (
            DATA_MODULE,
            INV_SCALER,
            SCALER,
            NULL_VAL,
            PLOT_VAR_IDXS,
            PLOT_VAR_NAMES,
            PAD_VAL,
        )

自定义数据集类CUSTOMDset

class CUSTOMDset(Dataset):
    def __init__(self, data_dir, context_points, transform=None):
        self.data_dir = data_dir
        self.context_points = context_points
        self.transform = transform
        self.file_list = os.listdir(self.data_dir)

    def __len__(self):
        return len(self.file_list)

    def __getitem__(self, i):
        img_name = os.path.join(self.data_dir, self.file_list[i])
        img = Image.open(img_name)
        y = img.convert("RGB")

        y_c = y.crop((0, 0, self.context_points, y.size[1]))
        y_t = y.crop((self.context_points, 0, y.size[0], y.size[1]))

        x = torch.arange(y.size[1]).view(-1, 1).float() / y.size[1]
        x_c = x[: self.context_points]
        x_t = x[self.context_points:]

        if self.transform:
            y_c = self.transform(y_c)
            y_t = self.transform(y_t)

        return x_c, y_c, x_t, y_t

错误信息

调用create_dset获取data_module后,执行test_dataloader = data_module.test_dataloader()时触发错误:

TypeError: 'CUSTOMDset' object is not callable

完整报错栈:

Traceback (most recent call last):
  File "train_my.py", line 525, in <module>
    main(args)
  File "train_my.py", line 435, in main
    test_dataloader = data_module.test_dataloader()
  File "/home/spacetimeformer/spacetimeformer/data/datamodule.py", line 37, in test_dataloader
    return self._make_dloader("test", shuffle=shuffle)
  File "/home/spacetimeformer/spacetimeformer/data/datamodule.py", line 44, in _make_dloader
    self.datasetCls(**self.dataset_kwargs, split=split),
TypeError: 'CUSTOMDset' object is not callable

环境信息:python=3.8,torch=2.0.1


解决思路

核心问题

在自定义数据集分支中,你将实例化后的CUSTOMDset对象传给了DataModule的datasetCls参数,但该参数需要接收的是数据集类本身,而非类的实例。原代码中MNIST/CIFAR分支都是传入类,DataModule内部会自动调用类并传入参数创建训练/验证/测试集实例,而传入实例后,DataModule尝试调用实例会触发"不可调用"错误。

具体修改步骤

  1. 修正DataModule参数传递
    去掉自定义数据集实例化的代码,直接传入CUSTOMDset类,并将所有初始化参数放入dataset_kwargs:

    elif config.dset == "custom":
        data_dir = "./spacetimeformer/mydata"
    
        transform = transforms.Compose([
            transforms.Resize((256, 256)),
            transforms.ToTensor(),
            transforms.Normalize((0.5, 0.5, 0.5), (0.5, 0.5, 0.5)),
        ])
        
        # 传入类而非实例
        datasetCls = CUSTOMDset
        # 整理所有需要的初始化参数
        dataset_kwargs = {
            "data_dir": data_dir,
            "context_points": config.context_points,
            "transform": transform
        }
        
        DATA_MODULE = stf.data.DataModule(
            datasetCls=datasetCls,
            dataset_kwargs=dataset_kwargs,
            batch_size=config.batch_size,
            workers=config.workers,
            overfit=args.overfit,
        )
    
        PLOT_VAR_IDXS = [0]
        PLOT_VAR_NAMES = ["Reds"]
        PAD_VAL = None
    
        return (
            DATA_MODULE,
            INV_SCALER,
            SCALER,
            NULL_VAL,
            PLOT_VAR_IDXS,
            PLOT_VAR_NAMES,
            PAD_VAL,
        )
    
  2. 修改CUSTOMDset支持split参数
    从报错栈可知,DataModule会自动传入split参数(区分train/val/test),因此需要在CUSTOMDset的__init__方法中添加该参数,可根据需要实现数据集拆分逻辑:

    class CUSTOMDset(Dataset):
        def __init__(self, data_dir, context_points, transform=None, split=None):
            self.data_dir = data_dir
            self.context_points = context_points
            self.transform = transform
            self.split = split
            
            # 可选:根据split加载不同子集,比如按文件夹划分或比例拆分
            if split:
                # 示例:假设数据按train/test文件夹存放
                self.file_list = os.listdir(os.path.join(data_dir, split))
            else:
                self.file_list = os.listdir(data_dir)
    
        # 其余方法保持不变...
    

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.14 18:05:56