FastAI视觉模型训练报错:标签‘755’未包含在训练数据集中
问题:FastAI蘑菇图像分类训练时触发KeyError:标签不在训练数据集中
报错信息
File /opt/conda/lib/python3.10/site-packages/fastai/data/transforms.py:263, in Categorize.encodes(self, o) 261 return TensorCategory(self.vocab.o2i[o]) 262 except KeyError as e: --> 263 raise KeyError(f"Label '{o}' was not included in the training dataset") from e KeyError: "Label '755' was not included in the training dataset"
问题代码片段
数据加载部分
bs = 64 path = Path("../input/mushrooms/Mushrooms/") fnames = [] for fpath in class_names: print(path/f'{fpath}/') fnames += get_image_files(path/f'{fpath}/') # 输出的文件夹示例: # ../input/mushrooms/Mushrooms/Entoloma # ../input/mushrooms/Mushrooms/Suillus # ...
数据集构建部分
np.random.seed(2) pat = r"(\d+)_([a-zA-Z0-9-_]+)\.jpg$" item_tfms = Resize(224) batch_tfms = [*aug_transforms(), Normalize.from_stats(*imagenet_stats)] data = ImageDataLoaders.from_name_re( path='.', fnames=fnames, pat=pat, item_tfms=item_tfms, batch_tfms=batch_tfms, bs=bs, num_workers=0 ) # 检查输出: # Number of classes: 2046 # Training dataset size: 5372 # Validation dataset size: 1342
模型训练代码
learn = vision_learner(data, models.resnet50, metrics=error_rate, lr=0.001) learn.fit(n_epochs = 5, start_epoch=0)
问题根源
核心错误是标签提取逻辑完全错误:
- 你使用的正则表达式
r"(\d+)_([a-zA-Z0-9-_]+)\.jpg$"会捕获两个组:第一个是文件名中的数字(如755),第二个是蘑菇属名(如Amanita)。 - FastAI的
from_name_re默认会用第一个捕获组作为标签,导致模型把图片的编号当成了类别。 - 这就造成两个问题:
- 类别数量高达2046(实际是不同的图片编号,而非真实的蘑菇类别);
- 验证集中出现了训练集未包含的编号,触发KeyError。
解决方法
有两种可靠的修正方案,优先推荐方案二(更简洁不易出错):
方案一:修正正则的标签提取逻辑
明确指定使用正则的第二个捕获组作为蘑菇类别标签:
import re np.random.seed(2) pat = r"(\d+)_([a-zA-Z0-9-_]+)\.jpg$" # 定义标签提取函数,取第二个捕获组 def get_label(fn): match = re.search(pat, fn.name) return match.group(2) if match else None item_tfms = Resize(224) batch_tfms = [*aug_transforms(), Normalize.from_stats(*imagenet_stats)] data = ImageDataLoaders.from_name_re( path='.', fnames=fnames, pat=pat, label_func=get_label, # 指定自定义标签提取函数 item_tfms=item_tfms, batch_tfms=batch_tfms, bs=bs, num_workers=0 ) # 验证:此时data.vocab应该是蘑菇属名列表,类别数与你实际的蘑菇类别一致(比如你打印的9个属) print(data.vocab)
方案二:直接从文件夹路径提取标签(更推荐)
由于你的图片已经按蘑菇属名分文件夹存放,直接用from_folder自动从父文件夹名提取标签,无需手写正则:
bs = 64 path = Path("../input/mushrooms/Mushrooms/") item_tfms = Resize(224) batch_tfms = [*aug_transforms(), Normalize.from_stats(*imagenet_stats)] # 自动按子文件夹分类,默认划分20%数据为验证集 data = ImageDataLoaders.from_folder( path, valid_pct=0.2, # 可根据需求调整验证集比例 item_tfms=item_tfms, batch_tfms=batch_tfms, bs=bs, num_workers=0 ) # 验证类别数量与实际蘑菇属一致 print(f"Number of classes: {len(data.vocab)}") # 应该等于你class_names的长度(比如9)
验证修正结果
运行修正后的代码后:
- 检查
data.vocab应显示真实的蘑菇属名(如Entoloma、Suillus等); - 类别数量应与你实际的蘑菇类别数一致(而非2046);
- 重新训练模型即可避免KeyError。
内容的提问来源于stack exchange,提问作者spoolito
相关产品推荐
相关产品推荐

