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()