250 lines
11 KiB
Python
250 lines
11 KiB
Python
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() |