460 lines
15 KiB
Rust
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);
|
|
}
|