LLaVA模型multi_modal_projector dtype在forward阶段意外变为float32报错
解决LLaVA模型multi_modal_projector前向传播时dtype意外从float16变为float32的问题
问题描述
使用修改后的llama-8b-llava模型,加载时指定torch.float16数据类型:
model = AttentionCaptureModel.from_pretrained( "xtuner/llava-llama-3-8b-v1_1-transformers", torch_dtype=torch.float16, output_attentions=True ).to(device)
但前向传播时触发如下错误:
RuntimeError: mat1 and mat2 must have the same dtype, but got Half and Float.
经排查:初始化阶段(包括__init__和post_init后)打印multi_modal_projector参数均为torch.float16,但前向传播时该模块参数意外变为float32。
原因分析
结合配置与代码细节,核心诱因有两点:
- 视觉塔输出特征的dtype影响:配置文件中
vision_config的torch_dtype设为float32,自定义视觉塔输出的图像特征为float32类型。虽PyTorch不会主动转换模型参数dtype,但部分隐式逻辑(如框架内部参数对齐、混合精度自动处理)可能强制将projector参数转为float32以匹配输入特征。 - 初始化后隐式参数修改:模型加载完成到前向传播的过程中,可能存在未被注意到的代码逻辑(如工具函数、第三方库操作)篡改了projector的dtype。
解决方案
方案1:强制锁定multi_modal_projector的dtype
在模型初始化末尾添加显式转换逻辑,固定projector为float16:
def __init__(self, config: CustomLlavaConfig): super().__init__(config) # 原有初始化代码 self.vision_tower = CustomCLIPVisionModel(config.vision_config) self.multi_modal_projector = LlavaMultiModalProjector(config) self.vocab_size = config.text_config.vocab_size self.language_model = AutoModelForCausalLM.from_config( config.text_config, attn_implementation=config._attn_implementation ) self.pad_token_id = self.config.pad_token_id if self.config.pad_token_id is not None else -1 self.post_init() # 新增:强制将projector转为float16 self.multi_modal_projector = self.multi_modal_projector.to(torch.float16) print("Multi-modal projector parameters dtype after forced conversion:") for name, param in self.multi_modal_projector.named_parameters(): print(f"{name}: {param.dtype}")
或在前向传播函数开头添加强制转换:
def forward(...): # 原有代码逻辑 # 新增:强制将projector转回float16 self.multi_modal_projector.to(torch.float16) print("Multi-modal projector parameters dtype:") for name, param in self.multi_modal_projector.named_parameters(): print(f"{name}: {param.dtype}") image_features = self.multi_modal_projector(selected_image_feature)
方案2:统一视觉特征与projector的dtype
修改配置文件中vision_config的torch_dtype为float16,从源头对齐dtype:
"vision_config": { "architectures": [ "CLIPVisionModel" ], "dropout": 0.0, "hidden_size": 1024, "image_size": 336, "intermediate_size": 4096, "model_type": "clip_vision_model", "num_attention_heads": 16, "num_hidden_layers": 24, "patch_size": 14, "projection_dim": 768, "torch_dtype": "float16" }
注意:修改前需确认自定义CLIP视觉模型支持float16精度,避免数值不稳定。
方案3:临时转换输入特征dtype
在前向传播时将视觉特征转为float16后再传入projector,作为临时 workaround:
# 调用projector前添加转换逻辑 selected_image_feature = selected_image_feature.to(torch.float16) image_features = self.multi_modal_projector(selected_image_feature)
内容的提问来源于stack exchange,提问作者ILOT
相关产品推荐
相关产品推荐

