Implement random tie breadking for cache_aware sglang router policy (#33138)
Co-authored-by: Patrick Toulme <ptoulme@google.com>
This commit is contained in:
co-authored by
Patrick Toulme
parent
cf7923615c
commit
735995e7bd
@@ -34,7 +34,8 @@
|
|||||||
2. Load Balancing (Shortest Queue)
|
2. Load Balancing (Shortest Queue)
|
||||||
-------------------------------------------
|
-------------------------------------------
|
||||||
This strategy tracks pending request counts per worker and routes new requests
|
This strategy tracks pending request counts per worker and routes new requests
|
||||||
to the least busy worker when the system is detected to be imbalanced.
|
to the least busy worker when the system is detected to be imbalanced. Ties
|
||||||
|
are randomly broken.
|
||||||
|
|
||||||
Configuration Parameters:
|
Configuration Parameters:
|
||||||
------------------------
|
------------------------
|
||||||
@@ -63,7 +64,7 @@ use std::sync::Arc;
|
|||||||
|
|
||||||
use async_trait::async_trait;
|
use async_trait::async_trait;
|
||||||
use dashmap::DashMap;
|
use dashmap::DashMap;
|
||||||
use rand::Rng;
|
use rand::{seq::IteratorRandom, Rng};
|
||||||
use smg_mesh::{tree_ops::TreeOperation, OptionalMeshSyncManager};
|
use smg_mesh::{tree_ops::TreeOperation, OptionalMeshSyncManager};
|
||||||
use tracing::{debug, warn};
|
use tracing::{debug, warn};
|
||||||
|
|
||||||
@@ -325,11 +326,21 @@ impl CacheAwarePolicy {
|
|||||||
);
|
);
|
||||||
}
|
}
|
||||||
|
|
||||||
// Use shortest queue when imbalanced
|
// Use shortest queue when imbalanced. Tie break randomly.
|
||||||
let min_load_idx = healthy_indices
|
// Snapshot load() (live atomic count of load). Without snapshot
|
||||||
|
// there could be no workers found matching min_load because of
|
||||||
|
// load update.
|
||||||
|
let loads: Vec<(usize, usize)> = healthy_indices
|
||||||
.iter()
|
.iter()
|
||||||
.min_by_key(|&&idx| workers[idx].load())
|
.map(|&idx| (idx, workers[idx].load()))
|
||||||
.copied()?;
|
.collect();
|
||||||
|
let min_load = loads.iter().map(|&(_, load)| load).min()?;
|
||||||
|
let min_load_idx = loads
|
||||||
|
.iter()
|
||||||
|
.copied()
|
||||||
|
.filter(|&(_, load)| load == min_load)
|
||||||
|
.map(|(idx, _)| idx)
|
||||||
|
.choose(&mut rand::rng())?;
|
||||||
|
|
||||||
// Even in imbalanced mode, update the tree to maintain cache state
|
// Even in imbalanced mode, update the tree to maintain cache state
|
||||||
if let Some(text) = request_text {
|
if let Some(text) = request_text {
|
||||||
@@ -438,11 +449,21 @@ impl LoadBalancingPolicy for CacheAwarePolicy {
|
|||||||
.position(|w| w.url() == tenant_url)
|
.position(|w| w.url() == tenant_url)
|
||||||
.filter(|&idx| workers[idx].is_healthy())
|
.filter(|&idx| workers[idx].is_healthy())
|
||||||
} else {
|
} else {
|
||||||
// Low cache match: use worker with minimum load
|
// Low cache match: use worker with minimum load. Tie break randomly.
|
||||||
healthy_indices
|
// Snapshot load() (live atomic count of load). Without snapshot
|
||||||
|
// there could be no workers found matching min_load because of
|
||||||
|
// load update.
|
||||||
|
let loads: Vec<(usize, usize)> = healthy_indices
|
||||||
|
.iter()
|
||||||
|
.map(|&idx| (idx, workers[idx].load()))
|
||||||
|
.collect();
|
||||||
|
let min_load = loads.iter().map(|&(_, load)| load).min()?;
|
||||||
|
loads
|
||||||
.iter()
|
.iter()
|
||||||
.min_by_key(|&&idx| workers[idx].load())
|
|
||||||
.copied()
|
.copied()
|
||||||
|
.filter(|&(_, load)| load == min_load)
|
||||||
|
.map(|(idx, _)| idx)
|
||||||
|
.choose(&mut rand::rng())
|
||||||
};
|
};
|
||||||
|
|
||||||
if let Some(idx) = selected_idx {
|
if let Some(idx) = selected_idx {
|
||||||
@@ -642,6 +663,113 @@ mod tests {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// In imbalanced mode the overloaded worker must be avoided AND the remaining
|
||||||
|
// tied (min-load) workers must be spread across via random tie-breaking.
|
||||||
|
#[tokio::test]
|
||||||
|
async fn test_cache_aware_imbalanced_random_tie_break() {
|
||||||
|
let policy = CacheAwarePolicy::with_config(CacheAwareConfig {
|
||||||
|
cache_threshold: 0.5,
|
||||||
|
balance_abs_threshold: 5,
|
||||||
|
balance_rel_threshold: 2.0,
|
||||||
|
eviction_interval_secs: 0,
|
||||||
|
max_tree_size: 10000,
|
||||||
|
});
|
||||||
|
|
||||||
|
let num_workers = 5;
|
||||||
|
let mut workers: Vec<Arc<dyn Worker>> = Vec::new();
|
||||||
|
for j in 0..num_workers {
|
||||||
|
workers.push(Arc::new(
|
||||||
|
BasicWorkerBuilder::new(&format!("http://w{}:8000", j))
|
||||||
|
.worker_type(WorkerType::Regular)
|
||||||
|
.build(),
|
||||||
|
));
|
||||||
|
}
|
||||||
|
|
||||||
|
// Overload worker 0: max=50, min=0 => (50-0) > 5 AND 50 > 0*2.0 => imbalanced.
|
||||||
|
for _ in 0..50 {
|
||||||
|
workers[0].increment_load();
|
||||||
|
}
|
||||||
|
policy.init_workers(&workers);
|
||||||
|
|
||||||
|
// Reuse the SAME prompt: the imbalanced branch bypasses cache affinity,
|
||||||
|
// so even a guaranteed cache hit must not pin all traffic to one worker.
|
||||||
|
let mut selection_counts = vec![0; num_workers];
|
||||||
|
for _ in 0..100 {
|
||||||
|
let info = SelectWorkerInfo {
|
||||||
|
request_text: Some("same_prompt"),
|
||||||
|
..Default::default()
|
||||||
|
};
|
||||||
|
let idx = policy.select_worker(&workers, &info).await.unwrap();
|
||||||
|
selection_counts[idx] += 1;
|
||||||
|
}
|
||||||
|
|
||||||
|
// The overloaded worker is never selected in imbalanced mode.
|
||||||
|
assert_eq!(
|
||||||
|
selection_counts[0], 0,
|
||||||
|
"overloaded worker was selected in imbalanced mode: {:?}",
|
||||||
|
selection_counts
|
||||||
|
);
|
||||||
|
// Every tied (load-0) worker is selected at least once.
|
||||||
|
for idx in 1..num_workers {
|
||||||
|
assert!(
|
||||||
|
selection_counts[idx] > 0,
|
||||||
|
"tied worker {} was never selected in imbalanced mode: {:?}",
|
||||||
|
idx,
|
||||||
|
selection_counts
|
||||||
|
);
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Verify random tie breaking for cache misses and low load situations.
|
||||||
|
// Important for when concurrency is lower than the number of workers.
|
||||||
|
// Without random tie breaking there will be multiple workers with 0 load
|
||||||
|
// and all requests will be sent to the same replica which doesn't utilize
|
||||||
|
// the available memory on all workers for the KV cache.
|
||||||
|
#[tokio::test]
|
||||||
|
async fn test_cache_aware_random_tie_break_cold_start() {
|
||||||
|
let config = CacheAwareConfig {
|
||||||
|
eviction_interval_secs: 0,
|
||||||
|
..Default::default()
|
||||||
|
};
|
||||||
|
let policy = CacheAwarePolicy::with_config(config);
|
||||||
|
|
||||||
|
// Create 5 workers
|
||||||
|
let num_workers = 5;
|
||||||
|
let mut workers: Vec<Arc<dyn Worker>> = Vec::new();
|
||||||
|
for j in 0..num_workers {
|
||||||
|
workers.push(Arc::new(
|
||||||
|
BasicWorkerBuilder::new(&format!("http://w{}:8000", j))
|
||||||
|
.worker_type(WorkerType::Regular)
|
||||||
|
.build(),
|
||||||
|
));
|
||||||
|
}
|
||||||
|
policy.init_workers(&workers);
|
||||||
|
|
||||||
|
// Send 100 requests with unique prompts to simulate 100 cache misses.
|
||||||
|
let mut selection_counts = vec![0; num_workers];
|
||||||
|
for i in 0..100 {
|
||||||
|
let prompt = format!("{}_unique_request", i);
|
||||||
|
let info = SelectWorkerInfo {
|
||||||
|
request_text: Some(&prompt),
|
||||||
|
..Default::default()
|
||||||
|
};
|
||||||
|
let idx = policy.select_worker(&workers, &info).await.unwrap();
|
||||||
|
selection_counts[idx] += 1;
|
||||||
|
}
|
||||||
|
|
||||||
|
// Verify all workers were selected at least once when prompts are
|
||||||
|
// unique.
|
||||||
|
for (idx, &count) in selection_counts.iter().enumerate() {
|
||||||
|
assert!(
|
||||||
|
count > 0,
|
||||||
|
"Worker {} was never selected; tie-breaking failed to spread across \
|
||||||
|
all tied workers: {:?}",
|
||||||
|
idx,
|
||||||
|
selection_counts
|
||||||
|
);
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
#[tokio::test]
|
#[tokio::test]
|
||||||
async fn test_cache_aware_worker_removal() {
|
async fn test_cache_aware_worker_removal() {
|
||||||
let config = CacheAwareConfig {
|
let config = CacheAwareConfig {
|
||||||
@@ -941,6 +1069,7 @@ mod tests {
|
|||||||
async fn test_pd_pool_isolation_two_policies() {
|
async fn test_pd_pool_isolation_two_policies() {
|
||||||
let config = CacheAwareConfig {
|
let config = CacheAwareConfig {
|
||||||
eviction_interval_secs: 0,
|
eviction_interval_secs: 0,
|
||||||
|
cache_threshold: 0.0,
|
||||||
..Default::default()
|
..Default::default()
|
||||||
};
|
};
|
||||||
let prefill_policy = CacheAwarePolicy::with_config(config.clone());
|
let prefill_policy = CacheAwarePolicy::with_config(config.clone());
|
||||||
@@ -1005,6 +1134,7 @@ mod tests {
|
|||||||
async fn test_pd_pool_isolation_shared_policy_regression() {
|
async fn test_pd_pool_isolation_shared_policy_regression() {
|
||||||
let config = CacheAwareConfig {
|
let config = CacheAwareConfig {
|
||||||
eviction_interval_secs: 0,
|
eviction_interval_secs: 0,
|
||||||
|
cache_threshold: 0.0,
|
||||||
..Default::default()
|
..Default::default()
|
||||||
};
|
};
|
||||||
let policy = CacheAwarePolicy::with_config(config);
|
let policy = CacheAwarePolicy::with_config(config);
|
||||||
|
|||||||
Reference in New Issue
Block a user