论文信息

  • 标题:Visual Instruction Tuning
  • 会议:NeurIPS 2023
  • 单位:威斯康星大学麦迪逊分校、微软研究院、哥伦比亚大学
  • 代码:https://llava-vl.github.io
  • 论文:https://arxiv.org/pdf/2304.08485.pdf

引言:当大语言模型睁开眼睛

想象一下,你给ChatGPT看一张你家冰箱的照片,问它"我能用这些食材做什么晚餐?“,它能立刻给你列出几个菜谱,还能告诉你每道菜的做法。这在2023年之前还是科幻小说里的场景——虽然大语言模型(LLM)能说会道,但它们天生是"瞎子”,根本看不懂图片。

之前的多模态模型比如BLIP-2和Flamingo虽然能处理图像,但它们都是用简单的图文对训练的,只能完成固定任务,比如"描述这张图片"或者"回答这个问题"。你不能像跟朋友聊天一样跟它们说话,更不能让它们帮你解决复杂的实际问题。

LLaVA的作者们想出了一个天才的解决方案:让GPT-4当老师,教一个开源的多模态模型学会听懂人类的视觉指令。他们只用了158K条GPT-4生成的指令数据,就把一个连接了CLIP视觉编码器和Vicuna语言模型的简单模型,训练成了一个能跟你自然对话的视觉助手。它不仅能看懂图片,还能进行复杂的视觉推理,甚至能根据你的手绘草图生成完整的HTML网页!


一、GPT-4当老师:自动生成多模态指令数据

训练一个能听懂指令的模型,首先需要大量高质量的指令数据。但人工标注多模态指令数据既昂贵又耗时,怎么办呢?作者们想到了一个绝妙的主意:用语言-only的GPT-4来生成多模态指令数据

1.1 把图片变成GPT-4能懂的语言

GPT-4虽然能处理文本,但它看不到图片。所以作者们把图片转换成了两种GPT-4能理解的符号表示:

  • 标题(Captions):用多个句子从不同角度描述图片内容
  • 边界框(Boxes):用坐标和类别名标注图片中的物体位置

表格1:指令数据生成示例

上下文类型 内容
标题 一群人站在一辆黑色汽车外面,带着各种行李
行李围绕着地下停车场里的一辆车
人们试图把所有行李塞进一辆SUV
这辆SUV停在公共车库里,正在为旅行打包
一些带着行李的人在一辆正在运输行李的货车旁边
边界框 person: [0.681, 0.242, 0.774, 0.694], backpack: [0.384, 0.696, 0.485, 0.914], suitcase: …
响应类型1:对话 问:图片里是什么类型的车?
答:图片里是一辆黑色的运动型多用途汽车(SUV)
问:车停在哪里?
答:车停在地下停车场,可能是公共车库
问:图片里的人在做什么?
答:他们正在把所有行李塞进SUV,准备去旅行
响应类型2:详细描述 这是一个地下停车场,停着一辆黑色的SUV。场景中有三个人,一个站在车的左侧,一个在中间,一个在右侧。他们正在一起把行李装进SUV准备旅行。停车场里散落着各种行李,包括两个背包和两个行李箱,车的左侧还能看到一辆自行车。周围还停着其他几辆车。
响应类型3:复杂推理 问:这些人面临什么挑战?
答:他们面临的挑战是如何把所有行李都塞进这辆黑色SUV。有多个行李箱和背包需要打包,这说明他们有很多东西要带。他们可能需要巧妙地安排行李的位置,确保所有东西都能放进去。同时,他们还要考虑乘客的舒适度和驾驶时的视野,行李不能挡住司机的视线,也不能让乘客坐得不舒服。
出处:论文表14

通俗解释:这就像你给一个盲人描述一张图片,告诉他图片里有什么东西,每个东西在哪里。然后这个盲人(也就是GPT-4)根据你的描述,生成各种关于这张图片的问题和答案。

1.2 三种类型的指令数据

作者们让GPT-4生成了三种不同类型的指令数据,覆盖了不同的应用场景:

  1. 对话型:模拟人类和AI助手之间的多轮对话,包括物体识别、计数、位置关系等简单问题
  2. 详细描述型:让AI助手对图片进行全面、细致的描述
  3. 复杂推理型:需要结合视觉信息和背景知识进行推理的问题,比如"这些人接下来会做什么?"、“这个场景可能发生在什么地方?”

最终,他们一共生成了158K条独特的语言-图像指令样本,其中包括58K对话、23K详细描述和77K复杂推理。


二、LLaVA模型架构:简单就是美

LLaVA的架构非常简单,它由三个部分组成:一个冻结的视觉编码器、一个可训练的投影层和一个冻结的大语言模型。
在这里插入图片描述

图片1:LLaVA网络架构

出处:论文图1

2.1 视觉编码器

LLaVA使用了预训练的CLIP ViT-L/14作为视觉编码器。对于一张输入图片XvX_vXv,视觉编码器会输出一个视觉特征图ZvZ_vZv
Zv=g(Xv)Z_v = g(X_v)Zv=g(Xv)

  • ZvZ_vZv:视觉特征图,形状为(N,D)(N, D)(N,D),其中NNN是特征图的token数量,DDD是特征维度
  • g(⋅)g(\cdot)g():CLIP视觉编码器函数
  • XvX_vXv:输入图片,形状为(3,H,W)(3, H, W)(3,H,W)

2.2 投影层

视觉特征的维度和语言模型的词嵌入维度通常是不一样的,所以需要一个投影层把视觉特征转换成语言模型能懂的格式:
Hv=W⋅ZvH_v = W \cdot Z_vHv=WZv

  • HvH_vHv:投影后的视觉token序列,形状为(N,d)(N, d)(N,d),其中ddd是语言模型的词嵌入维度
  • WWW:可训练的投影矩阵,形状为(D,d)(D, d)(D,d)
  • ZvZ_vZv:视觉编码器输出的特征图

通俗解释:这就像一个翻译官,把视觉编码器说的"视觉语言"翻译成语言模型能听懂的"自然语言"。

2.3 大语言模型

LLaVA使用了Vicuna作为语言模型,它是一个基于LLaMA的开源聊天机器人,性能接近GPT-3.5。投影后的视觉tokenHvH_vHv会被拼接到输入文本的前面,作为软视觉提示,然后一起输入到语言模型中。

语言模型会自回归地生成回答,其概率可以表示为:
p(Xa∣Xv,Xinstruct)=∏i=1Lpθ(xi∣Xv,Xinstruct,<i,Xa,<i)p(X_a | X_v, X_{instruct}) = \prod_{i=1}^{L} p_{\theta}(x_i | X_v, X_{instruct,<i}, X_{a,<i})p(XaXv,Xinstruct)=i=1Lpθ(xiXv,Xinstruct,<i,Xa,<i)

  • XaX_aXa:模型生成的回答序列
  • XinstructX_{instruct}Xinstruct:用户的指令序列
  • Xinstruct,<iX_{instruct,<i}Xinstruct,<i:第iii个token之前的所有指令token
  • Xa,<iX_{a,<i}Xa,<i:第iii个token之前的所有回答token
  • pθ(⋅)p_{\theta}(\cdot)pθ():语言模型的概率分布函数
  • LLL:回答序列的长度

三、两阶段训练策略:先对齐,再微调

LLaVA采用了两阶段的训练策略,这样可以在保证性能的同时,大大减少训练成本。

3.1 第一阶段:特征对齐预训练

在这个阶段,作者们使用了过滤后的CC3M数据集(595K条图文对),只训练投影层WWW,视觉编码器和语言模型都保持冻结。

这个阶段的目标是让投影层学会把视觉特征映射到语言模型的词嵌入空间,让语言模型能够"看懂"视觉特征。训练只需要1个epoch,在8张A100上大约需要4小时。

3.2 第二阶段:端到端指令微调

在这个阶段,作者们使用了之前生成的158K条指令数据,同时训练投影层和语言模型,视觉编码器仍然保持冻结。

训练分为两个具体的场景:

  1. 多模态聊天机器人:在全部158K条数据上微调,让模型学会跟人类进行自然的视觉对话
  2. 科学问答:在ScienceQA数据集上微调,让模型学会解答多模态科学问题

这个阶段训练3个epoch,在8张A100上大约需要10小时。


四、实验结果:用最少的数据打最好的仗

LLaVA在多个任务上都取得了令人惊艳的结果,特别是在零样本指令跟随能力上,远超之前的多模态模型。

4.1 多模态聊天能力:媲美GPT-4

作者们首先用GPT-4论文中的经典例子测试了LLaVA的视觉推理能力。

表格2:极端熨衣服例子对比

模型 回答
用户 这张图片有什么不寻常的地方?
LLaVA 这张图片不寻常的地方是一个男人在一辆小型货车的后面熨衣服。这不是一个典型的熨衣服的地方,人们通常会在更固定和安全的地方,比如家里,用普通的熨衣板熨衣服。这个场景很奇特,因为它涉及到一个临时搭建在车辆上的熨衣服装置,这既不安全也不符合常规。而且,不清楚这个男人在这样不稳定的环境中如何保持平衡和稳定。
GPT-4 这张图片不寻常的地方是一个男人在一辆行驶中的出租车的车顶上熨衣服。
BLIP-2 一个男人坐在一辆黄色出租车的后面。
OpenFlamingo 这个男人正在他的汽车引擎盖上晾衣服。
出处:论文表3

分析:可以看到,BLIP-2和OpenFlamingo只是简单地描述了图片内容,完全没有理解用户的问题。而LLaVA不仅准确地回答了问题,还给出了比GPT-4更详细的解释。

4.2 定量评估:85.1%的GPT-4水平

为了定量评估模型的指令跟随能力,作者们提出了一个新颖的评估方法:用GPT-4作为评委,给不同模型的回答打分。

表格3:LLaVA-Bench(COCO)不同训练数据的消融实验

训练数据 对话 详细描述 复杂推理 总体
全部数据 83.1 75.3 96.5 85.1
详细+复杂 81.5 (-1.6) 73.3 (-2.0) 90.8 (-5.7) 81.9 (-3.2)
对话+5%详细+10%复杂 81.0 (-2.1) 68.4 (-7.1) 91.5 (-5.0) 80.5 (-4.4)
只有对话 76.5 (-6.6) 59.8 (-16.2) 84.9 (-12.4) 73.8 (-11.3)
没有指令微调 22.0 (-61.1) 24.0 (-51.3) 18.5 (-78.0) 21.5 (-63.6)
出处:论文表4

分析

  1. 指令微调的效果非常显著,没有指令微调的模型得分只有21.5%
  2. 三种类型的数据都很重要,全部数据一起训练能得到最好的效果
  3. 复杂推理类型的数据对提升模型的整体能力贡献最大

表格4:与其他模型的对比(LLaVA-Bench In-the-Wild)

模型 对话 详细描述 复杂推理 总体
OpenFlamingo 19.3 ± 0.5 19.0 ± 0.5 19.1 ± 0.7 19.1 ± 0.4
BLIP-2 54.6 ± 1.4 29.1 ± 1.2 32.9 ± 0.7 38.1 ± 1.0
LLaVA 57.3 ± 1.9 52.5 ± 6.3 81.7 ± 1.8 67.3 ± 2.0
出处:论文表5

分析:LLaVA的总体得分比BLIP-2高29%,比OpenFlamingo高48%,特别是在复杂推理任务上,LLaVA达到了81.7%的高分,这说明它具备了很强的视觉推理能力。

4.3 ScienceQA:刷新SOTA记录

ScienceQA是一个大规模的多模态科学问答数据集,包含21K个问题,覆盖自然科学、社会科学和语言科学三个学科。

表格5:ScienceQA结果对比

方法 平均准确率
人类 88.40
GPT-3.5 73.97
GPT-3.5 w/ CoT 75.17
LLaMA-Adapter 85.19
MM-CoT Large 91.68
GPT-4 (文本-only) 82.69
LLaVA 90.92
LLaVA+GPT-4 (互补) 90.97
LLaVA+GPT-4 (法官) 92.53
出处:论文表7

分析

  1. LLaVA单独使用就达到了90.92%的准确率,非常接近之前的SOTA(91.68%)
  2. 当把LLaVA和GPT-4结合起来,用GPT-4作为法官来裁决两个模型的不同答案时,达到了92.53%的新SOTA,甚至超过了人类的平均水平!

4.4 消融实验:哪些设计最重要?

作者们还做了详细的消融实验,分析了各个设计选择对性能的影响。

表格6:设计选择消融实验

设计选择 准确率 变化
最佳变体(倒数第二层特征) 90.92 -
最后一层特征 89.96 -0.96
先预测答案再推理 89.77 -1.15
不进行预训练直接训练 85.81 -5.11
7B模型 89.84 -1.08
出处:论文表8

分析

  1. 使用视觉编码器倒数第二层的特征比最后一层效果更好,因为倒数第二层保留了更多的局部细节信息
  2. 先进行特征对齐预训练非常重要,能带来5.11%的性能提升
  3. 模型规模越大,性能越好,13B模型比7B模型高1.08%

五、核心代码实现

下面是LLaVA的核心代码实现,包括模型架构和前向传播过程:

import torch
import torch.nn as nn
from transformers import CLIPVisionModel, CLIPImageProcessor, AutoTokenizer, AutoModelForCausalLM

class LLaVA(nn.Module):
    def __init__(
        self,
        vision_model_name="openai/clip-vit-large-patch14",
        language_model_name="lmsys/vicuna-7b-v1.5",
        freeze_vision=True,
        freeze_language=True
    ):
        super().__init__()
        
        # 视觉编码器和图像处理器
        self.vision_encoder = CLIPVisionModel.from_pretrained(vision_model_name)
        self.image_processor = CLIPImageProcessor.from_pretrained(vision_model_name)
        
        # 投影层:将视觉特征映射到语言模型的词嵌入维度
        self.vision_proj = nn.Linear(
            self.vision_encoder.config.hidden_size,
            AutoModelForCausalLM.from_pretrained(language_model_name).config.hidden_size
        )
        
        # 语言模型和分词器
        self.tokenizer = AutoTokenizer.from_pretrained(language_model_name)
        self.language_model = AutoModelForCausalLM.from_pretrained(language_model_name)
        
        # 冻结模型参数
        if freeze_vision:
            for param in self.vision_encoder.parameters():
                param.requires_grad = False
        
        if freeze_language:
            for param in self.language_model.parameters():
                param.requires_grad = False
        
        # 特殊token
        self.image_token = "<image>"
        self.image_token_id = self.tokenizer.convert_tokens_to_ids(self.image_token)
        
        # 添加image token到分词器
        if self.image_token_id == self.tokenizer.unk_token_id:
            self.tokenizer.add_tokens([self.image_token])
            self.language_model.resize_token_embeddings(len(self.tokenizer))
            self.image_token_id = self.tokenizer.convert_tokens_to_ids(self.image_token)

    def encode_images(self, images):
        """编码图像为视觉特征"""
        # 预处理图像
        pixel_values = self.image_processor(images, return_tensors="pt").pixel_values
        pixel_values = pixel_values.to(self.vision_encoder.device, dtype=self.vision_encoder.dtype)
        
        # 提取视觉特征(使用倒数第二层)
        with torch.no_grad():
            vision_outputs = self.vision_encoder(pixel_values, output_hidden_states=True)
            # 取倒数第二层的特征
            image_features = vision_outputs.hidden_states[-2]
            # 去掉CLS token
            image_features = image_features[:, 1:, :]
        
        # 投影到语言模型维度
        image_embeds = self.vision_proj(image_features)
        
        return image_embeds

    def prepare_inputs(self, conversations, images=None):
        """准备模型输入"""
        input_ids = []
        attention_masks = []
        labels = []
        
        for conv in conversations:
            # 构建提示词
            prompt = ""
            for turn in conv:
                if turn["role"] == "user":
                    prompt += f"USER: {turn['content']} "
                elif turn["role"] == "assistant":
                    prompt += f"ASSISTANT: {turn['content']}</s>"
            
            # 替换<image>为实际的token
            if images is not None and "<image>" in prompt:
                prompt = prompt.replace("<image>", self.image_token)
            
            # 分词
            encoded = self.tokenizer(
                prompt,
                return_tensors="pt",
                padding="max_length",
                truncation=True,
                max_length=2048
            )
            
            input_ids.append(encoded["input_ids"])
            attention_masks.append(encoded["attention_mask"])
            
            # 构建标签:只计算assistant回答部分的损失
            label = encoded["input_ids"].clone()
            # 找到所有USER:的位置
            user_positions = (label == self.tokenizer.encode("USER:", add_special_tokens=False)[0]).nonzero()
            for pos in user_positions:
                start = pos[1] + len("USER:")
                # 找到下一个ASSISTANT:的位置
                assistant_pos = (label[0, start:] == self.tokenizer.encode("ASSISTANT:", add_special_tokens=False)[0]).nonzero()
                if len(assistant_pos) > 0:
                    end = start + assistant_pos[0][1] + len("ASSISTANT:")
                    # 将USER部分的标签设为-100(忽略)
                    label[0, :end] = -100
            
            labels.append(label)
        
        input_ids = torch.cat(input_ids, dim=0)
        attention_masks = torch.cat(attention_masks, dim=0)
        labels = torch.cat(labels, dim=0)
        
        return input_ids, attention_masks, labels

    def forward(self, input_ids, attention_mask, labels=None, images=None):
        """前向传播"""
        # 获取词嵌入
        inputs_embeds = self.language_model.get_input_embeddings()(input_ids)
        
        # 如果有图像,替换<image> token的嵌入为视觉特征
        if images is not None:
            image_embeds = self.encode_images(images)
            # 找到所有<image> token的位置
            image_positions = (input_ids == self.image_token_id).nonzero()
            for i, pos in enumerate(image_positions):
                batch_idx, token_idx = pos
                # 替换为对应的视觉特征
                inputs_embeds[batch_idx, token_idx:token_idx+image_embeds.size(1), :] = image_embeds[i]
        
        # 前向传播
        outputs = self.language_model(
            inputs_embeds=inputs_embeds,
            attention_mask=attention_mask,
            labels=labels
        )
        
        return outputs

    def generate(self, input_ids, attention_mask, images=None, **kwargs):
        """生成回答"""
        # 获取词嵌入
        inputs_embeds = self.language_model.get_input_embeddings()(input_ids)
        
        # 如果有图像,替换<image> token的嵌入为视觉特征
        if images is not None:
            image_embeds = self.encode_images(images)
            # 找到所有<image> token的位置
            image_positions = (input_ids == self.image_token_id).nonzero()
            for i, pos in enumerate(image_positions):
                batch_idx, token_idx = pos
                # 替换为对应的视觉特征
                inputs_embeds[batch_idx, token_idx:token_idx+image_embeds.size(1), :] = image_embeds[i]
        
        # 生成回答
        outputs = self.language_model.generate(
            inputs_embeds=inputs_embeds,
            attention_mask=attention_mask,
            **kwargs
        )
        
        return outputs

# 使用示例
if __name__ == "__main__":
    # 加载模型
    model = LLaVA()
    
    # 准备输入
    conversations = [
        [
            {"role": "user", "content": "<image> 这张图片里有什么?"}
        ]
    ]
    
    # 假设我们有一张图片
    from PIL import Image
    image = Image.open("example.jpg")
    
    # 准备输入
    input_ids, attention_mask, _ = model.prepare_inputs(conversations, images=[image])
    
    # 生成回答
    outputs = model.generate(
        input_ids,
        attention_mask,
        images=[image],
        max_new_tokens=100,
        temperature=0.7
    )
    
    # 解码回答
    print(model.tokenizer.decode(outputs[0], skip_special_tokens=True))

六、LLaVA的神奇能力

LLaVA不仅在学术基准上表现出色,还展现出了很多令人惊喜的 emergent 能力。

6.1 根据手绘草图生成HTML代码

LLaVA能看懂你的手绘草图,然后生成完整的HTML/JS/CSS代码,实现一个可交互的网页!
在这里插入图片描述

图片2:根据手绘草图生成HTML代码

出处:论文图2

6.2 识别名人表情包

LLaVA能识别出埃隆·马斯克,即使他穿着狗狗币的服装出现在表情包里!
在这里插入图片描述

图片3:识别埃隆·马斯克表情包

出处:论文图6

6.3 理解网络迷因

LLaVA能看懂网络迷因,理解其中的幽默之处。

表格7:鸡肉块地图迷因理解

模型 回答
用户 你能详细解释这个迷因吗?
LLaVA 这个迷因是对食物的一种创造性和幽默的表达,以鸡肉块为宇宙的中心。迷因开头说"有时候我只是看着从太空拍的地球照片,惊叹于它的美丽…"。主图是一盘煮熟的鸡肉块,鸡肉块的位置看起来像世界地图上的大陆和岛屿。迷因开玩笑地说这些鸡肉块代表地球,照片里的不同地方其实是不同地方的鸡肉块版本。迷因的笑点在底部的小字:“我的意思是,这不是真正的地球,但它是多么美丽啊。”
GPT-4 这个迷因是一个笑话,它把两个不相关的东西结合在一起:从太空拍的地球照片和鸡肉块。迷因的文字让你以为下面的图片是一张美丽的地球太空照,但实际上是一盘鸡肉块,它们的排列模糊地像一张世界地图。幽默来自于文字和图片的意外并置。
BLIP-2 有时候我只是看着从太空拍的地球照片,惊叹于它的美丽
OpenFlamingo 这是一张国际空间站上的鸡肉块照片。
出处:论文表9

七、局限性与未来方向

虽然LLaVA已经非常强大,但它仍然有一些局限性:

  1. 知识局限性:它继承了Vicuna的知识,对于一些专业领域的问题可能回答不准确
  2. 幻觉问题:和所有大语言模型一样,LLaVA有时会生成不符合事实的内容
  3. 细粒度理解不足:对于一些需要非常精确的细节识别的任务,比如识别冰箱里酸奶的品牌,LLaVA可能会出错
  4. 上下文学习能力有限:目前的LLaVA还不支持少样本上下文学习

未来的研究方向包括:

  • 使用更大、更强的基础模型,比如LLaMA 3和CLIP ViT-G
  • 生成更多样、更高质量的指令数据
  • 支持多轮对话和上下文学习
  • 结合工具使用能力,让LLaVA能调用外部工具来解决更复杂的问题

总结

LLaVA是多模态大模型发展史上的一个里程碑式的工作。它证明了用语言-only的大模型生成多模态指令数据,然后微调一个简单的多模态模型是一种非常有效且低成本的方法。

LLaVA的成功告诉我们:

  1. 简单就是美:有时候最简单的架构反而能取得最好的效果
  2. 数据为王:高质量的指令数据比复杂的模型架构更重要
  3. 站在巨人的肩膀上:充分利用已经训练好的单模态模型,可以大大降低多模态模型的训练成本

LLaVA的开源也极大地推动了多模态大模型的研究和应用。现在,任何人都可以在自己的电脑上运行一个强大的视觉助手,这为很多新的应用场景打开了大门。


更多推荐