Skip to main content

aisimulate_core/replay/
engine.rs

1// SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
2// SPDX-License-Identifier: Apache-2.0
3
4//! Default generalized-engine construction for offline replay.
5//!
6//! This module contains configuration and conversion helpers only. Scheduler
7//! state lives in `aisimulate_core::engine`; virtual time and worker lifecycle live
8//! in the moved aggregated/disaggregated replay runtimes.
9
10use std::num::NonZeroU32;
11use std::sync::Arc;
12
13use crate::engine::generalized::EngineIdentity;
14use crate::engine::{Backend, Engine, EngineConfig, EngineFactory, TimingModel, WorkerType};
15use serde::{Deserialize, Serialize};
16use serde_json::Value;
17
18use crate::replay::OfflineDisaggReplayConfig;
19use crate::replay::components::{
20    AdmissionQueue, NoReplayMetadata, ReplayEngineObservation, ReplayMode,
21};
22use crate::replay::core::EngineEventBatch;
23use crate::replay::core::round_robin::PoolRoundRobinPlacement;
24use crate::replay::disagg::DisaggRuntimeImpl;
25use crate::replay::error::runtime_error;
26use crate::replay::protocol::DirectRequest;
27use crate::replay::{
28    ReplayError, ReplayReport, ReplayResult, ReplaySpec, ReplayTopology, Replayer, WorkerStage,
29};
30
31fn default_dp_size() -> u32 {
32    1
33}
34
35fn default_tensor_parallel_size() -> u32 {
36    1
37}
38
39/// Serializable execution-time descriptor for the AISimulate engine.
40#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
41#[serde(default, deny_unknown_fields)]
42pub struct ReplayEngineConfig {
43    #[serde(default = "default_dp_size")]
44    pub dp_size: u32,
45    #[serde(default = "default_tensor_parallel_size")]
46    pub tensor_parallel_size: u32,
47    /// Whether `rank.num_gpu_blocks` came from the user rather than an
48    /// upstream capacity estimator.
49    #[serde(default, skip_serializing_if = "Option::is_none")]
50    pub num_gpu_blocks_is_explicit: Option<bool>,
51    pub rank: EngineConfig,
52    #[serde(default, skip_serializing_if = "Option::is_none")]
53    pub prefill: Option<ReplayRoleConfig>,
54    #[serde(default, skip_serializing_if = "Option::is_none")]
55    pub decode: Option<ReplayRoleConfig>,
56}
57
58impl Default for ReplayEngineConfig {
59    fn default() -> Self {
60        Self {
61            dp_size: 1,
62            tensor_parallel_size: 1,
63            num_gpu_blocks_is_explicit: None,
64            rank: EngineConfig::default(),
65            prefill: None,
66            decode: None,
67        }
68    }
69}
70
71/// Rank-group descriptor for one disaggregated role.
72#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
73#[serde(default, deny_unknown_fields)]
74pub struct ReplayRoleConfig {
75    #[serde(default = "default_dp_size")]
76    pub dp_size: u32,
77    #[serde(default = "default_tensor_parallel_size")]
78    pub tensor_parallel_size: u32,
79    /// Whether `rank.num_gpu_blocks` came from the user rather than an
80    /// upstream capacity estimator.
81    #[serde(default, skip_serializing_if = "Option::is_none")]
82    pub num_gpu_blocks_is_explicit: Option<bool>,
83    pub rank: EngineConfig,
84}
85
86impl Default for ReplayRoleConfig {
87    fn default() -> Self {
88        Self {
89            dp_size: 1,
90            tensor_parallel_size: 1,
91            num_gpu_blocks_is_explicit: None,
92            rank: EngineConfig::default(),
93        }
94    }
95}
96
97impl ReplayEngineConfig {
98    pub(crate) fn parse(value: &Value) -> ReplayResult<Self> {
99        if value.is_null() {
100            return Ok(Self::default());
101        }
102        serde_json::from_value(value.clone()).map_err(|error| {
103            ReplayError::InvalidSpec(format!("invalid native engine descriptor: {error}"))
104        })
105    }
106
107    pub(crate) fn role(&self, stage: WorkerStage) -> ReplayRoleConfig {
108        let mut role = match stage {
109            WorkerStage::Aggregated => ReplayRoleConfig {
110                dp_size: self.dp_size,
111                tensor_parallel_size: self.tensor_parallel_size,
112                num_gpu_blocks_is_explicit: self.num_gpu_blocks_is_explicit,
113                rank: self.rank.clone(),
114            },
115            WorkerStage::Prefill => self.prefill.clone().unwrap_or_else(|| ReplayRoleConfig {
116                dp_size: self.dp_size,
117                tensor_parallel_size: self.tensor_parallel_size,
118                num_gpu_blocks_is_explicit: self.num_gpu_blocks_is_explicit,
119                rank: self.rank.clone(),
120            }),
121            WorkerStage::Decode => self.decode.clone().unwrap_or_else(|| ReplayRoleConfig {
122                dp_size: self.dp_size,
123                tensor_parallel_size: self.tensor_parallel_size,
124                num_gpu_blocks_is_explicit: self.num_gpu_blocks_is_explicit,
125                rank: self.rank.clone(),
126            }),
127        };
128        role.rank.worker_type = match stage {
129            WorkerStage::Aggregated => WorkerType::Aggregated,
130            WorkerStage::Prefill => WorkerType::Prefill,
131            WorkerStage::Decode => WorkerType::Decode,
132        };
133        role
134    }
135
136    pub(crate) fn validate_topology(&self, topology: &ReplayTopology) -> ReplayResult<()> {
137        match topology {
138            ReplayTopology::Aggregated { .. } => {
139                if self.rank.native_host_offload.is_some() && self.dp_size != 1 {
140                    return Err(ReplayError::InvalidSpec(
141                        "native_host_offload supports only dp_size=1 in the initial implementation"
142                            .to_string(),
143                    ));
144                }
145            }
146            ReplayTopology::Disaggregated { .. } => {
147                for stage in [WorkerStage::Prefill, WorkerStage::Decode] {
148                    let role = self.role(stage);
149                    if role.rank.native_host_offload.is_some() {
150                        return Err(ReplayError::InvalidSpec(
151                            "native_host_offload supports only aggregated replay in the initial implementation"
152                                .to_string(),
153                        ));
154                    }
155                    if role.dp_size != 1 {
156                        // TODO(#12965): Carry logical-worker plus DP-rank identity through
157                        // disaggregated handoff before removing this fail-fast guard.
158                        let role_name = match stage {
159                            WorkerStage::Prefill => "prefill",
160                            WorkerStage::Decode => "decode",
161                            WorkerStage::Aggregated => unreachable!(),
162                        };
163                        return Err(ReplayError::InvalidSpec(format!(
164                            "disaggregated replay requires {role_name} dp_size=1; attention-DP handoff identity is not implemented"
165                        )));
166                    }
167                }
168            }
169        }
170        Ok(())
171    }
172}
173
174/// Reusable construction state for one worker role.
175#[doc(hidden)]
176#[derive(Clone)]
177pub struct ReplayRoleFactory {
178    factory: EngineFactory,
179    dp_size: NonZeroU32,
180    tensor_parallel_size: u32,
181    backend: Backend,
182    total_blocks: u64,
183}
184
185impl ReplayRoleFactory {
186    #[doc(hidden)]
187    pub fn build(&self, worker_id: usize) -> ReplayResult<Engine> {
188        let worker_id = u64::try_from(worker_id).map_err(|_| {
189            ReplayError::Engine(format!(
190                "worker id {worker_id} exceeds the native engine range"
191            ))
192        })?;
193        self.factory
194            .build(EngineIdentity::new(worker_id), self.dp_size)
195            .map_err(engine_error)
196    }
197
198    #[doc(hidden)]
199    pub fn dp_size(&self) -> u32 {
200        self.dp_size.get()
201    }
202
203    #[doc(hidden)]
204    pub fn gpus_per_worker(&self) -> ReplayResult<usize> {
205        usize::try_from(self.dp_size.get())
206            .ok()
207            .and_then(|dp| {
208                usize::try_from(self.tensor_parallel_size)
209                    .ok()
210                    .and_then(|tp| dp.checked_mul(tp))
211            })
212            .ok_or_else(|| ReplayError::InvalidSpec("engine GPU count overflows usize".into()))
213    }
214
215    #[doc(hidden)]
216    pub fn backend(&self) -> Backend {
217        self.backend
218    }
219
220    #[doc(hidden)]
221    pub fn total_blocks(&self) -> u64 {
222        self.total_blocks
223    }
224}
225
226/// Resolves built-in or Runner-provided timing once, then creates role factories.
227#[derive(Clone, Default)]
228pub struct ReplayEngineFactory {
229    timing: Option<Arc<dyn TimingModel>>,
230    prefill_timing: Option<Arc<dyn TimingModel>>,
231    decode_timing: Option<Arc<dyn TimingModel>>,
232}
233
234impl ReplayEngineFactory {
235    pub const fn new() -> Self {
236        Self {
237            timing: None,
238            prefill_timing: None,
239            decode_timing: None,
240        }
241    }
242
243    pub fn with_timing_model(timing: Arc<dyn TimingModel>) -> Self {
244        Self {
245            timing: Some(timing),
246            prefill_timing: None,
247            decode_timing: None,
248        }
249    }
250
251    pub fn with_optional_role_timing_models(
252        prefill: Option<Arc<dyn TimingModel>>,
253        decode: Option<Arc<dyn TimingModel>>,
254    ) -> Self {
255        Self {
256            timing: None,
257            prefill_timing: prefill,
258            decode_timing: decode,
259        }
260    }
261
262    #[doc(hidden)]
263    pub fn role_factory(
264        &self,
265        config: &ReplayEngineConfig,
266        stage: WorkerStage,
267        emit_kv_events: bool,
268    ) -> ReplayResult<ReplayRoleFactory> {
269        let mut role = config.role(stage);
270        role.rank.emit_kv_events = emit_kv_events;
271        let dp_size = NonZeroU32::new(role.dp_size).ok_or_else(|| {
272            ReplayError::InvalidSpec("native engine dp_size must be positive".into())
273        })?;
274        if role.tensor_parallel_size == 0 {
275            return Err(ReplayError::InvalidSpec(
276                "native tensor_parallel_size must be positive".into(),
277            ));
278        }
279        let timing = match stage {
280            WorkerStage::Aggregated => self.timing.as_ref(),
281            WorkerStage::Prefill => self.prefill_timing.as_ref().or(self.timing.as_ref()),
282            WorkerStage::Decode => self.decode_timing.as_ref().or(self.timing.as_ref()),
283        };
284        let backend = role.rank.backend;
285        let total_blocks = u64::try_from(role.rank.num_gpu_blocks).map_err(|_| {
286            ReplayError::InvalidSpec("engine KV block count exceeds the metrics range".into())
287        })?;
288        let factory = match timing {
289            Some(timing) => EngineFactory::with_timing_model(role.rank, Arc::clone(timing)),
290            None => EngineFactory::new(role.rank),
291        }
292        .map_err(engine_error)?;
293        Ok(ReplayRoleFactory {
294            factory,
295            dp_size,
296            tensor_parallel_size: role.tensor_parallel_size,
297            backend,
298            total_blocks,
299        })
300    }
301}
302
303pub fn run_engine_replay(spec: ReplaySpec) -> ReplayResult<ReplayReport> {
304    Replayer::new(spec, ReplayEngineFactory::new())?.run()
305}
306
307pub fn run_engine_replay_with_timing(
308    spec: ReplaySpec,
309    timing: Arc<dyn TimingModel>,
310) -> ReplayResult<ReplayReport> {
311    Replayer::new(spec, ReplayEngineFactory::with_timing_model(timing))?.run()
312}
313
314pub fn run_engine_replay_with_optional_role_timing(
315    spec: ReplaySpec,
316    prefill: Option<Arc<dyn TimingModel>>,
317    decode: Option<Arc<dyn TimingModel>>,
318) -> ReplayResult<ReplayReport> {
319    Replayer::new(
320        spec,
321        ReplayEngineFactory::with_optional_role_timing_models(prefill, decode),
322    )?
323    .run()
324}
325
326#[derive(Debug, Default)]
327struct KvEventBatch(Vec<crate::engine::KvEvent>);
328
329impl EngineEventBatch for KvEventBatch {
330    fn is_empty(&self) -> bool {
331        self.0.is_empty()
332    }
333
334    fn append(&mut self, mut other: Self) {
335        self.0.append(&mut other.0);
336    }
337}
338
339#[derive(Debug, Default)]
340struct KvEventObservation;
341
342impl ReplayEngineObservation for KvEventObservation {
343    type Batch = KvEventBatch;
344
345    const CAPTURE_ENGINE_KV_EVENTS: bool = true;
346
347    fn observe_engine_events(
348        _stage: WorkerStage,
349        _worker_id: usize,
350        _dp_rank: u32,
351        events: Vec<crate::engine::KvEvent>,
352    ) -> Self::Batch {
353        KvEventBatch(events)
354    }
355
356    fn stored_hashes(batch: &Self::Batch) -> Vec<u64> {
357        batch
358            .0
359            .iter()
360            .flat_map(|event| match &event.data {
361                crate::engine::KvEventData::Stored(stored) => stored.blocks.as_slice(),
362                crate::engine::KvEventData::Removed { .. } => &[],
363            })
364            .map(|block| block.tokens_hash)
365            .collect()
366    }
367}
368
369/// Run the engine-neutral half of Dynamo's live/offline handoff conformance
370/// fixture without importing or recompiling Replay implementation sources.
371#[doc(hidden)]
372pub fn run_engine_handoff_conformance(
373    config: ReplayEngineConfig,
374    factory: ReplayEngineFactory,
375    request: DirectRequest,
376) -> ReplayResult<crate::replay::NormalizedHandoffConformance> {
377    let prefill_factory = factory.role_factory(&config, WorkerStage::Prefill, true)?;
378    let decode_factory = factory.role_factory(&config, WorkerStage::Decode, true)?;
379    let backend = prefill_factory.backend();
380    if backend != decode_factory.backend() {
381        return Err(ReplayError::InvalidSpec(
382            "handoff conformance requires matching prefill/decode backends".into(),
383        ));
384    }
385    let runtime_config = OfflineDisaggReplayConfig {
386        prefill_factory,
387        decode_factory,
388        prefill_startup_time_ms: None,
389        decode_startup_time_ms: None,
390        num_prefill_workers: 1,
391        num_decode_workers: 1,
392        handoff_latency_ms: 0.0,
393    };
394    DisaggRuntimeImpl::<
395        PoolRoundRobinPlacement<KvEventBatch>,
396        KvEventObservation,
397        NoReplayMetadata,
398    >::new_composed(
399        &runtime_config,
400        AdmissionQueue::new_requests(
401            std::collections::VecDeque::from([request]),
402            ReplayMode::Trace,
403        ),
404        true,
405        |_, prefill_topology, _, decode_topology| {
406            Ok((
407                PoolRoundRobinPlacement::new(prefill_topology),
408                PoolRoundRobinPlacement::new(decode_topology),
409            ))
410        },
411    )
412    .map_err(runtime_error)?
413    .run_handoff_conformance(backend)
414    .map_err(runtime_error)
415}
416
417fn engine_error(error: impl std::fmt::Display) -> ReplayError {
418    ReplayError::Engine(error.to_string())
419}