前言
在 LLM 浪潮席卷全球的今天,如何高效、稳定地将大模型部署到生产环境,成为了每个 AI 工程师必须面对的课题。vLLM 作为伯克利大学开源的高性能推理框架,凭借其 PagedAttention 技术和 Continuous Batching 机制,将推理吞吐量提升了 2-10 倍,成为业界首选的推理引擎。
本文将深入剖析 vLLM 的核心原理,详细讲解 Docker/K8s 部署配置,探讨多 GPU 并行策略,并重点分享笔者在生产环境中踩过的坑及解决方案。全文约 6000 字,建议收藏阅读。
一、vLLM 核心原理
1.1 传统推理的瓶颈
在介绍 vLLM 之前,我们先回顾一下传统 LLM 推理面临的核心问题:

传统推理存在两个主要瓶颈:
KV Cache 内存碎片化:每个请求的 KV Cache 需要连续内存分配,请求完成后释放,产生大量碎片
Static Batching 低效:必须等待批次内所有请求完成才能处理下一个批次,造成 GPU 资源浪费
1.2 PagedAttention:虚拟内存管理思想
vLLM 的核心创新是引入了操作系统虚拟内存的思路来管理 KV Cache:
┌─────────────────────────────────────────────────────────────────┐
│ PagedAttention 原理 │
├─────────────────────────────────────────────────────────────────┤
│ │
│ Physical GPU Memory (Block-based) │
│ ┌────────┬────────┬────────┬────────┬────────┬────────┐ │
│ │ Block 0│ Block 1│ Block 2│ Block 3│ Block 4│ Block 5│ ... │
│ │ token │ token │ token │ token │ token │ token │ │
│ │ tokens │ tokens │ tokens │ tokens │ tokens │ tokens │ │
│ └────────┴────────┴────────┴────────┴────────┴────────┘ │
│ │
│ Virtual KV Cache (Logical) │
│ │
│ Request A: [Block 0] → [Block 2] → [Block 5] │
│ Request B: [Block 0] → [Block 1] → [Block 3] │
│ Request C: [Block 4] → [Block 2] │
│ │
│ 特点: │
│ ✅ 非连续块分配,按需申请 │
│ ✅ 块级共享(共享前缀) │
│ ✅ 动态扩缩容,无碎片 │
│ ✅ 类似 OS 页表的映射机制 │
│ │
└─────────────────────────────────────────────────────────────────┘PagedAttention 的核心思想:
Block 化管理:将 KV Cache 按固定大小的块(Block)管理,默认 16 个 token/块
逻辑-物理映射:每个请求有自己的逻辑视图,映射到物理内存块
按需分配:只有需要时才分配新块,避免预分配浪费
共享前缀:相同前缀的请求可以共享物理块,实现 prefix caching
1.3 Continuous Batching:动态组装批次
┌─────────────────────────────────────────────────────────────────┐
│ Continuous Batching vs Static Batching │
├─────────────────────────────────────────────────────────────────┤
│ │
│ Static Batching (传统方式) │
│ ┌─────────────────────────────────────────────────────────┐ │
│ │ T=0 │ [Req A] [Req B] [Req C] │ │
│ │ T=1 │ [Req A●] [Req B●] [Req C●] │ │
│ │ T=2 │ [Req A●●] [Req B●●] [Req C●●] │ │
│ │ T=3 │ [Req A●●●] [Req B●●●] [Req C●●●] ⏳ 等待 │ │
│ │ T=4 │ [Req A●●●●] [Req B●●●●] [Req C●●●●] 完成 ✅ │ │
│ │ T=5 │ ⏳ GPU 空闲等待 │ │
│ │ T=6 │ [Req D] [Req E] [Req F] ← 新批次开始 │ │
│ └─────────────────────────────────────────────────────────┘ │
│ │
│ Continuous Batching (vLLM 方式) │
│ ┌─────────────────────────────────────────────────────────┐ │
│ │ T=0 │ [Req A] [Req B] [Req C] │ │
│ │ T=1 │ [Req A●] [Req B] [Req C] [Req D 插入] │ │
│ │ T=2 │ [Req A●●] [Req B●] [Req C] [Req E 插入] │ │
│ │ T=3 │ [Req A●●●] [Req D] [Req C●] [Req F 插入] │ │
│ │ T=4 │ [Req A●●●●] 完成 ✅ [Req D●] [Req C●●] [Req G] │ │
│ │ T=5 │ ✅ GPU 持续工作,无空闲 │ │
│ └─────────────────────────────────────────────────────────┘ │
│ │
│ 效果: 吞吐量提升 2-10 倍 │
│ │
└─────────────────────────────────────────────────────────────────┘Continuous Batching 的核心是动态迭代级调度(Iteration-level Scheduling):
每生成一个 token 后,检查是否有请求完成
完成的请求立即退出,释放资源
新请求立即插入,充分利用 GPU
二、环境准备与依赖
2.1 硬件要求
2.2 基础镜像选择
推荐使用 NVIDIA 官方 PyTorch 镜像作为基础:
# Dockerfile.vllm
FROM nvidia/cuda:12.1.0-devel-ubuntu22.04
# 设置环境变量
ENV DEBIAN_FRONTEND=noninteractive
ENV PYTHONUNBUFFERED=1
ENV TRANSFORMERS_CACHE=/model_cache/transformers
ENV HF_HOME=/model_cache/huggingface
# 安装基础依赖
RUN apt-get update && apt-get install -y \
python3.10 \
python3-pip \
git \
curl \
wget \
vim \
htop \
&& rm -rf /var/lib/apt/lists/*
# 设置 Python 链接
RUN ln -sf /usr/bin/python3 /usr/bin/python
# 安装 vLLM(指定版本以保证稳定性)
RUN pip3 install --no-cache-dir \
vllm==0.4.3 \
transformers==4.40.0 \
accelerate==0.28.0 \
sentencepiece==0.1.99 \
tiktoken==0.6.0
# 安装监控工具
RUN pip3 install --no-cache-dir \
prometheus-client \
psutil \
GPUtil
# 创建工作目录
WORKDIR /app
# 复制应用代码
COPY app/ ./app/
# 下载模型(可选,或使用 volume mount)
# RUN python -c "from vllm import LLM; LLM.from_pretrained('meta-llama/Llama-2-7b-hf')"
EXPOSE 8000 8001
CMD ["python", "/app/server.py"]2.3 模型下载与管理
在部署 vLLM 之前,需要先下载模型文件。以下是几种常用方法:
方式 1: 使用 Hugging Face CLI
# 安装 huggingface-cli
pip install huggingface-hub
# 登录(如果是私有模型)
huggingface-cli login
# 下载模型
huggingface-cli download \
meta-llama/Llama-2-7b-chat-hf \
--local-dir /data/models/llama-2-7b-chat \
--local-dir-use-symlinks False
# 查看下载进度
du -sh /data/models/llama-2-7b-chat方式 2: 使用 Python 脚本(支持断点续传)
#!/usr/bin/env python3
"""
模型下载脚本
支持断点续传、并行下载
"""
from huggingface_hub import snapshot_download
import os
from pathlib import Path
def download_model(
model_id: str,
local_dir: str,
token: str = None,
max_workers: int = 4
):
"""
下载 Hugging Face 模型
Args:
model_id: 模型 ID,如 "meta-llama/Llama-2-7b-chat-hf"
local_dir: 本地存储路径
token: HF Token(私有模型需要)
max_workers: 并行下载线程数
"""
print(f"开始下载模型: {model_id}")
print(f"目标路径: {local_dir}")
snapshot_download(
repo_id=model_id,
local_dir=local_dir,
local_dir_use_symlinks=False,
token=token,
max_workers=max_workers,
resume_download=True, # 支持断点续传
ignore_patterns=["*.msgpack", "*.h5"] # 忽略不需要的文件
)
print(f"✅ 模型下载完成")
print(f" 大小: {get_dir_size(local_dir)}")
def get_dir_size(path: str) -> str:
"""获取目录大小"""
total = 0
for dirpath, dirnames, filenames in os.walk(path):
for f in filenames:
fp = os.path.join(dirpath, f)
if os.path.exists(fp):
total += os.path.getsize(fp)
# 转换为人类可读格式
for unit in ['B', 'KB', 'MB', 'GB', 'TB']:
if total < 1024.0:
return f"{total:.2f} {unit}"
total /= 1024.0
return f"{total:.2f} PB"
if __name__ == "__main__":
download_model(
model_id="meta-llama/Llama-2-7b-chat-hf",
local_dir="/data/models/llama-2-7b-chat",
token=os.getenv("HF_TOKEN"),
max_workers=8
)方式 3: 使用 ModelScope(国内镜像,速度更快)
# 安装 modelscope
pip install modelscope
# 下载模型(国内速度更快)
python -c "
from modelscope import snapshot_download
snapshot_download('qwen/Qwen2.5-7B-Instruct', cache_dir='/data/models')
"模型文件结构
/data/models/llama-2-7b-chat/
├── config.json # 模型配置
├── generation_config.json # 生成配置
├── pytorch_model.bin.index.json
├── pytorch_model-00001-of-00002.bin # 模型权重
├── pytorch_model-00002-of-00002.bin
├── special_tokens_map.json
├── tokenizer.json # 分词器
├── tokenizer.model
└── tokenizer_config.json三、Docker Compose 单机部署
3.1 完整配置
# docker-compose.vllm.yml
version: '3.8'
services:
# vLLM 推理服务
vllm-server:
image: vllm/vllm-openai:v0.4.3
container_name: vllm-inference
restart: unless-stopped
ports:
- "8000:8000" # API 端口
- "8001:8001" # Metrics 端口
# 使用 command 而不是 environment 来传递 vLLM 参数
command: >
python -m vllm.entrypoints.openai.api_server
--model ${MODEL_NAME:-meta-llama/Llama-2-7b-chat-hf}
--host 0.0.0.0
--port 8000
--tensor-parallel-size 4
--pipeline-parallel-size 1
--trust-remote-code
--dtype auto
--max-model-len 8192
--gpu-memory-utilization 0.92
--max-num-seqs 256
--max-num-batched-tokens 32768
--enable-prefix-caching
--worker-use-ray
environment:
# GPU 配置
- CUDA_VISIBLE_DEVICES=0,1,2,3
# HuggingFace Token(如果需要)
- HF_TOKEN=${HF_TOKEN}
volumes:
# 模型存储
- /data/models:/model:ro
- model_cache:/root/.cache/huggingface
# 应用配置
- ./config:/app/config:ro
- ./app:/app
# 监控数据
- vllm_logs:/var/log/vllm
deploy:
resources:
reservations:
devices:
- driver: nvidia
count: 4
capabilities: [gpu]
limits:
memory: 64G
reservations:
memory: 32G
healthcheck:
test: ["CMD", "curl", "-f", "http://localhost:8000/health"]
interval: 30s
timeout: 10s
retries: 3
start_period: 300s
networks:
- vllm-net
shm_size: '64g' # 共享内存,用于 Ray workers
# Redis 缓存(用于请求去重/限流)
redis:
image: redis:7-alpine
container_name: vllm-redis
restart: unless-stopped
ports:
- "6379:6379"
command: redis-server --maxmemory 512mb --maxmemory-policy allkeys-lru
volumes:
- redis_data:/data
networks:
- vllm-net
# Prometheus 监控
prometheus:
image: prom/prometheus:v2.48.0
container_name: vllm-prometheus
restart: unless-stopped
ports:
- "9090:9090"
volumes:
- ./config/prometheus.yml:/etc/prometheus/prometheus.yml:ro
- prometheus_data:/prometheus
command:
- '--config.file=/etc/prometheus/prometheus.yml'
- '--storage.tsdb.path=/prometheus'
- '--web.console.libraries=/usr/share/prometheus/console_libraries'
- '--web.console.templates=/usr/share/prometheus/consoles'
networks:
- vllm-net
# Grafana 可视化
grafana:
image: grafana/grafana:10.2.2
container_name: vllm-grafana
restart: unless-stopped
ports:
- "3000:3000"
environment:
- GF_SECURITY_ADMIN_USER=admin
- GF_SECURITY_ADMIN_PASSWORD=${GRAFANA_PASSWORD:-admin123}
- GF_USERS_ALLOW_SIGN_UP=false
volumes:
- grafana_data:/var/lib/grafana
- ./config/grafana/provisioning:/etc/grafana/provisioning:ro
depends_on:
- prometheus
networks:
- vllm-net
networks:
vllm-net:
driver: bridge
volumes:
model_cache:
vllm_logs:
redis_data:
prometheus_data:
grafana_data:3.2 启动命令
# 启动服务
docker-compose -f docker-compose.vllm.yml up -d
# 查看日志
docker-compose -f docker-compose.vllm.yml logs -f vllm-server
# 查看资源使用
docker stats vllm-inference
# 测试 API
curl http://localhost:8000/v1/models四、Kubernetes 分布式部署
4.1 整体架构
┌─────────────────────────────────────────────────────────────────┐
│ K8s vLLM 部署架构 │
├─────────────────────────────────────────────────────────────────┤
│ │
│ ┌──────────────┐ │
│ │ Ingress │ │
│ │ (NGINX) │ │
│ └──────┬───────┘ │
│ │ │
│ ┌──────▼───────┐ │
│ │ Service │ │
│ │ (L4 Load) │ │
│ └──────┬───────┘ │
│ │ │
│ ┌───────────────────────────┼───────────────────────────┐ │
│ │ │ │ │
│ ┌──▼─────────┐ ┌──────▼──────┐ ┌─────────▼┐ │
│ │ Pod-1 │ │ Pod-2 │ │ Pod-N │ │
│ │ ┌────┐ │ │ ┌────┐ │ │ ┌────┐ │ │
│ │ │Ray │ │◄─gRPC───►│ │Ray │ │◄─gRPC───►│ │Ray │ │ │
│ │ │Wkr1│ │ │ │Wkr1│ │ │ │Wkr1│ │ │
│ │ └────┘ │ │ └────┘ │ │ └────┘ │ │
│ │ ┌────┐ │ │ ┌────┐ │ │ ┌────┐ │ │
│ │ │Ray │ │ │ │Ray │ │ │ │Ray │ │ │
│ │ │Wkr2│ │ │ │Wkr2│ │ │ │Wkr2│ │ │
│ │ └────┘ │ │ └────┘ │ │ └────┘ │ │
│ └─────┬──────┘ └──────┬───────┘ └─────┬─────┘ │
│ │ │ │ │
│ └───────────────────────┼────────────────────────┘ │
│ │ │
│ ┌──────────▼──────────┐ │
│ │ PVC (Model) │ │
│ │ NFS / S3 │ │
│ └─────────────────────┘ │
│ │
└─────────────────────────────────────────────────────────────────┘4.2 Namespace 与 RBAC
# k8s/namespace.yaml
apiVersion: v1
kind: Namespace
metadata:
name: vllm-inference
labels:
app: vllm
environment: production
---
# k8s/service-account.yaml
apiVersion: v1
kind: ServiceAccount
metadata:
name: vllm-inference-sa
namespace: vllm-inference
---
# k8s/rbac.yaml
apiVersion: rbac.authorization.k8s.io/v1
kind: Role
metadata:
name: vllm-pod-reader
namespace: vllm-inference
rules:
- apiGroups: [""]
resources: ["pods"]
verbs: ["get", "list", "watch"]
- apiGroups: [""]
resources: ["services"]
verbs: ["get", "list", "watch", "create", "update", "patch"]
---
apiVersion: rbac.authorization.k8s.io/v1
kind: RoleBinding
metadata:
name: vllm-pod-reader-binding
namespace: vllm-inference
subjects:
- kind: ServiceAccount
name: vllm-inference-sa
namespace: vllm-inference
roleRef:
kind: Role
name: vllm-pod-reader
apiGroup: rbac.authorization.k8s.io4.3 ConfigMap 配置
# k8s/configmap.yaml
apiVersion: v1
kind: ConfigMap
metadata:
name: vllm-config
namespace: vllm-inference
data:
vllm-config.yaml: |
# 模型配置
model:
name: "meta-llama/Llama-2-7b-chat-hf"
path: "/model"
trust_remote_code: true
dtype: "auto"
# 并行策略
parallel:
tensor_parallel_size: 4
pipeline_parallel_size: 1
world_size: 4
# 资源限制
resources:
gpu_memory_utilization: 0.92
max_model_len: 8192
max_num_seqs: 256
max_num_batched_tokens: 32768
# Prefix Caching
cache:
enable_prefix_caching: true
enable_chunked_prefill: true
# 服务配置
server:
host: "0.0.0.0"
port: 8000
worker_use_ray: true
trust_remote_code: true
# 调度配置
scheduler:
enable_chunked_prefill: true
max_num_batched_tokens: 32768
max_num_seqs: 256
preemption_mode: "swap"
# 健康检查
health:
check_period: 30
check_timeout: 60
failure_threshold: 3
success_threshold: 1
initial_delay: 3004.4 PV/PVC 存储配置
# k8s/storage.yaml
apiVersion: v1
kind: PersistentVolumeClaim
metadata:
name: model-storage
namespace: vllm-inference
spec:
accessModes:
- ReadOnlyMany
resources:
requests:
storage: 200Gi
storageClassName: nfs-model
---
# 如果使用 S3 存储模型
# k8s/s3-secret.yaml
apiVersion: v1
kind: Secret
metadata:
name: s3-model-secret
namespace: vllm-inference
type: Opaque
stringData:
AWS_ACCESS_KEY_ID: "${AWS_ACCESS_KEY_ID}"
AWS_SECRET_ACCESS_KEY: "${AWS_SECRET_ACCESS_KEY}"
AWS_DEFAULT_REGION: "us-east-1"4.5 StatefulSet 部署(多 GPU)
# k8s/statefulset.yaml
apiVersion: apps/v1
kind: StatefulSet
metadata:
name: vllm-inference
namespace: vllm-inference
labels:
app: vllm
component: inference
spec:
serviceName: vllm-inference-headless
replicas: 1 # 使用 Ray 分布式,一个 Pod 管理多 GPU
podManagementPolicy: Parallel
selector:
matchLabels:
app: vllm
component: inference
template:
metadata:
labels:
app: vllm
component: inference
annotations:
prometheus.io/scrape: "true"
prometheus.io/port: "8001"
prometheus.io/path: "/metrics"
spec:
serviceAccountName: vllm-inference-sa
restartPolicy: Always
terminationGracePeriodSeconds: 60
# 亲和性调度 - 确保在同一节点
affinity:
nodeAffinity:
requiredDuringSchedulingIgnoredDuringExecution:
nodeSelectorTerms:
- matchExpressions:
- key: nvidia.com/gpu.count
operator: Gte
values:
- "4"
podAntiAffinity:
preferredDuringSchedulingIgnoredDuringExecution:
- weight: 100
podAffinityTerm:
labelSelector:
matchExpressions:
- key: app
operator: In
values:
- vllm
topologyKey: kubernetes.io/hostname
# 容忍污点
tolerations:
- key: "nvidia.com/gpu"
operator: "Exists"
effect: "NoSchedule"
containers:
- name: vllm
image: vllm/vllm-openai:v0.4.3
imagePullPolicy: IfNotPresent
ports:
- name: api
containerPort: 8000
protocol: TCP
- name: metrics
containerPort: 8001
protocol: TCP
- name: ray-dashboard
containerPort: 8265
protocol: TCP
# 环境变量
env:
# 模型配置
- name: MODEL_NAME
value: "meta-llama/Llama-2-7b-chat-hf"
- name: MODEL_PATH
value: "/model"
# vLLM 启动参数
- name: TP_SIZE
value: "4"
- name: MAX_MODEL_LEN
value: "8192"
- name: GPU_MEMORY_UTILIZATION
value: "0.92"
- name: MAX_NUM_SEQS
value: "256"
- name: MAX_NUM_BATCHED_TOKENS
value: "32768"
- name: ENABLE_PREFIX_CACHING
value: "true"
- name: ENABLE_CHUNKED_PREFILL
value: "true"
# Ray 配置
- name: RAY_DASHBOARD_PORT
value: "8265"
- name: RAY_memory_monitor_refresh_ms
value: "0"
# 安全配置
- name: HF_TOKEN
valueFrom:
secretKeyRef:
name: hf-secret
key: token
optional: true
# 启动命令(使用 command 而不是依赖环境变量)
command:
- python
- -m
- vllm.entrypoints.openai.api_server
args:
- --model
- $(MODEL_PATH)
- --host
- "0.0.0.0"
- --port
- "8000"
- --tensor-parallel-size
- $(TP_SIZE)
- --max-model-len
- $(MAX_MODEL_LEN)
- --gpu-memory-utilization
- $(GPU_MEMORY_UTILIZATION)
- --max-num-seqs
- $(MAX_NUM_SEQS)
- --max-num-batched-tokens
- $(MAX_NUM_BATCHED_TOKENS)
- --enable-prefix-caching
- --enable-chunked-prefill
- --trust-remote-code
- --worker-use-ray
resources:
requests:
cpu: "16"
memory: "64Gi"
nvidia.com/gpu: 4
limits:
cpu: "32"
memory: "128Gi"
nvidia.com/gpu: 4
volumeMounts:
- name: model-storage
mountPath: /model
readOnly: true
- name: config
mountPath: /app/config/vllm-config.yaml
subPath: vllm-config.yaml
- name: shm
mountPath: /dev/shm
# 健康检查
readinessProbe:
httpGet:
path: /health
port: 8000
initialDelaySeconds: 300
periodSeconds: 30
timeoutSeconds: 10
failureThreshold: 3
livenessProbe:
httpGet:
path: /health
port: 8000
initialDelaySeconds: 300
periodSeconds: 60
timeoutSeconds: 30
failureThreshold: 5
lifecycle:
preStop:
exec:
command:
- /bin/sh
- -c
- |
echo "Graceful shutdown initiated..."
# 可以添加优雅停止逻辑
sleep 10
# 安全上下文
securityContext:
runAsNonRoot: false
runAsUser: 0
allowPrivilegeEscalation: true
readOnlyRootFilesystem: false
volumes:
- name: model-storage
persistentVolumeClaim:
claimName: model-storage
- name: config
configMap:
name: vllm-config
- name: shm
emptyDir:
medium: Memory
sizeLimit: 64Gi
# 拓扑分布约束
topologySpreadConstraints:
- maxSkew: 1
topologyKey: topology.kubernetes.io/zone
whenUnsatisfiable: ScheduleAnyway
labelSelector:
matchLabels:
app: vllm4.6 Service 配置
# k8s/service.yaml
apiVersion: v1
kind: Service
metadata:
name: vllm-inference
namespace: vllm-inference
labels:
app: vllm
component: inference
annotations:
# AWS ALB 配置(如使用 AWS)
service.beta.kubernetes.io/aws-load-balancer-type: "nlb"
spec:
type: ClusterIP
ports:
- name: api
port: 8000
targetPort: 8000
protocol: TCP
- name: metrics
port: 8001
targetPort: 8001
protocol: TCP
selector:
app: vllm
component: inference
---
# 无头服务(用于 Ray 内部通信)
apiVersion: v1
kind: Service
metadata:
name: vllm-inference-headless
namespace: vllm-inference
spec:
clusterIP: None
ports:
- name: ray-dashboard
port: 8265
targetPort: 8265
selector:
app: vllm
component: inference4.7 HPA 自动扩缩容
# k8s/hpa.yaml
apiVersion: autoscaling/v2
kind: HorizontalPodAutoscaler
metadata:
name: vllm-inference-hpa
namespace: vllm-inference
spec:
scaleTargetRef:
apiVersion: apps/v1
kind: StatefulSet
name: vllm-inference
minReplicas: 1
maxReplicas: 3
metrics:
- type: Resource
resource:
name: gpu-memory
target:
type: Utilization
averageUtilization: 85
- type: Pods
pods:
metric:
name: vllm_pending_requests
target:
type: AverageValue
averageValue: "10"
behavior:
scaleDown:
stabilizationWindowSeconds: 600
policies:
- type: Percent
value: 50
periodSeconds: 60
scaleUp:
stabilizationWindowSeconds: 0
policies:
- type: Percent
value: 100
periodSeconds: 15
- type: Pods
value: 1
periodSeconds: 15
selectPolicy: Max五、生产环境踩坑记录
5.1 OOM 问题排查
问题现象:GPU 显存溢出,Pod 被 OOM Killer 终止
原因分析:
KV Cache 超出限制
max_model_len 设置过大
单请求 token 数过多
Batch Size 配置不当
max_num_seqs 过大
max_num_batched_tokens 过大
共享内存不足
shm_size 设置过小
模型加载时显存峰值
dtype=float16 vs dtype=bfloat16
量化未正确应用
解决方案:
# 问题排查脚本:check_gpu_memory.py
#!/usr/bin/env python3
"""
vLLM GPU 显存分析工具
帮助诊断 OOM 问题的原因
"""
import os
import sys
import argparse
import subprocess
from typing import Dict, List, Tuple
from dataclasses import dataclass
@dataclass
class MemoryInfo:
"""GPU 显存信息"""
total: int # 总显存 (MB)
used: int # 已使用 (MB)
free: int # 空闲 (MB)
utilization: float # 使用率
def get_gpu_memory_info(device_id: int = 0) -> MemoryInfo:
"""
获取指定 GPU 的显存信息
Args:
device_id: GPU 设备 ID
Returns:
MemoryInfo 对象
"""
try:
result = subprocess.run(
['nvidia-smi', '--query-gpu=memory.total,memory.used,memory.free',
'--format=csv,noheader,nounits', f'-i={device_id}'],
capture_output=True,
text=True,
check=True
)
total, used, free = map(int, result.stdout.strip().split(','))
return MemoryInfo(
total=total,
used=used,
free=free,
utilization=used / total * 100
)
except Exception as e:
print(f"获取 GPU 信息失败: {e}")
return MemoryInfo(0, 0, 0, 0)
def calculate_required_memory(
model_name: str,
tensor_parallel_size: int = 1,
max_model_len: int = 8192,
dtype: str = "float16",
enable_prefix_caching: bool = False
) -> Dict[str, int]:
"""
计算 vLLM 运行所需显存
Args:
model_name: 模型名称
tensor_parallel_size: Tensor 并行数
max_model_len: 最大模型长度
dtype: 数据类型
enable_prefix_caching: 是否启用前缀缓存
Returns:
显存使用分解字典
"""
# 模型权重估算(基于参数量)
# 这是一个估算值,实际需要根据具体模型调整
model_weights_estimate = {
"7b": 14 * 1024, # 7B 模型约 14GB (FP16)
"13b": 26 * 1024,
"34b": 68 * 1024,
"70b": 140 * 1024,
}
# 估算 KV Cache 需求
# 每个 token 的 KV Cache 大小 = 2 * layers * hidden_size * dtype_size
# 以 Llama 7B 为例: 2 * 32 * 4096 * 2 bytes = 0.5 MB/token
kv_cache_per_token_mb = 0.5 # FP16
# 计算总 KV Cache 需求
kv_cache_total = kv_cache_per_token_mb * max_model_len
# Prefix caching 可以减少部分缓存需求
if enable_prefix_caching:
kv_cache_total *= 0.7 # 假设节省 30%
# 激活值估算
activation_mb = max_model_len * 0.1 # 估算
results = {
"model_weights_mb": 0, # 需要根据实际模型查询
"kv_cache_mb": int(kv_cache_total),
"activation_mb": int(activation_mb),
"overhead_mb": 2048, # 其他开销
}
# 每个 GPU 的分配
results["per_gpu_mb"] = (
results["model_weights_mb"] // tensor_parallel_size +
results["kv_cache_mb"] +
results["activation_mb"] +
results["overhead_mb"]
)
return results
def check_oom_risk(
total_memory_mb: int,
required_memory_mb: int,
utilization_threshold: float = 0.95
) -> Tuple[bool, str]:
"""
检查 OOM 风险
Args:
total_memory_mb: 总显存
required_memory_mb: 需求显存
utilization_threshold: 风险阈值
Returns:
(是否有风险, 风险描述)
"""
risk = required_memory_mb > total_memory_mb * utilization_threshold
if required_memory_mb > total_memory_mb:
return True, "⚠️ 严重风险: 所需显存超过总显存,会导致 OOM"
elif risk:
deficit = required_memory_mb - total_memory_mb * 0.9
return True, f"⚠️ 中等风险: 显存使用率将超过 90%,建议增加 {int(deficit)} MB"
else:
safe_margin = total_memory_mb - required_memory_mb
return False, f"✅ 正常: 预估需要 {required_memory_mb} MB,剩余 {safe_margin} MB"
def recommend_config(
total_memory_mb: int,
model_size_gb: int = 14
) -> Dict[str, Dict]:
"""
根据可用显存推荐配置
Args:
total_memory_mb: 总显存 (MB)
model_size_gb: 模型大小 (GB)
Returns:
推荐配置字典
"""
model_memory_mb = model_size_gb * 1024
available_for_cache = total_memory_mb - model_memory_mb - 2048 # 预留开销
recommendations = []
# 计算不同 max_model_len 的可行性
for max_len in [2048, 4096, 8192, 16384]:
kv_per_token = 0.5 # MB
required = max_len * kv_per_token
if required <= available_for_cache * 0.7: # 保留 30% 余量
recommendations.append({
"max_model_len": max_len,
"kv_cache_mb": int(required),
"feasible": True,
"recommendation": "推荐" if max_len in [4096, 8192] else "可用"
})
else:
# 计算安全的 max_model_len
safe_len = int(available_for_cache * 0.7 / kv_per_token)
recommendations.append({
"max_model_len": max_len,
"feasible": False,
"safe_alternative": safe_len
})
return {
"total_memory_mb": total_memory_mb,
"model_memory_mb": model_memory_mb,
"available_for_cache_mb": available_for_cache,
"configurations": recommendations
}
def main():
"""主函数"""
parser = argparse.ArgumentParser(description="vLLM GPU 显存分析工具")
parser.add_argument("--device", "-d", type=int, default=0, help="GPU 设备 ID")
parser.add_argument("--model", "-m", type=str, default="llama-7b", help="模型大小")
parser.add_argument("--tp-size", "-t", type=int, default=1, help="Tensor 并行数")
parser.add_argument("--max-len", type=int, default=8192, help="最大序列长度")
parser.add_argument("--recommend", "-r", action="store_true", help="输出配置建议")
args = parser.parse_args()
print("=" * 60)
print("vLLM GPU 显存分析报告")
print("=" * 60)
# 获取 GPU 信息
print(f"\n📊 GPU {args.device} 状态:")
mem_info = get_gpu_memory_info(args.device)
print(f" 总显存: {mem_info.total} MB ({mem_info.total/1024:.1f} GB)")
print(f" 已使用: {mem_info.used} MB ({mem_info.used/1024:.1f} GB)")
print(f" 空闲: {mem_info.free} MB ({mem_info.free/1024:.1f} GB)")
print(f" 使用率: {mem_info.utilization:.1f}%")
# 计算所需显存
print(f"\n📋 vLLM 显存需求计算:")
required = calculate_required_memory(
model_name=args.model,
tensor_parallel_size=args.tp_size,
max_model_len=args.max_len
)
print(f" 模型权重: ~{required['model_weights_mb']} MB")
print(f" KV Cache: {required['kv_cache_mb']} MB")
print(f" 激活值: {required['activation_mb']} MB")
print(f" 系统开销: {required['overhead_mb']} MB")
print(f" 每 GPU 需求: {required['per_gpu_mb']} MB")
# 检查风险
print(f"\n🔍 OOM 风险评估:")
has_risk, message = check_oom_risk(mem_info.total, required['per_gpu_mb'])
print(f" {message}")
# 推荐配置
if args.recommend:
print(f"\n💡 配置建议:")
model_size_map = {"llama-7b": 14, "llama-13b": 26, "llama-70b": 140}
model_size = model_size_map.get(args.model, 14)
rec = recommend_config(mem_info.free, model_size)
for config in rec["configurations"]:
if config.get("feasible"):
print(f" max_model_len={config['max_model_len']}: "
f"{config['kv_cache_mb']} MB - {config['recommendation']}")
else:
print(f" max_model_len={config['max_model_len']}: "
f"不可用,建议使用 max_model_len={config.get('safe_alternative', 0)}")
print("\n" + "=" * 60)
if __name__ == "__main__":
main()5.2 长文本截断问题
问题:生成长文本时意外截断
原因:max_model_len 设置过小,或生成长度达到限制
解决方案:
# 长文本处理中间件:long_text_handler.py
#!/usr/bin/env python3
"""
vLLM 长文本处理中间件
处理生成长度限制和文本截断问题
"""
import time
import json
import logging
from typing import Dict, Optional, List, Any, Callable
from dataclasses import dataclass, field
from enum import Enum
from functools import wraps
import asyncio
# 配置日志
logging.basicConfig(level=logging.INFO)
logger = logging.getLogger(__name__)
class TruncationStrategy(Enum):
"""截断策略"""
TRUNCATE = "truncate" # 直接截断
STOP_SEQUENCE = "stop_sequence" # 遇到停止符截断
SMART_TRUNCATE = "smart_truncate" # 智能截断(保留完整句子)
@dataclass
class GenerationConfig:
"""生成配置"""
max_tokens: int = 4096
max_model_len: int = 8192
truncation_strategy: TruncationStrategy = TruncationStrategy.SMART_TRUNCATE
stop_sequences: List[str] = field(default_factory=lambda: ["\n\n", "---", "##"])
min_tokens_before_truncate: int = 100 # 最小截断token数
@dataclass
class GenerationRequest:
"""生成请求"""
prompt: str
config: GenerationConfig
request_id: str
metadata: Dict[str, Any] = field(default_factory=dict)
@dataclass
class GenerationResponse:
"""生成响应"""
text: str
prompt: str
num_tokens: int
finish_reason: str
truncation_detected: bool
metadata: Dict[str, Any]
class LongTextHandler:
"""
长文本处理器
解决长文本生成分段、截断问题
"""
def __init__(
self,
api_base: str = "http://localhost:8000/v1",
default_config: Optional[GenerationConfig] = None
):
"""
初始化长文本处理器
Args:
api_base: vLLM API 地址
default_config: 默认生成配置
"""
self.api_base = api_base.rstrip("/")
self.default_config = default_config or GenerationConfig()
self._session_stats = {}
def _check_truncation(self, text: str, config: GenerationConfig) -> bool:
"""
检查文本是否被截断
Args:
text: 生成的文本
config: 生成配置
Returns:
是否被截断
"""
# 检查是否在句子中间截断
for stop_seq in config.stop_sequences:
if text.endswith(stop_seq):
# 停止符在文本末尾,可能是正常停止
return False
# 检查最后是否为空格结尾(可能在单词中间截断)
if text and not text[-1].isspace() and len(text) > 50:
# 最后一个词可能不完整
last_words = text.split()[-3:]
if any(len(w) < 3 for w in last_words):
return True
return False
def _smart_truncate(self, text: str, config: GenerationConfig) -> str:
"""
智能截断,保留完整句子
Args:
text: 原始文本
config: 生成配置
Returns:
截断后的文本
"""
if not text.strip():
return text
# 尝试找到最后一个句号、问号或感叹号
for seq in ["。", ". ", "?", "!", "\n\n"]:
last_pos = text.rfind(seq)
if last_pos > len(text) * 0.7: # 在后30%位置找到
return text[:last_pos + len(seq)].strip()
# 尝试找逗号、换行等次级断点
for seq in [", ", ",", "\n"]:
last_pos = text.rfind(seq)
if last_pos > len(text) * 0.8:
return text[:last_pos + len(seq)].strip()
# 无法找到合适断点,返回原文
return text.strip()
async def generate_stream(
self,
prompt: str,
config: Optional[GenerationConfig] = None,
callback: Optional[Callable[[str], None]] = None
) -> GenerationResponse:
"""
流式生成文本
Args:
prompt: 输入提示词
config: 生成配置
callback: 流式回调函数
Returns:
GenerationResponse 对象
"""
config = config or self.default_config
# 预估是否需要分段处理
estimated_tokens = len(prompt.split()) * 1.3 + config.max_tokens
needs_chunking = estimated_tokens > config.max_model_len * 0.8
if needs_chunking:
logger.warning(
f"请求可能超出限制,预估 {estimated_tokens:.0f} tokens "
f"(限制: {config.max_model_len})"
)
full_text = []
total_tokens = 0
truncated = False
try:
import aiohttp
async with aiohttp.ClientSession() as session:
payload = {
"prompt": prompt,
"max_tokens": config.max_tokens,
"temperature": 0.7,
"top_p": 0.95,
"stream": True,
}
async with session.post(
f"{self.api_base}/completions",
json=payload,
timeout=aiohttp.ClientTimeout(total=300)
) as response:
if response.status != 200:
error_text = await response.text()
raise Exception(f"API 错误: {response.status} - {error_text}")
async for line in response.content:
line = line.decode('utf-8').strip()
if not line or not line.startswith('data: '):
continue
data = line[6:] # 去掉 "data: " 前缀
if data == "[DONE]":
break
try:
chunk = json.loads(data)
token = chunk.get("choices", [{}])[0].get("text", "")
full_text.append(token)
total_tokens += 1
if callback:
callback(token)
except json.JSONDecodeError:
continue
result_text = "".join(full_text)
# 检查截断
if self._check_truncation(result_text, config):
truncated = True
if config.truncation_strategy == TruncationStrategy.SMART_TRUNCATE:
result_text = self._smart_truncate(result_text, config)
return GenerationResponse(
text=result_text,
prompt=prompt,
num_tokens=total_tokens,
finish_reason="length" if truncated else "stop",
truncation_detected=truncated,
metadata={
"estimated_total": estimated_tokens,
"needs_chunking": needs_chunking
}
)
except asyncio.TimeoutError:
logger.error("生成超时")
return GenerationResponse(
text="".join(full_text),
prompt=prompt,
num_tokens=total_tokens,
finish_reason="timeout",
truncation_detected=True,
metadata={"error": "timeout"}
)
except Exception as e:
logger.error(f"生成失败: {e}")
raise
def generate_chunked(
self,
prompt: str,
config: Optional[GenerationConfig] = None,
overlap_tokens: int = 50
) -> List[GenerationResponse]:
"""
分段生成(用于超长文本)
Args:
prompt: 输入提示词
config: 生成配置
overlap_tokens: 分段重叠 token 数
Returns:
分段生成结果列表
"""
config = config or self.default_config
# 简单实现:分段落处理
paragraphs = prompt.split("\n\n")
results = []
for i, para in enumerate(paragraphs):
logger.info(f"处理段落 {i+1}/{len(paragraphs)}")
# 这里应该调用实际的 vLLM API
# 简化实现
response = GenerationResponse(
text=f"[Generated content for paragraph {i}]",
prompt=para,
num_tokens=len(para.split()),
finish_reason="stop",
truncation_detected=False,
metadata={"paragraph_index": i}
)
results.append(response)
return results
# 使用示例
async def main():
"""使用示例"""
handler = LongTextHandler(
api_base="http://localhost:8000/v1",
default_config=GenerationConfig(
max_tokens=4096,
max_model_len=8192,
truncation_strategy=TruncationStrategy.SMART_TRUNCATE,
stop_sequences=["\n\n", "---", "## 下一节"]
)
)
# 测试生成
response = await handler.generate_stream(
prompt="请详细介绍一下人工智能的发展历史,包括:\n1. 早期发展\n2. 机器学习时代\n3. 深度学习革命\n4. 大模型时代",
callback=lambda token: print(token, end="", flush=True)
)
print(f"\n\n统计信息:")
print(f" 生成 Token 数: {response.num_tokens}")
print(f" 是否截断: {response.truncation_detected}")
print(f" 结束原因: {response.finish_reason}")
if __name__ == "__main__":
asyncio.run(main())5.3 Prefix Caching 配置
# prefix-caching 配置示例
# 在 vLLM 启动参数中添加:
# --enable-prefix-caching
# --enable-chunked-prefill
# 或在 ConfigMap 中配置:
apiVersion: v1
kind: ConfigMap
metadata:
name: vllm-prefix-cache-config
data:
cache-config.yaml: |
# Prefix Caching 配置
prefix_caching:
enabled: true
cache_engine: "auto"
# Chunked Prefill 配置(可以更早开始生成)
chunked_prefill:
enabled: true
max_num_batched_tokens: 8192 start_chunk_size: 512
# 共享前缀示例
# system_prompt = "You are a helpful assistant."
#
# Request 1: system_prompt + "Hello"
# Request 2: system_prompt + "How are you?"
#
# 启用后,system_prompt 的 KV Cache 可以共享
# 节省约 30-50% 的显存5.4 并发控制
# 并发控制实现:rate_limiter.py
#!/usr/bin/env python3
"""
vLLM 并发控制与限流实现
解决高并发下的资源竞争问题
"""
import time
import asyncio
import logging
from typing import Dict, Optional, Set
from dataclasses import dataclass, field
from collections import defaultdict
from datetime import datetime, timedelta
import hashlib
logging.basicConfig(level=logging.INFO)
logger = logging.getLogger(__name__)
@dataclass
class RateLimitConfig:
"""限流配置"""
max_concurrent_requests: int = 100 # 最大并发请求数
max_requests_per_minute: int = 1000 # 每分钟最大请求数
max_tokens_per_minute: int = 100000 # 每分钟最大 token 数
queue_size: int = 500 # 队列大小
queue_timeout: int = 300 # 队列超时时间(秒)
@dataclass
class RequestRecord:
"""请求记录"""
request_id: str
user_id: str
tokens: int
timestamp: datetime = field(default_factory=datetime.now)
class TokenBucket:
"""令牌桶算法实现"""
def __init__(self, capacity: int, refill_rate: float):
"""
初始化令牌桶
Args:
capacity: 桶容量
refill_rate: 每秒补充速率
"""
self.capacity = capacity
self.refill_rate = refill_rate
self.tokens = capacity
self.last_refill = time.time()
self.lock = asyncio.Lock()
async def acquire(self, tokens: int = 1, timeout: float = 30) -> bool:
"""
获取令牌
Args:
tokens: 需要获取的令牌数
timeout: 超时时间
Returns:
是否成功获取
"""
start_time = time.time()
while True:
async with self.lock:
self._refill()
if self.tokens >= tokens:
self.tokens -= tokens
return True
# 检查超时
if time.time() - start_time >= timeout:
return False
# 等待后重试
await asyncio.sleep(0.1)
def _refill(self):
"""补充令牌"""
now = time.time()
elapsed = now - self.last_refill
new_tokens = elapsed * self.refill_rate
self.tokens = min(self.capacity, self.tokens + new_tokens)
self.last_refill = now
class RateLimiter:
"""
限流器
支持多种限流策略组合
"""
def __init__(self, config: RateLimitConfig):
"""
初始化限流器
Args:
config: 限流配置
"""
self.config = config
self.active_requests: Dict[str, RequestRecord] = {}
self.request_history: list = []
# 令牌桶
self.concurrent_bucket = TokenBucket(
capacity=config.max_concurrent_requests,
refill_rate=config.max_concurrent_requests
)
self.minute_bucket = TokenBucket(
capacity=config.max_requests_per_minute,
refill_rate=config.max_requests_per_minute / 60
)
self.token_bucket = TokenBucket(
capacity=config.max_tokens_per_minute,
refill_rate=config.max_tokens_per_minute / 60
)
# 用户级限流
self.user_buckets: Dict[str, TokenBucket] = defaultdict(
lambda: TokenBucket(100, 100)
)
# 用户历史记录
self.user_history: Dict[str, list] = defaultdict(list)
# 锁
self._lock = asyncio.Lock()
def _get_user_id(self, api_key: Optional[str] = None) -> str:
"""从 API Key 提取用户 ID"""
if api_key:
return hashlib.md5(api_key.encode()).hexdigest()[:8]
return "anonymous"
def _cleanup_old_records(self):
"""清理过期记录"""
cutoff = datetime.now() - timedelta(minutes=2)
self.request_history = [
r for r in self.request_history
if r.timestamp > cutoff
]
# 清理用户历史
for user_id in self.user_history:
self.user_history[user_id] = [
r for r in self.user_history[user_id]
if r.timestamp > cutoff
]
async def check_limit(
self,
user_id: Optional[str] = None,
estimated_tokens: int = 100,
api_key: Optional[str] = None
) -> tuple[bool, str]:
"""
检查是否允许请求
Args:
user_id: 用户 ID
estimated_tokens: 预估 token 数
api_key: API Key
Returns:
(是否允许, 拒绝原因)
"""
if user_id is None:
user_id = self._get_user_id(api_key)
# 清理过期记录
self._cleanup_old_records()
# 检查当前并发数
if len(self.active_requests) >= self.config.max_concurrent_requests:
return False, f"并发数超限 (当前: {len(self.active_requests)})"
# 检查队列大小
queue_wait = len(self.active_requests)
if queue_wait >= self.config.queue_size:
return False, f"队列已满 (当前等待: {queue_wait})"
# 尝试获取令牌
# 并发令牌
if not await self.concurrent_bucket.acquire(1, timeout=0):
return False, "并发限流"
# 分钟级请求数
if not await self.minute_bucket.acquire(1, timeout=0):
return False, "请求频率超限"
# Token 数限制
if not await self.token_bucket.acquire(estimated_tokens, timeout=0):
self.concurrent_bucket.tokens += 1 # 归还
self.minute_bucket.tokens += 1
return False, f"Token 数量超限 (预估: {estimated_tokens})"
# 用户级限流
user_bucket = self.user_buckets[user_id]
if not await user_bucket.acquire(1, timeout=0):
return False, "用户请求频率超限"
return True, ""
async def acquire(self, request: RequestRecord) -> bool:
"""
获取请求许可
Args:
request: 请求记录
Returns:
是否成功
"""
async with self._lock:
allowed, reason = await self.check_limit(
user_id=request.user_id,
estimated_tokens=request.tokens
)
if allowed:
self.active_requests[request.request_id] = request
return True
logger.warning(f"请求 {request.request_id} 被限流: {reason}")
return False
def release(self, request_id: str):
"""
释放请求
Args:
request_id: 请求 ID
"""
if request_id in self.active_requests:
request = self.active_requests.pop(request_id)
self.request_history.append(request)
logger.info(f"请求 {request_id} 完成,耗时: "
f"{(datetime.now() - request.timestamp).total_seconds():.2f}s")
def get_stats(self) -> Dict:
"""获取统计信息"""
return {
"active_requests": len(self.active_requests),
"total_requests": len(self.request_history),
"requests_in_last_minute": len([
r for r in self.request_history
if r.timestamp > datetime.now() - timedelta(minutes=1)
]),
"max_concurrent": self.config.max_concurrent_requests,
"concurrent_usage": len(self.active_requests) / self.config.max_concurrent_requests
}
# 中间件实现
class RateLimitMiddleware:
"""限流中间件"""
def __init__(self, rate_limiter: RateLimiter):
self.rate_limiter = rate_limiter
async def __call__(self, request, call_next):
"""处理请求"""
# 提取用户信息
api_key = request.headers.get("Authorization", "").replace("Bearer ", "")
user_id = self.rate_limiter._get_user_id(api_key)
# 预估 token 数
prompt = request.json().get("prompt", "")
estimated_tokens = len(prompt.split()) * 1.3 + \
request.json().get("max_tokens", 100)
# 创建请求记录
request_record = RequestRecord(
request_id=request.id,
user_id=user_id,
tokens=int(estimated_tokens)
)
# 尝试获取许可
if not await self.rate_limiter.acquire(request_record):
return Response(
status_code=429,
content={"error": "Rate limit exceeded"}
)
try:
response = await call_next(request)
return response
finally:
self.rate_limiter.release(request.request_id)
# 使用示例
async def example():
"""使用示例"""
config = RateLimitConfig(
max_concurrent_requests=50,
max_requests_per_minute=500,
max_tokens_per_minute=50000
)
limiter = RateLimiter(config)
# 模拟请求
for i in range(60):
record = RequestRecord(
request_id=f"req-{i}",
user_id=f"user-{i % 10}",
tokens=100
)
allowed = await limiter.acquire(record)
print(f"Request {i}: {'✓' if allowed else '✗'}")
if allowed and i % 5 == 0:
# 模拟完成
limiter.release(record.request_id)
# 输出统计
print("\n统计信息:")
stats = limiter.get_stats()
for key, value in stats.items():
print(f" {key}: {value}")
if __name__ == "__main__":
asyncio.run(example())六、性能调优参数详解
6.1 核心参数表
6.2 量化部署(降低显存占用)
量化可以显著降低显存占用,同时保持较高的推理精度。
AWQ 量化(推荐)
# 安装 AutoAWQ
pip install autoawq
# 量化脚本
python << 'EOF'
from awq import AutoAWQForCausalLM
from transformers import AutoTokenizer
model_path = '/data/models/llama-2-7b-chat'
quant_path = '/data/models/llama-2-7b-chat-awq'
# 加载模型
model = AutoAWQForCausalLM.from_pretrained(model_path)
tokenizer = AutoTokenizer.from_pretrained(model_path)
# 量化配置
quant_config = {
'zero_point': True,
'q_group_size': 128,
'w_bit': 4,
'version': 'GEMM'
}
# 执行量化
model.quantize(tokenizer, quant_config=quant_config)
# 保存量化模型
model.save_quantized(quant_path)
tokenizer.save_pretrained(quant_path)
print(f'✅ 量化完成,模型保存至: {quant_path}')
EOFvLLM 加载量化模型
# 启动 vLLM 时指定量化方法
python -m vllm.entrypoints.openai.api_server \
--model /data/models/llama-2-7b-chat-awq \
--quantization awq \
--dtype auto \
--tensor-parallel-size 2 \
--gpu-memory-utilization 0.95量化效果对比
6.3 多模型部署(模型路由)
在生产环境中,通常需要同时部署多个模型,根据请求复杂度路由到不同模型。
K8s 多模型部署架构
# k8s/multi-model-deployment.yaml
apiVersion: apps/v1
kind: Deployment
metadata:
name: vllm-multi-model
namespace: vllm-inference
spec:
replicas: 1
selector:
matchLabels:
app: vllm-multi
template:
metadata:
labels:
app: vllm-multi
spec:
containers:
# 小模型:快速响应(简单任务)
- name: vllm-small
image: vllm/vllm-openai:v0.4.3
command:
- python
- -m
- vllm.entrypoints.openai.api_server
args:
- --model
- /models/qwen-1.5b
- --host
- "0.0.0.0"
- --port
- "8000"
- --max-model-len
- "4096"
- --gpu-memory-utilization
- "0.5"
ports:
- name: small-api
containerPort: 8000
resources:
limits:
nvidia.com/gpu: 1
volumeMounts:
- name: models
mountPath: /models
readOnly: true
# 大模型:复杂任务
- name: vllm-large
image: vllm/vllm-openai:v0.4.3
command:
- python
- -m
- vllm.entrypoints.openai.api_server
args:
- --model
- /models/qwen-72b
- --host
- "0.0.0.0"
- --port
- "8001"
- --tensor-parallel-size
- "4"
- --max-model-len
- "8192"
- --gpu-memory-utilization
- "0.92"
ports:
- name: large-api
containerPort: 8001
resources:
limits:
nvidia.com/gpu: 4
volumeMounts:
- name: models
mountPath: /models
readOnly: true
# 路由器(基于请求复杂度路由)
- name: model-router
image: nginx:alpine
ports:
- name: http
containerPort: 80
volumeMounts:
- name: nginx-config
mountPath: /etc/nginx/nginx.conf
subPath: nginx.conf
volumes:
- name: models
persistentVolumeClaim:
claimName: model-storage
- name: nginx-config
configMap:
name: model-router-config
---
# 路由配置
apiVersion: v1
kind: ConfigMap
metadata:
name: model-router-config
namespace: vllm-inference
data:
nginx.conf: |
events {
worker_connections 1024;
}
http {
upstream small_model {
server localhost:8000;
}
upstream large_model {
server localhost:8001;
}
# 日志格式
log_format main '$remote_addr - $remote_user [$time_local] "$request" '
'$status $body_bytes_sent "$http_referer" '
'"$http_user_agent" "$http_x_forwarded_for" '
'model=$upstream_addr';
access_log /var/log/nginx/access.log main;
server {
listen 80;
location /v1/chat/completions {
# 根据请求头路由
set $backend "small_model";
# 检查 X-Model-Size 头
if ($http_x_model_size = "large") {
set $backend "large_model";
}
# 检查 max_tokens(大于 1000 使用大模型)
if ($request_body ~ "max_tokens\":\s*([0-9]+)") {
set $max_tokens $1;
}
proxy_pass http://$backend;
proxy_set_header Host $host;
proxy_set_header X-Real-IP $remote_addr;
proxy_set_header X-Forwarded-For $proxy_add_x_forwarded_for;
proxy_read_timeout 300s;
proxy_connect_timeout 10s;
}
location /health {
return 200 "OK\n";
add_header Content-Type text/plain;
}
}
}
---
# Service
apiVersion: v1
kind: Service
metadata:
name: vllm-multi-model
namespace: vllm-inference
spec:
type: ClusterIP
ports:
- name: http
port: 80
targetPort: 80
selector:
app: vllm-multi智能路由 Python 实现
#!/usr/bin/env python3
"""
智能模型路由器
根据请求复杂度自动选择合适的模型
"""
from fastapi import FastAPI, Request, HTTPException
from fastapi.responses import StreamingResponse
import httpx
import asyncio
from typing import Dict, Optional
import logging
logging.basicConfig(level=logging.INFO)
logger = logging.getLogger(__name__)
app = FastAPI()
# 模型端点配置
MODELS = {
"small": {
"url": "http://localhost:8000/v1/chat/completions",
"max_tokens": 2048,
"cost_per_1k": 0.001
},
"medium": {
"url": "http://localhost:8001/v1/chat/completions",
"max_tokens": 4096,
"cost_per_1k": 0.003
},
"large": {
"url": "http://localhost:8002/v1/chat/completions",
"max_tokens": 8192,
"cost_per_1k": 0.010
}
}
def estimate_complexity(messages: list, max_tokens: int) -> str:
"""
估算请求复杂度
Args:
messages: 对话消息列表
max_tokens: 最大生成 token 数
Returns:
模型大小:small/medium/large
"""
# 计算输入长度
total_input_length = sum(len(msg.get("content", "")) for msg in messages)
# 规则 1: 根据输入长度
if total_input_length > 2000:
return "large"
# 规则 2: 根据输出长度
if max_tokens > 1000:
return "large"
elif max_tokens > 500:
return "medium"
# 规则 3: 根据消息轮数(多轮对话用大模型)
if len(messages) > 10:
return "medium"
# 规则 4: 检查是否包含代码(代码生成用大模型)
for msg in messages:
content = msg.get("content", "")
if "```" in content or "def " in content or "class " in content:
return "large"
# 默认使用小模型
return "small"
@app.post("/v1/chat/completions")
async def chat_completions(request: Request):
"""
智能路由的对话接口
"""
try:
# 解析请求
body = await request.json()
messages = body.get("messages", [])
max_tokens = body.get("max_tokens", 256)
# 用户可以通过 model 参数强制指定模型
user_model = body.get("model", "")
if user_model in MODELS:
selected_model = user_model
else:
# 自动选择模型
selected_model = estimate_complexity(messages, max_tokens)
model_config = MODELS[selected_model]
# 检查是否超出模型限制
if max_tokens > model_config["max_tokens"]:
# 自动升级到更大的模型
for model_name, config in MODELS.items():
if config["max_tokens"] >= max_tokens:
selected_model = model_name
model_config = config
break
logger.info(f"路由到模型: {selected_model}, max_tokens: {max_tokens}")
# 转发请求
async with httpx.AsyncClient(timeout=300.0) as client:
response = await client.post(
model_config["url"],
json=body,
headers={"Content-Type": "application/json"}
)
# 添加路由信息到响应头
headers = dict(response.headers)
headers["X-Selected-Model"] = selected_model
headers["X-Model-Cost-Per-1k"] = str(model_config["cost_per_1k"])
return StreamingResponse(
response.aiter_bytes(),
status_code=response.status_code,
headers=headers,
media_type="application/json"
)
except Exception as e:
logger.error(f"路由失败: {e}")
raise HTTPException(status_code=500, detail=str(e))
@app.get("/health")
async def health():
"""健康检查"""
return {"status": "ok"}
if __name__ == "__main__":
import uvicorn
uvicorn.run(app, host="0.0.0.0", port=8080)6.4 不同场景配置
# 场景1: 高吞吐短文本 (客服场景)
tensor-parallel-size: 4
max-model-len: 2048
max-num-seqs: 256
max-num-batched-tokens: 16384
gpu-memory-utilization: 0.95
enable-prefix-caching: true
enable-chunked-prefill: true
---
# 场景2: 长上下文 (文档分析)
tensor-parallel-size: 4
max-model-len: 16384
max-num-seqs: 64
max-num-batched-tokens: 8192
gpu-memory-utilization: 0.85
enable-prefix-caching: true
enable-chunked-prefill: true
---
# 场景3: 低延迟实时响应
tensor-parallel-size: 2
max-model-len: 4096
max-num-seqs: 32
max-num-batched-tokens: 4096
gpu-memory-utilization: 0.9
enable-prefix-caching: false # 避免缓存污染
enable-chunked-prefill: true七、健康检查与自动恢复
7.1 健康检查脚本
#!/bin/bash
# health_check.sh - vLLM 健康检查脚本
VLLM_HOST="${VLLM_HOST:-localhost}"
VLLM_PORT="${VLLM_PORT:-8000}"
TIMEOUT="${TIMEOUT:-10}"
# 检查 API 是否响应
response=$(curl -s -o /dev/null -w "%{http_code}" \
--max-time $TIMEOUT \
"http://${VLLM_HOST}:${VLLM_PORT}/health" 2>/dev/null)
if [ "$response" = "200" ]; then
echo "✓ vLLM 健康检查通过"
exit 0
else
echo "✗ vLLM 健康检查失败 (HTTP $response)"
exit 1
fi7.2 自动恢复机制
# K8s Pod 重启策略
spec:
restartPolicy: Always
# 使用 LivenessProbe 检测进程崩溃
containers:
- name: vllm
livenessProbe:
httpGet:
path: /health
port: 8000
initialDelaySeconds: 300
periodSeconds: 60
timeoutSeconds: 30
failureThreshold: 5 # 连续5次失败后重启
# ReadinessProbe 检测服务不可用
readinessProbe:
httpGet:
path: /health
port: 8000
initialDelaySeconds: 300
periodSeconds: 30
failureThreshold: 37.3 监控告警规则
# config/prometheus/alert_rules.yml
groups:
- name: vllm_alerts
interval: 30s
rules:
# GPU 显存告警
- alert: GPUMemoryHigh
expr: |
(nvidia_gpu_memory_used_bytes / nvidia_gpu_memory_total_bytes) > 0.95
for: 5m
labels:
severity: warning
annotations:
summary: "GPU 显存使用率过高"
description: "GPU {{ $labels.gpu }} 显存使用率 {{ $value | humanizePercentage }}"
# 请求延迟告警
- alert: HighLatency
expr: |
histogram_quantile(0.95,
rate(vllm_request_duration_seconds_bucket[5m])
) > 5
for: 10m
labels:
severity: warning
annotations:
summary: "P95 延迟过高"
description: "P95 延迟 {{ $value }}s,超过 5s 阈值"
# 错误率告警
- alert: HighErrorRate
expr: |
rate(vllm_request_total{status="error"}[5m]) /
rate(vllm_request_total[5m]) > 0.05
for: 5m
labels:
severity: critical
annotations:
summary: "错误率过高"
description: "错误率 {{ $value | humanizePercentage }},超过 5%"
# OOM 预警
- alert: OOMRisk
expr: |
(nvidia_gpu_memory_used_bytes / nvidia_gpu_memory_total_bytes) > 0.90
for: 2m
labels:
severity: warning
annotations:
summary: "OOM 风险"
description: "GPU 显存使用率 {{ $value | humanizePercentage }},接近上限"
# 服务不可用
- alert: vLLMDown
expr: up{job="vllm"} == 0
for: 1m
labels:
severity: critical
annotations:
summary: "vLLM 服务不可用"
description: "vLLM 实例 {{ $labels.instance }} 已停止响应"
# 队列积压告警
- alert: RequestQueueBacklog
expr: vllm_queue_size > 100
for: 5m
labels:
severity: warning
annotations:
summary: "请求队列积压"
description: "当前队列大小 {{ $value }},可能需要扩容"
# Token 使用量告警
- alert: HighTokenUsage
expr: |
rate(vllm_token_usage_total[1h]) > 1000000
for: 10m
labels:
severity: info
annotations:
summary: "Token 使用量过高"
description: "过去 1 小时 Token 使用量 {{ $value }},请注意成本"AlertManager 配置(企业微信/飞书通知)
# config/alertmanager/alertmanager.yml
global:
resolve_timeout: 5m
route:
group_by: ['alertname', 'cluster', 'service']
group_wait: 10s
group_interval: 10s
repeat_interval: 12h
receiver: 'wechat'
routes:
- match:
severity: critical
receiver: 'wechat-critical'
continue: true
receivers:
# 企业微信通知
- name: 'wechat'
wechat_configs:
- corp_id: 'your-corp-id'
to_party: '1'
agent_id: 'your-agent-id'
api_secret: 'your-api-secret'
send_resolved: true
message: |
{{ range .Alerts }}
告警: {{ .Labels.alertname }}
级别: {{ .Labels.severity }}
实例: {{ .Labels.instance }}
描述: {{ .Annotations.description }}
时间: {{ .StartsAt.Format "2006-01-02 15:04:05" }}
{{ end }}
# 飞书通知
- name: 'feishu'
webhook_configs:
- url: 'https://open.feishu.cn/open-apis/bot/v2/hook/your-webhook-token'
send_resolved: true八、性能测试与压测
8.1 压测工具实现
#!/usr/bin/env python3
"""
vLLM 性能压测工具
支持并发测试、延迟统计、吞吐量测试
"""
import asyncio
import time
import statistics
import argparse
from typing import List, Dict, Tuple
from dataclasses import dataclass, field
import aiohttp
import json
@dataclass
class BenchmarkResult:
"""压测结果"""
total_requests: int
successful_requests: int
failed_requests: int
total_time: float
latencies: List[float] = field(default_factory=list)
tokens_generated: List[int] = field(default_factory=list)
@property
def qps(self) -> float:
"""每秒请求数"""
return self.successful_requests / self.total_time if self.total_time > 0 else 0
@property
def avg_latency(self) -> float:
"""平均延迟"""
return statistics.mean(self.latencies) if self.latencies else 0
@property
def p50_latency(self) -> float:
"""P50 延迟"""
return statistics.median(self.latencies) if self.latencies else 0
@property
def p95_latency(self) -> float:
"""P95 延迟"""
if not self.latencies:
return 0
sorted_latencies = sorted(self.latencies)
index = int(len(sorted_latencies) * 0.95)
return sorted_latencies[index]
@property
def p99_latency(self) -> float:
"""P99 延迟"""
if not self.latencies:
return 0
sorted_latencies = sorted(self.latencies)
index = int(len(sorted_latencies) * 0.99)
return sorted_latencies[index]
@property
def avg_tokens(self) -> float:
"""平均生成 token 数"""
return statistics.mean(self.tokens_generated) if self.tokens_generated else 0
@property
def tokens_per_second(self) -> float:
"""每秒生成 token 数"""
total_tokens = sum(self.tokens_generated)
return total_tokens / self.total_time if self.total_time > 0 else 0
class vLLMBenchmark:
"""vLLM 压测工具"""
def __init__(self, base_url: str, concurrency: int = 10):
self.base_url = base_url.rstrip("/")
self.concurrency = concurrency
self.semaphore = asyncio.Semaphore(concurrency)
async def single_request(
self,
session: aiohttp.ClientSession,
prompt: str,
max_tokens: int = 256
) -> Tuple[bool, float, int]:
"""
单次请求
Returns:
(是否成功, 延迟, 生成token数)
"""
async with self.semaphore:
start = time.time()
try:
async with session.post(
f"{self.base_url}/v1/completions",
json={
"prompt": prompt,
"max_tokens": max_tokens,
"temperature": 0.7
},
timeout=aiohttp.ClientTimeout(total=60)
) as response:
data = await response.json()
latency = time.time() - start
# 提取生成的 token 数
tokens = 0
if response.status == 200:
usage = data.get("usage", {})
tokens = usage.get("completion_tokens", 0)
return response.status == 200, latency, tokens
except Exception as e:
latency = time.time() - start
print(f"请求失败: {e}")
return False, latency, 0
async def run_benchmark(
self,
num_requests: int,
prompt: str = "请详细介绍一下人工智能的发展历史。",
max_tokens: int = 256
) -> BenchmarkResult:
"""运行压测"""
print(f"开始压测: {num_requests} 个请求,并发数 {self.concurrency}")
print(f"Prompt: {prompt[:50]}...")
print(f"Max tokens: {max_tokens}")
print("-" * 60)
result = BenchmarkResult(
total_requests=num_requests,
successful_requests=0,
failed_requests=0,
total_time=0
)
start_time = time.time()
async with aiohttp.ClientSession() as session:
tasks = [
self.single_request(session, prompt, max_tokens)
for _ in range(num_requests)
]
# 显示进度
completed = 0
for coro in asyncio.as_completed(tasks):
resp = await coro
completed += 1
if completed % 10 == 0 or completed == num_requests:
print(f"进度: {completed}/{num_requests} ({completed/num_requests*100:.1f}%)")
if isinstance(resp, tuple):
success, latency, tokens = resp
if success:
result.successful_requests += 1
result.tokens_generated.append(tokens)
else:
result.failed_requests += 1
result.latencies.append(latency)
else:
result.failed_requests += 1
result.total_time = time.time() - start_time
return result
def print_result(self, result: BenchmarkResult):
"""打印结果"""
print("\n" + "=" * 60)
print("压测结果")
print("=" * 60)
print(f"总请求数: {result.total_requests}")
print(f"成功请求数: {result.successful_requests}")
print(f"失败请求数: {result.failed_requests}")
print(f"成功率: {result.successful_requests/result.total_requests*100:.2f}%")
print(f"总耗时: {result.total_time:.2f}s")
print("-" * 60)
print(f"QPS: {result.qps:.2f} req/s")
print(f"平均延迟: {result.avg_latency:.3f}s")
print(f"P50 延迟: {result.p50_latency:.3f}s")
print(f"P95 延迟: {result.p95_latency:.3f}s")
print(f"P99 延迟: {result.p99_latency:.3f}s")
print("-" * 60)
print(f"平均生成 tokens: {result.avg_tokens:.1f}")
print(f"总生成 tokens: {sum(result.tokens_generated)}")
print(f"Tokens/秒: {result.tokens_per_second:.1f}")
print("=" * 60)
async def main():
parser = argparse.ArgumentParser(description="vLLM 性能压测工具")
parser.add_argument("--url", default="http://localhost:8000", help="vLLM 服务地址")
parser.add_argument("--requests", "-n", type=int, default=100, help="总请求数")
parser.add_argument("--concurrency", "-c", type=int, default=10, help="并发数")
parser.add_argument("--max-tokens", type=int, default=256, help="最大生成 token 数")
parser.add_argument("--prompt", type=str, default="请详细介绍一下人工智能的发展历史。", help="测试 Prompt")
args = parser.parse_args()
benchmark = vLLMBenchmark(args.url, args.concurrency)
result = await benchmark.run_benchmark(
args.requests,
prompt=args.prompt,
max_tokens=args.max_tokens
)
benchmark.print_result(result)
if __name__ == "__main__":
asyncio.run(main())8.2 使用示例
# 基础压测
python benchmark.py --url http://localhost:8000 -n 1000 -c 50
# 高并发压测
python benchmark.py -n 5000 -c 200
# 长文本压测
python benchmark.py -n 100 -c 10 --max-tokens 2048
# 自定义 Prompt 压测
python benchmark.py -n 500 -c 50 --prompt "写一篇关于机器学习的文章"8.3 压测结果分析
示例输出
开始压测: 1000 个请求,并发数 50
Prompt: 请详细介绍一下人工智能的发展历史。...
Max tokens: 256
------------------------------------------------------------
进度: 10/1000 (1.0%)
进度: 20/1000 (2.0%)
...
进度: 1000/1000 (100.0%)
============================================================
压测结果
============================================================
总请求数: 1000
成功请求数: 998
失败请求数: 2
成功率: 99.80%
总耗时: 45.23s
------------------------------------------------------------
QPS: 22.07 req/s
平均延迟: 2.265s
P50 延迟: 2.180s
P95 延迟: 3.450s
P99 延迟: 4.120s
------------------------------------------------------------
平均生成 tokens: 245.3
总生成 tokens: 244834
Tokens/秒: 5412.8
============================================================性能指标解读
九、Python 客户端封装
9.1 完整客户端实现
#!/usr/bin/env python3
"""
vLLM Python 客户端封装
提供易用的推理接口,支持同步/异步调用
"""
import os
import time
import json
import asyncio
import logging
from typing import Dict, List, Optional, Union, Any, Callable
from dataclasses import dataclass, field
from enum import Enum
from functools import wraps
import hashlib
import requests
import aiohttp
from tenacity import retry, stop_after_attempt, wait_exponential
# 配置日志
logging.basicConfig(level=logging.INFO)
logger = logging.getLogger(__name__)
class Role(Enum):
"""对话角色"""
SYSTEM = "system"
USER = "user"
ASSISTANT = "assistant"
@dataclass
class ChatMessage:
"""对话消息"""
role: Union[Role, str]
content: str
def to_dict(self) -> Dict:
role = self.role.value if isinstance(self.role, Role) else self.role
return {"role": role, "content": self.content}
@classmethod
def from_dict(cls, data: Dict) -> "ChatMessage":
return cls(role=data["role"], content=data["content"])
@dataclass
class CompletionRequest:
"""补全请求"""
prompt: str
model: Optional[str] = None
max_tokens: int = 256
temperature: float = 0.7
top_p: float = 0.95
top_k: int = 50
frequency_penalty: float = 0.0
presence_penalty: float = 0.0
stop: Optional[Union[str, List[str]]] = None
stream: bool = False
echo: bool = False
def to_payload(self) -> Dict:
"""转换为 API payload"""
payload = {
"prompt": self.prompt,
"max_tokens": self.max_tokens,
"temperature": self.temperature,
"top_p": self.top_p,
"top_k": self.top_k,
"frequency_penalty": self.frequency_penalty,
"presence_penalty": self.presence_penalty,
"stream": self.stream,
"echo": self.echo,
}
if self.stop:
payload["stop"] = self.stop
return payload
@dataclass
class ChatRequest:
"""对话请求"""
messages: List[ChatMessage]
model: Optional[str] = None
max_tokens: int = 256
temperature: float = 0.7
top_p: float = 0.95
frequency_penalty: float = 0.0
presence_penalty: float = 0.0
stop: Optional[Union[str, List[str]]] = None
stream: bool = False
def to_payload(self) -> Dict:
"""转换为 API payload"""
payload = {
"messages": [m.to_dict() for m in self.messages],
"max_tokens": self.max_tokens,
"temperature": self.temperature,
"top_p": self.top_p,
"frequency_penalty": self.frequency_penalty,
"presence_penalty": self.presence_penalty,
"stream": self.stream,
}
if self.stop:
payload["stop"] = self.stop
return payload
@dataclass
class UsageInfo:
"""用量信息"""
prompt_tokens: int
completion_tokens: int
total_tokens: int
@classmethod
def from_dict(cls, data: Dict) -> "UsageInfo":
return cls(
prompt_tokens=data.get("prompt_tokens", 0),
completion_tokens=data.get("completion_tokens", 0),
total_tokens=data.get("total_tokens", 0)
)
@dataclass
class CompletionResponse:
"""补全响应"""
id: str
text: str
finish_reason: str
usage: UsageInfo
latency_ms: float
raw_response: Dict = field(default_factory=dict)
@dataclass
class ChatResponse:
"""对话响应"""
id: str
message: ChatMessage
finish_reason: str
usage: UsageInfo
latency_ms: float
raw_response: Dict = field(default_factory=dict)
@property
def content(self) -> str:
"""快捷获取内容"""
return self.message.content
def retry_on_rate_limit(func):
"""速率限制重试装饰器"""
@wraps(func)
@retry(
stop=stop_after_attempt(3),
wait=wait_exponential(multiplier=1, min=2, max=10)
)
async def wrapper(*args, **kwargs):
try:
return await func(*args, **kwargs)
except aiohttp.ClientResponseError as e:
if e.status == 429:
logger.warning("遇到速率限制,等待后重试...")
raise
raise
return wrapper
class vLLMClient:
"""
vLLM Python 客户端
支持同步/异步调用,自动重试,错误处理
"""
def __init__(
self,
base_url: str = "http://localhost:8000/v1",
api_key: Optional[str] = None,
timeout: int = 300,
max_retries: int = 3,
default_model: Optional[str] = None
):
"""
初始化客户端
Args:
base_url: vLLM 服务地址
api_key: API 密钥(可选)
timeout: 请求超时时间(秒)
max_retries: 最大重试次数
default_model: 默认模型名称
"""
self.base_url = base_url.rstrip("/")
self.timeout = timeout
self.default_model = default_model
# 设置请求头
self.headers = {"Content-Type": "application/json"}
if api_key:
self.headers["Authorization"] = f"Bearer {api_key}"
# 同步会话
self._session: Optional[requests.Session] = None
# 统计信息
self._stats = {
"total_requests": 0,
"failed_requests": 0,
"total_tokens": 0,
"total_latency_ms": 0
}
def _get_session(self) -> requests.Session:
"""获取或创建同步会话"""
if self._session is None:
self._session = requests.Session()
self._session.headers.update(self.headers)
self._session.timeout = self.timeout
return self._session
@property
def completions_url(self) -> str:
"""补全 API URL"""
return f"{self.base_url}/completions"
@property
def chat_url(self) -> str:
"""对话 API URL"""
return f"{self.base_url}/chat/completions"
@property
def models_url(self) -> str:
"""模型列表 URL"""
return f"{self.base_url}/models"
def list_models(self) -> List[Dict]:
"""
获取可用模型列表
Returns:
模型列表
"""
session = self._get_session()
response = session.get(self.models_url)
response.raise_for_status()
data = response.json()
return data.get("data", [])
def complete(self, request: CompletionRequest) -> CompletionResponse:
"""
同步补全请求
Args:
request: 补全请求
Returns:
补全响应
"""
start_time = time.time()
try:
session = self._get_session()
payload = request.to_payload()
if request.model or self.default_model:
payload["model"] = request.model or self.default_model
response = session.post(self.completions_url, json=payload)
response.raise_for_status()
data = response.json()
latency_ms = (time.time() - start_time) * 1000
# 更新统计
self._update_stats(data, latency_ms)
return CompletionResponse(
id=data.get("id", ""),
text=data["choices"][0]["text"],
finish_reason=data["choices"][0].get("finish_reason", ""),
usage=UsageInfo.from_dict(data.get("usage", {})),
latency_ms=latency_ms,
raw_response=data
)
except requests.RequestException as e:
self._stats["failed_requests"] += 1
logger.error(f"补全请求失败: {e}")
raise
def chat(self, request: ChatRequest) -> ChatResponse:
"""
同步对话请求
Args:
request: 对话请求
Returns:
对话响应
"""
start_time = time.time()
try:
session = self._get_session()
payload = request.to_payload()
if request.model or self.default_model:
payload["model"] = request.model or self.default_model
response = session.post(self.chat_url, json=payload)
response.raise_for_status()
data = response.json()
latency_ms = (time.time() - start_time) * 1000
# 更新统计
self._update_stats(data, latency_ms)
choice = data["choices"][0]
return ChatResponse(
id=data.get("id", ""),
message=ChatMessage.from_dict(choice["message"]),
finish_reason=choice.get("finish_reason", ""),
usage=UsageInfo.from_dict(data.get("usage", {})),
latency_ms=latency_ms,
raw_response=data
)
except requests.RequestException as e:
self._stats["failed_requests"] += 1
logger.error(f"对话请求失败: {e}")
raise
def complete_stream(
self,
request: CompletionRequest,
callback: Optional[Callable[[str], None]] = None
) -> CompletionResponse:
"""
流式补全请求
Args:
request: 补全请求
callback: 回调函数,处理每个 token
Returns:
完整响应
"""
request.stream = True
session = self._get_session()
payload = request.to_payload()
if request.model or self.default_model:
payload["model"] = request.model or self.default_model
response = session.post(
self.completions_url,
json=payload,
stream=True
)
response.raise_for_status()
full_text = []
start_time = time.time()
for line in response.iter_lines():
line = line.decode('utf-8').strip()
if not line or not line.startswith('data: '):
continue
data = line[6:]
if data == "[DONE]":
break
try:
chunk = json.loads(data)
token = chunk["choices"][0].get("text", "")
full_text.append(token)
if callback:
callback(token)
except json.JSONDecodeError:
continue
latency_ms = (time.time() - start_time) * 1000
return CompletionResponse(
id="stream",
text="".join(full_text),
finish_reason="stop",
usage=UsageInfo(0, len(full_text), len(full_text)),
latency_ms=latency_ms
)
def _update_stats(self, data: Dict, latency_ms: float):
"""更新统计信息"""
self._stats["total_requests"] += 1
self._stats["total_latency_ms"] += latency_ms
usage = data.get("usage", {})
self._stats["total_tokens"] += usage.get("total_tokens", 0)
def get_stats(self) -> Dict:
"""获取统计信息"""
stats = self._stats.copy()
if stats["total_requests"] > 0:
stats["avg_latency_ms"] = stats["total_latency_ms"] / stats["total_requests"]
stats["success_rate"] = (
(stats["total_requests"] - stats["failed_requests"]) /
stats["total_requests"] * 100
)
return stats
def close(self):
"""关闭会话"""
if self._session:
self._session.close()
self._session = None
def __enter__(self):
return self
def __exit__(self, exc_type, exc_val, exc_tb):
self.close()
class AsyncvLLMClient:
"""
vLLM 异步客户端
支持高并发调用
"""
def __init__(
self,
base_url: str = "http://localhost:8000/v1",
api_key: Optional[str] = None,
timeout: int = 300,
max_concurrent: int = 10,
default_model: Optional[str] = None
):
"""
初始化异步客户端
Args:
base_url: vLLM 服务地址
api_key: API 密钥
timeout: 请求超时时间
max_concurrent: 最大并发数
default_model: 默认模型
"""
self.base_url = base_url.rstrip("/")
self.timeout = aiohttp.ClientTimeout(total=timeout)
self.default_model = default_model
self.max_concurrent = max_concurrent
# Semaphore 控制并发
self._semaphore: Optional[asyncio.Semaphore] = None
# Headers
self.headers = {"Content-Type": "application/json"}
if api_key:
self.headers["Authorization"] = f"Bearer {api_key}"
# Session
self._session: Optional[aiohttp.ClientSession] = None
async def _get_session(self) -> aiohttp.ClientSession:
"""获取或创建异步会话"""
if self._session is None or self._session.closed:
self._session = aiohttp.ClientSession(
headers=self.headers,
timeout=self.timeout
)
self._semaphore = asyncio.Semaphore(self.max_concurrent)
return self._session
async def close(self):
"""关闭会话"""
if self._session and not self._session.closed:
await self._session.close()
self._session = None
@retry_on_rate_limit
async def chat(self, request: ChatRequest) -> ChatResponse:
"""
异步对话请求
Args:
request: 对话请求
Returns:
对话响应
"""
session = await self._get_session()
payload = request.to_payload()
if request.model or self.default_model:
payload["model"] = request.model or self.default_model
async with self._semaphore:
start_time = time.time()
async with session.post(
f"{self.base_url}/chat/completions",
json=payload
) as response:
response.raise_for_status()
data = await response.json()
latency_ms = (time.time() - start_time) * 1000
choice = data["choices"][0]
return ChatResponse(
id=data.get("id", ""),
message=ChatMessage.from_dict(choice["message"]),
finish_reason=choice.get("finish_reason", ""),
usage=UsageInfo.from_dict(data.get("usage", {})),
latency_ms=latency_ms,
raw_response=data
)
async def batch_chat(
self,
requests: List[ChatRequest]
) -> List[ChatResponse]:
"""
批量异步对话
Args:
requests: 请求列表
Returns:
响应列表
"""
tasks = [self.chat(req) for req in requests]
return await asyncio.gather(*tasks, return_exceptions=True)
# 使用示例
def main():
"""使用示例"""
# 同步客户端
with vLLMClient(
base_url="http://localhost:8000/v1",
default_model="meta-llama/Llama-2-7b-chat-hf"
) as client:
# 获取模型列表
models = client.list_models()
print(f"可用模型: {[m['id'] for m in models]}")
# 简单对话
response = client.chat(ChatRequest(
messages=[
ChatMessage(role=Role.SYSTEM, content="你是一个有帮助的助手。"),
ChatMessage(role=Role.USER, content="你好,请介绍一下你自己。")
],
max_tokens=256,
temperature=0.7
))
print(f"回复: {response.content}")
print(f"Token 使用: {response.usage.total_tokens}")
print(f"延迟: {response.latency_ms:.2f}ms")
# 统计信息
stats = client.get_stats()
print(f"统计: {stats}")
# 异步使用示例
async def async_example():
"""异步使用示例"""
async with AsyncvLLMClient(
base_url="http://localhost:8000/v1"
) as client:
# 并发请求
requests = [
ChatRequest(
messages=[ChatMessage(role=Role.USER, content=f"第{i}个问题")],
max_tokens=128
)
for i in range(10)
]
responses = await client.batch_chat(requests)
for i, resp in enumerate(responses):
if isinstance(resp, ChatResponse):
print(f"Response {i}: {resp.content[:50]}...")
if __name__ == "__main__":
main()
# asyncio.run(async_example())十、总结
本文系统性地介绍了 vLLM 的核心原理、Docker/Kubernetes 部署配置、生产环境踩坑经验以及完整的 Python 客户端实现。
以下是关键要点回顾:
核心要点
PagedAttention:通过虚拟内存思想管理 KV Cache,将显存利用率提升 2-4 倍
Continuous Batching:动态批次调度,消除 GPU 空闲等待,吞吐量提升 2-10 倍
并行策略:Tensor Parallel 适合单节点多卡,Pipeline Parallel 适合多节点
OOM 预防:合理设置
gpu-memory-utilization、max-model-len、max-num-seqsPrefix Caching:共享系统提示词,减少重复计算
健康检查:多层次探针(Readiness/Liveness),配合优雅关闭实现零停机部署
最佳实践
下一步建议
集成 Prometheus + Grafana 监控实时性能
配置 Multi-Instance 部署实现水平扩展
尝试 AWQ/Q4_K_M 量化进一步降低显存占用
探索 Speculative Decoding 降低延迟
参考文档:
评论区