引言
大语言模型在医疗推理中展现出巨大潜力,但现有微调方法往往依赖海量标注数据,不仅计算成本高昂,且大量冗余、低质量的样本反而会稀释模型的临床推理能力。如何在有限数据下实现高效、精准的医学推理,成为亟待解决的核心难题。
近日,发表于CVPR的一项研究中,来自华东师范大学、MBZUAI等机构的研究团队提出了一个名为DIQ(Difficulty-Influence Quadrant)的数据选择框架。不同于传统方法仅基于样本难度或梯度影响单一维度进行筛选,DIQ创新性地将二者结合,优先选取“高难度-高影响力”的样本。实验表明,仅用1%的DIQ精选数据微调,性能即可媲美全量数据;使用10%数据时,则持续超越现有基线方法,为构建高效、可扩展的医学推理模型提供了全新路径。
基本信息
•文章标题:Towards Efficient Medical Reasoning with Minimal Fine-Tuning Data
•发表时间:2026年3月14日
•研究单位:华东师范大学、穆罕默德·本·扎耶德人工智能大学(MBZUAI)、莫纳什大学、上海人工智能实验室
•Github地址:https://github.com/mihara-bot/DIQ
•论文地址:https://arxiv.org/abs/2508.01450v3
•算力描述:所有训练在搭载4块NVIDIA A800 GPU的Ubuntu 22.04服务器上进行;采用LoRA微调,LoRA目标模块包含query、key、value投影,LoRA秩设为8,最大上下文长度为8192 tokens,学习率为1×10−4并使用余弦衰减调度,训练3个epoch。
研究内容与方法
一、医疗推理数据集预处理
针对现有医疗推理数据集(如Huatuo、FineMed、MedReason等),进行统一格式转换,将样本处理为包含问题文本、推理过程与答案的结构化数据,确保后续难度评估与影响力计算的输入一致性。
代码片段(来自process_data.py):
defprocess_medical_data(raw_data):processed_samples=[]foriteminraw_data:sample={'text':item['question']+'\n'+item['reasoning'],'answer':item['answer']}processed_samples.append(sample)returnprocessed_samples二、DIQ框架概述
DIQ是一种面向医疗推理的高效数据选择框架,通过联合评估样本的难度得分与Dot影响力得分,在二维空间划分象限并优先选择兼具临床复杂度与优化效用的样本,实现小样本微调下的高效医疗推理。框架流程分为三步:样本得分标注、象限划分、优先级选择。
【DIQ框架流程图】
三、难度得分计算
1. 核心目的
评估医疗样本的知识复杂度、推理复杂度与整体难度,筛选具有临床价值的复杂推理样本。
2. 实现逻辑
- 难度分类器训练
采用BiomedBERT作为基础模型,在标注了知识(Knowledge)、推理(Reasoning)、整体(Overall)三个维度难度的医疗样本上微调,每个维度难度为1-5分,对应不同临床复杂度层级。
代码片段(来自difficulty.py):defpredict_difficulty(model,tokenizer,text):inputs=tokenizer(text,return_tensors="pt",padding=True,truncation=True)outputs=model(**inputs)logits=outputs.logits scores=torch.softmax(logits,dim=1).tolist()[0]# 返回知识、推理、整体三个维度的得分returnscores[0],scores[1],scores[2] - 难度得分提取
从分类器输出中选择指定维度(默认选择Overall)作为样本的最终难度得分,公式为:
D(z)=Dϕ(z),ϕ∈{K,R,O}D(z) = D_\phi(z), \phi \in \{K, R, O\}D(z)=Dϕ(z),ϕ∈{K,R,O}
其中ϕ\phiϕ为选择的难度维度(K=知识、R=推理、O=整体)。
代码片段(来自selection.py):defget_difficulty_scores(difficulty_model,tokenizer,samples,dimension='overall'):difficulty_scores=[]forsampleinsamples:k_score,r_score,o_score=predict_difficulty(difficulty_model,tokenizer,sample['text'])ifdimension=='knowledge':difficulty_scores.append(k_score)elifdimension=='reasoning':difficulty_scores.append(r_score)else:difficulty_scores.append(o_score)returndifficulty_scores
四、Dot影响力得分计算
1. 核心目的
评估样本对模型验证损失的优化效用,筛选能有效降低验证损失的高价值训练样本。
2. 实现逻辑
- 梯度内积计算
计算每个训练样本的梯度与验证集平均梯度的内积,作为Dot影响力得分,公式为:
Dot(z)=1∣Dval∣∑z′∈Dvalg(z;θ^)⊤g(z′;θ^)\text{Dot}(z) = \frac{1}{|D_{val}|} \sum_{z' \in D_{val}} g(z; \hat{\theta})^\top g(z'; \hat{\theta})Dot(z)=∣Dval∣1z′∈Dval∑g(z;θ^)⊤g(z′;θ^)
其中g(z;θ^)g(z; \hat{\theta})g(z;θ^)是样本zzz在预训练模型参数θ^\hat{\theta}θ^下的梯度,DvalD_{val}Dval为验证集。 - 高效降维优化
采用随机投影降低梯度维度,减少计算成本同时保留梯度排序信息,确保大规模数据集下的计算效率。
代码片段(来自influence.py):defcompute_dot_influence(model,train_samples,val_samples,tokenizer,proj_dim=4096):# 计算验证集平均梯度val_grads=[]forsampleinval_samples:inputs=tokenizer(sample['text'],return_tensors="pt")outputs=model(**inputs,labels=inputs["input_ids"])loss=outputs.loss loss.backward()grad=torch.cat([p.grad.flatten()forpinmodel.parameters()ifp.gradisnotNone])val_grads.append(grad)model.zero_grad()avg_val_grad=torch.stack(val_grads).mean(dim=0)# 随机投影矩阵,降低梯度维度proj_matrix=torch.randn(avg_val_grad.shape[0],proj_dim)/np.sqrt(proj_dim)# 计算训练样本的Dot得分dot_scores=[]forsampleintrain_samples:inputs=tokenizer(sample['text'],return_tensors="pt")outputs=model(**inputs,labels=inputs["input_ids"])loss=outputs.loss loss.backward()grad=torch.cat([p.grad.flatten()forpinmodel.parameters()ifp.gradisnotNone])# 投影后计算内积proj_grad=grad @ proj_matrix proj_avg_val_grad=avg_val_grad @ proj_matrix dot_score=torch.dot(proj_grad,proj_avg_val_grad).item()dot_scores.append(dot_score)model.zero_grad()returndot_scores
五、象限划分与优先级选择
1. 核心目的
结合难度与影响力得分,筛选兼具临床推理复杂度与模型优化效用的样本子集。
2. 实现逻辑
- 象限划分
设定难度阈值τd\tau_dτd(训练集难度得分的百分位数)和Dot得分中位数mdotm_{dot}mdot,将样本分为四个象限:- Q1:高难度高影响力(D(z)≥τdD(z) \geq \tau_dD(z)≥τd且Dot(z)≥mdot\text{Dot}(z) \geq m_{dot}Dot(z)≥mdot)
- Q2:低难度高影响力(D(z)<τdD(z) < \tau_dD(z)<τd且Dot(z)≥mdot\text{Dot}(z) \geq m_{dot}Dot(z)≥mdot)
- Q3:高难度低影响力(D(z)≥τdD(z) \geq \tau_dD(z)≥τd且Dot(z)<mdot\text{Dot}(z) < m_{dot}Dot(z)<mdot)
- Q4:低难度低影响力(D(z)<τdD(z) < \tau_dD(z)<τd且Dot(z)<mdot\text{Dot}(z) < m_{dot}Dot(z)<mdot)
代码片段(来自selection.py):
defpartition_quadrants(difficulty_scores,dot_scores,tau_d=0.7,m_dot=None):ifm_dotisNone:m_dot=np.median(dot_scores)# 难度阈值取tau_d对应的百分位数tau_d_val=np.percentile(difficulty_scores,tau_d*100)quadrants={'Q1':[],'Q2':[],'Q3':[],'Q4':[]}foridx,(d,dot)inenumerate(zip(difficulty_scores,dot_scores)):ifd>=tau_d_valanddot>=m_dot:quadrants['Q1'].append(idx)elifd<tau_d_valanddot>=m_dot:quadrants['Q2'].append(idx)elifd>=tau_d_valanddot<m_dot:quadrants['Q3'].append(idx)else:quadrants['Q4'].append(idx)returnquadrants - 优先级样本选择
按照Q1→Q2→Q3→Q4的优先级顺序选择样本,每个象限内按Dot得分降序排序,若Dot得分相同则按难度得分降序排序,直到达到目标样本比例rrr。
代码片段(来自selection.py):defselect_samples(quadrants,difficulty_scores,dot_scores,target_ratio):total_samples=sum(len(q)forqinquadrants.values())target_num=int(total_samples*target_ratio)selected=[]# 优先级顺序:Q1 > Q2 > Q3 > Q4forq_namein['Q1','Q2','Q3','Q4']:iflen(selected)>=target_num:breakq_indices=quadrants[q_name]# 按Dot降序、难度降序排序sorted_indices=sorted(q_indices,key=lambdax:(-dot_scores[x],-difficulty_scores[x]))take_num=min(len(sorted_indices),target_num-len(selected))selected.extend(sorted_indices[:take_num])returnselected
实验结果分析
1. DIQ数据选择策略的总体性能表现
DIQ方法在多个数据集和保留比例下均显著优于随机选择和基线方法。实验表明,仅使用1%-10%的DIQ选择数据即可匹配甚至超越全量数据微调的性能。
- 数据效率极高:在FineMed数据集上,仅1%的DIQ选择数据(约170个样本)就使Llama3.1-8B-Instruct的平均准确率(AvgA)达到41.46%,远超全量数据训练的38.57%(提升2.89个百分点)。
- 持续超越基线:在Huatuo数据集上,10%的DIQ选择数据(约1,900个样本)达到44.04%的AvgA,超过所有基线方法(如LESS的40.54%、Similarity的41.73%),甚至优于全量数据训练的43.44%。
- 标准任务提升显著:在FineMed数据集1%保留比例下,DIQ在标准任务平均准确率(AvgS)上达到58.14%,相比全量数据训练的42.52%提升15.62个百分点。
【Llama3.1-8B-Instruct在不同数据集和保留比例下的下游任务性能对比,DIQ方法在多数设置下取得最佳结果】
2. 临床推理质量与专家对齐评估
DIQ选择的数据不仅提升了模型性能,还显著改善了临床推理质量,使其更贴近专家实践。通过LLM-as-a-judge评估,DIQ在鉴别诊断、安全检查、证据引用三个关键维度均表现优异。
- 数据层面质量提升:DIQ-1%选择的数据子集在鉴别诊断(DDx)评分上达到4.39,远高于全量数据剩余的3.59(提升0.80);证据引用(EC)评分4.77 vs 4.31(提升0.46)。
- 模型推理对齐增强:使用DIQ-1%数据训练的模型在鉴别诊断(3.71 vs 3.66)、安全检查(3.30 vs 3.14)和证据引用(4.90 vs 4.75)上均优于全量数据训练模型。
- 案例验证:在MedBullets-option5的临床病例中,DIQ-1%训练的Qwen3-8B模型能够系统性地进行鉴别诊断(DDx)、安全检查(SC)和证据引用(EC),其推理过程与专家临床思维高度一致。
【DIQ-1%数据与剩余数据、DIQ-1%模型与全量数据模型的临床价值对比,DDx、SC、EC三个指标均按5分制评分】
3. 消融实验与泛化能力验证
DIQ的二维选择策略(难度+影响力)优于任何单一维度的选择,并且在不同模型规模和模型家族间展现出良好的泛化能力。
- 二维策略优于单维度:在1%保留比例下,DIQ(42.78% AvgA)优于仅使用影响力(42.67%)、仅使用总体难度(41.36%)、仅使用知识难度(41.59%)和仅使用推理难度(40.70%)的单维度选择方法。10%保留比例下同样保持优势(44.04% vs 43.16%/40.01%/41.89%/40.70%)。
- 跨模型家族泛化:使用Llama3.1-8B计算的影响力分数应用于Qwen3-8B时,在10%保留比例下仍取得47.62%的AvgA,优于随机选择的45.51%(提升2.11个百分点)。
- 跨模型规模泛化:使用Qwen3-8B计算的影响力分数应用于Qwen3-14B时,在1%/10%/50%保留比例下分别提升1.42/1.89/1.56个百分点;应用于Qwen3-32B时提升0.45/2.41/2.64个百分点。
【不同消融设置下Llama3.1-8B-Instruct的平均准确率对比,DIQ在1%和10%保留比例下均优于单一维度选择方法】
【使用不同来源影响力分数的DIQ在Qwen3系列模型上的下游任务性能,DIQ能够跨模型规模和模型家族泛化】
优势与局限
优势
- 高效数据选择:DIQ通过联合评估样本难度与梯度影响,仅需1%–10%的精选数据即可匹配或超越全量微调性能,显著降低训练成本。
- 临床推理对齐:DIQ优先选择高难度-高影响力样本,提升模型在鉴别诊断、安全检查和证据引用上的表现,生成更贴近专家实践的推理过程。
- 跨模型泛化性强:DIQ在不同规模(8B–32B)和不同家族(Llama、Qwen)的模型上均有效,且支持跨模型迁移和偏好学习(DPO)扩展。
局限
- 计算资源依赖:DIQ需要为每个训练样本计算梯度内积,尽管已采用随机投影降维,但对大规模模型(≥70B)的计算开销仍需进一步验证。
- 验证集依赖:DIQ的性能受验证集大小影响,较小验证集可能导致影响力估计不稳定,需在成本与稳定性间权衡。
- 跨家族迁移有限:DIQ的跨家族迁移(如Llama→Qwen)在极低数据预算下可能表现不稳定,需结合模型特定影响力或难度重加权来弥补差距。
参考文献
Supervised Fine-tuning (SFT) of the language backbone plays a pivotal role in adapting Vision-Language Models to specialized domains such as medical reasoning.Zhuang et al., 2025:该论文提出了DIQ框架,通过联合评估样本难度与梯度影响,实现了仅用1%-10%精选数据即可匹配或超越全量微调性能,是本研究数据高效微调策略的核心方法。
Huatuogpt-o1, towards medical complex reasoning with llms.Chen et al., 2024:该论文构建了大规模医学推理数据集HuatuoGPT-o1,是本研究训练与评估所使用的核心数据来源之一,为验证DIQ的数据选择有效性提供了基准。
Lima: Less is more for alignment.Zhou et al., 2023:该论文证明了少量高质量数据即可有效激发大模型的对齐能力,启发了本研究“少即是多”的数据选择理念,并作为DIQ框架的理论基础之一。
LESS: Selecting influential data for targeted instruction tuning.Xia et al., 2024:该论文提出了基于TracIn影响分数的数据选择方法LESS,是本研究在实验部分对比的主要基线方法之一,用于验证DIQ在医学推理场景下的优越性。
Disentangling reasoning and knowledge in medical large language models.Thapa et al., 2025:该论文提出解耦医学推理中的知识与推理难度,是本研究DIQ框架中难度评估维度的直接来源,为样本难度的多维度量化提供了方法论支持。