以下是使用 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 文档

七、注意事项

  1. 显存要求:vLLM 后端需要约 6-8GB 显存(3B 模型 + 8k 上下文)
  2. 启动等待:vLLM 首次加载模型需要约 5 分钟,请耐心等待
  3. GPU 数量:可通过 --tensor_parallel_size 参数设置多卡并行
  4. 防火墙:确保服务器 8000 端口已开放
  5. 图像格式:支持 JPG、PNG 等常见格式,自动转换为 RGB

八、可用的任务类型

Task 说明
detection 目标检测,返回边界框
pointing 目标指向,返回坐标点
ocr_box OCR 文本框检测
ocr_polygon OCR 多边形检测
gui_grounding GUI 元素定位
gui_pointing GUI 元素指向
keypoint 关键点检测(人体/手部/动物)
visual_prompting 视觉提示(需提供参考框)

脚本已集成所有核心功能,启动后即可通过标准 HTTP 接口调用 Rex-Omni 模型。

Logo

Agent 垂直技术社区,欢迎活跃、内容共建。

更多推荐