add vllm
This commit is contained in:
@@ -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 как основного
|
||||
|
||||
@@ -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"
|
||||
|
||||
|
||||
119
app/main.py
119
app/main.py
@@ -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__":
|
||||
|
||||
@@ -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
|
||||
Reference in New Issue
Block a user