Views
No views yet

注: 元の LLM の性能を維持するため、Stage‑2(Projector と LLM の両方を学習可能にする段階)をスキップ し、Stage‑2 で使用予定だったデータを Stage‑1 に組み込みました。
transformers >= 4.35.3 をインストールしてください。USER: xxxASSISTANT:)を守り、画像を問い合わせる位置に <image> トークンを挿入してください。bfloat16 精度で生成を行うサンプルスクリプトです。1import requests
2from PIL import Image
3
4import torch
5from transformers import AutoProcessor, LlavaForConditionalGeneration
6
7model_id = "turing-motors/llava-1.5-sarashina2.2-1.7b-instruct"
8model = LlavaForConditionalGeneration.from_pretrained(
9 model_id,
10 torch_dtype=torch.bfloat16,
11 low_cpu_mem_usage=True,
12).to("cuda")
13
14processor = AutoProcessor.from_pretrained(model_id, use_fast=True)
15
16# チャット履歴を定義し、apply_chat_template でフォーマット済みプロンプトを作成
17# "content" 内の各値は ("text", "image") 型の dict のリスト
18conversation = [
19 {
20 "role": "user",
21 "content": [
22 {"type": "text", "text": "猫は何匹いますか?"},
23 {"type": "image"},
24 ],
25 },
26]
27prompt = processor.apply_chat_template(conversation, add_generation_prompt=True)
28
29image_file = "http://images.cocodataset.org/val2017/000000039769.jpg"
30raw_image = Image.open(requests.get(image_file, stream=True).raw)
31inputs = processor(images=raw_image, text=prompt, return_tensors='pt').to("cuda")
32
33generated_ids = model.generate(**inputs, max_new_tokens=128, do_sample=False)
34generated_texts = processor.batch_decode(
35 generated_ids,
36 skip_special_tokens=True,
37)
38print(generated_texts[0])
39# USER:
40# 猫は何匹いますか?ASSISTANT: 画像には2匹の猫がいます。transformers v4.48 以降では、画像 URL またはローカルパスを会話履歴に直接渡し、チャットテンプレートに処理を任せることも可能です。
テンプレートが画像を読み込み、torch.Tensor 形式で返すので、そのまま model.generate() に渡せます。1messages = [
2 {
3 "role": "user",
4 "content": [
5 {"type": "image", "url": "https://www.ilankelman.org/stopsigns/australia.jpg"},
6 {"type": "text", "text": "画像を非常に短く説明して。"},
7 ],
8 },
9]
10
11inputs = processor.apply_chat_template(
12 messages,
13 add_generation_prompt=True,
14 tokenize=True,
15 return_dict=True,
16 return_tensors="pt"
17).to("cuda")
18
19output = model.generate(**inputs, max_new_tokens=128)
20generated_texts = processor.batch_decode(
21 output,
22 skip_special_tokens=True,
23)
24print(generated_texts[0])
25# USER:
26# 画像を非常に短く説明して。ASSISTANT: 画像は、赤い停止標識と、その横にある赤い門を持つ伝統的な中国の門の2つの標識が写っています。