解决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尝试调用实例会触发"不可调用"错误。
具体修改步骤
修正
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, )修改
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
相关产品推荐
相关产品推荐

