【VLM】——vlm计算ppl损失

计算vlm模型的ppl损失。

代码:

python 复制代码
from transformers import Qwen2VLForConditionalGeneration, AutoProcessor
import torch
from torch.nn import CrossEntropyLoss
from PIL import Image


# 配置
DEVICE = "cuda:0"
MODEL_NAME = "/data1/chenjun/huf/Qwen2-VL-2B-Instruct"
IMAGE_SIZE = 384


def resize_image(path, max_side=384):
    """调整图片大小,保持宽高比"""
    image = Image.open(path).convert("RGB")
    width, height = image.size
    if width > height:
        new_width = max_side
        new_height = int(height * (max_side / width))
    else:
        new_height = max_side
        new_width = int(width * (max_side / height))
    return [image.resize((new_width, new_height), Image.Resampling.LANCZOS)]


def main():
    # 加载模型和处理器
    model = Qwen2VLForConditionalGeneration.from_pretrained(
        MODEL_NAME, dtype=torch.float32, device_map=DEVICE
    )
    processor = AutoProcessor.from_pretrained(MODEL_NAME)

    # 构建消息
    file = 'outputs/ppl_vlm_qwen3-vl-2b-axera-384/vit/0000.png'
    messages = [
        {
            "role": "user",
            "content": [
                {"type": "image", "image": file},
                {"type": "text", "text": "描述这张图片"},
            ],
        }
    ]

    # 应用chat模板
    text = processor.apply_chat_template(
        messages, tokenize=False, add_generation_prompt=True
    )

    # 处理图片
    image_inputs = resize_image(file, IMAGE_SIZE)
    inputs = processor(text=[text], images=image_inputs, return_tensors="pt").to(DEVICE)
    gen_idx = inputs['input_ids'].shape[1]

    # 生成文本
    generated_ids = model.generate(**inputs, max_new_tokens=256)
    generated_ids_trimmed = [
        out_ids[len(in_ids):]
        for in_ids, out_ids in zip(inputs.input_ids, generated_ids)
    ]
    output_text = processor.batch_decode(
        generated_ids_trimmed, skip_special_tokens=True, clean_up_tokenization_spaces=False
    )[0]

    # 计算PPL
    text_with_response = text + output_text
    image_inputs = resize_image(file, IMAGE_SIZE)
    inputs2 = processor(text=[text_with_response], images=image_inputs, return_tensors="pt").to(DEVICE)


    with torch.no_grad():
        outputs = model(**inputs2, max_new_tokens=1)
        logits = outputs.logits

        # 计算交叉熵损失
        shift_labels = inputs2['input_ids'][..., gen_idx+1:].contiguous().to(DEVICE)
        shift_logits = logits[..., gen_idx:-1, :].contiguous().to(dtype=torch.float32)
        loss_fct = CrossEntropyLoss()
        ce_loss = loss_fct(
            shift_logits.view(-1, shift_logits.size(-1)),
            shift_labels.view(-1)
        )
        print(f"ce_loss: {ce_loss:.3f}, ppl: {ce_loss.exp():.3f}")


if __name__ == "__main__":
    main()
相关推荐
编程一生4 分钟前
DevOps从持续开发到持续部署
人工智能
byte轻骑兵6 分钟前
【LE Audio】PBP精讲[5]: 广播音频的身份标识与元数据交互法则
人工智能·音视频·le audio·低功耗蓝牙音频
智慧物业老杨2 小时前
物业日常巡查的数智化重构:从“打卡式巡检“到“闭环式风控“
android·java·人工智能·系统架构·rxjava
吴佳浩5 小时前
Skill 为什么不同于 Tool?Agent 技能库的自演进与动态加载机制
人工智能·agent·ai编程
大模型任我行8 小时前
谷歌:“课程学习”融入扩散模型强化学习
人工智能·语言模型·自然语言处理·论文笔记
AIGCmagic社区8 小时前
具身智能专题:机器人也有Scaling Law?智元GE-Act 2.0用3万小时真机数据给出答案
人工智能·aigc·具身智能
Rosanci9 小时前
谷歌浏览器插件开发实战指南:从 Hello World 到上架发布
大数据·人工智能·chrome·程序人生
明月_清风9 小时前
AI 越来越强,程序员真正的价值到底是什么?
人工智能·后端
m0_4665252910 小时前
云从科技上线云起ModelHub:AI团队时代的模型算力基础设施
大数据·人工智能·科技
火山引擎开发者社区10 小时前
OpenViking:给 Codex 加上长期记忆
人工智能