diff --git a/Cargo.lock b/Cargo.lock index 7e5e92b85..40c4861a5 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -3927,8 +3927,11 @@ version = "0.1.0" dependencies = [ "anyhow", "clap", + "log", + "logforth", "pegainfer-frontend", "reqwest 0.12.28", + "serde", "serde_json", "tempfile", "tokio", diff --git a/docs/subsystems/frontend/simulated-inference-engine.md b/docs/subsystems/frontend/simulated-inference-engine.md index 9092219f1..3a4c7fdf8 100644 --- a/docs/subsystems/frontend/simulated-inference-engine.md +++ b/docs/subsystems/frontend/simulated-inference-engine.md @@ -1,8 +1,8 @@ # Simulated Inference Engine -> **TL;DR:** `pegainfer-sim` is a CPU-only `Scheduler` that serves through the vLLM/OpenAI frontend with configurable TTFT/TPOT. It launches as `LaunchedEngine::Stepped`. It is a frontend/bench harness, not a real-model performance path. +> **TL;DR:** `pegainfer-sim` is a CPU-only `Scheduler` that serves through the vLLM/OpenAI frontend with profile-priced engine steps or legacy-compatible timing flags. It launches as `LaunchedEngine::Stepped`. It is a frontend/bench harness, not a real-model performance path. > -> **Last touched:** 2026-08 +> **Last touched:** 2026-09 ## Scope @@ -16,9 +16,21 @@ Out of scope: ## Behavior -CLI knobs: model id, port, max model length, base TTFT, prefill throughput, TPOT, fallback token id. - -Timing: TTFT is `base_ttft_ms + prompt_len / prefill_tokens_per_ms`; TPOT is a fixed delay between generated tokens. `SimScheduler::step` emits at most one token per request per step and parks up to 1ms while waiting, so a CPU-only sim does not spin a core the way a GPU scheduler can. +CLI knobs: model identity, optional local metadata path, port, max model length, +legacy base TTFT/prefill throughput/TPOT, fallback token id, profile path and +strict out-of-domain handling. `--profile ` loads a versioned engine profile; +its scheduler limits and model context are authoritative. `--model-path ` +selects the local tokenizer/config directory used by the frontend when the +profile's target model identity is not itself a local path. Legacy timing flags +cannot be combined with an explicit profile. + +Timing: with a profile, `SimScheduler::step` prices one worker step from its +decode/prefill shape and commits progress only after that step duration. Without +an explicit profile, the CLI keeps the legacy per-request scheduler, preserving +`base_ttft_ms + prompt_len / prefill_tokens_per_ms` for prefill and fixed +`tpot_ms` for subsequent decode independently of batch width. The Rust API has +the same legacy behavior when callers do not attach a profile. The standalone +CLI initializes stderr logging so out-of-domain profile fallbacks remain visible. Output token ids cycle through the prompt tokens, or replay a scripted sequence (tool-call tests). Empty prompts use the fallback id. diff --git a/pegainfer-frontend/src/engine/step.rs b/pegainfer-frontend/src/engine/step.rs index d36d80e46..c993ec3e4 100644 --- a/pegainfer-frontend/src/engine/step.rs +++ b/pegainfer-frontend/src/engine/step.rs @@ -173,6 +173,8 @@ pub enum RejectReason { max_tokens: usize, limit: usize, }, + /// Whole-prefill scheduling cannot fit the request in one scheduler step. + PrefillStepBudget { prompt_tokens: usize, limit: usize }, /// Echo needs all-position logits in one forward pass, so the prompt must /// fit the profiled prefill bound. EchoPrefillTokens { prompt_tokens: usize, limit: usize }, @@ -201,6 +203,13 @@ impl fmt::Display for RejectReason { requested {} (prompt={prompt_tokens} + max_tokens={max_tokens})", prompt_tokens.saturating_add(*max_tokens) ), + Self::PrefillStepBudget { + prompt_tokens, + limit, + } => write!( + f, + "request prompt has {prompt_tokens} tokens but the whole-prefill step budget is {limit} tokens" + ), Self::EchoPrefillTokens { prompt_tokens, limit, diff --git a/pegainfer-sim/Cargo.toml b/pegainfer-sim/Cargo.toml index 215eab309..f30dd2fce 100644 --- a/pegainfer-sim/Cargo.toml +++ b/pegainfer-sim/Cargo.toml @@ -11,12 +11,15 @@ path = "src/main.rs" [dependencies] anyhow = { workspace = true } clap = { workspace = true } +log = { workspace = true } +logforth = { workspace = true } pegainfer-frontend = { workspace = true } +serde = { workspace = true } +serde_json = { workspace = true } tokio = { workspace = true, features = ["full"] } [dev-dependencies] reqwest = { workspace = true, features = ["json"] } -serde_json = { workspace = true } tempfile = { workspace = true } tokio-util = { workspace = true } diff --git a/pegainfer-sim/src/lib.rs b/pegainfer-sim/src/lib.rs index d7dae39cf..351a8d9b7 100644 --- a/pegainfer-sim/src/lib.rs +++ b/pegainfer-sim/src/lib.rs @@ -1,12 +1,18 @@ +use std::collections::HashMap; use std::time::Duration; use std::time::Instant; +use anyhow::Context; use anyhow::Result; +use anyhow::bail; use anyhow::ensure; use pegainfer_frontend::engine::Engine; use pegainfer_frontend::engine::EngineInfo; use pegainfer_frontend::engine::FinishReason; +use pegainfer_frontend::engine::PromptEcho; use pegainfer_frontend::engine::QueuedRequest; +use pegainfer_frontend::engine::RejectReason; +use pegainfer_frontend::engine::Request; use pegainfer_frontend::engine::RequestId; use pegainfer_frontend::engine::RequestLedger; use pegainfer_frontend::engine::Scheduler; @@ -14,6 +20,20 @@ use pegainfer_frontend::engine::SchedulerMetrics; use pegainfer_frontend::engine::SpecDecodeCounters; use pegainfer_frontend::engine::spawn_scheduler; +pub mod profile; +pub mod worker; + +use profile::EngineProfile; +use profile::OutOfDomainPolicy; +use worker::CancelResult; +use worker::GeneratedToken; +use worker::RequestRejection; +use worker::StepId; +use worker::StepOutcome; +use worker::SubmissionResult; +use worker::WorkerRequest; +use worker::WorkerState; + mod logprobs; /// Cap on how long `step` parks while waiting for the next due token. New @@ -32,6 +52,13 @@ pub struct SimulatedEngineConfig { scripted_completion: Vec, /// A pretend drafter: `(K, accepted per verify step)`; `None` is no drafter. spec_decode: Option<(usize, usize)>, + profile: Option, +} + +#[derive(Clone, Debug)] +struct ProfileConfig { + profile: EngineProfile, + out_of_domain: OutOfDomainPolicy, } impl SimulatedEngineConfig { @@ -61,6 +88,7 @@ impl SimulatedEngineConfig { fallback_token_id, scripted_completion: Vec::new(), spec_decode: None, + profile: None, }) } @@ -88,11 +116,34 @@ impl SimulatedEngineConfig { self } - fn ttft(&self, prompt_tokens: usize) -> Duration { + /// Set the token id used for empty prompts in either scheduler mode. + #[must_use] + pub fn with_fallback_token_id(mut self, fallback_token_id: u32) -> Self { + self.fallback_token_id = fallback_token_id; + self + } + + /// Use a validated engine profile for online step timing and scheduling. + /// CLI loading belongs to the following commit; this entry point keeps + /// the profiled path injectable for the scheduler and focused tests. + pub fn with_engine_profile( + mut self, + profile: EngineProfile, + out_of_domain: OutOfDomainPolicy, + ) -> Result { + profile.validate()?; + self.profile = Some(ProfileConfig { + profile, + out_of_domain, + }); + Ok(self) + } + + fn ttft(&self, prompt_tokens: usize) -> Result { duration_from_ms(self.base_ttft_ms + prompt_tokens as f64 / self.prefill_tokens_per_ms) } - fn tpot(&self) -> Duration { + fn tpot(&self) -> Result { duration_from_ms(self.tpot_ms) } } @@ -106,6 +157,7 @@ impl Default for SimulatedEngineConfig { fallback_token_id: 0, scripted_completion: Vec::new(), spec_decode: None, + profile: None, } } } @@ -143,6 +195,7 @@ struct SimScheduler { queued: Vec, running: Vec, spec_decode: Option, + profiled: Option, } struct RunningRequest { @@ -154,16 +207,327 @@ struct RunningRequest { logprobs: Option, } +struct ProfiledRuntime { + profile: EngineProfile, + out_of_domain: OutOfDomainPolicy, + worker: WorkerState, + requests: HashMap, + in_flight: Option, +} + +struct ProfiledRequest { + completion_tokens: Vec, + finish_reason: FinishReason, + logprobs: Option, + prompt_echo: Option, +} + +#[derive(Clone, Copy)] +struct ProfiledInFlight { + step_id: StepId, + ready_at: Instant, + decode_reqs: usize, +} + +impl ProfiledRuntime { + fn new(config: &ProfileConfig) -> Self { + let worker = WorkerState::new(config.profile.scheduler.clone()) + .expect("profile was validated before scheduler construction"); + Self { + profile: config.profile.clone(), + out_of_domain: config.out_of_domain, + worker, + requests: HashMap::new(), + in_flight: None, + } + } + + fn step( + &mut self, + config: &SimulatedEngineConfig, + queued: &mut Vec, + ledger: &mut RequestLedger, + spec_decode: &mut Option, + ) -> Result<()> { + self.drain_submissions(config, queued, ledger)?; + self.cancel_aborted(ledger)?; + + if let Some(in_flight) = self.in_flight { + let now = Instant::now(); + if in_flight.ready_at > now { + std::thread::sleep((in_flight.ready_at - now).min(WAIT_SLICE)); + return Ok(()); + } + let outcome = self.worker.complete_step(in_flight.step_id)?; + self.apply_outcome(outcome, in_flight.decode_reqs, config, ledger, spec_decode)?; + self.in_flight = None; + return Ok(()); + } + + let Some(plan) = self.worker.plan_step()? else { + return Ok(()); + }; + let step_id = plan.id(); + let shape = plan.shape(); + let estimate = self.profile.estimate_step(shape, self.out_of_domain)?; + for &request_id in plan.admitted() { + let request = self + .requests + .get_mut(&request_id) + .with_context(|| format!("missing profiled request {request_id}"))?; + ledger.admit(request_id); + if request + .prompt_echo + .as_ref() + .is_some_and(|echo| echo.ids.is_empty()) + { + ledger.echo_prompt( + request_id, + request + .prompt_echo + .take() + .expect("empty prompt echo was checked above"), + ); + } + } + let ready_at = Instant::now() + .checked_add(duration_from_us(estimate.duration_us)) + .context("profiled worker step deadline overflow")?; + self.in_flight = Some(ProfiledInFlight { + step_id, + ready_at, + decode_reqs: plan.decode().len(), + }); + Ok(()) + } + + fn drain_submissions( + &mut self, + config: &SimulatedEngineConfig, + queued: &mut Vec, + ledger: &mut RequestLedger, + ) -> Result<()> { + for QueuedRequest { id, request } in std::mem::take(queued) { + if ledger.is_aborted(id) { + ledger.retire(id); + continue; + } + let prompt_tokens = u32::try_from(request.prompt_tokens.len()) + .context("profiled request prompt length exceeds u32")?; + let output_token_count = planned_completion_len(config, request.max_tokens); + let output_token_count_u64 = u64::try_from(output_token_count) + .context("profiled output length overflows u64")?; + if let Some(rejection) = self.worker.preflight(prompt_tokens, output_token_count_u64) { + ledger.reject(id, reject_reason(rejection, &request)); + continue; + } + let (mut completion_tokens, finish_reason) = + planned_completion(config, &request.prompt_tokens, request.max_tokens); + completion_tokens.reverse(); + let output_tokens = + u32::try_from(completion_tokens.len()).context("profiled output exceeds u32")?; + match self.worker.submit(WorkerRequest { + id, + prompt_tokens, + output_tokens, + })? { + SubmissionResult::Queued => { + let previous = self.requests.insert( + id, + ProfiledRequest { + completion_tokens, + finish_reason, + logprobs: request.logprobs, + prompt_echo: request + .prompt_logprobs + .map(|top_k| logprobs::prompt(&request.prompt_tokens, top_k)), + }, + ); + ensure!( + previous.is_none(), + "profiled request metadata was duplicated" + ); + } + SubmissionResult::Finished => { + ledger.admit(id); + if let Some(top_k) = request.prompt_logprobs { + ledger.echo_prompt(id, logprobs::prompt(&request.prompt_tokens, top_k)); + } + ledger.finish(id, finish_reason); + } + SubmissionResult::Rejected(rejection) => { + ledger.reject(id, reject_reason(rejection, &request)); + } + } + } + Ok(()) + } + + fn cancel_aborted(&mut self, ledger: &mut RequestLedger) -> Result<()> { + let ids: Vec<_> = self.requests.keys().copied().collect(); + for id in ids { + if !ledger.is_aborted(id) { + continue; + } + match self.worker.cancel(id) { + CancelResult::Cancelled => { + ledger.retire(id); + self.requests.remove(&id); + } + CancelResult::Deferred | CancelResult::AlreadyRequested => {} + CancelResult::NotFound => { + bail!("profiled request {id} disappeared before cancellation") + } + } + } + Ok(()) + } + + fn apply_outcome( + &mut self, + outcome: StepOutcome, + decode_reqs: usize, + config: &SimulatedEngineConfig, + ledger: &mut RequestLedger, + spec_decode: &mut Option, + ) -> Result<()> { + let mut cancelled = outcome.cancelled; + let generated_ids: Vec<_> = outcome + .generated + .iter() + .map(|token| token.request_id) + .collect(); + for request_id in generated_ids { + if ledger.is_aborted(request_id) && !cancelled.contains(&request_id) { + cancelled.push(request_id); + } + } + self.cancel_worker_requests(&cancelled)?; + for &request_id in &cancelled { + if ledger.is_active(request_id) { + ledger.retire(request_id); + } + self.requests.remove(&request_id); + } + + for progress in outcome.prefill { + if progress.remaining_tokens != 0 || cancelled.contains(&progress.request_id) { + continue; + } + let request = self + .requests + .get_mut(&progress.request_id) + .with_context(|| format!("missing profiled request {}", progress.request_id))?; + if let Some(prompt_echo) = request.prompt_echo.take() { + ledger.echo_prompt(progress.request_id, prompt_echo); + } + } + for GeneratedToken { + request_id, + token_index, + } in outcome.generated + { + if cancelled.contains(&request_id) { + continue; + } + let request = self + .requests + .get(&request_id) + .with_context(|| format!("missing profiled request {request_id}"))?; + let token_index = token_index + .checked_sub(1) + .context("profiled worker returned an invalid zero token index")?; + let token = *request + .completion_tokens + .get(usize::try_from(token_index).context("token index overflow")?) + .with_context(|| { + format!("missing token {token_index} for profiled request {request_id}") + })?; + let logprob = request + .logprobs + .map(|top_k| logprobs::completion(token, top_k)); + let logprobs = match logprob { + Some(logprob) => vec![Some(logprob)], + None => Vec::new(), + }; + ledger.push_tokens(request_id, &[token], &logprobs); + } + for request_id in outcome.finished { + if cancelled.contains(&request_id) { + continue; + } + let request = self + .requests + .remove(&request_id) + .with_context(|| format!("missing profiled request {request_id}"))?; + ledger.finish(request_id, request.finish_reason); + } + if let (Some(counters), Some((k, accepted))) = (spec_decode.as_mut(), config.spec_decode) { + for _ in 0..decode_reqs { + counters.observe_draft(k, accepted); + } + } + Ok(()) + } + + fn cancel_worker_requests(&mut self, request_ids: &[RequestId]) -> Result<()> { + for &request_id in request_ids { + match self.worker.cancel(request_id) { + // `Cancelled` is the late-abort case: complete_step has + // already cleared the worker's in-flight plan, so the request + // must be removed from running here before its ledger account + // and metadata are retired below. + CancelResult::Cancelled | CancelResult::NotFound => {} + CancelResult::Deferred | CancelResult::AlreadyRequested => { + bail!("profiled request {request_id} remained in flight during cancellation") + } + } + } + Ok(()) + } + + fn metrics(&self, spec_decode: Option<&SpecDecodeCounters>) -> SchedulerMetrics { + SchedulerMetrics { + num_running_reqs: self.worker.running_len() as u64, + num_waiting_reqs: self.worker.waiting_len() as u64, + spec_decode: spec_decode.copied(), + ..SchedulerMetrics::default() + } + } +} + +fn reject_reason(rejection: RequestRejection, request: &Request) -> RejectReason { + match rejection { + RequestRejection::ModelLengthExceeded { + total_tokens: _, + max_model_len, + } => RejectReason::ContextLength { + prompt_tokens: request.prompt_tokens.len(), + max_tokens: request.max_tokens, + limit: max_model_len as usize, + }, + RequestRejection::WholePrefillExceedsStepBudget { + prompt_tokens, + max_num_batched_tokens, + } => RejectReason::PrefillStepBudget { + prompt_tokens: prompt_tokens as usize, + limit: max_num_batched_tokens as usize, + }, + } +} + impl SimScheduler { fn new(config: SimulatedEngineConfig) -> Self { let spec_decode = config .spec_decode .map(|(k, _)| SpecDecodeCounters::new(k).expect("K checked at config time")); + let profiled = config.profile.as_ref().map(ProfiledRuntime::new); Self { config, queued: Vec::new(), running: Vec::new(), spec_decode, + profiled, } } @@ -185,6 +549,38 @@ impl Scheduler for SimScheduler { } fn step(&mut self, ledger: &mut RequestLedger) -> Result<()> { + if self.profiled.is_some() { + let mut profiled = self + .profiled + .take() + .expect("profiled runtime presence was checked above"); + let result = profiled.step( + &self.config, + &mut self.queued, + ledger, + &mut self.spec_decode, + ); + self.profiled = Some(profiled); + return result; + } + self.step_legacy(ledger) + } + + fn metrics(&self) -> SchedulerMetrics { + if let Some(profiled) = &self.profiled { + return profiled.metrics(self.spec_decode.as_ref()); + } + SchedulerMetrics { + num_running_reqs: self.running.len() as u64, + num_waiting_reqs: self.queued.len() as u64, + spec_decode: self.spec_decode, + ..SchedulerMetrics::default() + } + } +} + +impl SimScheduler { + fn step_legacy(&mut self, ledger: &mut RequestLedger) -> Result<()> { for QueuedRequest { id, request } in self.queued.drain(..) { if ledger.is_aborted(id) { ledger.retire(id); @@ -204,7 +600,7 @@ impl Scheduler for SimScheduler { self.running.push(RunningRequest { id, pending, - next_token_at: Instant::now() + self.config.ttft(prompt_len), + next_token_at: Instant::now() + self.config.ttft(prompt_len)?, finish_reason, logprobs: request.logprobs, }); @@ -242,7 +638,7 @@ impl Scheduler for SimScheduler { if running.pending.is_empty() { ledger.finish(running.id, running.finish_reason); } else { - running.next_token_at = Instant::now() + self.config.tpot(); + running.next_token_at = Instant::now() + self.config.tpot()?; still_running.push(running); } } @@ -250,15 +646,6 @@ impl Scheduler for SimScheduler { self.park_if_waiting(); Ok(()) } - - fn metrics(&self) -> SchedulerMetrics { - SchedulerMetrics { - num_running_reqs: self.running.len() as u64, - num_waiting_reqs: self.queued.len() as u64, - spec_decode: self.spec_decode, - ..SchedulerMetrics::default() - } - } } /// Remaining tokens (reversed) plus the terminal reason. Empty pending means @@ -290,6 +677,14 @@ fn planned_completion( (pending, finish_reason) } +fn planned_completion_len(config: &SimulatedEngineConfig, max_tokens: usize) -> usize { + if config.scripted_completion.is_empty() { + max_tokens + } else { + max_tokens.min(config.scripted_completion.len()) + } +} + fn fake_token_id(prompt_tokens: &[u32], index: usize, fallback_token_id: u32) -> u32 { if prompt_tokens.is_empty() { return fallback_token_id; @@ -297,8 +692,20 @@ fn fake_token_id(prompt_tokens: &[u32], index: usize, fallback_token_id: u32) -> prompt_tokens[index % prompt_tokens.len()] } -fn duration_from_ms(ms: f64) -> Duration { - Duration::from_secs_f64(ms / 1000.0) +fn duration_from_ms(ms: f64) -> Result { + ensure!( + ms.is_finite() && ms >= 0.0, + "timing value must be finite and non-negative" + ); + Duration::try_from_secs_f64(ms / 1000.0) + .context("timing value is not representable as a Duration") +} + +fn duration_from_us(microseconds: u64) -> Duration { + Duration::new( + microseconds / 1_000_000, + ((microseconds % 1_000_000) * 1_000) as u32, + ) } #[cfg(test)] @@ -308,6 +715,14 @@ mod tests { use pegainfer_frontend::sampler::SamplingParams; use super::*; + use crate::profile::ENGINE_PROFILE_SCHEMA_VERSION; + use crate::profile::ParametricFallback; + use crate::profile::PrefillPolicy; + use crate::profile::ProfileProvenance; + use crate::profile::SchedulerPolicy; + use crate::profile::SchedulerProfile; + use crate::profile::StepTimingProfile; + use crate::profile::TimingGrid; fn request(prompt_tokens: Vec, max_tokens: usize, logprobs: usize) -> Request { Request { @@ -353,6 +768,43 @@ mod tests { (tokens, prompt_tokens, terminal.expect("terminal")) } + fn zero_cost_profile() -> EngineProfile { + EngineProfile { + schema_version: ENGINE_PROFILE_SCHEMA_VERSION, + profile_id: "online-test".to_string(), + provenance: ProfileProvenance { + target_engine: "test-engine".to_string(), + engine_version: "0.0.0".to_string(), + model_id: "test-model".to_string(), + model_revision: "test-revision".to_string(), + model_config_sha256: "00".repeat(32), + gpu: "test-gpu".to_string(), + server_flags: Vec::new(), + }, + scheduler: SchedulerProfile { + policy: SchedulerPolicy::VllmV1, + max_num_seqs: 2, + max_num_batched_tokens: 4, + max_model_len: 32, + prefill: PrefillPolicy::Whole, + }, + timing: StepTimingProfile { + grid: TimingGrid { + decode_reqs: vec![0, 2], + sum_decode_ctx_tokens: vec![0, 8], + prefill_tokens_in_step: vec![0, 4], + step_duration_us: vec![0; 8], + }, + fallback: ParametricFallback { + t0_us: 0.0, + prefill_token_us: 0.0, + decode_request_us: 0.0, + decode_context_token_us: 0.0, + }, + }, + } + } + #[test] fn fake_token_id_cycles_prompt_tokens() { assert_eq!(fake_token_id(&[7, 9], 0, 42), 7); @@ -395,6 +847,68 @@ mod tests { )); } + #[test] + fn profiled_scheduler_replays_scripted_completion() { + let config = SimulatedEngineConfig::default() + .with_engine_profile(zero_cost_profile(), OutOfDomainPolicy::Strict) + .unwrap() + .with_scripted_completion(vec![11, 22, 33]); + let (tokens, prompt_tokens, terminal) = + collect_completion(&config, request(vec![7, 9], 3, 0)); + assert_eq!(prompt_tokens, Some(2)); + assert_eq!(tokens, [11, 22, 33]); + assert!(matches!( + terminal, + Terminal::Finished { + reason: FinishReason::Stop, + prompt_tokens: 2, + completion_tokens: 3, + } + )); + } + + #[test] + fn late_abort_cleanup_removes_nonterminal_worker_request() { + let profile = zero_cost_profile(); + let mut runtime = ProfiledRuntime::new(&ProfileConfig { + profile, + out_of_domain: OutOfDomainPolicy::Strict, + }); + let request_id = RequestId::new(7); + let survivor_id = RequestId::new(8); + runtime + .worker + .submit(WorkerRequest { + id: request_id, + prompt_tokens: 0, + output_tokens: 2, + }) + .unwrap(); + runtime + .worker + .submit(WorkerRequest { + id: survivor_id, + prompt_tokens: 0, + output_tokens: 2, + }) + .unwrap(); + + let step_id = runtime.worker.plan_step().unwrap().unwrap().id(); + let outcome = runtime.worker.complete_step(step_id).unwrap(); + assert_eq!(outcome.generated.len(), 2); + assert!(outcome.finished.is_empty()); + assert!(runtime.worker.request(request_id).is_some()); + assert!(runtime.worker.request(survivor_id).is_some()); + + runtime.cancel_worker_requests(&[request_id]).unwrap(); + assert!(runtime.worker.request(request_id).is_none()); + + let survivor_step_id = runtime.worker.plan_step().unwrap().unwrap().id(); + let survivor_outcome = runtime.worker.complete_step(survivor_step_id).unwrap(); + assert_eq!(survivor_outcome.finished, [survivor_id]); + assert!(runtime.worker.request(survivor_id).is_none()); + } + #[test] fn config_rejects_invalid_timing_values() { assert!(SimulatedEngineConfig::new(-1.0, 100.0, 12.0, 0).is_err()); diff --git a/pegainfer-sim/src/main.rs b/pegainfer-sim/src/main.rs index ff7d33120..b1bec8aac 100644 --- a/pegainfer-sim/src/main.rs +++ b/pegainfer-sim/src/main.rs @@ -1,11 +1,24 @@ -use std::path::Path; +use std::path::PathBuf; +use std::sync::Once; +use anyhow::Context; use anyhow::Result; +use anyhow::bail; +use anyhow::ensure; use clap::Parser; use pegainfer_sim::SimulatedEngineConfig; +use pegainfer_sim::profile::EngineProfile; +use pegainfer_sim::profile::OutOfDomainPolicy; use pegainfer_sim::start_engine; const DEFAULT_MODEL_ID: &str = "Qwen/Qwen3-0.6B"; +const DEFAULT_MAX_MODEL_LEN: u32 = 8192; +const DEFAULT_BASE_TTFT_MS: f64 = 5.0; +const DEFAULT_PREFILL_TOKENS_PER_MS: f64 = 100.0; +const DEFAULT_TPOT_MS: f64 = 12.0; +const DEFAULT_FALLBACK_TOKEN_ID: u32 = 0; + +static LOGGING_INIT: Once = Once::new(); #[derive(Parser, Debug)] #[command( @@ -13,53 +26,253 @@ const DEFAULT_MODEL_ID: &str = "Qwen/Qwen3-0.6B"; about = "CPU-only simulated inference server for OpenAI/vLLM serving benchmarks" )] struct Args { - /// Tokenizer/model metadata id used by the vLLM frontend. No weights are loaded. - #[arg(long, default_value = DEFAULT_MODEL_ID)] - model_id: String, + /// Model identity. In legacy mode it also remains the metadata path; with + /// --profile it must match the profile's target model id. + #[arg(long)] + model_id: Option, + + /// Local tokenizer/model metadata directory used by the vLLM frontend. + /// With --profile, this can differ from the profile's target model id. + #[arg(long, value_name = "PATH")] + model_path: Option, /// Port to listen on. #[arg(long, default_value_t = 8000)] port: u16, /// Max context length reported to the vLLM frontend. - #[arg(long, default_value_t = 8192)] - max_model_len: u32, + #[arg(long)] + max_model_len: Option, /// Fixed TTFT floor before the first fake token. - #[arg(long, default_value_t = 5.0)] - base_ttft_ms: f64, + #[arg(long)] + base_ttft_ms: Option, /// Simulated prefill throughput used as prompt_len / throughput. - #[arg(long, default_value_t = 100.0)] - prefill_tokens_per_ms: f64, + #[arg(long)] + prefill_tokens_per_ms: Option, /// Fixed delay between generated fake tokens. - #[arg(long, default_value_t = 12.0)] - tpot_ms: f64, + #[arg(long)] + tpot_ms: Option, /// Token id used when a request has an empty prompt-token list. - #[arg(long, default_value_t = 0)] + #[arg(long, default_value_t = DEFAULT_FALLBACK_TOKEN_ID)] fallback_token_id: u32, + + /// Versioned timing and scheduler profile generated from a target engine. + #[arg(long, value_name = "FILE")] + profile: Option, + + /// Reject step shapes outside the profile timing grid instead of using + /// the profile's parametric fallback. + #[arg(long)] + strict: bool, +} + +#[derive(Debug)] +struct RuntimeConfig { + engine: SimulatedEngineConfig, + model_path: PathBuf, + served_model_name: Vec, + max_model_len: u32, + profile: Option, + out_of_domain: Option, +} + +fn build_runtime(args: &Args) -> Result { + if let Some(path) = &args.profile { + ensure_legacy_timing_flags_are_absent(args)?; + let bytes = std::fs::read(path) + .with_context(|| format!("failed to read engine profile {}", path.display()))?; + let profile = EngineProfile::from_json_slice(&bytes) + .with_context(|| format!("failed to load engine profile {}", path.display()))?; + let model_id = profile.provenance.model_id.clone(); + if let Some(requested) = &args.model_id { + ensure!( + requested == &profile.provenance.model_id, + "--model-id '{}' conflicts with profile model_id '{}'", + requested, + profile.provenance.model_id + ); + } + if let Some(requested) = args.max_model_len { + ensure!( + requested == profile.scheduler.max_model_len, + "--max-model-len {} conflicts with profile max_model_len {}", + requested, + profile.scheduler.max_model_len + ); + } + let out_of_domain = if args.strict { + OutOfDomainPolicy::Strict + } else { + OutOfDomainPolicy::WarnAndFallback + }; + let engine = SimulatedEngineConfig::default() + .with_fallback_token_id(args.fallback_token_id) + .with_engine_profile(profile.clone(), out_of_domain)?; + let model_path = args.model_path.clone().unwrap_or_else(|| { + args.model_id + .as_deref() + .map_or_else(|| PathBuf::from(&model_id), PathBuf::from) + }); + return Ok(RuntimeConfig { + engine, + model_path, + served_model_name: vec![model_id], + max_model_len: profile.scheduler.max_model_len, + profile: Some(profile), + out_of_domain: Some(out_of_domain), + }); + } + + ensure!( + !args.strict, + "--strict requires --profile; legacy timing has no profile domain to validate" + ); + let model_id = args + .model_id + .clone() + .or_else(|| { + args.model_path + .as_deref() + .map(|path| path.to_string_lossy().into_owned()) + }) + .unwrap_or_else(|| DEFAULT_MODEL_ID.to_string()); + let model_path = args + .model_path + .clone() + .unwrap_or_else(|| PathBuf::from(&model_id)); + let max_model_len = args.max_model_len.unwrap_or(DEFAULT_MAX_MODEL_LEN); + ensure!(max_model_len > 0, "max_model_len must be positive"); + let base_ttft_ms = args.base_ttft_ms.unwrap_or(DEFAULT_BASE_TTFT_MS); + let prefill_tokens_per_ms = args + .prefill_tokens_per_ms + .unwrap_or(DEFAULT_PREFILL_TOKENS_PER_MS); + let tpot_ms = args.tpot_ms.unwrap_or(DEFAULT_TPOT_MS); + let engine = SimulatedEngineConfig::new( + base_ttft_ms, + prefill_tokens_per_ms, + tpot_ms, + args.fallback_token_id, + )?; + Ok(RuntimeConfig { + engine, + model_path, + served_model_name: if args.model_path.is_some() && args.model_id.is_some() { + vec![model_id] + } else { + Vec::new() + }, + max_model_len, + profile: None, + out_of_domain: None, + }) +} + +fn ensure_legacy_timing_flags_are_absent(args: &Args) -> Result<()> { + let provided = [ + ("--base-ttft-ms", args.base_ttft_ms.is_some()), + ( + "--prefill-tokens-per-ms", + args.prefill_tokens_per_ms.is_some(), + ), + ("--tpot-ms", args.tpot_ms.is_some()), + ]; + if let Some((name, true)) = provided.into_iter().find(|(_, present)| *present) { + bail!("{name} cannot be combined with --profile; timing comes from the profile"); + } + Ok(()) +} + +fn report_profile(runtime: &RuntimeConfig) { + let Some(profile) = &runtime.profile else { + eprintln!("active engine profile: legacy fixed TTFT/TPOT scheduler"); + return; + }; + let scheduler = &profile.scheduler; + eprintln!( + "active engine profile: id={} target={} version={} model={} revision={} gpu={} scheduler={:?} max_num_seqs={} max_num_batched_tokens={} max_model_len={} timing_domain={:?} out_of_domain={:?}", + profile.profile_id, + profile.provenance.target_engine, + profile.provenance.engine_version, + profile.provenance.model_id, + profile.provenance.model_revision, + profile.provenance.gpu, + scheduler.policy, + scheduler.max_num_seqs, + scheduler.max_num_batched_tokens, + scheduler.max_model_len, + profile.timing.grid.domain(), + runtime.out_of_domain, + ); +} + +fn init_logging() { + LOGGING_INIT.call_once(|| { + let filter_spec = std::env::var("RUST_LOG").unwrap_or_else(|_| "info".to_string()); + let filter = logforth::filter::env_filter::EnvFilterBuilder::from_spec(filter_spec).build(); + logforth::starter_log::builder() + .dispatch(|dispatch| { + dispatch + .filter(filter) + .append(logforth::append::Stderr::default()) + }) + .apply(); + }); } #[tokio::main] async fn main() -> Result<()> { + init_logging(); let args = Args::parse(); - let config = SimulatedEngineConfig::new( - args.base_ttft_ms, - args.prefill_tokens_per_ms, - args.tpot_ms, - args.fallback_token_id, - )?; - let engine = start_engine(&config); + let runtime = build_runtime(&args)?; + report_profile(&runtime); + let engine = start_engine(&runtime.engine); pegainfer_frontend::vllm::serve( std::future::ready(Ok(engine.into())), - Path::new(&args.model_id), - Vec::new(), + &runtime.model_path, + runtime.served_model_name, args.port, - Some(args.max_model_len), + Some(runtime.max_model_len), pegainfer_frontend::vllm::shutdown_token_from_ctrl_c(), ) .await } + +#[cfg(test)] +mod tests { + use super::*; + + fn legacy_args() -> Args { + Args { + model_id: Some("test-model".to_string()), + model_path: None, + port: 8000, + max_model_len: None, + base_ttft_ms: None, + prefill_tokens_per_ms: None, + tpot_ms: None, + fallback_token_id: DEFAULT_FALLBACK_TOKEN_ID, + profile: None, + strict: false, + } + } + + #[test] + fn legacy_cli_keeps_fixed_timing_scheduler() { + let runtime = build_runtime(&legacy_args()).expect("legacy runtime should build"); + + assert!(runtime.profile.is_none()); + assert!(runtime.out_of_domain.is_none()); + } + + #[test] + fn sim_logger_accepts_warn_records() { + init_logging(); + + assert!(log::log_enabled!(target: "pegainfer_sim::profile", log::Level::Warn)); + } +} diff --git a/pegainfer-sim/src/profile.rs b/pegainfer-sim/src/profile.rs new file mode 100644 index 000000000..73b48182b --- /dev/null +++ b/pegainfer-sim/src/profile.rs @@ -0,0 +1,729 @@ +use std::fmt::Display; + +use anyhow::Context; +use anyhow::Result; +use anyhow::bail; +use anyhow::ensure; +use serde::Deserialize; +use serde::Serialize; + +pub const ENGINE_PROFILE_SCHEMA_VERSION: u32 = 1; + +#[derive(Clone, Debug, Deserialize, PartialEq, Serialize)] +#[serde(deny_unknown_fields)] +pub struct EngineProfile { + pub schema_version: u32, + pub profile_id: String, + pub provenance: ProfileProvenance, + pub scheduler: SchedulerProfile, + pub timing: StepTimingProfile, +} + +impl EngineProfile { + pub fn from_json_slice(bytes: &[u8]) -> Result { + let profile: Self = + serde_json::from_slice(bytes).context("failed to parse engine profile JSON")?; + profile.validate()?; + Ok(profile) + } + + pub fn validate(&self) -> Result<()> { + ensure!( + self.schema_version == ENGINE_PROFILE_SCHEMA_VERSION, + "unsupported engine profile schema_version {}; expected {}", + self.schema_version, + ENGINE_PROFILE_SCHEMA_VERSION + ); + ensure_nonempty("profile_id", &self.profile_id)?; + self.provenance.validate()?; + self.scheduler.validate()?; + self.timing.validate(&self.scheduler) + } + + pub fn estimate_step( + &self, + shape: StepShape, + out_of_domain: OutOfDomainPolicy, + ) -> Result { + self.scheduler.validate_shape(shape)?; + if let Some(duration_us) = self.timing.grid.interpolate(shape)? { + return Ok(StepTimingEstimate { + duration_us, + source: StepTimingSource::GridInterpolation, + }); + } + + let domain = self.timing.grid.domain(); + if out_of_domain == OutOfDomainPolicy::Strict { + bail!( + "timing profile '{}' does not cover step shape {shape:?}; supported grid domain: {domain:?}", + self.profile_id + ); + } + log::warn!( + "timing profile '{}' does not cover step shape {shape:?}; using parametric fallback outside grid domain {domain:?}", + self.profile_id + ); + Ok(StepTimingEstimate { + duration_us: self.timing.fallback.evaluate(shape)?, + source: StepTimingSource::ParametricFallback, + }) + } +} + +#[derive(Clone, Debug, Deserialize, PartialEq, Serialize)] +#[serde(deny_unknown_fields)] +pub struct ProfileProvenance { + pub target_engine: String, + pub engine_version: String, + pub model_id: String, + pub model_revision: String, + pub model_config_sha256: String, + pub gpu: String, + pub server_flags: Vec, +} + +impl ProfileProvenance { + fn validate(&self) -> Result<()> { + ensure_nonempty("provenance.target_engine", &self.target_engine)?; + ensure_nonempty("provenance.engine_version", &self.engine_version)?; + ensure_nonempty("provenance.model_id", &self.model_id)?; + ensure_nonempty("provenance.model_revision", &self.model_revision)?; + ensure_nonempty("provenance.gpu", &self.gpu)?; + ensure!( + is_sha256(&self.model_config_sha256), + "provenance.model_config_sha256 must be 64 hexadecimal characters" + ); + for flag in &self.server_flags { + ensure_nonempty("provenance.server_flags entry", flag)?; + } + Ok(()) + } +} + +#[derive(Clone, Copy, Debug, Deserialize, Eq, PartialEq, Serialize)] +#[serde(rename_all = "snake_case")] +pub enum SchedulerPolicy { + VllmV1, +} + +#[derive(Clone, Debug, Deserialize, Eq, PartialEq, Serialize)] +#[serde(deny_unknown_fields)] +pub struct SchedulerProfile { + pub policy: SchedulerPolicy, + pub max_num_seqs: u32, + pub max_num_batched_tokens: u32, + pub max_model_len: u32, + pub prefill: PrefillPolicy, +} + +impl SchedulerProfile { + pub(crate) fn validate(&self) -> Result<()> { + ensure!( + self.max_num_seqs > 0, + "scheduler.max_num_seqs must be positive" + ); + ensure!( + self.max_num_batched_tokens > 0, + "scheduler.max_num_batched_tokens must be positive" + ); + ensure!( + self.max_model_len > 0, + "scheduler.max_model_len must be positive" + ); + ensure!( + self.max_num_seqs <= self.max_num_batched_tokens, + "scheduler.max_num_seqs cannot exceed max_num_batched_tokens because every decoding request consumes one token per step" + ); + if let PrefillPolicy::Chunked { max_chunk_tokens } = self.prefill { + ensure!( + max_chunk_tokens > 0, + "scheduler.prefill.max_chunk_tokens must be positive" + ); + ensure!( + max_chunk_tokens <= self.max_num_batched_tokens, + "scheduler.prefill.max_chunk_tokens cannot exceed max_num_batched_tokens" + ); + } + Ok(()) + } + + fn validate_shape(&self, shape: StepShape) -> Result<()> { + ensure!( + shape.decode_reqs > 0 || shape.prefill_tokens_in_step > 0, + "step shape must contain prefill or decode work" + ); + ensure!( + shape.decode_reqs <= self.max_num_seqs, + "step shape decode_reqs {} exceed scheduler max_num_seqs {}", + shape.decode_reqs, + self.max_num_seqs + ); + let step_tokens = shape + .decode_reqs + .checked_add(shape.prefill_tokens_in_step) + .context("step token count overflow")?; + ensure!( + step_tokens <= self.max_num_batched_tokens, + "step shape token count {step_tokens} exceeds scheduler max_num_batched_tokens {}", + self.max_num_batched_tokens + ); + if shape.decode_reqs == 0 { + ensure!( + shape.sum_decode_ctx_tokens == 0, + "sum_decode_ctx_tokens must be zero when decode_reqs is zero" + ); + } + let max_decode_ctx = u64::from(shape.decode_reqs) + .checked_mul(u64::from(self.max_model_len)) + .context("maximum decode context overflow")?; + ensure!( + shape.sum_decode_ctx_tokens <= max_decode_ctx, + "step shape sum_decode_ctx_tokens {} exceed the per-request context ceiling {}", + shape.sum_decode_ctx_tokens, + max_decode_ctx + ); + Ok(()) + } +} + +#[derive(Clone, Copy, Debug, Deserialize, Eq, PartialEq, Serialize)] +#[serde(tag = "mode", rename_all = "snake_case", deny_unknown_fields)] +pub enum PrefillPolicy { + Whole, + Chunked { max_chunk_tokens: u32 }, +} + +#[derive(Clone, Copy, Debug, Eq, PartialEq)] +pub struct StepShape { + pub decode_reqs: u32, + pub sum_decode_ctx_tokens: u64, + pub prefill_tokens_in_step: u32, +} + +#[derive(Clone, Debug, Deserialize, PartialEq, Serialize)] +#[serde(deny_unknown_fields)] +pub struct StepTimingProfile { + pub grid: TimingGrid, + pub fallback: ParametricFallback, +} + +impl StepTimingProfile { + fn validate(&self, scheduler: &SchedulerProfile) -> Result<()> { + self.grid.validate(scheduler)?; + self.fallback.validate() + } +} + +#[derive(Clone, Debug, Deserialize, Eq, PartialEq, Serialize)] +#[serde(deny_unknown_fields)] +pub struct TimingGrid { + pub decode_reqs: Vec, + pub sum_decode_ctx_tokens: Vec, + pub prefill_tokens_in_step: Vec, + /// Row-major order: decode request, decode context, then prefill token axis. + pub step_duration_us: Vec, +} + +impl TimingGrid { + fn validate(&self, scheduler: &SchedulerProfile) -> Result<()> { + validate_axis("timing.grid.decode_reqs", &self.decode_reqs)?; + validate_axis( + "timing.grid.sum_decode_ctx_tokens", + &self.sum_decode_ctx_tokens, + )?; + validate_axis( + "timing.grid.prefill_tokens_in_step", + &self.prefill_tokens_in_step, + )?; + let expected_values = self + .decode_reqs + .len() + .checked_mul(self.sum_decode_ctx_tokens.len()) + .and_then(|value| value.checked_mul(self.prefill_tokens_in_step.len())) + .context("timing grid dimensions overflow")?; + ensure!( + self.step_duration_us.len() == expected_values, + "timing.grid.step_duration_us contains {} values; expected {expected_values}", + self.step_duration_us.len() + ); + ensure!( + self.decode_reqs + .last() + .copied() + .expect("validated non-empty axis") + <= scheduler.max_num_seqs, + "timing.grid.decode_reqs exceed scheduler.max_num_seqs" + ); + ensure!( + self.prefill_tokens_in_step + .last() + .copied() + .expect("validated non-empty axis") + <= scheduler.max_num_batched_tokens, + "timing.grid.prefill_tokens_in_step exceed scheduler.max_num_batched_tokens" + ); + let max_context = u64::from(scheduler.max_num_seqs) + .checked_mul(u64::from(scheduler.max_model_len)) + .context("scheduler context domain overflow")?; + ensure!( + self.sum_decode_ctx_tokens + .last() + .copied() + .expect("validated non-empty axis") + <= max_context, + "timing.grid.sum_decode_ctx_tokens exceed the scheduler context domain" + ); + Ok(()) + } + + pub fn domain(&self) -> Option { + Some(TimingDomain { + min_decode_reqs: *self.decode_reqs.first()?, + max_decode_reqs: *self.decode_reqs.last()?, + min_sum_decode_ctx_tokens: *self.sum_decode_ctx_tokens.first()?, + max_sum_decode_ctx_tokens: *self.sum_decode_ctx_tokens.last()?, + min_prefill_tokens_in_step: *self.prefill_tokens_in_step.first()?, + max_prefill_tokens_in_step: *self.prefill_tokens_in_step.last()?, + }) + } + + fn interpolate(&self, shape: StepShape) -> Result> { + let Some(decode) = bracket(&self.decode_reqs, shape.decode_reqs) else { + return Ok(None); + }; + let Some(context) = bracket(&self.sum_decode_ctx_tokens, shape.sum_decode_ctx_tokens) + else { + return Ok(None); + }; + let Some(prefill) = bracket(&self.prefill_tokens_in_step, shape.prefill_tokens_in_step) + else { + return Ok(None); + }; + + let d0c0 = interpolate_bracket( + self.value_at(decode.lower, context.lower, prefill.lower)?, + self.value_at(decode.lower, context.lower, prefill.upper)?, + prefill, + )?; + let d0c1 = interpolate_bracket( + self.value_at(decode.lower, context.upper, prefill.lower)?, + self.value_at(decode.lower, context.upper, prefill.upper)?, + prefill, + )?; + let d1c0 = interpolate_bracket( + self.value_at(decode.upper, context.lower, prefill.lower)?, + self.value_at(decode.upper, context.lower, prefill.upper)?, + prefill, + )?; + let d1c1 = interpolate_bracket( + self.value_at(decode.upper, context.upper, prefill.lower)?, + self.value_at(decode.upper, context.upper, prefill.upper)?, + prefill, + )?; + let d0 = interpolate_bracket(d0c0, d0c1, context)?; + let d1 = interpolate_bracket(d1c0, d1c1, context)?; + Ok(Some(interpolate_bracket(d0, d1, decode)?)) + } + + fn value_at(&self, decode: usize, context: usize, prefill: usize) -> Result { + let index = decode + .checked_mul(self.sum_decode_ctx_tokens.len()) + .and_then(|value| value.checked_add(context)) + .and_then(|value| value.checked_mul(self.prefill_tokens_in_step.len())) + .and_then(|value| value.checked_add(prefill)) + .context("timing grid index overflow")?; + self.step_duration_us + .get(index) + .copied() + .context("timing grid is missing a duration value") + } +} + +#[derive(Clone, Copy, Debug, Eq, PartialEq)] +pub struct TimingDomain { + pub min_decode_reqs: u32, + pub max_decode_reqs: u32, + pub min_sum_decode_ctx_tokens: u64, + pub max_sum_decode_ctx_tokens: u64, + pub min_prefill_tokens_in_step: u32, + pub max_prefill_tokens_in_step: u32, +} + +#[derive(Clone, Copy, Debug, Deserialize, PartialEq, Serialize)] +#[serde(deny_unknown_fields)] +pub struct ParametricFallback { + pub t0_us: f64, + pub prefill_token_us: f64, + pub decode_request_us: f64, + pub decode_context_token_us: f64, +} + +impl ParametricFallback { + fn validate(&self) -> Result<()> { + validate_coefficient("timing.fallback.t0_us", self.t0_us)?; + validate_coefficient("timing.fallback.prefill_token_us", self.prefill_token_us)?; + validate_coefficient("timing.fallback.decode_request_us", self.decode_request_us)?; + validate_coefficient( + "timing.fallback.decode_context_token_us", + self.decode_context_token_us, + ) + } + + fn evaluate(&self, shape: StepShape) -> Result { + let duration_us = self.t0_us + + self.prefill_token_us * f64::from(shape.prefill_tokens_in_step) + + self.decode_request_us * f64::from(shape.decode_reqs) + + self.decode_context_token_us * shape.sum_decode_ctx_tokens as f64; + let rounded = duration_us.round(); + ensure!( + rounded.is_finite() && rounded >= 0.0 && rounded < u64::MAX as f64, + "parametric timing fallback overflow for step shape {shape:?}" + ); + Ok(rounded as u64) + } +} + +#[derive(Clone, Copy, Debug, Eq, PartialEq)] +pub enum OutOfDomainPolicy { + WarnAndFallback, + Strict, +} + +#[derive(Clone, Copy, Debug, Eq, PartialEq)] +pub enum StepTimingSource { + GridInterpolation, + ParametricFallback, +} + +#[derive(Clone, Copy, Debug, Eq, PartialEq)] +pub struct StepTimingEstimate { + pub duration_us: u64, + pub source: StepTimingSource, +} + +#[derive(Clone, Copy)] +struct AxisBracket { + lower: usize, + upper: usize, + numerator: u64, + denominator: u64, +} + +fn bracket(axis: &[T], value: T) -> Option +where + T: Copy + Ord + Into, +{ + match axis.binary_search(&value) { + Ok(index) => Some(AxisBracket { + lower: index, + upper: index, + numerator: 0, + denominator: 1, + }), + Err(0) => None, + Err(index) if index == axis.len() => None, + Err(index) => { + let lower_value = axis[index - 1].into(); + let upper_value = axis[index].into(); + Some(AxisBracket { + lower: index - 1, + upper: index, + numerator: value.into() - lower_value, + denominator: upper_value - lower_value, + }) + } + } +} + +fn interpolate_bracket(lower: u64, upper: u64, bracket: AxisBracket) -> Result { + if bracket.lower == bracket.upper || lower == upper { + return Ok(lower); + } + let delta = lower.abs_diff(upper); + let scaled = u128::from(delta) + .checked_mul(u128::from(bracket.numerator)) + .context("timing interpolation multiplication overflow")?; + let rounded = scaled + .checked_add(u128::from(bracket.denominator / 2)) + .context("timing interpolation rounding overflow")? + / u128::from(bracket.denominator); + let adjustment = u64::try_from(rounded).context("timing interpolation result overflow")?; + if upper >= lower { + lower + .checked_add(adjustment) + .context("timing interpolation addition overflow") + } else { + lower + .checked_sub(adjustment) + .context("timing interpolation subtraction overflow") + } +} + +fn ensure_nonempty(field: &str, value: &str) -> Result<()> { + ensure!(!value.trim().is_empty(), "{field} must not be empty"); + Ok(()) +} + +fn is_sha256(value: &str) -> bool { + value.len() == 64 && value.bytes().all(|byte| byte.is_ascii_hexdigit()) +} + +fn validate_axis(name: &str, axis: &[T]) -> Result<()> +where + T: Copy + Ord + Display, +{ + ensure!(!axis.is_empty(), "{name} must not be empty"); + for pair in axis.windows(2) { + ensure!( + pair[0] < pair[1], + "{name} must be strictly increasing; found {} then {}", + pair[0], + pair[1] + ); + } + Ok(()) +} + +fn validate_coefficient(name: &str, value: f64) -> Result<()> { + ensure!( + value.is_finite() && value >= 0.0, + "{name} must be finite and non-negative" + ); + Ok(()) +} + +#[cfg(test)] +mod tests { + use serde_json::Value; + use serde_json::json; + + use super::*; + + fn profile_json() -> Value { + json!({ + "schema_version": 1, + "profile_id": "vllm-test", + "provenance": { + "target_engine": "vllm", + "engine_version": "0.27.1", + "model_id": "Qwen/Qwen3-4B", + "model_revision": "test-revision", + "model_config_sha256": "00".repeat(32), + "gpu": "NVIDIA RTX 5090", + "server_flags": [ + "--max-num-batched-tokens=16", + "--max-num-seqs=4" + ] + }, + "scheduler": { + "policy": "vllm_v1", + "max_num_seqs": 4, + "max_num_batched_tokens": 16, + "max_model_len": 32, + "prefill": { + "mode": "chunked", + "max_chunk_tokens": 8 + } + }, + "timing": { + "grid": { + "decode_reqs": [0, 2], + "sum_decode_ctx_tokens": [0, 20], + "prefill_tokens_in_step": [0, 10], + "step_duration_us": [10, 80, 70, 140, 210, 280, 270, 340] + }, + "fallback": { + "t0_us": 1.0, + "prefill_token_us": 2.0, + "decode_request_us": 3.0, + "decode_context_token_us": 0.5 + } + } + }) + } + + fn parse(value: &Value) -> Result { + EngineProfile::from_json_slice(&serde_json::to_vec(value)?) + } + + #[test] + fn profile_rejects_invalid_scheduler_and_grid() { + let mut bad_chunk = profile_json(); + bad_chunk["scheduler"]["prefill"]["max_chunk_tokens"] = json!(17); + assert!( + parse(&bad_chunk) + .unwrap_err() + .to_string() + .contains("cannot exceed max_num_batched_tokens") + ); + + let mut unordered_axis = profile_json(); + unordered_axis["timing"]["grid"]["decode_reqs"] = json!([2, 0]); + assert!( + parse(&unordered_axis) + .unwrap_err() + .to_string() + .contains("must be strictly increasing") + ); + + let mut missing_value = profile_json(); + missing_value["timing"]["grid"]["step_duration_us"] = json!([1, 2]); + assert!( + parse(&missing_value) + .unwrap_err() + .to_string() + .contains("contains 2 values; expected 8") + ); + } + + #[test] + fn grid_returns_exact_points_and_deterministic_interpolation() { + let profile = parse(&profile_json()).unwrap(); + let exact = profile + .estimate_step( + StepShape { + decode_reqs: 2, + sum_decode_ctx_tokens: 20, + prefill_tokens_in_step: 10, + }, + OutOfDomainPolicy::Strict, + ) + .unwrap(); + assert_eq!( + exact, + StepTimingEstimate { + duration_us: 340, + source: StepTimingSource::GridInterpolation, + } + ); + + let interpolated = profile + .estimate_step( + StepShape { + decode_reqs: 1, + sum_decode_ctx_tokens: 10, + prefill_tokens_in_step: 5, + }, + OutOfDomainPolicy::Strict, + ) + .unwrap(); + assert_eq!(interpolated.duration_us, 175); + assert_eq!(interpolated.source, StepTimingSource::GridInterpolation); + } + + #[test] + fn out_of_domain_shape_warns_and_falls_back_or_fails_strict() { + let profile = parse(&profile_json()).unwrap(); + let shape = StepShape { + decode_reqs: 1, + sum_decode_ctx_tokens: 10, + prefill_tokens_in_step: 12, + }; + let fallback = profile + .estimate_step(shape, OutOfDomainPolicy::WarnAndFallback) + .unwrap(); + assert_eq!( + fallback, + StepTimingEstimate { + duration_us: 33, + source: StepTimingSource::ParametricFallback, + } + ); + + let error = profile + .estimate_step(shape, OutOfDomainPolicy::Strict) + .unwrap_err(); + assert!(error.to_string().contains("does not cover step shape")); + } + + #[test] + fn invalid_step_shapes_fail_before_fallback() { + let profile = parse(&profile_json()).unwrap(); + let no_work = profile + .estimate_step( + StepShape { + decode_reqs: 0, + sum_decode_ctx_tokens: 0, + prefill_tokens_in_step: 0, + }, + OutOfDomainPolicy::WarnAndFallback, + ) + .unwrap_err(); + assert!( + no_work + .to_string() + .contains("must contain prefill or decode") + ); + + let over_budget = profile + .estimate_step( + StepShape { + decode_reqs: 4, + sum_decode_ctx_tokens: 16, + prefill_tokens_in_step: 13, + }, + OutOfDomainPolicy::WarnAndFallback, + ) + .unwrap_err(); + assert!(over_budget.to_string().contains("exceeds scheduler")); + + let context_without_decode = profile + .estimate_step( + StepShape { + decode_reqs: 0, + sum_decode_ctx_tokens: 1, + prefill_tokens_in_step: 1, + }, + OutOfDomainPolicy::WarnAndFallback, + ) + .unwrap_err(); + assert!( + context_without_decode + .to_string() + .contains("must be zero when decode_reqs is zero") + ); + } + + #[test] + fn zero_cost_profile_is_valid_and_fallback_overflow_is_rejected() { + let mut zero_cost = profile_json(); + zero_cost["timing"]["grid"]["step_duration_us"] = json!([0, 0, 0, 0, 0, 0, 0, 0]); + zero_cost["timing"]["fallback"] = json!({ + "t0_us": 0.0, + "prefill_token_us": 0.0, + "decode_request_us": 0.0, + "decode_context_token_us": 0.0 + }); + let profile = parse(&zero_cost).unwrap(); + assert_eq!( + profile + .estimate_step( + StepShape { + decode_reqs: 1, + sum_decode_ctx_tokens: 10, + prefill_tokens_in_step: 12, + }, + OutOfDomainPolicy::WarnAndFallback, + ) + .unwrap() + .duration_us, + 0 + ); + + let mut overflow = profile_json(); + overflow["timing"]["fallback"]["t0_us"] = json!(1.0e308); + let profile = parse(&overflow).unwrap(); + let error = profile + .estimate_step( + StepShape { + decode_reqs: 1, + sum_decode_ctx_tokens: 10, + prefill_tokens_in_step: 12, + }, + OutOfDomainPolicy::WarnAndFallback, + ) + .unwrap_err(); + assert!(error.to_string().contains("fallback overflow")); + } +} diff --git a/pegainfer-sim/src/worker.rs b/pegainfer-sim/src/worker.rs new file mode 100644 index 000000000..72ed59918 --- /dev/null +++ b/pegainfer-sim/src/worker.rs @@ -0,0 +1,749 @@ +use std::collections::VecDeque; +use std::fmt::Display; + +use anyhow::Context; +use anyhow::Result; +use anyhow::ensure; + +use crate::profile::PrefillPolicy; +use crate::profile::SchedulerProfile; +use crate::profile::StepShape; + +#[derive(Clone, Copy, Debug, Eq, PartialEq)] +pub enum RequestPhase { + Waiting, + Prefill, + Decode, +} + +#[derive(Clone, Copy, Debug, Eq, PartialEq)] +pub struct WorkerRequest { + pub id: I, + pub prompt_tokens: u32, + pub output_tokens: u32, +} + +#[derive(Clone, Copy, Debug, Eq, PartialEq)] +pub enum RequestRejection { + ModelLengthExceeded { + total_tokens: u64, + max_model_len: u32, + }, + WholePrefillExceedsStepBudget { + prompt_tokens: u32, + max_num_batched_tokens: u32, + }, +} + +impl Display for RequestRejection { + fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + match self { + Self::ModelLengthExceeded { + total_tokens, + max_model_len, + } => write!( + formatter, + "request total tokens {total_tokens} exceed max_model_len {max_model_len}" + ), + Self::WholePrefillExceedsStepBudget { + prompt_tokens, + max_num_batched_tokens, + } => write!( + formatter, + "whole prefill has {prompt_tokens} tokens but max_num_batched_tokens is {max_num_batched_tokens}" + ), + } + } +} + +#[derive(Clone, Copy, Debug, Eq, PartialEq)] +pub enum SubmissionResult { + Queued, + Finished, + Rejected(RequestRejection), +} + +#[derive(Clone, Copy, Debug, Eq, PartialEq)] +pub enum CancelResult { + Cancelled, + Deferred, + AlreadyRequested, + NotFound, +} + +#[derive(Clone, Copy, Debug, Eq, PartialEq)] +pub struct RequestSnapshot { + pub phase: RequestPhase, + pub remaining_prefill_tokens: u32, + pub generated_tokens: u32, + /// Tokens whose KV has been computed before the next decode query. + pub context_tokens: u32, + pub output_tokens: u32, + pub cancel_requested: bool, +} + +#[derive(Clone, Copy, Debug, Eq, PartialEq)] +pub struct StepId(u64); + +impl StepId { + #[must_use] + pub fn get(self) -> u64 { + self.0 + } +} + +#[derive(Clone, Copy, Debug, Eq, PartialEq)] +pub struct PrefillWork { + pub request_id: I, + pub tokens: u32, +} + +#[derive(Clone, Copy, Debug, Eq, PartialEq)] +pub struct DecodeWork { + pub request_id: I, + /// Existing KV length before this decode query is computed. + pub context_tokens: u32, +} + +#[derive(Debug, Eq, PartialEq)] +pub struct StepPlan { + id: StepId, + shape: StepShape, + admitted: Vec, + prefill: Vec>, + decode: Vec>, +} + +impl StepPlan { + #[must_use] + pub fn id(&self) -> StepId { + self.id + } + + #[must_use] + pub fn shape(&self) -> StepShape { + self.shape + } + + #[must_use] + pub fn admitted(&self) -> &[I] { + &self.admitted + } + + #[must_use] + pub fn prefill(&self) -> &[PrefillWork] { + &self.prefill + } + + #[must_use] + pub fn decode(&self) -> &[DecodeWork] { + &self.decode + } +} + +#[derive(Clone, Copy, Debug, Eq, PartialEq)] +pub struct PrefillProgress { + pub request_id: I, + pub processed_tokens: u32, + pub remaining_tokens: u32, + pub context_tokens: u32, +} + +#[derive(Clone, Copy, Debug, Eq, PartialEq)] +pub struct GeneratedToken { + pub request_id: I, + pub token_index: u32, +} + +#[derive(Debug, Eq, PartialEq)] +pub struct StepOutcome { + pub prefill: Vec>, + pub generated: Vec>, + pub finished: Vec, + pub cancelled: Vec, +} + +struct RequestState { + id: I, + phase: RequestPhase, + remaining_prefill_tokens: u32, + generated_tokens: u32, + context_tokens: u32, + output_tokens: u32, + cancel_requested: bool, +} + +impl RequestState { + fn new(request: WorkerRequest) -> Self { + Self { + id: request.id, + phase: RequestPhase::Waiting, + remaining_prefill_tokens: request.prompt_tokens, + generated_tokens: 0, + context_tokens: 0, + output_tokens: request.output_tokens, + cancel_requested: false, + } + } + + fn snapshot(&self) -> RequestSnapshot { + RequestSnapshot { + phase: self.phase, + remaining_prefill_tokens: self.remaining_prefill_tokens, + generated_tokens: self.generated_tokens, + context_tokens: self.context_tokens, + output_tokens: self.output_tokens, + cancel_requested: self.cancel_requested, + } + } +} + +pub struct WorkerState { + scheduler: SchedulerProfile, + waiting: VecDeque>, + running: Vec>, + in_flight: Option>, + next_step_id: u64, +} + +impl WorkerState +where + I: Copy + Eq, +{ + pub fn new(scheduler: SchedulerProfile) -> Result { + scheduler.validate()?; + Ok(Self { + scheduler, + waiting: VecDeque::new(), + running: Vec::new(), + in_flight: None, + next_step_id: 0, + }) + } + + #[must_use] + pub fn scheduler(&self) -> &SchedulerProfile { + &self.scheduler + } + + pub(crate) fn preflight( + &self, + prompt_tokens: u32, + output_tokens: u64, + ) -> Option { + let total_tokens = u64::from(prompt_tokens).saturating_add(output_tokens); + if total_tokens > u64::from(self.scheduler.max_model_len) { + return Some(RequestRejection::ModelLengthExceeded { + total_tokens, + max_model_len: self.scheduler.max_model_len, + }); + } + if output_tokens > 0 + && self.scheduler.prefill == PrefillPolicy::Whole + && prompt_tokens > self.scheduler.max_num_batched_tokens + { + return Some(RequestRejection::WholePrefillExceedsStepBudget { + prompt_tokens, + max_num_batched_tokens: self.scheduler.max_num_batched_tokens, + }); + } + None + } + + pub fn submit(&mut self, request: WorkerRequest) -> Result { + ensure!( + self.request(request.id).is_none(), + "worker request id is already active" + ); + if let Some(rejection) = + self.preflight(request.prompt_tokens, u64::from(request.output_tokens)) + { + return Ok(SubmissionResult::Rejected(rejection)); + } + if request.output_tokens == 0 { + return Ok(SubmissionResult::Finished); + } + self.waiting.push_back(RequestState::new(request)); + Ok(SubmissionResult::Queued) + } + + #[must_use] + pub fn waiting_len(&self) -> usize { + self.waiting.len() + } + + #[must_use] + pub fn running_len(&self) -> usize { + self.running.len() + } + + #[must_use] + pub fn is_idle(&self) -> bool { + self.waiting.is_empty() && self.running.is_empty() && self.in_flight.is_none() + } + + #[must_use] + pub fn in_flight(&self) -> Option<&StepPlan> { + self.in_flight.as_ref() + } + + #[must_use] + pub fn request(&self, id: I) -> Option { + self.running + .iter() + .chain(self.waiting.iter()) + .find(|request| request.id == id) + .map(RequestState::snapshot) + } + + pub fn cancel(&mut self, id: I) -> CancelResult { + if let Some(index) = self.waiting.iter().position(|request| request.id == id) { + self.waiting.remove(index); + return CancelResult::Cancelled; + } + let Some(index) = self.running.iter().position(|request| request.id == id) else { + return CancelResult::NotFound; + }; + if self.running[index].cancel_requested { + return CancelResult::AlreadyRequested; + } + if self.in_flight.is_some() { + self.running[index].cancel_requested = true; + CancelResult::Deferred + } else { + self.running.remove(index); + CancelResult::Cancelled + } + } + + pub fn plan_step(&mut self) -> Result>> { + ensure!( + self.in_flight.is_none(), + "cannot plan a worker step while another step is in flight" + ); + if self.running.is_empty() && self.waiting.is_empty() { + return Ok(None); + } + + let step_id = StepId(self.next_step_id); + let next_step_id = self + .next_step_id + .checked_add(1) + .context("worker step id overflow")?; + let mut remaining_budget = self.scheduler.max_num_batched_tokens; + let mut decode = Vec::new(); + let mut sum_decode_ctx_tokens = 0_u64; + + // vLLM V1 keeps every running decode scheduled before spending the + // remaining token budget on prefill work. + for request in &self.running { + if request.phase != RequestPhase::Decode || request.cancel_requested { + continue; + } + ensure!(remaining_budget > 0, "running decode exceeded token budget"); + remaining_budget -= 1; + sum_decode_ctx_tokens = sum_decode_ctx_tokens + .checked_add(u64::from(request.context_tokens)) + .context("decode context sum overflow")?; + decode.push(DecodeWork { + request_id: request.id, + context_tokens: request.context_tokens, + }); + } + + let mut prefill = Vec::new(); + let mut prefill_tokens_in_step = 0_u32; + let mut prefill_blocked = false; + for request in &self.running { + if request.phase != RequestPhase::Prefill || request.cancel_requested { + continue; + } + let tokens = scheduled_prefill_tokens( + self.scheduler.prefill, + request.remaining_prefill_tokens, + remaining_budget, + ); + if tokens == 0 { + prefill_blocked = true; + break; + } + remaining_budget -= tokens; + prefill_tokens_in_step += tokens; + prefill.push(PrefillWork { + request_id: request.id, + tokens, + }); + } + + let mut admitted = Vec::new(); + while !prefill_blocked + && remaining_budget > 0 + && self.running.len() < self.scheduler.max_num_seqs as usize + && !self.waiting.is_empty() + { + let request = self + .waiting + .front() + .expect("admission loop checked that a waiting request exists"); + let (phase, prefill_tokens) = if request.remaining_prefill_tokens == 0 { + (RequestPhase::Decode, 0) + } else { + let tokens = scheduled_prefill_tokens( + self.scheduler.prefill, + request.remaining_prefill_tokens, + remaining_budget, + ); + if tokens == 0 { + break; + } + (RequestPhase::Prefill, tokens) + }; + + let mut request = self + .waiting + .pop_front() + .context("waiting request disappeared during admission")?; + request.phase = phase; + admitted.push(request.id); + if phase == RequestPhase::Decode { + remaining_budget -= 1; + decode.push(DecodeWork { + request_id: request.id, + context_tokens: 0, + }); + } else { + remaining_budget -= prefill_tokens; + prefill_tokens_in_step += prefill_tokens; + prefill.push(PrefillWork { + request_id: request.id, + tokens: prefill_tokens, + }); + } + self.running.push(request); + } + + let decode_reqs = u32::try_from(decode.len()).context("decode request count overflow")?; + let shape = StepShape { + decode_reqs, + sum_decode_ctx_tokens, + prefill_tokens_in_step, + }; + ensure!( + shape.decode_reqs > 0 || shape.prefill_tokens_in_step > 0, + "worker with queued or running requests produced an empty step" + ); + ensure!( + shape.decode_reqs <= self.scheduler.max_num_seqs, + "planned decode requests exceed sequence capacity" + ); + ensure!( + shape.decode_reqs + shape.prefill_tokens_in_step + <= self.scheduler.max_num_batched_tokens, + "planned step exceeds token capacity" + ); + + self.next_step_id = next_step_id; + self.in_flight = Some(StepPlan { + id: step_id, + shape, + admitted, + prefill, + decode, + }); + Ok(self.in_flight.as_ref()) + } + + pub fn complete_step(&mut self, step_id: StepId) -> Result> { + let current_id = self + .in_flight + .as_ref() + .map(StepPlan::id) + .context("cannot complete a worker step when no step is in flight")?; + ensure!( + current_id == step_id, + "worker step id {} does not match in-flight step {}", + step_id.get(), + current_id.get() + ); + let plan = self.in_flight.take().expect("in-flight step was checked"); + let mut outcome = StepOutcome { + prefill: Vec::with_capacity(plan.prefill.len()), + generated: Vec::with_capacity(plan.prefill.len() + plan.decode.len()), + finished: Vec::new(), + cancelled: Vec::new(), + }; + + for work in plan.prefill { + let request = self + .running + .iter_mut() + .find(|request| request.id == work.request_id) + .context("planned prefill request is no longer running")?; + if request.cancel_requested { + continue; + } + ensure!( + request.phase == RequestPhase::Prefill + && work.tokens > 0 + && work.tokens <= request.remaining_prefill_tokens, + "invalid prefill work in completed step" + ); + request.remaining_prefill_tokens -= work.tokens; + request.context_tokens = request + .context_tokens + .checked_add(work.tokens) + .context("prefill context length overflow")?; + outcome.prefill.push(PrefillProgress { + request_id: request.id, + processed_tokens: work.tokens, + remaining_tokens: request.remaining_prefill_tokens, + context_tokens: request.context_tokens, + }); + if request.remaining_prefill_tokens == 0 { + request.phase = RequestPhase::Decode; + // The final prefill logits produce the first output token; a + // separate decode step here would add a spurious TTFT step. + request.generated_tokens = request + .generated_tokens + .checked_add(1) + .context("generated token count overflow")?; + outcome.generated.push(GeneratedToken { + request_id: request.id, + token_index: request.generated_tokens, + }); + if request.generated_tokens == request.output_tokens { + outcome.finished.push(request.id); + } + } + } + + for work in plan.decode { + let request = self + .running + .iter_mut() + .find(|request| request.id == work.request_id) + .context("planned decode request is no longer running")?; + if request.cancel_requested { + continue; + } + ensure!( + request.phase == RequestPhase::Decode + && request.context_tokens == work.context_tokens + && request.generated_tokens < request.output_tokens, + "invalid decode work in completed step" + ); + request.context_tokens = request + .context_tokens + .checked_add(1) + .context("decode context length overflow")?; + request.generated_tokens = request + .generated_tokens + .checked_add(1) + .context("generated token count overflow")?; + outcome.generated.push(GeneratedToken { + request_id: request.id, + token_index: request.generated_tokens, + }); + if request.generated_tokens == request.output_tokens { + outcome.finished.push(request.id); + } + } + + outcome.cancelled.extend( + self.running + .iter() + .filter(|request| request.cancel_requested) + .map(|request| request.id), + ); + self.running + .retain(|request| !request.cancel_requested && !outcome.finished.contains(&request.id)); + Ok(outcome) + } +} + +fn scheduled_prefill_tokens( + policy: PrefillPolicy, + remaining_tokens: u32, + remaining_budget: u32, +) -> u32 { + match policy { + PrefillPolicy::Whole => { + if remaining_tokens <= remaining_budget { + remaining_tokens + } else { + 0 + } + } + PrefillPolicy::Chunked { max_chunk_tokens } => { + remaining_tokens.min(max_chunk_tokens).min(remaining_budget) + } + } +} + +#[cfg(test)] +mod tests { + use super::*; + use crate::profile::SchedulerPolicy; + + fn scheduler( + max_num_seqs: u32, + max_num_batched_tokens: u32, + max_model_len: u32, + prefill: PrefillPolicy, + ) -> SchedulerProfile { + SchedulerProfile { + policy: SchedulerPolicy::VllmV1, + max_num_seqs, + max_num_batched_tokens, + max_model_len, + prefill, + } + } + + fn request(id: u32, prompt_tokens: u32, output_tokens: u32) -> WorkerRequest { + WorkerRequest { + id, + prompt_tokens, + output_tokens, + } + } + + #[test] + fn impossible_and_duplicate_requests_are_rejected() { + let mut whole = WorkerState::new(scheduler(2, 4, 8, PrefillPolicy::Whole)).unwrap(); + assert_eq!( + whole.submit(request(1, 7, 2)).unwrap(), + SubmissionResult::Rejected(RequestRejection::ModelLengthExceeded { + total_tokens: 9, + max_model_len: 8, + }) + ); + assert_eq!( + whole.submit(request(2, 5, 1)).unwrap(), + SubmissionResult::Rejected(RequestRejection::WholePrefillExceedsStepBudget { + prompt_tokens: 5, + max_num_batched_tokens: 4, + }) + ); + assert_eq!( + whole.submit(request(3, 1, 1)).unwrap(), + SubmissionResult::Queued + ); + assert!(whole.submit(request(3, 1, 1)).is_err()); + assert_eq!( + whole.submit(request(4, 4, 0)).unwrap(), + SubmissionResult::Finished + ); + } + + #[test] + fn context_and_generated_progress_end_in_one_terminal_outcome() { + let mut worker = WorkerState::new(scheduler(1, 4, 16, PrefillPolicy::Whole)).unwrap(); + worker.submit(request(1, 3, 3)).unwrap(); + + let prefill_id = worker.plan_step().unwrap().unwrap().id(); + let prefill = worker.complete_step(prefill_id).unwrap(); + assert_eq!( + prefill.generated, + [GeneratedToken { + request_id: 1, + token_index: 1 + }] + ); + assert_eq!(worker.request(1).unwrap().context_tokens, 3); + + let decode_one = worker.plan_step().unwrap().unwrap(); + assert_eq!(decode_one.shape().sum_decode_ctx_tokens, 3); + let decode_one_id = decode_one.id(); + let outcome = worker.complete_step(decode_one_id).unwrap(); + assert_eq!(outcome.generated[0].token_index, 2); + assert_eq!(worker.request(1).unwrap().context_tokens, 4); + + let decode_two = worker.plan_step().unwrap().unwrap(); + assert_eq!(decode_two.shape().sum_decode_ctx_tokens, 4); + let decode_two_id = decode_two.id(); + let terminal = worker.complete_step(decode_two_id).unwrap(); + assert_eq!(terminal.finished, [1]); + assert!(worker.request(1).is_none()); + assert!(worker.plan_step().unwrap().is_none()); + } + + #[test] + fn cancellation_after_nonterminal_step_completion_removes_running_request() { + let mut worker = WorkerState::new(scheduler(1, 4, 16, PrefillPolicy::Whole)).unwrap(); + worker.submit(request(1, 0, 3)).unwrap(); + + let step_id = worker.plan_step().unwrap().unwrap().id(); + let outcome = worker.complete_step(step_id).unwrap(); + assert_eq!( + outcome.generated, + [GeneratedToken { + request_id: 1, + token_index: 1 + }] + ); + assert!(outcome.finished.is_empty()); + assert!(worker.request(1).is_some()); + + assert_eq!(worker.cancel(1), CancelResult::Cancelled); + assert!(worker.request(1).is_none()); + assert!(worker.plan_step().unwrap().is_none()); + } + + #[test] + fn generated_plans_never_exceed_sequence_or_token_capacity() { + for max_num_seqs in 1..=4 { + for max_num_batched_tokens in max_num_seqs..=6 { + for max_chunk_tokens in 1..=max_num_batched_tokens { + let mut worker = WorkerState::new(scheduler( + max_num_seqs, + max_num_batched_tokens, + 32, + PrefillPolicy::Chunked { max_chunk_tokens }, + )) + .unwrap(); + let mut terminal_requests = 0; + for id in 0..12 { + let prompt_tokens = id % 8; + let output_tokens = id % 3 + 1; + assert_eq!( + worker + .submit(request(id, prompt_tokens, output_tokens)) + .unwrap(), + SubmissionResult::Queued + ); + } + + for _ in 0..256 { + let Some(plan) = worker.plan_step().unwrap() else { + break; + }; + let shape = plan.shape(); + assert!(shape.decode_reqs <= max_num_seqs); + assert!( + shape.decode_reqs + shape.prefill_tokens_in_step + <= max_num_batched_tokens + ); + assert!( + plan.prefill() + .iter() + .all(|work| work.tokens <= max_chunk_tokens) + ); + assert_eq!( + plan.decode() + .iter() + .map(|work| u64::from(work.context_tokens)) + .sum::(), + shape.sum_decode_ctx_tokens + ); + let step_id = plan.id(); + assert!(worker.running_len() <= max_num_seqs as usize); + terminal_requests += worker.complete_step(step_id).unwrap().finished.len(); + } + assert!(worker.is_idle()); + assert_eq!(terminal_requests, 12); + } + } + } + } +} diff --git a/pegainfer-sim/tests/fixtures/online-step-gate.json b/pegainfer-sim/tests/fixtures/online-step-gate.json new file mode 100644 index 000000000..8c1a91376 --- /dev/null +++ b/pegainfer-sim/tests/fixtures/online-step-gate.json @@ -0,0 +1,52 @@ +{ + "schema_version": 1, + "profile_id": "online-step-gate", + "provenance": { + "target_engine": "vllm", + "engine_version": "fixture", + "model_id": "pegainfer-sim-online-gate", + "model_revision": "fixture-revision", + "model_config_sha256": "0000000000000000000000000000000000000000000000000000000000000000", + "gpu": "fixture-cpu", + "server_flags": [ + "--max-num-seqs=2", + "--max-num-batched-tokens=4", + "--max-model-len=16", + "--enable-chunked-prefill" + ] + }, + "scheduler": { + "policy": "vllm_v1", + "max_num_seqs": 2, + "max_num_batched_tokens": 4, + "max_model_len": 16, + "prefill": { + "mode": "chunked", + "max_chunk_tokens": 2 + } + }, + "timing": { + "grid": { + "decode_reqs": [0, 1, 2], + "sum_decode_ctx_tokens": [0, 2, 4], + "prefill_tokens_in_step": [0, 2, 4], + "step_duration_us": [ + 50000, 150000, 250000, + 60000, 160000, 260000, + 70000, 170000, 270000, + 75000, 175000, 275000, + 85000, 185000, 285000, + 95000, 195000, 295000, + 100000, 200000, 300000, + 110000, 210000, 310000, + 120000, 220000, 320000 + ] + }, + "fallback": { + "t0_us": 10000.0, + "prefill_token_us": 10000.0, + "decode_request_us": 5000.0, + "decode_context_token_us": 1000.0 + } + } +} diff --git a/pegainfer-sim/tests/fixtures/online-zero-cost.json b/pegainfer-sim/tests/fixtures/online-zero-cost.json new file mode 100644 index 000000000..dec3c48f8 --- /dev/null +++ b/pegainfer-sim/tests/fixtures/online-zero-cost.json @@ -0,0 +1,52 @@ +{ + "schema_version": 1, + "profile_id": "online-zero-cost", + "provenance": { + "target_engine": "vllm", + "engine_version": "fixture", + "model_id": "pegainfer-sim-zero-cost", + "model_revision": "fixture-revision", + "model_config_sha256": "0000000000000000000000000000000000000000000000000000000000000000", + "gpu": "fixture-cpu", + "server_flags": [ + "--max-num-seqs=2", + "--max-num-batched-tokens=4", + "--max-model-len=16", + "--enable-chunked-prefill" + ] + }, + "scheduler": { + "policy": "vllm_v1", + "max_num_seqs": 2, + "max_num_batched_tokens": 4, + "max_model_len": 16, + "prefill": { + "mode": "chunked", + "max_chunk_tokens": 2 + } + }, + "timing": { + "grid": { + "decode_reqs": [0, 1, 2], + "sum_decode_ctx_tokens": [0, 2, 4], + "prefill_tokens_in_step": [0, 2, 4], + "step_duration_us": [ + 0, 0, 0, + 0, 0, 0, + 0, 0, 0, + 0, 0, 0, + 0, 0, 0, + 0, 0, 0, + 0, 0, 0, + 0, 0, 0, + 0, 0, 0 + ] + }, + "fallback": { + "t0_us": 0.0, + "prefill_token_us": 0.0, + "decode_request_us": 0.0, + "decode_context_token_us": 0.0 + } + } +} diff --git a/pegainfer-sim/tests/frontend_e2e.rs b/pegainfer-sim/tests/frontend_e2e.rs index e26b72845..12a262907 100644 --- a/pegainfer-sim/tests/frontend_e2e.rs +++ b/pegainfer-sim/tests/frontend_e2e.rs @@ -1,14 +1,20 @@ use std::fs; use std::net::TcpListener; use std::time::Duration; +use std::time::Instant; use anyhow::Context; use anyhow::Result; use anyhow::anyhow; use anyhow::bail; use pegainfer_sim::SimulatedEngineConfig; +use pegainfer_sim::profile::EngineProfile; +use pegainfer_sim::profile::OutOfDomainPolicy; +use pegainfer_sim::profile::StepShape; use pegainfer_sim::start_engine; use pegainfer_sim::start_engine_with_partitions; +use pegainfer_sim::worker::WorkerRequest; +use pegainfer_sim::worker::WorkerState; use reqwest::Client; use serde_json::Value; use serde_json::json; @@ -21,6 +27,8 @@ const MODEL_NAME: &str = "pegainfer-sim-e2e"; const METRICS_MODEL_NAME: &str = "pegainfer-sim-e2e-metrics"; const SLOW_METRICS_MODEL_NAME: &str = "pegainfer-sim-e2e-slow-metrics"; const SPEC_METRICS_MODEL_NAME: &str = "pegainfer-sim-e2e-spec-metrics"; +const PROFILE_GATE_MODEL_NAME: &str = "pegainfer-sim-online-gate"; +const ZERO_COST_MODEL_NAME: &str = "pegainfer-sim-zero-cost"; /// The pretend drafter the spec-metrics server runs: `K` and how many of those /// draft tokens each verify step accepts. const SPEC_K: usize = 3; @@ -185,6 +193,94 @@ struct StartedSimServer { task: JoinHandle>, } +#[derive(Debug, Eq, PartialEq)] +struct PlanTrace { + shape: StepShape, + duration_us: u64, + admitted: Vec, + prefill: Vec<(u64, u32)>, + decode: Vec<(u64, u32)>, +} + +fn profile_fixture(name: &str) -> Result { + let bytes: &[u8] = match name { + "online-step-gate.json" => include_bytes!("fixtures/online-step-gate.json"), + "online-zero-cost.json" => include_bytes!("fixtures/online-zero-cost.json"), + other => bail!("unknown profile fixture {other}"), + }; + EngineProfile::from_json_slice(bytes) +} + +fn replay_profiled_worker(profile: &EngineProfile) -> Result> { + let mut worker = WorkerState::new(profile.scheduler.clone())?; + for request in [ + WorkerRequest { + id: 1_u64, + prompt_tokens: 4, + output_tokens: 2, + }, + WorkerRequest { + id: 2, + prompt_tokens: 2, + output_tokens: 2, + }, + WorkerRequest { + id: 3, + prompt_tokens: 1, + output_tokens: 2, + }, + ] { + assert!(matches!( + worker.submit(request)?, + pegainfer_sim::worker::SubmissionResult::Queued + )); + } + + let mut trace = Vec::new(); + let mut observed_waiting = false; + while !worker.is_idle() { + observed_waiting |= worker.waiting_len() > 0; + let plan = worker + .plan_step()? + .context("worker with active requests produced no step")?; + let step_id = plan.id(); + let shape = plan.shape(); + let estimate = profile.estimate_step(shape, OutOfDomainPolicy::Strict)?; + assert!(matches!( + estimate.source, + pegainfer_sim::profile::StepTimingSource::GridInterpolation + )); + assert!( + shape.decode_reqs <= profile.scheduler.max_num_seqs, + "decode request count exceeded scheduler capacity: {shape:?}" + ); + assert!( + shape.decode_reqs + shape.prefill_tokens_in_step + <= profile.scheduler.max_num_batched_tokens, + "step token count exceeded scheduler capacity: {shape:?}" + ); + let entry = PlanTrace { + shape, + duration_us: estimate.duration_us, + admitted: plan.admitted().to_vec(), + prefill: plan + .prefill() + .iter() + .map(|work| (work.request_id, work.tokens)) + .collect(), + decode: plan + .decode() + .iter() + .map(|work| (work.request_id, work.context_tokens)) + .collect(), + }; + worker.complete_step(step_id)?; + trace.push(entry); + } + assert!(observed_waiting, "replay must exercise sequence admission"); + Ok(trace) +} + fn empty_model_dir() -> Result { tempfile::tempdir().context("failed to create temp model dir") } @@ -221,6 +317,201 @@ async fn simulated_engine_serves_openai_completions_over_http() -> Result<()> { server.shutdown().await } +#[test] +fn profiled_worker_replay_is_deterministic_and_bounded() -> Result<()> { + let profile = profile_fixture("online-step-gate.json")?; + let first = replay_profiled_worker(&profile)?; + let second = replay_profiled_worker(&profile)?; + + assert_eq!( + first, second, + "same profile and arrivals must replay identically" + ); + assert!( + first + .iter() + .any(|step| step.shape.prefill_tokens_in_step > 0), + "replay must include prefill work" + ); + assert!( + first.iter().any(|step| step.shape.decode_reqs > 0), + "replay must include decode work" + ); + assert!( + first + .iter() + .map(|step| step.duration_us) + .min() + .is_some_and(|minimum| { + first + .iter() + .map(|step| step.duration_us) + .any(|duration| duration > minimum) + }), + "grid pricing must vary across the replayed step shapes" + ); + + Ok(()) +} + +#[test] +fn zero_cost_profile_fixture_is_valid() -> Result<()> { + let profile = profile_fixture("online-zero-cost.json")?; + assert!( + profile + .timing + .grid + .step_duration_us + .iter() + .all(|duration| *duration == 0), + "zero-cost fixture must not add synthetic engine delay" + ); + let estimate = profile.estimate_step( + StepShape { + decode_reqs: 1, + sum_decode_ctx_tokens: 2, + prefill_tokens_in_step: 0, + }, + OutOfDomainPolicy::Strict, + )?; + assert_eq!(estimate.duration_us, 0); + + Ok(()) +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 2)] +async fn profiled_online_worker_serves_multi_request_workload() -> Result<()> { + let profile = profile_fixture("online-step-gate.json")?; + let config = + SimulatedEngineConfig::default().with_engine_profile(profile, OutOfDomainPolicy::Strict)?; + let server = SimServer::spawn_with_config( + model_dir_with_minimal_metadata()?, + 1, + PROFILE_GATE_MODEL_NAME, + config, + ) + .await?; + let client = test_client()?; + + assert_models_endpoint(&client, &server.base_url, PROFILE_GATE_MODEL_NAME).await?; + + let base_url = server.base_url.clone(); + let post_request = move |client: Client| { + let url = format!("{base_url}/v1/completions"); + tokio::spawn(async move { + client + .post(url) + .header(reqwest::header::CONTENT_TYPE, "application/json") + .json(&json!({ + "model": PROFILE_GATE_MODEL_NAME, + "prompt": [1, 2], + "max_tokens": 2, + "temperature": 0.0, + "ignore_eos": true + })) + .send() + .await? + .error_for_status()? + .json::() + .await + .context("failed to parse profiled completion") + }) + }; + + let mut requests = vec![post_request(client.clone()), post_request(client.clone())]; + wait_for_metrics( + &client, + &server.base_url, + &[("vllm:num_requests_running", "0", 2.0)], + PROFILE_GATE_MODEL_NAME, + ) + .await?; + + // Capacity ordering is asserted by the deterministic worker replay; the + // HTTP gate checks stable post-drain counters instead of a transient scrape. + requests.push(post_request(client.clone())); + + for request in requests { + let response = request + .await + .context("profiled completion task panicked")??; + assert_eq!( + response["choices"][0]["finish_reason"].as_str(), + Some("length"), + "profiled completion must consume the requested output budget: {response}" + ); + } + wait_for_metrics( + &client, + &server.base_url, + &[ + ("vllm:num_requests_running", "0", 0.0), + ("vllm:num_requests_waiting", "0", 0.0), + ("vllm:prompt_tokens_total", "0", 6.0), + ("vllm:generation_tokens_total", "0", 6.0), + ], + PROFILE_GATE_MODEL_NAME, + ) + .await?; + wait_for_labeled_metrics( + &client, + &server.base_url, + &[( + "vllm:request_success_total", + "0", + &[("finished_reason", "length")][..], + 3.0, + )], + PROFILE_GATE_MODEL_NAME, + ) + .await?; + + server.shutdown().await +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 2)] +async fn zero_cost_profile_measures_frontend_baseline() -> Result<()> { + let profile = profile_fixture("online-zero-cost.json")?; + let config = + SimulatedEngineConfig::default().with_engine_profile(profile, OutOfDomainPolicy::Strict)?; + let server = SimServer::spawn_with_config( + model_dir_with_minimal_metadata()?, + 1, + ZERO_COST_MODEL_NAME, + config, + ) + .await?; + let client = test_client()?; + let started = Instant::now(); + let response = client + .post(format!("{}/v1/completions", server.base_url)) + .header(reqwest::header::CONTENT_TYPE, "application/json") + .json(&json!({ + "model": ZERO_COST_MODEL_NAME, + "prompt": [1, 2], + "max_tokens": 2, + "temperature": 0.0, + "ignore_eos": true + })) + .send() + .await? + .error_for_status()? + .json::() + .await?; + let elapsed = started.elapsed(); + assert_eq!( + response["choices"][0]["finish_reason"].as_str(), + Some("length"), + "zero-cost completion must still traverse the normal frontend: {response}" + ); + eprintln!( + "zero-cost profile frontend baseline: {:.3} ms", + elapsed.as_secs_f64() * 1_000.0 + ); + + server.shutdown().await +} + #[tokio::test(flavor = "multi_thread", worker_threads = 2)] async fn one_http_endpoint_exports_per_engine_scheduler_metrics() -> Result<()> { let server = SimServer::spawn_partitioned().await?;