Files
rustytorch/crates/production/rtx-inference/tests/speculative_decoding.rs
T
2026-03-04 00:08:42 +00:00

460 lines
15 KiB
Rust

//! 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<u32>, Vec<Token>>,
}
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<u32>, tokens: Vec<Token>) {
self.predictions.insert(context, tokens);
}
}
#[async_trait::async_trait]
impl DraftModel for MockDraftModel {
async fn generate_draft(&self, context: &[u32], k: usize) -> InferenceResult<Vec<Token>> {
// 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<u32>, Vec<u32>), Vec<bool>>,
}
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<u32>,
draft_tokens: Vec<u32>,
acceptance: Vec<bool>,
) {
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<Vec<bool>> {
let draft_ids: Vec<u32> = 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<Token> {
// 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<u32> = (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);
}