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