Files
llm/app/main.py
2025-08-03 08:30:12 +00:00

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