1use async_trait::async_trait;
7use ferrum_types::{
8 BatchId, InferenceRequest, InferenceResponse, Priority, RequestId, RequestState, Result,
9 SchedulerConfig as TypesSchedulerConfig, SchedulerStats,
10};
11use serde::{Deserialize, Serialize};
12use std::{collections::HashMap, time::Duration};
13
14#[async_trait]
16pub trait Scheduler: Send + Sync {
17 async fn submit(&self, request: InferenceRequest) -> Result<RequestId>;
19
20 async fn next_batch(&self, hint: BatchHint) -> Option<BatchPlan>;
22
23 async fn complete(&self, request_id: RequestId, response: &InferenceResponse) -> Result<()>;
25
26 async fn cancel(&self, request_id: RequestId) -> Result<bool>;
28
29 async fn update_priority(&self, request_id: RequestId, priority: Priority) -> Result<()>;
31
32 fn metrics(&self) -> SchedulerMetrics;
34
35 fn config(&self) -> &TypesSchedulerConfig;
37
38 fn request_state(&self, request_id: &RequestId) -> Option<RequestState> {
40 let _ = request_id;
41 None
42 }
43
44 async fn preempt(&self, _request_id: RequestId) -> Result<PreemptionResult> {
46 Err(ferrum_types::FerrumError::unsupported(
48 "Preemption not supported",
49 ))
50 }
51
52 async fn resume(&self, _request_id: RequestId) -> Result<()> {
54 Err(ferrum_types::FerrumError::unsupported(
56 "Resumption not supported",
57 ))
58 }
59}
60
61#[derive(Debug, Clone)]
63pub struct BatchHint {
64 pub max_batch_size: usize,
66 pub max_tokens: usize,
68 pub target_latency_ms: Option<u64>,
70 pub available_memory: Option<u64>,
72 pub resource_constraints: ResourceConstraints,
74}
75
76impl BatchHint {
77 pub fn simple(max_batch_size: usize) -> Self {
79 Self {
80 max_batch_size,
81 max_tokens: max_batch_size * 2048, target_latency_ms: None,
83 available_memory: None,
84 resource_constraints: ResourceConstraints::default(),
85 }
86 }
87}
88
89#[derive(Debug, Clone, Serialize, Deserialize, Default)]
91pub struct ResourceConstraints {
92 pub max_gpu_memory: Option<u64>,
94 pub max_cpu_memory: Option<u64>,
96 pub max_recurrent_state_bytes: Option<u64>,
98 pub max_recurrent_state_slots: Option<usize>,
100 pub max_compute_units: Option<usize>,
102 pub required_devices: Vec<ferrum_types::Device>,
104}
105
106#[derive(Debug, Clone)]
108pub struct BatchPlan {
109 pub batch_id: BatchId,
111 pub requests: Vec<ScheduledRequest>,
113 pub max_sequence_length: usize,
115 pub estimated_time_ms: Option<u64>,
117 pub resource_requirements: BatchResourceRequirements,
119 pub created_at: chrono::DateTime<chrono::Utc>,
121}
122
123impl BatchPlan {
124 pub fn total_tokens(&self) -> usize {
126 self.requests
127 .iter()
128 .map(|req| {
129 req.tokens_to_process
130 .unwrap_or(req.request.sampling_params.max_tokens)
131 })
132 .sum()
133 }
134
135 pub fn size(&self) -> usize {
137 self.requests.len()
138 }
139
140 pub fn is_empty(&self) -> bool {
142 self.requests.is_empty()
143 }
144
145 pub fn max_priority(&self) -> Priority {
147 self.requests
148 .iter()
149 .map(|req| req.request.priority)
150 .max()
151 .unwrap_or(Priority::Low)
152 }
153}
154
155#[derive(Debug, Clone)]
157pub struct ScheduledRequest {
158 pub request: InferenceRequest,
160 pub state: RequestState,
162 pub queue_position: Option<usize>,
164 pub estimated_wait_time: Option<Duration>,
166 pub tokens_processed: usize,
168 pub tokens_to_process: Option<usize>,
173 pub allocated_resources: AllocatedResources,
175 pub submitted_at: chrono::DateTime<chrono::Utc>,
177 pub started_at: Option<chrono::DateTime<chrono::Utc>>,
179}
180
181impl ScheduledRequest {
182 pub fn new(request: InferenceRequest) -> Self {
184 Self {
185 request,
186 state: RequestState::Waiting,
187 queue_position: None,
188 estimated_wait_time: None,
189 tokens_processed: 0,
190 tokens_to_process: None,
191 allocated_resources: AllocatedResources::default(),
192 submitted_at: chrono::Utc::now(),
193 started_at: None,
194 }
195 }
196
197 pub fn age(&self) -> Duration {
199 (chrono::Utc::now() - self.submitted_at)
200 .to_std()
201 .unwrap_or_default()
202 }
203
204 pub fn processing_time(&self) -> Option<Duration> {
206 self.started_at
207 .map(|start| (chrono::Utc::now() - start).to_std().unwrap_or_default())
208 }
209}
210
211#[derive(Debug, Clone, Default)]
213pub struct AllocatedResources {
214 pub kv_cache_blocks: Vec<ferrum_types::BlockId>,
216 pub gpu_memory: u64,
218 pub cpu_memory: u64,
220 pub recurrent_state_bytes: u64,
222 pub recurrent_state_slots: usize,
224 pub compute_units: usize,
226}
227
228#[derive(Debug, Clone, Default)]
230pub struct BatchResourceRequirements {
231 pub gpu_memory: u64,
233 pub cpu_memory: u64,
235 pub kv_cache_blocks: usize,
237 pub recurrent_state_bytes: u64,
239 pub recurrent_state_slots: usize,
241 pub compute_units: usize,
243}
244
245#[derive(Debug, Clone)]
247pub struct PreemptionResult {
248 pub success: bool,
250 pub saved_state: Option<PreemptionState>,
252 pub freed_resources: AllocatedResources,
254}
255
256#[derive(Debug, Clone)]
258pub struct PreemptionState {
259 pub kv_cache_checkpoint: Vec<u8>,
261 pub tokens_processed: usize,
263 pub generation_state: HashMap<String, serde_json::Value>,
265}
266
267pub type SchedulerConfig = TypesSchedulerConfig;
269
270#[derive(Debug, Clone, Copy, Serialize, Deserialize)]
272pub enum SchedulingPolicy {
273 FCFS,
275 Priority,
277 FairShare,
279 SJF,
281 ResourceAware,
283 SlaAware,
285}
286
287#[derive(Debug, Clone, Serialize, Deserialize)]
289pub struct FairShareConfig {
290 pub client_shares: HashMap<String, f32>,
292 pub default_share: f32,
294 pub enforcement_strictness: f32,
296}
297
298#[derive(Debug, Clone, Serialize, Deserialize)]
300pub struct SlaConfig {
301 pub enabled: bool,
303 pub default_sla: SlaRequirements,
305 pub client_slas: HashMap<String, SlaRequirements>,
307}
308
309#[derive(Debug, Clone, Serialize, Deserialize)]
311pub struct SlaRequirements {
312 pub max_latency_p95_ms: u64,
314 pub max_latency_p99_ms: u64,
316 pub min_throughput_rps: f32,
318 pub availability_percent: f32,
320}
321
322#[derive(Debug, Clone, Serialize, Deserialize, Default)]
324pub struct ResourceLimits {
325 pub max_gpu_memory: Option<u64>,
327 pub max_cpu_memory: Option<u64>,
329 pub max_kv_cache_blocks: Option<usize>,
331 pub max_recurrent_state_bytes: Option<u64>,
333 pub max_recurrent_state_slots: Option<usize>,
335 pub per_client_limits: HashMap<String, ClientResourceLimits>,
337}
338
339#[derive(Debug, Clone, Serialize, Deserialize)]
341pub struct ClientResourceLimits {
342 pub max_concurrent_requests: usize,
344 pub max_gpu_memory: Option<u64>,
346 pub max_recurrent_state_bytes: Option<u64>,
348 pub max_requests_per_minute: Option<u32>,
350}
351
352pub type SchedulerMetrics = SchedulerStats;
353
354#[async_trait]
356pub trait AdvancedScheduler: Scheduler {
357 async fn enable_resource_awareness(&mut self, config: ResourceAwarenessConfig) -> Result<()>;
359
360 async fn set_admission_policy(&mut self, policy: Box<dyn AdmissionPolicy>) -> Result<()>;
362
363 async fn configure_dynamic_batching(&mut self, config: DynamicBatchingConfig) -> Result<()>;
365
366 fn queue_analysis(&self) -> QueueAnalysis;
368
369 async fn simulate_load(
371 &self,
372 workload: &SimulatedWorkload,
373 ) -> Result<SchedulingSimulationResult>;
374}
375
376#[derive(Debug, Clone, Serialize, Deserialize)]
378pub struct ResourceAwarenessConfig {
379 pub enable_memory_awareness: bool,
381 pub enable_compute_awareness: bool,
383 pub prediction_horizon_ms: u64,
385 pub safety_margin: f32,
387}
388
389pub trait AdmissionPolicy: Send + Sync {
391 fn should_admit(
393 &self,
394 request: &InferenceRequest,
395 current_metrics: &SchedulerMetrics,
396 ) -> AdmissionDecision;
397
398 fn name(&self) -> &str;
400}
401
402#[derive(Debug, Clone)]
404pub enum AdmissionDecision {
405 Accept,
407 Reject(String),
409 AcceptWithDelay(Duration),
411}
412
413#[derive(Debug, Clone, Serialize, Deserialize)]
415pub struct DynamicBatchingConfig {
416 pub min_batch_size: usize,
418 pub max_batch_size: usize,
420 pub batch_timeout_ms: u64,
422 pub enable_adaptive_sizing: bool,
424 pub target_utilization: f32,
426}
427
428#[derive(Debug, Clone)]
430pub struct QueueAnalysis {
431 pub queue_depth_history: Vec<(chrono::DateTime<chrono::Utc>, usize)>,
433 pub wait_time_distribution: WaitTimeDistribution,
435 pub request_patterns: RequestPatternAnalysis,
437 pub bottlenecks: Vec<BottleneckAnalysis>,
439}
440
441#[derive(Debug, Clone)]
443pub struct WaitTimeDistribution {
444 pub p50_ms: f64,
446 pub p95_ms: f64,
448 pub p99_ms: f64,
450 pub max_ms: f64,
452 pub mean_ms: f64,
454}
455
456#[derive(Debug, Clone)]
458pub struct RequestPatternAnalysis {
459 pub peak_times: Vec<chrono::DateTime<chrono::Utc>>,
461 pub rate_trend: RateTrend,
463 pub seasonality: SeasonalityPattern,
465}
466
467#[derive(Debug, Clone, Copy)]
469pub enum RateTrend {
470 Increasing,
471 Decreasing,
472 Stable,
473 Volatile,
474}
475
476#[derive(Debug, Clone)]
478pub struct SeasonalityPattern {
479 pub hourly_pattern: Vec<f32>,
481 pub daily_pattern: Vec<f32>,
483 pub weekly_pattern: Vec<f32>,
485}
486
487#[derive(Debug, Clone)]
489pub struct BottleneckAnalysis {
490 pub bottleneck_type: BottleneckType,
492 pub severity: f32,
494 pub description: String,
496 pub mitigation: String,
498}
499
500#[derive(Debug, Clone, Copy)]
502pub enum BottleneckType {
503 Memory,
505 Compute,
507 IO,
509 Scheduling,
511 Network,
513}
514
515#[derive(Debug, Clone)]
517pub struct SimulatedWorkload {
518 pub arrival_pattern: ArrivalPattern,
520 pub size_distribution: SizeDistribution,
522 pub duration_seconds: u64,
524}
525
526#[derive(Debug, Clone)]
528pub enum ArrivalPattern {
529 Constant { rate_rps: f32 },
531 Poisson { lambda: f32 },
533 Bursty {
535 burst_rate: f32,
536 quiet_rate: f32,
537 burst_duration_s: f32,
538 },
539 Seasonal {
541 base_rate: f32,
542 peaks: Vec<(f32, f32)>,
543 }, }
545
546#[derive(Debug, Clone)]
548pub enum SizeDistribution {
549 Fixed { tokens: usize },
551 Uniform {
553 min_tokens: usize,
554 max_tokens: usize,
555 },
556 Normal { mean: f32, std_dev: f32 },
558 LogNormal { mu: f32, sigma: f32 },
560}
561
562#[derive(Debug, Clone)]
564pub struct SchedulingSimulationResult {
565 pub total_requests: u64,
567 pub successful_requests: u64,
569 pub failed_requests: u64,
571 pub avg_latency_ms: f64,
573 pub p95_latency_ms: f64,
575 pub p99_latency_ms: f64,
577 pub throughput_rps: f32,
579 pub resource_utilization: Option<ResourceStats>,
581 pub bottlenecks: Vec<BottleneckAnalysis>,
583}
584
585#[derive(Debug, Clone, Serialize, Deserialize, Default)]
586pub struct ResourceStats {
587 pub gpu_memory_bytes: Option<u64>,
588 pub cpu_memory_bytes: Option<u64>,
589 pub compute_utilization: Option<f32>,
590}