最近很多朋友问我,训练一个32B多模态医疗模型到底需要什么条件?说实话,比起网上那些动辄几百B的模型刷榜消息,32B这个规模在医疗场景里反而更接近真实落地:既有足够的常识和推理储备,又不至于让大多数团队直接放弃。这篇文章是我对一个32B多模态医疗模型项目从零到一的完整复盘,重点不是展示一个漂亮的benchmark,而是把我们怎样整理数据、选架构、跑训练、排坑的过程如实写出来。规划分两篇,这篇先讲上半程的决定性工作,也就是数据、模型骨架和训练工程;评测、对齐和部署放到下篇再聊。
1. 项目立项:为什么选32B,而不是更大或更小
1.1 医疗场景对模型规模的现实约束
多模态医疗模型要面对的东西很杂:胸部X光片、CT序列、病理切片、病历文本、检验指标、历史诊断结论。模型既要做视觉理解,又要做医学推理,还要能生成结构化的诊断意见,单靠一个纯文本大模型或一个小号视觉模型都很难覆盖。
最开始我们也考虑过直接从70B以上的底座出发,毕竟通用大模型的能力已经在很多场景被验证过。但真到医疗场景里看,至少有三个约束把规模往下拉:
第一是数据量。医学数据尤其是高质量图文对数据,远不如互联网通用图文数据那么丰富。模型参数越大,对数据多样性的需求越高,在医学数据总量有限的情况下,强行上大参数很容易把记忆训练得很好,但推理泛化并没有同步提升,甚至会出现更严重的过拟合。
第二是训练和推理成本。一个34B左右模型的全参数训练,单是优化器状态就是一笔不小的显存开销,如果没有多机多卡的环境,根本跑不动。就算训练出来了,落地到医院或科研机构时,人家不一定会买8卡甚至16卡的高端服务器,单卡能扛住推理是硬门槛。32B这个级别配合量化,还能勉强进到单卡A100或H100里跑部署,70B就基本不可能了。
第三是长尾任务的可控性。医疗场景对错误容忍度极低,模型不能只是"能聊",而是要能准确引用图像中的位置信息、区分左右、判断病灶大小,这些能力在中等规模模型上可以通过针对性的训练数据做得比较好,反而过大模型在指令跟随上容易发散,可控性下降。
所以,32B并不是一个拍脑袋的数字,而是在"模型容量够用"和"团队资源够得着"之间取的点。如果你今天要复现类似项目,我建议先把资源预算算明白,再决定底座规模,不要一上来就追求大。
1.2 从7B/14B升到32B的动机和成本测算
我们在项目启动前的第一版方案其实是从7B/14B开始的,原因是团队里有人担心医疗数据太少,大模型训不动。但跑完一轮小规模实验后,差距很明显:7B模型在简单影像描述任务上还可以,一旦涉及多病种对比、跨模态信息综合判断,输出就开始出现逻辑混乱。比如给一张既有肺气肿又有少量胸腔积液的片子,7B经常只能说出其中一个,14B偶尔能说出两个,但很难准确判断谁先谁后。
于是我们做了一轮成本测算,决定直接上32B。
简单算笔账。假设模型参数量是32B,用BF16存储权重,光参数就要约64GB。训练时开启AdamW优化器,梯度约64GB,优化器状态在FP32下需要约128GB。这意味着最基本的参数、梯度和优化器状态加在一起已经超过256GB,这还不算中间激活值、通信缓冲和临时显存。单机8卡A100 80G,总共640GB显存,如果不开任何优化,理论上勉强能塞下,但一旦把batch size提上去就会立刻爆掉。
我们的实际配置是8个节点、每节点8张H800,总计64卡,走DeepSpeed ZeRO-3将参数、梯度和优化器状态全部切分到所有卡上。这样单卡显存压力大幅下降,但跨节点通信开销明显上升,对网络带宽要求很高。我们训练数据约为30万条图文对,约15万条纯医学文本,再混入少量通用指令数据,global batch size设为128,序列长度4096,训练一个epoch大约需要2000多步。最后整个预训练对齐加指令微调阶段,累计跑了大约20天,中途经历过三次训练中断和一次数据污染回滚。
这里想强调的是,成本测算不能只看显存,还要看训练吞吐和稳定性。32B模型在多机环境下的通信占比非常高,如果网络不是InfiniBand而是普通万兆以太网,训练速度可能只有前者的三分之一。
需要模型API调用? 免费领10W Token,多模型网关一键接入 Claude、DeepSeek 等主流模型。
2. 医疗多模态数据:最难啃的硬骨头
2.1 数据来源与合规清洗
如果说模型训练是盖楼,那医疗多模态数据就是地基,同时还是最容易出安全事故的一层。医疗数据的特殊性在于,它不只是技术问题,更是合规和伦理问题。
我们项目的数据主要来自几个公开数据集:MIMIC-CXR、CheXpert、MedMNIST系列,以及一部分和合作医院签署过授权协议的脱敏数据。公开数据集的好处是版权和伦理问题相对清晰,很多学术研究可以直接用,但它们的格式非常不统一。MIMIC-CXR是胸部X光片加对应放射报告,CheXpert是影像加标签,MedMNIST是多个医学图像分类任务的集合。要把这些数据揉进同一个训练管线,第一步就是统一格式。
我强烈建议在数据清洗阶段做一次完整的去标识化处理。有人说公开数据集是不是就干净了?不是。很多影像文件里还嵌着DICOM头信息,里面可能包含患者姓名、检查日期、医院ID等元数据,直接拿来训练会留下隐私隐患。我们用pydicom库批量处理,把所有非图像字段全部抹掉,只保留尺寸、像素间距、窗宽窗位这些和图像内容相关的属性。对于文本报告,则需要用规则匹配和NER模型筛掉姓名、生日、联系方式等实体,这一步宁可多过滤,也不能漏。
另外一个容易被忽略的是许可协议。不同公开数据集的持续使用条款不一样,有些允许学术研究但不允许商业落地,有些要求我们发布模型时必须带上数据来源声明。我们在立项时就专门拉了一张合规清单,把每个数据集的用途范围、引用方式、是否可以商用逐项写清楚,避免模型训到一半才发现不能上线。
2.2 图文对的构建策略
多模态模型的核心输入是图文对,但医疗图文对天然不是现成的。我们原始数据里,影像和报告是一对多或者多对一的关系,一份CT序列可能有几百张切片,每个切片都没有独立报告;一份病历里可能同时提到了几次检查的信息。直接做全局配对,模型学到的是模糊对应关系,很容易产生幻觉。
我们采取的方法是"影像单元"和"文本单元"的双层切分。影像单元不是单张图,而是根据临床逻辑切出来的子集:比如胸部CT按解剖位置分成肺尖、肺门、膈顶等区域,每个区域取2-3张代表性切片组成一个输入视窗。文本单元则是从一段长报告中切出来的语义块,规则上以句号或"印象:"为边界,尽量保证每段只描述一个观察点。
切完之后,还要做对齐。早期我们试过直接用文本分类模型把每段报告映射到"是否有病灶"、"病灶位置"等标签,再根据标签和影像单元做匹配,效果一般。后来改用规则+预训练CLIP模型的组合策略:先用疾病标签粗筛一遍,再用CLIP计算每个影像单元和每段文本的相似度,取相似度最高的配对作为训练样本。CLIP模型本身对医疗图像不一定敏感,但它能提供一个先验排序,配合人工抽检,可以把明显错配的样本筛掉。
这个过程中一定会有数据损失。我们最初切出大约60万对候选,经过对齐、清洗、抽检后,最终保留到30万对左右。别觉得浪费,脏数据训练出来的模型,在医疗场景里犯的错误可能是致命的,损失一些数据比模型学会胡说八道要划算得多。
2.3 数据配比与质量过滤
医疗数据天然存在严重的类别不平衡。肺炎、肺结节、胸腔积液这类常见病样本非常多,而气胸、心包积液、少见肿瘤的样本可能只有几十例。如果直接按原始分布训练,模型对常见病会非常熟练,对罕见病几乎完全失效。我们需要做配比控制。
具体做法是,先把所有样本按主诊断标签划分,然后对数量超过1万例的常见病做下采样,对数量少于200例的罕见病做过采样。过采样不是简单复制,而是结合图像增强:旋转、翻转、随机裁剪、灰度扰动,这样能在不改变疾病语义的前提下增加多样性。不过过采样的倍率不能太高,我们控制在2倍以内,否则模型会记住增强后的噪声。
质量过滤也要分图像质量和文本质量两条线。图像方面,我们用简单的统计指标过滤曝光异常、纯黑边、伪影严重的图,再人工抽检几千张确认过滤阈值。文本方面,重点过滤乱码、英文缩写滥用、无临床意义的套话报告。我们甚至写了一个小分类器,专门识别"报告是否包含实质医学描述",凡是空泛的结论都被踢掉。
数据配比最终定成:影像-报告对占60%,纯医学文本占30%,通用指令数据占10%。这个配比的思路是,影像-报告对让模型学会跨模态理解和生成,纯医学文本保持和强化模型对医学知识的记忆,通用指令数据则减少模型在对话能力上的退化。后面做指令微调时,我们又重新调过配比,但那是以任务为导向的另一套逻辑。
3. 模型架构与基础模型选型
3.1 视觉编码器、语言底座与连接器怎么选
模型架构看起来是老三样:视觉编码器、语言底座、连接器,但每个部分在医疗场景下都有讲究。
视觉编码器我们对比过ViT-L/14、CLIP ViT-H、SigLIP-SO400M,最后选的是SigLIP-SO400M。原因很简单:它在遮挡、模糊、小目标上的视觉特征更稳,对于医学影像里那种低对比度、病灶占比小的图片,比CLIP训练出来的特征更能保留细节。另一个选择是直接用已经在医学图像上预训练过的编码器,比如一些基于RadImageNet的模型,但我们试下来发现,它的通用迁移能力略弱,在多样化数据集上的表现反而不如SigLIP。
语言底座,我们用了Qwen2.5系列的一个32B版本。选择它不是因为刷榜分数最高,而是因为它在多语言指令跟随、结构化输出和长文本方面的能力比较稳,而且社区资料多,遇到问题容易排查。医疗场景里,报告生成经常要求输出模板化结构,这对语言模型的指令跟随能力要求比通用对话更高。
连接器我们最初试过Q-Former和C-Abstractor这类池化方案,但最终换成了两层MLP。MLP看着简单,但对图文对齐来说足够直接,尤其在高分辨率医学图像上,池化会丢掉很多空间细节,MLP则逐token投影,保留位置信息的能力更强。这里有一个经验:在医疗影像任务里,空间位置细节往往决定了模型能不能区分出"左上肺结节"和"左下肺结节",任何会压缩空间信息的连接器都要非常谨慎。
视觉分辨率也值得单独说。我们最终把图像输入统一到448x448,部分切片在训练时会随机裁剪到336x336做增强。这个数字不是随意定的,而是兼顾了模型输入限制和显存开销。医学图像中有些病灶很小,分辨率太低完全看不出来,但一味提高分辨率会让视觉token数量暴涨,拖慢训练速度,所以448是一个平衡点。
3.2 训练策略:预训练+指令微调两阶段
很多人拿到一张图文对数据集,第一反应就是直接拿去微调大模型。这样做不是不行,但很容易出现两个问题:一是视觉编码器输出分布和语言模型的文本空间还没对齐,模型学得很慢;二是直接更新全部参数,会覆盖掉语言模型原有的医学常识,导致灾难性遗忘。
我们把训练明确分成两个阶段。
第一阶段是图文对齐预训练。这个阶段把语言底座和视觉塔基本冻结,只更新连接器,用大量图文对数据让视觉特征映射到文本语义空间。学习率可以稍微高一点,我们用1e-4,batch size也可以大一些,因为不更新大参数,显存压力小很多。跑3000到5000步之后,连接器基本能对齐两种模态,生成质量虽然还不好,但模型已经能根据图像产生相关文本。
第二阶段是医疗指令微调。这个阶段才真正更新全部参数。数据换成以指令形式组织的医疗任务样本,比如"请根据这张胸部X光片描述影像所见"、"请判断是否存在胸腔积液,并给出位置和程度"。学习率要降下来,我们用1e-5左右,配合warmup和cosine衰减,让模型在已有对齐基础上慢慢进入医学任务状态。
两阶段分开还有一个好处:排查问题容易。第一阶段loss不降,大概率是数据对齐或连接器问题;第二阶段效果差,则更多是任务数据质量和微调策略问题。直接一步到位,出了问题很难定位。
3.3 全参数微调 vs 参数高效微调的选择
和很多团队一样,我们也认真考虑过是不是用LoRA就够了。LoRA的好处很明显:显存占用小,训练速度快,一张A100也能跑起32B模型。在早期消融实验里,LoRA在简单报告生成上确实能追到接近全参数微调的效果,但一旦遇到复杂病例,比如需要综合影像特征和病历信息做鉴别诊断,LoRA版本的回答明显更浅,容易漏掉关键发现。
我们的判断是:医疗场景对输出质量的要求高于大多数通用场景,全参数微调值得付出额外成本。当数据量只有几十万级,全参数微调可以更充分地利用数据去调整模型内部表征,LoRA相当于在原始权重旁边加了一个低秩补丁,表达上限受限于秩的大小。当然,如果你的目标只是做一个demo或者验证想法,先用QLoRA把流程跑通是完全合理的,我们甚至建议所有新手都先走这个路线。
全参数微调带来的工程挑战是巨大的,必须配合分布式训练框架。我们用DeepSpeed ZeRO-3把参数、梯度、优化器状态全部切分到64卡上,开启BF16混合精度和FlashAttention,再配合activation checkpointing,才把单卡显存压到可用范围内。这一部分如果你不打算踩一遍坑,可以直接拿我们后面要说的配置做起点。
4. 训练实施:从单机调试到多机分布式
4.1 软硬件环境与分布式方案
我们最终训练环境是8节点、每节点8张H800,共64卡。卡间通信走NVLink和InfiniBand,节点间用RoCE网络。这套配置对32B模型来说不是富余,而是刚刚够用。如果你用8张A100 80G单机训练,也不是完全不能跑,但global batch size会被压得很小,训练稳定性和效果都会受影响。
软件栈用的是PyTorch 2.1 + DeepSpeed 0.12 + FlashAttention-2。DeepSpeed的ZeRO-3负责把模型参数、梯度和优化器状态分片到每张卡上,这样每张卡只需要持有一小部分参数,计算时再通过通信把需要的参数收集起来。这个机制对显存非常友好,但会显著增加通信量,所以网络不好会非常痛苦。
我们启动训练的关键DeepSpeed配置是这样的(简化版):
json复制{
"zero_optimization": {
"stage": 3,
"offload_optimizer": {
"device": "none"
},
"overlap_comm": true,
"contiguous_gradients": true,
"reduce_bucket_size": 5e8,
"stage3_prefetch_bucket_size": 5e8,
"stage3_param_persistence_threshold": 1e6
},
"bf16": {
"enabled": true
},
"activation_checkpointing": {
"partition_activations": true
},
"train_batch_size": 128,
"gradient_accumulation_steps": 4
}
这里最值得注意的两个参数是overlap_comm和activation_checkpointing。前者让梯度通信和反向计算重叠,能明显提高GPU利用率;后者用额外的重计算换显存,开启后我们才把每卡batch size从1提升到2,global batch从64提到128。没有开activation checkpointing之前,稍微把序列调长一点就会OOM,开了之后稳定很多。
4.2 关键训练超参数与损失函数细节
训练超参数不是随便填的,下面这张表是我们经过几轮消融后确定的最终值:
| 超参数 | 对齐阶段 | 指令微调阶段 |
|---|---|---|
| 优化器 | AdamW | AdamW |
| 学习率峰值 | 1e-4 | 1e-5 |
| 学习率调度 | cosine,warmup 3% | cosine,warmup 5% |
| 权重衰减 | 0.01 | 0.01 |
| 图像分辨率 | 448x448 | 448x448 |
| 最大序列长度 | 2048 | 4096 |
| 每卡batch size | 4 | 2 |
| 梯度累积步数 | 2 | 4 |
| 有效global batch | 512 | 128 |
global batch的计算公式是:每卡batch size乘以卡数再乘以梯度累积步数。对齐阶段我们用了512,因为只更新连接器,可以承受大batch;指令微调阶段降到128,避免大batch在小数据上过早收敛到局部最优。学习率从对齐到微调降了一个数量级,这也是全参数微调必须要做的,否则很容易让预训练好的语言表征被带偏。
损失函数我们用的是主损失加辅助损失的组合。主损失是文本生成的标准交叉熵,只计算文本token的loss,不计算图像patch的loss。辅助损失是一个图文对比损失,用CLIP loss的形式把视觉特征和文本特征拉近。两个损失的权重是0.9和0.1,辅助损失只在预训练对齐阶段启用,指令微调阶段我们把它关掉了。原因是指令微调阶段我们希望模型更多依赖语言指令完成任务,而不是一直受视觉-文本对齐信号约束,否则生成灵活性会下降。
4.3 训练过程监控与稳定性
在长周期训练里,你不能只盯着一张loss曲线看,那是练到第三天才发现问题的时候才做的事。我们从第一天起就搭了一个训练看板,记录train loss、grad norm、learning rate、吞吐量、每卡显存、GPU利用率和网络通信时间。其中我最喜欢看的是grad norm曲线,它对训练稳定性非常敏感。
grad norm可以理解为模型梯度的大小变化。正常训练里,它会在一个范围内波动,如果突然出现一个比正常值大10倍的尖峰,基本可以判断出了问题。我们遇到过两次,一次是某个batch里混入了一张极端异常的图像,另一次是一段报告里出现了大量重复的乱码token。通过定位到具体样本后删除,问题就消失了。建议你在训练代码里加一个sample-level loss的钩子,定期对每个样本的损失排序,把异常样本挑出来,不要等整个loss都崩了再去查。
另一个容易出问题的是学习率调度。全程使用cosine decay时,到训练后半段学习率会变得非常小,模型更新几乎停滞。我们后来改成在训练到60%和80%时各做一次手动step decay,每个阶段降一半,效果比纯cosine更稳。医疗数据量不大,模型很容易过拟合,学习率提前降下来能有效抑制后面几个epoch的抖动。
5. 训练中的常见问题与排坑实录
5.1 loss不降或震荡
训练一开始,最让人慌的就是loss纹丝不动或者上蹿下跳。我总结了一下,通常逃不过这几个原因:
- 学习率太高:大模型参数更新幅度大,一下跑出最优区域,loss反复震荡。
- 数据噪声太大:图文对错配、文本本身乱码、标签错标的样本会影响梯度方向。
- 连接器欠拟合:第一阶段还没对齐好就进入第二阶段,模型接收了视觉特征却不知道怎么映射到文本。
- 图像预处理问题:有些图像是12位深度的DICOM转出来的,如果不做像素值归一化,模型看到的是完全不同的分布。
最简单的排查步骤是:先冻结视觉塔和语言模型,只训连接器,看loss是否能下降;如果连这样都不降,基本可以断定数据有问题。我们有一次跑了两天loss都没动静,最后抽样检查数据才发现,有接近10%的图文对是"图是某患者,报告是另一个患者",这种硬错位在样本量小的时候影响特别大。
从实操角度,我建议每2小时人工看一眼loss曲线,同时记录每一个checkpoint对应的loss值。很多问题并不是突然出现的,而是小到肉眼难察的异常累积出来的,定期归档很重要。
5.2 多卡通信与显存溢出
多机训练大面积翻车基本都发生在通信和显存上。下面几个问题我认为最值得提:
- allreduce超时:在8节点环境里,偶尔会出现某个节点响应慢,导致训练hang住。原因有时候不是程序问题,而是机房散热不均,某张卡温度过高降频了。我们后来加入了一项监控,每隔10分钟记录所有节点的温度、功耗和PCIe/NVLink错误计数,能提前预警。
- 显存溢出:最常见的情况是序列长度拉长后,attention带来的显存增长远超预期。解决办法依次是:开启activation checkpointing、减少每卡batch size、降低图像分辨率、关闭CPU offload。很多人一遇到OOM就想offload到CPU,我建议不到万不得已不要开,因为CPU offload会大幅度降低吞吐量,尤其在图文混合数据上更明显。
- token长度不统一:医疗报告长短差异很大,如果按固定长度padding,很多短报告会浪费算力。我们用sequence packing把小样本拼在一起训练,但一定要改attention mask,否则模型会跨样本乱串上下文。这个坑我们踩过,刚开始packing时loss没问题,但生成文本的语义会出现跨样本跳跃,后来才发现是mask没有隔离。
5.3 模型幻觉与医疗术语错误
训练过程中loss正常下降,不代表模型输出没问题。我们在checkpoint推理时发现模型经常把病灶位置写反,特别是"左肺"和"右肺"这类空间词乱用。这类问题在训练loss曲线上完全看不出来,因为文本生成loss只会衡量token预测概率,不会理解"左右的解剖学含义"。
我们需要在训练数据层面做约束。方法是在报告文本中尽量统一位置描述格式,比如"右肺下叶"不要转写成"右下肺叶","左侧胸腔"不要和"左肺"混用。另一个办法是让数据里的空间关系显式存在。用规则解析报告时,把"左右"、"上下"、"内外侧"这些位置词抽出来加特征标记,虽然会增加一些数据预处理的复杂度,但能显著减少空间幻觉。
另外,医疗术语的拼写错误也出现过。原因是训练集中某些OCR识别错了单词,比如把"pneumothorax"识别成"peumothorax",模型学会了错误拼写。我们后来加入了一个医学术语词典校验流程,凡是在词典里找不到的词汇,要么修正,要么标记为低置信度样本降到训练权重里。这个动作看起来很小,但非常有用。
5.4 评估指标的陷阱
训练阶段我们还在每个checkpoint上跑一套小规模的放射学报告评估,用的指标是CheXbert标签的F1和BLEU、Rouge-L。但我要提醒后来者:这些指标在医疗场景里都不完美,尤其BLEU和Rouge对同义改写非常不敏感,模型写出一堆正确但空泛的话也能拿到高分。真正有用的还是人工抽读报告,尤其是病灶位置、大小变化、否定词判断这几类关键信息。
我们要在下篇讲完整的评测体系,但在训练过程中有一件事现在就值得做:不要只看最后一个checkpoint的效果,而是要保留训练不同阶段的checkpoint,用同一套评估样本去跑一遍。我们发现最优checkpoint往往出现在训练中后期,而不是训练结束时刻。评估指标与训练loss不同步在医疗任务里尤其明显,模型可能在loss上继续下降,但输出变得越来越"模板化",术语越来越枯燥,这时就该停下训练,回到上一个好的checkpoint。
写在后面
这篇先停在这里。训练一个32B多模态医疗模型,数据、架构和工程三件事做到位,后面的评估和迭代才有意义。我个人最大的体会是,医疗模型不能只看通用benchmark,必须把临床场景的错例拿出来逐条看,很多问题在loss曲线上根本看不出来。我们在训练过程中反复做过数据回滚、checkpoint回退、超参数调整,整个过程没有想象中那么"自动化",反而充满了人工干预。下一篇再说评测、RLHF/DPO对齐和部署落地的细节,到时候会把这部分踩过的坑一起补上。
