You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

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默认会用第一个捕获组作为标签,导致模型把图片的编号当成了类别。
  • 这就造成两个问题:
    1. 类别数量高达2046(实际是不同的图片编号,而非真实的蘑菇类别);
    2. 验证集中出现了训练集未包含的编号,触发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)

验证修正结果

运行修正后的代码后:

  1. 检查data.vocab应显示真实的蘑菇属名(如Entoloma、Suillus等);
  2. 类别数量应与你实际的蘑菇类别数一致(而非2046);
  3. 重新训练模型即可避免KeyError。

内容的提问来源于stack exchange,提问作者spoolito

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.06.22 22:09:56