Deeplabv3+四类语义分割:非连续标签引发图执行错误求修复
Deeplabv3+非标准标签映射修改方案
核心逻辑
把你的原始标签值[2,4,6,8]一对一映射为标准索引[0,1,2,3],在数据加载阶段完成转换,同时将num_classes设回4,就能解决Graph execution error。
具体修改步骤及代码
1. 添加标签转换函数
找到代码中负责加载/预处理标签图像的部分(比如数据生成器、自定义数据集类的__getitem__方法,或者单独的加载函数),插入以下转换逻辑:
# 定义原始标签到标准索引的映射 LABEL_MAPPING = {2: 0, 4: 1, 6: 2, 8: 3} def convert_label_values(mask): # 创建与原标签同形状的空数组 converted_mask = np.zeros_like(mask, dtype=np.int32) # 遍历映射规则替换标签 for orig_val, target_val in LABEL_MAPPING.items(): converted_mask[mask == orig_val] = target_val return converted_mask
2. 在标签加载后立即调用转换函数
比如你原本加载标签的代码是这样的:
# 示例:加载灰度标签图 mask = cv2.imread(mask_file_path, cv2.IMREAD_GRAYSCALE)
在这行代码后面加上转换步骤:
mask = convert_label_values(mask)
⚠️ 注意:训练集、验证集的标签加载都要加这一步,漏了会导致部分数据仍有非法标签值。
3. 修正num_classes参数
找到代码中设置num_classes的位置(通常在模型定义、损失函数初始化或训练参数里),把值改回4:
num_classes = 4 # 替换之前设的9
4. 损失函数适配(按需调整)
- 如果你用的是
SparseCategoricalCrossentropy(适合整数索引标签):转换后的标签直接可用,无需额外修改。 - 如果你用的是
CategoricalCrossentropy(需要one-hot编码):在转换后添加one-hot处理:
from tensorflow.keras.utils import to_categorical mask = to_categorical(mask, num_classes=num_classes)
验证修改效果
重启训练后,若之前的Graph execution error消失,且训练损失正常波动下降,说明修改生效。
排查要点
- 检查所有标签加载的分支(比如训练/验证/测试集)是否都应用了转换函数
- 确认标签图像的 dtype 是整数类型(如uint8),避免浮点类型导致匹配失败
内容的提问来源于stack exchange,提问作者Enigma
相关产品推荐
相关产品推荐

