This commit is contained in:
2025-08-04 02:06:07 +00:00
parent 370a11a463
commit 711ad223a5
4 changed files with 80 additions and 53 deletions

View File

@@ -15,6 +15,10 @@ RUN apt-get update && \
wget \
libgl1 \
libglib2.0-0 \
build-essential \
python3.10-dev \
python3.10-distutils \
libpython3-dev gcc g++ make \
&& rm -rf /var/lib/apt/lists/*
# Установка Python 3.10 как основного

View File

@@ -9,9 +9,9 @@ class Config(BaseSettings):
PARSED_JSON_PATH: str = os.path.join(BASE_DIR, "data", "parsed_json")
DOCS_CHROMA_PATH: str = os.path.join(BASE_DIR, "data", "chroma_db")
DOCS_COLLECTION_NAME: str = "docs"
MAX_CHUNK_SIZE: int = 512
MAX_CHUNK_SIZE: int = 2048
CHUNK_OVERLAP: int = 50
LM_MODEL_NAME: str = "/models/paraphrase-multilingual-MiniLM-L12-v2"
LM_MODEL_NAME: str = "/models/Qwen3-Embedding-0.6B"
LOCAL_LLM_NAME: str = "/models/Qwen3-8B"
QWEN_MODEL_NAME: str = "Qwen3-8B"

View File

@@ -1,13 +1,13 @@
from typing import Any, Dict, List, Optional
from typing import Any, Dict, List, Optional, AsyncGenerator
import gradio as gr
import torch
import asyncio
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 vllm import AsyncLLMEngine, SamplingParams
from vllm.engine.arg_utils import AsyncEngineArgs
from config import settings
@@ -22,21 +22,22 @@ class ChatWithAI:
if provider == "qwen3":
model_name = getattr(settings, "LOCAL_LLM_NAME", "/models/Qwen3-8B")
logger.info(f"Загрузка локальной модели: {model_name}")
logger.info(f"Инициализация vLLM движка с моделью: {model_name}")
self.tokenizer = AutoTokenizer.from_pretrained(model_name)
self.model = AutoModelForCausalLM.from_pretrained(
model_name,
torch_dtype=torch.float16,
device_map="cuda",
)
# Streamer для потоковой генерации
self.streamer = TextIteratorStreamer(
self.tokenizer,
skip_prompt=True,
skip_special_tokens=True
# Аргументы для асинхронного движка vLLM
engine_args = AsyncEngineArgs(
model=model_name,
tokenizer=model_name,
tokenizer_mode="auto",
trust_remote_code=True,
dtype=torch.float16, # float16
gpu_memory_utilization=0.85,
max_model_len=16384,
enable_prefix_caching=True,
tensor_parallel_size=1, # Увеличь, если у тебя несколько GPU
)
self.engine = AsyncLLMEngine.from_engine_args(engine_args)
self.tokenizer = self.engine.engine.tokenizer.tokenizer # доступ к токенизатору
else:
raise ValueError(f"Неподдерживаемый провайдер: {provider}")
@@ -46,7 +47,7 @@ class ChatWithAI:
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 = 5) -> List[Dict[str, Any]]:
"""Получение релевантного контекста из базы данных."""
try:
results = self.chroma_db.similarity_search(query, k=k)
@@ -71,8 +72,8 @@ class ChatWithAI:
)
return "\n---\n".join(formatted_context)
def generate_response_stream(self, query: str):
"""Генерация ответа с потоковой передачей токенов."""
async def generate_response_stream(self, query: str) -> AsyncGenerator[str, None]:
"""Генерация ответа с потоковой передачей токенов через vLLM."""
try:
logger.info(f"Пользовательский запрос: {query}")
context = self.get_relevant_context(query)
@@ -115,37 +116,36 @@ class ChatWithAI:
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
# Параметры генерации
sampling_params = SamplingParams(
temperature=0.3,
top_p=0.9,
max_tokens=2048,
stop_token_ids=[], # можно добавить, если нужно
)
generate_kwargs = {
"input_ids": inputs["input_ids"],
"max_new_tokens": 2048,
"temperature": 0.3,
"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()
# Потоковая передача токенов
# Генерация через vLLM
final_output = ""
buffer = ""
for token in self.streamer:
buffer += token
yield buffer # Отправляем частичный ответ
# Генерируем асинхронно
request_id = f"request-{hash(query)}"
try:
async for output in self.engine.generate(prompt, sampling_params, request_id):
if output.outputs:
text = output.outputs[0].text
# Отправляем только новые токены
new_text = text[len(final_output):]
final_output = text
if new_text:
buffer += new_text
yield buffer # Постепенно возвращаем накопленный текст
except Exception as e:
logger.error(f"Ошибка при генерации vLLM: {e}")
yield "Произошла ошибка при генерации ответа."
except Exception as e:
logger.error(f"Ошибка при генерации ответа: {e}")
logger.error(f"Неожиданная ошибка: {e}")
yield "Произошла ошибка при генерации ответа."
@@ -153,10 +153,27 @@ class ChatWithAI:
def main():
chat = ChatWithAI(provider="qwen3")
# Gradio требует синхронную функцию, но мы можем обернуть асинхронный вызов
def respond(message, history):
# Генерируем ответ по частям
for token in chat.generate_response_stream(message):
yield token
# Запускаем асинхронный генератор в синхронном контексте
async def async_generate():
async for token in chat.generate_response_stream(message):
yield token
# Используем asyncio.new_event_loop() для запуска внутри потока
loop = asyncio.new_event_loop()
asyncio.set_event_loop(loop)
try:
gen = async_generate()
while True:
try:
token = loop.run_until_complete(gen.__anext__())
yield token
except StopAsyncIteration:
break
finally:
loop.close()
demo = gr.ChatInterface(
fn=respond,
@@ -169,7 +186,9 @@ def main():
"Установка Unraid и первоначальная настройка"
],
)
demo.queue(max_size=20, default_concurrency_limit=10).launch(server_name="0.0.0.0", server_port=8080, share=False)
demo.queue(max_size=20, default_concurrency_limit=10).launch(
server_name="0.0.0.0", server_port=8080, share=False
)
if __name__ == "__main__":

View File

@@ -11,4 +11,8 @@ chromadb==0.6.3
sentence-transformers==3.4.1
langchain-chroma==0.2.2
pydantic-settings==2.8.1
langchain-text-splitters==0.3.7
langchain-text-splitters==0.3.7
vllm>=0.4.2
aiohttp>=3.8.0
fastapi>=0.100.0
starlette>=0.27.0