You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

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。

原因分析

结合配置与代码细节,核心诱因有两点:

  1. 视觉塔输出特征的dtype影响:配置文件中vision_config的torch_dtype设为float32,自定义视觉塔输出的图像特征为float32类型。虽PyTorch不会主动转换模型参数dtype,但部分隐式逻辑(如框架内部参数对齐、混合精度自动处理)可能强制将projector参数转为float32以匹配输入特征。
  2. 初始化后隐式参数修改:模型加载完成到前向传播的过程中,可能存在未被注意到的代码逻辑(如工具函数、第三方库操作)篡改了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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.06.12 10:43:17