1use super::chunked_prefill::{ChunkedPrefillConfig, ChunkedPrefillExecutor};
7use ferrum_interfaces::{
8 model_executor::DecodeInput, sampler::SamplingRng, KvCacheHandle, ModelExecutor, Sampler,
9 TensorRef,
10};
11use ferrum_models::CandleTensorWrapper;
12use ferrum_types::{FerrumError, Priority, Result, SamplingParams, TokenId};
13use parking_lot::RwLock;
14use std::collections::HashMap;
15use std::sync::atomic::{AtomicU64, Ordering};
16use std::sync::Arc;
17use std::time::Instant;
18use tracing::{debug, info};
19
20#[derive(Debug, Clone, Copy, PartialEq, Eq)]
22pub enum ExecutionPhase {
23 Prefill,
25 Decode,
27 Done,
29}
30
31#[derive(Debug, Clone)]
33pub struct PipelineConfig {
34 pub chunked_prefill: ChunkedPrefillConfig,
36 pub max_decode_batch: usize,
38 pub max_prefill_batch: usize,
40 pub target_tps: Option<f32>,
42 pub enable_interleaving: bool,
44 pub prefill_priority: f32,
46}
47
48impl Default for PipelineConfig {
49 fn default() -> Self {
50 Self {
51 chunked_prefill: ChunkedPrefillConfig::default(),
52 max_decode_batch: 256,
53 max_prefill_batch: 8,
54 target_tps: None,
55 enable_interleaving: true,
56 prefill_priority: 0.5,
57 }
58 }
59}
60
61#[derive(Debug)]
63pub struct PipelineSequence {
64 pub id: u64,
66 pub phase: ExecutionPhase,
68 pub input_tokens: Vec<TokenId>,
70 pub generated_tokens: Vec<TokenId>,
72 pub kv_cache: Option<Arc<dyn KvCacheHandle>>,
74 pub sampling_params: SamplingParams,
76 pub priority: Priority,
78 pub rng: SamplingRng,
80 pub prefill_start: Option<Instant>,
82 pub first_token_time: Option<Instant>,
84 pub completion_time: Option<Instant>,
86}
87
88impl PipelineSequence {
89 pub fn new(
91 id: u64,
92 input_tokens: Vec<TokenId>,
93 sampling_params: SamplingParams,
94 priority: Priority,
95 ) -> Self {
96 let seed = sampling_params.seed.unwrap_or(42);
97 Self {
98 id,
99 phase: ExecutionPhase::Prefill,
100 input_tokens,
101 generated_tokens: Vec::new(),
102 kv_cache: None,
103 sampling_params,
104 priority,
105 rng: SamplingRng::seeded(seed),
106 prefill_start: None,
107 first_token_time: None,
108 completion_time: None,
109 }
110 }
111
112 pub fn total_tokens(&self) -> usize {
114 self.input_tokens.len() + self.generated_tokens.len()
115 }
116
117 pub fn should_stop(&self, vocab_size: usize) -> bool {
119 if self.generated_tokens.len() >= self.sampling_params.max_tokens {
121 return true;
122 }
123
124 if let Some(&last_token) = self.generated_tokens.last() {
126 if last_token.get() >= (vocab_size.saturating_sub(10)) as u32 {
127 return true;
128 }
129 }
130
131 false
132 }
133
134 pub fn ttft_ms(&self) -> Option<u64> {
136 match (self.prefill_start, self.first_token_time) {
137 (Some(start), Some(first)) => Some(first.duration_since(start).as_millis() as u64),
138 _ => None,
139 }
140 }
141}
142
143pub struct PipelineExecutor {
145 config: PipelineConfig,
147 model_executor: Arc<dyn ModelExecutor + Send + Sync>,
149 sampler: Arc<dyn Sampler + Send + Sync>,
151 chunked_prefill: ChunkedPrefillExecutor,
153 sequences: RwLock<HashMap<u64, PipelineSequence>>,
155 next_id: AtomicU64,
157 stats: PipelineStats,
159}
160
161#[derive(Debug, Default)]
163struct PipelineStats {
164 total_prefill_tokens: AtomicU64,
165 total_decode_tokens: AtomicU64,
166 total_prefill_time_us: AtomicU64,
167 total_decode_time_us: AtomicU64,
168 completed_sequences: AtomicU64,
169}
170
171impl PipelineExecutor {
172 pub fn new(
174 model_executor: Arc<dyn ModelExecutor + Send + Sync>,
175 sampler: Arc<dyn Sampler + Send + Sync>,
176 config: PipelineConfig,
177 ) -> Self {
178 info!(
179 "Creating PipelineExecutor: max_decode_batch={}, max_prefill_batch={}",
180 config.max_decode_batch, config.max_prefill_batch
181 );
182
183 let chunked_prefill =
184 ChunkedPrefillExecutor::new(model_executor.clone(), config.chunked_prefill.clone());
185
186 Self {
187 config,
188 model_executor,
189 sampler,
190 chunked_prefill,
191 sequences: RwLock::new(HashMap::new()),
192 next_id: AtomicU64::new(0),
193 stats: PipelineStats::default(),
194 }
195 }
196
197 pub fn submit(
199 &self,
200 input_tokens: Vec<TokenId>,
201 sampling_params: SamplingParams,
202 priority: Priority,
203 ) -> u64 {
204 let id = self.next_id.fetch_add(1, Ordering::Relaxed);
205 let sequence = PipelineSequence::new(id, input_tokens, sampling_params, priority);
206
207 self.sequences.write().insert(id, sequence);
208 debug!("Submitted sequence {} for processing", id);
209
210 id
211 }
212
213 pub async fn run_prefill(&self, sequence_id: u64) -> Result<()> {
215 let tokens = {
216 let mut sequences = self.sequences.write();
217 let seq = sequences
218 .get_mut(&sequence_id)
219 .ok_or_else(|| FerrumError::internal("Sequence not found"))?;
220
221 seq.phase = ExecutionPhase::Prefill;
222 seq.prefill_start = Some(Instant::now());
223 seq.input_tokens.clone()
224 };
225
226 let start = Instant::now();
227 let output = self.chunked_prefill.execute(tokens.clone()).await?;
228 let prefill_time = start.elapsed();
229
230 let logits_vec = output.logits.to_vec_f32()?;
232 let first_token = {
233 let mut sequences = self.sequences.write();
234 let seq = sequences
235 .get_mut(&sequence_id)
236 .ok_or_else(|| FerrumError::internal("Sequence not found"))?;
237
238 let token = self.sampler.sample(&logits_vec, &mut seq.rng)?;
239 seq.generated_tokens.push(token);
240 seq.kv_cache = Some(output.kv_cache);
241 seq.phase = ExecutionPhase::Decode;
242 seq.first_token_time = Some(Instant::now());
243
244 token
245 };
246
247 self.stats
249 .total_prefill_tokens
250 .fetch_add(tokens.len() as u64, Ordering::Relaxed);
251 self.stats
252 .total_prefill_time_us
253 .fetch_add(prefill_time.as_micros() as u64, Ordering::Relaxed);
254
255 debug!(
256 "Prefill complete for seq {}: {} tokens in {:?}, first token: {}",
257 sequence_id,
258 tokens.len(),
259 prefill_time,
260 first_token.get()
261 );
262
263 Ok(())
264 }
265
266 pub async fn run_decode_step(&self, sequence_id: u64) -> Result<Option<TokenId>> {
268 let (last_token, kv_cache) = {
269 let sequences = self.sequences.read();
270 let seq = sequences
271 .get(&sequence_id)
272 .ok_or_else(|| FerrumError::internal("Sequence not found"))?;
273
274 if seq.phase != ExecutionPhase::Decode {
275 return Err(FerrumError::internal("Sequence not in decode phase"));
276 }
277
278 let last_token = seq
279 .generated_tokens
280 .last()
281 .copied()
282 .ok_or_else(|| FerrumError::internal("No tokens generated"))?;
283
284 let kv_cache = seq
285 .kv_cache
286 .as_ref()
287 .ok_or_else(|| FerrumError::internal("No KV cache"))?
288 .clone();
289
290 (last_token, kv_cache)
291 };
292
293 let start = Instant::now();
294
295 let tensor = candle_core::Tensor::new(&[last_token.get()], &candle_core::Device::Cpu)
297 .map_err(|e| FerrumError::model(format!("Tensor error: {}", e)))?
298 .unsqueeze(0)
299 .map_err(|e| FerrumError::model(format!("Unsqueeze error: {}", e)))?;
300
301 let tensor_ref: TensorRef = Arc::new(CandleTensorWrapper::new(tensor));
302 let decode_input = DecodeInput::new(tensor_ref, kv_cache);
303
304 let output = self.model_executor.decode(&decode_input).await?;
306 let decode_time = start.elapsed();
307
308 let logits_vec = output.logits.to_vec_f32()?;
310 let (next_token, should_stop) = {
311 let mut sequences = self.sequences.write();
312 let seq = sequences
313 .get_mut(&sequence_id)
314 .ok_or_else(|| FerrumError::internal("Sequence not found"))?;
315
316 let token = self.sampler.sample(&logits_vec, &mut seq.rng)?;
317 seq.generated_tokens.push(token);
318 seq.kv_cache = Some(output.kv_cache);
319
320 let vocab_size = self.model_executor.info().vocab_size;
321 let stop = seq.should_stop(vocab_size);
322
323 if stop {
324 seq.phase = ExecutionPhase::Done;
325 seq.completion_time = Some(Instant::now());
326 }
327
328 (token, stop)
329 };
330
331 self.stats
333 .total_decode_tokens
334 .fetch_add(1, Ordering::Relaxed);
335 self.stats
336 .total_decode_time_us
337 .fetch_add(decode_time.as_micros() as u64, Ordering::Relaxed);
338
339 if should_stop {
340 self.stats
341 .completed_sequences
342 .fetch_add(1, Ordering::Relaxed);
343 debug!(
344 "Decode complete for seq {}: generated {} tokens",
345 sequence_id,
346 {
347 let sequences = self.sequences.read();
348 sequences
349 .get(&sequence_id)
350 .map(|s| s.generated_tokens.len())
351 .unwrap_or(0)
352 }
353 );
354 return Ok(None);
355 }
356
357 Ok(Some(next_token))
358 }
359
360 pub async fn generate(&self, sequence_id: u64) -> Result<Vec<TokenId>> {
362 self.run_prefill(sequence_id).await?;
364
365 loop {
367 match self.run_decode_step(sequence_id).await? {
368 Some(_token) => continue,
369 None => break,
370 }
371 }
372
373 let sequences = self.sequences.read();
375 let seq = sequences
376 .get(&sequence_id)
377 .ok_or_else(|| FerrumError::internal("Sequence not found"))?;
378
379 Ok(seq.generated_tokens.clone())
380 }
381
382 pub fn get_sequence(&self, sequence_id: u64) -> Option<PipelineSequence> {
384 self.sequences.read().get(&sequence_id).cloned()
385 }
386
387 pub fn remove_sequence(&self, sequence_id: u64) -> Option<PipelineSequence> {
389 self.sequences.write().remove(&sequence_id)
390 }
391
392 pub fn active_count(&self) -> usize {
394 self.sequences.read().len()
395 }
396
397 pub fn prefill_count(&self) -> usize {
399 self.sequences
400 .read()
401 .values()
402 .filter(|s| s.phase == ExecutionPhase::Prefill)
403 .count()
404 }
405
406 pub fn decode_count(&self) -> usize {
408 self.sequences
409 .read()
410 .values()
411 .filter(|s| s.phase == ExecutionPhase::Decode)
412 .count()
413 }
414
415 pub fn get_stats(&self) -> PipelineStatsSnapshot {
417 let total_prefill = self.stats.total_prefill_tokens.load(Ordering::Relaxed);
418 let total_decode = self.stats.total_decode_tokens.load(Ordering::Relaxed);
419 let prefill_time_us = self.stats.total_prefill_time_us.load(Ordering::Relaxed);
420 let decode_time_us = self.stats.total_decode_time_us.load(Ordering::Relaxed);
421 let completed = self.stats.completed_sequences.load(Ordering::Relaxed);
422
423 PipelineStatsSnapshot {
424 total_prefill_tokens: total_prefill,
425 total_decode_tokens: total_decode,
426 avg_prefill_tokens_per_sec: if prefill_time_us > 0 {
427 (total_prefill as f64 * 1_000_000.0) / prefill_time_us as f64
428 } else {
429 0.0
430 },
431 avg_decode_tokens_per_sec: if decode_time_us > 0 {
432 (total_decode as f64 * 1_000_000.0) / decode_time_us as f64
433 } else {
434 0.0
435 },
436 completed_sequences: completed,
437 active_sequences: self.active_count() as u64,
438 }
439 }
440
441 pub fn config(&self) -> &PipelineConfig {
443 &self.config
444 }
445}
446
447#[derive(Debug, Clone)]
449pub struct PipelineStatsSnapshot {
450 pub total_prefill_tokens: u64,
451 pub total_decode_tokens: u64,
452 pub avg_prefill_tokens_per_sec: f64,
453 pub avg_decode_tokens_per_sec: f64,
454 pub completed_sequences: u64,
455 pub active_sequences: u64,
456}
457
458impl std::fmt::Debug for PipelineExecutor {
459 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
460 f.debug_struct("PipelineExecutor")
461 .field("active_sequences", &self.active_count())
462 .field("prefill_count", &self.prefill_count())
463 .field("decode_count", &self.decode_count())
464 .finish()
465 }
466}
467
468impl Clone for PipelineSequence {
470 fn clone(&self) -> Self {
471 Self {
472 id: self.id,
473 phase: self.phase,
474 input_tokens: self.input_tokens.clone(),
475 generated_tokens: self.generated_tokens.clone(),
476 kv_cache: self.kv_cache.clone(),
477 sampling_params: self.sampling_params.clone(),
478 priority: self.priority,
479 rng: SamplingRng::seeded(self.sampling_params.seed.unwrap_or(42)),
480 prefill_start: self.prefill_start,
481 first_token_time: self.first_token_time,
482 completion_time: self.completion_time,
483 }
484 }
485}
486
487#[cfg(test)]
492mod tests {
493 use super::*;
494
495 #[test]
496 fn test_pipeline_sequence_creation() {
497 let tokens: Vec<TokenId> = (0..10).map(|i| TokenId::new(i as u32)).collect();
498 let params = SamplingParams::default();
499 let seq = PipelineSequence::new(1, tokens, params, Priority::Normal);
500
501 assert_eq!(seq.id, 1);
502 assert_eq!(seq.phase, ExecutionPhase::Prefill);
503 assert_eq!(seq.input_tokens.len(), 10);
504 assert!(seq.generated_tokens.is_empty());
505 }
506
507 #[test]
508 fn test_pipeline_sequence_stop_conditions() {
509 let tokens: Vec<TokenId> = vec![TokenId::new(0)];
510 let mut params = SamplingParams::default();
511 params.max_tokens = 5;
512
513 let mut seq = PipelineSequence::new(1, tokens, params, Priority::Normal);
514
515 seq.generated_tokens.push(TokenId::new(1));
517 seq.generated_tokens.push(TokenId::new(2));
518 assert!(!seq.should_stop(32000));
519
520 seq.generated_tokens.push(TokenId::new(3));
522 seq.generated_tokens.push(TokenId::new(4));
523 seq.generated_tokens.push(TokenId::new(5));
524 assert!(seq.should_stop(32000));
525 }
526
527 #[test]
528 fn test_pipeline_config_defaults() {
529 let config = PipelineConfig::default();
530
531 assert_eq!(config.max_decode_batch, 256);
532 assert_eq!(config.max_prefill_batch, 8);
533 assert!(config.enable_interleaving);
534 }
535
536 #[test]
537 fn test_execution_phase() {
538 assert_eq!(ExecutionPhase::Prefill, ExecutionPhase::Prefill);
539 assert_ne!(ExecutionPhase::Prefill, ExecutionPhase::Decode);
540 }
541}