vLLM 框架部署 Rex-Omni 模型并提供标准 OpenAI 兼容接口的完整脚本
·
以下是使用 vLLM 框架部署 Rex-Omni 模型并提供标准 OpenAI 兼容接口的完整脚本。
一、环境准备与依赖安装
# 创建虚拟环境
conda create -n rexomni python=3.10 -y
conda activate rexomni
# 安装 PyTorch (CUDA 12.4)
pip install torch==2.6.0 torchvision==0.21.0 --index-url https://download.pytorch.org/whl/cu124
# 安装 vLLM
pip install vllm
# 克隆并安装 Rex-Omni
git clone https://github.com/IDEA-Research/Rex-Omni.git
cd Rex-Omni
pip install -v -e .
# 安装 FastAPI 相关依赖
pip install fastapi uvicorn pillow python-multipart
二、部署脚本
创建 deploy_rexomni_vllm.py:
#!/usr/bin/env python3
"""
Rex-Omni vLLM 部署脚本
提供标准 OpenAI 兼容接口
服务器 IP: 192.168.0.153
"""
import base64
import io
import json
import logging
from typing import Optional, List, Dict, Any
from pathlib import Path
from fastapi import FastAPI, HTTPException
from fastapi.responses import JSONResponse
from pydantic import BaseModel, Field
from PIL import Image
import uvicorn
# 配置日志
logging.basicConfig(level=logging.INFO)
logger = logging.getLogger(__name__)
# ==================== 请求/响应模型定义 ====================
class ChatMessage(BaseModel):
role: str
content: str
class ChatCompletionRequest(BaseModel):
"""OpenAI Chat Completion 请求格式"""
model: str = "IDEA-Research/Rex-Omni"
messages: List[ChatMessage]
temperature: Optional[float] = 0.0
max_tokens: Optional[int] = 2048
top_p: Optional[float] = 0.05
top_k: Optional[int] = 1
repetition_penalty: Optional[float] = 1.05
stream: Optional[bool] = False
class DetectionRequest(BaseModel):
"""Rex-Omni 专用检测请求格式"""
image_base64: str = Field(..., description="Base64编码的图像")
categories: List[str] = Field(..., description="要检测的类别列表")
task: str = Field(default="detection", description="任务类型: detection, pointing, ocr_box等")
temperature: Optional[float] = 0.0
max_tokens: Optional[int] = 2048
top_p: Optional[float] = 0.05
top_k: Optional[int] = 1
repetition_penalty: Optional[float] = 1.05
class ChatCompletionResponse(BaseModel):
"""OpenAI Chat Completion 响应格式"""
id: str
object: str = "chat.completion"
created: int
model: str
choices: List[Dict[str, Any]]
usage: Dict[str, int]
class DetectionResponse(BaseModel):
"""Rex-Omni 检测响应格式"""
success: bool
raw_output: Optional[str] = None
extracted_predictions: Optional[Dict] = None
error: Optional[str] = None
# ==================== 全局模型实例 ====================
rex_wrapper = None
def init_model(
model_path: str = "IDEA-Research/Rex-Omni",
backend: str = "vllm",
tensor_parallel_size: int = 1,
gpu_memory_utilization: float = 0.9,
max_model_len: int = 8192,
**kwargs
):
"""
初始化 Rex-Omni 模型
Args:
model_path: HuggingFace 模型 ID 或本地路径
backend: 推理后端 ("transformers" 或 "vllm")
tensor_parallel_size: vLLM 张量并行大小
gpu_memory_utilization: GPU 显存利用率
max_model_len: 最大模型长度
"""
global rex_wrapper
try:
from rex_omni import RexOmniWrapper
logger.info(f"正在加载 Rex-Omni 模型: {model_path}")
logger.info(f"后端: {backend}, tensor_parallel_size: {tensor_parallel_size}")
rex_wrapper = RexOmniWrapper(
model_path=model_path,
backend=backend,
max_tokens=kwargs.get('max_tokens', 2048),
temperature=kwargs.get('temperature', 0.0),
top_p=kwargs.get('top_p', 0.05),
top_k=kwargs.get('top_k', 1),
repetition_penalty=kwargs.get('repetition_penalty', 1.05),
# vLLM 特有参数
tokenizer_mode="auto",
limit_mm_per_prompt={"image": 1},
max_model_len=max_model_len,
gpu_memory_utilization=gpu_memory_utilization,
tensor_parallel_size=tensor_parallel_size,
trust_remote_code=True,
)
logger.info("模型加载成功")
except Exception as e:
logger.error(f"模型加载失败: {e}")
raise
# ==================== FastAPI 应用 ====================
app = FastAPI(
title="Rex-Omni API",
description="Rex-Omni 对象检测服务 - OpenAI 兼容接口",
version="1.0.0"
)
def decode_base64_image(base64_str: str) -> Image.Image:
"""解码 Base64 图像为 PIL Image"""
try:
# 去除可能的 data:image 前缀
if ',' in base64_str:
base64_str = base64_str.split(',')[1]
image_bytes = base64.b64decode(base64_str)
image = Image.open(io.BytesIO(image_bytes)).convert("RGB")
return image
except Exception as e:
raise HTTPException(status_code=400, detail=f"图像解码失败: {str(e)}")
@app.get("/health")
async def health_check():
"""健康检查接口"""
if rex_wrapper is None:
return JSONResponse(status_code=503, content={"status": "unhealthy", "message": "模型未加载"})
return {"status": "healthy", "model": "IDEA-Research/Rex-Omni", "backend": "vllm"}
@app.get("/v1/models")
async def list_models():
"""OpenAI 兼容的模型列表接口"""
return {
"object": "list",
"data": [
{
"id": "IDEA-Research/Rex-Omni",
"object": "model",
"created": 1700000000,
"owned_by": "IDEA-Research"
}
]
}
@app.post("/v1/chat/completions")
async def chat_completions(request: ChatCompletionRequest):
"""
OpenAI 兼容的 Chat Completion 接口
注意:此接口用于文本对话,Rex-Omni 的多模态能力需使用 /v1/detect 接口
"""
return JSONResponse(
status_code=400,
content={
"error": {
"message": "Rex-Omni 是多模态检测模型。请使用 /v1/detect 接口传递图像和检测类别。",
"type": "invalid_request_error"
}
}
)
@app.post("/v1/detect", response_model=DetectionResponse)
async def detect_objects(request: DetectionRequest):
"""
Rex-Omni 对象检测接口
请求示例:
{
"image_base64": "/9j/4AAQSkZJRg...",
"categories": ["person", "car", "dog"],
"task": "detection",
"temperature": 0.0
}
"""
if rex_wrapper is None:
raise HTTPException(status_code=503, detail="模型未加载")
try:
# 解码图像
image = decode_base64_image(request.image_base64)
logger.info(f"接收到检测请求,类别: {request.categories}, 任务: {request.task}")
# 执行推理
results = rex_wrapper.inference(
images=image,
task=request.task,
categories=request.categories,
)
result = results[0]
return DetectionResponse(
success=True,
raw_output=result.get("raw_output"),
extracted_predictions=result.get("extracted_predictions")
)
except HTTPException:
raise
except Exception as e:
logger.error(f"推理失败: {e}")
return DetectionResponse(
success=False,
error=str(e)
)
@app.get("/v1/detect/example")
async def get_example():
"""获取检测接口示例"""
return {
"example_request": {
"image_base64": "base64_encoded_image_string",
"categories": ["person", "cup", "laptop"],
"task": "detection",
"temperature": 0.0,
"max_tokens": 2048
},
"available_tasks": [
"detection",
"pointing",
"visual_prompting",
"keypoint",
"ocr_box",
"ocr_polygon",
"gui_grounding",
"gui_pointing"
],
"note": "图像需转换为 Base64 编码,支持 data:image 前缀"
}
# ==================== 主入口 ====================
if __name__ == "__main__":
import argparse
parser = argparse.ArgumentParser(description="Rex-Omni vLLM 部署脚本")
parser.add_argument("--model_path", type=str, default="IDEA-Research/Rex-Omni",
help="模型路径或 HuggingFace ID")
parser.add_argument("--host", type=str, default="192.168.0.153",
help="服务监听地址")
parser.add_argument("--port", type=int, default=8000,
help="服务端口")
parser.add_argument("--tensor_parallel_size", type=int, default=1,
help="vLLM 张量并行大小 (GPU数量)")
parser.add_argument("--gpu_memory_utilization", type=float, default=0.9,
help="GPU显存利用率")
parser.add_argument("--max_model_len", type=int, default=8192,
help="最大模型长度")
parser.add_argument("--temperature", type=float, default=0.0,
help="采样温度")
parser.add_argument("--max_tokens", type=int, default=2048,
help="最大输出token数")
args = parser.parse_args()
# 初始化模型
init_model(
model_path=args.model_path,
backend="vllm",
tensor_parallel_size=args.tensor_parallel_size,
gpu_memory_utilization=args.gpu_memory_utilization,
max_model_len=args.max_model_len,
temperature=args.temperature,
max_tokens=args.max_tokens,
)
# 启动服务
logger.info(f"启动服务: http://{args.host}:{args.port}")
logger.info(f"API 文档: http://{args.host}:{args.port}/docs")
logger.info(f"检测接口: POST http://{args.host}:{args.port}/v1/detect")
uvicorn.run(app, host=args.host, port=args.port, log_level="info")
三、启动脚本
创建 start_service.sh:
#!/bin/bash
# Rex-Omni 服务启动脚本
# 服务器 IP: 192.168.0.153
cd /path/to/Rex-Omni # 请替换为实际路径
# 激活虚拟环境
source ~/miniconda3/etc/profile.d/conda.sh
conda activate rexomni
# 启动服务
python deploy_rexomni_vllm.py \
--model_path IDEA-Research/Rex-Omni \
--host 192.168.0.153 \
--port 8000 \
--tensor_parallel_size 1 \
--gpu_memory_utilization 0.9 \
--max_model_len 8192 \
--temperature 0.0 \
--max_tokens 2048
赋予执行权限并启动:
chmod +x start_service.sh
./start_service.sh
四、使用 systemd 管理服务(可选)
创建 /etc/systemd/system/rexomni.service:
[Unit]
Description=Rex-Omni vLLM Service
After=network.target
[Service]
Type=simple
User=your_username
WorkingDirectory=/path/to/Rex-Omni
Environment="PATH=/home/your_username/miniconda3/envs/rexomni/bin:/usr/local/sbin:/usr/local/bin:/usr/sbin:/usr/bin:/sbin:/bin"
ExecStart=/home/your_username/miniconda3/envs/rexomni/bin/python /path/to/Rex-Omni/deploy_rexomni_vllm.py --host 192.168.0.153 --port 8000
Restart=always
RestartSec=10
[Install]
WantedBy=multi-user.target
启用服务:
sudo systemctl daemon-reload
sudo systemctl enable rexomni.service
sudo systemctl start rexomni.service
sudo systemctl status rexomni.service
五、测试调用
5.1 健康检查
curl http://192.168.0.153:8000/health
5.2 检测接口调用
import base64
import requests
from PIL import Image
import io
# 读取图像并转换为 Base64
def image_to_base64(image_path):
with open(image_path, "rb") as f:
return base64.b64encode(f.read()).decode("utf-8")
# 调用检测接口
url = "http://192.168.0.153:8000/v1/detect"
payload = {
"image_base64": image_to_base64("test.jpg"),
"categories": ["person", "cup", "laptop"],
"task": "detection",
"temperature": 0.0
}
response = requests.post(url, json=payload)
print(response.json())
或使用 curl:
curl -X POST http://192.168.0.153:8000/v1/detect \
-H "Content-Type: application/json" \
-d '{
"image_base64": "base64_encoded_image_string",
"categories": ["person", "cup"],
"task": "detection"
}'
六、支持的 API 端点
| 端点 | 方法 | 说明 |
|---|---|---|
/health |
GET | 健康检查 |
/v1/models |
GET | 模型列表(OpenAI兼容) |
/v1/chat/completions |
POST | OpenAI兼容接口(提示使用检测接口) |
/v1/detect |
POST | Rex-Omni 对象检测接口 |
/v1/detect/example |
GET | 获取检测接口示例 |
/docs |
GET | Swagger API 文档 |
七、注意事项
- 显存要求:vLLM 后端需要约 6-8GB 显存(3B 模型 + 8k 上下文)
- 启动等待:vLLM 首次加载模型需要约 5 分钟,请耐心等待
- GPU 数量:可通过
--tensor_parallel_size参数设置多卡并行 - 防火墙:确保服务器 8000 端口已开放
- 图像格式:支持 JPG、PNG 等常见格式,自动转换为 RGB
八、可用的任务类型
| Task | 说明 |
|---|---|
detection |
目标检测,返回边界框 |
pointing |
目标指向,返回坐标点 |
ocr_box |
OCR 文本框检测 |
ocr_polygon |
OCR 多边形检测 |
gui_grounding |
GUI 元素定位 |
gui_pointing |
GUI 元素指向 |
keypoint |
关键点检测(人体/手部/动物) |
visual_prompting |
视觉提示(需提供参考框) |
脚本已集成所有核心功能,启动后即可通过标准 HTTP 接口调用 Rex-Omni 模型。
更多推荐
所有评论(0)