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