Compare commits
2 Commits
main
...
change-mod
| Author | SHA1 | Date | |
|---|---|---|---|
| a868025a9e | |||
| 9e86008ba0 |
@@ -37,4 +37,4 @@ RUN pip install --no-cache-dir -r /app/requirements.txt
|
|||||||
COPY app /app
|
COPY app /app
|
||||||
|
|
||||||
# Запуск приложения
|
# Запуск приложения
|
||||||
CMD ["python", "main.py"]
|
#CMD ["python", "main.py"]
|
||||||
147
app/main.py
147
app/main.py
@@ -12,39 +12,45 @@ from config import settings
|
|||||||
|
|
||||||
|
|
||||||
class ChatWithAI:
|
class ChatWithAI:
|
||||||
def __init__(self, provider: str = "qwen3"):
|
def __init__(self, model_name: str):
|
||||||
self.provider = provider
|
self.model_name = model_name
|
||||||
self.embeddings = HuggingFaceEmbeddings(
|
self.embeddings = HuggingFaceEmbeddings(
|
||||||
model_name=settings.LM_MODEL_NAME,
|
model_name=settings.LM_MODEL_NAME,
|
||||||
model_kwargs={"device": "cuda"},
|
model_kwargs={"device": "cuda"},
|
||||||
encode_kwargs={"normalize_embeddings": True},
|
encode_kwargs={"normalize_embeddings": True},
|
||||||
)
|
)
|
||||||
|
|
||||||
if provider == "qwen3":
|
self.chroma_db = Chroma(
|
||||||
model_name = getattr(settings, "LOCAL_LLM_NAME", "/models/Qwen3-4B")
|
persist_directory=settings.DOCS_CHROMA_PATH,
|
||||||
logger.info(f"Загрузка локальной модели: {model_name}")
|
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.tokenizer = AutoTokenizer.from_pretrained(model_name)
|
||||||
self.model = AutoModelForCausalLM.from_pretrained(
|
self.model = AutoModelForCausalLM.from_pretrained(
|
||||||
model_name,
|
model_name,
|
||||||
torch_dtype=torch.float32,
|
torch_dtype=torch.float32,
|
||||||
device_map="cuda",
|
device_map="cuda",
|
||||||
)
|
)
|
||||||
|
logger.info(f"Модель {model_name} успешно загружена.")
|
||||||
# Streamer для потоковой генерации
|
except Exception as e:
|
||||||
self.streamer = TextIteratorStreamer(
|
logger.error(f"Ошибка при загрузке модели {model_name}: {e}")
|
||||||
self.tokenizer,
|
raise
|
||||||
skip_prompt=True,
|
|
||||||
skip_special_tokens=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 = 3) -> List[Dict[str, Any]]:
|
def get_relevant_context(self, query: str, k: int = 3) -> List[Dict[str, Any]]:
|
||||||
"""Получение релевантного контекста из базы данных."""
|
"""Получение релевантного контекста из базы данных."""
|
||||||
@@ -74,7 +80,6 @@ class ChatWithAI:
|
|||||||
def generate_response_stream(self, query: str):
|
def generate_response_stream(self, query: str):
|
||||||
"""Генерация ответа с потоковой передачей токенов."""
|
"""Генерация ответа с потоковой передачей токенов."""
|
||||||
try:
|
try:
|
||||||
logger.info(f"Пользовательский запрос: {query}")
|
|
||||||
context = self.get_relevant_context(query)
|
context = self.get_relevant_context(query)
|
||||||
if not context:
|
if not context:
|
||||||
yield "Извините, не удалось найти релевантный контекст для ответа."
|
yield "Извините, не удалось найти релевантный контекст для ответа."
|
||||||
@@ -85,7 +90,7 @@ class ChatWithAI:
|
|||||||
messages = [
|
messages = [
|
||||||
{
|
{
|
||||||
"role": "system",
|
"role": "system",
|
||||||
"content": """Ты — внутренний менеджер компании Mazban. Отвечаешь по делу без лишних вступлений.
|
"content": """Ты — внутренний менеджер компании Amvera Cloud. Отвечаешь по делу без лишних вступлений.
|
||||||
|
|
||||||
Правила:
|
Правила:
|
||||||
1. Сразу переходи к сути, без фраз типа "На основе контекста"
|
1. Сразу переходи к сути, без фраз типа "На основе контекста"
|
||||||
@@ -118,13 +123,14 @@ class ChatWithAI:
|
|||||||
# Подготавливаем вход
|
# Подготавливаем вход
|
||||||
inputs = self.tokenizer(prompt, return_tensors="pt").to("cuda")
|
inputs = self.tokenizer(prompt, return_tensors="pt").to("cuda")
|
||||||
|
|
||||||
# Очищаем streamer и запускаем генерацию в отдельном потоке
|
# Создаём новый streamer
|
||||||
self.streamer = TextIteratorStreamer(
|
self.streamer = TextIteratorStreamer(
|
||||||
self.tokenizer,
|
self.tokenizer,
|
||||||
skip_prompt=True,
|
skip_prompt=True,
|
||||||
skip_special_tokens=True
|
skip_special_tokens=True
|
||||||
)
|
)
|
||||||
|
|
||||||
|
# Параметры генерации
|
||||||
generate_kwargs = {
|
generate_kwargs = {
|
||||||
"input_ids": inputs["input_ids"],
|
"input_ids": inputs["input_ids"],
|
||||||
"max_new_tokens": 512,
|
"max_new_tokens": 512,
|
||||||
@@ -135,6 +141,7 @@ class ChatWithAI:
|
|||||||
"streamer": self.streamer,
|
"streamer": self.streamer,
|
||||||
}
|
}
|
||||||
|
|
||||||
|
# Запускаем генерацию в отдельном потоке
|
||||||
thread = threading.Thread(target=self.model.generate, kwargs=generate_kwargs)
|
thread = threading.Thread(target=self.model.generate, kwargs=generate_kwargs)
|
||||||
thread.start()
|
thread.start()
|
||||||
|
|
||||||
@@ -142,32 +149,100 @@ class ChatWithAI:
|
|||||||
buffer = ""
|
buffer = ""
|
||||||
for token in self.streamer:
|
for token in self.streamer:
|
||||||
buffer += token
|
buffer += token
|
||||||
yield buffer # Отправляем частичный ответ
|
yield buffer
|
||||||
|
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
logger.error(f"Ошибка при генерации ответа: {e}")
|
logger.error(f"Ошибка при генерации ответа: {e}")
|
||||||
yield "Произошла ошибка при генерации ответа."
|
yield "Произошла ошибка при генерации ответа."
|
||||||
|
|
||||||
|
|
||||||
# === Gradio интерфейс ===
|
# === Gradio интерфейс с выбором модели ===
|
||||||
def main():
|
def main():
|
||||||
chat = ChatWithAI(provider="qwen3")
|
# Список доступных моделей (можно вынести в 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):
|
def respond(message, history):
|
||||||
# Генерируем ответ по частям
|
# Генерируем ответ по частям
|
||||||
for token in chat.generate_response_stream(message):
|
for token in chat.generate_response_stream(message):
|
||||||
yield token
|
yield token
|
||||||
|
|
||||||
demo = gr.ChatInterface(
|
with gr.Blocks(title="Amvera Cloud Assistant") as demo:
|
||||||
fn=respond,
|
gr.Markdown("# 🤖 Amvera Cloud Assistant")
|
||||||
title="Помощник настройки Marzban",
|
gr.Markdown("Выберите модель и задайте вопрос — получите ответ с контекстом.")
|
||||||
description="Задайте вопрос — получите ответ от внутреннего менеджера.",
|
|
||||||
examples=[
|
with gr.Row():
|
||||||
"Как подключить Marzban?",
|
with gr.Column(scale=1):
|
||||||
"Как настроить telegram бота?",
|
model_dropdown = gr.Dropdown(
|
||||||
"Что такое marzban?"
|
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)
|
demo.launch(server_name="0.0.0.0", server_port=8080, share=False)
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@@ -20,7 +20,7 @@ services:
|
|||||||
- ./cache/cache:/root/.cache
|
- ./cache/cache:/root/.cache
|
||||||
- ./cache/site-packages:/usr/local/lib/python3.10/site-packages
|
- ./cache/site-packages:/usr/local/lib/python3.10/site-packages
|
||||||
- ./offline_packages:/offline_packages
|
- ./offline_packages:/offline_packages
|
||||||
# entrypoint: sleep 1000000 #./app/entrypoint.sh
|
entrypoint: sleep 1000000 #./app/entrypoint.sh
|
||||||
ports:
|
ports:
|
||||||
- "8080:8080"
|
- "8080:8080"
|
||||||
networks:
|
networks:
|
||||||
|
|||||||
Reference in New Issue
Block a user