恒美微站
首页
关于我们
建站服务
主题模板
案例展示
资讯中心
联系我们
Pokemon小样本图像分类实战:数据增强与端侧部署指南
首页
资讯中心
/
Pokemon小样本图像分类实战:数据增强与端侧部署指南
Pokemon小样本图像分类实战:数据增强与端侧部署指南
发布时间:2026/9/14 16:24:07
简介面向机器学习与深度学习研究者的宝可梦图像分类数据集共包含 1000 个物种类别、26,539 张图像所有图片统一缩放为 128×128 像素并保存为 PNG 格式。数据按类别分子目录存放加载时可直接用子目录名作为类别标签省去额外的标注整理工作每类约 40 张样本尤其适合在样本量受限的场景下测试模型性能或用于数据增强策略的对比研究。资源包共 2000 个文件其中 1999 张 PNG 图像构成主要数据另附 1 个 Python 脚本便于批量预览或辅助处理图像整体压缩后约 672.13MB目录结构清晰可快速接入常用深度学习训练流程。由于每个类别图像数量有限训练时建议配合旋转、裁剪等数据增强方法以提升模型的泛化能力该数据最初为 Flutter 的 Pokedex 项目设计也可用于其他多类别图像分类任务。目前已有 97 人学习适合正在开展图像分类实验、需要小样本多类别数据的研究者与开发者。1. 26K 张图、1000 个类这个 Pokemon 数据集值得拿来训什么先说一个反直觉的结论这份 Pokemon 数据集最值钱的不是它有 26,539 张图而是它只有 1000 个类、每类约 40 张、统一缩放到 128x128 的“小”设定。这个规模恰好卡在“直接训 CNN 容易过拟合、但不做任何处理又完全够用”的临界点上——很适合用来验证数据增强策略、迁移学习效果以及端侧模型在长尾分布下的真实表现。如果你是做图像分类的初学者拿它入门比用 ImageNet 友好得多如果是五年以上经验的人这套数据的价值在于类别多、样本少、文件名噪声大正好用来测试你的数据管线是否足够健壮。它最初是给 Flutter 版 Pokedex 做识别的所以从目录结构到 PNG 格式都不是随便拍的而是经过一遍预处理和标签化的产物。下面我会从数据组织方式讲起一路写到训练配置和端侧部署最后给一个一键审计脚本。2. 数据集结构与标签体系目录即标签CSV 负责映射2.1 子目录与文件名命名规则从 gyarados_18.png 可以看出什么打开数据集根目录你会看到 1000 个子目录每个目录名就是精灵的英文名比如gyarados/、arbok/、ho-oh/。里面放的是一批物种名_编号.png文件例如gyarados_18.png。这种命名方式有一个直接好处不需要额外读标签文件os.listdir就能拿到类别名。但要注意ho-oh_27.png这种带连字符的名字在部分框架里会被当成特殊符号比如某些版本的数据加载器会按_切分文件名导致ho和oh被拆开。所以稳妥的做法是直接从父目录取标签不要依赖文件名解析。类别数量是 1000但每类不一定是严格的 40 张。用下面这段代码可以快速统计每类图片数并找出最少的类from pathlib import Path import collections root Path(pokemon_dataset) counts {} for cls_dir in root.iterdir(): if cls_dir.is_dir(): counts[cls_dir.name] len(list(cls_dir.glob(*.png))) counter collections.Counter(counts) print(总类别数:, len(counter)) print(最少样本类:, counter.most_common()[-5:]) print(最多样本类:, counter.most_common()[:5])这段代码遍历所有子目录统计每个目录下的 PNG 数量。输出结果能直接告诉你类别分布是否均衡以及有没有空目录或图片数异常少的类。拿到这个分布后你才好决定后面是用WeightedRandomSampler还是直接做类别重采样。2.2 从目录生成 CSV用 Python 脚本一次搞定原数据集本身并没有把标签集中放在一个 CSV 里标题里的“CSV”更多是告诉你应该自己生成一份方便后续用 Pandas 做分析或在 Flutter 里读取。我习惯生成一个三列的pokemon_labels.csvfilename相对路径、label_id从 0 到 999 的整数、species精灵名。import csv from pathlib import Path root Path(pokemon_dataset) rows [] for label_id, cls_dir in enumerate(sorted(p for p in root.iterdir() if p.is_dir())): for img_path in sorted(cls_dir.glob(*.png)): rows.append([str(img_path.relative_to(root)), label_id, cls_dir.name]) with open(pokemon_labels.csv, w, newline, encodingutf-8-sig) as f: writer csv.writer(f) writer.writerow([filename, label_id, species]) writer.writerows(rows) print(写入完成共, len(rows), 条记录)这里有一个容易踩的坑如果你用encodingutf-8写 CSV然后在 PyCharm 里直接双击打开可能会看到乱码。因为 PyCharm 默认用 UTF-8而 Windows 下的 Excel 或某些旧工具用 GBK。utf-8-sig会在文件头加上 BOMExcel 和 PyCharm 都能正确识别。这也是为什么网上经常有人问“pycharm中生成的csv文件打开是乱码”——就是编码没带 BOM 导致的。2.3 CSV 完整性校验与拆分MD5 与行数核对拿到 CSV 后第一件事不是急着训练而是校验它和真实目录是否对得上。常见校验包括CSV 里的每个filename是否真实存在、图片是否能被解码、尺寸是否是 128x128。对于 CSV 本身可以用 MD5 做完整性校验防止传输过程中被截断。md5sum pokemon_labels.csv # 输出形如: 3f2e4b1c... pokemon_labels.csv如果你需要把这份 CSV 拆成多个部分比如按类别拆成 train/val 两份不要用split直接按行切因为那样会把同一个类切散导致验证集里出现训练时见过的类别。推荐用 Pandas 或sklearn的分层切分后面第 4 章会给出具体代码。如果只是想快速拆分给不同同事处理可以按行数均分split -l 5000 pokemon_labels.csv part_但这样每个子文件都会带上表头吗不会只有第一个文件有表头。所以更规范的做法是用tail -n 2去掉表头再拆拆完再手动加。这里推荐你直接用 Python 脚本做按类拆分避免污染标签映射。3. 图像预处理与数据增强128x128 小图喂给 CNN 前必须做的事3.1 为什么每类只有 40 张时数据增强不是可选项而是必选项1000 类、每类 40 张意味着平均每个类别只有 40 个样本。而一个标准的 ResNet50 分类头就要输出 1000 个概率参数量集中在最后的全连接层。如果没有增强模型很容易把每个类的 40 张图背下来在训练集上拿到 99% 的准确率验证集却只有 20%。更麻烦的是Pokemon 图像本身是同一角色在不同姿态、不同光影下的渲染图类内差异不小类间差异有时候很小——比如ninetales和vulpix的早期形态。这种任务天然需要模型学到“纹理轮廓”的鲁棒特征而不是死记背景。原数据集作者也特别强调“强烈建议使用数据增强”。所以这里的数据增强不是随便翻转一下而是要覆盖几何变换、颜色扰动和遮挡。下面给出一个我在类似小样本任务里常用的 PyTorch 增强管线。3.2 PyTorch 增强管线旋转、裁剪、色彩抖动、Cutoutimport torch from torchvision import transforms train_transform transforms.Compose([ transforms.RandomResizedCrop(size(128, 128), scale(0.7, 1.0)), transforms.RandomHorizontalFlip(p0.5), transforms.RandomRotation(degrees15), transforms.ColorJitter(brightness0.3, contrast0.3, saturation0.2, hue0.05), transforms.ToTensor(), transforms.Normalize(mean[0.5, 0.5, 0.5], std[0.5, 0.5, 0.5]), ])参数选择的逻辑RandomResizedCrop的scale设为(0.7, 1.0)意思是随机裁剪原图的 70% 到 100% 区域再缩放到 128x128。这比固定中心裁剪更有用因为 Pokemon 主体在画面中的占比并不固定。RandomRotation(15)控制在 ±15 度超过这个角度会引入过多背景反而干扰。ColorJitter的 hue 只给 0.05因为精灵原色比较鲜明色相变化太大容易改变物种属性。如果还想更强一点可以加 Cutout 或 RandomErasingfrom torchvision.transforms import RandomErasing train_transform transforms.Compose([ # ... 前面的变换 ... transforms.ToTensor(), transforms.RandomErasing(p0.3, scale(0.02, 0.15), ratio(0.3, 3.3)), transforms.Normalize(mean[0.5, 0.5, 0.5], std[0.5, 0.5, 0.5]), ])RandomErasing会在图上随机遮挡一块矩形区域模拟物体被部分遮挡的情况。p0.3表示 30% 的图会被擦除scale是擦除区域占原图面积的比例。对于这种小尺寸图像擦除区域不要超过 15%否则主体都没了。3.3 常见加载误区CSV 里存 base64 PNG 的坑有些人为了方便传输会把图片转成data:image/png;base64,xxxx这种字符串存进 CSV。如果你拿到的是这种格式千万不要直接把它当路径喂给Image.open()。你需要先解码成字节流import base64 import io from PIL import Image def load_base64_png(s: str): if s.startswith(data:image/png;base64,): s s.split(,, 1)[1] img Image.open(io.BytesIO(base64.b64decode(s))) return img.convert(RGB)base64 编码会让体积膨胀约 33%而且读取时需要额外解码内存开销更大。如果你的 CSV 里有这种字段建议直接用上一章的方法把图片落盘成.png文件然后 CSV 里只存路径。实际训练时直接从文件系统读取比从内存解码快得多也方便 PyTorch 的DataLoader做多进程预取。3.4 从 PNG 到 Tensor批量读取与缓存优化128x128 的 PNG 本身不大单张约 10-50KB但 26K 张图在训练时如果每轮都从磁盘读IO 会成为瓶颈。常见优化是先做一次批量预加载到内存或者用lmdb打包。不过更省事的做法是让DataLoader的num_workers设为 4 或 8并把pin_memoryTrue在 GPU 训练时。还可以用torchdata或WebDataset把图片打包成 tar 分片但这个数据集规模还没到那个程度。如果你用的是ImageFolder只需要把根目录传进去它会自动递归读取子目录作为类别。这样连 CSV 都不用但前面生成的 CSV 可以用来做 Stratified Split 和后续部署时的映射。4. 模型选型与训练配置1000 类小样本的分类实战4.1 模型选择ResNet50 还是 MobileNetV3在 1000 类、每类 40 张的设定下直接从头训练一个深层 CNN 基本不可行除非你把数据增强开到很大并配合强正则化。我的建议是优先用在 ImageNet 上预训练过的模型只替换最后的全连接层。下表是几种常见模型的对比模型参数量输入尺寸在 1000 类小样本上的表现倾向适用场景ResNet5025.6M224x224需要大量正则化容易过拟合可迁移精度上限高EfficientNet-B05.3M224x224参数效率高收敛快中等算力推荐试试MobileNetV3-Large5.4M224x224轻量适合端侧部署Flutter 场景ViT-Tiny/DeiT5.7M224x224数据少时不如 CNN不推荐除非加蒸馏这里要注意数据集本身是 128x128但预训练模型要求 224 输入。你可以把输入 resize 到 224或者用torchvision.models.resnet50(weightsResNet50_Weights.IMAGENET1K_V2)然后修改conv1的 stride 来适配小图。常见做法是直接 resize 到 224虽然损失一点原始分辨率但能直接复用预训练权重。如果你想保持 128 输入需要修改模型的第一层卷积并且加载权重时忽略不匹配的层这样多少会损失一些性能。4.2 训练脚本核心标签平滑、MixUp、余弦退火下面这段代码是训练循环里最关键的部分包含了标签平滑、MixUp 和余弦退火三个技巧。这三个技巧对 1000 类小样本非常有效import torch import torch.nn as nn import torch.optim as optim from torch.optim.lr_scheduler import CosineAnnealingLR def mixup_data(x, y, alpha0.2): lam torch.distributions.Beta(alpha, alpha).sample().item() index torch.randperm(x.size(0)).to(x.device) mixed_x lam * x (1 - lam) * x[index] return mixed_x, y, y[index], lam # 标签平滑不用硬标签防止模型对训练集过于自信 class LabelSmoothingCrossEntropy(nn.Module): def __init__(self, classes1000, smoothing0.1): super().__init__() self.conf 1.0 - smoothing self.smooth smoothing / classes def forward(self, pred, target): pred torch.log_softmax(pred, dim-1) with torch.no_grad(): true_dist torch.zeros_like(pred).fill_(self.smooth) true_dist.scatter_(1, target.unsqueeze(1), self.conf) return torch.mean(torch.sum(-true_dist * pred, dim-1)) criterion LabelSmoothingCrossEntropy(classes1000, smoothing0.1) optimizer optim.SGD(model.parameters(), lr0.01, momentum0.9, weight_decay1e-4) scheduler CosineAnnealingLR(optimizer, T_max50)mixup_data在每批次中随机把两张图按比例混合标签也按同样比例混合相当于一种免费的数据扩充。alpha0.2是比较温和的选择太小混合效果弱太大容易让模型学不到清晰边界。标签平滑的smoothing0.1是经验值对于 1000 类问题可以防止全连接层的输出分数无限趋近于正确类别。余弦退火T_max50表示 50 个 epoch 内学习率从 0.01 降到接近 0配合warmup效果更好。训练时还要注意model的最后一层输出维度必须改成 1000。如果你用ResNet50代码是model.fc nn.Linear(2048, 1000)用MobileNetV3是model.classifier[-1] nn.Linear(1024, 1000)。4.3 分层划分与评估保证每个类都有验证样本由于类别多且每类样本少普通随机划分很可能让某个类在训练集或验证集中缺失。必须用分层采样import pandas as pd from sklearn.model_selection import train_test_split df pd.read_csv(pokemon_labels.csv) train_df, val_df train_test_split( df, test_size0.2, stratifydf[label_id], random_state42 ) print(train_df[label_id].nunique(), val_df[label_id].nunique())stratifydf[label_id]确保每个类别在训练集和验证集中所占比例相同。random_state42固定随机种子保证可复现。验证集每类大约 8 张足够计算 Top-1 和 Top-5 准确率。评估时用classification_report能直接看到每个类的精确率和召回率from sklearn.metrics import classification_report y_true val_labels y_pred model_predictions(val_loader) print(classification_report(y_true, y_pred, target_namesclass_names, zero_division0))重点关注那些样本少且准确率低的类比如某些闪光形态或进化前后差别很大的精灵。如果某个类的召回率明显低于其他类说明它和某个相近类混淆严重这时可以针对该类额外采集数据或者调整损失函数的类权重。5. 端侧部署到 Flutter Pokedex模型转换与资源路径修复5.1 导出 ONNX/TorchScript 到移动端训练完模型后如果要集成到 Flutter 的 Pokedex 项目里常见做法是导出为 ONNX 或 TorchScript然后用onnxruntime或libtorch在移动端推理。这里给出 ONNX 导出代码import torch import torch.onnx model.eval() dummy_input torch.randn(1, 3, 128, 128) # 注意输入尺寸要和预处理一致 torch.onnx.export( model, dummy_input, pokemon_model.onnx, input_names[input], output_names[output], dynamic_axes{input: {0: batch}, output: {0: batch}} )导出时最容易忽略的是输入尺寸。训练时如果用了RandomResizedCrop(128)但导出时你传了 224那么 ONNX 的输入 shape 就是 224Flutter 端也必须 resize 到 224。建议固定为 128因为原数据集就是 128x128这样端侧也少一次缩放操作。另外dynamic_axes让 batch 维度可变方便单张图片推断。5.2 Flutter 侧图像预处理128x128 PNG 转 Tensor在 Flutter 里加载 PNG 并转成模型的输入最直接的方式是使用image包解码然后按 RGB 顺序填充Float32Listimport package:image/image.dart as img; import dart:typed_data; Float32List preprocessPng(Uint8List bytes, {int size 128}) { final image img.decodeImage(bytes)!; final resized img.copyResize(image, width: size, height: size); final mean [0.5, 0.5, 0.5], std [0.5, 0.5, 0.5]; final tensor Float32List(1 * 3 * size * size); var idx 0; for (var c 0; c 3; c) { for (var y 0; y size; y) { for (var x 0; x size; x) { final pixel resized.getPixel(x, y); final value (c 0) ? pixel.r / 255.0 : (c 1 ? pixel.g / 255.0 : pixel.b / 255.0); tensor[idx] (value - mean[c]) / std[c]; } } } return tensor; }这里采用NCHW布局和 PyTorch 一致。注意image包的getPixel返回的pixel.r的取值范围可能是 0-1 或 0-255取决于包版本。我这里假设是 0-255如果你用的版本是 0-1就不要再除以 255。这就是为什么端侧推理经常出现“结果全错但代码看起来没问题”的原因——数值范围不一致。5.3 资源路径问题failed to resolve import ../assets/grenade...的排查很多人把数据集和模型放进 Flutter 项目后会遇到failed to resolve import ../assets/grenade (1024x128)[frames8].png这类报错。这个报错通常不是模型的问题而是pubspec.yaml中 assets 路径没有匹配到资源。Flutter 加载 assets 时使用相对路径../assets/...这种写法在打包后根本不存在因为 assets 会被平铺到AssetManifest.json。正确做法是在pubspec.yaml里声明目录flutter: assets: - assets/images/ - assets/models/pokemon_model.onnx然后在代码里用rootBundle.load(assets/images/gyarados_18.png)加载不要用../。另外报错里的(1024x128)[frames8]表明这是一个 SpriteSheet 或序列帧但 Pokemon 数据集是单张 PNG说明你可能把精灵动画资源也塞进了同一个 assets 目录。这是两套东西尽量分开目录存放避免 Flutter 把序列帧图也当成普通图片处理。6. 搞一个一键质量审计脚本图像、CSV、类别分布一次查完最后分享一个我每次拿到新图像数据集都会写的审计脚本。它把前面第 2、3 章提到的问题全部合并成一次检查输出一份文本报告并用退出码表示是否通过。对于 26K 张图运行时间约十几秒但能省掉后面训练时的大量排错时间。import csv import hashlib from pathlib import Path from PIL import Image root Path(pokemon_dataset) csv_path Path(pokemon_labels.csv) report [] errors [] # 1. 检查 CSV 完整性与 MD5 if csv_path.exists(): md5 hashlib.md5(csv_path.read_bytes()).hexdigest() report.append(fCSV MD5: {md5}) with csv_path.open(encodingutf-8-sig) as f: rows list(csv.DictReader(f)) report.append(fCSV 行数: {len(rows)}) else: errors.append(缺少 pokemon_labels.csv) # 2. 检查 CSV 中每个文件是否存在、尺寸是否符合 128x128 for row in rows: img_path root / row[filename] if not img_path.exists(): errors.append(f文件不存在: {row[filename]}) continue try: with Image.open(img_path) as img: if img.size ! (128, 128): errors.append(f尺寸异常 {img.size}: {row[filename]}) if img.format ! PNG: errors.append(f格式非 PNG: {row[filename]} ({img.format})) except Exception as e: errors.append(f解码失败 {row[filename]}: {e}) # 3. 检查类别样本数分布 from collections import Counter class_counts Counter(row[species] for row in rows) report.append(f类别数: {len(class_counts)}) report.append(f最少样本类: {class_counts.most_common()[-1]}) print(\n.join(report)) if errors: print(\n发现错误:) for e in errors[:50]: print( -, e) exit(1) else: print(\n全部通过可以开始训练。)这段脚本的核心是三层校验CSV 自身完整性、文件系统一致性、类别分布合理性。第一层通过 MD5 检查 CSV 是否在传输中损坏第二层遍历每个条目确认 PNG 存在且尺寸是 128x128避免训练到一半出现decode error或尺寸不一致导致collate_fn崩溃第三层用Counter找出样本数最少的类别方便你决定是否要做重采样。脚本的退出码可以接入 CI/CD让每次新增数据后自动跑一次比人工检查靠谱得多。你可以把这段代码保存成audit_dataset.py以后无论拿到什么图像数据集只要改成自己的root和csv_path就能直接复用。对于这份 Pokemon 数据集运行结束后如果所有检查都通过就可以放心开始训练了。本文还有配套的精品资源点击获取