
简介视觉TransformerViT凭借全局建模能力刷新了图像分类精度上限但高计算开销也限制了其在资源受限场景下的落地。如何在保证精度的同时降低注意力计算量成为工程优化的重要方向。DeBiFormer提出的可变形双级路由注意力DBRA机制通过粗粒度区域路由筛选关键区域再以可变形偏移精调采样位置实现了高效稀疏注意力。相较于固定窗口的Swin TransformerDBRA对不同形态目标更具适应性在较小模型规模下也能获得更高准确率。针对植物幼苗分类任务选用debi_tiny模型并优化训练策略可将验证集准确率稳定提升至82%以上。本文从环境配置、数据增强、混合精度训练到常见避坑点系统梳理了一套可复现的DeBiFormer实战流程为图像分类模型选型与调优提供参考。1. DeBiFormer实战把图像分类准确率从80%拉到82%以上我只做了这三件事做图像分类的同行应该都有这种感觉Vision TransformerViT系列虽然精度上限高但计算量也高得吓人尤其是在嵌入式设备或者 GPU 资源有限的环境下经常面临“模型跑不动、精度又上不去”的两难。我最近在做一个植物幼苗分类项目试过 ResNet、ViT、Swin Transformer效果都不太理想直到换上了 DeBiFormer 的 debi_tiny 模型才把验证集准确率稳定在了 82% 以上。这篇文章就基于我实际跑通的这套流程把数据准备、模型结构、训练配置和避坑点完整拆给你。DebiFormer 的核心创新在于一个叫可变形双级路由注意力DBRA的机制它能把注意力计算聚焦到真正有判别力的区域而不是像标准 ViT 那样在全局均匀分配算力这也是它能在较小模型规模下拿到更高精度的根本原因。如果你正在纠结怎么选分类模型、怎么部署 Transformer 到自己的数据集上这篇笔记应该能帮你省掉不少弯路。2. 为什么是DeBiFormerDBRA注意力机制与模型选型2.1 从Swin到DeBiFormer稀疏注意力到底在优化什么要理解 DeBiFormer 的优势得先回顾一下 Swin Transformer 的做法。Swin 把特征图划分成固定大小的窗口只在窗口内部做自注意力通过 shift 操作让不同窗口之间产生信息交互。这种设计的好处是计算复杂度从全局注意力的 O(N²) 降到了 O(N×W²)N 是 token 数量W 是窗口边长。但它的局限在于窗口是固定的、规则的几何划分并不一定贴合图像中物体的实际形状和分布。打个比方一棵幼苗的叶片是细长弯曲的如果恰好被切到两个不同的窗口里窗口内的注意力就很难捕捉到叶片的完整语义。DeBiFormer 的思路是在路由注意力BiFormer 的核心机制的基础上引入了可变形偏移。它的做法是先在粗粒度层面通过路由机制筛选出少量相关区域然后在这些区域内利用可变形偏移来微调采样位置让注意力 token 能自适应地落在更关键的像素位置上。换句话说DBRA 有两级筛选第一级决定“看哪几个区域”第二级决定“在每个区域内具体看哪个位置”。相比 Swin 的固定窗口这种机制对不同形状、不同尺度的目标更友好。2.2 debi_tiny vs debi_small vs debi_base显存与精度怎么权衡DeBiFormer 提供了 tiny、small、base 几个不同规模的变体主要区别在于 embedding 维度、Transformer block 层数和注意力头数。我这次使用的是 debi_tiny因为植物幼苗分类任务本身类别数不多我用的数据集是 12 类图像分辨率也不算高224×224tiny 级别的模型容量已经够用。模型变体Embedding维度Block层数注意力头数参数量约适用场景debi_tiny644212M小规模分类、资源受限debi_small1286443M中等规模分类、需更高精度debi_base256128108M大规模数据集、密集预测如果你用的是 4090 或者 A100可以尝试 debi_small训练速度并不会慢到无法接受。但如果你像我一样手头只有一块 3080 或者租的云 GPU 显存只有 10G 左右debi_tiny 会更稳妥。实测 debi_tiny 在 batch size 为 32、分辨率 224×224 的条件下显存占用大约 6~7G预留了一些余量给数据加载和混合精度训练。2.3 路由机制的实现从公式到代码DBRA 的完整实现涉及 Top-k 路由和可变形偏移两个阶段。这里我用伪代码表示其核心计算逻辑import torch import torch.nn.functional as F def dbrm_attention(q, k, v, region_num4, deform_scale2.0): q: (B, H, N, d) 查询向量 k: (B, H, N, d) 键向量 v: (B, H, N, d) 值向量 region_num: 每个查询选择的区域数 deform_scale: 可变形偏移的缩放系数 B, H, N, d q.shape region_size int(N ** 0.5) # 假设特征图是正方形 region_h region_w region_size // 4 # 将特征图划分为 4x4 区域 # 第一步区域级路由计算每个区域的平均 query/key 相似度 q_region q.reshape(B, H, region_h, region_w, d).mean(dim(2, 3)) k_region k.reshape(B, H, region_h, region_w, d).mean(dim(2, 3)) # 计算区域相似度矩阵简化为点积 region_scores torch.matmul(q_region, k_region.transpose(-2, -1)) region_scores region_scores / (d ** 0.5) # 每个查询区域选择 top-k 个相关区域 topk_indices torch.topk(region_scores, kregion_num, dim-1).indices # 第二步在选中的区域内对采样位置施加可变形偏移 offset torch.zeros_like(k) # 实际由子网络预测 offset torch.tanh(offset) * deform_scale # 根据偏移后的位置采样 key/value此处简化为伪代码 sampled_k k offset sampled_v v offset return torch.matmul(q, sampled_k.transpose(-2, -1)) sampled_v逻辑上分两步先计算区域间相关性选出 top-k 区域再在选中的区域内部做更精细的偏移和采样。offset 在实际代码中由一个轻量级卷积子网络预测训练过程中会自动学习到目标物形状的先验。参数说明region_num4表示每个查询区域只和 4 个区域做精细注意力其余区域直接忽略这是控制计算量的关键。deform_scale2.0是偏移幅度上限太大会导致采样位置偏离目标区域太小则退化成普通路由注意力适中的值是 1.5~2.5。在 PyTorch 中你不需要自己实现这套注意力直接用官方仓库里的models/debiformer.py即可但理解这个机制对调参非常有帮助。3. 环境搭建与数据集准备以植物幼苗分类为例3.1 依赖安装与项目结构DeBiFormer 基于 PyTorch 实现要求 PyTorch 1.9 及以上版本。我使用的是 PyTorch 2.0.1 CUDA 11.8实测可以正常编译运行。仓库里还依赖timm库注意版本需要大于 0.6。# 创建虚拟环境 conda create -n debi python3.9 -y conda activate debi # 安装 PyTorch以 CUDA 11.8 为例 pip install torch2.0.1 torchvision0.15.2 --index-url https://download.pytorch.org/whl/cu118 # 克隆项目并安装依赖 git clone https://github.com/rayleizhu/DeBiFormer.git cd DeBiFormer pip install timm0.9.2 tensorboard matplotlib # 如果有 apex 混合精度训练需求可选 pip install apex一个容易踩的坑DeBiFormer 中使用了自定义的 CUDA 算子来实现可变形偏移这部分需要编译。如果你用的 PyTorch 版本和仓库作者编译时的版本不一致会报undefined symbol之类的错误。解决办法是直接设置环境变量TORCH_CUDA_ARCH_LIST为你的显卡算力版本后重新编译比如 3080 是 8.6export TORCH_CUDA_ARCH_LIST8.6 python setup.py develop3.2 数据集划分与目录组织本次使用的数据集是植物幼苗分类共 12 个类别原始图像大小不一。我先统一 Resize 到 224×224并按 8:1:1 划分训练集、验证集、测试集。目录结构如下plant_seedlings/ ├── train/ │ ├── class1/ │ │ ├── img_001.jpg │ │ └── img_002.jpg │ ├── class2/ │ └── ... ├── val/ │ └── ... └── test/ └── ...划分脚本python split_dataset.py --data_dir plant_seedlings_raw --output_dir plant_seedlings --split 0.8 0.1 0.1# split_dataset.py 核心逻辑 import os import random import shutil from glob import glob def split_dataset(data_dir, output_dir, split_ratio(0.8, 0.1, 0.1)): classes [d for d in os.listdir(data_dir) if os.path.isdir(os.path.join(data_dir, d))] for cls in classes: cls_dir os.path.join(data_dir, cls) images glob(os.path.join(cls_dir, *.jpg)) glob(os.path.join(cls_dir, *.png)) random.shuffle(images) n_train int(len(images) * split_ratio[0]) n_val int(len(images) * split_ratio[1]) for i, img_path in enumerate(images): if i n_train: target os.path.join(output_dir, train, cls) elif i n_train n_val: target os.path.join(output_dir, val, cls) else: target os.path.join(output_dir, test, cls) os.makedirs(target, exist_okTrue) shutil.copy(img_path, os.path.join(target, os.path.basename(img_path)))这里注意一点随机划分前最好设置固定的random.seed(42)否则每次跑出来的结果不一致后续对比实验就失去了参考意义。另外如果你的类别数量很少比如只有 3~4 类建议split_ratio调整为 0.7:0.15:0.15并且开启数据增强来弥补数据量不足。3.3 配置 JSON 与数据加载项目使用.json文件配置数据路径这一点容易忽略。需要手动创建一个class.json格式如下{ train_root: data/plant_seedlings/train, val_root: data/plant_seedlings/val, test_root: data/plant_seedlings/test, num_classes: 12, input_size: 224, batch_size: 32, num_workers: 8, model_name: debi_tiny, pretrained: true, lr: 5e-4, epochs: 60, warmup_epochs: 5 }数据加载部分使用标准的 PyTorchImageFolder和DataLoader。注意num_workers的设置在 Windows 上不宜超过 4否则容易报DataLoader worker (pid 12345) exited unexpectedly的错误这通常是内存不足或 worker 数量过高导致的。from torchvision import datasets, transforms from torch.utils.data import DataLoader train_transform transforms.Compose([ transforms.RandomResizedCrop(224, scale(0.6, 1.0)), transforms.RandomHorizontalFlip(), transforms.ColorJitter(brightness0.3, contrast0.3, saturation0.2), transforms.ToTensor(), transforms.Normalize(mean[0.485, 0.456, 0.406], std[0.229, 0.224, 0.225]), ]) val_transform transforms.Compose([ transforms.Resize(256), transforms.CenterCrop(224), transforms.ToTensor(), transforms.Normalize(mean[0.485, 0.456, 0.406], std[0.229, 0.224, 0.225]), ]) train_dataset datasets.ImageFolder(data/plant_seedlings/train, transformtrain_transform) val_dataset datasets.ImageFolder(data/plant_seedlings/val, transformval_transform) train_loader DataLoader(train_dataset, batch_size32, shuffleTrue, num_workers8, pin_memoryTrue) val_loader DataLoader(val_dataset, batch_size32, shuffleFalse, num_workers8, pin_memoryTrue)这里有个细节训练集的RandomResizedCrop的scale参数我设置的是(0.6, 1.0)比默认的(0.08, 1.0)裁剪范围更保守。因为植物幼苗图像中幼苗主体占整张图的比例通常较大如果裁剪太狠模型可能会学到残缺的叶片特征。4. 训练实现与参数调优从损失函数到学习率策略4.1 完整训练脚本下面是我实际使用的训练脚本核心部分去掉了断点续训等干扰项保留了主干逻辑import torch import torch.nn as nn from torch.cuda.amp import GradScaler, autocast from models.debiformer import debi_tiny import json # 加载配置 with open(class.json, r) as f: cfg json.load(f) device torch.device(cuda if torch.cuda.is_available() else cpu) model debi_tiny(pretrainedcfg[pretrained], num_classescfg[num_classes]) model.to(device) criterion nn.CrossEntropyLoss(label_smoothing0.1) optimizer torch.optim.AdamW(model.parameters(), lrcfg[lr], weight_decay0.05) # 学习率预热 余弦退火 def warmup_cosine_lr(epoch, warmup_epochs, total_epochs, base_lr): if epoch warmup_epochs: return base_lr * (epoch 1) / warmup_epochs progress (epoch - warmup_epochs) / (total_epochs - warmup_epochs) return base_lr * 0.5 * (1 torch.cos(torch.tensor(progress) * 3.14159)) scaler GradScaler() best_acc 0.0 for epoch in range(cfg[epochs]): model.train() running_loss 0.0 for images, labels in train_loader: images, labels images.to(device), labels.to(device) optimizer.zero_grad() with autocast(): outputs model(images) loss criterion(outputs, labels) scaler.scale(loss).backward() scaler.step(optimizer) scaler.update() running_loss loss.item() * images.size(0) # 验证 model.eval() correct 0 total 0 with torch.no_grad(): for images, labels in val_loader: images, labels images.to(device), labels.to(device) outputs model(images) _, predicted torch.max(outputs, 1) total labels.size(0) correct (predicted labels).sum().item() acc 100.0 * correct / total print(fEpoch {epoch1}/{cfg[epochs]} | Loss: {running_loss/total:.4f} | Acc: {acc:.2f}%) if acc best_acc: best_acc acc torch.save(model.state_dict(), best_debi_tiny.pth)几个关键点逐个解释label_smoothing0.1这是提点的一个小技巧。植物幼苗分类中部分类别之间形态接近比如某种杂草和某种幼苗的叶片颜色几乎一样标签平滑可以让模型对错误预测不会过度自信实测约提升 0.5~1 个百分点。weight_decay0.05AdamW 配合相对大的 weight decay 对 ViT 类模型非常重要。默认的 0.01 稍显保守0.05 在一些开源代码中较为常见这个参数可以根据过拟合程度适当调整不必过度纠结。混合精度训练autocastGradScaler的组合能让训练速度提升约 30%且精度基本无损。debi_tiny 参数量 12M显存占用本身不高但混合精度可以让你把 batch size 调大如果 batch size 从 32 提升到 48 甚至 64训练的稳定性会更好。4.2 学习率策略与优化器选择我使用 warmup 5 个 epoch 余弦退火的学习率曲线初始学习率5e-4这个配置来自之前跑 Swin Transformer 的经验。需要说明的是ViT 模型的训练对学习率比较敏感学习率太高容易梯度爆炸太低则收敛缓慢。如果训练过程中出现准确率在 20% 附近震荡不上升的情况极有可能是学习率设置过大。建议将初始学习率调整为2e-4再试。反过来如果 loss 下降非常平稳但准确率提升缓慢可以适当增大学习率到8e-4。优化器方面AdamW 是首选几乎没有争议。不要使用普通的 SGD动量实验证明对 Transformer 类模型收敛速度明显偏慢。4.3 数据增强策略RandAugment 与 CutMix 的效果对比我在训练中对比了两种数据增强策略。第一种是简单的 RandomCrop Flip ColorJitter第二种是在此基础上叠加 RandAugmenttimm 提供。# 使用 timm 的 RandAugment from timm.data.auto_augment import rand_augment_transform train_transform_aug transforms.Compose([ transforms.RandomResizedCrop(224, scale(0.6, 1.0)), transforms.RandomHorizontalFlip(), rand_augment_transform(magnitude9, num_layers2), transforms.ToTensor(), transforms.Normalize(mean[0.485, 0.456, 0.406], std[0.229, 0.224, 0.225]), ])实验结果单独使用 RandAugment 相比 ColorJitter 提升约 1.2 个百分点的验证集准确率但代价是训练时间增加了约 15%因为数据增强的 CPU 开销变大。如果你在训练时发现 GPU 利用率不高CPU 反而成了瓶颈优先检查num_workers是否设置过低或者数据加载过程中是否存在磁盘 I/O 瓶颈。CutMix 和 MixUp 这一类输入混合增强我试过但对 debi_tiny 提升并不明显推测是这个数据集本身类别间区分度不高混出来的样本反而模糊了类别边界。如果换用 ImageNet 这种大粒度数据集CutMix 可能更有效。4.4 训练过程中的监控与调参判断训练时我主要看三个指标训练准确率、验证准确率和 loss 下降趋势。如果训练准确率已经到 95% 以上验证集只有 80%说明过拟合了。优先尝试增大 weight decay 到 0.08或者减少训练轮数。如果训练准确率和验证准确率差距不大但两者都在 70% 左右停滞说明欠拟合需要增加模型容量或优化数据增强策略。曾经有一次我用 debi_small 在相同配置下训练结果验证集准确率只有 79%反而比 debi_tiny 低了 3 个百分点。原因很简单小模型在这个数据规模下过拟合了而且 debi_small 需要更长的训练轮数才能收敛60 个 epoch 根本不够。不要盲目追求大模型先看数据量级。5. 训练避坑指南六个让我差点放弃的报错与玄学问题5.1 编译自定义算子报错 undefined symbol现象执行python setup.py develop编译成功后导入模型的瞬间报ImportError: /.../cdpr_deform_attn.so: undefined symbol: _ZN2at6TensorC1ERKS_。原因这是典型的 PyTorch 版本不匹配问题。仓库在编译时使用的是 PyTorch 1.12而我本地环境是 2.0.1导致二进制接口对不上。解决执行python setup.py clean --all后重新编译。如果仍然报错检查当前环境里是否有多个 PyTorch 版本conda 环境容易混用pip list | grep torch确认唯一的 torch 版本后再python setup.py develop。从那以后我每次换环境都会强制走一遍python -c import torch; print(torch.__version__)来确认版本。5.2 训练时 loss 突然变为 NaN现象正常训练到第 15 个 epochloss 突然从 1.2 跳变到 NaN验证集准确率也直接崩掉。原因混合精度训练下梯度值因为 FP16 的表示范围有限而溢出经典原因有两个一是学习率过大梯度更新过猛二是某些样本的数值范围异常比如输入图像有全黑的异常值。解决先关闭autocast跑 10 个 epoch 确认是否是精度问题。如果是降低初始学习率到2e-4同时在损失函数之前对 logits 做一次数值稳定处理。我这里干脆加了梯度裁剪torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm5.0)这个max_norm5.0意味着梯度向量的 L2 范数会被裁剪到 5 以内既不影响模型正常收敛又能防止梯度爆炸带来 NaN。5.3 加载预训练权重时报 shape mismatch现象使用pretrainedTrue加载 ImageNet 权重时报错说最后一层全连接层的权重形状不匹配1000 vs 12。原因debi_tiny的 ImageNet 预训练权重是为 1000 类设计的最后一层head.weight的形状和新的分类任务不一致。解决在加载权重时显式忽略head.weight和head.biaspretrained_dict torch.load(debi_tiny_imagenet.pth, map_locationcpu) model_dict model.state_dict() pretrained_dict {k: v for k, v in pretrained_dict.items() if k in model_dict and model_dict[k].shape v.shape} model_dict.update(pretrained_dict) model.load_state_dict(model_dict)这样只加载主干部分的预训练参数分类头从头训练。实际上对于植物幼苗这种和 ImageNet 分布差异较大的数据集预训练权重的收益主要在前几层边缘、纹理特征全连接层重训反而是常规操作。5.4 验证集准确率一直在 50% 左右徘徊现象训练流程没问题loss 也正常下降但验证集准确率始终上不去就在 50% 附近震荡。数据集有 12 个类别50% 意味着模型只学到了一点皮毛。原因检查代码后发现验证集做的是transforms.Resize(256) transforms.CenterCrop(224)但训练时是RandomResizedCrop。两者做了相同的缩放但是验证集 CenterCrop 时截取的区域可能不是目标所在的位置。植物幼苗图像中目标通常不在正中心CenterCrop 正好把幼苗截掉了这属于典型的训练验证分布不一致。解决换用FiveCrop或者TenCrop增强验证集评估的鲁棒性val_transform transforms.Compose([ transforms.Resize(224), transforms.FiveCrop(224), transforms.Lambda(lambda crops: torch.stack([transforms.ToTensor()(crop) for crop in crops])), transforms.Normalize(mean[0.485, 0.456, 0.406], std[0.229, 0.224, 0.225]), ])不过这样推理时需要特殊处理 output 形状更简单的方案是放弃 CenterCrop 直接 Resize 到 224×224虽然可能轻微变形但至少不会丢目标。5.5 多 GPU 训练时 batch_size 没缩放现象从单卡切到两张卡训练loss 下降速度明显变慢准确率上去得更慢。原因batch_size在 DataLoader 里写的是 32但两张卡各自 32实际全局 batch 变成了 64。学习率仍然保持5e-4相当于学习率相对 batch 变小了。解决使用DataParallel或DistributedDataParallel时学习率需要按照 batch 变化进行缩放。常见做法是线性缩放规则学习率 × (新 batch / 旧 batch)。如果 batch 从 32 变到 64学习率调到1e-3这张卡 64 张卡 32 同理。lr cfg[lr] * (total_batch_size / 32.0)5.6 分类结果中某一类准确率极低现象训练结束后查看每类准确率发现blackgrass这一类的准确率只有 35%比平均准确率低了 50 个百分点。原因一方面是这个类别的样本数确实少只占总样本的 5%另一方面是这个类别的图像颜色纹理和另一类loose silky bent高度相似模型很难区分。解决使用类别平衡采样from torch.utils.data import WeightedRandomSampler class_counts [len(train_dataset.imgs) for cls in train_dataset.classes] target_samples max(class_counts) class_weights [target_samples/cnt for cnt in class_counts] sample_weights [0] * len(train_dataset) for idx, (_, label) in enumerate(train_dataset.imgs): sample_weights[idx] class_weights[label] sampler WeightedRandomSampler(sample_weights, num_sampleslen(train_dataset), replacementTrue) train_loader DataLoader(train_dataset, batch_size32, samplersampler)使用WeightedRandomSampler后少数类样本每个 epoch 被采样到的概率会更高能有效缓解类别不平衡导致的单类准确率过低。实测 blackgrass 的准确率从 35% 提升到了 58%整体准确率没有明显下降。6. 进阶技巧类别激活图可视化与混淆矩阵分析模型训练完成后准确率达标只是第一步。实际工程中我还要确认模型是真的学到了叶片纹理特征还是靠背景信息来判断类别。这一点在植物幼苗分类场景里尤其重要因为即使同一类幼苗在不同光照、不同土壤背景下拍摄的图像可能差异很大。如果模型是靠背景猜类别换一套拍摄环境就全崩了。我常用的工具是 Grad-CAM它能生成类别激活图告诉我们模型做出决策时主要关注图像的哪些区域。实现代码如下import cv2 import numpy as np from pytorch_grad_cam import GradCAM from pytorch_grad_cam.utils.model_targets import ClassifierOutputTarget from pytorch_grad_cam.utils.image import show_cam_on_image, preprocess_image def visualize_gradcam(model, img_path, target_class): img cv2.imread(img_path) img cv2.cvtColor(img, cv2.COLOR_BGR2RGB) img_resized cv2.resize(img, (224, 224)) img_normalized img_resized / 255.0 input_tensor preprocess_image(img_resized, mean[0.485, 0.456, 0.406], std[0.229, 0.224, 0.225]) target_layers [model.blocks[-1].norm1] cam GradCAM(modelmodel, target_layerstarget_layers) grayscale_cam cam(input_tensorinput_tensor, targets[ClassifierOutputTarget(target_class)]) grayscale_cam grayscale_cam[0, :] visualization show_cam_on_image(img_normalized, grayscale_cam, use_rgbTrue) cv2.imwrite(gradcam_output.jpg, visualization)我在项目中挑选了三张不同类别的幼苗图像做可视化发现模型对叶片边缘和叶脉位置的激活值明显高于背景区域这说明模型学到了有判别力的结构特征。如果你生成的热力图集中在图像角落或背景上那就要警惕了模型大概率是在偷懒需要重新设计数据增强或检查数据标注质量。另一种诊断工具是混淆矩阵。它能直观地看出哪些类别之间容易混淆帮助我判断是模型能力不足还是数据标注本身有问题。import matplotlib.pyplot as plt import seaborn as sns from sklearn.metrics import confusion_matrix all_preds [] all_labels [] with torch.no_grad(): for images, labels in val_loader: images images.to(device) outputs model(images) _, preds torch.max(outputs, 1) all_preds.extend(preds.cpu().tolist()) all_labels.extend(labels.tolist()) cm confusion_matrix(all_labels, all_preds) plt.figure(figsize(12, 10)) sns.heatmap(cm, annotTrue, fmtd, cmapBlues, xticklabelstrain_dataset.classes, yticklabelstrain_dataset.classes) plt.xlabel(Predicted) plt.ylabel(True) plt.savefig(confusion_matrix.png, dpi150)通过混淆矩阵我发现chickweed和cleavers这两类经常被互相误判。翻看原始图像后发现这两类幼苗在早期生长阶段的外形高度相似连人眼都很难区分。这种情况属于标注本身的模糊性处理方案有两个一是合并相似类别降低分类粒度二是引入更细的类别细分标注需要领域专家参与。最终我在项目中保留原分类粒度但在训练时加大这两类的采样权重才把互相误判的比例降下来。从那次之后我每次训练完模型都会强制执行一遍 Grad-CAM 可视化和混淆矩阵分析的流程确认模型学到的特征可靠后再交付。希望这套方法也能帮到你少走一些我当年踩过的弯路。本文还有配套的精品资源点击获取