[Nature 2025]只有25%的标注数据,医学大模型如何反超GPT-4o?一篇用GAN破解“数据荒”与“偏见”的佳作
做医疗AI的同行们,大概都经历过这样一种无奈:医院系统里躺着千万级别的电子病历(EHR),看似是一座数据金矿,但真要拿来训练模型时却发现,这座金矿根本挖不动。
为什么?因为医生随手写的临床笔记是非结构化的,里面充满了极其丰富但也极度随意的病情推演。要想让模型看懂,就需要顶级的专科医生来做人工标注。但现实是,专科医生的时间比黄金还贵,海量标注根本不现实。退而求其次用系统里现成的ICD诊断编码吧,噪声又大得惊人,拿青光眼来说,ICD编码的特异性甚至不到50%。
用充满偏见、标签稀缺又带着粗糙噪声的数据去喂养一个监督学习模型,结果只能是灾难。模型不仅学不会精准诊断,还会把现实中对特定性别、种族或年龄群体的偏见成倍放大。
面对这个几乎无解的死局,最近发表在《npj Digital Medicine》上的一篇论文给出了一种极其优雅的破局思路。这篇名为“Clinically informed semi-supervised learning improves disease annotation and equity from electronic health records”的研究,巧妙地借用了生成对抗网络(GAN)的思想,不仅用极少的标注数据打败了全监督的BERT基线,还在零样本测试中把GPT-4o甩在了身后。

左上角的对比揭示了医疗AI面临的真实困境——海量沉睡的未标注电子病历与极度稀缺的专家标注数据。Ci-SSGAN 的核心思路正是利用这一巨大落差,通过融合人口统计学信息的生成器,在半监督框架下最大限度榨取无标签临床文本的潜在价值。
不生成文字,而是生成“高维灵魂”
一听到GAN,很多人的第一反应是用它来生成虚假的医疗影像,或者是编造假病历。但在自然语言处理领域,直接让GAN生成长篇大论的临床笔记是非常吃力不讨好的,语法断裂和语义崩塌是常态。
这篇论文的聪明之处在于,它完全放弃了在“文本空间”里玩泥巴,而是把战场转移到了“特征空间”。

论文提出了一个叫 Ci-SSGAN(临床先验驱动的半监督GAN)的框架。这个框架的生成器不吐文字,它吃进去的是三样东西的组合:100维的随机噪声、3维的人口统计学信息(年龄、性别、种族),以及一条实实在在的、高达768维的未标注临床文本特征(由BioClinical BERT提取)。
有了这768维的真实临床语义作为“锚点”,生成器就像是有了一副骨架。它不再是凭空捏造数据,而是在真实患者特征的流形空间附近进行合理的微扰与探索。它输出的是高度逼真的“合成患者特征向量”。这就完美规避了生成文本带来的语义破碎问题。
重点看左侧生成器的输入端。与常规 GAN 仅仅依赖随机噪声不同,这里的输入是一个由 100 维噪声、3 维人口统计学特征以及 768 维无标签文本特征拼接而成的 871 维张量。这 768 维的真实临床特征犹如给生成器装上了一副“临床骨架”,从底层逻辑上杜绝了凭空捏造带来的语义崩塌。
给判别器装上“双头大脑”
传统GAN的判别器就是个保安,只负责看大门,分辨来的人是真实的还是伪造的。
但在这项研究里,判别器是个十足的“双头打工仔”。它共享了底层的特征提取层,但长出了两个脑袋:一个脑袋继续玩对抗游戏,分辨特征向量是来自真实病历还是生成器捏造的;另一个脑袋则是一个专注的分类器,专门盯着那些仅有的、带有医生真实标注的病历,去精准分类6种青光眼亚型。
通过这种半监督的对抗训练,哪怕只有极少量的真实标签,判别器也能在海量无标签数据和生成器不断造出的“逼真特征”的对决中,被迫学到泛化能力极强的决策边界。

这是一个极具临床视觉冲击力的对比。左侧是 Ci-SSGAN 的预测结果,热力图对角线异常干净,跨亚型误判率极低;右侧则是直接调用医院系统 ICD 诊断编码的结果,大面积的错位色块直接暴露了传统结构化标签在精细亚型分类上的粗糙与不可靠。
用非广延熵给“少数派”撑腰
医疗数据永远是长尾的。有些罕见的青光眼亚型,或者某些边缘人口特征的患者,在数据集里少得可怜。普通的模型一跑,很容易就陷入了“模式崩溃”,只认得最常见的白人开角型青光眼患者。
为了对冲这种不公平,作者不仅在采样阶段做了极其严格的分层对抗平衡,在不确定性评估上更是下了一步险棋。
论文抛弃了大家最爱用的香农熵,转而使用了非广延的 Tsallis 熵,并且把熵参数 q 设定在了极端的 0.25。
在数学意义上,当 q 远小于 1 时,这个函数会对那些概率极小的罕见事件给予指数级的放大补偿。这就像是在模型的损失评估中安装了一个巨大的放大镜。即使模型对某个黑人女性患有罕见继发性青光眼的初始预测概率极低,它产生的不确定性信号也会被急剧放大。模型被这种机制死死盯住,再也无法糊弄过去,必须对边缘群体给出确定的、可靠的预测。
这也直接催生了论文中提出的全新公平性指标——PV Score(均等违背分数),严格约束模型在阳性和阴性预测值上的群体差异。

这张特征空间的可视化图是模型打破“模式崩溃”的有力铁证。可以看到,在 Ci-SSGAN 学习到的高维空间中,不同种族、性别和年龄段的患者特征并没有模糊重叠,而是形成了边界清晰、且保留了真实类内方差的多样性特征簇。
降维打击般的实验结果
这种架构上的降维打击,在实验数据上体现得淋漓尽致。
研究团队仅仅使用了 25% 的专家标注数据,Ci-SSGAN 的准确率就达到了惊人的 0.871,AUROC 飙升至 0.956,全面碾压了使用 100% 数据进行全监督微调的 Base BERT 和 Bio BERT 模型。更关键的是,它在不同种族、性别和年龄段上的表现出奇地均衡。
在和当今顶流大语言模型的 Zero-shot(零样本)对比中,Ci-SSGAN 也毫不手软。面对复杂的临床推演,GPT-4o 的准确率仅为 0.641,而 Ci-SSGAN 稳稳站上了 0.840,直接反超。这再次证明了,在高度垂直的医学细分领域,通过小数据+生成式半监督架构精心雕琢的领域模型,依然有着通用大模型难以企及的穿透力。

即使只用 25% 的标注数据,Ci-SSGAN(灰色柱形)在所有族裔、性别和年龄亚组中的准确率都稳稳压制了完全监督的基线模型。更关键的是图表右侧——代表算法偏见的 PV 分数条带被压缩到了极低值,真正实现了分类精度与弱势群体公平性的双赢。
一些可能深入的地方
尽管该研究在架构设计和临床价值上表现出色,但从深层模型优化和特征解释性的角度来看,仍有几个值得探讨和深入的维度:
1.论文采用了 UMAP 聚类和基于梯度的 token 归因分析(Gradient-weighted token attribution)来展示模型关注了临床实体 。然而,对于神经网络的视觉化和可解释性来说,仅仅展示 token 权重的重新归一化 尚显薄弱。
2.作者在讨论部分也坦诚,当前模型缺乏显式的推理机制 。在未来的迭代中,将潜在文本表示映射到可解释的概念或自然语言解释(Auxiliary reasoning head),才是打通黑盒医疗 AI 的关键 。
3.论文开源的代码只包含纯推理的部分,并没有包含模型最核心的训练循环(Training Loop)。论文中至关重要的对抗损失(Adversarial Loss)、多样性损失(Diversity Loss)、焦点损失(Focal Loss)以及为了公平性设计的对齐采样逻辑,并没有在这个开源出来的推理文件中体现。在当前提供的代码库中,并没有包含 Tsallis 熵的实现。预测逻辑中,模型仅仅是使用标准的 Softmax 函数将 Logits 转化为概率,并直接提取最大概率值作为置信度(np.max(probabilities, axis=1))。
代码解析和补充
由于开源代码不完整,这里就核心模型部分进行注释,还是比较简单的:
import torch
import torch.nn as nn
import torch.nn.functional as F
from transformers import AutoModel
# ============================================================================
# 1. 文本编码器 (Text Encoder)
# ============================================================================
class TextEncoder(nn.Module):
def __init__(self, model_name='emilyalsentzer/Bio_ClinicalBERT', hidden_dim=768, dropout=0.3):
super().__init__()
# 加载预训练的医疗垂直领域 BERT (Bio_ClinicalBERT)
self.bert = AutoModel.from_pretrained(model_name)
self.dropout = nn.Dropout(dropout)
# 核心设计:自定义注意力池化层 (Attention Pooling)
# 放弃了直接使用 BERT 的 [CLS] token,转而构建一个独立的注意力机制。
# 这种设计大幅提升了模型的可解释性,使其能够对输入序列中的关键临床实体(如“色素性”、“假性剥脱”)分配更高的权重。
self.attention_pool = nn.Sequential(
nn.Linear(hidden_dim, 768),
nn.Tanh(),
nn.Dropout(0.1),
nn.Linear(768, 1)
)
def forward(self, input_ids, attention_mask):
# 提取逐词的隐藏状态
outputs = self.bert(input_ids=input_ids, attention_mask=attention_mask)
last_hidden = outputs.last_hidden_state
# 计算每个 token 的注意力权重,并使用 softmax 归一化
attention_weights = self.attention_pool(last_hidden).squeeze(-1)
attention_weights = F.softmax(attention_weights * attention_mask, dim=1)
# 依据注意力权重对所有 token 的特征进行加权求和,输出 768 维的句子级临床表征
pooled_output = torch.sum(last_hidden * attention_weights.unsqueeze(-1), dim=1)
return self.dropout(pooled_output)
# ============================================================================
# 2. 临床先验驱动生成器 (Clinically Informed Generator)
# ============================================================================
class ClinicallyInformedGenerator(nn.Module):
def __init__(self, noise_dim=100, demographic_dim=3, text_embed_dim=768, hidden_dims=[1024, 768]):
super().__init__()
layers = []
# 维度设计:100(高斯噪声) + 3(年龄/种族/性别) + 768(未标注的真实临床文本特征) = 871维
input_dim = noise_dim + demographic_dim + text_embed_dim
# 构建全连接网络,加入 BatchNorm1d 稳定对抗训练的梯度流
for hidden_dim in hidden_dims:
layers.extend([
nn.Linear(input_dim, hidden_dim),
nn.BatchNorm1d(hidden_dim),
nn.LeakyReLU(0.2, inplace=True),
nn.Dropout(0.2)
])
input_dim = hidden_dim
# 输出层:将特征映射回 768 维空间,使用 Tanh 激活函数将输出边界约束在 [-1, 1]
# 以完美对齐 BERT 提取出的真实特征流形
layers.append(nn.Linear(input_dim, text_embed_dim))
layers.append(nn.Tanh())
self.generator = nn.Sequential(*layers)
def forward(self, noise, demographics, unlabeled_text_embeddings):
# 核心创新操作:在第一维度 (Feature Dimension) 进行硬拼接 (Concatenation)
# 这迫使生成器必须以真实的临床语义特征(unlabeled_text_embeddings)为骨架来进行合成
x = torch.cat([noise, demographics, unlabeled_text_embeddings], dim=1)
return self.generator(x)
# ============================================================================
# 3. 双头判别器 (Dual-head Discriminator)
# ============================================================================
class Discriminator(nn.Module):
def __init__(self, text_embed_dim=768, demographic_dim=3, num_classes=6, hidden_dims=[512, 256]):
super().__init__()
layers = []
# 维度设计:768(真实或生成的文本特征) + 3(条件:人口统计学特征) = 771维
input_dim = text_embed_dim + demographic_dim
# 共享特征提取层 (Shared Feature Extractor)
# 这部分网络同时服务于分类任务和真假辨别任务,充当底层表征学习器
for hidden_dim in hidden_dims:
layers.extend([
nn.Linear(input_dim, hidden_dim),
nn.LeakyReLU(0.2, inplace=True),
nn.Dropout(0.3)
])
input_dim = hidden_dim
self.feature_extractor = nn.Sequential(*layers)
# 【输出头 1】 分类头 (Classifier Head):负责输出 6 种青光眼亚型的 Logits
self.classifier = nn.Linear(input_dim, num_classes)
# 【输出头 2】 对抗头 (Source Head):单神经元输出,负责判断输入特征是真实的(1)还是生成的(0)
self.source_classifier = nn.Linear(input_dim, 1)
def forward(self, text_embeddings, demographics):
# 强制特征与人口统计学变量进行晚期融合 (Late Fusion)
# 相当于对判别器下达指令:“不仅要看这份病历像不像真的,还要看它像不像属于这个特定种族/性别/年龄的真病历”
x = torch.cat([text_embeddings, demographics], dim=1)
features = self.feature_extractor(x)
# 同步输出多分类任务的 Logits 与对抗任务的 Logits
class_logits = self.classifier(features)
source_logits = self.source_classifier(features)
# 返回提取到的潜在特征(features)供后续计算多样性损失或进行 UMAP 聚类分析
return class_logits, source_logits, features
就论文中提到的Tsallis 熵多说两句。
标准的香农熵(Shannon Entropy)基于一个假设:系统是“广延的”(Extensive),即局部信息的简单相加等于整体。但现实世界中的医疗数据往往是“长尾分布”的,罕见病理特征的权重如果在简单的线性系统中会被多数类(如健康样本)淹没。
Tsallis 熵是香农熵的非广延(Non-extensive)推广版本,它引入了一个极其关键的熵参数(Entropic parameter) qqq 。
当 q→1q \to 1q→1 时,Tsallis 熵退化为标准的香农熵。
当 q<1q < 1q<1 时(核心机制):函数会对那些概率极小(Pi≈0P_i \approx 0Pi≈0)的事件给予指数级的放大补偿。这意味着,即使模型对某个罕见类别的预测概率很低,它产生的不确定性惩罚或信号也会被显著放大。
在评估每个样本的预测不确定性时,论文使用 Tsallis 熵的具体计算公式如下 :
1q−1×(1−∑i=0n=5Piq)\frac{1}{q-1} \times \left(1-\sum_{i=0}^{n=5}P_i^q\right)q−11×(1−i=0∑n=5Piq)
(注:论文原文给出的公式排版中包含了一个 log 符号 ,即 (1q−1)×(1−Σi=0n=5log(Piq)(\frac{1}{q-1})\times(1-\Sigma_{i=0}^{n=5}log({P_{i}}^{q})(q−11)×(1−Σi=0n=5log(Piq),但从 Tsallis 熵的标准物理学定义推导来看,经典形式不包含对数操作。我以其核心的非线性幂次机制为准。)
在这个公式中:PiP_iPi 代表模型对第 iii 个类别(共有 6 个类别,即非青光眼和 5 种青光眼亚型,因此 n=5n=5n=5)输出的预测概率 。这篇论文的临床场景是青光眼亚型分类,面临着严重的类别不平衡(Class imbalance) 。例如,继发性青光眼(SGL)的样本量远远少于非青光眼(Non-GL)或开角型青光眼(OAG/S)。作者特意将熵参数设定为 q=0.25q=0.25q=0.25 。通过设置这样一个远小于 1 的 qqq 值,模型极大地强调了在罕见类别上的不确定性 。这种机制使得模型对代表性不足的青光眼亚型的预测提供了更高的敏感度 。如果用传统的香农熵,由于罕见类别的初始预测概率极低,其不确定性会被掩盖;而 q=0.25q=0.25q=0.25 的 Tsallis 熵就像一个放大镜,强迫模型在那些样本量极少、极容易出错的边缘群体(例如特定人口统计学背景下的罕见亚型)上给出更确定、更可靠的预测。
开源代码没有给出实现,我这里给出一个简单的实现和测试:
import torch
import torch.nn.functional as F
def tsallis_entropy(logits, q=0.25):
"""
计算非广延 Tsallis 熵,用于衡量不平衡类别的预测不确定性。
Args:
logits (torch.Tensor): 模型的原始输出 (Batch_size, Num_classes)
q (float): 熵参数。q < 1 时会放大罕见类别(低概率事件)的权重。
Returns:
torch.Tensor: 每个样本的 Tsallis 熵值
"""
# 1. 将 logits 转化为概率分布
probs = F.softmax(logits, dim=1)
# 2. 增加微小扰动 (epsilon) 防止底层计算 0 的幂次出现数值不稳定
probs = torch.clamp(probs, min=1e-9)
# 3. 计算 Tsallis 熵: (1 / (q - 1)) * (1 - sum(P_i^q))
# 注意:论文中 q=0.25,所以 (1 / (0.25 - 1)) = -1.333
sum_probs_q = torch.sum(probs ** q, dim=1)
entropy = (1.0 / (q - 1.0)) * (1.0 - sum_probs_q)
return entropy
# === 测试代码 ===
# 假设我们有一个极其不确定的样本(类别概率平均)和一个极其确定的样本
mock_logits = torch.tensor([
[0.1, 0.1, 0.1, 0.1, 0.1, 0.5], # 不确定性高
[10.0, -2.0, -2.0, -2.0, -2.0, -2.0] # 确定性极高 (属于第0类)
])
uncertainties = tsallis_entropy(mock_logits, q=0.25)
print("Tsallis Entropies:", uncertainties)
更多推荐


所有评论(0)