Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
5 changes: 5 additions & 0 deletions bindings/python/src/lib.rs
Original file line number Diff line number Diff line change
Expand Up @@ -397,6 +397,7 @@ struct Router {
overload_token_usage_threshold: f32,
prefix_token_count: usize,
prefix_hash_load_factor: f64,
prefix_hash_balance_abs_threshold: usize,
least_load_kv_pressure_weight: f64,
least_load_default_throughput: f64,
least_load_mean_prefill_tokens: u32,
Expand Down Expand Up @@ -610,6 +611,7 @@ impl Router {
PolicyType::PrefixHash => ConfigPolicyConfig::PrefixHash {
prefix_token_count: self.prefix_token_count,
load_factor: self.prefix_hash_load_factor,
balance_abs_threshold: self.prefix_hash_balance_abs_threshold,
},
})
};
Expand Down Expand Up @@ -1006,6 +1008,7 @@ impl Router {
zmq_engine_count = None,
prefix_token_count = 256,
prefix_hash_load_factor = 1.25,
prefix_hash_balance_abs_threshold = 10,
))]
#[expect(clippy::too_many_arguments)]
#[expect(
Expand Down Expand Up @@ -1138,6 +1141,7 @@ impl Router {
zmq_engine_count: Option<usize>,
prefix_token_count: usize,
prefix_hash_load_factor: f64,
prefix_hash_balance_abs_threshold: usize,
) -> PyResult<Self> {
let mut all_urls = worker_urls.clone();

Expand Down Expand Up @@ -1179,6 +1183,7 @@ impl Router {
overload_token_usage_threshold,
prefix_token_count,
prefix_hash_load_factor,
prefix_hash_balance_abs_threshold,
least_load_kv_pressure_weight,
least_load_default_throughput,
least_load_mean_prefill_tokens,
Expand Down
4 changes: 4 additions & 0 deletions bindings/python/src/smg/router.py
Original file line number Diff line number Diff line change
Expand Up @@ -172,6 +172,10 @@ class Router:
prefix_hash_load_factor: Load factor above which the prefix_hash policy
walks the ring instead of using the hashed worker (multiple of the
average load). Default: 1.25
prefix_hash_balance_abs_threshold: Absolute load difference over average
a worker must also exceed before prefix_hash treats it as
overloaded. Guards the load factor against sampling noise when each
router replica sees only a share of a worker's load. Default: 10
balance_abs_threshold: Load balancing is triggered when (max_load - min_load) >
abs_threshold AND max_load > min_load * rel_threshold. Otherwise, use cache
aware. Default: 32
Expand Down
10 changes: 10 additions & 0 deletions bindings/python/src/smg/router_args.py
Original file line number Diff line number Diff line change
Expand Up @@ -208,6 +208,7 @@ class RouterArgs:
zmq_engine_count: int | None = None
prefix_token_count: int = 256
prefix_hash_load_factor: float = 1.25
prefix_hash_balance_abs_threshold: int = 10

@staticmethod
def add_cli_args(
Expand Down Expand Up @@ -404,6 +405,15 @@ def add_cli_args(
default=RouterArgs.prefix_hash_load_factor,
help="Load factor above which prefix_hash walks the ring (multiple of average load)",
)
routing_group.add_argument(
f"--{prefix}prefix-hash-balance-abs-threshold",
type=int,
default=RouterArgs.prefix_hash_balance_abs_threshold,
help=(
"Absolute load difference over average a worker must also "
"exceed before prefix_hash treats it as overloaded"
),
)
Comment thread
coderabbitai[bot] marked this conversation as resolved.
routing_group.add_argument(
f"--{prefix}least-load-kv-pressure-weight",
type=float,
Expand Down
13 changes: 11 additions & 2 deletions model_gateway/src/config/types.rs
Original file line number Diff line number Diff line change
Expand Up @@ -555,17 +555,22 @@ pub enum PolicyConfig {
/// A lightweight alternative to cache_aware radix tree.
/// Routes requests based on prefix token hash for cache locality.
/// - Uses consistent hash ring with bounded load balancing
/// - Walks ring if worker is overloaded (load > avg * load_factor)
/// - Diverts to the least loaded worker when the hashed one is overloaded
/// - O(log n) lookup instead of O(prefix_len) radix tree traversal
#[serde(rename = "prefix_hash")]
PrefixHash {
/// Number of prefix tokens to hash, or four times as many characters
/// of the prompt when the request is untokenized (default: 256)
#[serde(default = "default_prefix_token_count")]
prefix_token_count: usize,
/// Load factor threshold - walk ring if load > avg * factor (default: 1.25)
/// Relative load threshold - a worker is overloaded when its load
/// exceeds both avg * factor and the absolute margin (default: 1.25)
#[serde(default = "default_load_factor")]
load_factor: f64,
/// Absolute load difference over average a worker must also exceed
/// before it counts as overloaded (default: 10)
#[serde(default = "default_prefix_hash_balance_abs_threshold")]
balance_abs_threshold: usize,
Comment thread
coderabbitai[bot] marked this conversation as resolved.
},
}

Expand All @@ -585,6 +590,10 @@ fn default_load_factor() -> f64 {
1.25
}

fn default_prefix_hash_balance_abs_threshold() -> usize {
10
}

fn default_manual_eviction_interval_secs() -> u64 {
60
}
Expand Down
1 change: 1 addition & 0 deletions model_gateway/src/config/validation.rs
Original file line number Diff line number Diff line change
Expand Up @@ -593,6 +593,7 @@ impl ConfigValidator {
PolicyConfig::PrefixHash {
prefix_token_count,
load_factor,
balance_abs_threshold: _,
} => {
if *prefix_token_count == 0 {
return Err(ConfigError::InvalidValue {
Expand Down
6 changes: 6 additions & 0 deletions model_gateway/src/main.rs
Original file line number Diff line number Diff line change
Expand Up @@ -254,6 +254,11 @@ struct CliArgs {
#[arg(long, default_value_t = 1.25, help_heading = "Routing Policy")]
prefix_hash_load_factor: f64,

/// Absolute load difference over average a worker must also exceed before
/// the prefix_hash policy treats it as overloaded
#[arg(long, default_value_t = 10, help_heading = "Routing Policy")]
prefix_hash_balance_abs_threshold: usize,

/// 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,
Expand Down Expand Up @@ -1152,6 +1157,7 @@ impl CliArgs {
"prefix_hash" => PolicyConfig::PrefixHash {
prefix_token_count: self.prefix_token_count,
load_factor: self.prefix_hash_load_factor,
balance_abs_threshold: self.prefix_hash_balance_abs_threshold,
},
"consistent_hashing" => PolicyConfig::ConsistentHashing,
"manual" => PolicyConfig::Manual {
Expand Down
3 changes: 3 additions & 0 deletions model_gateway/src/policies/factory.rs
Original file line number Diff line number Diff line change
Expand Up @@ -80,10 +80,12 @@ impl PolicyFactory {
PolicyConfig::PrefixHash {
prefix_token_count,
load_factor,
balance_abs_threshold,
} => {
let config = PrefixHashConfig {
prefix_token_count: *prefix_token_count,
load_factor: *load_factor,
balance_abs_threshold: *balance_abs_threshold,
};
Arc::new(PrefixHashPolicy::new(config))
}
Expand Down Expand Up @@ -162,6 +164,7 @@ mod tests {
let policy = PolicyFactory::create_from_config(&PolicyConfig::PrefixHash {
prefix_token_count: 100,
load_factor: 0.8,
balance_abs_threshold: 10,
});
assert_eq!(policy.name(), "prefix_hash");
}
Expand Down
65 changes: 57 additions & 8 deletions model_gateway/src/policies/prefix_hash.rs
Original file line number Diff line number Diff line change
Expand Up @@ -10,7 +10,8 @@
//! or the equivalent span of routing text when the request is untokenized
//! 2. Hash the prefix using xxhash for fast, stable hashing
//! 3. Use consistent hash ring to find the target worker
//! 4. If worker is overloaded (load > avg * load_factor), find least loaded
//! 4. If worker is overloaded (load above both the relative and absolute
//! margins over average), find least loaded
//! 5. Return least loaded worker that passes load check, or initial if all overloaded
//!
//! ## Complexity
Expand Down Expand Up @@ -45,18 +46,28 @@ pub struct PrefixHashConfig {
/// Default: 256 tokens (~1 paragraph of text)
pub prefix_token_count: usize,

/// Load factor threshold for walking the ring.
/// If a worker's load > (total_load / num_workers) * load_factor,
/// walk clockwise to the next worker.
/// Relative load threshold for the overload check.
/// A worker counts as overloaded once its load exceeds the average by
/// this multiple as well as by `balance_abs_threshold`, at which point
/// the request goes to the least loaded worker that is not overloaded.
/// Default: 1.25 (125% of average load)
pub load_factor: f64,

/// Absolute load difference threshold for the overload check.
/// A worker only counts as overloaded once its load exceeds the average
/// by this many requests as well as by `load_factor`. Without it the
/// check is purely relative, so its false-positive rate grows as each
/// router replica observes a smaller share of a worker's true load.
/// Default: 10 requests
pub balance_abs_threshold: usize,
}

impl Default for PrefixHashConfig {
fn default() -> Self {
Self {
prefix_token_count: 256,
load_factor: 1.25,
balance_abs_threshold: 10,
}
}
}
Expand Down Expand Up @@ -150,6 +161,9 @@ impl PrefixHashPolicy {
}

/// Check if a worker's load is acceptable
///
/// Overload requires clearing both the relative and the absolute margin,
/// mirroring the imbalance test cache_aware uses.
#[inline]
fn load_ok(&self, worker_load: usize, total_load: usize, num_workers: usize) -> bool {
if total_load == 0 || num_workers == 0 {
Expand All @@ -158,7 +172,8 @@ impl PrefixHashPolicy {

// Average load per worker (with +1 to simulate incoming request)
let avg_load = (total_load + 1) as f64 / num_workers as f64;
let threshold = avg_load * self.config.load_factor;
let threshold = (avg_load * self.config.load_factor)
.max(avg_load + self.config.balance_abs_threshold as f64);

(worker_load as f64) <= threshold
}
Expand Down Expand Up @@ -529,18 +544,52 @@ mod tests {
fn test_load_ok_calculation() {
let policy = PrefixHashPolicy::new(PrefixHashConfig {
load_factor: 1.25,
balance_abs_threshold: 0,
..Default::default()
});

// Total load 100, 4 workers -> avg 25, threshold 31.25
assert!(policy.load_ok(30, 100, 4)); // 30 <= 31.25
assert!(!policy.load_ok(35, 100, 4)); // 35 > 31.25
// Total load 100, 4 workers -> avg 25.25, threshold 31.5625
assert!(policy.load_ok(30, 100, 4));
assert!(!policy.load_ok(35, 100, 4));

// Edge cases
assert!(policy.load_ok(0, 0, 4)); // No load = OK
assert!(policy.load_ok(100, 0, 0)); // No workers = OK (shouldn't happen)
}

#[test]
fn test_absolute_margin_absorbs_small_count_noise() {
let policy = PrefixHashPolicy::new(PrefixHashConfig {
load_factor: 1.25,
balance_abs_threshold: 10,
..Default::default()
});

// Average 10.25 across 4 workers: the relative margin alone flags 13,
// but 13 is within 10 requests of average so it stays acceptable.
assert!(policy.load_ok(13, 40, 4));
assert!(!policy.load_ok(21, 40, 4));
}

#[test]
fn test_relative_margin_still_binds_at_high_load() {
let policy = PrefixHashPolicy::new(PrefixHashConfig {
load_factor: 1.25,
balance_abs_threshold: 10,
..Default::default()
});

// Average 200.25: the relative margin (250.3) now exceeds the absolute
// one (210.25), so it is the binding constraint.
assert!(policy.load_ok(240, 800, 4));
assert!(!policy.load_ok(260, 800, 4));
}

#[test]
fn test_absolute_margin_defaults_on() {
assert_eq!(PrefixHashConfig::default().balance_abs_threshold, 10);
}

#[test]
fn test_policy_name() {
let policy = PrefixHashPolicy::with_defaults();
Expand Down
1 change: 1 addition & 0 deletions model_gateway/tests/common/test_config.rs
Original file line number Diff line number Diff line change
Expand Up @@ -122,6 +122,7 @@ impl TestRouterConfig {
.policy(PolicyConfig::PrefixHash {
prefix_token_count,
load_factor: 1.25,
balance_abs_threshold: 10,
})
.host(defaults::HOST)
.port(port)
Expand Down
Loading