简介:英文字母手语图像分类数据集包含约两万六千张已标注的手语字母图像,覆盖二十八个类别,并划分好训练集与测试集,适用于图像分类模型训练、迁移学习以及算法效果对比。压缩包内共两千个文件,其中一千九百九十八个jpg图像按类别存放,一个py脚本用于可视化数据集,一个json文件提供类别映射与标签配置,整体约851.4MB。已有两百三十二人学习浏览,训练集与测试集分别按类别存放,可直接按路径加载使用。借助自带show脚本可快速预览样本,配合CNN分类网络改进思路,可开展网络结构对比、参数调优等实验。JSON文件中的类别映射便于核对标签关系,整体目录结构清晰,适合需要标准数据集支撑手语识别研究的学生与工程师。
1. 英文字母手语图像分类数据集:26类、26,000张,开箱能做什么
一个手语字母识别的活儿,要不要从零开始拍数据?绝大多数人卡在这一步。英文字母手语图像分类数据集正好把这步省了——26个类别,A到Z每个字母对应一类,总计约26,000张图片,而且已经标注好,不需要自己框、自己打标签。它的价值对刚入图像分类的人尤其明显:你可以直接跨过数据清洗,把精力放在模型训练、参数调节和避坑上,用最短时间把“训练自己的数据集”这条流程完整跑通。对工程师来说,这个规模的数据也适合做算法预实验,或者当迁移学习、数据增强策略的验证场。平均每个类别约一千张,属于“够用但并不宽裕”的典型中小规模图像分类数据集,能训出可用的模型,同时也逼着你去处理过拟合、类别不均衡这些非常现实的问题。
2. 拿到数据集先别急着训练:目录结构、标签划分与数据体检
很多人拿到已标注数据集,解压完就开训,结果后期排查问题时才发现标签对不上、图片尺寸不统一、某个类别样本特别少。磨刀不误砍柴工,先用半小时做一次数据体检,后面能省几天的排查时间。
2.1 拿到数据集做三件事:数文件、看标签、查尺寸
这类数据集常见的组织方式有两种:一种是按类别分目录,比如A/xxx.jpg、B/xxx.jpg;另一种是给你一个 CSV 或 JSON 清单,里面每行是文件名,标签。不管哪种形式,第一步先跑个脚本统计一下分布。
import os from collections import Counter data_root = "path/to/dataset" # 方式一:按文件夹组织 class_counts = {} for cls in sorted(os.listdir(data_root)): cls_path = os.path.join(data_root, cls) if os.path.isdir(cls_path): class_counts[cls] = len(os.listdir(cls_path)) # 方式二:CSV 清单 # import pandas as pd # df = pd.read_csv("labels.csv") # class_counts = df["label"].value_counts().to_dict() print(f"类别数: {len(class_counts)}") print(f"总样本数: {sum(class_counts.values())}") for cls, cnt in sorted(class_counts.items()): print(f"{cls}: {cnt}")这段逻辑不复杂,关键是输出两个信息:类别数和总样本数。标题说约 26,000 张,那你解压后第一件事就是核对这两个数字。若总数和预期差得远,先别往下走,可能下载不完整或归档有问题。
接下来查图片尺寸,顺便清理损坏文件。常见的坑是数据集中混入了非 RGB 图片、灰度图、甚至是无法解码的坏文件,训练时Image.open()会直接报错。
from PIL import Image import os bad_files = [] size_counter = Counter() for root, _, files in os.walk(data_root): for f in files: if not f.lower().endswith((".jpg", ".jpeg", ".png")): continue p = os.path.join(root, f) try: with Image.open(p) as im: im.verify() # 先验证文件完整性 size_counter[im.size] += 1 except Exception: bad_files.append(p) print(f"损坏文件数: {len(bad_files)}") print(f"最常见的尺寸: {size_counter.most_common(5)}")Image.verify()是 PIL 自带的快速校验接口,能识别大部分截断损坏的图片。尺寸统计则帮你确定后续要不要统一做 resize。如果 90% 以上是同一个尺寸,后面预处理会省事很多;如果大小很离散,就得靠Resize和RandomResizedCrop兜底。
2.2 划分训练/验证/测试集:按类别比例做分层抽样
图像分类里划分数据要遵守一个原则:按类别做分层抽样,而不是直接random.shuffle后切分。原因很直接——如果某个字母的图片恰好都排在前面,简单切分可能导致这个类别全部进了训练集,验证集里完全见不到这个字母,指标当然很难看。
下面这个脚本把数据按 8:1:1 划分成 train / val / test,并照顾到每个类别的比例。
import os import random import shutil random.seed(42) split_ratio = {"train": 0.8, "val": 0.1, "test": 0.1} src_root = "path/to/dataset" out_root = "path/to/split_dataset" for cls in sorted(os.listdir(src_root)): cls_src = os.path.join(src_root, cls) if not os.path.isdir(cls_src): continue imgs = [f for f in os.listdir(cls_src) if f.lower().endswith((".jpg", ".jpeg", ".png"))] random.shuffle(imgs) n_train = int(len(imgs) * split_ratio["train"]) n_val = int(len(imgs) * split_ratio["val"]) for split, subset in zip( ["train", "val", "test"], [imgs[:n_train], imgs[n_train:n_train + n_val], imgs[n_train + n_val:]] ): out_dir = os.path.join(out_root, split, cls) os.makedirs(out_dir, exist_ok=True) for f in subset: shutil.copy2(os.path.join(cls_src, f), os.path.join(out_dir, f)) print("划分完成,各集合样本数:") for split in split_ratio: total = sum(len(os.listdir(os.path.join(out_root, split, cls))) for cls in os.listdir(os.path.join(out_root, split))) print(f"{split}: {total}")两个参数要留意。第一是random.seed(42),固定随机种子才能保证别人用同一份数据能复现你的结果,这个习惯建议从第一次实验就养成。第二是split_ratio,8:1:1 是中小数据集的常用比例;如果总样本量更少,可以放宽到 7:2:1,把更多数据留给验证集,避免验证指标波动太大。
2.3 标签正确性抽查:肉眼核对放在训练前
已标注不等于标注一定完美。26 个字母里,像 M 和 N、U 和 V、S 和 T 这些手语形状非常接近,标注员自己也容易搞混。建议训练前从每个类别随机抽 5 到 10 张图,拼成一张网格图,花几分钟肉眼过一遍。
import matplotlib.pyplot as plt from PIL import Image import random random.seed(0) fig, axes = plt.subplots(26, 5, figsize=(10, 55)) for row, cls in enumerate(sorted(os.listdir(src_root))): cls_path = os.path.join(src_root, cls) imgs = os.listdir(cls_path) for col in range(5): img_path = os.path.join(cls_path, random.choice(imgs)) ax = axes[row][col] ax.imshow(Image.open(img_path)) ax.axis("off") axes[row][0].set_ylabel(cls, fontsize=14) plt.tight_layout() plt.savefig("label_check_grid.png", dpi=100)这个环节看起来原始,但价值极大。你肉眼能发现三类问题:明显的错标、手势方向不一致、以及某个字母的图片风格和其他类差异过大。这些发现会影响后面的数据增强策略——比如手语是有方向性的,如果发现数据集中左右手混用,那水平翻转这类增强就要慎用,否则会把“左手 A”翻成“右手 A”,反而引入错误样本。
3. 用 ResNet18 在 PyTorch 上跑通 26 字母分类:训练脚本与日志解读
数据准备好之后,下一步是用一个相对保守的模型把基线跑出来。为什么强调先跑基线?因为你得知道一个不加花哨技巧的模型能做到什么水平,后面做数据增强、调参才有对照物。
3.1 为什么从 ResNet18 起步,而不是一上来就上大模型
26,000 张图、26 个类别,这是一个典型的中小规模图像分类任务。最常见的做法是用 ResNet18 这类轻量级 CNN 先打底,原因是它足够小、训练速度快,不容易在数据量不够的情况下严重过拟合,而且预训练权重随处可得。你的目标不是用最深的网络刷新榜单,而是先拿到一个稳定的基线。
很多初学者一上来就用 EfficientNet-B4 甚至 ViT,结果在这个数据量上训练时间翻倍,验证集准确率反而没有明显优势。最新的图像分类模型确实在 ImageNet 这类大数据集上表现优异,但 26,000 张图对 ViT 这类需要海量数据的结构来说偏少了。先用 ResNet18 跑通全流程,把数据、代码、评估链路都验证好,再谈换更强模型也不迟。
3.2 最小可复现训练脚本
下面这个脚本按最常见的工程习惯组织:继承Dataset写数据加载、用torchvision加载预训练模型、最后保存验证精度最高的权重。
import os import torch import torch.nn as nn from torch.utils.data import Dataset, DataLoader from torchvision import transforms, models from PIL import Image class AlphabetDataset(Dataset): def __init__(self, root, transform=None): self.samples = [] self.transform = transform classes = sorted(os.listdir(root)) self.class_to_idx = {cls: i for i, cls in enumerate(classes)} for cls in classes: cls_dir = os.path.join(root, cls) for f in os.listdir(cls_dir): if f.lower().endswith((".jpg", ".jpeg", ".png")): self.samples.append((os.path.join(cls_dir, f), self.class_to_idx[cls])) def __len__(self): return len(self.samples) def __getitem__(self, idx): path, label = self.samples[idx] img = Image.open(path).convert("RGB") if self.transform: img = self.transform(img) return img, label transform_train = transforms.Compose([ transforms.Resize((256, 256)), transforms.RandomResizedCrop(224, scale=(0.7, 1.0)), transforms.RandomRotation(10), transforms.ToTensor(), transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]), ]) transform_val = transforms.Compose([ transforms.Resize((256, 256)), transforms.CenterCrop(224), transforms.ToTensor(), transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]), ]) train_ds = AlphabetDataset("path/to/split_dataset/train", transform_train) val_ds = AlphabetDataset("path/to/split_dataset/val", transform_val) train_loader = DataLoader(train_ds, batch_size=64, shuffle=True, num_workers=4) val_loader = DataLoader(val_ds, batch_size=64, shuffle=False, num_workers=4) model = models.resnet18(weights=models.ResNet18_Weights.DEFAULT) model.fc = nn.Linear(model.fc.in_features, 26) criterion = nn.CrossEntropyLoss() optimizer = torch.optim.AdamW(model.parameters(), lr=1e-4, weight_decay=5e-4) device = torch.device("cuda" if torch.cuda.is_available() else "cpu") model = model.to(device) best_acc = 0.0 for epoch in range(30): model.train() running_loss = 0.0 for imgs, labels in train_loader: imgs, labels = imgs.to(device), labels.to(device) optimizer.zero_grad() outputs = model(imgs) loss = criterion(outputs, labels) loss.backward() optimizer.step() running_loss += loss.item() model.eval() correct = 0 total = 0 with torch.no_grad(): for imgs, labels in val_loader: imgs, labels = imgs.to(device), labels.to(device) outputs = model(imgs) _, preds = torch.max(outputs, 1) correct += (preds == labels).sum().item() total += labels.size(0) val_acc = correct / total print(f"Epoch {epoch+1}/30 | Loss: {running_loss/len(train_loader):.4f} | Val Acc: {val_acc:.4f}") if val_acc > best_acc: best_acc = val_acc torch.save(model.state_dict(), "best_alphabet_model.pth") print(f"Best Val Acc: {best_acc:.4f}")几个参数值得展开说。Resize((256, 256))再RandomResizedCrop(224)是常见做法:先放大到 256,再做随机裁剪到 224,等于给模型提供了一些空间平移和尺度变化的样本。scale=(0.7, 1.0)控制随机裁剪保留原图的比例,手语手势的关键信息集中在手部,裁剪采样范围设得太小容易把手指裁掉,0.7 是相对保守的设置。batch_size=64在 ResNet18 上显存占用约 2 到 3GB,一般单卡都能跑;显存小就降到 32。lr=1e-4配合AdamW是迁移学习场景下比较稳妥的起点,比从头训练的常用学习率要低,因为预训练权重不需要大步伐更新。代码里保存模型用的是验证集准确率作为标准,而不是最后一个 epoch 的权重,这是避免遗忘“最佳模型在中间某个 epoch 出现”这个事实的后悔药。
3.3 训练日志里哪些数字该盯紧
训练时会输出两类关键指标:训练 Loss 和验证准确率。Loss 持续下降说明模型在拟合;验证准确率逐步上升说明泛化有效。你需要警惕两种异常。第一种是 Loss 正常下降但验证准确率停滞——大概率是数据增强不够或者模型容量不足。第二种是训练 Loss 降到接近 0 而验证准确率不涨反跌——这是经典的过拟合信号,后面会专门展开。
另一件值得做的是记录每个 epoch 的训练时长和显存占用。如果单个 epoch 超过两分钟,就应该考虑减小 batch size 或换更轻的模型;如果 GPU 利用率一直上不去,瓶颈可能在num_workers线程太少或 CPU 解码太慢。
4. 把验证准确率从 90 上下推到 98:数据增强、学习率与类别不均衡
基线模型能跑通以后,真正的调参才刚开始。多数人会在这一阶段发现准确率停在 90 左右上不去,这是正常的,接下来的四个方向能把指标继续往上顶。
4.1 数据增强参数:裁剪范围、旋转角度与颜色抖动怎么设
手语识别的数据增强和普通物体分类有个显著区别:手势有方向性。水平翻转会把左手比划的 A 变成右手比划的 A,对某些类别来说这就是错误标签。实际项目里我会刻意不开RandomHorizontalFlip,或者只在确认数据集全部是右手且语义对方向不敏感时才开。
推荐一套偏保守的增强配置:
transform_train = transforms.Compose([ transforms.Resize((256, 256)), transforms.RandomResizedCrop(224, scale=(0.7, 1.0), ratio=(0.9, 1.1)), transforms.RandomRotation(degrees=10), transforms.ColorJitter(brightness=0.2, contrast=0.2, saturation=0.1, hue=0.05), transforms.RandomErasing(p=0.2, scale=(0.02, 0.08)), transforms.ToTensor(), transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]), ])逐个参数解释。RandomRotation(10)只允许正负 10 度的旋转,手语的识别对手指朝向敏感,转多了会把 U 和 V 的区分度抹掉。ColorJitter的亮度对比度抖动帮助模型抵抗不同光照条件——手语数据的拍摄环境往往不统一,有的室内打光、有的户外自然光。RandomErasing是让模型不要过度依赖局部像素,随机遮挡一小块区域让模型学会看整体手势结构,scale=(0.02, 0.08)控制遮挡面积占比,太大容易把关键的手指部位整个遮掉。ratio=(0.9, 1.1)限制 RandomResizedCrop 的长宽比变化幅度,避免拉伸变形。
注意这套增强只用于训练集,验证集和测试集要保留干净的Resize + CenterCrop,否则评估指标会被增强噪声污染,没法公平对比实验组。
4.2 优化器与学习率:AdamW 配合 cosine 退火
基线上的lr=1e-4固定学习率能跑,但不是最优解。实际项目中常见做法是配合 cosine 退火,让学习率在训练后期自动衰减到接近 0,这样能在损失函数曲面的谷底附近精细收敛,减少震荡。
from torch.optim.lr_scheduler import CosineAnnealingLR optimizer = torch.optim.AdamW(model.parameters(), lr=1e-4, weight_decay=5e-4) scheduler = CosineAnnealingLR(optimizer, T_max=30, eta_min=1e-6) # 训练循环里每个 epoch 结束后加上 scheduler.step()T_max设为总 epoch 数,表示学习率在 30 个 epoch 内从 1e-4 余弦下降到 1e-6。eta_min是退火最低值,设得太低会浪费后面的训练时间,太低的学习率几乎不再更新参数。想再激进一点,可以把前两个 epoch 设成 warmup:先用 1e-5 跑两个 epoch 让预训练权重适应新数据分布,再切回 1e-4。在 PyTorch 里可以简单用torch.optim.lr_scheduler.LinearLR叠加实现,但中小数据集上不做 warmup 差别也不会特别大。
相比固定学习率,cosine 退火通常能让验证准确率提升 1 到 3 个百分点,代价几乎为零,是我会优先推荐的默认配置。
4.3 类别不均衡的应对:加权采样与焦点损失
先统计一下每个类别的样本数,如果发现某些字母只有六百多张、另一些有一千二百多张,就属于轻度类别不均衡。26 个字母里,J、Q、Z 这类手语形状复杂或者不常用的字母确实可能样本偏少。
最直接的解决方法是按类别频率的倒数给每个样本配权重,让采样器在每次迭代里多抽少数类。PyTorch 里用WeightedRandomSampler实现:
import numpy as np from torch.utils.data import WeightedRandomSampler labels = [label for _, label in train_ds.samples] class_counts = np.bincount(labels) sample_weights = 1.0 / class_counts[labels] sampler = WeightedRandomSampler(sample_weights, num_samples=len(labels), replacement=True) train_loader = DataLoader(train_ds, batch_size=64, sampler=sampler, num_workers=4)关键参数是replacement=True,它表示允许重复采样同一张图,少数类会被反复抽到,同时多数类样本在整个 epoch 里不会全部被覆盖。代价是训练速度略有下降,且如果某个类样本特别少,模型可能对这几张图过拟合。另一个替代方案是改用Focal Loss,它通过调制因子让模型把注意力放在难分类的样本上,但需要多调一个gamma超参数。对于 26 类、每类约千张的规模,WeightedRandomSampler的收益往往更直接。
4.4 微调还是从头训:预训练权重与冻结策略
有了预训练权重,最常见的做法是全量微调,即加载resnet18(weights=...)后让它所有层都参与训练。这在 26,000 张图的数据集上是可行的,因为数据量已经足够支撑微调。如果你的数据量更少,比如只有每类两三百张,那最好冻结前面的卷积层,只训练最后的全连接层——防止数据量不足导致前面层学到的通用特征被破坏。
实践中具体怎么操作:先把fc层换掉,只训练fc两三个 epoch,观察验证损失是否下降;再解冻最后两层卷积参与训练。这套 freeze 到 unfreeze 的策略,是迁移学习里最稳的节奏。真正从零训练 ResNet18 我一般不建议——在这个数据集规模下,从零训的收敛速度明显更慢,最终精度也更依赖随机种子,属于吃力不讨好。
说到最新的图像分类模型,如果你后续想尝试 Transformer,ViT-small 或 Swin-T 也值得试,但对这个数据量要格外注意:它们的数据饥饿特性比 CNN 更明显,26,000 张图往往不够,需要更强的正则化手段兜底。我的建议是先把手头的 ResNet18 压榨到极限,再考虑更换架构,否则很难分清是模型的问题还是数据的问题。
5. 手语分类最常见的五个坑:现象、原因、解决办法
这章写的每条都是真实项目里会碰到的、能快速排查和解决的问题。按“现象 → 原因 → 解决”来梳理,方便你对照排查。
5.1 验证集上 M/N、U/V 总是互认
现象:模型的整体准确率在 95% 以上,但看混淆矩阵时 M 和 N、U 和 V 之间频繁互认,错误样本高度集中在几对字母上。
原因:手语字母里 M 和 N 的手形非常接近——同样是四指弯曲,只是弯曲的指节位置不同;U 和 V 更是只有食指和中指是否分开的细微差别。对 224x224 的输入分辨率来说,区分这些差异需要模型关注非常细小的局部纹理,而标准数据增强里的随机裁剪和旋转会进一步模糊这些细节。
解决:最有效的手段是把输入分辨率提到 256 或 320,给模型更多像素去分辨细微的指间差异。其次是去掉RandomRotation或把角度降到 5 度以内——M/N 的差异对旋转特别敏感。如果还不行,就针对易混淆类别对收集额外样本,做局部关键点裁剪后单独微调一个二分类器。
5.2 训练准确率 99.8%,验证集只有 73%
现象:训练 Loss 降到 0.01 以下,训练集准确率接近满分,验证集却一直在 70% 上下。
原因:这就是典型的过拟合。26,000 张图对 26 个类别来说不算少,但对 ResNet18 这类模型容量来说,如果数据增强太弱、正则化不足,模型完全可以把训练集背下来而不是学到手势的通用特征。
解决:先确认数据增强有没有生效,别让transform_train里的RandomHorizontalFlip这类不明显但有害的增强污染训练。然后补上三件套:更强的基础增强(多尺度裁剪 + 颜色抖动)、weight_decay调到 1e-3、以及早停——当验证集准确率连续 8 个 epoch 不提升时,直接保存当前最优权重并停止训练。也可以临时把model.fc层的输出前面加一个nn.Dropout(0.3),限制全连接层的拟合能力。
5.3 模型记住背景而不是手
现象:训练时准确率很高,但把图片里手以外的区域用纯色块遮挡后,模型预测完全乱了;或者测试集换成新的背景环境后准确率骤降。
原因:数据集中如果大部分图片的背景都比较相似,比如都是同一种桌面纹理或同一面墙,模型很容易拿背景当捷径。这是图像分类在中小数据集上的老问题,手语数据尤其容易中招——因为手势拍摄往往是固定机位、固定背景。
解决:最实用的办法是训练时用随机裁剪和多尺度变换去打乱背景和手势的耦合关系。更彻底的方法是引入手部检测或分割模型,先把手部区域抠出来再送入分类器,这时你实际上用上了目标检测的思路——对这类结构化任务,先把主体位置固定住,分类器的负担会小很多。
5.4 某些字母类别样本量明显偏少
现象:统计类别分布后,Q、J、Z 各自只有几百张,而 A、B、C 等字母超过一千张;这几个类别的验证准确率也明显偏低。
原因:数据收集阶段很难保证均衡,不常用的字母在现实中出现频率更低。模型对这些类别见过太少,自然学不稳。
解决:优先用WeightedRandomSampler兜底(代码见 4.3),让每个 batch 里少数类出现的概率近均衡。其次对少数类额外做增强,比如把RandomRotation的角度放宽到 15 度、RandomErasing的遮挡比例调大,等于用增广把每个少数类“复制”出更多变体。注意不要给多数类也开同样强度的增强,否则数据量更大的多数类会进一步挤压少数类的学习空间。
5.5 别人复现你的精度差一截
现象:把训练脚本和数据划分代码发给同事,对方在同样数据上跑出来的结果比你低两三个百分点。
原因:最隐蔽的原因是随机性没被完全固定。random.seed只控制了你脚本里的随机数,PyTorch 的 CUDA 卷积算子、DataLoader 的 worker 线程都会引入额外随机性;另外,数据增强里的裁剪和旋转随机性也对最终精度有影响。
解决:彻底固定随机源,在训练脚本最前面加一行:
torch.manual_seed(42) torch.cuda.manual_seed_all(42)再把划分好的训练/验证/测试集文件清单导出成一个split.json或 CSV,和代码一起提交。这样别人无论在哪台机器上跑,用的都是同一批训练样本、同一批验证样本,增强的随机性也受控。最后在 README 里写清楚预处理全链路——包括统一的分辨率、归一化参数、增强开关——这是让结果可复现的最后一步。
6. 用混淆矩阵决定下一轮迭代:易混淆类别分析与模型导出
6.1 输出混淆矩阵,找出模型最犹豫的类别对
不要只盯着整体准确率看,混淆矩阵才是诊断模型短板的核心工具。在测试集上跑完一轮预测后,用下面的代码生成 26x26 的混淆矩阵:
import numpy as np from sklearn.metrics import confusion_matrix import torch model.eval() all_labels, all_preds = [], [] with torch.no_grad(): for imgs, labels in test_loader: imgs = imgs.to(device) outputs = model(imgs) _, preds = torch.max(outputs, 1) all_labels.extend(labels.cpu().numpy()) all_preds.extend(preds.cpu().numpy()) cm = confusion_matrix(all_labels, all_preds, labels=range(26)) # 找出混淆最严重的类别对:把对角线置0,取矩阵Top5 np.fill_diagonal(cm, 0) top_pairs = np.argwhere(cm > 0) top_pairs = sorted(top_pairs, key=lambda x: -cm[x[0], x[1]])[:5] class_names = sorted(os.listdir(test_root)) for i, j in top_pairs: print(f"{class_names[i]} -> {class_names[j]}: {cm[i, j]} 次")这个脚本输出的就是模型最常搞错的几对字母,对应前面 5.1 里说的易混淆问题。有了这张表,迭代方向就清楚了——你不需要对所有 26 个类别盲目优化,集中火力处理最挤的那两三对。
6.2 针对易混淆对的三个进阶手段
拿到混淆对之后,常见的做法有三个。第一个是把输入分辨率继续往上抬,从 224 到 256 甚至 320,增加手指缝隙的像素,U/V 这种只差几像素开合的边缘情况会显著改善。第二个是引入手部区域定位,用轻量的手部检测模型先裁剪出掌心区域,把分类模型的输入从“含背景的整图”换成“聚焦手部的局部图”,模型就不用自己隐式学习空间注意力。第三个是给两个混淆类单独训一个小分类器,比如只拿 M 和 N 的样本训一个 2 类 ResNet,判断阈值可以调得很精细,再回掺到主分类器的结果里做二次修正。
6.3 把训练好的分类器导出为 ONNX 做落地
模型调好后,常见的落地方式是导出 ONNX,便于部署到 CPU 服务器或边缘设备。PyTorch 导出 ONNX 只要几行:
model.eval() dummy_input = torch.randn(1, 3, 224, 224).to(device) torch.onnx.export( model, dummy_input, "alphabet_resnet18.onnx", input_names=["input"], output_names=["logits"], dynamic_axes={"input": {0: "batch"}, "logits": {0: "batch"}} )导出后建议先用onnxruntime验证一遍输出与 PyTorch 结果一致:
import onnxruntime as ort import numpy as np sess = ort.InferenceSession("alphabet_resnet18.onnx") x = np.random.randn(1, 3, 224, 224).astype(np.float32) onnx_out = sess.run(None, {"input": x})[0]dynamic_axes让导出的模型支持可变 batch,部署时传入单张或多张图都能跑。这里有个容易踩的细节:ONNX 导出时模型必须处于eval模式,否则 BN 层和 Dropout 层会用训练行为,导出的模型在推理时表现异常。我自己第一次跑这个任务的时候,就是没看混淆矩阵,闷头调了三天增强参数,结果 M/N 的互认率一点没变——后来才发现问题不在增强不够,而是模型根本没有足够的像素去看到手指缝的差别,把输入分辨率提上去后立竿见影。这个教训让我养成了“每次迭代先看混淆矩阵再做决定”的习惯,希望帮到你。
本文还有配套的精品资源,点击获取