//! Comprehensive tests for speculative decoding functionality //! //! Tests cover: //! - Draft model integration and loading //! - Token verification logic and acceptance decisions #![cfg(feature = "disabled_tests")] //! - Performance improvement tracking and measurement //! - Error handling and edge cases use rtx_inference::{AcceptanceDecision, DraftModel, SpeculativeDecoder, TargetModel, Token}; use rtx_inference::{InferenceError, InferenceResult}; use std::collections::HashMap; use tokio::sync::{mpsc, oneshot}; /// Mock draft model for testing #[derive(Clone)] pub struct MockDraftModel { pub model_name: String, pub generation_speed: f32, // tokens/sec pub accuracy: f32, // acceptance rate pub predictions: HashMap, Vec>, } impl MockDraftModel { pub fn new(name: &str, speed: f32, accuracy: f32) -> Self { Self { model_name: name.to_string(), generation_speed: speed, accuracy, predictions: HashMap::new(), } } pub fn add_prediction(&mut self, context: Vec, tokens: Vec) { self.predictions.insert(context, tokens); } } #[async_trait::async_trait] impl DraftModel for MockDraftModel { async fn generate_draft(&self, context: &[u32], k: usize) -> InferenceResult> { // Simulate generation latency tokio::time::sleep(tokio::time::Duration::from_micros( (k as f32 / self.generation_speed * 1_000_000.0) as u64, )) .await; if let Some(tokens) = self.predictions.get(context) { Ok(tokens.iter().take(k).cloned().collect()) } else { // Generate synthetic tokens for testing let mut tokens = Vec::new(); for i in 0..k { let token_id = 1000 + i as u32; let logits = vec![0.1, 0.2, 0.3, 0.4]; // Dummy logits tokens.push(Token::new(token_id, format!("token_{}", i), logits)); } Ok(tokens) } } fn model_name(&self) -> &str { &self.model_name } fn generation_speed(&self) -> f32 { self.generation_speed } } /// Mock target model for verification #[derive(Clone)] pub struct MockTargetModel { pub model_name: String, pub verification_predictions: HashMap<(Vec, Vec), Vec>, } impl MockTargetModel { pub fn new(name: &str) -> Self { Self { model_name: name.to_string(), verification_predictions: HashMap::new(), } } pub fn add_verification( &mut self, context: Vec, draft_tokens: Vec, acceptance: Vec, ) { self.verification_predictions .insert((context, draft_tokens), acceptance); } } #[async_trait::async_trait] impl TargetModel for MockTargetModel { async fn verify_draft( &self, context: &[u32], draft_tokens: &[Token], ) -> InferenceResult> { let draft_ids: Vec = draft_tokens.iter().map(|t| t.id).collect(); if let Some(acceptance) = self .verification_predictions .get(&(context.to_vec(), draft_ids)) { Ok(acceptance.clone()) } else { // Default verification based on mock accuracy let mut acceptance = Vec::new(); for (i, _token) in draft_tokens.iter().enumerate() { // Simulate decreasing acceptance probability with position let accept_prob = 0.8 - (i as f32 * 0.1); acceptance.push(fastrand::f32() < accept_prob); } Ok(acceptance) } } async fn generate_continuation(&self, context: &[u32]) -> InferenceResult { // Generate a continuation token let token_id = 2000 + context.len() as u32; let logits = vec![0.5, 0.3, 0.2]; Ok(Token::new( token_id, format!("cont_{}", context.len()), logits, )) } fn model_name(&self) -> &str { &self.model_name } } /// Test speculative decoding configuration #[tokio::test] async fn test_speculative_config_creation() { use rtx_inference::speculative::{DraftModelType, SpeculativeConfig}; let config = SpeculativeConfig::new() .with_draft_model(DraftModelType::SmallModel("gpt2-small".to_string())) .with_max_draft_tokens(4) .with_acceptance_threshold(0.7) .with_lookahead_tokens(2) .with_performance_tracking(true); assert_eq!(config.max_draft_tokens, 4); assert_eq!(config.acceptance_threshold, 0.7); assert_eq!(config.lookahead_tokens, 2); assert!(config.performance_tracking); // Test validation let invalid_config = SpeculativeConfig::new().with_max_draft_tokens(0); // Should be > 0 assert!(invalid_config.validate().is_err()); } /// Test draft model integration #[tokio::test] async fn test_draft_model_integration() { let mut draft_model = MockDraftModel::new("test-draft", 100.0, 0.8); // Setup test predictions let context = vec![1, 2, 3]; let expected_tokens = vec![ Token::new(10, "hello".to_string(), vec![0.5, 0.3, 0.2]), Token::new(11, "world".to_string(), vec![0.4, 0.4, 0.2]), ]; draft_model.add_prediction(context.clone(), expected_tokens.clone()); // Test generation let generated = draft_model.generate_draft(&context, 2).await.unwrap(); assert_eq!(generated.len(), 2); assert_eq!(generated[0].id, 10); assert_eq!(generated[0].text, "hello"); assert_eq!(generated[1].id, 11); assert_eq!(generated[1].text, "world"); } /// Test token verification logic #[tokio::test] async fn test_token_verification() { let mut target_model = MockTargetModel::new("test-target"); // Setup verification predictions let context = vec![1, 2, 3]; let draft_tokens = vec![10, 11, 12]; let expected_acceptance = vec![true, true, false]; // First two accepted, third rejected target_model.add_verification( context.clone(), draft_tokens.clone(), expected_acceptance.clone(), ); // Test verification let tokens = vec![ Token::new(10, "hello".to_string(), vec![0.5, 0.3, 0.2]), Token::new(11, "world".to_string(), vec![0.4, 0.4, 0.2]), Token::new(12, "test".to_string(), vec![0.3, 0.3, 0.4]), ]; let acceptance = target_model.verify_draft(&context, &tokens).await.unwrap(); assert_eq!(acceptance, expected_acceptance); } /// Test acceptance/rejection decisions #[tokio::test] async fn test_acceptance_rejection_decisions() { use rtx_inference::speculative::{AcceptanceDecision, SpeculativeDecoder}; let draft_model = MockDraftModel::new("draft", 150.0, 0.8); let target_model = MockTargetModel::new("target"); let config = rtx_inference::speculative::SpeculativeConfig::new() .with_max_draft_tokens(3) .with_acceptance_threshold(0.75); let decoder = SpeculativeDecoder::new(Box::new(draft_model), Box::new(target_model), config); let context = vec![1, 2, 3]; let draft_tokens = vec![ Token::new(10, "hello".to_string(), vec![0.8, 0.1, 0.1]), // High confidence Token::new(11, "world".to_string(), vec![0.7, 0.2, 0.1]), // Medium confidence Token::new(12, "test".to_string(), vec![0.4, 0.3, 0.3]), // Low confidence ]; let verification = vec![true, true, false]; let decision = decoder .make_acceptance_decision(draft_tokens.as_slice(), &verification) .await .unwrap(); match decision { AcceptanceDecision::AcceptAll => panic!("Should not accept all with rejection"), AcceptanceDecision::AcceptPrefix(n) => { assert_eq!(n, 2); // Should accept first two tokens } AcceptanceDecision::RejectAll => panic!("Should accept some tokens"), } } /// Test performance improvement tracking #[tokio::test] async fn test_performance_tracking() { use rtx_inference::speculative::{PerformanceMetrics, SpeculativeDecoder}; let draft_model = MockDraftModel::new("draft", 200.0, 0.85); let target_model = MockTargetModel::new("target"); let config = rtx_inference::speculative::SpeculativeConfig::new().with_performance_tracking(true); let mut decoder = SpeculativeDecoder::new(Box::new(draft_model), Box::new(target_model), config); // Simulate several decoding rounds for i in 0..10 { let context = vec![1, 2, i]; let _result = decoder.decode_step(&context).await.unwrap(); } let metrics = decoder.get_performance_metrics().await; assert!(metrics.total_steps > 0); assert!(metrics.tokens_accepted >= 0); assert!(metrics.acceptance_rate >= 0.0 && metrics.acceptance_rate <= 1.0); assert!(metrics.speedup_ratio >= 0.0); assert!(metrics.draft_generation_time.as_nanos() > 0); assert!(metrics.verification_time.as_nanos() > 0); } /// Test speculative decoding with different draft lengths #[tokio::test] async fn test_variable_draft_lengths() { use rtx_inference::speculative::SpeculativeDecoder; let draft_model = MockDraftModel::new("draft", 100.0, 0.8); let target_model = MockTargetModel::new("target"); for k in 1..=8 { let config = rtx_inference::speculative::SpeculativeConfig::new().with_max_draft_tokens(k); let decoder = SpeculativeDecoder::new( Box::new(draft_model.clone()), Box::new(target_model.clone()), config, ); let context = vec![1, 2, 3]; let result = decoder.decode_step(&context).await.unwrap(); assert!(!result.accepted_tokens.is_empty()); assert!(result.accepted_tokens.len() <= k); } } /// Test error handling in speculative decoding #[tokio::test] async fn test_error_handling() { use rtx_inference::speculative::{SpeculativeDecoder, SpeculativeError}; // Test with invalid configuration let draft_model = MockDraftModel::new("draft", 100.0, 0.8); let target_model = MockTargetModel::new("target"); let invalid_config = rtx_inference::speculative::SpeculativeConfig::new().with_max_draft_tokens(0); // Invalid let result = SpeculativeDecoder::try_new( Box::new(draft_model), Box::new(target_model), invalid_config, ); assert!(result.is_err()); // Test with empty context let draft_model = MockDraftModel::new("draft", 100.0, 0.8); let target_model = MockTargetModel::new("target"); let config = rtx_inference::speculative::SpeculativeConfig::new(); let decoder = SpeculativeDecoder::new(Box::new(draft_model), Box::new(target_model), config); let result = decoder.decode_step(&[]).await; assert!(result.is_err()); if let Err(InferenceError::InvalidRequest { message }) = result { assert!(message.contains("empty context")); } else { panic!("Expected InvalidRequest error"); } } /// Test concurrent speculative decoding #[tokio::test] async fn test_concurrent_decoding() { use rtx_inference::speculative::SpeculativeDecoder; use std::sync::Arc; let draft_model = MockDraftModel::new("draft", 150.0, 0.8); let target_model = MockTargetModel::new("target"); let config = rtx_inference::speculative::SpeculativeConfig::new(); let decoder = Arc::new(SpeculativeDecoder::new( Box::new(draft_model), Box::new(target_model), config, )); // Launch concurrent decoding tasks let mut handles = Vec::new(); for i in 0..5 { let decoder = Arc::clone(&decoder); let handle = tokio::spawn(async move { let context = vec![i, i + 1, i + 2]; decoder.decode_step(&context).await }); handles.push(handle); } // Wait for all tasks to complete for handle in handles { let result = handle.await.unwrap(); assert!(result.is_ok()); } } /// Test memory efficiency with large contexts #[tokio::test] async fn test_memory_efficiency() { use rtx_inference::speculative::SpeculativeDecoder; let draft_model = MockDraftModel::new("draft", 100.0, 0.8); let target_model = MockTargetModel::new("target"); let config = rtx_inference::speculative::SpeculativeConfig::new().with_max_draft_tokens(4); let decoder = SpeculativeDecoder::new(Box::new(draft_model), Box::new(target_model), config); // Test with progressively larger contexts for size in [100, 1000, 10000] { let large_context: Vec = (0..size).collect(); let result = decoder.decode_step(&large_context).await.unwrap(); assert!(!result.accepted_tokens.is_empty()); } } /// Test performance benchmarking #[tokio::test] async fn test_performance_benchmarks() { use rtx_inference::speculative::SpeculativeDecoder; use std::time::Instant; let draft_model = MockDraftModel::new("fast-draft", 500.0, 0.9); // Very fast draft let target_model = MockTargetModel::new("target"); let config = rtx_inference::speculative::SpeculativeConfig::new() .with_max_draft_tokens(4) .with_performance_tracking(true); let mut decoder = SpeculativeDecoder::new(Box::new(draft_model), Box::new(target_model), config); let context = vec![1, 2, 3, 4, 5]; let num_steps = 50; let start = Instant::now(); for _ in 0..num_steps { let _result = decoder.decode_step(&context).await.unwrap(); } let elapsed = start.elapsed(); let metrics = decoder.get_performance_metrics().await; // Verify performance targets assert!( metrics.acceptance_rate > 0.5, "Acceptance rate should be reasonable" ); assert!( metrics.speedup_ratio > 1.0, "Should show speedup over baseline" ); let tokens_per_sec = (metrics.tokens_accepted as f64) / elapsed.as_secs_f64(); println!("Achieved {} tokens/sec", tokens_per_sec); // Target: Should achieve reasonable throughput assert!(tokens_per_sec > 10.0, "Should achieve minimum throughput"); } /// Test adaptive draft length optimization #[tokio::test] async fn test_adaptive_draft_length() { use rtx_inference::speculative::{AdaptiveConfig, SpeculativeDecoder}; let draft_model = MockDraftModel::new("draft", 200.0, 0.75); let target_model = MockTargetModel::new("target"); let config = rtx_inference::speculative::SpeculativeConfig::new().with_adaptive_draft_length( AdaptiveConfig { min_length: 1, max_length: 8, target_acceptance_rate: 0.8, adjustment_factor: 0.1, }, ); let mut decoder = SpeculativeDecoder::new(Box::new(draft_model), Box::new(target_model), config); let context = vec![1, 2, 3]; // Run several steps to allow adaptation for _ in 0..20 { let _result = decoder.decode_step(&context).await.unwrap(); } let metrics = decoder.get_performance_metrics().await; let current_length = decoder.get_current_draft_length_sync(); // Verify adaptive behavior assert!(current_length >= 1 && current_length <= 8); assert!(metrics.total_steps > 0); }