from typing import Any, Dict, List, Optional import gradio as gr import torch 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 config import settings class ChatWithAI: def __init__(self, model_name: str): self.model_name = model_name self.embeddings = HuggingFaceEmbeddings( model_name=settings.LM_MODEL_NAME, model_kwargs={"device": "cuda"}, encode_kwargs={"normalize_embeddings": True}, ) self.chroma_db = Chroma( persist_directory=settings.DOCS_CHROMA_PATH, embedding_function=self.embeddings, collection_name=settings.DOCS_COLLECTION_NAME, ) self.tokenizer = None self.model = None self.streamer = None self.load_model(model_name) def load_model(self, model_name: str): """Загружает модель и токенизатор, освобождает предыдущие ресурсы.""" if hasattr(self, "model") and self.model is not None: del self.model torch.cuda.empty_cache() logger.info("Предыдущая модель удалена.") logger.info(f"Загрузка модели: {model_name}") try: self.tokenizer = AutoTokenizer.from_pretrained(model_name) self.model = AutoModelForCausalLM.from_pretrained( model_name, torch_dtype=torch.float32, device_map="cuda", ) logger.info(f"Модель {model_name} успешно загружена.") except Exception as e: logger.error(f"Ошибка при загрузке модели {model_name}: {e}") raise def get_relevant_context(self, query: str, k: int = 3) -> 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) def generate_response_stream(self, query: str): """Генерация ответа с потоковой передачей токенов.""" try: context = self.get_relevant_context(query) if not context: yield "Извините, не удалось найти релевантный контекст для ответа." return formatted_context = self.format_context(context) messages = [ { "role": "system", "content": """Ты — внутренний менеджер компании Amvera Cloud. Отвечаешь по делу без лишних вступлений. Правила: 1. Сразу переходи к сути, без фраз типа "На основе контекста" 2. Используй только факты. Если точных данных нет — отвечай общими фразами об Marzban, но не придумывай конкретику 3. Используй обычный текст без форматирования 4. Включай ссылки только если они есть в контексте 5. Говори от первого лица множественного числа: "Мы предоставляем", "У нас есть" 6. При упоминании файлов делай это естественно, например: "Я прикреплю инструкцию, где подробно описаны шаги" 7. На приветствия отвечай доброжелательно, на негатив — с легким юмором 8. Можешь при ответах использовать общую информацию из открытых источников по Marzban, но опирайся на контекст 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 ) # Подготавливаем вход inputs = self.tokenizer(prompt, return_tensors="pt").to("cuda") # Создаём новый streamer self.streamer = TextIteratorStreamer( self.tokenizer, skip_prompt=True, skip_special_tokens=True ) # Параметры генерации generate_kwargs = { "input_ids": inputs["input_ids"], "max_new_tokens": 512, "temperature": 0.7, "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() # Потоковая передача токенов buffer = "" for token in self.streamer: buffer += token yield buffer except Exception as e: logger.error(f"Ошибка при генерации ответа: {e}") yield "Произошла ошибка при генерации ответа." # === Gradio интерфейс с выбором модели === def main(): # Список доступных моделей (можно вынести в settings) MODEL_OPTIONS = { "Qwen3-4B (локальная)": getattr(settings, "LOCAL_LLM_NAME", "/models/Qwen3-4B"), "Alibaba/Qwen2.5-1.8B-Chat": "Alibaba/Qwen2.5-1.8B-Chat", "HuggingFaceH4/zephyr-7b-beta": "HuggingFaceH4/zephyr-7b-beta", # Добавь другие модели по желанию } chat = None # Будет инициализирована при выборе модели def initialize_model(selected_model_key): """Создаёт или перезагружает экземпляр ChatWithAI с выбранной моделью.""" nonlocal chat if chat is not None: logger.info(f"Смена модели") # Удаляем экземпляр del chat torch.cuda.empty_cache() model_path = MODEL_OPTIONS[selected_model_key] try: chat = ChatWithAI(model_path) return gr.update(interactive=True, placeholder="Введите свой вопрос..."), \ gr.update(visible=True), \ f"✅ Модель загружена: {selected_model_key}" except Exception as e: logger.error(f"Не удалось загрузить модель {selected_model_key}: {e}") return gr.update(interactive=False), \ gr.update(visible=False), \ f"❌ Ошибка загрузки модели: {str(e)}" def respond(message, history): # Генерируем ответ по частям for token in chat.generate_response_stream(message): yield token with gr.Blocks(title="Amvera Cloud Assistant") as demo: gr.Markdown("# 🤖 Amvera Cloud Assistant") gr.Markdown("Выберите модель и задайте вопрос — получите ответ с контекстом.") with gr.Row(): with gr.Column(scale=1): model_dropdown = gr.Dropdown( choices=list(MODEL_OPTIONS.keys()), value=list(MODEL_OPTIONS.keys())[0], label="Выберите модель", interactive=True ) load_button = gr.Button("Загрузить модель", variant="primary") status_text = gr.Textbox( label="Состояние", value="Выберите модель и нажмите 'Загрузить модель'", interactive=False ) with gr.Column(scale=3): chatbot = gr.Chatbot( label="Чат", height=600, bubble_full_width=False, ) msg = gr.Textbox( label="Сообщение", placeholder="Введите ваш вопрос...", interactive=False ) clear = gr.Button("Очистить чат") # Логика выбора модели load_button.click( initialize_model, inputs=[model_dropdown], outputs=[msg, chatbot, status_text] ) # Отправка сообщения msg.submit( respond, inputs=[msg, chatbot], outputs=[chatbot] ) # Очистка чата clear.click(lambda: None, None, chatbot, queue=False) demo.launch(server_name="0.0.0.0", server_port=8080, share=False) if __name__ == "__main__": main()