use crate::error::MlError;
#[derive(Debug, Clone)]
pub struct PredictionRequest {
pub id: u64,
pub inputs: Vec<Vec<f32>>,
pub input_shapes: Vec<Vec<usize>>,
}
#[derive(Debug, Clone)]
pub struct PredictionResult {
pub id: u64,
pub outputs: Vec<Vec<f32>>,
pub output_shapes: Vec<Vec<usize>>,
pub latency_ms: f64,
}
#[derive(Debug, Clone)]
pub struct AdaptiveBatchConfig {
pub min_batch_size: usize,
pub max_batch_size: usize,
pub target_latency_ms: f64,
pub adaptation_rate: f64,
}
impl Default for AdaptiveBatchConfig {
fn default() -> Self {
Self {
min_batch_size: 1,
max_batch_size: 64,
target_latency_ms: 50.0,
adaptation_rate: 0.1,
}
}
}
impl AdaptiveBatchConfig {
pub fn validate(&self) -> Result<(), MlError> {
if self.min_batch_size == 0 {
return Err(MlError::InvalidConfig(
"min_batch_size must be at least 1".into(),
));
}
if self.max_batch_size < self.min_batch_size {
return Err(MlError::InvalidConfig(
"max_batch_size must be >= min_batch_size".into(),
));
}
if !(0.0..=1.0).contains(&self.adaptation_rate) {
return Err(MlError::InvalidConfig(
"adaptation_rate must be in [0.0, 1.0]".into(),
));
}
if self.target_latency_ms <= 0.0 {
return Err(MlError::InvalidConfig(
"target_latency_ms must be positive".into(),
));
}
Ok(())
}
}
pub struct AdaptiveBatcher {
config: AdaptiveBatchConfig,
current_batch_size: usize,
recent_latencies: Vec<f64>,
total_batches: u64,
total_items: u64,
window_size: usize,
}
impl AdaptiveBatcher {
pub fn new(config: AdaptiveBatchConfig) -> Self {
let start = config.min_batch_size;
Self {
config,
current_batch_size: start,
recent_latencies: Vec::new(),
total_batches: 0,
total_items: 0,
window_size: 10,
}
}
pub fn recommended_batch_size(&self) -> usize {
self.current_batch_size
}
pub fn update_latency(&mut self, latency_ms: f64, batch_size: usize) {
self.recent_latencies.push(latency_ms);
if self.recent_latencies.len() > self.window_size {
self.recent_latencies.remove(0);
}
self.total_batches += 1;
self.total_items += batch_size as u64;
let avg = self.average_latency_ms();
let target = self.config.target_latency_ms;
let rate = self.config.adaptation_rate;
let min_bs = self.config.min_batch_size as f64;
let max_bs = self.config.max_batch_size as f64;
let current = self.current_batch_size as f64;
let new_size = if avg > target {
let reduction = (current * rate * (avg - target) / target).max(1.0);
(current - reduction).max(min_bs)
} else {
let gain = (current * rate * (target - avg) / target).max(1.0);
(current + gain).min(max_bs)
};
self.current_batch_size = (new_size.round() as usize)
.max(self.config.min_batch_size)
.min(self.config.max_batch_size);
}
pub fn create_batches(&self, requests: Vec<PredictionRequest>) -> Vec<Vec<PredictionRequest>> {
if requests.is_empty() {
return Vec::new();
}
let bs = self.current_batch_size.max(1);
requests.chunks(bs).map(|chunk| chunk.to_vec()).collect()
}
pub fn average_latency_ms(&self) -> f64 {
if self.recent_latencies.is_empty() {
return 0.0;
}
self.recent_latencies.iter().sum::<f64>() / self.recent_latencies.len() as f64
}
pub fn total_batches(&self) -> u64 {
self.total_batches
}
pub fn total_items(&self) -> u64 {
self.total_items
}
pub fn config(&self) -> &AdaptiveBatchConfig {
&self.config
}
}
#[cfg(test)]
mod tests {
use super::*;
fn default_batcher() -> AdaptiveBatcher {
AdaptiveBatcher::new(AdaptiveBatchConfig::default())
}
fn make_request(id: u64) -> PredictionRequest {
PredictionRequest {
id,
inputs: vec![vec![1.0, 2.0, 3.0]],
input_shapes: vec![vec![3]],
}
}
#[test]
fn test_construction_with_default_config() {
let batcher = default_batcher();
assert_eq!(
batcher.recommended_batch_size(),
AdaptiveBatchConfig::default().min_batch_size
);
}
#[test]
fn test_recommended_batch_size_starts_at_min() {
let config = AdaptiveBatchConfig {
min_batch_size: 4,
max_batch_size: 64,
..Default::default()
};
let batcher = AdaptiveBatcher::new(config);
assert_eq!(batcher.recommended_batch_size(), 4);
}
#[test]
fn test_update_latency_adjusts_up_when_fast() {
let mut batcher = AdaptiveBatcher::new(AdaptiveBatchConfig {
min_batch_size: 1,
max_batch_size: 128,
target_latency_ms: 100.0,
adaptation_rate: 0.5,
});
let initial = batcher.recommended_batch_size();
batcher.update_latency(10.0, initial);
assert!(
batcher.recommended_batch_size() > initial,
"batch size should grow when latency is well below target"
);
}
#[test]
fn test_update_latency_adjusts_down_when_slow() {
let mut batcher = AdaptiveBatcher::new(AdaptiveBatchConfig {
min_batch_size: 1,
max_batch_size: 64,
target_latency_ms: 50.0,
adaptation_rate: 0.5,
});
for _ in 0..10 {
let sz = batcher.recommended_batch_size();
batcher.update_latency(10.0, sz);
}
let high = batcher.recommended_batch_size();
batcher.update_latency(9999.0, high);
assert!(
batcher.recommended_batch_size() < high,
"batch size should shrink when latency exceeds target"
);
}
#[test]
fn test_batch_size_does_not_exceed_max() {
let mut batcher = AdaptiveBatcher::new(AdaptiveBatchConfig {
min_batch_size: 1,
max_batch_size: 8,
target_latency_ms: 1000.0, adaptation_rate: 1.0,
});
for _ in 0..100 {
let sz = batcher.recommended_batch_size();
batcher.update_latency(0.001, sz);
}
assert!(batcher.recommended_batch_size() <= 8);
}
#[test]
fn test_batch_size_does_not_go_below_min() {
let mut batcher = AdaptiveBatcher::new(AdaptiveBatchConfig {
min_batch_size: 4,
max_batch_size: 64,
target_latency_ms: 1.0, adaptation_rate: 1.0,
});
for _ in 0..100 {
let sz = batcher.recommended_batch_size();
batcher.update_latency(99999.0, sz);
}
assert!(batcher.recommended_batch_size() >= 4);
}
#[test]
fn test_create_batches_splits_correctly() {
let mut batcher = AdaptiveBatcher::new(AdaptiveBatchConfig {
min_batch_size: 3,
max_batch_size: 3,
..Default::default()
});
batcher.current_batch_size = 3;
let requests: Vec<PredictionRequest> = (0..7).map(make_request).collect();
let batches = batcher.create_batches(requests);
assert_eq!(batches.len(), 3, "7 items / 3 = 3 batches (3, 3, 1)");
assert_eq!(batches[0].len(), 3);
assert_eq!(batches[1].len(), 3);
assert_eq!(batches[2].len(), 1);
}
#[test]
fn test_create_batches_fewer_than_batch_size() {
let batcher = AdaptiveBatcher::new(AdaptiveBatchConfig {
min_batch_size: 16,
max_batch_size: 64,
..Default::default()
});
let requests: Vec<PredictionRequest> = (0..5).map(make_request).collect();
let batches = batcher.create_batches(requests);
assert_eq!(batches.len(), 1);
assert_eq!(batches[0].len(), 5);
}
#[test]
fn test_create_batches_empty_input() {
let batcher = default_batcher();
let batches = batcher.create_batches(vec![]);
assert!(batches.is_empty());
}
#[test]
fn test_average_latency_ms_no_observations() {
let batcher = default_batcher();
assert_eq!(batcher.average_latency_ms(), 0.0);
}
#[test]
fn test_average_latency_ms_single_observation() {
let mut batcher = default_batcher();
batcher.update_latency(42.0, 1);
assert!((batcher.average_latency_ms() - 42.0).abs() < 1e-9);
}
#[test]
fn test_average_latency_ms_multiple_observations() {
let mut batcher = default_batcher();
batcher.update_latency(10.0, 1);
batcher.update_latency(20.0, 1);
batcher.update_latency(30.0, 1);
assert!((batcher.average_latency_ms() - 20.0).abs() < 1e-9);
}
#[test]
fn test_total_batches_and_items_tracking() {
let mut batcher = default_batcher();
batcher.update_latency(50.0, 8);
batcher.update_latency(50.0, 4);
assert_eq!(batcher.total_batches(), 2);
assert_eq!(batcher.total_items(), 12);
}
#[test]
fn test_config_validation_invalid_min_batch() {
let config = AdaptiveBatchConfig {
min_batch_size: 0,
..Default::default()
};
assert!(config.validate().is_err());
}
#[test]
fn test_config_validation_max_less_than_min() {
let config = AdaptiveBatchConfig {
min_batch_size: 10,
max_batch_size: 5,
..Default::default()
};
assert!(config.validate().is_err());
}
#[test]
fn test_config_validation_invalid_adaptation_rate() {
let config = AdaptiveBatchConfig {
adaptation_rate: 1.5,
..Default::default()
};
assert!(config.validate().is_err());
}
#[test]
fn test_prediction_result_fields() {
let result = PredictionResult {
id: 42,
outputs: vec![vec![0.9, 0.1]],
output_shapes: vec![vec![2]],
latency_ms: 12.5,
};
assert_eq!(result.id, 42);
assert!((result.latency_ms - 12.5).abs() < 1e-9);
}
}