训练ConvNeXt时触发TypeError:__init__()收到意外参数'split'
问题:MMpretrain训练ConvNeXt时触发TypeError参数错误
问题场景
使用mmpretrain仓库训练ConvNeXt模型,执行训练命令:
python tools/train.py configs/convnext/convnext-tiny_32xb128_custom.py
错误日志
Traceback (most recent call last): File "tools/train.py", line 162, in <module> main() File "tools/train.py", line 158, in main runner.train() File "C:\ProgramData\Anaconda3\envs\open-mmlab\lib\site-packages\mmengine\runner\runner.py", line 1728, in train self._train_loop = self.build_train_loop( File "C:\ProgramData\Anaconda3\envs\open-mmlab\lib\site-packages\mmengine\runner\runner.py", line 1527, in build_train_loop loop = EpochBasedTrainLoop( File "C:\ProgramData\Anaconda3\envs\open-mmlab\lib\site-packages\mmengine\runner\loops.py", line 44, in __init__ super().__init__(runner, dataloader) File "C:\ProgramData\Anaconda3\envs\open-mmlab\lib\site-packages\mmengine\runner\base_loop.py", line 26, in __init__ self.dataloader = runner.build_dataloader( File "C:\ProgramData\Anaconda3\envs\open-mmlab\lib\site-packages\mmengine\runner\runner.py", line 1370, in build_dataloader dataset = DATASETS.build(dataset_cfg) File "C:\ProgramData\Anaconda3\envs\open-mmlab\lib\site-packages\mmengine\registry\registry.py", line 570, in build return self.build_func(cfg, *args, **kwargs, registry=self) File "C:\ProgramData\Anaconda3\envs\open-mmlab\lib\site-packages\mmengine\registry\build_functions.py", line 121, in build_from_cfg obj = obj_cls(**args) # type: ignore File "d:\working\ai\convnext\convnext\mmpretrain\mmpretrain\datasets\custom.py", line 207, in __init__ super().__init__( TypeError: __init__() got an unexpected keyword argument 'split'
使用的配置文件
_base_ = [ '../_base_/models/convnext/convnext-tiny.py', '../_base_/datasets/imagenet_bs64_swin_224.py', '../_base_/schedules/imagenet_bs1024_adamw_swin.py', '../_base_/default_runtime.py', ] # dataset setting train_dataloader = dict(batch_size=128, dataset=dict(type="CustomDataset")) # schedule setting optim_wrapper = dict( optimizer=dict(lr=4e-3), clip_grad=None, ) # runtime setting custom_hooks = [dict(type='EMAHook', momentum=1e-4, priority='ABOVE_NORMAL')] # NOTE: `auto_scale_lr` is for automatically scaling LR # based on the actual training batch size. # base_batch_size = (32 GPUs) x (128 samples per GPU) auto_scale_lr = dict(base_batch_size=4096)
原因分析
你继承的基础数据集配置../_base_/datasets/imagenet_bs64_swin_224.py中,给数据集定义了split参数,但CustomDataset的初始化方法并不支持该参数,导致构建数据集实例时抛出参数不匹配的错误。
解决方法
方案1:覆盖split参数
在自定义配置的train_dataloader.dataset中显式设置split=None,覆盖基础配置中的该参数:
train_dataloader = dict(batch_size=128, dataset=dict(type="CustomDataset", split=None))
方案2:替换基础数据集配置
将_base_中的数据集配置替换为适用于CustomDataset的模板(比如仓库中提供的custom_dataset.py基础配置),避免继承到split参数。
注意:确保
CustomDataset所需的其他必要参数(如data_root、ann_file等)已正确配置,否则后续可能会出现其他数据集加载错误。
内容的提问来源于stack exchange,提问作者WonderCreator
相关产品推荐
相关产品推荐

