1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
| from fastapi import FastAPI, UploadFile, File
from pydantic import BaseModel
from typing import Optional
app = FastAPI()
class MultimodalChatBot:
"""多模态聊天机器人"""
def __init__(self):
self.vlm_model = LlavaForConditionalGeneration.from_pretrained(
"llava-hf/llava-1.5-7b-hf",
torch_dtype=torch.float16,
device_map="auto"
)
self.processor = AutoProcessor.from_pretrained(
"llava-hf/llava-1.5-7b-hf"
)
# 对话历史
self.conversation_history = {}
async def chat(
self,
user_id: str,
message: str,
image: Optional[UploadFile] = None
) -> str:
"""多模态对话"""
# 获取历史
history = self.conversation_history.get(user_id, [])
# 准备输入
if image:
# 有图像
image_bytes = await image.read()
image_pil = Image.open(io.BytesIO(image_bytes)).convert("RGB")
prompt = self._build_prompt_with_image(history, message)
inputs = self.processor(
text=prompt,
images=image_pil,
return_tensors="pt"
).to(self.vlm_model.device)
# 生成
with torch.no_grad():
outputs = self.vlm_model.generate(
**inputs,
max_new_tokens=500,
do_sample=True,
temperature=0.7,
)
response = self.processor.decode(outputs[0], skip_special_tokens=True)
else:
# 纯文本
prompt = self._build_prompt(history, message)
response = self.text_llm.generate(prompt)
# 更新历史
history.append({"role": "user", "content": message})
history.append({"role": "assistant", "content": response})
self.conversation_history[user_id] = history[-10:] # 保留最近10轮
return response
def _build_prompt_with_image(self, history, message):
prompt = "USER: <image>\n"
for h in history:
prompt += f"{h['role'].upper()}: {h['content']}\n"
prompt += f"USER: {message}\nASSISTANT:"
return prompt
def _build_prompt(self, history, message):
prompt = ""
for h in history:
prompt += f"{h['role'].upper()}: {h['content']}\n"
prompt += f"USER: {message}\nASSISTANT:"
return prompt
chatbot = MultimodalChatBot()
@app.post("/chat/{user_id}")
async def chat_endpoint(
user_id: str,
message: str = Form(...),
image: UploadFile = File(None)
):
response = await chatbot.chat(user_id, message, image)
return {"response": response}
|