Skip to main content

aisimulate_core/engine/
timing.rs

1// SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
2// SPDX-License-Identifier: Apache-2.0
3
4//! Runtime-neutral forward-pass timing models.
5
6use 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/// Serializable timing-provider selection.
16#[derive(Debug, Clone, Default, PartialEq, Serialize, Deserialize)]
17#[serde(rename_all = "snake_case", tag = "type")]
18pub enum TimingModelConfig {
19    /// Current-main polynomial fallback.
20    #[default]
21    Polynomial,
22    /// Deterministic constant latency, primarily useful for parity fixtures.
23    Fixed { prefill_ms: f64, decode_ms: f64 },
24    /// Process-local provider loaded by a Runner or binding.
25    External {
26        provider: String,
27        #[serde(default)]
28        config: Value,
29    },
30}
31
32/// Runtime latency model injected at the engine boundary.
33///
34/// Implementations may call AIC, interpolate profiler data, or use another
35/// provider without adding that dependency to `aisimulate-core`.
36pub trait TimingModel: Send + Sync {
37    /// Predict one prefill batch's latency in milliseconds.
38    fn predict_prefill_ms(
39        &self,
40        batch_size: usize,
41        mean_isl: usize,
42        mean_prefix: usize,
43    ) -> Result<f64>;
44
45    /// Predict one decode batch's latency in milliseconds.
46    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}