1use std::sync::Arc;
7use std::time::Duration;
8
9use anyhow::{Result, bail, ensure};
10use serde::{Deserialize, Serialize};
11use serde_json::Value;
12
13use crate::engine::common::perf_model::{polynomial_decode_time, polynomial_prefill_time};
14
15#[derive(Debug, Clone, Default, PartialEq, Serialize, Deserialize)]
17#[serde(rename_all = "snake_case", tag = "type")]
18pub enum TimingModelConfig {
19 #[default]
21 Polynomial,
22 Fixed { prefill_ms: f64, decode_ms: f64 },
24 External {
26 provider: String,
27 #[serde(default)]
28 config: Value,
29 },
30}
31
32pub trait TimingModel: Send + Sync {
37 fn predict_prefill_ms(
39 &self,
40 batch_size: usize,
41 mean_isl: usize,
42 mean_prefix: usize,
43 ) -> Result<f64>;
44
45 fn predict_decode_ms(
47 &self,
48 batch_size: usize,
49 active_kv_tokens: usize,
50 mean_context_length: usize,
51 total_kv_tokens: usize,
52 ) -> Result<f64>;
53}
54
55struct PolynomialTimingModel;
56
57impl TimingModel for PolynomialTimingModel {
58 fn predict_prefill_ms(
59 &self,
60 batch_size: usize,
61 mean_isl: usize,
62 mean_prefix: usize,
63 ) -> Result<f64> {
64 Ok(polynomial_prefill_time(
65 batch_size,
66 mean_isl.saturating_sub(mean_prefix),
67 ))
68 }
69
70 fn predict_decode_ms(
71 &self,
72 batch_size: usize,
73 active_kv_tokens: usize,
74 _mean_context_length: usize,
75 total_kv_tokens: usize,
76 ) -> Result<f64> {
77 if batch_size == 0 {
78 return Ok(0.0);
79 }
80 Ok(polynomial_decode_time(active_kv_tokens, total_kv_tokens))
81 }
82}
83
84struct FixedTimingModel {
85 prefill_ms: f64,
86 decode_ms: f64,
87}
88
89impl TimingModel for FixedTimingModel {
90 fn predict_prefill_ms(
91 &self,
92 batch_size: usize,
93 _mean_isl: usize,
94 _mean_prefix: usize,
95 ) -> Result<f64> {
96 Ok(if batch_size == 0 {
97 0.0
98 } else {
99 self.prefill_ms
100 })
101 }
102
103 fn predict_decode_ms(
104 &self,
105 batch_size: usize,
106 _active_kv_tokens: usize,
107 _mean_context_length: usize,
108 _total_kv_tokens: usize,
109 ) -> Result<f64> {
110 Ok(if batch_size == 0 { 0.0 } else { self.decode_ms })
111 }
112}
113
114pub(crate) fn built_in_timing_model(config: &TimingModelConfig) -> Result<Arc<dyn TimingModel>> {
115 match config {
116 TimingModelConfig::Polynomial => Ok(Arc::new(PolynomialTimingModel)),
117 TimingModelConfig::Fixed {
118 prefill_ms,
119 decode_ms,
120 } => Ok(Arc::new(FixedTimingModel {
121 prefill_ms: *prefill_ms,
122 decode_ms: *decode_ms,
123 })),
124 TimingModelConfig::External { provider, .. } => {
125 bail!("timing provider '{provider}' requires EngineFactory::with_timing_model")
126 }
127 }
128}
129
130pub(crate) fn modeled_duration_ms(raw_ms: f64, speedup_ratio: f64) -> Result<f64> {
131 ensure!(
132 raw_ms.is_finite() && raw_ms >= 0.0,
133 "timing provider returned invalid duration {raw_ms}ms"
134 );
135 ensure!(
136 speedup_ratio.is_finite() && speedup_ratio >= 0.0,
137 "modeled speedup ratio must be finite and non-negative, got {speedup_ratio}"
138 );
139 let unscaled = Duration::try_from_secs_f64(raw_ms / 1_000.0)
140 .map_err(|error| anyhow::anyhow!("timing duration {raw_ms}ms is out of range: {error}"))?;
141 let modeled = if speedup_ratio > 0.0 && unscaled > Duration::ZERO {
142 Duration::try_from_secs_f64(unscaled.as_secs_f64() / speedup_ratio).map_err(|error| {
143 anyhow::anyhow!(
144 "scaled timing duration is out of range for speedup {speedup_ratio}: {error}"
145 )
146 })?
147 } else {
148 unscaled
149 };
150 Ok(modeled.as_secs_f64() * 1_000.0)
151}
152
153#[cfg(test)]
154mod tests {
155 use super::*;
156
157 #[test]
158 fn modeled_duration_applies_speedup_and_zero_means_unscaled() {
159 assert_eq!(modeled_duration_ms(12.0, 3.0).unwrap(), 4.0);
160 assert_eq!(modeled_duration_ms(12.0, 0.0).unwrap(), 12.0);
161 assert_eq!(modeled_duration_ms(0.0, 3.0).unwrap(), 0.0);
162 }
163
164 #[test]
165 fn modeled_duration_rejects_invalid_provider_values() {
166 for raw_ms in [f64::NAN, f64::INFINITY, -1.0] {
167 assert!(modeled_duration_ms(raw_ms, 1.0).is_err());
168 }
169 for speedup in [f64::NAN, f64::INFINITY, -1.0] {
170 assert!(modeled_duration_ms(1.0, speedup).is_err());
171 }
172 }
173
174 #[test]
175 fn fixed_model_returns_zero_for_empty_batches() {
176 let model = built_in_timing_model(&TimingModelConfig::Fixed {
177 prefill_ms: 7.0,
178 decode_ms: 3.0,
179 })
180 .unwrap();
181 assert_eq!(model.predict_prefill_ms(0, 128, 0).unwrap(), 0.0);
182 assert_eq!(model.predict_decode_ms(0, 128, 64, 1024).unwrap(), 0.0);
183 assert_eq!(model.predict_prefill_ms(2, 128, 0).unwrap(), 7.0);
184 assert_eq!(model.predict_decode_ms(2, 128, 64, 1024).unwrap(), 3.0);
185 }
186
187 #[test]
188 fn external_provider_requires_runner_resolution() {
189 let error = built_in_timing_model(&TimingModelConfig::External {
190 provider: "example".to_string(),
191 config: Value::Null,
192 })
193 .err()
194 .expect("external descriptor cannot be materialized in the neutral crate");
195 assert!(
196 error
197 .to_string()
198 .contains("EngineFactory::with_timing_model")
199 );
200 }
201}