use super::queue_item::QueueItem;
use crate::autoregressive::Autoregressive;
use crate::backend::Backend;
use crate::backend::Unsqueezable;
use crate::core::handler::BatchHandler;
use crate::tensor::operations::{
add_sequence_to_outside_of_slot,
concat_output,
pad_all_sequences,
pad_single_sequence,
pop_sequence_from_slot,
slice_tensor_by_batch_dimension,
trim_sequence,
where_equals_stop_token
};
use async_trait::async_trait;
use std::collections::HashSet;
use std::sync::Arc;
use tokio::sync::Mutex;
pub struct AutoregressiveHandler<M, B> {
pub model: M,
pub padding_token: B,
pub stop_token: B
}
#[async_trait]
impl<M, B> BatchHandler for AutoregressiveHandler<M, B>
where
B: Backend + Unsqueezable,
M: Autoregressive<B> + 'static + Sync + Send
{
type Request = QueueItem<B>;
type ModelInput = B::Unsqueezed;
type ModelOutput = B;
async fn make_batch_input(&self, model_input: &mut Option<Self::ModelInput>, requests: &[Self::Request]) {
for item in requests.iter() {
let request_tensor = item.input();
if model_input.is_none() {
*model_input = Some(request_tensor.unsqueeze(0));
continue;
}
if let Some(active_tensor) = model_input {
let sequence_dims = request_tensor.shape();
let mut request_tensor_copy = request_tensor.clone();
let sequence_length = sequence_dims[0];
let active_dims = active_tensor.shape();
let batch_length = active_dims[1];
if sequence_length > batch_length {
let padding_amount = sequence_length - batch_length;
*active_tensor = pad_all_sequences(active_tensor, padding_amount, &self.padding_token);
} else if batch_length > sequence_length {
let padding_amount = batch_length - sequence_length;
request_tensor_copy = pad_single_sequence(&request_tensor_copy, padding_amount, &self.padding_token);
}
*active_tensor = add_sequence_to_outside_of_slot(active_tensor, &request_tensor_copy);
}
}
}
async fn forward(&self, model_input: &Self::ModelInput) -> Self::ModelOutput {
self.model.forward(model_input.clone()).await
}
async fn handle_outputs(
&self,
batch: &mut Vec<Self::Request>,
input: &mut Option<Self::ModelInput>,
output: Self::ModelOutput,
active_count: Arc<Mutex<usize>>,
) {
if let Some(input_val) = input.as_mut() {
*input_val = concat_output(input_val, &output);
}
for batch_item in batch.iter_mut() {
batch_item.increment_sequence_length(1);
}
let split_outputs = slice_tensor_by_batch_dimension(output.clone());
for (sender, slice) in batch.iter().zip(split_outputs.iter()) {
sender.sender().send(slice.clone()).unwrap();
}
let completed_indices: HashSet<usize> = where_equals_stop_token(&output, &self.stop_token)
.into_iter()
.collect();
let num_completed = completed_indices.len();
if num_completed == 0 {
return;
}
let mut idx = 0;
batch.retain(|_| {
let keep = !completed_indices.contains(&idx);
idx += 1;
keep
});
let mut sorted_indices: Vec<usize> = completed_indices.into_iter().collect();
sorted_indices.sort_by(|a, b| b.cmp(a));
for &idx in &sorted_indices {
*input = pop_sequence_from_slot(input, idx);
}
let max_sequence_length = QueueItem::max_seq_len_for_batch_items(batch);
*input = trim_sequence(input, max_sequence_length);
{
let mut sequence_count = active_count.lock().await;
*sequence_count = sequence_count.saturating_sub(num_completed);
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::backend::{
mock_tensor::MockTensor
,
Backend
};
use async_trait::async_trait;
use std::sync::Arc;
use tokio::sync::mpsc;
use tokio::test;
struct MockAutoregressive;
#[async_trait]
impl Autoregressive<MockTensor> for MockAutoregressive {
async fn forward(&self, input: MockTensor) -> MockTensor {
let mut output_shape = input.shape();
output_shape.remove(0);
MockTensor::new(output_shape, 42)
}
}
#[test]
async fn test_make_batch_input_creates_new_batch() {
let handler = AutoregressiveHandler {
model: MockAutoregressive,
padding_token: MockTensor::new(vec![10], 0),
stop_token: MockTensor::new(vec![10], 1)
};
let (tx, _rx) = mpsc::unbounded_channel();
let request = QueueItem::new(MockTensor::new(vec![10], 42), 10, tx);
let mut model_input = None;
handler.make_batch_input(&mut model_input, &[request]).await;
assert!(model_input.is_some());
if let Some(input) = &model_input {
let shape = input.shape();
assert_eq!(shape.len(), 2); assert_eq!(shape[0], 1); assert_eq!(shape[1], 10); }
}
#[test]
async fn test_make_batch_input_adds_to_existing_batch() {
let handler = AutoregressiveHandler {
model: MockAutoregressive,
padding_token: MockTensor::new(vec![10], 0),
stop_token: MockTensor::new(vec![10], 1)
};
let mut model_input = Some(MockTensor::new(vec![1, 10], 42));
let (tx, _rx) = mpsc::unbounded_channel();
let request = QueueItem::new(MockTensor::new(vec![10], 42), 10, tx);
handler.make_batch_input(&mut model_input, &[request]).await;
assert!(model_input.is_some());
if let Some(input) = &model_input {
let shape = input.shape();
assert_eq!(shape[0], 2); }
}
#[test]
async fn test_forward_returns_expected_output() {
let handler = AutoregressiveHandler {
model: MockAutoregressive,
padding_token: MockTensor::new(vec![10], 0),
stop_token: MockTensor::new(vec![10], 1)
};
let model_input = MockTensor::new(vec![2, 10], 42);
let output = handler.forward(&model_input).await;
let shape = output.shape();
assert_eq!(shape.len(), 1); assert_eq!(shape[0], 10); assert_eq!(output.value, 42); }
#[test]
async fn test_handle_outputs_processes_completions() {
let handler = AutoregressiveHandler {
model: MockAutoregressive,
padding_token: MockTensor::new(vec![10], 0),
stop_token: MockTensor::new(vec![10], 42) };
let active_count = Arc::new(Mutex::new(2));
let (tx1, mut rx1) = mpsc::unbounded_channel();
let (tx2, mut rx2) = mpsc::unbounded_channel();
let mut batch = vec![
QueueItem::new(MockTensor::new(vec![10], 42), 10, tx1),
QueueItem::new(MockTensor::new(vec![10], 42), 10, tx2)
];
let mut model_input = Some(MockTensor::new(vec![2, 10], 42));
let output = MockTensor::new(vec![2], 42);
handler.handle_outputs(&mut batch, &mut model_input, output, active_count.clone()).await;
assert!(rx1.try_recv().is_ok());
assert!(rx2.try_recv().is_ok());
assert_eq!(batch.len(), 0);
let count = *active_count.lock().await;
assert_eq!(count, 0);
}
#[test]
async fn test_handle_outputs_continues_generation() {
let handler = AutoregressiveHandler {
model: MockAutoregressive,
padding_token: MockTensor::new(vec![10], 0),
stop_token: MockTensor::new(vec![10], 99) };
let active_count = Arc::new(Mutex::new(2));
let (tx1, mut rx1) = mpsc::unbounded_channel();
let (tx2, mut rx2) = mpsc::unbounded_channel();
let mut batch = vec![
QueueItem::new(MockTensor::new(vec![10], 42), 10, tx1),
QueueItem::new(MockTensor::new(vec![10], 42), 10, tx2)
];
let mut model_input = Some(MockTensor::new(vec![2, 10], 42));
let output = MockTensor::new(vec![2], 42);
handler.handle_outputs(&mut batch, &mut model_input, output, active_count.clone()).await;
assert!(rx1.try_recv().is_ok());
assert!(rx2.try_recv().is_ok());
assert_eq!(batch.len(), 2);
assert_eq!(batch[0].len(), 11);
assert_eq!(batch[1].len(), 11);
let count = *active_count.lock().await;
assert_eq!(count, 2);
}
#[test]
async fn test_end_to_end_processing() {
let mut handler = AutoregressiveHandler {
model: MockAutoregressive,
padding_token: MockTensor::new(vec![10], 0),
stop_token: MockTensor::new(vec![10], 99) };
let active_count = Arc::new(Mutex::new(1));
let (tx, mut rx) = mpsc::unbounded_channel();
let mut batch = vec![QueueItem::new(MockTensor::new(vec![10], 21), 10, tx)];
let mut model_input = None;
handler.make_batch_input(&mut model_input, &batch).await;
assert!(model_input.is_some());
let output = handler.forward(model_input.as_ref().unwrap()).await;
handler.handle_outputs(&mut batch, &mut model_input, output, active_count.clone()).await;
assert!(rx.try_recv().is_ok());
assert_eq!(batch.len(), 1);
{
let count = *active_count.lock().await;
assert_eq!(count, 1);
}
handler.stop_token = MockTensor::new(vec![10], 42);
let output = handler.forward(model_input.as_ref().unwrap()).await;
handler.handle_outputs(&mut batch, &mut model_input, output, active_count.clone()).await;
assert!(rx.try_recv().is_ok());
assert_eq!(batch.len(), 0);
let count = *active_count.lock().await;
assert_eq!(count, 0);
}
}