1use 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#[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#[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 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#[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#[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#[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}