资讯动态

LLaVA架构与训练全解析:从CLIP对齐到多模态模型复现

发布时间:2026/9/25 1:15:16 来源:尧图企业网站定制
多模态大模型这两年从论文走向工程落地的速度比我预想的快得多。2023年那会儿大家还在讨论“图文对齐到底有没有用”到了2024年随便一个做RAG的团队都会顺手接一个视觉编码器进来。而在这条演进路线上LLaVALarge Language-and-Vision Assistant几乎是绕不开的一个节点——它用极简的架构把视觉能力“嫁接”到了语言模型上训练成本低、复现难度小、效果又足够能打成了很多人入门多模态的第一站。而它背后站着的CLIP则是整个视觉-语言对齐范式的奠基者。这篇文章我想把LLaVA的架构和训练过程拆开讲透顺带把CLIP这个“地基”也讲清楚适合刚接触多模态的工程师、想复现LLaVA的学生以及需要给团队做技术选型的架构师。读完你应该能自己搭出一个最小可用的多模态模型并且知道每一步为什么这么做。1. 从CLIP到LLaVA多模态架构的演进逻辑1.1 为什么多模态需要先解决“对齐”问题要理解LLaVA得先理解它解决的是什么问题。语言模型本身只吃token它不知道“猫”这个字对应的像素长什么样。视觉编码器比如ViT只吃图像它也不知道“猫”这个词在语义空间里的位置。两个模态各自在自己的空间里自说自话中间隔着一道墙。多模态的核心任务就是把这堵墙拆掉让图像特征和文本特征落在同一个语义空间里或者说让语言模型能“读懂”图像特征。CLIP做的事情就是拆墙。它的思路非常直接拿4亿对图文数据用对比学习的方式让匹配的图文对在特征空间里靠近不匹配的推远。训练完之后一张猫的图片经过图像编码器得到的向量和“a photo of a cat”经过文本编码器得到的向量余弦相似度会很高。这就意味着图像特征和文本特征共享了一个语义空间。这个空间就是后来所有多模态模型的“公共语言”。我经常用一个类比来解释CLIP就像给图像和文本配了一个翻译官翻译官不负责理解内容只负责把两种语言映射到同一套坐标系里。至于翻译完之后谁来“理解”那是语言模型的事。1.2 CLIP的双塔结构与对比学习目标CLIP的架构是典型的双塔dual-encoder结构。图像侧用ViT或者ResNet文本侧用Transformer两边各自输出一个特征向量然后做L2归一化再计算相似度矩阵。训练目标是InfoNCE损失本质是一个batch内的对比学习对于batch里N个图文对第i张图应该和第i段文本最相似和其他N-1段文本都不相似。这里有个关键细节很多人会忽略CLIP的文本编码器有最大77个token的限制而且它用的是因果掩码还是双向注意力不同实现有差异。OpenAI原版用的是带因果掩码的Transformer取最后一个token的特征作为文本表示。这个细节在复现时如果搞错对齐效果会明显下降。CLIP训练完之后图像编码器就成了一个非常强的视觉特征提取器。它的输出是图像的高级语义特征而不是像素级细节。这一点很重要因为LLaVA正是拿这个编码器当“眼睛”用的。1.3 LLaVA的核心思路用投影层连接两个世界LLaVA的架构可以用一句话概括CLIP视觉编码器 投影层 语言模型。图像经过CLIP的ViT编码成一系列patch特征然后通过一个投影层通常是线性层或MLP映射到语言模型的词嵌入空间最后和文本token拼在一起送进语言模型。这个设计最妙的地方在于“极简”。它没有设计复杂的跨模态注意力没有搞花哨的融合模块就是一个投影层。为什么这么简单还能work因为CLIP已经把视觉特征对齐到了语义空间语言模型本身又有强大的语言理解能力中间只需要一个“适配器”把维度对上就行。这就像两个说不同方言的人其实语言底层是相通的只需要一个简单的翻译就能沟通。LLaVA有两个版本值得关注LLaVA-1.0用的是单层线性投影LLaVA-1.5升级成了两层MLP。别小看这个改动1.5的效果比1.0提升明显尤其是在视觉细节理解上。原因我后面会详细讲。2. LLaVA架构拆解三个组件与数据流2.1 视觉编码器CLIP ViT的选型与输出格式LLaVA用的视觉编码器是CLIP ViT-L/14输入分辨率224×224patch大小14×14所以一张图会被切成16×16256个patch。每个patch经过ViT编码后得到一个1024维的特征向量ViT-L的隐藏维度是1024。最终输出是256×1024的特征矩阵。这里有个工程上的坑CLIP ViT在训练时是带class token的但LLaVA用的是patch特征不用class token。因为class token是全局特征丢失了空间信息而LLaVA需要保留patch级别的细节才能让语言模型理解“左上角有什么”“右下角有什么”。另外224×224的分辨率对于很多实际场景是不够的。LLaVA-1.5后来支持了336×336patch数量变成24×24576视觉token数量翻倍细节理解能力明显提升。但token多了也有代价语言模型的上下文长度压力变大推理速度变慢。这是一个典型的精度和效率的权衡。2.2 投影层线性层还是MLP这是个问题投影层的作用是把视觉特征从CLIP的维度空间映射到语言模型的词嵌入空间。LLaVA-1.0用的是单层线性层LLaVA-1.5用的是两层MLP加GELU激活。为什么MLP更好我的理解是线性层只能做线性变换而视觉特征和文本嵌入之间的关系可能是非线性的。MLP提供了更强的表达能力能学到更复杂的映射。实测下来MLP版本在OCR、细粒度识别等任务上提升明显。投影层的参数量其实很小。以LLaVA-1.5为例输入1024维ViT-L输出4096维Vicuna-7B的隐藏维度两层MLP的参数量大约是1024×4096 4096×4096 ≈ 2100万参数。相比语言模型的70亿参数这个量级几乎可以忽略。但就是这2100万参数承担了整个多模态对齐的关键任务。2.3 语言模型Vicuna的选择与token拼接方式LLaVA-1.0和1.5用的都是Vicuna一个基于LLaMA微调的对话模型。为什么选Vicuna因为它的指令跟随能力强而且开源可商用。后来也有用Mistral、Qwen的变体思路是一样的。token拼接方式很直接把256个视觉token或576个直接拼在文本token前面。比如输入是“这张图里有什么”实际送进语言模型的序列是[IMG_TOKEN_1, ..., IMG_TOKEN_256, 这张图里有什么]。语言模型看到的就是一个长序列它不知道哪些是图像哪些是文本但通过训练它能学会区分。这里有个细节视觉token和文本token共享同一个嵌入空间但视觉token没有对应的位置编码。LLaVA的做法是给视觉token也分配位置编码位置从0开始文本token接着往后排。这个设计是否最优有争议但实测能work。2.4 完整数据流从像素到回答的每一步把整个流程串起来一张224×224的图 → CLIP ViT切成256个patch → 每个patch编码成1024维向量 → 投影层映射到4096维 → 得到256个视觉token → 和文本token拼接 → 送进Vicuna → 自回归生成回答。整个过程里CLIP ViT是冻结的不参与训练。投影层和语言模型参与训练。这个设计选择很关键我下一节会详细讲。3. 训练过程全解析两阶段策略与参数冻结逻辑3.1 第一阶段特征对齐预训练LLaVA的训练分两个阶段。第一阶段叫特征对齐feature alignment目标是让投影层学会把视觉特征映射到语言模型的语义空间。这个阶段只训练投影层CLIP和语言模型都冻结。训练数据是CC3M过滤后的595K图文对。过滤逻辑是用CLIP计算图文相似度去掉相似度低的噪声数据。这个过滤步骤很重要因为CC3M本身噪声很大不过滤的话投影层学不到干净的对齐关系。这个阶段的训练目标就是标准的自回归语言建模损失也就是预测下一个token。输入是图像问题输出是答案。但注意这个阶段的问题和答案都是比较简单的描述性文本比如“描述这张图片”。因为目标是学对齐不是学复杂推理。为什么只训练投影层因为投影层参数量小训练快而且冻结其他部分能防止灾难性遗忘。如果这个阶段就解冻语言模型语言模型可能会被视觉数据带偏丢失原有的语言能力。3.2 第二阶段指令微调第二阶段叫指令微调instruction tuning目标是让模型学会按照指令回答问题。这个阶段解冻语言模型和投影层一起训练CLIP仍然冻结。训练数据是GPT-4生成的158K多模态指令数据包括对话、详细描述、复杂推理三类。这些数据的生成方式是把图像的文字描述比如COCO的caption和边界框信息喂给GPT-4让GPT-4生成多轮对话和推理问题。这个“用语言模型生成多模态指令数据”的思路是LLaVA的一大贡献解决了多模态指令数据稀缺的问题。这个阶段的学习率要调小通常是第一阶段的一半甚至更低。因为语言模型已经预训练好了大学习率会破坏原有能力。实测下来2e-5是比较稳的选择。3.3 参数冻结策略背后的考量整个训练过程里CLIP始终冻结。为什么因为CLIP已经在大规模图文对上训练过了它的视觉特征提取能力足够强再训练反而可能过拟合到小规模数据上。而且冻结CLIP能省大量显存和计算让训练在单卡上也能跑。语言模型在第一阶段冻结、第二阶段解冻这个策略也是经过权衡的。第一阶段如果解冻语言模型它会被大量简单描述数据带偏而且计算成本高。第二阶段解冻是为了让语言模型适应多模态输入的分布学会利用视觉token。投影层始终训练因为它是随机初始化的必须从头学。3.4 训练成本与硬件配置参考LLaVA-1.5的完整训练在8张A100 80G上大约需要1天。第一阶段595K数据batch size 128训练1个epoch。第二阶段158K数据batch size 16训练3个epoch。如果资源有限可以只做第一阶段得到一个能描述图像的模型但指令跟随能力会差很多。也可以用小一点的语言模型比如Vicuna-7B换成3B显存需求能降到单张24G卡就能跑。我实测过用单张3090 24G跑LLaVA-7B的第二阶段batch size只能开到2用梯度累积模拟大batch训练时间大约3天。效果和官方配置有差距但作为学习复现是够用的。4. 实操复现从环境搭建到推理验证4.1 环境准备与依赖安装复现LLaVA的环境不算复杂但版本兼容性是个坑。我推荐用Python 3.10PyTorch 2.0以上CUDA 11.8。transformers库要用4.31以上因为LLaVA用到了新的tokenizer接口。关键依赖包括torch、transformers、accelerate、peft、bitsandbytes如果做量化、flash-attn可选能加速。flash-attn安装比较麻烦需要匹配CUDA版本如果装不上可以先跳过不影响功能。模型权重需要下载三个部分CLIP ViT-L/14的权重、Vicuna的权重、投影层的权重。前两个可以从HuggingFace下载投影层权重官方有提供。如果要做训练还需要下载训练数据。4.2 数据准备595K预训练数据与158K指令数据595K预训练数据是CC3M过滤后的子集官方提供了下载脚本。数据格式是图像路径文本描述。158K指令数据是GPT-4生成的官方也开源了格式是JSON包含图像ID、多轮对话。这里有个实操细节图像要预先处理成224×224或336×336并且做归一化。CLIP的归一化参数是mean[0.48145466, 0.4578275, 0.40821073]std[0.26862954, 0.26130258, 0.27577711]。这个参数如果搞错视觉特征会完全跑偏。数据加载用PyTorch的Dataset和DataLoader就行但要注意图像解码是CPU密集型的建议用多进程加载num_workers开到8以上。4.3 关键代码解析投影层实现与token拼接投影层的实现很简单LLaVA-1.5的版本是class LlavaProjector(nn.Module): def __init__(self, vision_hidden_size1024, llm_hidden_size4096): super().__init__() self.linear_1 nn.Linear(vision_hidden_size, llm_hidden_size) self.linear_2 nn.Linear(llm_hidden_size, llm_hidden_size) self.act nn.GELU() def forward(self, vision_features): hidden self.linear_1(vision_features) hidden self.act(hidden) hidden self.linear_2(hidden) return hiddentoken拼接的逻辑是先把文本token的embedding查出来然后把视觉token拼在前面。注意视觉token不需要查embedding因为投影层输出的已经是embedding维度的向量了。text_embeds llm.get_input_embeddings()(input_ids) combined_embeds torch.cat([vision_embeds, text_embeds], dim1) outputs llm(inputs_embedscombined_embeds, attention_maskcombined_mask)这里有个坑attention mask也要对应拼接视觉token的位置都是1可见文本token按原来的mask。如果mask搞错模型会看到不该看的位置训练会发散。4.4 推理验证用一张图测试模型效果训练完之后推理流程是加载CLIP、投影层、语言模型 → 图像过CLIP得到视觉特征 → 过投影层 → 文本tokenize → 拼接 → 生成。我建议先用官方权重做推理验证确认环境没问题再自己训练。推理时注意temperature和top_p的设置LLaVA官方推荐temperature0.2top_p0.9。temperature太高回答会发散太低会重复。测试用例建议覆盖简单描述、OCR、空间关系、计数。这几个维度能快速暴露模型的能力边界。如果OCR完全不行可能是分辨率不够如果空间关系混乱可能是视觉token太少。5. 常见问题与排查技巧实录5.1 训练不收敛的典型原因训练不收敛最常见的原因是学习率太大。第二阶段如果学习率超过5e-5语言模型很容易发散loss会先降后升。我的经验是第二阶段用2e-5第一阶段用1e-3因为只训练投影层可以大一点。第二个原因是数据格式错误。视觉token和文本token的拼接顺序、attention mask的对齐、label的mask只对答案计算loss不对问题和图像计算这三处任何一处出错都会导致不收敛。建议先用小批量数据过拟合确认能过拟合再上全量数据。第三个原因是CLIP归一化参数搞错。这个错误很隐蔽loss会降但效果很差。建议把归一化后的图像可视化出来确认看起来正常。5.2 显存不足的优化方案显存不足是复现时最常遇到的问题。优化手段有几个层次优化手段显存节省效果影响实现难度梯度累积线性节省无低混合精度约40%极小低梯度检查点约60%训练慢20%中LoRA微调约70%略有下降中4bit量化约75%明显下降中冻结语言模型约50%指令能力差低我一般先用混合精度梯度累积还不够就上梯度检查点再不够就上LoRA。4bit量化我一般不推荐用于训练推理可以用。5.3 视觉token数量与效果的权衡视觉token数量直接影响效果和效率。224分辨率是256个token336是576个448是1024个。token越多细节理解越好但推理越慢上下文压力越大。我的实测数据256 token在简单描述上够用OCR基本不行576 token OCR有明显提升能识别中等大小的文字1024 token OCR很好但推理速度慢一倍以上。如果任务不涉及细粒度识别256或576就够了。还有一个技巧可以对视觉token做池化比如每4个token合并成1个减少token数量。但这样会丢失空间信息需要重新训练投影层。5.4 多模态幻觉的成因与缓解多模态幻觉是指模型描述了图中不存在的东西。成因有几个一是视觉特征不够强模型靠语言先验“脑补”二是训练数据里有偏差某些物体共现频率高模型学会了走捷径三是解码策略问题temperature太高容易发散。缓解手段提高视觉分辨率、增加视觉token、在训练数据里加入负样本告诉模型什么不在图里、降低temperature。我试过在指令数据里加入“图中没有X”这类样本幻觉率能降不少。6. 从LLaVA延伸多模态工程的几个方向6.1 视觉编码器的替换与升级LLaVA的架构是模块化的视觉编码器可以换。除了CLIP ViT还可以用SigLIP、EVA-CLIP、InternViT等。SigLIP用sigmoid损失替代softmax训练更稳定小batch下效果更好。EVA-CLIP在细粒度识别上更强。InternViT支持更高分辨率。换编码器要注意维度匹配投影层的输入维度要跟着改。另外不同编码器的归一化参数不同要对应调整。6.2 投影层的进阶设计投影层除了MLP还有几种进阶设计。Q-FormerBLIP-2用的用一组可学习的query token去“抽取”视觉特征能把256个token压缩到32个大幅减少token数量。Perceiver ResamplerFlamingo用的也是类似思路。这些设计的 trade-off 是压缩了token数量但可能丢失细节。对于需要细粒度理解的任务直接MLP可能更好。对于需要长上下文的任务压缩更有优势。6.3 多模态RAG与Agent场景的落地LLaVA这类模型在实际落地时很少单独用通常是作为多模态RAG或Agent的一个组件。比如文档问答场景PDF转图像 → LLaVA提取文字和图表信息 → 文本送进RAG检索 → 语言模型生成回答。Agent场景里LLaVA可以作为“眼睛”把看到的界面截图转成文字描述再交给Agent决策。这种用法对OCR和空间关系理解要求很高需要高分辨率和大token数量。我在实际项目里发现LLaVA的输出稳定性是个问题同样的图多次推理可能给出不同描述。解决方案是降低temperature或者在后面接一个校验模块用规则或另一个模型检查输出的一致性。6.4 训练数据的质量比数量更重要LLaVA的成功很大程度归功于GPT-4生成的高质量指令数据。158K数据比很多百万级数据效果还好因为质量高、多样性好、覆盖了对话和推理。自己做多模态微调时数据质量是第一位。我的经验是1000条高质量人工标注数据比10000条噪声数据效果好。如果预算有限优先保证数据的准确性和多样性而不是数量。数据生成可以用GPT-4V或Claude但要注意成本。一个技巧是先用小模型生成候选再用大模型筛选能省不少钱。6.5 评估指标与benchmark选择多模态模型的评估是个难题。常用的benchmark有VQAv2、GQA、TextVQA、POPE测幻觉、MMBench、SEED-Bench。每个benchmark侧重点不同VQAv2偏通用问答TextVQA偏OCRPOPE专门测幻觉。我的建议是不要只看一个benchmark要组合看。而且benchmark分数高不代表实际场景好用最好自己构建一个贴近业务的测试集。我见过benchmark分数很高但实际用起来一塌糊涂的模型因为benchmark和真实场景的分布差异很大。评估时还要注意prompt的影响。同一个模型不同的prompt模板分数可能差好几个点。做对比实验时要固定prompt。7. 我踩过的坑与实操心得7.1 版本兼容性transformers与tokenizer的坑LLaVA早期版本用的tokenizer和后来Vicuna的tokenizer有差异导致token id对不上训练出来的模型完全不能用。我建议直接用官方提供的tokenizer不要自己替换。如果要用新的语言模型一定要确认tokenizer的special token和padding token设置正确。transformers版本也要注意4.31之前不支持LLaVA的某些接口4.36之后又有API变动。我一般锁定4.35或4.36比较稳定。7.2 数据加载图像解码的性能瓶颈训练时图像解码是CPU瓶颈如果num_workers不够GPU会等数据。我试过num_workers4和num_workers16后者吞吐量能提升一倍多。但num_workers也不是越大越好太大内存会爆。一般设为CPU核心数的一半比较合适。另一个技巧是预先把图像解码成numpy数组存成npy文件训练时直接读npy能省大量解码时间。代价是磁盘占用变大但训练速度提升明显。7.3 学习率调度warmup的重要性LLaVA训练用了warmup前3%的step学习率从0线性升到目标值。这个warmup很重要尤其是第二阶段解冻语言模型时没有warmup很容易在初期发散。我试过去掉warmuploss在前100步就炸了。warmup的比例可以调3%是官方配置我试过1%和5%差别不大。关键是不要没有。7.4 推理部署量化与加速的取舍推理部署时量化能大幅降低显存和加速。8bit量化基本无损4bit量化在简单任务上无损复杂任务有下降。我一般用8bit除非显存实在不够。加速方面flash-attn能提升20%左右vLLM能提升更多但配置复杂。如果只是做demo不用折腾这些原生推理够用。如果要上生产vLLM是值得投入的。7.5 一个容易被忽略的细节padding sideLLaVA训练时用的是left padding还是right padding这个细节很多人忽略。语言模型自回归生成时如果用right padding生成的位置会错。LLaVA用的是left padding确保生成从正确位置开始。这个坑我在复现时踩过loss正常但生成全是乱码排查了很久才发现是padding side的问题。多模态这条路LLaVA是一个很好的起点但绝不是终点。它的架构简单到几乎“简陋”但正是这种简单让它成为了一个优秀的baseline和教学样本。理解了LLaVA再去理解BLIP-2、Flamingo、Qwen-VL这些更复杂的架构会容易很多。我个人的体会是多模态的难点不在模型架构而在数据和评估——怎么构造高质量的对齐数据怎么评估模型在真实场景下的表现这两个问题比调模型结构难得多。如果你正在做多模态相关的项目建议先把数据 pipeline 和评估体系搭好再动模型。

读完文章,也想定制专属网站?

尧图设计师 24 小时内与您沟通定制方案

免费获取报价