Files
llm/app/main.py
2025-08-04 21:42:50 +07:00

185 lines
9.0 KiB
Python
Raw Permalink Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
from typing import Any, Dict, List, Optional, AsyncGenerator
import gradio as gr
import uuid
import torch
import asyncio
from transformers import AutoTokenizer as HFTokenizer
from loguru import logger
from langchain_chroma import Chroma
from langchain_huggingface import HuggingFaceEmbeddings
from vllm import AsyncLLMEngine, SamplingParams
from vllm.engine.arg_utils import AsyncEngineArgs
from config import settings
class ChatWithAI:
def __init__(self, provider: str = "qwen3"):
self.provider = provider
self.embeddings = HuggingFaceEmbeddings(
model_name=settings.LM_MODEL_NAME,
model_kwargs={"device": "cuda"},
encode_kwargs={"normalize_embeddings": True},
)
if provider == "qwen3":
model_name = getattr(settings, "LOCAL_LLM_NAME", "/models/Qwen3-8B")
logger.info(f"Инициализация vLLM движка с моделью: {model_name}")
# Аргументы для асинхронного движка vLLM
engine_args = AsyncEngineArgs(
model=model_name,
tokenizer=model_name,
tokenizer_mode="auto",
trust_remote_code=True,
dtype=torch.bfloat16, # float16
gpu_memory_utilization=0.65,
max_model_len=32768,
enable_prefix_caching=True,
tensor_parallel_size=1, # Увеличь, если у тебя несколько GPU
enable_chunked_prefill=False,
)
self.engine = AsyncLLMEngine.from_engine_args(engine_args)
self.tokenizer = HFTokenizer.from_pretrained(settings.LM_MODEL_NAME, trust_remote_code=True)
else:
raise ValueError(f"Неподдерживаемый провайдер: {provider}")
self.chroma_db = Chroma(
persist_directory=settings.DOCS_CHROMA_PATH,
embedding_function=self.embeddings,
collection_name=settings.DOCS_COLLECTION_NAME,
)
def get_relevant_context(self, query: str, k: int = 5) -> List[Dict[str, Any]]:
"""Получение релевантного контекста из базы данных."""
try:
results = self.chroma_db.similarity_search(query, k=k)
return [
{
"text": doc.page_content,
"metadata": doc.metadata,
}
for doc in results
]
except Exception as e:
logger.error(f"Ошибка при получении контекста: {e}")
return []
def format_context(self, context: List[Dict[str, Any]]) -> str:
"""Форматирование контекста для промпта."""
formatted_context = []
for item in context:
metadata_str = "\n".join(f"{k}: {v}" for k, v in item["metadata"].items())
formatted_context.append(
f"Текст: {item['text']}\nМетаданные:\n{metadata_str}\n"
)
return "\n---\n".join(formatted_context)
async def generate_response_stream(self, query: str) -> AsyncGenerator[str, None]:
"""Генерация ответа с потоковой передачей токенов через vLLM."""
try:
logger.info(f"Пользовательский запрос: {query}")
context = self.get_relevant_context(query)
if not context:
yield "Извините, не удалось найти релевантный контекст для ответа."
return
formatted_context = self.format_context(context)
messages = [
{
"role": "system",
"content": """Ты — внутренний менеджер помощи пользоваткля по вопросам настройки сервера. Отвечаешь по делу без лишних вступлений.
Правила:
1. Сразу переходи к сути, без фраз типа "На основе контекста"
2. Используй только факты. Если точных данных нет — отвечай общими фразами об настройки сервера, но не придумывай конкретику
3. Используй текст с форматированием.
4. Включай ссылки только если они есть в контексте
5. Говори от первого лица множественного числа: "Мы предоставляем", "У нас есть"
6. При упоминании файлов делай это естественно, например: "Я прикреплю инструкцию, где подробно описаны шаги"
7. На приветствия отвечай доброжелательно, на негатив — с легким юмором
8. Можешь при ответах использовать общую информацию из открытых источников по настройке сервера, но опирайся на контекст
9. Если пользователь спрашивает о ценах, планах или технических характеристиках — давай конкретные ответы из контекста
10. При технических вопросах предлагай практические решения
Персонализируй ответы, упоминая имя клиента если оно есть в контексте. Будь краток, информативен и полезен.""",
},
{
"role": "user",
"content": f"Вопрос: {query}\nКонтекст: {formatted_context}",
},
]
# Применяем шаблон чата
prompt = self.tokenizer.apply_chat_template(
messages,
tokenize=False,
add_generation_prompt=True,
enable_thinking=False
)
# Параметры генерации
sampling_params = SamplingParams(
temperature=0.3,
top_p=0.9,
max_tokens=2048,
stop_token_ids=[], # можно добавить, если нужно
)
# Генерация через vLLM
final_output = ""
buffer = ""
# Генерируем асинхронно
request_id = str(uuid.uuid4())
try:
async for output in self.engine.generate(prompt, sampling_params, request_id):
if output.outputs:
text = output.outputs[0].text
# Отправляем только новые токены
new_text = text[len(final_output):]
final_output = text
if new_text:
buffer += new_text
yield buffer # Постепенно возвращаем накопленный текст
except Exception as e:
logger.error(f"Ошибка при генерации vLLM: {e}")
yield "Произошла ошибка при генерации ответа."
except Exception as e:
logger.error(f"Неожиданная ошибка: {e}")
yield "Произошла ошибка при генерации ответа."
# === Gradio интерфейс ===
def main():
chat = ChatWithAI(provider="qwen3")
async def respond(message, history):
# Генерируем ответ по частям
try:
async for token in chat.generate_response_stream(message):
yield token
except Exception as e:
logger.error(f"Ошибка в respond: {e}")
yield "Извините, произошла ошибка при обработке запроса."
demo = gr.ChatInterface(
fn=respond,
title="Помощник настройки сервера и подбора железа",
description="Задайте вопрос — получите ответ от внутреннего менеджера.",
examples=[
"Как определить цели и требования для домашнего сервера?",
"Как выбрать ОС для домашнего сервера?",
"Проверка совместимости и выбор процессора",
"Установка Unraid и первоначальная настройка"
],
)
demo.queue(max_size=20, default_concurrency_limit=10).launch(
server_name="0.0.0.0", server_port=8080, share=False
)
if __name__ == "__main__":
main()