use std::sync::Arc;
use async_trait::async_trait;
use tokio::sync::{mpsc, Mutex};
use crate::autoregressive::handler::AutoregressiveHandler;
use crate::autoregressive::queue_item::QueueItem;
use super::item_stream::ItemStream;
use crate::communication::Pill;
use crate::backend::{Backend, Unsqueezable};
use crate::core::batch::batch_inference_loop;
use crate::core::worker::BatchWorkerHandle;
use super::{Autoregressive, AutoregressiveBatcher};
pub struct AutoregressiveBatchInference<B, const S: usize>
where
B: Backend + Unsqueezable
{
waiting_requests: Arc<Mutex<Vec<QueueItem<B>>>>,
handle: BatchWorkerHandle
}
impl<B, const S: usize> AutoregressiveBatchInference<B, S>
where
B: Backend + Unsqueezable,
{
pub fn new<M>(
model: M,
stop_token: &B,
padding_token: &B,
) -> Self
where
M: Autoregressive<B> + Send + Sync + 'static,
{
let padding_shape = padding_token.shape();
assert!(
!padding_shape.is_empty(),
"padding token must have rank 1 or higher"
);
assert_eq!(
padding_shape[0], 1,
"first dimension of padding token must be 1"
);
let waiting_requests = Arc::new(Mutex::new(vec![]));
let pill = Pill::new();
let worker_handle = BatchWorkerHandle::new({
let waiting_requests = waiting_requests.clone();
let padding_token_clone = padding_token.clone();
let stop_token_clone = stop_token.clone();
move |running, work_notifier| {
tokio::spawn(async move {
#[allow(unused_variables)]
let moved_pill = pill;
let inference_handler = AutoregressiveHandler {
model,
padding_token: padding_token_clone,
stop_token: stop_token_clone,
};
batch_inference_loop::<AutoregressiveHandler<M, B>, S>(
&inference_handler,
running,
work_notifier,
waiting_requests,
).await;
})
}
});
Self {
waiting_requests,
handle: worker_handle,
}
}
}
#[async_trait]
impl<B, const S: usize> AutoregressiveBatcher<B, B> for AutoregressiveBatchInference<B, S>
where
B: Backend + Unsqueezable,
{
async fn run(&self, item: B) -> ItemStream<B> {
let (tx, rx) = mpsc::unbounded_channel();
let sequence_length = item.shape()[0];
let queue_item = QueueItem::new(
item,
sequence_length,
tx,
);
{
let mut queue = self.waiting_requests.lock().await;
queue.push(queue_item);
}
self.handle.notify();
ItemStream::new(rx)
}
}
#[cfg(test)]
mod tests {
use super::*;
use tokio::test;
use async_trait::async_trait;
use futures::StreamExt;
use tokio::time::{sleep, Duration};
use crate::backend::mock_tensor::MockTensor;
struct MockModel {
stop_after: usize,
}
#[async_trait]
impl Autoregressive<MockTensor> for MockModel {
async fn forward(&self, input: MockTensor) -> MockTensor {
let mut output_shape = input.shape();
output_shape.remove(0);
let batch_size = input.shape()[0];
let seq_len = input.shape()[1];
let output_value = if seq_len >= self.stop_after {
99
} else {
42
};
MockTensor::new(vec![batch_size], output_value)
}
}
#[test]
async fn test_inference_creation() {
let model = MockModel { stop_after: 5 };
let stop_token = MockTensor::new(vec![1, 10], 99);
let padding_token = MockTensor::new(vec![1, 10], 0);
let inference = AutoregressiveBatchInference::<MockTensor, 4>::new(
model,
&stop_token,
&padding_token
);
let queue = inference.waiting_requests.lock().await;
assert_eq!(queue.len(), 0);
}
#[test]
#[should_panic(expected = "padding token must have rank 1 or higher")]
async fn test_empty_padding_shape_panics() {
let model = MockModel { stop_after: 5 };
let stop_token = MockTensor::new(vec![1, 10], 99);
let padding_token = MockTensor::new(vec![], 0);
let _inference = AutoregressiveBatchInference::<MockTensor, 4>::new(
model,
&stop_token,
&padding_token
);
}
#[test]
#[should_panic(expected = "first dimension of padding token must be 1")]
async fn test_invalid_padding_shape_panics() {
let model = MockModel { stop_after: 5 };
let stop_token = MockTensor::new(vec![1, 10], 99);
let padding_token = MockTensor::new(vec![2, 10], 0);
let _inference = AutoregressiveBatchInference::<MockTensor, 4>::new(
model,
&stop_token,
&padding_token
);
}
#[test]
async fn test_run_adds_request_to_queue() {
let model = MockModel { stop_after: 5 };
let stop_token = MockTensor::new(vec![1, 10], 99);
let padding_token = MockTensor::new(vec![1, 10], 0);
let inference = AutoregressiveBatchInference::<MockTensor, 4>::new(
model,
&stop_token,
&padding_token
);
let input = MockTensor::new(vec![3, 10], 42);
let _stream = inference.run(input).await;
let queue = inference.waiting_requests.lock().await;
assert_eq!(queue.len(), 1);
assert_eq!(queue[0].len(), 3); }
#[test]
async fn test_end_to_end_generation() {
let model = MockModel { stop_after: 3 };
let stop_token = MockTensor::new(vec![1, 10], 99);
let padding_token = MockTensor::new(vec![1, 10], 0);
let inference = AutoregressiveBatchInference::<MockTensor, 2>::new(
model,
&stop_token,
&padding_token
);
let input = MockTensor::new(vec![1, 10], 42);
let stream = inference.run(input).await;
sleep(Duration::from_millis(100)).await;
let mut received_tokens = 0;
let mut stream_clone = stream;
for _ in 0..5 {
match tokio::time::timeout(Duration::from_millis(100), stream_clone.next()).await {
Ok(Some(_token)) => {
received_tokens += 1;
}
Ok(None) => break, Err(_) => break, }
}
assert!(received_tokens > 0, "Should receive at least one token");
}
}