1use std::sync::Arc;
7
8use anyhow::{Result, ensure};
9use serde::{Deserialize, Deserializer, Serialize};
10
11use crate::engine::common::speculative::normalize_conditional_accept_rates;
12use crate::engine::handoff::TransferTimingMode;
13use crate::engine::timing::{TimingModel, TimingModelConfig, built_in_timing_model};
14
15const DEFAULT_MAX_PREFILL_TOKENS: usize = 16_384;
16const DEFAULT_CHUNKED_PREFILL_SIZE: usize = 8_192;
17const DEFAULT_CLIP_MAX_NEW_TOKENS: usize = 4_096;
18const DEFAULT_SCHEDULE_CONSERVATIVENESS: f64 = 1.0;
19
20fn default_num_gpu_blocks() -> usize {
21 16_384
22}
23
24fn default_block_size() -> usize {
25 64
26}
27
28fn default_max_num_seqs() -> usize {
29 256
30}
31
32fn default_max_num_batched_tokens() -> usize {
33 8_192
34}
35
36fn default_true() -> bool {
37 true
38}
39
40fn default_one() -> f64 {
41 1.0
42}
43
44fn default_aic_mtp_seed() -> u64 {
45 42
46}
47
48fn default_max_prefill_tokens() -> usize {
49 DEFAULT_MAX_PREFILL_TOKENS
50}
51
52fn default_chunked_prefill_size() -> usize {
53 DEFAULT_CHUNKED_PREFILL_SIZE
54}
55
56fn default_clip_max_new_tokens() -> usize {
57 DEFAULT_CLIP_MAX_NEW_TOKENS
58}
59
60fn default_schedule_conservativeness() -> f64 {
61 DEFAULT_SCHEDULE_CONSERVATIVENESS
62}
63
64#[derive(Debug, Clone, Copy, Default, PartialEq, Eq, Serialize, Deserialize)]
66#[serde(rename_all = "snake_case")]
67pub enum Backend {
68 #[default]
70 Vllm,
71 Sglang,
73 Trtllm,
75}
76
77impl Backend {
78 pub const fn default_block_size(self) -> usize {
80 match self {
81 Self::Vllm => 64,
82 Self::Sglang => 1,
83 Self::Trtllm => 32,
84 }
85 }
86}
87
88#[derive(Debug, Clone, Copy, Default, PartialEq, Eq, Serialize, Deserialize)]
90#[serde(rename_all = "snake_case")]
91pub enum WorkerType {
92 #[default]
94 Aggregated,
95 Prefill,
97 Decode,
99}
100
101#[derive(Debug, Clone, Copy, Default, PartialEq, Eq, Serialize, Deserialize)]
103#[serde(rename_all = "snake_case")]
104pub enum PreemptionMode {
105 #[default]
107 Lifo,
108 Fifo,
110}
111
112#[derive(Debug, Clone, Copy, Default, PartialEq, Eq, Serialize, Deserialize)]
114#[serde(rename_all = "snake_case")]
115pub enum SglangSchedulePolicy {
116 #[default]
118 Fifo,
119 Lpm,
121}
122
123#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
125#[serde(default, deny_unknown_fields)]
126pub struct SglangConfig {
127 pub schedule_policy: SglangSchedulePolicy,
129 #[serde(default = "default_max_prefill_tokens")]
131 pub max_prefill_tokens: usize,
132 #[serde(default = "default_chunked_prefill_size")]
134 pub chunked_prefill_size: usize,
135 #[serde(default = "default_clip_max_new_tokens")]
137 pub clip_max_new_tokens: usize,
138 #[serde(default = "default_schedule_conservativeness")]
140 pub schedule_conservativeness: f64,
141}
142
143impl Default for SglangConfig {
144 fn default() -> Self {
145 Self {
146 schedule_policy: SglangSchedulePolicy::Fifo,
147 max_prefill_tokens: default_max_prefill_tokens(),
148 chunked_prefill_size: default_chunked_prefill_size(),
149 clip_max_new_tokens: default_clip_max_new_tokens(),
150 schedule_conservativeness: default_schedule_conservativeness(),
151 }
152 }
153}
154
155impl SglangConfig {
156 pub(crate) fn validate(&self) -> Result<()> {
157 ensure!(
158 self.max_prefill_tokens > 0,
159 "sglang.max_prefill_tokens must be positive"
160 );
161 ensure!(
162 self.chunked_prefill_size > 0,
163 "sglang.chunked_prefill_size must be positive"
164 );
165 ensure!(
166 self.schedule_conservativeness.is_finite() && self.schedule_conservativeness >= 0.0,
167 "sglang.schedule_conservativeness must be finite and non-negative"
168 );
169 Ok(())
170 }
171}
172
173#[derive(Debug, Clone, Copy, Default, PartialEq, Eq, Serialize, Deserialize)]
179#[serde(rename_all = "snake_case")]
180pub enum TrtllmCapacityPolicy {
181 #[default]
183 GuaranteedNoEvict,
184}
185
186#[derive(Debug, Clone, Default, PartialEq, Eq, Serialize, Deserialize)]
188#[serde(default, deny_unknown_fields)]
189pub struct TrtllmConfig {
190 pub capacity_scheduler_policy: TrtllmCapacityPolicy,
192}
193
194#[derive(Debug, Clone, PartialEq, Serialize)]
206pub struct EngineConfig {
207 pub backend: Backend,
212 #[serde(default = "default_num_gpu_blocks")]
214 pub num_gpu_blocks: usize,
215 #[serde(default = "default_block_size")]
217 pub block_size: usize,
218 pub max_model_len: Option<usize>,
220 #[serde(default = "default_max_num_seqs")]
222 pub max_num_seqs: usize,
223 #[serde(default = "default_max_num_batched_tokens")]
225 pub max_num_batched_tokens: usize,
226 #[serde(default = "default_true")]
228 pub enable_prefix_caching: bool,
229 #[serde(default = "default_true")]
231 pub enable_chunked_prefill: bool,
232 #[serde(default = "default_one")]
234 pub speedup_ratio: f64,
235 #[serde(default = "default_one")]
237 pub decode_speedup_ratio: f64,
238 pub aic_nextn: Option<usize>,
241 pub aic_nextn_accept_rates: Option<String>,
246 #[serde(default = "default_aic_mtp_seed")]
248 pub aic_mtp_seed: u64,
249 pub worker_type: WorkerType,
251 pub preemption_mode: PreemptionMode,
253 pub emit_kv_events: bool,
255 pub emit_kv_token_ids: bool,
257 pub kv_bytes_per_token: Option<usize>,
259 pub kv_transfer_bandwidth: Option<f64>,
261 pub kv_transfer_timing_mode: TransferTimingMode,
263 pub timing_model: TimingModelConfig,
265 pub sglang: SglangConfig,
267 pub trtllm: TrtllmConfig,
269}
270
271#[derive(Deserialize)]
272#[serde(deny_unknown_fields)]
273struct EngineConfigWire {
274 #[serde(default)]
275 backend: Backend,
276 #[serde(default = "default_num_gpu_blocks")]
277 num_gpu_blocks: usize,
278 #[serde(default)]
279 block_size: Option<usize>,
280 #[serde(default)]
281 max_model_len: Option<usize>,
282 #[serde(default = "default_max_num_seqs")]
283 max_num_seqs: usize,
284 #[serde(default = "default_max_num_batched_tokens")]
285 max_num_batched_tokens: usize,
286 #[serde(default = "default_true")]
287 enable_prefix_caching: bool,
288 #[serde(default = "default_true")]
289 enable_chunked_prefill: bool,
290 #[serde(default = "default_one")]
291 speedup_ratio: f64,
292 #[serde(default = "default_one")]
293 decode_speedup_ratio: f64,
294 #[serde(default)]
295 aic_nextn: Option<usize>,
296 #[serde(default)]
297 aic_nextn_accept_rates: Option<String>,
298 #[serde(default = "default_aic_mtp_seed")]
299 aic_mtp_seed: u64,
300 #[serde(default)]
301 worker_type: WorkerType,
302 #[serde(default)]
303 preemption_mode: PreemptionMode,
304 #[serde(default)]
305 emit_kv_events: bool,
306 #[serde(default)]
307 emit_kv_token_ids: bool,
308 #[serde(default)]
309 kv_bytes_per_token: Option<usize>,
310 #[serde(default)]
311 kv_transfer_bandwidth: Option<f64>,
312 #[serde(default)]
313 kv_transfer_timing_mode: TransferTimingMode,
314 #[serde(default)]
315 timing_model: TimingModelConfig,
316 #[serde(default)]
317 sglang: SglangConfig,
318 #[serde(default)]
319 trtllm: TrtllmConfig,
320}
321
322impl<'de> Deserialize<'de> for EngineConfig {
323 fn deserialize<D>(deserializer: D) -> std::result::Result<Self, D::Error>
324 where
325 D: Deserializer<'de>,
326 {
327 let wire = EngineConfigWire::deserialize(deserializer)?;
328 Ok(Self {
329 backend: wire.backend,
330 num_gpu_blocks: wire.num_gpu_blocks,
331 block_size: wire
332 .block_size
333 .unwrap_or_else(|| wire.backend.default_block_size()),
334 max_model_len: wire.max_model_len,
335 max_num_seqs: wire.max_num_seqs,
336 max_num_batched_tokens: wire.max_num_batched_tokens,
337 enable_prefix_caching: wire.enable_prefix_caching,
338 enable_chunked_prefill: wire.enable_chunked_prefill,
339 speedup_ratio: wire.speedup_ratio,
340 decode_speedup_ratio: wire.decode_speedup_ratio,
341 aic_nextn: wire.aic_nextn,
342 aic_nextn_accept_rates: wire.aic_nextn_accept_rates,
343 aic_mtp_seed: wire.aic_mtp_seed,
344 worker_type: wire.worker_type,
345 preemption_mode: wire.preemption_mode,
346 emit_kv_events: wire.emit_kv_events,
347 emit_kv_token_ids: wire.emit_kv_token_ids,
348 kv_bytes_per_token: wire.kv_bytes_per_token,
349 kv_transfer_bandwidth: wire.kv_transfer_bandwidth,
350 kv_transfer_timing_mode: wire.kv_transfer_timing_mode,
351 timing_model: wire.timing_model,
352 sglang: wire.sglang,
353 trtllm: wire.trtllm,
354 })
355 }
356}
357
358impl Default for EngineConfig {
359 fn default() -> Self {
360 Self {
361 backend: Backend::Vllm,
362 num_gpu_blocks: default_num_gpu_blocks(),
363 block_size: default_block_size(),
364 max_model_len: None,
365 max_num_seqs: default_max_num_seqs(),
366 max_num_batched_tokens: default_max_num_batched_tokens(),
367 enable_prefix_caching: true,
368 enable_chunked_prefill: true,
369 speedup_ratio: 1.0,
370 decode_speedup_ratio: 1.0,
371 aic_nextn: None,
372 aic_nextn_accept_rates: None,
373 aic_mtp_seed: default_aic_mtp_seed(),
374 worker_type: WorkerType::Aggregated,
375 preemption_mode: PreemptionMode::Lifo,
376 emit_kv_events: false,
377 emit_kv_token_ids: false,
378 kv_bytes_per_token: None,
379 kv_transfer_bandwidth: None,
380 kv_transfer_timing_mode: TransferTimingMode::FullPrompt,
381 timing_model: TimingModelConfig::Polynomial,
382 sglang: SglangConfig::default(),
383 trtllm: TrtllmConfig::default(),
384 }
385 }
386}
387
388impl EngineConfig {
389 pub fn for_backend(backend: Backend) -> Self {
394 Self {
395 backend,
396 block_size: backend.default_block_size(),
397 ..Self::default()
398 }
399 }
400
401 pub(crate) fn validate(&self) -> Result<()> {
402 ensure!(self.num_gpu_blocks > 0, "num_gpu_blocks must be positive");
403 ensure!(self.block_size > 0, "block_size must be positive");
404 if matches!(self.backend, Backend::Vllm | Backend::Trtllm) {
405 ensure!(
406 self.block_size >= 2,
407 "vLLM/TRT-LLM block_size must be at least two"
408 );
409 }
410 ensure!(self.max_num_seqs > 0, "max_num_seqs must be positive");
411 ensure!(
412 self.max_num_batched_tokens > 0,
413 "max_num_batched_tokens must be positive"
414 );
415 ensure!(
416 self.max_model_len.is_none_or(|limit| limit > 0),
417 "max_model_len must be positive"
418 );
419 ensure!(
420 self.backend == Backend::Vllm || self.max_model_len.is_none(),
421 "max_model_len is supported only for backend=vllm"
422 );
423 ensure!(
424 self.speedup_ratio.is_finite() && self.speedup_ratio >= 0.0,
425 "speedup_ratio must be finite and non-negative"
426 );
427 ensure!(
428 self.decode_speedup_ratio.is_finite() && self.decode_speedup_ratio >= 0.0,
429 "decode_speedup_ratio must be finite and non-negative"
430 );
431 if let Some(nextn) = self.aic_nextn {
432 normalize_conditional_accept_rates(nextn, self.aic_nextn_accept_rates.as_deref())?;
433 ensure!(
434 self.decode_speedup_ratio == 1.0,
435 "aic_nextn requires decode_speedup_ratio=1.0 because MTP output acceleration is modeled by burst sampling"
436 );
437 } else {
438 ensure!(
439 self.aic_nextn_accept_rates.is_none(),
440 "aic_nextn_accept_rates requires aic_nextn"
441 );
442 }
443 if self.backend == Backend::Sglang {
444 ensure!(
445 !self.emit_kv_token_ids,
446 "emit_kv_token_ids=true is not supported for backend=sglang"
447 );
448 ensure!(
449 self.enable_chunked_prefill,
450 "enable_chunked_prefill=false is not supported for backend=sglang"
451 );
452 self.sglang.validate()?;
453 }
454 ensure!(
455 !self.emit_kv_token_ids || self.emit_kv_events,
456 "emit_kv_token_ids requires emit_kv_events"
457 );
458 ensure!(
459 self.kv_bytes_per_token.is_none_or(|bytes| bytes > 0),
460 "kv_bytes_per_token must be positive"
461 );
462 ensure!(
463 self.kv_transfer_bandwidth
464 .is_none_or(|bandwidth| bandwidth.is_finite() && bandwidth >= 0.0),
465 "kv_transfer_bandwidth must be finite and non-negative"
466 );
467 match &self.timing_model {
468 TimingModelConfig::Polynomial => {}
469 TimingModelConfig::Fixed {
470 prefill_ms,
471 decode_ms,
472 } => {
473 ensure!(
474 prefill_ms.is_finite() && *prefill_ms >= 0.0,
475 "fixed prefill latency must be finite and non-negative"
476 );
477 ensure!(
478 decode_ms.is_finite() && *decode_ms >= 0.0,
479 "fixed decode latency must be finite and non-negative"
480 );
481 }
482 TimingModelConfig::External { provider, .. } => {
483 ensure!(
484 !provider.trim().is_empty(),
485 "timing provider cannot be empty"
486 );
487 }
488 }
489 Ok(())
490 }
491
492 pub(crate) fn built_in_timing_model(&self) -> Result<Arc<dyn TimingModel>> {
493 built_in_timing_model(&self.timing_model)
494 }
495}
496
497#[cfg(test)]
498mod tests {
499 use super::*;
500
501 #[test]
502 fn deserialization_uses_backend_native_block_size() {
503 for (backend, expected) in [("vllm", 64), ("sglang", 1), ("trtllm", 32)] {
504 let config: EngineConfig =
505 serde_json::from_value(serde_json::json!({ "backend": backend })).unwrap();
506 assert_eq!(config.block_size, expected, "backend={backend}");
507 }
508 }
509
510 #[test]
511 fn for_backend_uses_backend_native_block_size() {
512 for backend in [Backend::Vllm, Backend::Sglang, Backend::Trtllm] {
513 let config = EngineConfig::for_backend(backend);
514 assert_eq!(config.backend, backend);
515 assert_eq!(config.block_size, backend.default_block_size());
516 }
517 }
518
519 #[test]
520 fn deserialization_preserves_an_explicit_block_size() {
521 let config: EngineConfig = serde_json::from_value(serde_json::json!({
522 "backend": "sglang",
523 "block_size": 17
524 }))
525 .unwrap();
526 assert_eq!(config.block_size, 17);
527 }
528
529 #[test]
530 fn deserialization_still_rejects_unknown_fields() {
531 let error = serde_json::from_value::<EngineConfig>(serde_json::json!({
532 "backend": "vllm",
533 "unknown": true
534 }))
535 .unwrap_err();
536 assert!(error.to_string().contains("unknown field"));
537 }
538
539 #[test]
540 fn serialization_round_trip_preserves_runtime_neutral_controls() {
541 let config = EngineConfig {
542 backend: Backend::Sglang,
543 block_size: 8,
544 num_gpu_blocks: 123,
545 max_num_seqs: 7,
546 max_num_batched_tokens: 456,
547 worker_type: WorkerType::Decode,
548 preemption_mode: PreemptionMode::Fifo,
549 emit_kv_events: true,
550 emit_kv_token_ids: true,
551 timing_model: TimingModelConfig::Fixed {
552 prefill_ms: 2.5,
553 decode_ms: 0.75,
554 },
555 ..EngineConfig::for_backend(Backend::Sglang)
556 };
557 let encoded = serde_json::to_value(&config).unwrap();
558 let decoded: EngineConfig = serde_json::from_value(encoded).unwrap();
559 assert_eq!(decoded, config);
560 }
561
562 #[test]
563 fn validation_rejects_zero_or_backend_invalid_capacity_fields() {
564 let config = EngineConfig {
565 num_gpu_blocks: 0,
566 ..EngineConfig::default()
567 };
568 assert!(
569 config
570 .validate()
571 .unwrap_err()
572 .to_string()
573 .contains("num_gpu_blocks")
574 );
575
576 let config = EngineConfig {
577 block_size: 1,
578 ..EngineConfig::default()
579 };
580 assert!(
581 config
582 .validate()
583 .unwrap_err()
584 .to_string()
585 .contains("at least two")
586 );
587
588 let config = EngineConfig {
589 max_model_len: Some(0),
590 ..EngineConfig::default()
591 };
592 assert!(
593 config
594 .validate()
595 .unwrap_err()
596 .to_string()
597 .contains("max_model_len")
598 );
599 }
600
601 #[test]
602 fn validation_accepts_sglang_page_size_one_and_rejects_invalid_controls() {
603 let mut config = EngineConfig::for_backend(Backend::Sglang);
604 config.validate().unwrap();
605
606 config.sglang.chunked_prefill_size = 0;
607 assert!(
608 config
609 .validate()
610 .unwrap_err()
611 .to_string()
612 .contains("chunked_prefill_size")
613 );
614
615 let mut config = EngineConfig::for_backend(Backend::Sglang);
616 config.sglang.schedule_conservativeness = f64::NAN;
617 assert!(
618 config
619 .validate()
620 .unwrap_err()
621 .to_string()
622 .contains("schedule_conservativeness")
623 );
624 }
625
626 #[test]
627 fn sglang_supports_disabled_prefix_caching() {
628 let config = EngineConfig {
629 enable_prefix_caching: false,
630 ..EngineConfig::for_backend(Backend::Sglang)
631 };
632 config.validate().unwrap();
633 crate::engine::EngineFactory::new(config).unwrap();
634 }
635
636 #[test]
637 fn sglang_rejects_remaining_unsupported_controls_at_validation_and_factory_boundaries() {
638 let cases = [
639 ("emit_kv_token_ids", true, true, true),
640 ("enable_chunked_prefill", false, true, false),
641 ];
642
643 for (field, emit_kv_token_ids, enable_prefix_caching, enable_chunked_prefill) in cases {
644 let config = EngineConfig {
645 emit_kv_events: emit_kv_token_ids,
646 emit_kv_token_ids,
647 enable_prefix_caching,
648 enable_chunked_prefill,
649 ..EngineConfig::for_backend(Backend::Sglang)
650 };
651 assert!(config.validate().unwrap_err().to_string().contains(field));
652 let error = match crate::engine::EngineFactory::new(config) {
653 Ok(_) => panic!("expected EngineFactory to reject {field}"),
654 Err(error) => error,
655 };
656 assert!(error.to_string().contains(field));
657 }
658 }
659
660 #[test]
661 fn max_model_len_is_vllm_only() {
662 for backend in [Backend::Sglang, Backend::Trtllm] {
663 let mut config = EngineConfig::for_backend(backend);
664 config.max_model_len = Some(128);
665 assert!(
666 config
667 .validate()
668 .unwrap_err()
669 .to_string()
670 .contains("backend=vllm")
671 );
672 }
673 }
674
675 #[test]
676 fn mtp_configuration_validates_rates_and_decode_scaling() {
677 let mut config = EngineConfig {
678 aic_nextn: Some(2),
679 aic_nextn_accept_rates: Some("0.8,0.5".to_string()),
680 ..EngineConfig::default()
681 };
682 config.validate().unwrap();
683
684 config.aic_nextn_accept_rates = Some("1.2".to_string());
685 assert!(config.validate().is_err());
686
687 config.aic_nextn_accept_rates = Some("0.8,0.5".to_string());
688 config.decode_speedup_ratio = 2.0;
689 assert!(
690 config
691 .validate()
692 .unwrap_err()
693 .to_string()
694 .contains("decode_speedup_ratio=1.0")
695 );
696 }
697
698 #[test]
699 fn mtp_rates_require_mtp_to_be_enabled() {
700 let config = EngineConfig {
701 aic_nextn_accept_rates: Some("0.5".to_string()),
702 ..EngineConfig::default()
703 };
704 assert!(
705 config
706 .validate()
707 .unwrap_err()
708 .to_string()
709 .contains("requires aic_nextn")
710 );
711 }
712
713 #[test]
714 fn kv_token_ids_require_kv_event_emission() {
715 let config = EngineConfig {
716 emit_kv_token_ids: true,
717 emit_kv_events: false,
718 ..EngineConfig::default()
719 };
720 assert!(
721 config
722 .validate()
723 .unwrap_err()
724 .to_string()
725 .contains("emit_kv_token_ids")
726 );
727 }
728
729 #[test]
730 fn timing_provider_descriptors_are_validated_without_loading_them() {
731 let config = EngineConfig {
732 timing_model: TimingModelConfig::External {
733 provider: " ".to_string(),
734 config: serde_json::Value::Null,
735 },
736 ..EngineConfig::default()
737 };
738 assert!(
739 config
740 .validate()
741 .unwrap_err()
742 .to_string()
743 .contains("provider cannot be empty")
744 );
745
746 let config = EngineConfig {
747 timing_model: TimingModelConfig::Fixed {
748 prefill_ms: f64::NAN,
749 decode_ms: 1.0,
750 },
751 ..EngineConfig::default()
752 };
753 assert!(config.validate().is_err());
754 }
755}