简介:基于视觉变换网络(ViT)的自闭症谱系障碍儿童脸部分析检测项目,面向医疗AI与计算机视觉研究者,解决ASD早期诊断中面部特征客观量化困难的问题。项目利用ViT的自我注意力机制,从面部表情、眼睛注视、头部姿态等维度学习自闭症相关表征,可作为临床辅助工具,适合有一定深度学习基础、希望将Transformer架构应用于医学图像分类的开发者。资源共39个文件,包括17个Python脚本,覆盖模型搭建、训练、评估、可视化流程;12个YAML配置对应ViTASD小、中、大不同规模模型;另有4张PNG结果图及说明文档,压缩包仅3.42MB,结构紧凑且易于部署。目前已有158人学习,代码与配置分离,附带数据加载与训练脚本,可直接复现项目实验,也可基于此进行迁移学习和二次开发,为医疗AI落地提供完整参考。
1. 基于ViT的自闭症谱系障碍脸部分析:这项技术到底解决了什么
自闭症谱系障碍(ASD)的诊断长期依赖行为量表和临床观察,医生通过ADOS、CARS等工具打分,主观性很强,而且一个孩子从初筛到确诊往往要等上数月。近些年有团队把目光转向面部分析——能不能用深度学习模型,从儿童面部图像直接给出ASD 风险概率?这个项目就是基于 Vision Transformer(ViT)实现的一套完整方案:它把面部图像切成patch序列,用自注意力机制学习ASD儿童面部特征,并附带了从数据加载、训练到OOD(分布外)检测、注意力可视化的全套工程文件。适合正在做医疗影像分类、ViT 微调落地,或想了解如何在真实场景中评估模型置信度的开发者。
2. ViTASD 模型架构与配置文件:三档参数和一个关键设计
2.1 为什么选 ViT 而不是 ResNet:局部 patch 与全局注意力的取舍
做儿童面部分析,常见的思路是直接用 ResNet 或 EfficientNet 提取全局特征后接全连接分类。这类 CNN 模型的问题是感受野逐步扩大,浅层特征偏向局部纹理,高层特征才具备全局语义,但 ASD 相关的面部特征往往是弥散的——眼动模式异常、特定肌肉紧张度、面部结构比例差异,这些信号分布在图像的不同区域,需要模型跨区域建立关联。
ViT 的做法是把一张图切成固定大小的 patch,例如 224x224 的图像切成 16x16 的 patch 序列,每个 patch 线性映射成 token 后送入 Transformer 编码器。自注意力机制让每个 patch 都能直接关注到其他所有 patch,跨区域的依赖建模是显式的。对 ASD 检测来说,模型可以在早期层就学会“嘴角区域和眼部区域的联合异常”这种跨区域特征,而不需要像 CNN 那样层层堆叠才能融合。
项目里的models/vitasd.py实现了这个结构。核心代码如下:
import torch import torch.nn as nn from timm.models.vision_transformer import VisionTransformer class ViTASD(nn.Module): def __init__(self, img_size=224, patch_size=16, embed_dim=768, depth=12, num_heads=12, num_classes=2, use_sngp=False): super().__init__() # 复用 timm 的 VisionTransformer 主干 self.backbone = VisionTransformer( img_size=img_size, patch_size=patch_size, embed_dim=embed_dim, depth=depth, num_heads=num_heads, num_classes=0, # 不接原分类头 ) self.use_sngp = use_sngp if use_sngp: # 光谱归一化高斯过程层,用于 OOD 不确定性估计 self.classifier = SNGPHead(embed_dim, num_classes) else: self.classifier = nn.Linear(embed_dim, num_classes) def forward(self, x): features = self.backbone(x) # [B, embed_dim] return self.classifier(features)这段代码的关键在于num_classes=0截断了 timm 预训练模型的分类头,只取特征向量;SNGPHead是项目自定义的分类层(lib/sngp.py),它用光谱归一化约束权重矩阵的谱范数,并引入高斯过程近似,使得模型对分布外样本能输出高不确定性分数而不是盲目给一个高置信度类别。这个设计直接影响后面的 OOD 评估效果。
2.2 配置文件拆解:small / base / large 到底差在哪
configs/目录下有三个 ViTASD 配置:config_vitasd_small.yaml、config_vitasd_base.yaml、config_vitasd_large.yaml,还有一个config_vitasd_base_attonly.yaml。前三个对应三档模型规模,attonly版本代表只保留注意力模块、去掉 SNGP 头的消融配置。
以config_vitasd_base.yaml为例:
model: name: vitasd img_size: 224 patch_size: 16 embed_dim: 768 depth: 12 num_heads: 12 use_sngp: true data: dataset: autism data_root: ./datasets/ASD train_batch_size: 32 eval_batch_size: 64 num_workers: 8 train: epochs: 100 lr: 3e-5 weight_decay: 0.01 warmup_epochs: 5 scheduler: cosine eval: metrics: [accuracy, precision, recall, f1, auc] ood_threshold: 0.5三个配置的差异集中体现在embed_dim、depth、num_heads三个参数上:small 版本是embed_dim=384, depth=6, num_heads=6,适合在单卡上快速跑通流程验证数据;base 版本对应 ViT-Base 结构,是项目默认的实验配置;large 版本是embed_dim=1024, depth=24, num_heads=16,需要至少两张 24GB 显存的卡。这里有个容易忽略的点:lr: 3e-5不是随意给的——医疗图像分类模型微调时,如果直接从 ImageNet 预训练权重开始,学习率过大会破坏低层特征,过小则收敛极慢,3e-5 配合 cosine 衰减是这类任务里比较稳妥的起点。
2.3 SNGP 模块:为什么医疗场景需要它
常规分类模型在推理时只有 softmax 概率,但 softmax 的置信度并不可靠——一个模型没见过旋转角度异常的面部图像,照样可能给出 0.95 的 ASD 概率。医疗辅助诊断场景里,这种过度自信是不能接受的。
SNGP(Spectral-normalized Neural Gaussian Process)的做法是:在最后一层全连接之前,对隐藏层权重做光谱归一化,让模型的预测函数满足高斯过程先验,然后通过拉普拉斯近似估计预测不确定性。实现要点在lib/sngp.py:
class SNGPHead(nn.Module): def __init__(self, in_features, num_classes, num_inducing=128): super().__init__() self.fc = nn.Linear(in_features, num_features) # 先升维 self.spec_norm = nn.utils.spectral_norm(self.fc) # 光谱归一化 self.gp_layer = GaussianProcessLayer( num_features, num_classes, num_inducing ) def forward(self, x): h = self.spec_norm(x) return self.gp_layer(h)实际训练时,这个模块的前向输出除了 logits 还会返回一个协方差矩阵,推理阶段用mean ± 2 * std作为置信区间。如果你的数据里混入了大量模糊图像或者不同采集设备拍摄的照片,这个不确定性分数比单纯的最大概率值更有参考价值。
3. 训练流程与数据准备:从 AffectNet 预训练到 ASD 微调
3.1 数据集加载:autism_dataset.py中做了什么
datasets/autism_dataset.py负责加载 ASD 儿童面部图像。它的核心逻辑是读取标注文件,把图像路径和类别标签对齐,然后做训练/验证集划分。源码里值得注意的一点是它支持两种标注格式:CSV 格式(两列:image_path,label)和文件夹格式(按类别分子目录)。此外,它还处理了类别不平衡问题——ASD 数据集普遍存在正常儿童样本远多于 ASD 样本的情况,项目里通过weights参数在采样器层面做了加权,而不是简单的过采样。
class AutismDataset(Dataset): def __init__(self, data_root, split='train', transform=None, class_weights=None): self.samples = self._load_samples(data_root, split) self.transform = transform self.class_weights = class_weights def _load_samples(self, data_root, split): # 读取 CSV 标注,过滤掉不存在的图像文件 # 返回 [(img_path, label), ...] ... def __getitem__(self, idx): img_path, label = self.samples[idx] img = Image.open(img_path).convert('RGB') if self.transform: img = self.transform(img) return img, label这里有个工程细节:_load_samples里做了文件存在性检查,训练前把损坏的图片直接过滤掉,而不是等训练中途报错。处理医疗图像数据时,不同来源的数据集经常混入灰度图、损坏文件或带水印的图,这一步能省很多排查时间。
3.2 数据增强策略:augment.py的边界设计
lib/augment.py里定义了训练时的数据增强管线。和常规的 ImageNet 训练不同,面部图像不能随便做水平翻转——人脸左右不对称性本身可能是特征,但 ASD 检测并不依赖左右不对称,所以翻转是允许的;真正要小心的是旋转角度,超过 15 度的旋转会引入非自然的面部姿态,模型学到的是姿态伪影而不是病理特征。
train_transform = transforms.Compose([ transforms.Resize((256, 256)), transforms.RandomResizedCrop(224, scale=(0.8, 1.0)), transforms.RandomHorizontalFlip(p=0.5), transforms.ColorJitter(brightness=0.2, contrast=0.2), transforms.RandomRotation(degrees=10), transforms.ToTensor(), transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]) ])RandomResizedCrop(scale=(0.8, 1.0))的取值是经过考量的:scale 低于 0.8 会裁掉眼睛或嘴部区域,破坏面部关键结构;RandomRotation(degrees=10)模拟的是儿童在采集时轻微转头的情况,低于常见的数据增强强度。医疗图像增强的第一原则是“不能制造现实中不存在的样本”,这里的每一组参数都对应着真实采集场景中的方差来源。
3.3 训练主脚本:train.py与train_affectnet.py的关系
项目提供了两个训练入口:train.py是直接在 ASD 数据集上微调,train_affectnet.py是先在 AffectNet 表情数据集上做预训练,再用预训练权重初始化 ASD 模型。
AffectNet 是面部表情识别领域的大规模数据集,包含约 45 万张标注了 8 类表情的图像。在它上面预训练的价值在于:模型先学会通用的面部表征(眼睛、嘴部动作单元的联动关系),再迁移到 ASD 检测时,微调数据需求量大幅下降。实际命令如下:
# 第一步:在 AffectNet 上预训练 python train_affectnet.py \ --config configs/config_affectnet_base.yaml \ --gpus 2 # 第二步:用预训练权重初始化,在 ASD 数据上微调 python train.py \ --config configs/config_vitasd_base.yaml \ --pretrained ./checkpoints/affectnet_base_best.ckpt \ --gpus 1第二步里有个参数值得解释:--gpus 1是因为微调阶段图像量小,单卡足够,但如果你用的是 large 配置,需要改成--gpus 2并把 batch size 减半。项目里train.py基于 PyTorch Lightning 实现,日志自动写入lightning_logs/目录,配合 TensorBoard 可以实时看训练曲线。
微调阶段的关键超参数是冻结策略。常见做法是前 10 个 epoch 冻结 backbone,只训练分类头;10 个 epoch 后再解冻全部参数,用 1/10 的学习率做全量微调。这个项目里没有显式做冻结,而是直接把全模型学习率设为3e-5——对 ViT-Base 来说这个值足够小,不会剧烈破坏预训练特征,但如果你在更大规模的数据集上从头训练,这个策略就不适用了。
3.4 训练时的监控指标
医疗分类不能只看准确率。ASD 检测中阴性样本比例高,一个把所有样本都预测为正常的模型也能拿到很高的 accuracy。项目里train.py默认记录了 precision、recall、F1 和 AUC,其中 AUC 是最值得关注的——它不受分类阈值影响,衡量的是模型对正负样本的区分能力。
# 训练过程中的典型输出 Epoch 10/100: loss=0.4832, acc=0.8214, precision=0.7638, recall=0.7952, f1=0.7792, auc=0.8735如果你看到 acc 在涨但 recall 在跌,说明模型在往“保守预测”偏——负样本容易分对,正样本开始漏。这时候需要调整类别权重或降低分类阈值。
4. 评估与 OOD 检测:模型准确率背后的置信度陷阱
4.1eval.py的完整评估流程
tools/eval.py承担模型评估职责,除了计算基础指标,还会输出每个类别的混淆矩阵、按性别和年龄段分层的结果。运行方式:
python tools/eval.py \ --ckpt ./lightning_logs/ViTASD-B/best.ckpt \ --config configs/config_vitasd_base.yaml \ --data_root ./datasets/ASD/test评估逻辑的代码骨架如下:
model = ViTASD.load_from_checkpoint(args.ckpt) model.eval() all_preds, all_labels, all_uncertainties = [], [], [] with torch.no_grad(): for batch in test_dataloader: images, labels = batch logits, uncertainty = model(images) probs = torch.softmax(logits, dim=-1) all_preds.extend(probs.argmax(dim=-1).cpu().numpy()) all_labels.extend(labels.cpu().numpy()) all_uncertainties.extend(uncertainty.cpu().numpy())这段代码里model(images)返回两个值而不是一个,这是 SNGP 头带来的变化。uncertainty代表模型对这个预测有多“没把握”,取值在 0 到 1 之间。评估时除了算 accuracy,还要算一个关键指标:排除高不确定性样本后,模型准确率提升了多少。
4.2 OOD 评估:ood_evaluator.py检测模型不知道什么
lib/ood_evaluator.py实现的是分布外(Out-of-Distribution)检测评估。它的任务是把已知类别图像(ASD数据集)和未知分布图像(比如健康儿童的照片、模糊图像、卡通人脸混在一起)区分开。实现中,OOD 评分函数使用不确定性分数和最大 softmax 概率的组合:
class OODEvaluator: def __init__(self, ood_threshold=0.5, score_type='uncertainty'): self.threshold = ood_threshold self.score_type = score_type def compute_ood_score(self, uncertainty, max_prob): # 组合不确定性分数和置信度 if self.score_type == 'uncertainty': return uncertainty elif self.score_type == 'combined': return 0.7 * uncertainty + 0.3 * (1 - max_prob) def evaluate(self, known_unc, unknown_unc, known_prob, unknown_prob): # 计算 AUROC、FPR@95TPR 等 OOD 检测指标 ...这里暴露了 SNGP 设计的价值:普通 ViT 模型没有不确定性输出,只能用最大 softmax 概率做 OOD 判断,但 softmax 概率在分布外样本上往往虚高。SNGP 的协方差估计让分布外样本的不确定性显著高于分布内样本,使得FPR@95TPR这个指标(在保持 95% 已知样本被正确保留的前提下,错误接收了多少分布外样本)大幅下降。
4.3 评估结果怎么看
拿到评估输出后,重点看三个数字:
| 指标 | 含义 | 可接受范围 |
|---|---|---|
| AUC | 区分 ASD 与非 ASD 的能力 | > 0.85 |
| FPR@95TPR | OOD 检测的误接率 | < 0.30 |
| 不确定性中位数(分布外) | 模型对未知样本的平均怀疑程度 | 显著高于分布内 |
如果最后一个指标不明显高于分布内的不确定性中位数,说明 SNGP 头没有正常工作,常见原因是学习率太大导致光谱归一化层训练不充分。
5. 常见问题与避坑:训练和评估中的六个实际问题
5.1 训练 loss 不下降,accuracy 一直卡在 0.7 左右
现象:微调 ViTASD-Base 时,训练 loss 在前 5 个 epoch 内从 0.8 降到 0.5 后就不再下降,验证集 accuracy 始终在 0.7 附近徘徊,precision 和 recall 差异很大。
原因:这个现象通常是类别不平衡导致的。ASD 数据集中正负样本比例可能达到 1:4 甚至 1:6,模型学到的最优策略是把所有样本预测为多数类(正常儿童),accuracy 自然停在多数类占比附近。另一个常见原因是学习率设置过低——3e-5对从头训练来说太小,对微调来说需要配合足够长的 warmup。
解决:先检查类别分布,在AutismDataset中打印torch.bincount(labels)确认比例。然后做两步调整:在损失函数中传入class_weights,给少数类更高的权重;或者把学习率提高到1e-4并加长 warmup 到 10 个 epoch。如果已经用了类别权重但情况没改善,检查数据增强里的ColorJitter强度是否过大——过强的颜色扰动会抹掉肤色和肌肉紧张度这些关键信号。
5.2 模型推理表现正常,但 OOD 分数完全没区分度
现象:把正常测试集中的图像换成模糊图片或不同设备拍摄的面部照片后,模型对所有输入的 uncertainty 输出都差不多,OOD 检测的 AUROC 在 0.5 附近,相当于随机猜测。
原因:这通常是 SNGP 头没有正确训练。项目里use_sngp: true时,train.py需要额外在损失中加一个 GP 的 KL 散度项,有些配置下这个项被遗漏了,导致高斯过程层退化成普通全连接层。另一个原因是从不用 SNGP 的 checkpoin 加载预训练权重时,classifier层的权重被随机初始化,需要更长训练时间才能适配。
解决:确认train.py的损失计算中包含kl_div = model.gp_layer.kl_divergence()并乘以一个较小的系数(如 0.1)加入总损失。其次,检查 SNGP 头是否在训练 30 个 epoch 之后再从 checkpoint 加载——如果是,务必确认 checkpoint 里use_sngp标记和当前配置一致,否则分类头参数不匹配。
5.3 CUDA OOM:large 配置单卡直接显存溢出
现象:用config_vitasd_large.yaml配置训练,batch size 设为 32,启动时报CUDA out of memory。
原因:ViT-Large 有 24 层 transformer,embedding 维度 1024,单张 24GB 显卡在 batch size 32 的场景下无法容纳。ViT 的内存占用和序列长度的平方成正比——224x224 图像切 16x16 patch 得到 196 个 token,序列长度为 196,如果你未来换到更高分辨率,显存增长会非常快。
解决:三选一:batch size 降到 8 并使用梯度累积(每 4 步累加一次,等效 batch size 仍为 32);换用config_vitasd_base.yaml完成实验后再用 large 做最终验证;开启混合精度训练(PyTorch Lightning 的--precision 16),显存占用直接减半。
5.4 微调 ASDA 数据时验证集波动大,指标忽高忽低
现象:验证集的 accuracy 在 0.75 和 0.85 之间大幅波动,AUC 也时好时坏,训练曲线像锯齿。
原因:ASD 数据集的采集标准不统一,不同机构提供的数据光照条件、拍摄角度差异大,验证集可能混入了一些低质量样本。更关键的是验证集本身太小(可能只有几十到一两百张图),每次评估的方差极大,不足以反映模型真实水平。
解决:不要每 epoch 都评估,改为每 5 个 epoch 评估一次,并取多次评估的平均值作为当前模型的真实表现。另外在AutismDataset划分数据时用分层采样,确保验证集中正负样本比例和训练集一致,避免某次划分后验证集里全是正常样本,acc 虚高。
5.5 注意力可视化结果看起来完全随机
现象:运行visualization_attention.py后输出的注意力热力图分布在整个图像上,面部区域没有任何重点,看不出模型关注了什么。
原因:模型未收敛或数据预处理不当。如果你在训练初期就做可视化,自注意力权重还没有被训练信号驱动,自然呈现接近均匀分布的状态。另一种可能是Normalize时的均值和标准差与你加载图像的实际分布不一致,导致模型输入分布偏离训练时的分布。
解决:确保加载的是训练完成且验证指标正常的 checkpoint。同时检查输入图像是否经过了相同的预处理管线——项目里visualization_attention.py默认从config_vitasd_base.yaml中读取预处理参数,如果你换过数据集但没换配置,就会出现输入分布偏移。
5.6 推理时模型对同一张图多次预测结果不一致
现象:同一张测试图片,多次运行eval.py得到不同的预测标签和概率。
原因:模型处于训练模式而不是评估模式。eval.py中如果遗漏了model.eval()或torch.no_grad(),BatchNorm 层和 dropout 层会持续更新和随机屏蔽,导致输出有随机性。ViT 虽然没有 BatchNorm(用的是 LayerNorm),但 SNGP 头中如果有 dropout 层,同样产生这个问题。
解决:检查推理脚本中是否调用了model.eval()。一个值得养成的习惯是在torch.inference_mode()上下文里做推理,它比torch.no_grad()更严格,会关闭所有与梯度追踪相关的机制,同时提升推理速度。
6. 注意力可视化的落实方法:从热力图到可解释的医疗依据
模型的预测结果要能被临床医生接受,不能只给一个“ASD 概率 0.87”的数字,必须告诉医生模型是依据什么特征做出的判断。visualization_attention.py就是干这个事的——它提取 ViT 最后一个 transformer block 的注意力权重,与 CLS token(分类标记)的注意力关联,叠加回原始图像生成热力图。
# visualization_attention.py 的核心流程 def visualize_attention(image_path, model, save_path): # 1. 预处理输入图像 image = load_image(image_path) # 读取并缩放到 224x224 tensor = preprocess(image) # Normalize + ToTensor # 2. 注册 forward hook 捕获注意力权重 attention_maps = [] def hook_fn(module, input, output): # output 包含 (attn_weights, attn_output) attention_maps.append(output[0].detach()) hook = model.backbone.blocks[-1].attn.register_forward_hook(hook_fn) # 3. 模型推理 logits, uncertainty = model(tensor.unsqueeze(0)) # 4. 取 CLS token 的注意力均值,reshape 回 14x14 的空间分辨率 attn = attention_maps[0][0, 0, 1:, :].mean(dim=0) # 224/16=14 attn_map = attn.reshape(14, 14) # 5. 上采样并叠加到原图 ...注意第 4 步的下标[0, 0, 1:, :]:第一个 0 是 batch 维度,第二个 0 是注意力头编号——如果你对多个注意力头做平均,特征会更平滑但可能丢失关键信息;1:的意思是去掉 CLS token 与自身计算的注意力权重,只看它对图像 patch 的关注程度。
实际运行中,你可以对比不同模型的注意力图:正常儿童的面部热力图通常均匀分布在眼睛、口鼻区域,ASD 儿童样本的热力图往往集中在某个局部区域且注意力熵更低——这符合 ASD 儿童对面部信息加工策略异于常人的临床观察。在你的论文或报告中放上这样一组对比图,说服力远大于单独放准确率数字。
让我补充一个实操层面的建议:部署时把visualization_attention.py封装成一个 HTTP 服务,每次预测连同热力图一起返回,方便医生在查看报告中确认模型依据。封装时要注意torch.inference_mode()上下文和model.eval()的调用顺序,确保服务端线程安全。
从那以后,我每次跑这类医疗图像项目都会强制走一遍“训练 — 评估 — 可视化 — 人工复核”的完整闭环,并在交付时把不确定性阈值写进接口文档,而不是只交出准确率。希望这个项目的拆解和这些经验能帮你更快地落地自己的 ViT 方案——别让模型在测试集上看起来很美,却在真实数据面前翻车。
本文还有配套的精品资源,点击获取