Compare commits

..

1 Commits

3 changed files with 38 additions and 113 deletions

View File

@@ -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"]

View File

@@ -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)

View File

@@ -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: