diff --git a/Dockerfile b/Dockerfile index a5edd0a..843d8c2 100644 --- a/Dockerfile +++ b/Dockerfile @@ -15,6 +15,10 @@ RUN apt-get update && \ wget \ libgl1 \ libglib2.0-0 \ + build-essential \ + python3.10-dev \ + python3.10-distutils \ + libpython3-dev gcc g++ make \ && rm -rf /var/lib/apt/lists/* # Установка Python 3.10 как основного diff --git a/app/config.py b/app/config.py index 4be0dea..4183633 100644 --- a/app/config.py +++ b/app/config.py @@ -9,9 +9,9 @@ class Config(BaseSettings): PARSED_JSON_PATH: str = os.path.join(BASE_DIR, "data", "parsed_json") DOCS_CHROMA_PATH: str = os.path.join(BASE_DIR, "data", "chroma_db") DOCS_COLLECTION_NAME: str = "docs" - MAX_CHUNK_SIZE: int = 512 + MAX_CHUNK_SIZE: int = 2048 CHUNK_OVERLAP: int = 50 - LM_MODEL_NAME: str = "/models/paraphrase-multilingual-MiniLM-L12-v2" + LM_MODEL_NAME: str = "/models/Qwen3-Embedding-0.6B" LOCAL_LLM_NAME: str = "/models/Qwen3-8B" QWEN_MODEL_NAME: str = "Qwen3-8B" diff --git a/app/main.py b/app/main.py index 4a9ce5f..4d23a6e 100644 --- a/app/main.py +++ b/app/main.py @@ -1,13 +1,13 @@ -from typing import Any, Dict, List, Optional +from typing import Any, Dict, List, Optional, AsyncGenerator import gradio as gr import torch +import asyncio import threading from loguru import logger from langchain_chroma import Chroma from langchain_huggingface import HuggingFaceEmbeddings -from transformers import AutoModelForCausalLM, AutoTokenizer, TextIteratorStreamer - -# Предполагается, что у тебя есть config.py с settings +from vllm import AsyncLLMEngine, SamplingParams +from vllm.engine.arg_utils import AsyncEngineArgs from config import settings @@ -22,21 +22,22 @@ class ChatWithAI: if provider == "qwen3": model_name = getattr(settings, "LOCAL_LLM_NAME", "/models/Qwen3-8B") - logger.info(f"Загрузка локальной модели: {model_name}") + logger.info(f"Инициализация vLLM движка с моделью: {model_name}") - self.tokenizer = AutoTokenizer.from_pretrained(model_name) - self.model = AutoModelForCausalLM.from_pretrained( - model_name, - torch_dtype=torch.float16, - device_map="cuda", - ) - - # Streamer для потоковой генерации - self.streamer = TextIteratorStreamer( - self.tokenizer, - skip_prompt=True, - skip_special_tokens=True + # Аргументы для асинхронного движка vLLM + engine_args = AsyncEngineArgs( + model=model_name, + tokenizer=model_name, + tokenizer_mode="auto", + trust_remote_code=True, + dtype=torch.float16, # float16 + gpu_memory_utilization=0.85, + max_model_len=16384, + enable_prefix_caching=True, + tensor_parallel_size=1, # Увеличь, если у тебя несколько GPU ) + self.engine = AsyncLLMEngine.from_engine_args(engine_args) + self.tokenizer = self.engine.engine.tokenizer.tokenizer # доступ к токенизатору else: raise ValueError(f"Неподдерживаемый провайдер: {provider}") @@ -46,7 +47,7 @@ class ChatWithAI: collection_name=settings.DOCS_COLLECTION_NAME, ) - def get_relevant_context(self, query: str, k: int = 3) -> List[Dict[str, Any]]: + def get_relevant_context(self, query: str, k: int = 5) -> List[Dict[str, Any]]: """Получение релевантного контекста из базы данных.""" try: results = self.chroma_db.similarity_search(query, k=k) @@ -71,8 +72,8 @@ class ChatWithAI: ) return "\n---\n".join(formatted_context) - def generate_response_stream(self, query: str): - """Генерация ответа с потоковой передачей токенов.""" + async def generate_response_stream(self, query: str) -> AsyncGenerator[str, None]: + """Генерация ответа с потоковой передачей токенов через vLLM.""" try: logger.info(f"Пользовательский запрос: {query}") context = self.get_relevant_context(query) @@ -115,37 +116,36 @@ class ChatWithAI: enable_thinking=False ) - # Подготавливаем вход - inputs = self.tokenizer(prompt, return_tensors="pt").to("cuda") - - # Очищаем streamer и запускаем генерацию в отдельном потоке - self.streamer = TextIteratorStreamer( - self.tokenizer, - skip_prompt=True, - skip_special_tokens=True + # Параметры генерации + sampling_params = SamplingParams( + temperature=0.3, + top_p=0.9, + max_tokens=2048, + stop_token_ids=[], # можно добавить, если нужно ) - generate_kwargs = { - "input_ids": inputs["input_ids"], - "max_new_tokens": 2048, - "temperature": 0.3, - "do_sample": True, - "top_p": 0.9, - "pad_token_id": self.tokenizer.eos_token_id, - "streamer": self.streamer, - } - - thread = threading.Thread(target=self.model.generate, kwargs=generate_kwargs) - thread.start() - - # Потоковая передача токенов + # Генерация через vLLM + final_output = "" buffer = "" - for token in self.streamer: - buffer += token - yield buffer # Отправляем частичный ответ + + # Генерируем асинхронно + request_id = f"request-{hash(query)}" + 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}") + logger.error(f"Неожиданная ошибка: {e}") yield "Произошла ошибка при генерации ответа." @@ -153,10 +153,27 @@ class ChatWithAI: def main(): chat = ChatWithAI(provider="qwen3") + # Gradio требует синхронную функцию, но мы можем обернуть асинхронный вызов def respond(message, history): - # Генерируем ответ по частям - for token in chat.generate_response_stream(message): - yield token + # Запускаем асинхронный генератор в синхронном контексте + async def async_generate(): + async for token in chat.generate_response_stream(message): + yield token + + # Используем asyncio.new_event_loop() для запуска внутри потока + loop = asyncio.new_event_loop() + asyncio.set_event_loop(loop) + + try: + gen = async_generate() + while True: + try: + token = loop.run_until_complete(gen.__anext__()) + yield token + except StopAsyncIteration: + break + finally: + loop.close() demo = gr.ChatInterface( fn=respond, @@ -169,7 +186,9 @@ def main(): "Установка Unraid и первоначальная настройка" ], ) - demo.queue(max_size=20, default_concurrency_limit=10).launch(server_name="0.0.0.0", server_port=8080, share=False) + demo.queue(max_size=20, default_concurrency_limit=10).launch( + server_name="0.0.0.0", server_port=8080, share=False + ) if __name__ == "__main__": diff --git a/requirements.txt b/requirements.txt index 9f11a23..4c550a8 100644 --- a/requirements.txt +++ b/requirements.txt @@ -11,4 +11,8 @@ chromadb==0.6.3 sentence-transformers==3.4.1 langchain-chroma==0.2.2 pydantic-settings==2.8.1 -langchain-text-splitters==0.3.7 \ No newline at end of file +langchain-text-splitters==0.3.7 +vllm>=0.4.2 +aiohttp>=3.8.0 +fastapi>=0.100.0 +starlette>=0.27.0 \ No newline at end of file