AI Engineering Deployment
You ran a model on your laptop, and it generated responses quickly and accurately. But when you put it on a server and let a hundred users call it simultaneously, the situation changes:Some users wait ten seconds for a reply, some requests time out and report errors, GPU memory fills up in no time, and the bill climbs to a painful level.
This is the problem that AI deployment needs to solve:Turn a runnable model into a usable service。
The bottleneck for traditional web services is usually the CPU and the database.
The bottleneck for AI services is mainly the GPU — whether memory is sufficient, whether computation is fast, and how concurrent requests are queued.
The core challenges of AI deployment can be summarized in three words: latency, throughput, and cost.
| Metric | Meaning | User experience |
|---|---|---|
| Latency | Time from when a user sends a request to receiving the first character | Fast or not |
| Throughput | How many requests can be processed per second | Whether it can serve many people at the same time |
| Cost | GPU/server cost for running the service | Expensive or not |
Good AI deployment is about finding a balance among these three: low enough latency, high enough throughput, and controllable cost.
Model Serving Framework
To wrap a trained model as an API service, you need a dedicated framework. This section introduces three mainstream options: vLLM, TGI, and Ollama.
vLLM: High-Performance Inference Driven by PagedAttention
vLLM is an inference engine developed by the UC Berkeley team, and its biggest feature is speed.
vLLM's core innovation isPagedAttention— a technique for efficiently managing GPU memory, inspired by the virtual memory paging of operating systems.
In traditional inference frameworks, each request's KV Cache (key-value cache) needs to occupy contiguous GPU memory space.
When requests have varying lengths, GPU memory becomes fragmented, resulting in low utilization.
PagedAttention divides GPU memory into fixed-size pages. Each request's KV Cache can be scattered across different pages, with a page table recording their locations.
This greatly improves GPU memory utilization and allows serving more requests concurrently.
Example
pip install vllm
# Start an OpenAI-compatible API service with vLLM
# --model specifies the model, --host and --port specify the listening address
# --tensor-parallel-size: number of GPUs for parallel inference
python -m vllm.entrypoints.openai.api_server \
--model Qwen/Qwen2.5-7B-Instruct \
--host 0.0.0.0 \
--port 8000 \
--tensor-parallel-size 1 \
--gpu-memory-utilization 0.9 \
--max-model-len 8192
After the service starts, you can directly call it using the OpenAI SDK:
Example
from openai import OpenAI
# Connect to the local vLLM service
client = OpenAI(
base_url="http://localhost:8000/v1",
api_key="example-demo-key" # vLLM does not require a real API key by default
)
# Call chat completion
response = client.chat.completions.create(
model="Qwen/Qwen2.5-7B-Instruct",
messages=[
{"role": "system", "content": "You are a helpful AI assistant."},
{"role": "user", "content": "Introduce the Python tutorial in one sentence."}
],
temperature=0.7,
max_tokens=500,
stream=True # Streaming output
)
print("AI reply:", end="", flush=True)
for chunk in response:
if chunk.choices[0].delta.content:
print(chunk.choices[0].delta.content, end="", flush=True)
print()
vLLM also supports writing custom services directly with Python code:
Example
from vllm import LLM, SamplingParams
# Initialize the model
llm = LLM(
model="Qwen/Qwen2.5-7B-Instruct",
gpu_memory_utilization=0.9, # GPU memory usage ratio
tensor_parallel_size=1, # Tensor parallelism (number of GPUs)
max_model_len=8192, # Maximum context length
)
# Configure sampling parameters
sampling_params = SamplingParams(
temperature=0.7,
top_p=0.9,
max_tokens=500,
stop=["</s>"],
)
# Batch inference
prompts = [
"Introduce Python",
"What is AI?",
"How to learn programming?",
]
# Generate reply
outputs = llm.generate(prompts, sampling_params)
# Print results
for output in outputs:
prompt = output.prompt
generated_text = output.outputs[0].text
print(f"Prompt: {prompt}")
print(f"Generated: {generated_text}")
print("-" * 50)
TGI: Hugging Face's Inference Engine
TGI (Text Generation Inference) is an inference service framework launched by Hugging Face.
It features a complete ecosystem, deep integration with Hugging Face Hub, and supports optimizations such as Flash Attention and dynamic batching.
Example
# Note: You need to install Docker and configure GPU support first
docker run -d \
--gpus all \
-p 8080:80 \
-v $PWD/data:/data \
--name tgi-server \
ghcr.io/huggingface/text-generation-inference:latest \
--model-id Qwen/Qwen2.5-7B-Instruct \
--max-input-length 4096 \
--max-total-tokens 8192 \
--max-batch-prefill-tokens 4096
# Or install directly with Python (development environment)
# pip install text-generation
Calling the TGI service:
Example
from text_generation import Client
# Connect to the TGI service
client = Client("http://localhost:8080", timeout=60)
# Non-streaming call
response = client.generate(
"Introduce the Python tutorial",
max_new_tokens=500,
temperature=0.7,
top_p=0.9,
)
print(response.generated_text)
print("-" * 50)
# Streaming call
print("Streaming output:")
for chunk in client.generate_stream(
"Write a Hello World in Python",
max_new_tokens=200,
):
if not chunk.token.special:
print(chunk.token.text, end="", flush=True)
print()
Ollama: From Local Demo to Production Deployment
Ollama is known for "one-click local model running", and many people use it for development and demos.
But in fact, Ollama is also suitable for small-scale production deployment—it is simple, stable, and has controllable resource usage.
Example
curl -fsSL https://ollama.com/install.sh | sh
# Configure Ollama to listen on all network interfaces (pay attention to firewall in production)
# Edit the systemd service file
sudo systemctl edit ollama.service
# Add the following content:
# [Service]
# Environment="OLLAMA_HOST=0.0.0.0:11434"
# Environment="OLLAMA_MODELS=/var/lib/ollama/models"
# Restart the service
sudo systemctl restart ollama
# Pull and run the model
ollama pull qwen2.5:7b
Call Ollama via the API:
Example
import requests
import json
# Ollama API address
base_url = "http://localhost:11434/api"
# Chat interface (streaming)
def chat_stream(model, messages):
url = f"{base_url}/chat"
data = {
"model": model,
"messages": messages,
"stream": True,
"options": {
"temperature": 0.7,
"num_ctx": 8192,
}
}
response = requests.post(url, json=data, stream=True)
full_response = ""
for line in response.iter_lines():
if line:
chunk = json.loads(line)
if "message" in chunk:
content = chunk["message"].get("content", "")
full_response += content
print(content, end="", flush=True)
if chunk.get("done", False):
break
print()
return full_response
# Example invocation
messages = [
{"role": "system", "content": "You are a helpful AI assistant."},
{"role": "user", "content": "Introduce the Python tutorial"}
]
print("AI reply:", end="")
result = chat_stream("qwen2.5:7b", messages)
Framework Selection Comparison
The three frameworks each have their own characteristics. How to choose? It depends on your needs.
| Framework | Advantages | Disadvantages | Applicable scenarios |
|---|---|---|---|
| vLLM | Strongest performance, highest throughput | Relatively complex configuration | High-concurrency production environments, pursuing performance |
| TGI | Hugging Face ecosystem, simple deployment | Performance slightly lower than vLLM | Already using the Hugging Face toolchain |
| Ollama | Minimal deployment, low maintenance cost | Not suitable for extremely large concurrency | Small and medium-scale services, quick launch |
Early stage: start with Ollama, simple and reliable. After traffic grows: migrate to vLLM, prioritizing performance.
Inference Optimization Techniques
The same model, run in different ways, can differ in speed and cost by several times.
This section introduces inference optimization techniques commonly used in production environments.
Continuous Batching
Traditional batching is static: gather enough requests, run inference together, and process the next batch only after all are completed.
The problem is: some requests are short (e.g., only generate 10 characters), and some are long (generate 500 characters).
Short requests have to wait for long requests to complete before resources can be released together, so GPU utilization cannot go up.
Continuous batchingIt is dynamic: when one request completes, new requests from the queue are immediately added, without waiting for the entire batch to finish.
This way the GPU is almost always at full load, significantly increasing throughput.
Both vLLM and TGI have built-in continuous batching. You don't need to implement it manually; just enable it.
Speculative Decoding
Generating one token with a large model requires a full forward pass, which is slow.
The idea of speculative decoding is:Use a small model to "guess" the next few tokens, then have the large model verify them all at once.。
The small model guesses quickly, and although it may not be entirely correct, the correctly guessed part can be output in batch.
The incorrect guesses are discarded and re-guessed.
This can improve overall speed by 2-3x while keeping output quality unchanged.
Example
# --speculative-model specifies the small model (draft model)
# --num-speculative-tokens specifies how many tokens to guess each time
python -m vllm.entrypoints.openai.api_server \
--model Qwen/Qwen2.5-7B-Instruct \
--speculative-model Qwen/Qwen2.5-0.5B-Instruct \
--num-speculative-tokens 5 \
--host 0.0.0.0 \
--port 8000
Quantized Inference: AWQ and GPTQ
Model quantization compresses model weights from 16-bit floating point (FP16) to 8-bit or 4-bit integers.
Memory usage is halved or even less, and inference can be faster.
The two most commonly used quantization methods are AWQ and GPTQ.
| Quantization method | Feature | Accuracy loss | Inference speed |
|---|---|---|---|
| AWQ | Activation-aware weight quantization | Minimal | Fast |
| GPTQ | Layer-wise optimized quantization | small | Fast |
| GGUF | Suitable for local deployment | Acceptable | Fast |
Example
# Many quantized models can be downloaded directly from Hugging Face
python -m vllm.entrypoints.openai.api_server \
--model TheBloke/Qwen2.5-7B-Instruct-AWQ \
--quantization awq \
--host 0.0.0.0 \
--port 8000
# Or run GGUF quantized models with Ollama
ollama run qwen2.5:7b-instruct-q4_0
Load an AWQ quantized model in Python:
Example
from transformers import AutoModelForCausalLM, AutoTokenizer
import torch
# Load AWQ quantized model
model_name = "TheBloke/Qwen2.5-7B-Instruct-AWQ"
# Note: autoawq library must be installed to run this
# pip install autoawq
tokenizer = AutoTokenizer.from_pretrained(model_name)
model = AutoModelForCausalLM.from_pretrained(
model_name,
torch_dtype=torch.float16,
device_map="auto",
trust_remote_code=True,
)
# Inference
messages = [
{"role": "system", "content": "You are a helpful AI assistant."},
{"role": "user", "content": "Introduce Python"}
]
text = tokenizer.apply_chat_template(
messages,
tokenize=False,
add_generation_prompt=True
)
inputs = tokenizer([text], return_tensors="pt").to(model.device)
with torch.no_grad():
outputs = model.generate(
**inputs,
max_new_tokens=500,
temperature=0.7,
)
response = tokenizer.decode(
outputs[0][inputs.input_ids.shape[1]:],
skip_special_tokens=True
)
print(response)
Model Distillation
Distillation uses a large model (teacher) to train a small model (student), allowing the small model to acquire the capabilities of the large model.
The small model is faster and cheaper, and although its performance is slightly worse, it is sufficient in many scenarios.
For example, using Qwen2.5-72B as the teacher to train a Qwen2.5-7B student may yield results close to the 72B model, but with speed and cost at the 7B level.
The benefits of distillation are typically:Model size reduced to 1/10, inference speed improved by 5-10x, with over 90% of the performance retained.。
API Service Design
Wrapping the model is only the first step; the production API also needs many features: authentication, rate limiting, streaming responses, and error handling.
FastAPI is an excellent choice for this.
FastAPI Wrapping LLM Services
Let's write a complete API service with authentication, rate limiting, and streaming responses.
Example
from fastapi import FastAPI, Depends, HTTPException, status
from fastapi.security import HTTPBearer, HTTPAuthorizationCredentials
from fastapi.responses import StreamingResponse
from pydantic import BaseModel, Field
from typing import List, Optional, AsyncGenerator
from enum import Enum
import time
import asyncio
import uuid
from datetime import datetime, timedelta
from collections import defaultdict
# ============================================
# 1. Initialize the application and configuration
# ============================================
app = FastAPI(title="EXAMPLE LLM API", version="1.0.0")
# Simple API Key authentication (a more robust solution is recommended for production)
VALID_API_KEYS = {
"sk-example-123456": {"user": "demo_user", "quota": 1000},
"sk-example-789012": {"user": "pro_user", "quota": 10000},
}
# Simple rate limiter (in-memory implementation; Redis is recommended for production)
class RateLimiter:
def __init__(self):
self.requests = defaultdict(list) # key -> [timestamp, ...]
async def is_allowed(self, key: str, max_requests: int, window_seconds: int) -> bool:
"""Check if the rate limit is exceeded"""
now = time.time()
# Clean up expired records
self.requests[key] = [t for t in self.requests[key] if now - t < window_seconds]
# Check if the limit is exceeded
if len(self.requests[key]) >= max_requests:
return False
# Record this request
self.requests[key].append(now)
return True
rate_limiter = RateLimiter()
# ============================================
# 2. Data models
# ============================================
class Role(str, Enum):
SYSTEM = "system"
USER = "user"
ASSISTANT = "assistant"
class Message(BaseModel):
role: Role = Field(..., description="Message role")
content: str = Field(..., description="Message content")
class ChatCompletionRequest(BaseModel):
model: str = Field(..., description="Model name")
messages: List[Message] = Field(..., description="Conversation message list")
temperature: Optional[float] = Field(default=0.7, ge=0, le=2, description="Sampling temperature")
top_p: Optional[float] = Field(default=0.9, ge=0, le=1, description="Top-P sampling")
max_tokens: Optional[int] = Field(default=500, ge=1, le=4096, description="Maximum number of generated tokens")
stream: Optional[bool] = Field(default=False, description="Whether to stream output")
class ChatCompletionResponse(BaseModel):
id: str
object: str = "chat.completion"
created: int
model: str
choices: List[dict]
usage: dict
# ============================================
# 3. Authentication dependency
# ============================================
security = HTTPBearer()
async def get_current_user(credentials: HTTPAuthorizationCredentials = Depends(security)):
"""Validate API Key and return user information"""
api_key = credentials.credentials
if api_key not in VALID_API_KEYS:
raise HTTPException(
status_code=status.HTTP_401_UNAUTHORIZED,
detail="Invalid API Key"
)
return VALID_API_KEYS[api_key]
# ============================================
# 4. Simulate LLM generation (replace with real model calls in actual projects)
# ============================================
async def mock_llm_generate(messages: List[Message], max_tokens: int) -> str:
"""Simulate the LLM generation process"""
# This is just a simulation; in a real project, call the actual model
response_text = (
"Hello! I am EXAMPLE's AI assistant.\n\n"
"EXAMPLE provides a wealth of programming tutorials, including Python, Java, C++, front-end development, and more.\n"
"You can visit https://www.example.com to learn more!\n\n"
"How can I help you?"
)
# Simulate generation delay
await asyncio.sleep(0.5)
return response_text
async def mock_llm_generate_stream(messages: List[Message], max_tokens: int) -> AsyncGenerator[str, None]:
"""Simulate LLM streaming generation"""
response_text = await mock_llm_generate(messages, max_tokens)
# Output character by character to simulate a streaming effect
for char in response_text:
yield char
await asyncio.sleep(0.03) # Delay 30ms per character
# ============================================
# 5. API endpoints
# ============================================
@app.post("/v1/chat/completions")
async def chat_completions(
request: ChatCompletionRequest,
user: dict = Depends(get_current_user)
):
"""Chat completion endpoint, compatible with OpenAI format"""
# Rate limiting check: each user can make at most 60 requests per minute
user_key = user["user"]
if not await rate_limiter.is_allowed(user_key, max_requests=60, window_seconds=60):
raise HTTPException(
status_code=status.HTTP_429_TOO_MANY_REQUESTS,
detail="Rate limit exceeded"
)
# Generate request ID
request_id = f"chatcmpl-{uuid.uuid4().hex[:24]}"
created = int(datetime.now().timestamp())
if request.stream:
# Streaming response
async def stream_generator():
async for token in mock_llm_generate_stream(request.messages, request.max_tokens):
# Construct an SSE-format chunk
chunk = {
"id": request_id,
"object": "chat.completion.chunk",
"created": created,
"model": request.model,
"choices": [{
"index": 0,
"delta": {"content": token},
"finish_reason": None
}]
}
yield f"data: {__import__('json').dumps(chunk)}\n\n"
# Send end marker
yield "data: [DONE]\n\n"
return StreamingResponse(
stream_generator(),
media_type="text/event-stream",
headers={"Cache-Control": "no-cache"}
)
else:
# Non-streaming response
response_text = await mock_llm_generate(request.messages, request.max_tokens)
return ChatCompletionResponse(
id=request_id,
created=created,
model=request.model,
choices=[{
"index": 0,
"message": {
"role": "assistant",
"content": response_text
},
"finish_reason": "stop"
}],
usage={
"prompt_tokens": sum(len(m.content) for m in request.messages),
"completion_tokens": len(response_text),
"total_tokens": sum(len(m.content) for m in request.messages) + len(response_text)
}
)
@app.get("/health")
async def health_check():
"""Health check endpoint"""
return {"status": "ok", "timestamp": datetime.now().isoformat()}
@app.get("/v1/models")
async def list_models(user: dict = Depends(get_current_user)):
"""List available models"""
return {
"object": "list",
"data": [
{
"id": "qwen2.5-7b-instruct",
"object": "model",
"created": 1699000000,
"owned_by": "example"
},
{
"id": "llama-3.1-8b-instruct",
"object": "model",
"created": 1699000000,
"owned_by": "example"
}
]
}
if __name__ == "__main__":
import uvicorn
uvicorn.run(app, host="0.0.0.0", port=8000)
Start the service:
Example
pip install fastapi uvicorn python-multipart pydantic
# Start the service
python llm_api_server.py
# Test in another terminal
curl -X POST http://localhost:8000/v1/chat/completions \
-H "Content-Type: application/json" \
-H "Authorization: Bearer sk-example-123456" \
-d '{
"model": "qwen2.5-7b-instruct",
"messages": [
{"role": "user", "content": "Hello"}
],
"stream": false
}'
Call with Python client:
Example
from openai import OpenAI
# Connect to our own API service
client = OpenAI(
base_url="http://localhost:8000/v1",
api_key="sk-example-123456"
)
# Test non-streaming
print("Non-streaming call:")
response = client.chat.completions.create(
model="qwen2.5-7b-instruct",
messages=[{"role": "user", "content": "Introduce Python"}],
stream=False
)
print(response.choices[0].message.content)
print("\n" + "="*50 + "\n")
# Test streaming
print("Streaming call:")
stream = client.chat.completions.create(
model="qwen2.5-7b-instruct",
messages=[{"role": "user", "content": "Introduce Python"}],
stream=True
)
for chunk in stream:
if chunk.choices[0].delta.content:
print(chunk.choices[0].delta.content, end="", flush=True)
print()
Benefits of OpenAI-Compatible Interfaces
Making your API OpenAI-compatible has several clear benefits:
Users don't need to change code, just change base_url and api_key.
Rich ecosystem — any SDK, tool, or application that supports OpenAI can directly use your service.
No need to design your own API format; OpenAI's design is already well thought out.
Authentication and Rate Limiting in Production Environments
The previous examples used simple in-memory rate limiting; for production, it's recommended to use Redis for distributed rate limiting.
Example
import redis.asyncio as redis
import time
class RedisRateLimiter:
"""Redis-based distributed rate limiter"""
def __init__(self, redis_url: str = "redis://localhost:6379"):
self.redis = redis.from_url(redis_url)
async def is_allowed(
self,
key: str,
max_requests: int,
window_seconds: int
) -> bool:
"""
Sliding window rate limiting algorithm
key: rate limiting key (e.g., user_id or api_key)
max_requests: maximum number of requests in the window
window_seconds: window size (seconds)
"""
now = time.time()
window_start = now - window_seconds
# Redis key
redis_key = f"rate_limit:{key}"
# Use pipeline atomic operations
async with self.redis.pipeline() as pipe:
# 1. Remove records outside the window
pipe.zremrangebyscore(redis_key, 0, window_start)
# 2. Count requests in the current window
pipe.zcard(redis_key)
# 3. Add the current request
pipe.zadd(redis_key, {str(now): now})
# 4. Set expiration time
pipe.expire(redis_key, window_seconds)
# Execute
_, current_count, _, _ = await pipe.execute()
return current_count < max_requests
async def close(self):
await self.redis.close()
# Usage example
async def main():
limiter = RedisRateLimiter("redis://localhost:6379")
# Test: same user at most 10 requests per minute
user_id = "user_123"
for i in range(15):
allowed = await limiter.is_allowed(user_id, max_requests=10, window_seconds=60)
print(f"Request {i+1}: {'allowed' if allowed else 'denied'}")
await limiter.close()
if __name__ == "__main__":
import asyncio
asyncio.run(main())
Authentication and rate limiting are the infrastructure of an API; better to make them a bit more complex than to regret after being overwhelmed.
Load Balancing and Scaling
When there are more and more users and one server can't handle it, you need to add machines and distribute traffic across multiple machines.
This is the problem that load balancing and auto-scaling need to solve.
Stateless Service Design
First, remember:AI services should be designed to be stateless。
Stateless means: requests do not share memory, and any server can handle any request.
This makes scaling simple: just add machines, no need to worry about data synchronization.
Session state (such as conversation history) can be stored on the client (have the user send it each time), or in Redis.
Nginx Reverse Proxy
Nginx is the most commonly used load balancer, distributing traffic to multiple backend servers.
Example
# Nginx configuration: load balancing for multiple LLM services
upstream llm_backends {
# Least connections algorithm: send requests to the server with the fewest current connections
least_conn;
# Backend server list
server 192.168.1.101:8000 weight=1 max_fails=3 fail_timeout=30s;
server 192.168.1.102:8000 weight=1 max_fails=3 fail_timeout=30s;
server 192.168.1.103:8000 weight=1 max_fails=3 fail_timeout=30s;
# If you have multiple GPU models, you can allocate different proportions of traffic by weight
# weight=2 means this server receives twice the traffic of others
}
server {
listen 80;
server_name api.example.com;
# API request forwarding
location /v1/ {
# Forward to backend
proxy_pass http://llm_backends;
# Timeout settings (AI generation can be slow)
proxy_connect_timeout 10s;
proxy_send_timeout 600s;
proxy_read_timeout 600s;
# Pass the client's real IP
proxy_set_header Host $host;
proxy_set_header X-Real-IP $remote_addr;
proxy_set_header X-Forwarded-For $proxy_add_x_forwarded_for;
# Settings required for streaming responses
proxy_buffering off;
proxy_cache off;
proxy_set_header Connection '';
proxy_http_version 1.1;
}
# Health check endpoint
location /health {
proxy_pass http://llm_backends/health;
proxy_connect_timeout 5s;
proxy_read_timeout 5s;
}
# Rate limiting: max 10 requests per second per IP
limit_req_zone $binary_remote_addr zone=api_limit:10m rate=10r/s;
limit_req zone=api_limit burst=20 nodelay;
}
Kubernetes HPA Auto-scaling
If you deploy with Kubernetes, you can use HPA (Horizontal Pod Autoscaler) to automatically scale in and out based on load.
Example
# Kubernetes deployment configuration for LLM service
# 1. Deployment: deploy LLM Pod
apiVersion: apps/v1
kind: Deployment
metadata:
name: llm-server
namespace: example
spec:
replicas: 3 # Initial number of replicas
selector:
matchLabels:
app: llm-server
template:
metadata:
labels:
app: llm-server
spec:
containers:
- name: llm-server
image: example/llm-server:latest
ports:
- containerPort: 8000
resources:
requests:
cpu: "4"
memory: "16Gi"
nvidia.com/gpu: 1 # Request 1 GPU
limits:
cpu: "8"
memory: "32Gi"
nvidia.com/gpu: 1 # Limit 1 GPU
env:
- name: MODEL_NAME
value: "Qwen/Qwen2.5-7B-Instruct"
livenessProbe:
httpGet:
path: /health
port: 8000
initialDelaySeconds: 300 # Model loading takes time
periodSeconds: 30
readinessProbe:
httpGet:
path: /health
port: 8000
initialDelaySeconds: 300
periodSeconds: 10
tolerations:
- key: "nvidia.com/gpu"
operator: "Exists"
effect: "NoSchedule"
---
# 2. Service: load balancing Service
apiVersion: v1
kind: Service
metadata:
name: llm-service
namespace: example
spec:
selector:
app: llm-server
ports:
- port: 80
targetPort: 8000
type: LoadBalancer
---
# 3. HPA: auto-scaling
apiVersion: autoscaling/v2
kind: HorizontalPodAutoscaler
metadata:
name: llm-hpa
namespace: example
spec:
scaleTargetRef:
apiVersion: apps/v1
kind: Deployment
name: llm-server
minReplicas: 3 # Minimum 3 Pods
maxReplicas: 20 # Maximum 20 Pods
metrics:
# Scale based on CPU usage
- type: Resource
resource:
name: cpu
target:
type: Utilization
averageUtilization: 70
# Scale based on GPU usage (requires custom metrics)
- type: Pods
pods:
metric:
name: gpu_utilization
target:
type: AverageValue
averageValue: 80
# Scale based on queue length (custom metric)
- type: Pods
pods:
metric:
name: pending_requests
target:
type: AverageValue
averageValue: 10
behavior:
scaleUp:
stabilizationWindowSeconds: 60 # Observe for 60 seconds before scaling out
policies:
- type: Percent
value: 50
periodSeconds: 60
scaleDown:
stabilizationWindowSeconds: 300 # Observe for 5 minutes before scaling down
policies:
- type: Percent
value: 20
periodSeconds: 120
GPU Resource Pool Management
GPUs are expensive, so maximize utilization.
A common strategy is: group by model type, placing Pods of the same model together.
For example:
| Resource pool | GPU model | Running model | Purpose |
|---|---|---|---|
| small-pool | RTX 3090 / L4 | 7B model | Daily traffic |
| medium-pool | A10 / L40 | 13B-70B model | Complex tasks |
| large-pool | A100 / H100 | 70B+ model | High-value users |
| spot-pool | Various machine types | Offline batch processing | Use low-priced Spot instances |
A/B Testing AI Models
You've released a new model version. How do you know it's better than the old one? Use A/B testing.
The core of A/B testing is:Split traffic into groups, use different models for different groups, and compare metrics。
Traffic Split Strategy
There are several common traffic splitting methods:
Example
import hashlib
import random
from typing import Dict, List, Tuple
from dataclasses import dataclass
@dataclass
class Variant:
"""A variant of A/B testing"""
name: str # Variant name
model_name: str # Model name
weight: int # Traffic weight
config: dict # Other configuration
class ABTestRouter:
"""A/B testing traffic router"""
def __init__(self, test_name: str, variants: List[Variant]):
self.test_name = test_name
self.variants = variants
# Compute total weight
total_weight = sum(v.weight for v in variants)
# Build segment ranges
self.ranges: List[Tuple[int, int, Variant]] = []
current = 0
for v in variants:
self.ranges.append((current, current + v.weight, v))
current += v.weight
self.total_weight = total_weight
def get_variant_by_user(self, user_id: str) -> Variant:
"""
Stable traffic splitting based on user ID
The same user is always assigned to the same group to avoid inconsistent experience
"""
# Compute the user's score using a hash
hash_input = f"{self.test_name}-{user_id}".encode()
hash_value = hashlib.md5(hash_input).hexdigest()
# Convert to a number from 0-999
score = int(hash_value[:4], 16) % 1000
# Find the corresponding variant
for start, end, variant in self.ranges:
if start <= score < end:
return variant
return self.variants[-1]
def get_variant_random(self) -> Variant:
"""Random splitting (suitable for scenarios that don't require stickiness)"""
score = random.randint(0, self.total_weight - 1)
for start, end, variant in self.ranges:
if start <= score < end:
return variant
return self.variants[-1]
# Usage example
def main():
# Define three variants
variants = [
Variant(
name="control",
model_name="Qwen2.5-7B-Instruct",
weight=50, # 50% traffic
config={"temperature": 0.7}
),
Variant(
name="treatment-v1",
model_name="Qwen2.5-14B-Instruct",
weight=30, # 30% traffic
config={"temperature": 0.7}
),
Variant(
name="treatment-v2",
model_name="Llama-3.1-8B-Instruct",
weight=20, # 20% traffic
config={"temperature": 0.8}
)
]
router = ABTestRouter("model-comparison-v1", variants)
# Test user splitting
test_users = [f"user_{i}" for i in range(10)]
for user_id in test_users:
variant = router.get_variant_by_user(user_id)
print(f"User {user_id} → Group {variant.name} → Model {variant.model_name}")
# Distribution statistics
count: Dict[str, int] = {}
for i in range(10000):
variant = router.get_variant_random()
count[variant.name] = count.get(variant.name, 0) + 1
print("\nRandom splitting statistics (10000 times):")
for name, cnt in count.items():
print(f" {name}: {cnt} times ({cnt/100}%)")
if __name__ == "__main__":
main()
Metric Definition and Collection
A/B testing cannot just look at "which model responds smarter"; it is necessary to define quantifiable metrics.
| Metric types | Specific metrics | Description |
|---|---|---|
| System metrics | Latency, throughput, error rate | Whether the model is fast and stable |
| User behavior | User satisfaction, number of dialogue turns, renewal rate | Whether users like it |
| Business metrics | Conversion rate, retention rate, revenue | Whether it helps the business |
| Quality metrics | Human ratings, harmful rate, factual accuracy | Whether output quality is good |
Example
import time
import json
from dataclasses import dataclass, asdict
from typing import Optional
from datetime import datetime
import uuid
@dataclass
class ChatMetrics:
"""Complete metrics for chat requests"""
# Basic information
request_id: str
user_id: str
variant_name: str # A/B test group name
model_name: str
# Time metrics
timestamp: str
ttft: float # Time to First Token
total_time: float # Total elapsed time
# Generation metrics
prompt_tokens: int
completion_tokens: int
total_tokens: int
# Business metrics
user_rating: Optional[int] = None # User rating (1-5)
is_error: bool = False
error_message: Optional[str] = None
class MetricsCollector:
"""Metric collector"""
def __init__(self):
# In a real project, this should be written to Kafka or a database
self.buffer = []
def collect(self, metrics: ChatMetrics):
"""Collect one metric"""
self.buffer.append(asdict(metrics))
# Simple output; in a real project, write to storage
print(f"[Metrics] {json.dumps(asdict(metrics), ensure_ascii=False)}")
# When the buffer is full, write in batches
if len(self.buffer) >= 100:
self.flush()
def flush(self):
"""Batch write to storage"""
if self.buffer:
# In a real project: write to ClickHouse, BigQuery, etc.
print(f"Flushed {len(self.buffer)} metrics")
self.buffer = []
# Usage example
collector = MetricsCollector()
def process_chat_request(user_id: str, prompt: str, router):
"""Process a chat request and record metrics"""
request_id = str(uuid.uuid4())
# Split traffic
variant = router.get_variant_by_user(user_id)
# Record start time
start_time = time.time()
ttft = 0.0
prompt_tokens = len(prompt)
completion_tokens = 0
try:
# Replace this with a real model call
# response = call_model(variant.model_name, prompt)
response = "This is the AI's reply"
completion_tokens = len(response)
# Simulate TTFT (assume the first token is received after 0.2 seconds)
ttft = 0.2
# Record success metrics
metrics = ChatMetrics(
request_id=request_id,
user_id=user_id,
variant_name=variant.name,
model_name=variant.model_name,
timestamp=datetime.now().isoformat(),
ttft=ttft,
total_time=time.time() - start_time,
prompt_tokens=prompt_tokens,
completion_tokens=completion_tokens,
total_tokens=prompt_tokens + completion_tokens,
is_error=False
)
collector.collect(metrics)
return response, request_id
except Exception as e:
# Record error metrics
metrics = ChatMetrics(
request_id=request_id,
user_id=user_id,
variant_name=variant.name,
model_name=variant.model_name,
timestamp=datetime.now().isoformat(),
ttft=ttft,
total_time=time.time() - start_time,
prompt_tokens=prompt_tokens,
completion_tokens=0,
total_tokens=prompt_tokens,
is_error=True,
error_message=str(e)
)
collector.collect(metrics)
raise
# User feedback
def submit_user_feedback(request_id: str, rating: int):
"""User submits rating feedback"""
# In a real project: update the corresponding record in the database
print(f"Request {request_id} got rating {rating}")
Statistical Significance Assessment
After obtaining the data, you cannot simply look at 52% for group A and 55% for group B and conclude that group B is better.
A statistical significance test must be performed to confirm that the difference is not caused by random fluctuation.
Example
import math
from typing import Tuple
def calculate_z_test(
control_conv: float,
treatment_conv: float,
control_total: int,
treatment_total: int
) -> Tuple[float, float]:
"""
Calculate the p-value of the Z-test
control_conv: control group conversion rate
treatment_conv: treatment group conversion rate
control_total: control group sample size
treatment_total: treatment group sample size
"""
# pooled conversion rate
p_pool = (control_conv * control_total + treatment_conv * treatment_total) / (control_total + treatment_total)
# standard error
se_pool = math.sqrt(p_pool * (1 - p_pool) * (1 / control_total + 1 / treatment_total))
# Z statistic
z_score = (treatment_conv - control_conv) / se_pool
# compute p-value (two-tailed test)
p_value = 2 * (1 - normal_cdf(abs(z_score)))
return z_score, p_value
def normal_cdf(x: float) -> float:
"""Cumulative distribution function of the standard normal distribution (approximate)"""
return (1.0 + math.erf(x / math.sqrt(2.0))) / 2.0
def analyze_ab_test(
control_conversions: int,
control_total: int,
treatment_conversions: int,
treatment_total: int,
alpha: float = 0.05
):
"""Analyze A/B test results"""
control_conv = control_conversions / control_total
treatment_conv = treatment_conversions / treatment_total
z_score, p_value = calculate_z_test(
control_conv, treatment_conv,
control_total, treatment_total
)
relative_lift = (treatment_conv - control_conv) / control_conv
print("A/B test results analysis")
print("=" * 50)
print(f"Control group: {control_conversions}/{control_total} ({control_conv:.2%})")
print(f"Treatment group: {treatment_conversions}/{treatment_total} ({treatment_conv:.2%})")
print(f"Relative lift: {relative_lift:+.2%}")
print("-" * 50)
print(f"Z statistic: {z_score:.4f}")
print(f"p-value: {p_value:.4f}")
print("-" * 50)
if p_value < alpha:
print(f"✅ Result is statistically significant (p < {alpha})")
if relative_lift > 0:
print(" Treatment group performs better!")
else:
print(" Control group performs better!")
else:
print(f"❌ Result is not significant (p >= {alpha})")
print(" Cannot determine which version is better; more data is needed")
# Example: assuming A/B test results
if __name__ == "__main__":
# Scenario 1: significant difference
print("Scenario 1: significant difference")
analyze_ab_test(
control_conversions=450, # Control group: 450 satisfied
control_total=1000, # Control group total: 1000 people
treatment_conversions=520, # Treatment group: 520 satisfied
treatment_total=1000 # Treatment group total: 1000 people
)
print("\n" + "="*50 + "\n")
# Scenario 2: no significant difference
print("Scenario 2: no significant difference")
analyze_ab_test(
control_conversions=48,
control_total=100,
treatment_conversions=52,
treatment_total=100
)
Statistical significance is important. An untested "improvement" could just be random fluctuation.
MLflow Experiment Tracking
You've trained dozens of model versions. Which one performs best? Which version is deployed?
MLflow helps you manage experiments, track models, and record metrics.
Experiment Management
Example
import mlflow
import mlflow.pyfunc
import mlflow.sklearn
from mlflow.tracking import MlflowClient
from sklearn.ensemble import RandomForestClassifier
from sklearn.datasets import load_iris
from sklearn.model_selection import train_test_split
from sklearn.metrics import accuracy_score, f1_score
import numpy as np
import time
# Set MLflow tracking URI (local or remote server)
mlflow.set_tracking_uri("http://localhost:5000")
# Create or set experiment
experiment_name = "example-llm-finetuning"
try:
experiment_id = mlflow.create_experiment(experiment_name)
except mlflow.exceptions.MlflowException:
experiment_id = mlflow.get_experiment_by_name(experiment_name).experiment_id
mlflow.set_experiment(experiment_name)
# Simulate an LLM fine-tuning experiment
def run_finetuning_experiment(
model_name: str,
learning_rate: float,
batch_size: int,
num_epochs: int,
lora_rank: int,
dataset_name: str = "example-demo"
):
"""Run a fine-tuning experiment and log it to MLflow"""
with mlflow.start_run(run_name=f"{model_name}-lr{learning_rate}-bs{batch_size}") as run:
run_id = run.info.run_id
# Record parameters
mlflow.log_param("model_name", model_name)
mlflow.log_param("learning_rate", learning_rate)
mlflow.log_param("batch_size", batch_size)
mlflow.log_param("num_epochs", num_epochs)
mlflow.log_param("lora_rank", lora_rank)
mlflow.log_param("dataset_name", dataset_name)
# Simulate the training process and record metrics
print(f"Starting training experiment {run_id}...")
for epoch in range(num_epochs):
# Simulate training loss (decreases with epochs)
train_loss = 2.0 - 0.3 * epoch + np.random.normal(0, 0.1)
val_loss = 2.2 - 0.25 * epoch + np.random.normal(0, 0.15)
# Record metrics
mlflow.log_metric("train_loss", train_loss, step=epoch)
mlflow.log_metric("val_loss", val_loss, step=epoch)
mlflow.log_metric("epoch", epoch, step=epoch)
print(f"Epoch {epoch}: train_loss={train_loss:.4f}, val_loss={val_loss:.4f}")
time.sleep(0.5)
# Final evaluation metrics
final_val_loss = 1.2 + np.random.normal(0, 0.1)
final_accuracy = 0.85 + np.random.normal(0, 0.03)
final_f1 = 0.83 + np.random.normal(0, 0.04)
mlflow.log_metric("final_val_loss", final_val_loss)
mlflow.log_metric("final_accuracy", final_accuracy)
mlflow.log_metric("final_f1", final_f1)
# Record files (e.g., training logs, sample outputs)
with open("training_log.txt", "w") as f:
f.write(f"Experiment run_id: {run_id}\n")
f.write(f"Model: {model_name}\n")
f.write(f"Final accuracy: {final_accuracy:.4f}\n")
mlflow.log_artifact("training_log.txt")
# Log the model (simulated with sklearn here; in real projects, log the actual LLM)
# In real projects, you can use mlflow.transformers.log_model
X, y = load_iris(return_X_y=True)
X_train, X_test, y_train, y_test = train_test_split(X, y, random_state=42)
model = RandomForestClassifier(n_estimators=100)
model.fit(X_train, y_train)
mlflow.sklearn.log_model(
model,
"model",
registered_model_name=f"{model_name.replace('/', '-')}-demo"
)
print(f"Experiment complete! run_id: {run_id}")
print(f"Final accuracy: {final_accuracy:.4f}")
return run_id, final_accuracy
# Run several experiments
if __name__ == "__main__":
print("Start MLflow UI (run in another terminal):")
print(" mlflow ui --port 5000")
print("\nStarting experiments...\n")
# Experiment 1: Basic configuration
run_id1, acc1 = run_finetuning_experiment(
model_name="Qwen2.5-7B-Instruct",
learning_rate=2e-5,
batch_size=16,
num_epochs=3,
lora_rank=8
)
print()
# Experiment 2: Larger LoRA rank
run_id2, acc2 = run_finetuning_experiment(
model_name="Qwen2.5-7B-Instruct",
learning_rate=2e-5,
batch_size=16,
num_epochs=3,
lora_rank=16
)
print()
# Experiment 3: Higher learning rate
run_id3, acc3 = run_finetuning_experiment(
model_name="Qwen2.5-7B-Instruct",
learning_rate=5e-5,
batch_size=16,
num_epochs=3,
lora_rank=8
)
print()
print("All experiments complete!")
print(f"Experiment 1 (run_id={run_id1}): {acc1:.4f}")
print(f"Experiment 2 (run_id={run_id2}): {acc2:.4f}")
print(f"Experiment 3 (run_id={run_id3}): {acc3:.4f}")
print("\n"View results: http://localhost:5000")
Model Registry
MLflow Model Registry manages the model lifecycle: from Staging (testing) to Production.
Example
import mlflow
from mlflow.tracking import MlflowClient
mlflow.set_tracking_uri("http://localhost:5000")
client = MlflowClient()
def promote_model_to_production(
model_name: str,
version: int,
description: str = ""
):
"""Promote the model to the Production stage"""
# Update version description
client.update_model_version(
name=model_name,
version=version,
description=description
)
# Set this version as Production
client.transition_model_version_stage(
name=model_name,
version=version,
stage="Production",
archive_existing_versions=True # Archive old versions
)
print(f"Model {model_name} version {version} has been promoted to Production")
def get_production_model(model_name: str):
"""Get the current Production model"""
versions = client.get_latest_versions(model_name, stages=["Production"])
if versions:
return versions[0]
return None
def list_all_models():
"""List all registered models"""
models = client.search_registered_models()
print("Registered models:")
for model in models:
print(f" - {model.name}")
# List all versions of this model
versions = client.search_model_versions(f"name='{model.name}'")
for v in versions:
print(f" Version {v.version}: {v.current_stage}")
# Usage example
if __name__ == "__main__":
model_name = "Qwen2.5-7B-Instruct-demo"
print("List all models:")
list_all_models()
print("\n" + "="*50)
# Assume we want to promote version 1 to production
# promote_model_to_production(
# model_name=model_name,
# version=1,
# description="Production environment model: 2024 optimized version"
# )
print("\nGet current production model: ")
prod_model = get_production_model(model_name)
if prod_model:
print(f" Model name: {prod_model.name}")
print(f" Version: {prod_model.version}")
print(f" Stage: {prod_model.current_stage}")
print(f" Run ID: {prod_model.run_id}")
else:
print(" No production model yet")
Monitoring and Alerting
Once the service is live, how do you know it's working properly? Will users encounter timeouts? Will the model become increasingly inaccurate?
You need a complete monitoring system.
Key Metrics: TTFT, TPS, Error Rate
AI services have several core monitoring metrics:
| Metric | Full name | Meaning | Example alert threshold |
|---|---|---|---|
| TTFT | Time to First Token | Time from user request to receiving the first token | Alert > 3 seconds |
| TPOT | Time per Output Token | Average time per output token | Alert > 100ms |
| TPS | Tokens Per Second | How many tokens generated per second | Alert < 100 |
| Error rate | Error Rate | Proportion of failed requests | Alert > 5% |
| GPU utilization | GPU Utilization | GPU usage rate | Alert < 30% or > 95% |
| VRAM usage | GPU Memory | GPU VRAM occupancy rate | Alert > 90% |
Prometheus + Grafana
Prometheus collects metrics, Grafana displays dashboards - the golden combination for monitoring.
Example
from prometheus_client import (
Counter, Histogram, Gauge,
start_http_server, generate_latest
)
import time
import random
# ============================================
# Define metrics
# ============================================
# Counter: a counter that only increases
REQUEST_COUNT = Counter(
"llm_requests_total",
"Total number of LLM requests",
["model", "variant", "status"] # Labels
)
# Histogram: statistical distribution (e.g., latency)
REQUEST_LATENCY = Histogram(
"llm_request_latency_seconds",
"LLM request latency",
["model", "variant"],
buckets=[0.1, 0.5, 1, 2, 5, 10, 30, 60] # Buckets
)
TTFT_HISTOGRAM = Histogram(
"llm_ttft_seconds",
"Time to first token",
["model", "variant"],
buckets=[0.05, 0.1, 0.2, 0.5, 1, 2, 5]
)
# Gauge: a gauge that can go up or down
ACTIVE_CONNECTIONS = Gauge(
"llm_active_connections",
"Number of active connections",
["model"]
)
GPU_UTILIZATION = Gauge(
"llm_gpu_utilization_percent",
"GPU utilization",
["gpu_id", "model"]
)
GPU_MEMORY_USED = Gauge(
"llm_gpu_memory_used_bytes",
"GPU memory used",
["gpu_id", "model"]
)
TOKENS_PER_SECOND = Gauge(
"llm_tokens_per_second",
"Tokens generated per second",
["model"]
)
# ============================================
# Simulate service
# ============================================
def process_request(model: str, variant: str):
"""Simulate processing a request and record metrics"""
start_time = time.time()
# Active connections +1
ACTIVE_CONNECTIONS.labels(model=model).inc()
try:
# Simulate TTFT (time to first token)
ttft = random.uniform(0.05, 0.3)
TTFT_HISTOGRAM.labels(model=model, variant=variant).observe(ttft)
# Simulate processing time
processing_time = random.uniform(0.5, 3.0)
time.sleep(processing_time)
# Simulate success
REQUEST_COUNT.labels(
model=model, variant=variant, status="success"
).inc()
# Record latency
latency = time.time() - start_time
REQUEST_LATENCY.labels(model=model, variant=variant).observe(latency)
# Simulate TPS
tps = random.uniform(50, 200)
TOKENS_PER_SECOND.labels(model=model).set(tps)
except Exception:
# Simulate error
REQUEST_COUNT.labels(
model=model, variant=variant, status="error"
).inc()
raise
finally:
# Active connections -1
ACTIVE_CONNECTIONS.labels(model=model).dec()
def update_gpu_metrics():
"""Simulate updating GPU metrics"""
# Simulate metrics for GPU 0
gpu0_util = random.uniform(40, 90)
gpu0_mem = random.uniform(10, 20) * 1024**3 # 10-20 GB
GPU_UTILIZATION.labels(gpu_id="0", model="Qwen2.5-7B").set(gpu0_util)
GPU_MEMORY_USED.labels(gpu_id="0", model="Qwen2.5-7B").set(gpu0_mem)
# ============================================
# Start service
# ============================================
if __name__ == "__main__":
# Start Prometheus metrics endpoint
start_http_server(8081)
print("Prometheus metrics server started on port 8081")
print("Metrics available at http://localhost:8081/metrics")
# Simulate traffic
models = ["Qwen2.5-7B-Instruct", "Llama-3.1-8B-Instruct"]
variants = ["control", "treatment-v1"]
while True:
model = random.choice(models)
variant = random.choice(variants)
process_request(model, variant)
update_gpu_metrics()
time.sleep(random.uniform(0.1, 0.5))
Prometheus configuration file:
Example
global:
scrape_interval: 15s
scrape_configs:
# Scrape LLM service metrics
- job_name: "llm-service"
static_configs:
- targets: ["localhost:8081"]
labels:
service: "llm-api"
environment: "production"
# Scrape node metrics (GPU, CPU, memory)
- job_name: "node"
static_configs:
- targets: ["localhost:9100"]
alerting:
alertmanagers:
- static_configs:
- targets: ["localhost:9093"]
# Alert rules
rule_files:
- "alerts.yml"
Alert rules:
Example
groups:
- name: llm_alerts
interval: 30s
rules:
# High error rate alert
- alert: HighErrorRate
expr: |
sum(rate(llm_requests_total{status="error"}[5m]))
/
sum(rate(llm_requests_total[5m])) > 0.05
for: 5m
labels:
severity: critical
annotations:
summary: "High error rate alert"
description: "Error rate exceeds 5%, current value: {{ $value | humanizePercentage }}"
# TTFT too high alert
- alert: SlowTTFT
expr: |
histogram_quantile(0.95, sum(rate(llm_ttft_seconds_bucket[5m])) by (le, model)) > 2
for: 5m
labels:
severity: warning
annotations:
summary: "TTFT too high"
description: "The 95th percentile TTFT of model {{ $labels.model }} exceeds 2 seconds"
# GPU utilization too low (resource waste)
- alert: LowGPUUtilization
expr: llm_gpu_utilization_percent < 30
for: 30m
labels:
severity: info
annotations:
summary: "GPU utilization too low"
description: "GPU {{ $labels.gpu_id }} utilization {{ $value }}%, possible resource waste"
# GPU utilization too high (bottleneck risk)
- alert: HighGPUUtilization
expr: llm_gpu_utilization_percent > 95
for: 5m
labels:
severity: warning
annotations:
summary: "GPU utilization too high"
description: "GPU {{ $labels.gpu_id }} utilization {{ $value }}%, may become a bottleneck"
AI-Specific Monitoring: Drift Detection
AI models not only "break", they also "drift": over time, the data distribution changes and model performance degrades.
Example
import numpy as np
from scipy import stats
from dataclasses import dataclass
from typing import List, Tuple
import json
@dataclass
class DistributionStats:
"""Distribution statistics"""
mean: float
std: float
min: float
max: float
percentiles: dict
class DataDriftDetector:
"""Data drift detector"""
def __init__(self, baseline_data: np.ndarray):
"""Initialize with baseline data"""
self.baseline_stats = self._compute_stats(baseline_data)
self.baseline_data = baseline_data
def _compute_stats(self, data: np.ndarray) -> DistributionStats:
"""Compute statistics for the data"""
return DistributionStats(
mean=float(np.mean(data)),
std=float(np.std(data)),
min=float(np.min(data)),
max=float(np.max(data)),
percentiles={
"p10": float(np.percentile(data, 10)),
"p25": float(np.percentile(data, 25)),
"p50": float(np.percentile(data, 50)),
"p75": float(np.percentile(data, 75)),
"p90": float(np.percentile(data, 90))
}
)
def detect(self, new_data: np.ndarray, threshold: float = 0.05) -> dict:
"""
Detect drift
Returns: whether drift exists, p-value, statistics
"""
# KS test: compare whether two distributions are the same
ks_statistic, ks_pvalue = stats.ks_2samp(self.baseline_data, new_data)
# Compute statistics for new data
new_stats = self._compute_stats(new_data)
# Compute mean difference (percentage)
mean_diff_pct = abs((new_stats.mean - self.baseline_stats.mean) / self.baseline_stats.mean * 100)
# Determine if drift occurred
is_drift = ks_pvalue < threshold
return {
"is_drift": is_drift,
"ks_statistic": float(ks_statistic),
"ks_pvalue": float(ks_pvalue),
"mean_diff_percent": float(mean_diff_pct),
"baseline_stats": self.baseline_stats,
"new_stats": new_stats,
"threshold": threshold
}
# Simulation: drift detection for user input length
def simulate_input_length_drift():
# Baseline data: average user input length 50
np.random.seed(42)
baseline = np.random.normal(loc=50, scale=15, size=1000)
baseline = np.clip(baseline, 5, 200)
detector = DataDriftDetector(baseline)
print("Data drift detection example")
print("="*50)
print(fBaseline: mean={detector.baseline_stats.mean:.2f}, std={detector.baseline_stats.std:.2f})
print()
# Scenario 1: Normal data (no drift)
print(Scenario 1: Normal data)
normal_data = np.random.normal(loc=50, scale=15, size=500)
normal_data = np.clip(normal_data, 5, 200)
result = detector.detect(normal_data)
print(fNew data mean: {result['new_stats'].mean:.2f})
print(f" KS p-value: {result['ks_pvalue']:.4f}")
print(fDrift? {'Yes' if result['is_drift'] else 'No'})
print()
# Scenario 2: Drift (input got longer)
print(Scenario 2: Input got longer)
drift_data = np.random.normal(loc=100, scale=20, size=500)
drift_data = np.clip(drift_data, 5, 300)
result = detector.detect(drift_data)
print(fNew data mean: {result['new_stats'].mean:.2f})
print(f" KS p-value: {result['ks_pvalue']:.4f}")
print(fDrift? {'Yes' if result['is_drift'] else 'No'})
print(fMean difference: {result['mean_diff_percent']:.1f}%)
if __name__ == "__main__":
simulate_input_length_drift()
Cost Optimization Strategies
GPUs are expensive. Optimizing costs isn't "being cheap"—it's "engineering ability".
Request Caching: Semantic Caching
Traditional caching is exact match: for exactly the same question, directly return the previous answer.
Semantic caching goes further: similar questions can also reuse previous answers or intermediate results.
Example
import numpy as np
from typing import Optional, Tuple
import time
class SemanticCache:
"""Simple semantic cache implementation"""
def __init__(self, threshold: float = 0.9):
self.threshold = threshold
self.queries = [] # Store the vectors of questions
self.answers = [] # Store corresponding answers
self.timestamps = [] # Store timestamps
def _simulate_embedding(self, text: str) -> np.ndarray:
"""
Simulated embedding; in real projects, use OpenAI Embeddings, Sentence-Transformers, etc.
"""
# Simple hash simulation; in real projects, use a real embedding model
np.random.seed(hash(text) % (2**32))
return np.random.randn(128) # 128-dimensional vector
def _cosine_similarity(self, a: np.ndarray, b: np.ndarray) -> float:
"""Calculate cosine similarity"""
return np.dot(a, b) / (np.linalg.norm(a) * np.linalg.norm(b))
def get(self, query: str) -> Optional[str]:
"""Query cache: if a similar question is found, return the answer"""
query_vec = self._simulate_embedding(query)
for i, (q_vec, answer, ts) in enumerate(zip(self.queries, self.answers, self.timestamps)):
# Expiration cleanup (more than 1 hour)
if time.time() - ts > 3600:
continue
similarity = self._cosine_similarity(query_vec, q_vec)
if similarity >= self.threshold:
return answer
return None
def put(self, query: str, answer: str):
"""Store into cache"""
query_vec = self._simulate_embedding(query)
self.queries.append(query_vec)
self.answers.append(answer)
self.timestamps.append(time.time())
# Simple capacity control: store at most 1000 entries
if len(self.queries) > 1000:
self.queries = self.queries[-1000:]
self.answers = self.answers[-1000:]
self.timestamps = self.timestamps[-1000:]
# Usage example
def main():
cache = SemanticCache(threshold=0.9)
# First request
query1 = How to read a CSV file in Python?
print(fQuery 1: {query1})
answer1 = cache.get(query1)
if answer1:
print(fCache hit: {answer1})
else:
print(Cache miss, calling model...)
answer1 = You can use pandas.read_csv('file.csv') to read CSV files.
cache.put(query1, answer1)
print(fStore in cache: {answer1})
print()
# Second time: similar question
query2 = How to read CSV in Python?
print(fQuery 2: {query2})
answer2 = cache.get(query2)
if answer2:
print(fCache hit: {answer2})
else:
print(Cache miss, calling model...)
print()
# Third time: completely different question
query3 = What is machine learning?
print(fQuery 3: {query3})
answer3 = cache.get(query3)
if answer3:
print(fCache hit: {answer3})
else:
print(Cache miss, calling model...)
if __name__ == "__main__":
main()
Model Routing: Small Model as Fallback
Not all requests need a large model. Use a small model for simple questions, and a large model for complex questions.
Example
from typing import List, Dict, Any
from dataclasses import dataclass
@dataclass
class ModelConfig:
name: str
cost_per_1k_tokens: float
max_tokens: int
capabilities: List[str] # Capability list
class SmartModelRouter:
"""Intelligent model router"""
def __init__(self):
self.models = {
"small": ModelConfig(
name="Qwen2.5-0.5B-Instruct",
cost_per_1k_tokens=0.001,
max_tokens=4096,
capabilities=["chat", "simple_qa", "summarization"]
),
"medium": ModelConfig(
name="Qwen2.5-7B-Instruct",
cost_per_1k_tokens=0.01,
max_tokens=8192,
capabilities=["chat", "qa", "coding", "reasoning"]
),
"large": ModelConfig(
name="Qwen2.5-72B-Instruct",
cost_per_1k_tokens=0.1,
max_tokens=32768,
capabilities=["complex_qa", "advanced_coding", "deep_reasoning"]
)
}
def _classify_query(self, query: str) -> str:
"""
Simple question classification (you can use a classifier in real projects)
Return: small/medium/large
"""
query_lower = query.lower()
# Keywords for simple questions
simple_keywords = [
"Hello", "hello", "Thank you", "hi", "Goodbye",
"Today's weather", "What time", "Date",
"Help translate", "translate",
"Brief introduction", "What is"
]
# Keywords for complex questions
complex_keywords = [
"Code", "Write a", "Implement", "code", "python",
"Analyze", "Explain why", "Derive", "Prove",
"Compare", "Contrast", "Difference",
"Detailed", "In-depth", "Complex"
]
# Check if it is a simple question
for keyword in simple_keywords:
if keyword in query_lower:
return "small"
# Check if it is a complex question
for keyword in complex_keywords:
if keyword in query_lower:
return "large"
# Default to medium model
return "medium"
def route(self, query: str, user_tier: str = "free") -> str:
"""
Routing decision
"""
# First classify the question
category = self._classify_query(query)
# Paid users can use a better model
if user_tier == "pro":
upgrade_map = {
"small": "small",
"medium": "medium",
"large": "large"
}
category = upgrade_map[category]
elif user_tier == "enterprise":
upgrade_map = {
"small": "medium",
"medium": "large",
"large": "large"
}
category = upgrade_map.get(category, category)
# Return model name
return self.models[category].name
def get_cost_estimate(self, model_name: str, input_tokens: int, output_tokens: int) -> float:
"""Estimate cost"""
for model in self.models.values():
if model.name == model_name:
return (input_tokens + output_tokens) * model.cost_per_1k_tokens / 1000
return 0.0
# Usage example
def main():
router = SmartModelRouter()
test_queries = [
"Hello, I'd like to learn more",
"What is Python?",
"Help me write a Python quicksort",
Provide a detailed analysis of the Transformer architecture.
]
print(Intelligent routing example)
print("="*50)
for query in test_queries:
model = router.route(query, user_tier="free")
cost = router.get_cost_estimate(model, 100, 200)
print(fQuery: {query})
print(fRoute to: {model})
print(fEstimated cost: ${cost:.4f})
print()
# Different user tiers
print(Comparison of different user tiers)
print("-"*50)
query = Help me write a Python quicksort
for tier in ["free", "pro", "enterprise"]:
model = router.route(query, user_tier=tier)
print(f{tier} user: {model})
if __name__ == "__main__":
main()
Spot Instances and Batch Processing
Use On-Demand instances for real-time requests (stable but expensive), and Spot instances for offline batch processing (cheap but may be preempted).
Summary of several cost optimization methods:
| Strategy | Scenario | Cost Savings |
|---|---|---|
| Semantic cache | Many repeated queries | 30-70% |
| Model routing | High proportion of simple queries | 50-80% |
| Model quantization | All scenarios | 30-50% |
| Spot instances | Offline batch processing | 60-90% |
| Auto scaling | Large traffic fluctuations | 30-60% |
Fault Handling and Degradation
In production environments, there is no such thing as 'no failures' — only 'how to handle failures gracefully'.
Degradation Strategy
Example
import time
from typing import Optional
from dataclasses import dataclass
from enum import Enum
class ServiceStatus(Enum):
HEALTHY = "healthy"
DEGRADED = "degraded"
FAILED = "failed"
@dataclass
class FallbackResponse:
content: str
is_fallback: bool
fallback_reason: Optional[str] = None
class FallbackHandler:
"""Fallback handler"""
def __init__(self):
self.error_count = 0
self.last_error_time = 0
self.circuit_open = False
self.circuit_open_time = 0
# Circuit breaker configuration
self.error_threshold = 5 # Trip after 5 consecutive errors
self.circuit_timeout = 30 # Attempt recovery after 30 seconds of open circuit
def _call_primary_service(self, query: str) -> str:
"""Simulate calling the primary service"""
# Replace this with a real model call
return "This is the primary service's reply"
def _call_fallback_service(self, query: str) -> str:
"""Simulate calling the fallback service (faster but lower quality)"""
# Could be a smaller model, a cached reply, or even a predefined template
return "Sorry, the service is temporarily busy. This is a simplified reply."
def _call_static_fallback(self, query: str) -> str:
"""Static fallback reply (last line of defense)"""
return (
"We are very sorry, our service is currently experiencing some issues.\n"
"Please try again later. For urgent issues, please contact [email protected]."
)
def process_request(self, query: str) -> FallbackResponse:
Handle requests with complete fallback logic
# 1. Check whether the circuit breaker is open
if self.circuit_open:
if time.time() - self.circuit_open_time > self.circuit_timeout:
# Timed out, try half-open state, allow a small number of requests through
print("Circuit breaker half-open, attempting recovery...")
self.circuit_open = False
else:
# Fall back directly
print("Circuit breaker is open, falling back directly")
return FallbackResponse(
content=self._call_fallback_service(query),
is_fallback=True,
fallback_reason="circuit_open"
)
# 2. Try primary service
try:
result = self._call_primary_service(query)
# Success: reset error count
self.error_count = 0
return FallbackResponse(
content=result,
is_fallback=False
)
except Exception as e:
print(f"Primary service call failed: {e}")
self.error_count += 1
self.last_error_time = time.time()
# 3. Check if circuit breaking is needed
if self.error_count >= self.error_threshold:
print("Too many errors, opening circuit breaker")
self.circuit_open = True
self.circuit_open_time = time.time()
# 4. Try fallback service
try:
print("Try to call fallback service")
result = self._call_fallback_service(query)
return FallbackResponse(
content=result,
is_fallback=True,
fallback_reason="primary_failed"
)
except Exception as e2:
print(f"Fallback service also failed: {e2}")
# 5. Final static reply
return FallbackResponse(
content=self._call_static_fallback(query),
is_fallback=True,
fallback_reason="all_failed"
)
# Usage example
def main():
handler = FallbackHandler()
print("Fallback handling example")
print("="*50)
# Normal request
print("Normal request:")
resp = handler.process_request("Hello")
print(f" Reply: {resp.content}")
print(f" Is fallback? {resp.is_fallback}")
print()
# Simulate error (in real projects, exceptions are actually thrown)
# Here we simply demonstrate the logic
print("Fallback strategy summary:")
print(" 1. Prefer primary service")
print(" 2. Primary service fails → use small model/cache")
print(" 3. Consecutive failures → circuit breaking to avoid avalanche")
print(" 4. During circuit breaking → quickly return fallback reply")
print(" 5. Last line of defense → static reply")
Canary Release and Rollback
Don't roll out the new version to all users at once. First give it to 1% of traffic, and if there are no issues, then 5%, 10%, 100%.
If problems occur, you need to be able to roll back quickly. Don't wait until users have complained for a long time before frantically changing configuration.
Other extensions