MMSegmentation处理Freiburg Forest数据集时的张量尺寸RuntimeError求助
语义分割训练RuntimeError问题排查与解决
我正在用不同分割模型对Freiburg Forest数据集进行语义分割,已参照MMSegmentation官方文档完成配置文件编写,但训练时触发RuntimeError,错误堆栈如下:
Traceback (most recent call last): File "tools/train.py", line 109, in <module> main() File "tools/train.py", line 105, in main runner.train() File "/home/deshpand/anaconda3/envs/openmmlab/lib/python3.8/site-packages/mmengine/runner/runner.py", line 1733, in train model = self.train_loop.run() # type: ignore File "/home/deshpand/anaconda3/envs/openmmlab/lib/python3.8/site-packages/mmengine/runner/loops.py", line 278, in run self.run_iter(data_batch) File "/home/deshpand/anaconda3/envs/openmmlab/lib/python3.8/site-packages/mmengine/runner/loops.py", line 301, in run_iter outputs = self.runner.model.train_step( File "/home/deshpand/anaconda3/envs/openmmlab/lib/python3.8/site-packages/mmengine/model/wrappers/distributed.py", line 121, in train_step losses = self._run_forward(data, mode='loss') File "/home/deshpand/anaconda3/envs/openmmlab/lib/python3.8/site-packages/mmengine/model/wrappers/distributed.py", line 161, in _run_forward results = self(**data, mode=mode) File "/home/deshpand/anaconda3/envs/openmmlab/lib/python3.8/site-packages/torch/nn/modules/module.py", line 1190, in _call_impl return forward_call(*input, **kwargs) File "/home/deshpand/anaconda3/envs/openmmlab/lib/python3.8/site-packages/torch/nn/parallel/distributed.py", line 1040, in forward output = self._run_ddp_forward(*inputs, **kwargs) File "/home/deshpand/anaconda3/envs/openmmlab/lib/python3.8/site-packages/torch/nn/parallel/distributed.py", line 1000, in _run_ddp_forward return module_to_run(*inputs[0], **kwargs[0]) File "/home/deshpand/anaconda3/envs/openmmlab/lib/python3.8/site-packages/torch/nn/modules/module.py", line 1190, in _call_impl return forward_call(*input, **kwargs) File "/home/deshpand/Thesis/semantic_segmentation_network/mmseg_new/mmsegmentation/mmseg/models/segmentors/base.py", line 94, in forward return self.loss(inputs, data_samples) File "/home/deshpand/Thesis/semantic_segmentation_network/mmseg_new/mmsegmentation/mmseg/models/segmentors/encoder_decoder.py", line 177, in loss loss_decode = self._decode_head_forward_train(x, data_samples) File "/home/deshpand/Thesis/semantic_segmentation_network/mmseg_new/mmsegmentation/mmseg/models/segmentors/encoder_decoder.py", line 138, in _decode_head_forward_train loss_decode = self.decode_head.loss(inputs, data_samples, File "/home/deshpand/Thesis/semantic_segmentation_network/mmseg_new/mmsegmentation/mmseg/models/decode_heads/decode_head.py", line 262, in loss losses = self.loss_by_feat(seg_logits, batch_data_samples) File "/home/deshpand/Thesis/semantic_segmentation_network/mmseg_new/mmsegmentation/mmseg/models/decode_heads/decode_head.py", line 324, in loss_by_feat loss[loss_decode.loss_name] = loss_decode( File "/home/deshpand/anaconda3/envs/openmmlab/lib/python3.8/site-packages/torch/nn/modules/module.py", line 1190, in _call_impl return forward_call(*input, **kwargs) File "/home/deshpand/Thesis/semantic_segmentation_network/mmseg_new/mmsegmentation/mmseg/models/losses/cross_entropy_loss.py", line 271, in forward loss_cls = self.loss_weight * self.cls_criterion( File "/home/deshpand/Thesis/semantic_segmentation_network/mmseg_new/mmsegmentation/mmseg/models/losses/cross_entropy_loss.py", line 45, in cross_entropy loss = F.cross_entropy( File "/home/deshpand/anaconda3/envs/openmmlab/lib/python3.8/site-packages/torch/nn/functional.py", line 3026, in cross_entropy return torch._C._nn.cross_entropy_loss(input, target, weight, _Reduction.get_enum(reduction), ignore_index, label_smoothing) RuntimeError: only batches of spatial targets supported (3D tensors) but got targets of size: : [2, 256, 256, 3] ERROR:torch.distributed.elastic.multiprocessing.api:failed (exitcode: 1) local_rank: 0 (pid: 1128219) of binary: /home/deshpand/anaconda3/envs/openmmlab/bin/python Traceback (most recent call last): File "/home/deshpand/anaconda3/envs/openmmlab/lib/python3.8/runpy.py", line 194, in _run_module_as_main return _run_code(code, main_globals, None, File "/home/deshpand/anaconda3/envs/openmmlab/lib/python3.8/runpy.py", line 87, in _run_code exec(code, run_globals) File "/home/deshpand/anaconda3/envs/openmmlab/lib/python3.8/site-packages/torch/distributed/launch.py", line 195, in <module> main() File "/home/deshpand/anaconda3/envs/openmmlab/lib/python3.8/site-packages/torch/distributed/launch.py", line 191, in main launch(args) File "/home/deshpand/anaconda3/envs/openmmlab/lib/python3.8/site-packages/torch/distributed/launch.py", line 176, in launch run(args) File "/home/deshpand/anaconda3/envs/openmmlab/lib/python3.8/site-packages/torch/distributed/run.py", line 753, in run elastic_launch( File "/home/deshpand/anaconda3/envs/openmmlab/lib/python3.8/site-packages/torch/distributed/launcher/api.py", line 132, in __call__ return launch_agent(self._config, self._entrypoint, list(args)) File "/home/deshpand/anaconda3/envs/openmmlab/lib/python3.8/site-packages/torch/distributed/launcher/api.py", line 246, in launch_agent raise ChildFailedError( torch.distributed.elastic.multiprocessing.errors.ChildFailedError: ============================================================ tools/train.py FAILED ------------------------------------------------------------ Failures: [1]: time : 2023-05-29_10:18:58 host : neptun.informatik.uni-kl.de rank : 1 (local_rank: 1) exitcode : 1 (pid: 1128220) error_file: <N/A> traceback : To enable traceback see: https://pytorch.org/docs/stable/elastic/errors.html ------------------------------------------------------------ Root Cause (first observed failure): [0]: time : 2023-05-29_10:18:58 host : neptun.informatik.uni-kl.de rank : 0 (local_rank: 0) exitcode : 1 (pid: 1128219) error_file: <N/A> traceback : To enable traceback see: https://pytorch.org/docs/stable/elastic/errors.html ============================================================
错误根源
核心错误提示RuntimeError: only batches of spatial targets supported (3D tensors) but got targets of size: : [2, 256, 256, 3],说明:
- MMSegmentation语义分割任务要求标签是3D张量(形状为[批量数, 高度, 宽度]),每个像素值对应一个类别索引
- 当前加载的标签是4D张量(多了3通道维度),意味着程序把Freiburg Forest的RGB彩色标签直接当成了输入,没有转换为单通道的类别索引
解决方法
- 添加RGB标签到类别索引的映射转换:
建立数据集的颜色-类别ID映射表(可参考数据集官方说明获取对应关系),在数据预处理pipeline中添加自定义转换步骤,将RGB格式的标签图像转换为单通道的类别索引张量。 - 修改数据加载配置:
在dataset的pipeline配置中,确保使用正确的标签加载逻辑,替换直接读取RGB图像的方式,改用能处理彩色掩码的加载模块(如自定义LoadColoredAnnotations)。 - 验证张量维度:
在训练前手动检查加载后的标签张量形状,确保输出为[batch_size, height, width]。可以在数据加载器中添加打印语句,查看标签的shape确认转换是否生效。 - 核对类别数量配置:
确保配置文件中num_classes参数和数据集实际类别数一致,避免因类别不匹配间接导致维度错误。
内容的提问来源于stack exchange,提问作者programmer_04_03
相关产品推荐
相关产品推荐

