Skip to main content

ferrum_engine/pipeline/
executor.rs

1//! Pipeline Executor
2//!
3//! This module provides a high-level pipeline executor that manages
4//! the separation between prefill and decode phases for optimal efficiency.
5
6use 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/// Current execution phase
21#[derive(Debug, Clone, Copy, PartialEq, Eq)]
22pub enum ExecutionPhase {
23    /// Processing initial prompt
24    Prefill,
25    /// Generating tokens
26    Decode,
27    /// Completed
28    Done,
29}
30
31/// Pipeline configuration
32#[derive(Debug, Clone)]
33pub struct PipelineConfig {
34    /// Chunked prefill configuration
35    pub chunked_prefill: ChunkedPrefillConfig,
36    /// Maximum decode batch size
37    pub max_decode_batch: usize,
38    /// Maximum prefill batch size
39    pub max_prefill_batch: usize,
40    /// Target tokens per second for scheduling
41    pub target_tps: Option<f32>,
42    /// Enable prefill-decode interleaving
43    pub enable_interleaving: bool,
44    /// Prefill priority ratio (0.0-1.0)
45    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/// Sequence being processed in the pipeline
62#[derive(Debug)]
63pub struct PipelineSequence {
64    /// Unique sequence ID
65    pub id: u64,
66    /// Current phase
67    pub phase: ExecutionPhase,
68    /// Input tokens (for prefill)
69    pub input_tokens: Vec<TokenId>,
70    /// Generated tokens
71    pub generated_tokens: Vec<TokenId>,
72    /// KV cache handle
73    pub kv_cache: Option<Arc<dyn KvCacheHandle>>,
74    /// Sampling parameters
75    pub sampling_params: SamplingParams,
76    /// Priority
77    pub priority: Priority,
78    /// Random number generator
79    pub rng: SamplingRng,
80    /// Prefill start time
81    pub prefill_start: Option<Instant>,
82    /// First token time (TTFT)
83    pub first_token_time: Option<Instant>,
84    /// Total time to complete
85    pub completion_time: Option<Instant>,
86}
87
88impl PipelineSequence {
89    /// Create new sequence
90    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    /// Total tokens in sequence
113    pub fn total_tokens(&self) -> usize {
114        self.input_tokens.len() + self.generated_tokens.len()
115    }
116
117    /// Check if generation should stop
118    pub fn should_stop(&self, vocab_size: usize) -> bool {
119        // Check max tokens
120        if self.generated_tokens.len() >= self.sampling_params.max_tokens {
121            return true;
122        }
123
124        // Check for EOS token (last ~10 tokens are usually special)
125        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    /// Get time to first token in milliseconds
135    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
143/// Pipeline executor for efficient prefill/decode separation
144pub struct PipelineExecutor {
145    /// Configuration
146    config: PipelineConfig,
147    /// Model executor
148    model_executor: Arc<dyn ModelExecutor + Send + Sync>,
149    /// Sampler
150    sampler: Arc<dyn Sampler + Send + Sync>,
151    /// Chunked prefill executor
152    chunked_prefill: ChunkedPrefillExecutor,
153    /// Active sequences
154    sequences: RwLock<HashMap<u64, PipelineSequence>>,
155    /// Next sequence ID
156    next_id: AtomicU64,
157    /// Statistics
158    stats: PipelineStats,
159}
160
161/// Pipeline statistics
162#[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    /// Create new pipeline executor
173    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    /// Submit a new sequence for processing
198    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    /// Execute prefill for a sequence
214    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        // Sample first token
231        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        // Update stats
248        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    /// Execute a single decode step for a sequence
267    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        // Create decode input
296        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        // Run decode
305        let output = self.model_executor.decode(&decode_input).await?;
306        let decode_time = start.elapsed();
307
308        // Sample next token
309        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        // Update stats
332        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    /// Run a full generation for a sequence
361    pub async fn generate(&self, sequence_id: u64) -> Result<Vec<TokenId>> {
362        // Run prefill
363        self.run_prefill(sequence_id).await?;
364
365        // Run decode loop
366        loop {
367            match self.run_decode_step(sequence_id).await? {
368                Some(_token) => continue,
369                None => break,
370            }
371        }
372
373        // Get generated tokens
374        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    /// Get sequence by ID
383    pub fn get_sequence(&self, sequence_id: u64) -> Option<PipelineSequence> {
384        self.sequences.read().get(&sequence_id).cloned()
385    }
386
387    /// Remove completed sequence
388    pub fn remove_sequence(&self, sequence_id: u64) -> Option<PipelineSequence> {
389        self.sequences.write().remove(&sequence_id)
390    }
391
392    /// Get number of active sequences
393    pub fn active_count(&self) -> usize {
394        self.sequences.read().len()
395    }
396
397    /// Get sequences in prefill phase
398    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    /// Get sequences in decode phase
407    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    /// Get statistics
416    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    /// Get configuration
442    pub fn config(&self) -> &PipelineConfig {
443        &self.config
444    }
445}
446
447/// Snapshot of pipeline statistics
448#[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
468// Implement Clone for PipelineSequence
469impl 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// ============================================================================
488// Tests
489// ============================================================================
490
491#[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        // Not yet at max tokens
516        seq.generated_tokens.push(TokenId::new(1));
517        seq.generated_tokens.push(TokenId::new(2));
518        assert!(!seq.should_stop(32000));
519
520        // At max tokens
521        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}