1. 从CLIP到LLava:多模态模型到底在解决什么问题
第一次接触LLava的人,脑子里往往有两个疑问:CLIP是什么,它跟LLava又是什么关系?我刚开始看这篇论文的时候也绕了不少弯路,后来把整个链路拆开才想明白——CLIP解决的是“图文对齐”,LLava解决的是“让语言模型能看图说话”。这两件事听起来接近,实际上是完全不同的两个技术层次。
先说CLIP。它的全称是Contrastive Language-Image Pre-training,核心思路非常朴素:给一批“图片-文本”配对数据,让模型学会判断哪段文字配哪张图。训练时,一个batch里有N对图文,模型分别用图像编码器和文本编码器把它们映射到同一个向量空间,然后计算N×N的相似度矩阵。对角线上的正样本对要拉近,非对角线上的负样本对要推远。就这么一个对比学习目标,在4亿对图文数据上跑下来,CLIP的零样本分类能力直接追平了有监督训练的ResNet。
但CLIP有个天生的局限:它只会“打分”,不会“生成”。你给它一张图和一句“这是一只猫”,它能告诉你这句话跟图搭不搭,但你让它描述图里发生了什么,它做不到。这就是LLava要补的那块拼图。
LLava的思路可以概括成一句话:把CLIP的视觉编码器接到一个开源大语言模型上,再用视觉指令数据做微调。视觉编码器负责把图片变成一串向量,一个投影层(通常是两层MLP)把这串向量翻译成语言模型能理解的token,语言模型再基于这些视觉token和文本token一起生成回答。整个架构没有太多花哨的东西,但效果出奇地好——LLava-1.5用不到2000条视觉指令数据微调,在多个多模态benchmark上就超过了同期很多更复杂的方案。
这篇文章我打算把CLIP和LLava这两块拆开讲透。从CLIP的对比学习目标函数、温度系数的作用,到LLava的架构细节、两阶段训练策略、投影层为什么用MLP而不是Cross-Attention,再到实际复现时数据格式怎么组织、训练超参怎么设、显存不够怎么降级。适合已经了解Transformer基础、想往多模态方向深入的人,也适合想快速跑通一个LLava复现的实验者。我会尽量把每个设计选择背后的“为什么”讲清楚,而不是只列一堆结论。
2. CLIP模型:对比学习如何让图文对齐
2.1 双塔架构的设计逻辑与信息瓶颈
CLIP的架构是典型的双塔结构:图像侧用ViT或ResNet,文本侧用Transformer。两边各自独立编码,最后在向量空间里做对比。这个设计看起来简单,但背后有一个关键取舍——双塔意味着图像和文本在编码阶段没有任何交互。
为什么这么设计?因为交互式架构(比如把图像特征和文本特征拼在一起过Cross-Attention)虽然表达能力强,但推理时没法预计算。你每来一个新的文本查询,都得把图像重新编码一遍。双塔的好处是图像编码一次就能存起来,文本侧编码完直接算余弦相似度,检索速度是交互式架构的几十倍甚至上百倍。CLIP的定位从一开始就是“可检索的图文对齐模型”,所以它选择了双塔。
但双塔的代价是信息瓶颈。图像编码器输出的那个向量(ViT-L/14是768维)必须承载整张图的所有信息,文本编码器输出的向量也必须承载整句话的所有信息。两边在编码阶段看不到对方,所有跨模态的细粒度对应关系都得压缩到这两个向量里。这就是为什么CLIP在需要精细空间推理的任务上表现一般——比如“图中左边的猫和右边的狗哪个更大”,这种问题双塔架构天然吃亏。
2.2 对比学习目标函数与温度系数的实际作用
CLIP的损失函数是InfoNCE的对称版本。假设一个batch有N对图文,图像编码得到$I_1,...,I_N$,文本编码得到$T_1,...,T_N$,先做L2归一化,然后计算相似度矩阵$S_{ij} = I_i \cdot T_j / \tau$,其中$\tau$是可学习的温度系数。损失分两个方向:对每一行做softmax,目标是让对角线上的值最大;对每一列做softmax,目标同样是对角线最大。两个方向的交叉熵取平均。
温度系数$\tau$在这里非常关键。它控制softmax的“尖锐程度”:$\tau$越小,softmax越尖锐,模型对难负样本的惩罚越大;$\tau$越大,分布越平滑,模型对所有负样本一视同仁。CLIP的做法是把$\tau$设成可学习参数,初始值0.07,训练过程中让它自己调整。我实测下来,这个初始值很关键——如果设成1.0,模型几乎学不动,因为softmax太平滑,正样本和负样本的梯度信号被稀释了。
注意:温度系数在实现时通常用logit_scale = log(1/τ)来参数化,并且会clamp到最大值100(对应τ=0.01),防止训练后期τ过小导致梯度爆炸。
2.3 图像编码器的选型对比:ViT vs ResNet
CLIP论文里试了多种图像编码器,最终ViT-L/14表现最好。但这里有个容易被忽略的细节:ViT在小数据上不如ResNet,在大数据上反超。CLIP用了4亿对数据,所以ViT的优势能发挥出来。如果你自己的数据只有几十万对,用ResNet可能更稳。
具体对比一下:ResNet-50的输出是2048维,经过一个投影层降到512或768;ViT-L/14的输出是1024维(patch 14x14,输入224x224,共256个patch),经过投影层降到768。ViT的优势在于注意力机制能捕捉长距离依赖,对全局结构的建模更好;ResNet的归纳偏置更强,局部特征提取更高效。CLIP最终开源的模型里,ViT-B/32、ViT-B/16、ViT-L/14三个规格最常用,B/32最快但精度最低,L/14最慢但精度最高。
| 模型规格 | 图像编码器 | 参数量 | 输出维度 | 推理速度(相对) | 典型用途 |
|---|---|---|---|---|---|
| ViT-B/32 | ViT-Base, patch 32 | 约86M | 512 | 最快 | 快速原型验证 |
| ViT-B/16 | ViT-Base, patch 16 | 约86M | 512 | 中等 | 平衡精度与速度 |
| ViT-L/14 | ViT-Large, patch 14 | 约304M | 768 | 最慢 | 高精度检索 |
2.4 文本编码器的细节与因果掩码的取舍
CLIP的文本编码器是一个12层、512宽、8头的Transformer。这里有一个设计选择值得注意:它用的是因果掩码(causal mask)还是双向注意力?答案是因果掩码。文本序列经过Transformer后,取最后一个token(EOT token)的输出作为整句话的表示。
为什么用因果掩码而不是双向?因为CLIP的文本编码器要跟GPT系列保持兼容,而且因果掩码在推理时可以做KV Cache加速。但因果掩码意味着前面的token看不到后面的token,对于“一只猫在沙发上”这种短文本影响不大,但对于长文本,最后一个token要承载所有信息,压力比较大。这也是CLIP在长文本检索上表现一般的原因之一。
文本侧还有一个细节:词表用的是BPE,大小49408,最大序列长度77。超过77个token的文本会被截断。实际用的时候,如果你的文本经常超过77个token,需要考虑分段编码再聚合,或者换用支持更长序列的模型。
3. LLava架构拆解:视觉token如何接入语言模型
3.1 整体架构:三件套的拼接方式
LLava的架构可以拆成三个部分:视觉编码器、投影层、语言模型。视觉编码器直接用CLIP的ViT-L/14(或者ViT-L/14-336px),投影层是一个两层MLP(Linear -> GELU -> Linear),语言模型用Vicuna或Llama。
数据流是这样的:一张图片经过ViT编码,得到256个patch token(每个1024维),投影层把这256个token映射到语言模型的词嵌入空间(比如Vicuna的4096维),然后这256个视觉token和文本token拼在一起,送入语言模型。语言模型看到的是“
这里的关键设计是投影层用MLP而不是Cross-Attention。为什么?因为MLP简单、训练快、参数量小。Cross-Attention虽然表达能力强,但会引入额外的计算开销和训练不稳定性。LLava论文里做了消融实验,MLP和Cross-Attention的效果差距不大,但MLP的训练速度快了将近一倍。对于“把视觉特征翻译成语言模型能理解的token”这个任务,MLP已经够用了。
3.2 视觉编码器的冻结策略与分辨率选择
LLava训练时,视觉编码器是冻结的。也就是说,ViT的权重从头到尾不更新,只训练投影层和语言模型(或者只训练投影层)。为什么冻结?两个原因:第一,CLIP的ViT已经在4亿对数据上训练过了,特征提取能力足够强,再微调容易过拟合;第二,冻结ViT能省大量显存,ViT-L/14有304M参数,如果参与训练,显存占用会翻倍。
分辨率方面,LLava-1.0用的是224x224,LLava-1.5升级到了336x336。336x336的ViT-L/14会产生576个patch token(24x24),比224的256个多了一倍多。更多的视觉token意味着更细粒度的视觉信息,但也意味着更长的序列和更大的显存占用。实测下来,336分辨率在OCR和细粒度识别任务上提升明显,但在一般对话任务上差距不大。
提示:如果你显存有限,可以用224分辨率的ViT-L/14,然后把投影层的输出维度对齐到语言模型的隐藏维度。显存占用能降40%左右,效果损失在5%以内。
3.3 投影层的初始化与训练稳定性
投影层虽然只有两层MLP,但初始化方式对训练稳定性影响很大。LLava官方实现里,第一层Linear的权重用正态分布初始化(std=0.02),第二层Linear的权重初始化为零。为什么第二层初始化为零?因为这样在训练初期,视觉token的输出全是零,语言模型完全忽略视觉输入,相当于从纯文本模型开始训练。随着训练进行,第二层权重逐渐更新,视觉信息慢慢注入。这个技巧叫“零初始化残差”,在很多多模态模型里都有用。
如果不做零初始化,训练初期视觉token的数值可能很大,跟文本token的嵌入尺度不匹配,导致loss震荡甚至发散。我试过用默认的Kaiming初始化,loss在前500步几乎不降,换成零初始化后loss平稳下降。
3.4 视觉token与文本token的拼接方式
拼接方式看起来简单,但有几个细节容易踩坑。LLava的做法是把256个视觉token放在文本token前面,中间不加任何分隔符。也就是说,序列变成[v1, v2, ..., v256, t1, t2, ..., tn]。语言模型的位置编码会自然地区分视觉和文本部分。
但这里有个问题:语言模型的注意力是因果的,视觉token只能看到自己前面的视觉token,看不到后面的文本token。这意味着视觉token之间可以互相注意,但文本token可以看到所有视觉token。这个设计是合理的——文本需要参考视觉信息来生成回答,但视觉token不需要参考文本。
另一个细节是padding。如果一个batch里有多张图,每张图的视觉token数量固定(256或576),所以不需要padding视觉部分。但文本部分的长度不一,需要padding到同一长度。padding token的label要设成-100,不参与loss计算。
4. LLava两阶段训练:从对齐到指令微调
4.1 阶段一:特征对齐预训练
第一阶段的目标是让投影层学会把CLIP的视觉特征翻译成语言模型能理解的token。这个阶段只训练投影层,视觉编码器和语言模型都冻结。数据用的是CC3M过滤后的595K对图文数据,训练目标就是标准的自回归语言建模loss——给定图片和对应的描述文本,让模型预测描述文本的下一个token。
为什么只训练投影层?因为投影层是随机初始化的,如果同时训练语言模型,随机初始化的投影层会产生噪声梯度,把语言模型已经学好的权重带偏。只训练投影层能让它快速收敛到一个合理的映射,同时不破坏语言模型的能力。
这个阶段通常训1个epoch,学习率2e-3,batch size 128。实测下来,595K数据训1个epoch大概需要8张A100跑20小时左右。如果显存不够,可以把batch size降到32,学习率相应降到5e-4,训练时间会拉长但效果差不多。
4.2 阶段二:视觉指令微调
第二阶段的目标是让模型学会按照指令回答问题。这个阶段训练投影层和语言模型,视觉编码器仍然冻结。数据用的是158K条视觉指令数据,包括对话、详细描述、复杂推理三类。训练目标还是自回归loss,但数据格式变成了多轮对话。
这个阶段的学习率要调小,通常2e-5,batch size 16或32。为什么学习率降这么多?因为语言模型已经预训练好了,大学习率会破坏它的语言能力。我试过用1e-4的学习率,模型在训练后期开始输出重复的、无意义的文本,降到2e-5后恢复正常。
注意:第二阶段的数据格式很关键。LLava用的是Vicuna的对话模板,系统提示词是“A chat between a curious human and an artificial intelligence assistant...”。如果你用自己的模板,需要确保训练和推理时一致,否则模型会困惑。
4.3 训练数据格式与对话模板
LLava的训练数据是JSON格式,每条数据包含id、image、conversations三个字段。conversations是一个列表,每个元素有from和value两个键,from是"human"或"gpt",value是文本内容。图片路径是相对路径,训练时根据image字段加载。
对话模板方面,Vicuna的模板是:
A chat between a curious human and an artificial intelligence assistant. The assistant gives helpful, detailed, and polite answers to the human's questions. USER: <image>\n{question} ASSISTANT: {answer}</s>注意<image>是一个特殊token,在tokenizer里会被替换成256个视觉token的占位符。实际实现时,通常把<image>替换成<im_start><im_patch>*256<im_end>这样的格式,然后在embedding层把<im_patch>的位置替换成投影层的输出。
4.4 显存优化与训练加速技巧
LLava训练最大的瓶颈是显存。7B模型+ViT-L/14,如果全量微调,需要至少8张A100 80G。如果显存不够,有几个降级方案:
第一,用LoRA微调语言模型。LoRA只训练低秩矩阵,参数量减少90%以上,显存占用大幅降低。实测7B模型用LoRA,单张A100 40G就能跑起来。
第二,用梯度检查点(gradient checkpointing)。这个技巧用计算换显存,能把激活值占用的显存降低60%左右,代价是训练速度慢20%-30%。
第三,用DeepSpeed ZeRO-2或ZeRO-3。ZeRO-2把优化器状态和梯度分片,ZeRO-3把参数也分片。8张A100用ZeRO-3可以训13B模型。
第四,降低分辨率。从336降到224,视觉token从576降到256,显存占用降低约35%。
| 优化方案 | 显存降低 | 速度影响 | 效果影响 | 适用场景 |
|---|---|---|---|---|
| LoRA | 约70% | 基本无影响 | 轻微下降 | 显存严重不足 |
| 梯度检查点 | 约60% | 慢20-30% | 无 | 显存中等不足 |
| DeepSpeed ZeRO-3 | 约50% | 慢10-20% | 无 | 多卡环境 |
| 降分辨率 | 约35% | 快30% | 5%以内 | 快速实验 |
5. 实操复现:从环境搭建到推理验证
5.1 环境准备与依赖安装
复现LLava的第一步是搭环境。我推荐用conda建一个干净的虚拟环境,Python 3.10,PyTorch 2.0以上,CUDA 11.8。依赖主要包括transformers、accelerate、bitsandbytes、peft、deepspeed。
conda create -n llava python=3.10 -y conda activate llava pip install torch==2.1.0 torchvision==0.16.0 --index-url https://download.pytorch.org/whl/cu118 pip install transformers==4.36.0 accelerate==0.25.0 bitsandbytes==0.41.0 peft==0.7.0 deepspeed==0.12.0版本兼容性是个大坑。transformers 4.36跟Vicuna的tokenizer兼容性最好,4.37以上有tokenizer加载的bug。bitsandbytes 0.41支持4bit量化,0.42以上有API变动。我踩过的坑是transformers版本太新导致LlamaTokenizer加载失败,报“tokenizer class not found”,降回4.36就好了。
5.2 模型权重下载与合并
LLava的权重分两部分:视觉编码器(CLIP ViT-L/14-336px)和语言模型(Vicuna-7B或Llama-2-7B)。视觉编码器可以从CLIP官方仓库下载,语言模型需要从对应的发布渠道获取。下载完后,需要把投影层的权重合并进去。
合并脚本的核心逻辑是:加载语言模型,加载视觉编码器,加载投影层权重,然后把投影层的state_dict加载到模型里。注意投影层的命名要跟模型定义一致,通常是model.mm_projector.0.weight和model.mm_projector.2.weight。
from llava.model import LlavaLlamaForCausalLM model = LlavaLlamaForCausalLM.from_pretrained( "lmsys/vicuna-7b-v1.5", torch_dtype=torch.float16, device_map="auto" ) model.get_model().mm_projector.load_state_dict(projector_weights)5.3 推理脚本与关键参数
推理时最关键的是<image>token的处理。LLava的tokenizer里,<image>是一个特殊token,id是32000(Vicuna的词表大小是32000)。推理时,先把<image>替换成256个<im_patch>token,然后在embedding层把<im_patch>的位置替换成投影层的输出。
from llava.mm_utils import tokenizer_image_token, process_images input_ids = tokenizer_image_token(prompt, tokenizer, IMAGE_TOKEN_INDEX, return_tensors='pt').unsqueeze(0).cuda() image_tensor = process_images([image], image_processor, model.config).half().cuda() output_ids = model.generate( input_ids, images=image_tensor, do_sample=True, temperature=0.2, max_new_tokens=512, use_cache=True )temperature设0.2比较稳,太高容易胡说,太低容易重复。max_new_tokens设512够用了,除非你要生成很长的描述。
5.4 推理效果验证与常见输出问题
跑通推理后,先拿几张标准测试图验证。我常用的测试集是COCO的几张图,问“Describe this image in detail”和“What is the person doing in this image”。正常情况下,模型应该能给出连贯、准确的描述。
常见问题有几个:第一,模型输出重复文本,比如“a cat a cat a cat”。这通常是temperature太低或者repetition_penalty没设。把repetition_penalty设成1.1-1.2能缓解。第二,模型忽略图片,输出跟图片无关的内容。这通常是投影层权重没加载对,或者<image>token没替换。检查一下embedding层是否真的把<im_patch>替换成了视觉特征。第三,模型输出乱码。这通常是tokenizer版本不匹配,检查一下tokenizer的vocab_size是否跟模型一致。
6. 踩坑记录与排查速查表
6.1 训练不收敛的典型原因
训练loss不降或者震荡,最常见的原因有三个。第一,学习率太大。第二阶段用2e-5,如果误用2e-4,loss会在前几百步震荡然后发散。第二,投影层初始化不对。第二层Linear必须零初始化,否则初期视觉token数值过大,语言模型被噪声梯度带偏。第三,数据格式不对。对话模板必须跟预训练时一致,如果系统提示词变了,模型需要额外时间适应,loss下降会慢很多。
还有一个隐蔽的坑:图片路径错误。如果图片加载失败,返回全黑图,模型学到的就是“不管什么图都输出同样的文本”。训练前一定要检查图片路径,确保每张图都能正常加载。
6.2 显存溢出与降级方案
显存溢出(OOM)是复现LLava最常见的障碍。排查顺序是:先看batch size是不是太大,7B模型全量微调batch size 16需要约60G显存,如果只有40G,降到8或4。再看序列长度,视觉token 576+文本token 512=1088,如果文本很长,序列长度可能超过2048,显存会暴涨。最后看是否开了梯度检查点,没开的话激活值占用很大。
降级方案按优先级排:第一,开梯度检查点,显存降60%,速度慢20%。第二,用LoRA,显存降70%,效果损失5%以内。第三,降分辨率到224,显存降35%。第四,用4bit量化加载语言模型,显存降50%,但训练时量化会引入噪声,效果损失10%左右。
6.3 推理阶段的高频异常
推理阶段的高频异常我整理了一个速查表:
| 异常现象 | 可能原因 | 排查方法 | 解决方案 |
|---|---|---|---|
| 输出重复文本 | temperature过低 | 检查temperature参数 | 调到0.2-0.7,加repetition_penalty |
| 忽略图片内容 | 视觉token未注入 | 检查embedding层替换逻辑 | 确认<im_patch>被替换为视觉特征 |
| 输出乱码 | tokenizer不匹配 | 对比vocab_size | 换用匹配的tokenizer版本 |
| 推理速度极慢 | 未用KV Cache | 检查use_cache参数 | 设use_cache=True |
| 显存溢出 | 序列过长 | 打印input_ids长度 | 截断文本或降分辨率 |
| 输出英文夹杂中文 | 训练数据语言混杂 | 检查训练数据 | 统一训练数据语言 |
6.4 我踩过的三个印象最深的坑
第一个坑是tokenizer的<image>token。LLava的tokenizer里,<image>是added_token,id是32000。但如果加载tokenizer时没调add_special_tokens,这个token不会被加进去,推理时<image>会被拆成普通字符,模型完全看不到图片。我在这卡了大半天,后来打印input_ids才发现问题。
第二个坑是投影层的device。投影层权重加载后,如果没跟语言模型放在同一个device上,推理时会报device mismatch。用device_map="auto"能自动处理,但手动加载时要注意把投影层也放到cuda上。
第三个坑是图片预处理。LLava用的CLIP image processor,归一化的mean和std是CLIP的默认值(mean=[0.481,0.458,0.408],std=[0.268,0.261,0.275])。如果用自己的预处理,数值范围不对,视觉特征会完全跑偏。我试过用ImageNet的mean和std,模型输出全是“I don't know”。
7. 多模态能力的扩展方向与个人实践体会
LLava的架构虽然简单,但扩展性很强。往小了说,你可以换不同的视觉编码器(比如换成SigLIP或EVA-CLIP),换不同的语言模型(比如换成Qwen或Mistral),换不同的投影层(比如换成Q-Former)。往大了说,你可以加视频输入(把多帧的视觉token拼起来),加音频输入(用Whisper编码音频再投影),加3D点云输入(用PointNet编码再投影)。LLava的论文里也提到了这些扩展方向,核心思路都是“把不同模态的输入编码成token,投影到语言模型的嵌入空间,然后拼在一起”。
我个人在实际操作中的体会是,多模态模型的效果上限取决于视觉编码器的质量,下限取决于投影层的训练。CLIP的ViT-L/14已经很强了,如果你换一个更弱的视觉编码器,投影层训得再好也补不回来。反过来,如果投影层没训好,视觉编码器再强,语言模型也看不到有用的信息。所以复现的时候,视觉编码器直接用CLIP的预训练权重,投影层认真训,基本不会太差。
最后分享一个小技巧:如果你只有单卡,想快速验证LLava的效果,可以用4bit量化加载7B模型,然后用LoRA微调投影层和语言模型的q_proj、v_proj。这样单张24G的卡就能跑起来,训练速度大概每小时1000步,训1万步就能看到明显效果。我试过用这个方法在单卡上复现LLava-1.5,在VQAv2上的准确率能到75%左右,跟全量微调的差距在3%以内。对于快速实验来说,这个性价比很高。