Load balancing isn't just "distribute requests to servers." For LLMs it's a strategic decision: which server for which request?
The Problem: Naive Round-Robin
LB distributes REQUESTS (not LOAD):
Request 1 (10 Tokens): β Server A
Request 2 (1000 Tokens): β Server B
Request 3 (50 Tokens): β Server C
Server A: 10 Tokens (1% CPU)
Server B: 1000 Tokens (100% CPU) β BOTTLENECK!
Server C: 50 Tokens (5% CPU)
Result: Server B overloaded, A and C idle.
Better: Token-Aware Load Balancing
Count tokens BEFORE routing:
Request 1 (10 Tokens): β Server C (lowest load)
Request 2 (1000 Tokens): β Distribute or queue
Request 3 (50 Tokens): β Server A
Result: All servers balanced.
1. Routing Strategies
Strategy 1: Least Connections
Simplest strategy:
class LeastConnectionsRouter:
def __init__(self, servers: list[Server]):
self.servers = servers
self.active_connections = {s.id: 0 for s in servers}
def route(self, request: dict) -> Server:
"""Route to server with fewest connections"""
best_server = min(
self.servers,
key=lambda s: self.active_connections[s.id]
)
self.active_connections[best_server.id] += 1
return best_server
def release(self, server_id: str):
"""Request doneβdecrement counter"""
self.active_connections[server_id] -= 1
Problem: Ignores that Request 2 uses 1000 tokens and Request 1 only 10.
Strategy 2: Token-Aware Routing
Better: Consider tokens needed:
class TokenAwareRouter:
def __init__(self, servers: list[Server]):
self.servers = servers
self.current_load = {s.id: 0 for s in servers}
self.max_token_load = 100000
def estimate_tokens(self, request: dict) -> int:
"""Estimate tokens for request"""
prompt_tokens = len(request["prompt"].split()) * 1.3
output_tokens = request.get("max_tokens", 200)
return int(prompt_tokens + output_tokens)
def route(self, request: dict) -> Server:
"""Route to server with lowest token load"""
tokens = self.estimate_tokens(request)
available = [
s for s in self.servers
if self.current_load[s.id] + tokens <= self.max_token_load
]
if not available:
return self.servers[0] # Fallback
best_server = min(available, key=lambda s: self.current_load[s.id])
self.current_load[best_server.id] += tokens
return best_server
def update_load(self, server_id: str, tokens_used: int):
"""Update after request completion"""
self.current_load[server_id] = max(0, self.current_load[server_id] - tokens_used)
Strategy 3: Latency-Aware Routing
Consider measured latency:
class LatencyAwareRouter:
def __init__(self, servers: list[Server]):
self.servers = servers
self.latency_history = {s.id: [] for s in servers}
self.window_size = 100
def get_p95_latency(self, server_id: str) -> float:
"""P95 latency for server"""
latencies = self.latency_history[server_id]
if not latencies:
return 0
sorted_lat = sorted(latencies)
p95_idx = int(len(sorted_lat) * 0.95)
return sorted_lat[p95_idx]
def route(self, request: dict) -> Server:
"""Route to fastest server"""
return min(self.servers, key=lambda s: self.get_p95_latency(s.id))
def record_latency(self, server_id: str, latency_ms: float):
"""Record latency"""
latencies = self.latency_history[server_id]
latencies.append(latency_ms)
if len(latencies) > self.window_size:
latencies.pop(0)
Strategy 4: Weighted Load Balancing
Different servers have different capacity:
class WeightedRouter:
def __init__(self, servers: list[tuple[Server, float]]):
"""
servers = [
(Server("gpu-1"), 0.5), # Weaker
(Server("gpu-2"), 1.0), # Reference
(Server("gpu-3"), 2.0), # Double capacity
]
"""
self.servers = servers
self.request_counts = {s.id: 0 for s, _ in servers}
def route(self, request: dict) -> Server:
"""Distribute based on weight"""
best_server = min(
self.servers,
key=lambda sw: (self.request_counts[sw[0].id] / sw[1])
)
server, _ = best_server
self.request_counts[server.id] += 1
return server
2. Multi-Model Deployment
Not all models on every server (VRAM limits):
class MultiModelRouter:
def __init__(self):
self.server_capabilities = {
"gpu-1": {
"models": ["llama-7b", "llama-13b"],
"vram_total": 16,
"vram_used": 14
},
"gpu-2": {
"models": ["llama-70b"],
"vram_total": 80,
"vram_used": 70
}
}
def find_server_for_model(self, model: str) -> str | None:
"""Find server with this model"""
for server_id, cap in self.server_capabilities.items():
if model in cap["models"]:
available_vram = cap["vram_total"] - cap["vram_used"]
if available_vram > 2:
return server_id
return None
def route(self, request: dict) -> str:
"""Route to correct server for model"""
model = request.get("model")
server = self.find_server_for_model(model)
if not server:
raise ValueError(f"Model {model} unavailable")
return server
3. A/B Testing on Inference Layer
Test two model versions in parallel:
import random
class ABTestingRouter:
def __init__(self, model_a: str, model_b: str, split_ratio: float = 0.5):
self.model_a = model_a
self.model_b = model_b
self.split_ratio = split_ratio
self.results = {"model_a": [], "model_b": []}
def route(self, request: dict) -> tuple[str, str]:
"""Decide which model and return test group"""
if random.random() < self.split_ratio:
return self.model_a, "A"
else:
return self.model_b, "B"
def record_result(self, test_group: str, latency_ms: float, quality_score: float):
"""Store test result"""
self.results[f"model_{test_group.lower()}"].append({
"latency": latency_ms,
"quality": quality_score
})
def get_stats(self) -> dict:
"""Compare models"""
stats = {}
for group in ["model_a", "model_b"]:
results = self.results[group]
if not results:
continue
latencies = [r["latency"] for r in results]
qualities = [r["quality"] for r in results]
stats[group] = {
"requests": len(results),
"avg_latency_ms": sum(latencies) / len(latencies),
"avg_quality": sum(qualities) / len(qualities)
}
return stats
4. Failover and Health Checks
import httpx
import asyncio
class ResilientRouter:
def __init__(self, servers: list[str]):
self.servers = servers
self.status = {s: "healthy" for s in servers}
self.consecutive_failures = {s: 0 for s in servers}
async def health_check(self, server: str) -> bool:
"""Check if server alive"""
try:
async with httpx.AsyncClient(timeout=2) as client:
response = await client.get(f"{server}/health")
return response.status_code == 200
except Exception:
return False
async def monitor_health(self):
"""Background: Monitor servers"""
while True:
tasks = [self.health_check(s) for s in self.servers]
results = await asyncio.gather(*tasks)
for server, is_healthy in zip(self.servers, results):
if is_healthy:
self.status[server] = "healthy"
self.consecutive_failures[server] = 0
else:
self.consecutive_failures[server] += 1
if self.consecutive_failures[server] >= 3:
self.status[server] = "dead"
await asyncio.sleep(10)
def route_with_failover(self, request: dict) -> str:
"""Route with failover chain"""
healthy = [
s for s in self.servers
if self.status[s] in ["healthy", "degraded"]
]
if not healthy:
raise Exception("All servers dead!")
return healthy[0]
5. Dynamic Scaling with Kubernetes
apiVersion: v1
kind: Service
metadata:
name: llm-service-lb
spec:
type: LoadBalancer
selector:
app: llm-service
sessionAffinity: ClientIP
---
apiVersion: autoscaling.k8s.io/v2
kind: HorizontalPodAutoscaler
metadata:
name: llm-service-hpa
spec:
scaleTargetRef:
apiVersion: apps/v1
kind: Deployment
name: llm-service
minReplicas: 2
maxReplicas: 20
metrics:
- type: Resource
resource:
name: memory
target:
type: Utilization
averageUtilization: 75
behavior:
scaleUp:
stabilizationWindowSeconds: 0 # Scale up immediately
policies:
- type: Percent
value: 100 # Double replicas
periodSeconds: 15
Summary: Routing Strategies
| Strategy | Complexity | Best For |
|---|---|---|
| Round Robin | Low | Homogeneous servers |
| Least Conn | Low | Many short requests |
| Token-Aware | Medium | LLM Inference |
| Latency-Aware | Medium | Performance-critical |
| Weighted | Medium | Heterogeneous hardware |
| A/B Testing | High | Model optimization |
Start with Token-Aware + Health Checks. That covers 90% of cases.
Sources and Links
- NGINX Load Balancing β Production Load Balancer
- Traefik AI Gateway β LLM-aware LB
- Kubernetes HPA β Auto Scaling
- LiteLLM Load Balancing β LLM Router
- Token-Aware Routing β AnyscaleResearch
- Consistent Hashing β Advanced Technique
