From 71c29bc5968bcf6bd9739534da9f3cd755c1b344 Mon Sep 17 00:00:00 2001 From: Simo Lin <25425177+slin1237@users.noreply.github.com> Date: Wed, 10 Jun 2026 15:08:49 -0700 Subject: [PATCH] feat(policies): route least_load by token-work expected-wait Score each healthy worker by its estimated time-to-drain plus a convex KV-pressure barrier, and route to the lowest (argmin): score = (queued_tokens + inflight_tokens) / throughput + kv_pressure_weight * k/(1-k) - queued_tokens: new num_waiting_uncached_tokens load signal, wired through the sglang scheduler proto + servicer and the shared SchedulerLoadSnapshot; defaults to 0 for backends that do not report it. - inflight_tokens: per-worker token-work this router dispatched since the last poll (stale-snapshot correction), reset on update_loads. - throughput: the worker's gen_throughput when reported, else the configurable default_throughput (backends without a live generation rate). - k/(1-k): M/M/1 KV-pressure barrier. Both terms are in seconds, so they add directly. Missing signals degrade to in-flight + barrier, then to JSQ. Token-work, not request count, reflects the load a worker actually carries, so size-skewed traffic is spread by work rather than by count. Tuning knobs are exposed with matching defaults via PolicyConfig, the smg CLI, and the Python router CLI/binding: kv_pressure_weight (0.15s), default_throughput (2000 tok/s), mean_prefill_tokens (1024). Signed-off-by: Simo Lin <25425177+slin1237@users.noreply.github.com> --- bindings/python/src/lib.rs | 16 +- bindings/python/src/smg/router_args.py | 27 ++ .../grpc_client/proto/sglang_scheduler.proto | 6 + crates/grpc_client/src/sglang_scheduler.rs | 1 + .../grpc_client/src/tokenspeed_scheduler.rs | 2 + crates/protocols/src/worker.rs | 16 + .../smg_grpc_servicer/sglang/servicer.py | 2 + model_gateway/src/config/types.rs | 30 +- model_gateway/src/config/validation.rs | 18 + model_gateway/src/main.rs | 18 +- model_gateway/src/policies/factory.rs | 9 +- model_gateway/src/policies/least_load.rs | 376 ++++++++++++++++-- model_gateway/src/policies/power_of_two.rs | 1 + 13 files changed, 477 insertions(+), 45 deletions(-) diff --git a/bindings/python/src/lib.rs b/bindings/python/src/lib.rs index 4e0da1cd21..6fe69b6ca0 100644 --- a/bindings/python/src/lib.rs +++ b/bindings/python/src/lib.rs @@ -377,6 +377,9 @@ struct Router { eviction_interval_secs: u64, max_tree_size: usize, block_size: usize, + least_load_kv_pressure_weight: f64, + least_load_default_throughput: f64, + least_load_mean_prefill_tokens: u32, max_idle_secs: u64, assignment_mode: String, max_payload_size: usize, @@ -518,7 +521,9 @@ impl Router { }, PolicyType::LeastLoad => ConfigPolicyConfig::LeastLoad { load_check_interval_secs: 5, - kv_pressure_weight: 1.5, + kv_pressure_weight: self.least_load_kv_pressure_weight, + mean_prefill_tokens: self.least_load_mean_prefill_tokens, + default_throughput: self.least_load_default_throughput, }, PolicyType::Bucket => ConfigPolicyConfig::Bucket { balance_abs_threshold: self.balance_abs_threshold, @@ -775,6 +780,9 @@ impl Router { eviction_interval_secs = 120, max_tree_size = 2usize.pow(26), block_size = 16, + least_load_kv_pressure_weight = 0.15, + least_load_default_throughput = 2000.0, + least_load_mean_prefill_tokens = 1024, max_idle_secs = 14400, assignment_mode = String::from("random"), max_payload_size = 512 * 1024 * 1024, @@ -885,6 +893,9 @@ impl Router { eviction_interval_secs: u64, max_tree_size: usize, block_size: usize, + least_load_kv_pressure_weight: f64, + least_load_default_throughput: f64, + least_load_mean_prefill_tokens: u32, max_idle_secs: u64, assignment_mode: String, max_payload_size: usize, @@ -1004,6 +1015,9 @@ impl Router { eviction_interval_secs, max_tree_size, block_size, + least_load_kv_pressure_weight, + least_load_default_throughput, + least_load_mean_prefill_tokens, max_idle_secs, assignment_mode, max_payload_size, diff --git a/bindings/python/src/smg/router_args.py b/bindings/python/src/smg/router_args.py index d1e6f56f5e..a42b3b9b52 100644 --- a/bindings/python/src/smg/router_args.py +++ b/bindings/python/src/smg/router_args.py @@ -51,6 +51,9 @@ class RouterArgs: eviction_interval_secs: int = 60 max_tree_size: int = 2**26 block_size: int = 16 + least_load_kv_pressure_weight: float = 0.15 + least_load_default_throughput: float = 2000.0 + least_load_mean_prefill_tokens: int = 1024 max_idle_secs: int = 4 * 3600 assignment_mode: str = "random" # Mode for manual policy new routing key assignment max_payload_size: int = 512 * 1024 * 1024 # 512MB default for large batches @@ -332,6 +335,30 @@ def add_cli_args( default=RouterArgs.cache_threshold, help="Cache threshold (0.0-1.0) for cache-aware routing", ) + routing_group.add_argument( + f"--{prefix}least-load-kv-pressure-weight", + type=float, + default=RouterArgs.least_load_kv_pressure_weight, + help="KV-pressure weight (seconds) for the least_load policy", + ) + routing_group.add_argument( + f"--{prefix}least-load-default-throughput", + type=float, + default=RouterArgs.least_load_default_throughput, + help=( + "Fallback generation throughput (tokens/s) for least_load when a" + " backend reports no live throughput" + ), + ) + routing_group.add_argument( + f"--{prefix}least-load-mean-prefill-tokens", + type=int, + default=RouterArgs.least_load_mean_prefill_tokens, + help=( + "Mean prefill tokens for least_load's in-flight estimate when a" + " request's token count is unknown at routing" + ), + ) routing_group.add_argument( f"--{prefix}balance-abs-threshold", type=int, diff --git a/crates/grpc_client/proto/sglang_scheduler.proto b/crates/grpc_client/proto/sglang_scheduler.proto index 1e092862c9..b93a56a5c9 100644 --- a/crates/grpc_client/proto/sglang_scheduler.proto +++ b/crates/grpc_client/proto/sglang_scheduler.proto @@ -518,6 +518,12 @@ message SchedulerLoad { double utilization = 10; int32 max_running_requests = 11; + // Queued token-work: waiting-queue tokens not yet served from cache + // (SGLang num_waiting_uncached_tokens). Field 17 is intentionally out of + // order — appended after the optional-section field numbers to preserve + // wire compatibility with existing readers. + int32 num_waiting_uncached_tokens = 17; + // Optional sections optional MemoryMetrics memory = 12; optional SpeculativeMetrics speculative = 13; diff --git a/crates/grpc_client/src/sglang_scheduler.rs b/crates/grpc_client/src/sglang_scheduler.rs index 88d2247163..531208c5cb 100644 --- a/crates/grpc_client/src/sglang_scheduler.rs +++ b/crates/grpc_client/src/sglang_scheduler.rs @@ -703,6 +703,7 @@ impl From for openai_protocol::worker::SchedulerLoadSnapsh dp_rank: load.dp_rank, num_running_reqs: load.num_running_reqs, num_waiting_reqs: load.num_waiting_reqs, + num_waiting_uncached_tokens: load.num_waiting_uncached_tokens, num_total_reqs: load.num_total_reqs, num_used_tokens: load.num_used_tokens, max_total_num_tokens: load.max_total_num_tokens, diff --git a/crates/grpc_client/src/tokenspeed_scheduler.rs b/crates/grpc_client/src/tokenspeed_scheduler.rs index ff3e8e18a3..aa1738c452 100644 --- a/crates/grpc_client/src/tokenspeed_scheduler.rs +++ b/crates/grpc_client/src/tokenspeed_scheduler.rs @@ -666,6 +666,8 @@ impl From for openai_protocol::worker::Schedule dp_rank: load.dp_rank, num_running_reqs: load.num_running_reqs, num_waiting_reqs: load.num_waiting_reqs, + // TokenSpeed does not report queued token-work; degrade to 0. + num_waiting_uncached_tokens: 0, num_total_reqs: load.num_total_reqs, num_used_tokens: load.num_used_tokens, max_total_num_tokens: load.max_total_num_tokens, diff --git a/crates/protocols/src/worker.rs b/crates/protocols/src/worker.rs index adc8d944d8..2c69e278f3 100644 --- a/crates/protocols/src/worker.rs +++ b/crates/protocols/src/worker.rs @@ -1017,6 +1017,9 @@ pub struct SchedulerLoadSnapshot { pub dp_rank: i32, pub num_running_reqs: i32, pub num_waiting_reqs: i32, + /// Queued token-work: waiting-queue tokens not yet served from cache. 0 when + /// the backend does not report it — callers degrade gracefully. + pub num_waiting_uncached_tokens: i32, pub num_total_reqs: i32, pub num_used_tokens: i32, pub max_total_num_tokens: i32, @@ -1051,6 +1054,19 @@ impl WorkerLoadResponse { self.loads.iter().map(|l| l.num_used_tokens as i64).sum() } + /// Total queued (waiting, uncached) tokens summed across all DP ranks. + pub fn total_waiting_uncached_tokens(&self) -> i64 { + self.loads + .iter() + .map(|l| l.num_waiting_uncached_tokens as i64) + .sum() + } + + /// Total generation throughput (tokens/s) summed across all DP ranks. + pub fn total_gen_throughput(&self) -> f64 { + self.loads.iter().map(|l| l.gen_throughput).sum() + } + pub fn dp_rank_loads(&self) -> HashMap { let mut map = HashMap::new(); for snapshot in &self.loads { diff --git a/grpc_servicer/smg_grpc_servicer/sglang/servicer.py b/grpc_servicer/smg_grpc_servicer/sglang/servicer.py index 33c5ad805b..51b19829d5 100644 --- a/grpc_servicer/smg_grpc_servicer/sglang/servicer.py +++ b/grpc_servicer/smg_grpc_servicer/sglang/servicer.py @@ -90,6 +90,8 @@ def _convert_loads_to_protobuf( cache_hit_rate=result.cache_hit_rate, utilization=result.utilization, max_running_requests=result.max_running_requests, + # Queued token-work: waiting-queue tokens not served from cache. + num_waiting_uncached_tokens=result.num_waiting_uncached_tokens, ) # Add optional sections using CopyFrom for proper protobuf assignment diff --git a/model_gateway/src/config/types.rs b/model_gateway/src/config/types.rs index f3c2813fbb..0bd077e070 100644 --- a/model_gateway/src/config/types.rs +++ b/model_gateway/src/config/types.rs @@ -312,17 +312,27 @@ pub enum PolicyConfig { #[serde(rename = "power_of_two")] PowerOfTwo { load_check_interval_secs: u64 }, - /// Least-load policy: routes to the worker minimizing - /// `in_flight + kv_pressure_weight * k/(1-k)` — the real-time in-flight request count plus - /// a convex KV-cache pressure term from the load monitor. - /// See `policies/least_load.rs`. + /// Least-(token-)work policy: routes to the worker minimizing the expected + /// wait `(queued_tokens + inflight_tokens) / throughput + kv_pressure_weight * k/(1-k)` + /// — token-work drain time plus a convex KV-cache pressure barrier, computed + /// from the load monitor with in-flight correction. See `policies/least_load.rs`. #[serde(rename = "least_load")] LeastLoad { #[serde(default = "default_least_load_interval")] load_check_interval_secs: u64, - /// KV-pressure weight (request-equivalents per unit of M/M/1 congestion). + /// KV-pressure weight `λ_t` (seconds): the time-cost of KV contention, + /// commensurate with the expected-queue-wait term. #[serde(default = "default_least_load_kv_pressure_weight")] kv_pressure_weight: f64, + /// Mean prefill length (tokens) used to estimate in-flight token-work + /// when a request's token count is unknown at routing time. + #[serde(default = "default_least_load_mean_prefill")] + mean_prefill_tokens: u32, + /// Fallback generation throughput (tokens/s) for the expected-wait term + /// when a backend reports no live `gen_throughput`. Set to the fleet's + /// per-replica generation rate; co-tunes with `kv_pressure_weight`. + #[serde(default = "default_least_load_throughput")] + default_throughput: f64, }, #[serde(rename = "bucket")] @@ -402,7 +412,15 @@ fn default_least_load_interval() -> u64 { } fn default_least_load_kv_pressure_weight() -> f64 { - 1.5 + 0.15 +} + +fn default_least_load_mean_prefill() -> u32 { + 1024 +} + +fn default_least_load_throughput() -> f64 { + 2000.0 } impl PolicyConfig { diff --git a/model_gateway/src/config/validation.rs b/model_gateway/src/config/validation.rs index 5ab5752584..643522af8a 100644 --- a/model_gateway/src/config/validation.rs +++ b/model_gateway/src/config/validation.rs @@ -303,6 +303,8 @@ impl ConfigValidator { PolicyConfig::LeastLoad { load_check_interval_secs, kv_pressure_weight, + mean_prefill_tokens, + default_throughput, } => { if *load_check_interval_secs == 0 { return Err(ConfigError::InvalidValue { @@ -319,6 +321,22 @@ impl ConfigValidator { reason: "Must be finite and >= 0.0".to_string(), }); } + + if *mean_prefill_tokens == 0 { + return Err(ConfigError::InvalidValue { + field: "mean_prefill_tokens".to_string(), + value: mean_prefill_tokens.to_string(), + reason: "Must be > 0".to_string(), + }); + } + + if !default_throughput.is_finite() || *default_throughput <= 0.0 { + return Err(ConfigError::InvalidValue { + field: "default_throughput".to_string(), + value: default_throughput.to_string(), + reason: "Must be finite and > 0.0".to_string(), + }); + } } PolicyConfig::Bucket { balance_abs_threshold: _, diff --git a/model_gateway/src/main.rs b/model_gateway/src/main.rs index 36d9c7111f..3708b17dc8 100644 --- a/model_gateway/src/main.rs +++ b/model_gateway/src/main.rs @@ -192,6 +192,20 @@ struct CliArgs { #[arg(long, default_value_t = 1.25, help_heading = "Routing Policy")] prefix_hash_load_factor: f64, + /// KV-pressure weight (seconds) for the least_load policy + #[arg(long, default_value_t = 0.15, help_heading = "Routing Policy")] + least_load_kv_pressure_weight: f64, + + /// Fallback generation throughput (tokens/s) for least_load when a backend + /// reports no live throughput + #[arg(long, default_value_t = 2000.0, help_heading = "Routing Policy")] + least_load_default_throughput: f64, + + /// Mean prefill tokens for least_load's in-flight estimate when a request's + /// token count is unknown at routing + #[arg(long, default_value_t = 1024, help_heading = "Routing Policy")] + least_load_mean_prefill_tokens: u32, + /// Enable data parallelism aware scheduling #[arg(long, default_value_t = false, help_heading = "Routing Policy")] dp_aware: bool, @@ -931,7 +945,9 @@ impl CliArgs { }, "least_load" => PolicyConfig::LeastLoad { load_check_interval_secs: 5, - kv_pressure_weight: 1.5, + kv_pressure_weight: self.least_load_kv_pressure_weight, + mean_prefill_tokens: self.least_load_mean_prefill_tokens, + default_throughput: self.least_load_default_throughput, }, "prefix_hash" => PolicyConfig::PrefixHash { prefix_token_count: self.prefix_token_count, diff --git a/model_gateway/src/policies/factory.rs b/model_gateway/src/policies/factory.rs index 479419d900..b694e28145 100644 --- a/model_gateway/src/policies/factory.rs +++ b/model_gateway/src/policies/factory.rs @@ -20,9 +20,14 @@ impl PolicyFactory { PolicyConfig::RoundRobin => Arc::new(RoundRobinPolicy::new()), PolicyConfig::PowerOfTwo { .. } => Arc::new(PowerOfTwoPolicy::new()), PolicyConfig::LeastLoad { - kv_pressure_weight, .. - } => Arc::new(LeastLoadPolicy::with_kv_pressure_weight( + kv_pressure_weight, + mean_prefill_tokens, + default_throughput, + .. + } => Arc::new(LeastLoadPolicy::with_params( *kv_pressure_weight, + *mean_prefill_tokens, + *default_throughput, )), PolicyConfig::CacheAware { cache_threshold, diff --git a/model_gateway/src/policies/least_load.rs b/model_gateway/src/policies/least_load.rs index 184d350d51..4067edec01 100644 --- a/model_gateway/src/policies/least_load.rs +++ b/model_gateway/src/policies/least_load.rs @@ -9,55 +9,184 @@ use tracing::debug; use super::{get_healthy_worker_indices, LoadBalancingPolicy, SelectWorkerInfo}; use crate::worker::Worker; -/// Default KV-pressure weight (request-equivalents per unit of M/M/1 congestion). -pub const DEFAULT_KV_PRESSURE_WEIGHT: f64 = 1.5; +/// Default KV-pressure weight `λ_t` (seconds): the time-cost of KV contention, +/// chosen commensurate with the expected-queue-wait term so the two add cleanly. +pub const DEFAULT_KV_PRESSURE_WEIGHT: f64 = 0.15; +/// Default mean prefill length (tokens), used to estimate in-flight token-work +/// for a dispatched request whose token count is unknown at routing time. +pub const DEFAULT_MEAN_PREFILL_TOKENS: u32 = 1024; + +/// Default fallback throughput (tokens/s) for the `/throughput` term when a +/// backend reports KV usage but no live `gen_throughput`. On a homogeneous +/// fleet its absolute value mainly sets the work-vs-barrier balance, so it +/// co-tunes with `kv_pressure_weight`. +pub const DEFAULT_THROUGHPUT: f64 = 2000.0; + +/// Least-(token-)work routing — route to the worker with the lowest estimated +/// time-to-drain plus a convex KV-pressure barrier (argmin, lower is better): +/// +/// ```text +/// score_i = (queued_tokens_i + inflight_tokens_i) / throughput_i +/// + kv_pressure_weight · k_i / (1 − k_i) +/// ``` +/// +/// - `queued_tokens` — the backend's waiting-queue token-work +/// (`num_waiting_uncached_tokens`). Token-work, not request count, is what +/// sets the wait under size-skewed traffic: a long prompt is far more work +/// than a short one, regardless of how many requests are queued. +/// - `inflight_tokens` — token-work this router has dispatched to the worker +/// since its last load poll. Polls are stale between intervals; without this +/// correction, plain argmin sends a whole interval's arrivals to one worker +/// (incast). Crediting each dispatch water-fills load across workers instead. +/// - `/ throughput` — normalizes work to *time*, comparing heterogeneous +/// workers by drain time rather than raw token count. +/// - `k / (1 − k)` — the M/M/1 expected-occupancy barrier on KV utilization +/// `k`; convex and divergent at the KV cliff, so routing avoids the +/// preemption/recompute that begins as KV fills. +/// +/// Both terms are in seconds, so they add directly. Missing signals degrade +/// gracefully and stay in time units: +/// - no queued-token report (backend doesn't expose waiting-queue tokens): +/// `queued_tokens = 0`, leaving in-flight-corrected drain time plus the barrier; +/// - zero/absent throughput (backend reports no generation rate): falls back to +/// the configured `default_throughput`, so the work term stays in seconds and +/// the KV barrier stays relevant; +/// - a worker with no fresh snapshot while peers report: its live in-flight is +/// converted to a drain-time estimate (`load · p̄ / fleet_nominal_throughput`) +/// so it is comparable to reporting workers, not scored on a raw count; +/// - the whole fleet dark (true cold start, or a backend that never reports +/// loads): join-shortest-queue on the live in-flight count. +/// +/// In-flight token-work is exact on the gRPC routing path (the request's token +/// count is known at selection); the HTTP path has no token count and falls +/// back to `p̄ · count`, which is weaker on size-skewed traffic. This policy is +/// therefore intended for gRPC workers. +/// +/// # Tuning knobs +/// +/// All are fields of `PolicyConfig::LeastLoad` with the defaults below: +/// - `kv_pressure_weight` (λ_t, default `0.15` s) — weight of the KV-pressure +/// barrier. Raise it to steer harder away from near-full KV; lower it to +/// weight raw drain time more. +/// - `default_throughput` (default `2000` tok/s) — drain rate used when a +/// backend reports no live `gen_throughput`. Set it to the fleet's measured +/// per-replica generation rate; it co-tunes with `kv_pressure_weight`. +/// - `mean_prefill_tokens` (p̄, default `1024`) — per-request token estimate for +/// the in-flight term when the request's token count is unknown at routing +/// (the HTTP path; ignored when tokens are known, i.e. gRPC). +/// - `load_check_interval_secs` (default `10`) — worker-load poll period; the +/// in-flight correction absorbs staleness between polls. #[derive(Debug)] pub struct LeastLoadPolicy { /// Cached load reports from the worker monitor (keyed by worker URL). cached_loads: RwLock>, - /// KV-pressure weight. + /// In-flight token-work dispatched per worker since its last load poll + /// (keyed by worker URL); reset when a fresh report arrives. + inflight_tokens: RwLock>, + /// KV-pressure weight `λ_t` (seconds). kv_pressure_weight: f64, + /// Mean prefill length (tokens) for estimating in-flight token-work when a + /// request's token count is unknown at routing time. + mean_prefill_tokens: u32, + /// Fallback throughput (tokens/s) for the `/throughput` term when a backend + /// reports no live `gen_throughput`. + default_throughput: f64, } impl LeastLoadPolicy { pub fn new() -> Self { - Self::with_kv_pressure_weight(DEFAULT_KV_PRESSURE_WEIGHT) + Self::with_params( + DEFAULT_KV_PRESSURE_WEIGHT, + DEFAULT_MEAN_PREFILL_TOKENS, + DEFAULT_THROUGHPUT, + ) } pub fn with_kv_pressure_weight(kv_pressure_weight: f64) -> Self { + Self::with_params( + kv_pressure_weight, + DEFAULT_MEAN_PREFILL_TOKENS, + DEFAULT_THROUGHPUT, + ) + } + + pub fn with_params( + kv_pressure_weight: f64, + mean_prefill_tokens: u32, + default_throughput: f64, + ) -> Self { Self { cached_loads: RwLock::new(HashMap::new()), + inflight_tokens: RwLock::new(HashMap::new()), kv_pressure_weight: if kv_pressure_weight.is_finite() && kv_pressure_weight >= 0.0 { kv_pressure_weight } else { DEFAULT_KV_PRESSURE_WEIGHT }, + mean_prefill_tokens: mean_prefill_tokens.max(1), + default_throughput: if default_throughput.is_finite() && default_throughput > 0.0 { + default_throughput + } else { + DEFAULT_THROUGHPUT + }, } } - /// Least-load score for a worker (lower is better). + /// Expected-wait score for a worker (lower is better). + /// + /// `inflight` maps worker URL -> token-work dispatched since its last poll. + /// `nominal_throughput` (a peer-derived mean) estimates drain rate for a + /// worker missing a fresh snapshot; `fleet_has_loads` is false only when no + /// worker reports at all, in which case we fall back to join-shortest-queue + /// on the live in-flight count (which, unlike the since-poll estimate, + /// reflects completions and so suits backends that never report loads). fn score( &self, worker: &Arc, loads: Option<&HashMap>, + inflight: &HashMap, + nominal_throughput: f64, + fleet_has_loads: bool, ) -> f64 { - let in_flight = worker.load() as f64; - // KV-cache utilization from the latest load report; absent -> 0 (no barrier). - let k = loads - .and_then(|m| m.get(worker.url())) - .map(|l| l.effective_token_usage().clamp(0.0, 0.999)) - .unwrap_or(0.0); - in_flight + self.kv_pressure_weight * k / (1.0 - k) + let url = worker.url(); + match loads.and_then(|m| m.get(url)) { + Some(load) => { + let inflight_tokens = inflight.get(url).copied().unwrap_or(0) as f64; + let queued_tokens = load.total_waiting_uncached_tokens() as f64; + let live_throughput = load.total_gen_throughput(); + let throughput = if live_throughput > 0.0 { + live_throughput + } else { + self.default_throughput + }; + let k = load.effective_token_usage().clamp(0.0, 0.999); + (queued_tokens + inflight_tokens) / throughput + + self.kv_pressure_weight * k / (1.0 - k) + } + // No fresh snapshot, but peers report: estimate this worker's drain + // time from its live in-flight (count × mean prefill) at the fleet's + // nominal throughput, keeping the same units as reporting workers. + None if fleet_has_loads => { + worker.load() as f64 * self.mean_prefill_tokens as f64 / nominal_throughput + } + // Whole fleet dark (cold start, or a backend that never reports + // loads): join-shortest-queue on live in-flight. + None => worker.load() as f64, + } + } + + /// Token-work the request being routed adds to the chosen worker's + /// in-flight estimate: its token count if known, else the mean prefill. + fn request_tokens(&self, info: &SelectWorkerInfo) -> u64 { + info.tokens + .map(|t| t.len() as u64) + .unwrap_or(self.mean_prefill_tokens as u64) } } impl LoadBalancingPolicy for LeastLoadPolicy { - fn select_worker( - &self, - workers: &[Arc], - _info: &SelectWorkerInfo, - ) -> Option { + fn select_worker(&self, workers: &[Arc], info: &SelectWorkerInfo) -> Option { let healthy = get_healthy_worker_indices(workers); if healthy.is_empty() { return None; @@ -66,19 +195,63 @@ impl LoadBalancingPolicy for LeastLoadPolicy { return Some(healthy[0]); } - let guard = self.cached_loads.read().ok(); - let loads = guard.as_deref(); + let loads_guard = self.cached_loads.read().ok(); + let loads = loads_guard.as_deref(); + + // Fleet-nominal throughput (mean of positive reports) stands in for a + // worker missing a fresh snapshot; `fleet_has_loads` distinguishes a + // partial gap (estimate that worker's drain time at the nominal rate) + // from a fully dark fleet (fall back to join-shortest-queue). + let (tp_sum, tp_count) = healthy + .iter() + .filter_map(|&i| loads.and_then(|m| m.get(workers[i].url()))) + .map(|l| l.total_gen_throughput()) + .filter(|t| *t > 0.0) + .fold((0.0, 0u32), |(s, n), t| (s + t, n + 1)); + let nominal_throughput = if tp_count > 0 { + tp_sum / tp_count as f64 + } else { + self.default_throughput + }; + let fleet_has_loads = loads + .map(|m| healthy.iter().any(|&i| m.contains_key(workers[i].url()))) + .unwrap_or(false); + + // Held across selection so the in-flight estimate stays consistent and + // the chosen worker can be credited before the guard is released. + let mut inflight = self + .inflight_tokens + .write() + .unwrap_or_else(|poisoned| poisoned.into_inner()); let mut best = healthy[0]; - let mut best_score = self.score(&workers[best], loads); + let mut best_score = self.score( + &workers[best], + loads, + &inflight, + nominal_throughput, + fleet_has_loads, + ); for &idx in &healthy[1..] { - let s = self.score(&workers[idx], loads); + let s = self.score( + &workers[idx], + loads, + &inflight, + nominal_throughput, + fleet_has_loads, + ); if s < best_score { best = idx; best_score = s; } } + // In-flight correction: credit the chosen worker with this request's + // token-work until its next poll refreshes the snapshot. + let req_tokens = self.request_tokens(info); + *inflight.entry(workers[best].url().to_string()).or_insert(0) += req_tokens; + drop(inflight); + debug!( "least_load selected {} (score {:.4}, in_flight {})", workers[best].url(), @@ -97,12 +270,22 @@ impl LoadBalancingPolicy for LeastLoadPolicy { if let Ok(mut cached) = self.cached_loads.write() { cached.extend(loads.iter().map(|(k, v)| (k.clone(), v.clone()))); } + // A fresh snapshot already reflects work up to the poll, so reset the + // since-poll in-flight estimate for the workers it covers. + if let Ok(mut inflight) = self.inflight_tokens.write() { + for url in loads.keys() { + inflight.insert(url.clone(), 0); + } + } } fn remove_worker(&self, url: &str) { if let Ok(mut cached) = self.cached_loads.write() { cached.remove(url); } + if let Ok(mut inflight) = self.inflight_tokens.write() { + inflight.remove(url); + } } fn as_any(&self) -> &dyn std::any::Any { @@ -130,7 +313,12 @@ mod tests { } } - fn make_load(token_usage: f64) -> WorkerLoadResponse { + /// One DP rank with the given queued tokens, KV utilization, and throughput. + fn make_load( + num_waiting_uncached_tokens: i32, + token_usage: f64, + gen_throughput: f64, + ) -> WorkerLoadResponse { WorkerLoadResponse { timestamp: String::new(), dp_rank_count: 1, @@ -138,11 +326,12 @@ mod tests { dp_rank: 0, num_running_reqs: 0, num_waiting_reqs: 0, + num_waiting_uncached_tokens, num_total_reqs: 0, num_used_tokens: 0, max_total_num_tokens: 0, token_usage, - gen_throughput: 0.0, + gen_throughput, cache_hit_rate: 0.0, utilization: 0.0, max_running_requests: 0, @@ -160,7 +349,8 @@ mod tests { } #[test] - fn picks_lowest_in_flight() { + fn cold_start_picks_lowest_in_flight() { + // No load reports yet -> join-shortest-queue on live in-flight count. let policy = LeastLoadPolicy::new(); let a = mk("http://a:8000"); let b = mk("http://b:8000"); @@ -168,31 +358,147 @@ mod tests { a.increment_load(); } let workers = vec![a, b]; - // No load reports -> pure in-flight; b (0) beats a (5). assert_eq!( policy.select_worker(&workers, &SelectWorkerInfo::default()), Some(1) ); } + #[test] + fn routes_to_lower_queued_token_work() { + // Equal KV/throughput; the worker with fewer queued tokens wins. + let policy = LeastLoadPolicy::new(); + let workers = vec![mk("http://a:8000"), mk("http://b:8000")]; + let mut loads = HashMap::new(); + loads.insert("http://a:8000".to_string(), make_load(8000, 0.2, 100.0)); + loads.insert("http://b:8000".to_string(), make_load(1000, 0.2, 100.0)); + policy.update_loads(&loads); + // a: 8000/100 = 80s ; b: 1000/100 = 10s -> pick b. + assert_eq!( + policy.select_worker(&workers, &SelectWorkerInfo::default()), + Some(1) + ); + } + + #[test] + fn throughput_normalization_prefers_faster_worker() { + // Same queued tokens; the faster worker (higher throughput) drains sooner. + let policy = LeastLoadPolicy::new(); + let workers = vec![mk("http://a:8000"), mk("http://b:8000")]; + let mut loads = HashMap::new(); + loads.insert("http://a:8000".to_string(), make_load(5000, 0.2, 50.0)); + loads.insert("http://b:8000".to_string(), make_load(5000, 0.2, 500.0)); + policy.update_loads(&loads); + // a: 5000/50 = 100s ; b: 5000/500 = 10s -> pick b. + assert_eq!( + policy.select_worker(&workers, &SelectWorkerInfo::default()), + Some(1) + ); + } + + #[test] + fn zero_throughput_falls_back_to_default() { + // A backend that reports no gen_throughput (0); the score must still + // discriminate via the configured default_throughput, not collapse. + let policy = LeastLoadPolicy::new(); // default_throughput = 2000 + let workers = vec![mk("http://a:8000"), mk("http://b:8000")]; + let mut loads = HashMap::new(); + loads.insert("http://a:8000".to_string(), make_load(10000, 0.2, 0.0)); + loads.insert("http://b:8000".to_string(), make_load(1000, 0.2, 0.0)); + policy.update_loads(&loads); + // gen_throughput=0 -> default 2000: a 10000/2000=5s ; b 1000/2000=0.5s -> pick b. + assert_eq!( + policy.select_worker(&workers, &SelectWorkerInfo::default()), + Some(1) + ); + } + + #[test] + fn missing_snapshot_estimated_in_time_units() { + // Worker a reports ~40s of queued work; worker b has no snapshot but 5 + // live in-flight. Scoring b on raw count (5) would wrongly beat a's 40s; + // scoring it as drain time (5 * p̄ / nominal ≈ 51s) keeps the lighter a. + let policy = LeastLoadPolicy::new(); // p̄ = 1024 + let a = mk("http://a:8000"); + let b = mk("http://b:8000"); + for _ in 0..5 { + b.increment_load(); + } + let workers = vec![a, b]; + let mut loads = HashMap::new(); + loads.insert("http://a:8000".to_string(), make_load(4000, 0.0, 100.0)); + policy.update_loads(&loads); + // a: 4000/100 = 40s ; b: 5 * 1024 / 100 ≈ 51.2s -> pick a. + assert_eq!( + policy.select_worker(&workers, &SelectWorkerInfo::default()), + Some(0) + ); + } + #[test] fn kv_barrier_avoids_full_worker() { + // No queued work; the convex KV barrier steers off the near-full worker. let policy = LeastLoadPolicy::with_kv_pressure_weight(2.0); - let a = mk("http://a:8000"); // idle but KV-full - let b = mk("http://b:8000"); // 1 in-flight but KV empty - b.increment_load(); - let workers = vec![a, b]; + let workers = vec![mk("http://a:8000"), mk("http://b:8000")]; let mut loads = HashMap::new(); - loads.insert("http://a:8000".to_string(), make_load(0.95)); // barrier 2*0.95/0.05 = 38 - loads.insert("http://b:8000".to_string(), make_load(0.0)); + loads.insert("http://a:8000".to_string(), make_load(0, 0.98, 100.0)); + loads.insert("http://b:8000".to_string(), make_load(0, 0.0, 100.0)); policy.update_loads(&loads); - // a: 0 + 38 = 38 ; b: 1 + 0 = 1 -> pick b. + // a: 0 + 2*0.98/0.02 = 98 ; b: 0 -> pick b. assert_eq!( policy.select_worker(&workers, &SelectWorkerInfo::default()), Some(1) ); } + #[test] + fn inflight_correction_spreads_within_poll_interval() { + // Two identical workers, no fresh poll between dispatches: the in-flight + // token credit must push the second request to the other worker rather + // than herding both onto the first. + let policy = LeastLoadPolicy::new(); + let workers = vec![mk("http://a:8000"), mk("http://b:8000")]; + let mut loads = HashMap::new(); + loads.insert("http://a:8000".to_string(), make_load(0, 0.1, 100.0)); + loads.insert("http://b:8000".to_string(), make_load(0, 0.1, 100.0)); + policy.update_loads(&loads); + + let info = SelectWorkerInfo::default(); // tokens unknown -> mean prefill + let first = policy.select_worker(&workers, &info).unwrap(); + let second = policy.select_worker(&workers, &info).unwrap(); + assert_ne!(first, second); + } + + #[test] + fn update_loads_resets_inflight() { + let policy = LeastLoadPolicy::new(); + let workers = vec![mk("http://a:8000"), mk("http://b:8000")]; + let mut loads = HashMap::new(); + loads.insert("http://a:8000".to_string(), make_load(0, 0.1, 100.0)); + loads.insert("http://b:8000".to_string(), make_load(0, 0.1, 100.0)); + policy.update_loads(&loads); + + let info = SelectWorkerInfo::default(); + for _ in 0..4 { + policy.select_worker(&workers, &info); + } + assert!(policy + .inflight_tokens + .read() + .unwrap() + .values() + .any(|&v| v > 0)); + + // A fresh poll clears the since-poll estimate. + policy.update_loads(&loads); + assert!(policy + .inflight_tokens + .read() + .unwrap() + .values() + .all(|&v| v == 0)); + } + #[test] fn single_worker_always_selected() { let policy = LeastLoadPolicy::new(); @@ -204,11 +510,11 @@ mod tests { } #[test] - fn remove_worker_prunes_cached_load() { + fn remove_worker_prunes_state() { let policy = LeastLoadPolicy::new(); let mut loads = HashMap::new(); - loads.insert("http://a:8000".to_string(), make_load(0.5)); - loads.insert("http://b:8000".to_string(), make_load(0.3)); + loads.insert("http://a:8000".to_string(), make_load(0, 0.5, 100.0)); + loads.insert("http://b:8000".to_string(), make_load(0, 0.3, 100.0)); policy.update_loads(&loads); assert_eq!(policy.cached_loads.read().unwrap().len(), 2); diff --git a/model_gateway/src/policies/power_of_two.rs b/model_gateway/src/policies/power_of_two.rs index 0d8c3c3a00..1180b7316d 100644 --- a/model_gateway/src/policies/power_of_two.rs +++ b/model_gateway/src/policies/power_of_two.rs @@ -159,6 +159,7 @@ mod tests { dp_rank: 0, num_running_reqs: 0, num_waiting_reqs: 0, + num_waiting_uncached_tokens: 0, num_total_reqs: 0, num_used_tokens: 0, max_total_num_tokens: 0,