ARTICLE DETAIL

资讯详情

深耕网站建设与运营推广的一线实战洞察。

从CLIP到LLava:多模态模型原理与复现指南

从CLIP到LLava:多模态模型原理与复现指南 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 ResNetCLIP论文里试了多种图像编码器最终ViT-L/14表现最好。但这里有个容易被忽略的细节ViT在小数据上不如ResNet在大数据上反超。CLIP用了4亿对数据所以ViT的优势能发挥出来。如果你自己的数据只有几十万对用ResNet可能更稳。具体对比一下ResNet-50的输出是2048维经过一个投影层降到512或768ViT-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/32ViT-Base, patch 32约86M512最快快速原型验证ViT-B/16ViT-Base, patch 16约86M512中等平衡精度与速度ViT-L/14ViT-Large, patch 14约304M768最慢高精度检索2.4 文本编码器的细节与因果掩码的取舍CLIP的文本编码器是一个12层、512宽、8头的Transformer。这里有一个设计选择值得注意它用的是因果掩码causal mask还是双向注意力答案是因果掩码。文本序列经过Transformer后取最后一个tokenEOT 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投影层是一个两层MLPLinear - GELU - Linear语言模型用Vicuna或Llama。数据流是这样的一张图片经过ViT编码得到256个patch token每个1024维投影层把这256个token映射到语言模型的词嵌入空间比如Vicuna的4096维然后这256个视觉token和文本token拼在一起送入语言模型。语言模型看到的是“256个视觉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用的是224x224LLava-1.5升级到了336x336。336x336的ViT-L/14会产生576个patch token24x24比224的256个多了一倍多。更多的视觉token意味着更细粒度的视觉信息但也意味着更长的序列和更大的显存占用。实测下来336分辨率在OCR和细粒度识别任务上提升明显但在一般对话任务上差距不大。提示如果你显存有限可以用224分辨率的ViT-L/14然后把投影层的输出维度对齐到语言模型的隐藏维度。显存占用能降40%左右效果损失在5%以内。3.3 投影层的初始化与训练稳定性投影层虽然只有两层MLP但初始化方式对训练稳定性影响很大。LLava官方实现里第一层Linear的权重用正态分布初始化std0.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-3batch size 128。实测下来595K数据训1个epoch大概需要8张A100跑20小时左右。如果显存不够可以把batch size降到32学习率相应降到5e-4训练时间会拉长但效果差不多。4.2 阶段二视觉指令微调第二阶段的目标是让模型学会按照指令回答问题。这个阶段训练投影层和语言模型视觉编码器仍然冻结。数据用的是158K条视觉指令数据包括对话、详细描述、复杂推理三类。训练目标还是自回归loss但数据格式变成了多轮对话。这个阶段的学习率要调小通常2e-5batch 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或gptvalue是文本内容。图片路径是相对路径训练时根据image字段加载。对话模板方面Vicuna的模板是A chat between a curious human and an artificial intelligence assistant. The assistant gives helpful, detailed, and polite answers to the humans questions. USER: image\n{question} ASSISTANT: {answer}/s注意image是一个特殊token在tokenizer里会被替换成256个视觉token的占位符。实际实现时通常把image替换成im_startim_patch*256im_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.10PyTorch 2.0以上CUDA 11.8。依赖主要包括transformers、accelerate、bitsandbytes、peft、deepspeed。conda create -n llava python3.10 -y conda activate llava pip install torch2.1.0 torchvision0.16.0 --index-url https://download.pytorch.org/whl/cu118 pip install transformers4.36.0 accelerate0.25.0 bitsandbytes0.41.0 peft0.7.0 deepspeed0.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_dtypetorch.float16, device_mapauto ) model.get_model().mm_projector.load_state_dict(projector_weights)5.3 推理脚本与关键参数推理时最关键的是imagetoken的处理。LLava的tokenizer里image是一个特殊tokenid是32000Vicuna的词表大小是32000。推理时先把image替换成256个im_patchtoken然后在embedding层把im_patch的位置替换成投影层的输出。from llava.mm_utils import tokenizer_image_token, process_images input_ids tokenizer_image_token(prompt, tokenizer, IMAGE_TOKEN_INDEX, return_tensorspt).unsqueeze(0).cuda() image_tensor process_images([image], image_processor, model.config).half().cuda() output_ids model.generate( input_ids, imagesimage_tensor, do_sampleTrue, temperature0.2, max_new_tokens512, use_cacheTrue )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能缓解。第二模型忽略图片输出跟图片无关的内容。这通常是投影层权重没加载对或者imagetoken没替换。检查一下embedding层是否真的把im_patch替换成了视觉特征。第三模型输出乱码。这通常是tokenizer版本不匹配检查一下tokenizer的vocab_size是否跟模型一致。6. 踩坑记录与排查速查表6.1 训练不收敛的典型原因训练loss不降或者震荡最常见的原因有三个。第一学习率太大。第二阶段用2e-5如果误用2e-4loss会在前几百步震荡然后发散。第二投影层初始化不对。第二层Linear必须零初始化否则初期视觉token数值过大语言模型被噪声梯度带偏。第三数据格式不对。对话模板必须跟预训练时一致如果系统提示词变了模型需要额外时间适应loss下降会慢很多。还有一个隐蔽的坑图片路径错误。如果图片加载失败返回全黑图模型学到的就是“不管什么图都输出同样的文本”。训练前一定要检查图片路径确保每张图都能正常加载。6.2 显存溢出与降级方案显存溢出OOM是复现LLava最常见的障碍。排查顺序是先看batch size是不是太大7B模型全量微调batch size 16需要约60G显存如果只有40G降到8或4。再看序列长度视觉token 576文本token 5121088如果文本很长序列长度可能超过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_cacheTrue显存溢出序列过长打印input_ids长度截断文本或降分辨率输出英文夹杂中文训练数据语言混杂检查训练数据统一训练数据语言6.4 我踩过的三个印象最深的坑第一个坑是tokenizer的imagetoken。LLava的tokenizer里image是added_tokenid是32000。但如果加载tokenizer时没调add_special_tokens这个token不会被加进去推理时image会被拆成普通字符模型完全看不到图片。我在这卡了大半天后来打印input_ids才发现问题。第二个坑是投影层的device。投影层权重加载后如果没跟语言模型放在同一个device上推理时会报device mismatch。用device_mapauto能自动处理但手动加载时要注意把投影层也放到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 dont 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%以内。对于快速实验来说这个性价比很高。
返回列表