Compare commits
1 Commits
change-mod
...
fd4eef8e7f
| Author | SHA1 | Date | |
|---|---|---|---|
| fd4eef8e7f |
@@ -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,45 +12,39 @@ from config import settings
|
|||||||
|
|
||||||
|
|
||||||
class ChatWithAI:
|
class ChatWithAI:
|
||||||
def __init__(self, model_name: str):
|
def __init__(self, provider: str = "qwen3"):
|
||||||
self.model_name = model_name
|
self.provider = provider
|
||||||
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},
|
||||||
)
|
)
|
||||||
|
|
||||||
self.chroma_db = Chroma(
|
if provider == "qwen3":
|
||||||
persist_directory=settings.DOCS_CHROMA_PATH,
|
model_name = getattr(settings, "LOCAL_LLM_NAME", "/models/Qwen3-4B")
|
||||||
embedding_function=self.embeddings,
|
logger.info(f"Загрузка локальной модели: {model_name}")
|
||||||
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} успешно загружена.")
|
|
||||||
except Exception as e:
|
# Streamer для потоковой генерации
|
||||||
logger.error(f"Ошибка при загрузке модели {model_name}: {e}")
|
self.streamer = TextIteratorStreamer(
|
||||||
raise
|
self.tokenizer,
|
||||||
|
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]]:
|
||||||
"""Получение релевантного контекста из базы данных."""
|
"""Получение релевантного контекста из базы данных."""
|
||||||
@@ -80,6 +74,7 @@ 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 "Извините, не удалось найти релевантный контекст для ответа."
|
||||||
@@ -90,7 +85,7 @@ class ChatWithAI:
|
|||||||
messages = [
|
messages = [
|
||||||
{
|
{
|
||||||
"role": "system",
|
"role": "system",
|
||||||
"content": """Ты — внутренний менеджер компании Amvera Cloud. Отвечаешь по делу без лишних вступлений.
|
"content": """Ты — внутренний менеджер компании Mazban. Отвечаешь по делу без лишних вступлений.
|
||||||
|
|
||||||
Правила:
|
Правила:
|
||||||
1. Сразу переходи к сути, без фраз типа "На основе контекста"
|
1. Сразу переходи к сути, без фраз типа "На основе контекста"
|
||||||
@@ -123,14 +118,13 @@ 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,
|
||||||
@@ -141,7 +135,6 @@ 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()
|
||||||
|
|
||||||
@@ -149,100 +142,32 @@ 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():
|
||||||
# Список доступных моделей (можно вынести в settings)
|
chat = ChatWithAI(provider="qwen3")
|
||||||
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
|
||||||
|
|
||||||
with gr.Blocks(title="Amvera Cloud Assistant") as demo:
|
demo = gr.ChatInterface(
|
||||||
gr.Markdown("# 🤖 Amvera Cloud Assistant")
|
fn=respond,
|
||||||
gr.Markdown("Выберите модель и задайте вопрос — получите ответ с контекстом.")
|
title="Помощник настройки Marzban",
|
||||||
|
description="Задайте вопрос — получите ответ от внутреннего менеджера.",
|
||||||
with gr.Row():
|
examples=[
|
||||||
with gr.Column(scale=1):
|
"Как подключить Marzban?",
|
||||||
model_dropdown = gr.Dropdown(
|
"Как настроить telegram бота?",
|
||||||
choices=list(MODEL_OPTIONS.keys()),
|
"Что такое marzban?"
|
||||||
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