今天把Senna模型的类间关系详细的梳理了一下,画了个uml类图,分享给大家,希望对大家也有所帮助。

以前写过关于Senna的文章:

Senna多模态大模型中关键数据及代码解析

上一篇文章相当于从核心/关键数据流转的角度来解读,可以直观的看到输入,输出的样例。本篇换个角度,从代码结构或类图的角度来进行详细的解读。

一,类图解析

1,SennaLlavaLlamaForCausalLM类是Senna模型的主类。

2,图中Extends标记代表继承关系,被三角形指向的类是父类。菱形代表组合关系。带删除线的成员函数或变量,是被子类覆盖掉,且子类也没有显式的调用父类中的成员函数或变量,相当于就是没有用的函数或变量。

3,把下图从中间分开的话,左边是Llama相关的类,右边是Llava相关类。llama可以简单理解为纯文本的大模型。llava是基于llama,添加了视觉输入的多模态大模型。类名中有Causal可以理解为有输入输出业务逻辑处理的类。

4,Senna与Llava是啥关系呢? Senna是在Llava基础之上的工作,是基于Llava的代码二次开发,网络结构没有变化,所以可以理解Senna==Llava,所以可以看到Senna中一方面把Llava的一些成员覆盖了(虽然网络结构没有变化,但业务意义不一样),另一方面 ,也复用了LlavaMetaModel类。

5,把上面的整体框架理清楚后,再看代码就比较清晰了。简单的说下各成员的用途:

  • vision_tower:就是一个图片特征提取器,直接使用了openai/clip-vit-large-patch14-336,训练过程中是冻结状态。
  • mm_projector和img_adapter:组合在一起相当于把上面的图片特征,转换成固定的维度(6, 128,4096),6是图片数量,128是每个图片转换成128个token,4096是每个token的维度,跟文字的维度是一样的。
  • encode_images,mm_project,prepare_inputs_labels_for_multimodal等函数就是调用上面的一些模块对输入进行转换处理,最终送到llama大模型中。

在这里插入图片描述

6,在SennaLlavaLlamaForCausalLM.from_pretrained()执行完成后,把SennaLlavaLlamaForCausalLM.state_dict().keys()打印出来看一下,如下,可以看到主要分为几部分:

  • model.xxx,其中model对应上面类图中的最上面的model成员变量。
  • model.layers.xxx,就是llama模型的32层layers,具体对应上面类图中的LlamaModel的成员(我没有在类图中画出来,因为它属于transformers库,类图中主要展现项目中自已开发的类)。
  • model.mm_projector和model.img_adapter,具体对应上面类图中的LlavaMetaModel中的mm_projector和img_adapter。

可以看到,SennaLlavaLlamaForCausalLM中的weights权重,通过from_pretrained加载后,layers+mm_projector+img_adapter是在一起的。

dict_keys(['model.embed_tokens.weight'
 'model.layers.0.self_attn.q_proj.weight'
 'model.layers.0.self_attn.k_proj.weight'
 'model.layers.0.self_attn.v_proj.weight'
 'model.layers.0.self_attn.o_proj.weight'
 'model.layers.0.self_attn.rotary_emb.inv_freq'
 'model.layers.0.mlp.gate_proj.weight'
 'model.layers.0.mlp.up_proj.weight'
 'model.layers.0.mlp.down_proj.weight'
 'model.layers.0.input_layernorm.weight'
 'model.layers.0.post_attention_layernorm.weight'
 'model.layers.1.self_attn.q_proj.weight'
 ......
 'model.layers.31.mlp.down_proj.weight'
 'model.layers.31.input_layernorm.weight'
 'model.layers.31.post_attention_layernorm.weight'
 'model.norm.weight'
 'model.mm_projector.0.weight'
 'model.mm_projector.0.bias'
 'model.mm_projector.2.weight'
 'model.mm_projector.2.bias'
 'model.img_adapter.adapter.0.weight'
 'model.img_adapter.adapter.0.bias'
 'model.img_adapter.adapter.2.weight'
 'model.img_adapter.adapter.2.bias'
 'model.img_adapter.adapter.4.weight'
 'model.img_adapter.adapter.4.bias'
 'lm_head.weight'])

7 ,在vision_tower.load_model()执行完成后,再打印一下SennaLlavaLlamaForCausalLM.state_dict().keys()后,可以发现多了vision towner的一些key,如下:

'model.vision_tower.vision_tower.vision_model.embeddings.class_embedding'

 'model.vision_tower.vision_tower.vision_model.embeddings.patch_embedding.weight'

 'model.vision_tower.......'

因为vision tower是直接使用的第三方的库openai/clip-vit-large-patch14-336,在训练过程中也是冻结状态,所以作者并没有将vision tower和主模型的weight放在一起,而是单独加载。这些逻辑在模型导出的时候应该会有所体现。

二,推理代码简介:

回到senna_plan_cmd_eval_multi_img.py文件的代码执行流程上,核心就2部分:

1,模型加载: load_senna_pretrained_model()函数。加载了四个组件:tokenizer, model, image_processor, vision_tower。image_processor其实是vision_tower的一部分,就是图片在提取特征之前需要进行一些预处理,例如pad, normalize等。具体是通过openai/clip-vit-large-patch14-336中有一个preprocessor_config.json配置文件定义的。

2,前向推理: eval_multi_img_model_wo_init()函数。分为load_images,process_images,tokenizer_image_token,SennaLlavaLlamaForCausalLM::generate几个主要部分。process_images就是调用了上面的image_processor进行前处理。tokenizer_image_token就是把文本转换成token id。generate就是把文本和处理后的图片送到SennaLlavaLlamaForCausalLM::prepare_inputs_labels_for_multimodal()进行整合之后,最终送入llama大语言模型进行推理。

思考: 感觉hugging face的组件化做的不是特别好?如果使用过mmdetection框架的同学应该能感受到框架的每一个部分都是组件,例如backbone,head,loss等等,训练的核心就是一个配置文件,通过配置文件把各部分扁平的组合在一起。而在上面senna的代码中的继承关系太过于复杂,需要把这么多类之间的相互关系理清楚才能着手进行开发。个人感受,欢迎大家一起讨论。

Logo

北京人形旗下天工造物具身智能开源社区,聚焦具身天工与慧思开物两大平台

更多推荐