Llava的结构
LLava的架构很简单就是在LLM的前面挂一个视觉encoder,这个encoder是预训练过的,LLava是用的CLIP的视觉encoder,但这里不是得到一个最终的图片向量,而是一个更细粒度的图片patch_token向量,然后再过一个Projection,就是一个FFN,一个线性层+激活函数再加一个线性层,但是维度变换只是将patch向量的维度和LLM中的隐藏层维度对齐,而没有升维再降维度的操作。 图片和文本模态的向量试直接拼在一起的。

这里用一个小的Llava模型(tiny-LlavaForConditionalGeneration)debug一下llava,走一遍前向传播。
代码的init函数如下:
class LlavaModel(LlavaPreTrainedModel):
def __init__(self, config: LlavaConfig):
super().__init__(config)
self.vision_tower = AutoModel.from_config(config.vision_config)
self.multi_modal_projector = LlavaMultiModalProjector(config)
self.language_model = AutoModel.from_config(config.text_config)
self.post_init()LlavaProcessor
在图文对输入进模型时,需要对图片和文本进行预处理,文本预处理就是tokenizer那套(但又有区别),比较陌生的是image_processor。
LlavaProcessor包含两个部分,分别是image_processor和tokenizer
image_processor
入口位置:src/transformers/models/llava/processing_llava.py:141
image_inputs = self.image_processor(images, **output_kwargs["images_kwargs"])而这里的image_processor为CLIPImageProcessor
前面初始化了一堆参数

if do_resize:
image = self.resize(image=image, size=size, resample=resample, input_data_format=input_data_format)
if do_center_crop:
image = self.center_crop(image=image, size=crop_size, input_data_format=input_data_format)
if do_rescale:
image = self.rescale(image=image, scale=rescale_factor, input_data_format=input_data_format)
if do_normalize:
image = self.normalize(
image=image, mean=image_mean, std=image_std, input_data_format=input_data_format
)在 image_processor 阶段,会对原始图片进行预处理,包括 resize、center_crop、rescale 和 normalize 操作,将图片转换成模型需要的固定格式:[Batch_size,C,H,W]
其中,resize、center_crop操作的原因是:
LLaVA 使用的视觉编码器是 CLIP 的 Vision Transformer(ViT)。ViT 与传统 CNN 不同,它需要固定大小的输入图片。因为 ViT 会将图片划分成固定大小的 patch,并将每个 patch 作为一个 token 输入 Transformer,而 Transformer 需要对应的位置编码,所以要求尺寸固定
当前模型中 CLIP ViT 的配置为:image_size: 336
因此输入任意一张图片,在进入 Vision Encoder 之前,image processor 会对输入图片进行 resize和center crop操作,使其转换为模型预期的336×336。
而resize操作的作用保持图片宽高比例,将图片短边缩放到指定大小,一般是将图片短边设置成image_size,这里是336,而长边要保持和原来短边一样的比例:long=image_size×(long/short)
center crop操作的作用是:从resize后的图片中心裁剪出模型需要的固定尺寸。就是取长边的最中间的地方。
这样的操作就意味着,llava(或者说VIT)对一些长宽相差非常大的图片的处理效果不会很好
而rescale 和 normalize操作就是CV里经典的操作了
rescale将图片像素从0~255缩放到0~1(原地除以255),而 normalize是进一步对图片进行归一化化。它会使用预先计算好的均值(mean)和标准差(std)进行归一化。这里的 mean 和 std 不是随便选的,而是 CLIP 在大规模图文数据预训练时统计得到的
image_token的预填充
经过image_processor之后,就会得到image_inputs['pixel_values']
image_inputs['pixel_values'].shape
torch.Size([1, 3, 336, 336])可以看到是一个[B,C,H,W]大小的图片像素了。在进入tokenizer之前,要进行一个预填充操作
pixel_values = image_inputs["pixel_values"]
height, width = get_image_size(to_numpy_array(pixel_values[0]))
num_image_tokens = (height // self.patch_size) * (
width // self.patch_size
) + self.num_additional_image_tokens
if self.vision_feature_select_strategy == "default":
num_image_tokens -= 1这里在算num_image_tokens,这里patch_size是14×14,那就是(336/14)×(336/14)=24×24=576,而这里的self.num_additional_image_tokens是VIT中的cls向量,这里也给他减去了,那就是576个image_token,然后就是填充:
prompt_strings = []
for sample in text:
sample = sample.replace(self.image_token, self.image_token * num_image_tokens)
prompt_strings.append(sample)这里text是我们的输入文本,原本是带了一个<image>占位符
['USER: <image>\nDescribe the image briefly. ASSISTANT:']这一步操作就是将这个<image>占位符复制576份,这里就是将imagetoken预填充好。
tokenizer
这里和LLM中的tokenizer一样,输出包含:input_ids,attention_mask。还有一个额外的图片的pixel_values。其中input_ids包含<image>作为一个特殊token,token_id为3200。后续会进入模型的视觉塔,视觉层会将pixel_values的patch_token向量替换到3200中。
LlavaModel
token_embedding
processor之后,就进入Llava模型内部了,核心输入就是上述的input_ids,attention_mask,pixel_values
if inputs_embeds is None:
inputs_embeds = self.get_input_embeddings()(input_ids)
if pixel_values is not None:
image_features = self.get_image_features(
pixel_values=pixel_values,
vision_feature_layer=vision_feature_layer,
vision_feature_select_strategy=vision_feature_select_strategy,
image_sizes=image_sizes,
)
image_features = torch.cat(image_features, dim=0).to(inputs_embeds.device, inputs_embeds.dtype)
special_image_mask = self.get_placeholder_mask(
input_ids, inputs_embeds=inputs_embeds, image_features=image_features
)
inputs_embeds = inputs_embeds.masked_scatter(special_image_mask, image_features) inputs_embeds = self.get_input_embeddings()(input_ids)这步操作首先会将inputs_id经过一个nn.Embedding,
input_ids.shape
torch.Size([1, 594])
inputs_embeds.shape
torch.Size([1, 594, 8])因为我用的是tiny-LlavaForConditionalGeneration这个专门用来调试的小模型,所以隐藏层的维度是8,正常这个维度会很大
这里通过调试控制台可以看到这一步将所有token,包含<image>占位符的token也都进行了embedding。
image_features = self.get_image_features(
pixel_values=pixel_values,
vision_feature_layer=vision_feature_layer,
vision_feature_select_strategy=vision_feature_select_strategy,
image_sizes=image_sizes,
)这一步就是将pixel_values映射成向量了,这个函数内部会将pixel_values进入视觉塔:
image_outputs = self.vision_tower(pixel_values, output_hidden_states=True, **kwargs)CLIPVisionModel(VIT和Transformer的区别)
而视觉塔是CLIPVisionModel,就是CLIP的视觉encoder,这里主要看一下怎么对pixel_values向量化(也就是VIT中的切成patch再向量化)核心代码如下:
def __init__(self, config: CLIPVisionConfig):
self.embed_dim = config.hidden_size
self.image_size = config.image_size
self.patch_size = config.patch_size
self.class_embedding = nn.Parameter(torch.randn(self.embed_dim))
self.patch_embedding = nn.Conv2d(
in_channels=config.num_channels,
out_channels=self.embed_dim,
kernel_size=self.patch_size,
stride=self.patch_size,
bias=False,
)
self.num_patches = (self.image_size // self.patch_size) ** 2
self.num_positions = self.num_patches + 1
self.position_embedding = nn.Embedding(self.num_positions, self.embed_dim)
def forward(self, pixel_values: torch.FloatTensor, interpolate_pos_encoding=False) -> torch.Tensor:
batch_size, _, height, width = pixel_values.shape
if not interpolate_pos_encoding and (height != self.image_size or width != self.image_size):
raise ValueError(
f"Input image size ({height}*{width}) doesn't match model ({self.image_size}*{self.image_size})."
)
target_dtype = self.patch_embedding.weight.dtype
patch_embeds = self.patch_embedding(pixel_values.to(dtype=target_dtype)) # shape = [*, width, grid, grid]
patch_embeds = patch_embeds.flatten(2).transpose(1, 2)
class_embeds = self.class_embedding.expand(batch_size, 1, -1)
embeddings = torch.cat([class_embeds, patch_embeds], dim=1)
if interpolate_pos_encoding:
embeddings = embeddings + self.interpolate_pos_encoding(embeddings, height, width)
else:
embeddings = embeddings + self.position_embedding(self.position_ids)
return embeddings主要就核心两步:
- 切patch:
- 将每个patch看作一个单独的token,进行embedding
这两部都被一个卷积操作实现了使用kernel_size=14,stride=14即卷积窗口为14×14,他在这张图片上每次移动14步,每个14×14像素被映射成embed_dim。代码如下:patch_embeds = self.patch_embedding(pixel_values.to(dtype=target_dtype)) # shape = [*, width, grid, grid]pixel_values.shape torch.Size([1, 3, 336, 336])#(B,C,H,W)
patch_embeds.shape
torch.Size([1, 8, 24, 24])#(B,embed_dim,patch_num,patch_num)
之后:
```python
patch_embeds = patch_embeds.flatten(2).transpose(1, 2)这里还要将patch×patch扁平化成为seq_len。最终形成:([1, 576, 8]):即[B,Seq_len,embed_dim]
紧接着就是
class_embeds = self.class_embedding.expand(batch_size, 1, -1)
embeddings = torch.cat([class_embeds, patch_embeds], dim=1)这里class_embeds就是CLIP中的cls向量,再往后就是加一下位置编码。
这个出来之后,后边的流程和普通的Transformer没什么区别,VIT和Transformer的区别就是怎么将图片切成token。
Projector
从视觉塔出来之后,得到每层的向量,Llava原文是取的倒数第二层,没有取最后一层因为他觉得CLIP 最后一层。更偏向图文对齐任务。倒数第二层保留更多视觉细节。所以取倒数第二层的向量,但是源码里也保留了可以拼接多层的操作。
得到该层的向量之后,这时候就是过Projector层了,整个projector类都非常简洁,就是两个线性层,中间套一个激活函数,就类似FFN层,不过没有FFN层的升维再降维
class LlavaMultiModalProjector(nn.Module):
def __init__(self, config: LlavaConfig):
super().__init__()
# We have hidden_size * the number of vision feature layers
num_feature_layers = 1 if isinstance(config.vision_feature_layer, int) else len(config.vision_feature_layer)
self.linear_1 = nn.Linear(
config.vision_config.hidden_size * num_feature_layers,
config.text_config.hidden_size,
bias=config.multimodal_projector_bias,
)
self.act = ACT2FN[config.projector_hidden_act]
self.linear_2 = nn.Linear(
config.text_config.hidden_size, config.text_config.hidden_size, bias=config.multimodal_projector_bias
)
def forward(self, image_features):
hidden_states = self.linear_1(image_features)
hidden_states = self.act(hidden_states)
hidden_states = self.linear_2(hidden_states)
return hidden_states这里num_feature_layers参数就是如果取了拼接多层的操作之后的这个多层的层数
Projector的作用说白了就是对齐,Llava的训练初期就是冻住其他层,只训练Projector这个层,先初步对齐视觉和文本向量。之后再冻住视觉塔,只激活LLM层进行指令微调。
替换预填充的<image>向量
上边说到我们目前的inputs_embeds还充斥着大量的通过<image>产生的向量,这时候我们就需要使用image_features替代这些向量
image_features = torch.cat(image_features, dim=0).to(inputs_embeds.device, inputs_embeds.dtype)#原本image_features是列表,这里将他们从第一个维度上cat一下,得到(B,S,H)
special_image_mask = self.get_placeholder_mask(
input_ids, inputs_embeds=inputs_embeds, image_features=image_features
)
inputs_embeds = inputs_embeds.masked_scatter(special_image_mask, image_features)get_placeholder_mask就是就是生成一个mask来代表哪些位置原来是<image>token。逻辑就是将token_id为image_token_id的标记一下,Llava里是3200
special_image_mask = input_ids == self.config.image_token_id最后那个masked_scatter是pytorch的自带操作,就是将这些mask标记好的位置替换成对应的image_feature向量。
Language Model
outputs = self.language_model(
attention_mask=attention_mask,
position_ids=position_ids,
past_key_values=past_key_values,
inputs_embeds=inputs_embeds,
cache_position=cache_position,
**kwargs,
)这里就是进入llama了。可以看到试直接输入 inputs_embeds而不是input_ids。之前总疑惑为啥要设置一个接口可以传入 inputs_embeds。现在明白了是为多模态模型留的。
这里就不细说了。已经很了解了。
Llava的训练
LLaVA采用两阶段训练策略。
第一阶段的目标是让 CLIP 提取的视觉特征能够进对齐LLM的语言特征空间,其中Vision Encoder和LLM层都是冻住的,只训练Projector层
第二阶段就是Visual Instruction Tuning,目标是让模型具备多模态对话和指令跟随能力,这个期间还是冻住Vision Encoder,但是解冻LLM层,让LLM层和Projector层一块训练。








