use std::marker::PhantomData;
use std::sync::Arc;
use async_trait::async_trait;
use crate::backend::Backend;
use tokio::sync::Mutex;
use crate::backend::Unsqueezable;
use crate::core::handler::BatchHandler;
use crate::feedforward::core_trait::Feedforward;
use crate::feedforward::queue_item::QueueItem;
use crate::tensor::operations::{
add_sequence_to_outside_of_slot,
slice_tensor_by_batch_dimension
};
pub struct FeedForwardHandler<M, B, O>
{
pub _marker: PhantomData<(B, O)>,
pub model: M
}
#[async_trait]
impl<M, B, O> BatchHandler for FeedForwardHandler<M, B, O>
where B: Backend + Unsqueezable, O: Backend,
M: Feedforward<B, O> + 'static + Sync + Send
{
type Request = QueueItem<B, O>;
type ModelInput = B::Unsqueezed;
type ModelOutput = O;
async fn make_batch_input(&self, model_input: &mut Option<Self::ModelInput>, requests: &[Self::Request]) {
for item in requests.iter() {
match model_input {
None => *model_input = Some(item.input().unsqueeze(0)),
Some(active_tensor) => {
*model_input = Some(add_sequence_to_outside_of_slot(active_tensor, &item.input().clone()));
}
}
}
}
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>>) {
*input = None;
*active_count.lock().await = 0;
let to_send = slice_tensor_by_batch_dimension(output);
for (sender, slice) in batch.drain(..).zip(to_send.iter()) {
sender.sender().send(slice.clone()).unwrap();
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use tokio::test;
use std::sync::Arc;
use tokio::sync::oneshot;
use crate::backend::{Backend};
use crate::backend::mock_tensor::MockTensor;
use crate::feedforward::core_trait::Feedforward;
use crate::feedforward::queue_item::QueueItem;
use async_trait::async_trait;
struct MockModel;
#[async_trait]
impl Feedforward<MockTensor, MockTensor> for MockModel {
async fn forward(&self, input: MockTensor) -> MockTensor {
let shape = input.shape();
let value = input.value * 2;
MockTensor::new(shape, value)
}
}
#[tokio::test]
async fn test_make_batch_input_empty() {
let handler = FeedForwardHandler {
_marker: PhantomData,
model: MockModel,
};
let mut model_input: Option<MockTensor> = None;
let requests: Vec<QueueItem<MockTensor, MockTensor>> = vec![];
handler.make_batch_input(&mut model_input, &requests).await;
assert!(model_input.is_none());
}
#[test]
async fn test_make_batch_input_single() {
let handler = FeedForwardHandler {
_marker: PhantomData,
model: MockModel,
};
let mut model_input: Option<MockTensor> = None;
let (tx, _rx) = oneshot::channel();
let input_tensor = MockTensor::new(vec![3, 4], 5);
let request = QueueItem::new(input_tensor.clone(), tx);
let requests = vec![request];
handler.make_batch_input(&mut model_input, &requests).await;
assert!(model_input.is_some());
let input = model_input.unwrap();
assert_eq!(input.shape(), vec![1, 3, 4]);
assert_eq!(input.value, 5);
}
#[test]
async fn test_make_batch_input_multiple() {
let handler = FeedForwardHandler {
_marker: PhantomData,
model: MockModel,
};
let mut model_input: Option<MockTensor> = None;
let (tx1, _rx1) = oneshot::channel();
let input_tensor1 = MockTensor::new(vec![3, 4], 5);
let request1 = QueueItem::new(input_tensor1.clone(), tx1);
let (tx2, _rx2) = oneshot::channel();
let input_tensor2 = MockTensor::new(vec![3, 4], 7);
let request2 = QueueItem::new(input_tensor2.clone(), tx2);
let requests = vec![request1, request2];
handler.make_batch_input(&mut model_input, &requests).await;
assert!(model_input.is_some());
let input = model_input.unwrap();
assert_eq!(input.shape(), vec![2, 3, 4]);
assert_eq!(input.value, 5);
}
#[test]
async fn test_forward() {
let handler = FeedForwardHandler {
_marker: PhantomData,
model: MockModel,
};
let input = MockTensor::new(vec![2, 3, 4], 5);
let output = handler.forward(&input).await;
assert_eq!(output.value, 10);
assert_eq!(output.shape(), vec![2, 3, 4]);
}
#[test]
async fn test_handle_outputs() {
let handler = FeedForwardHandler {
_marker: PhantomData,
model: MockModel,
};
let (tx1, rx1) = oneshot::channel();
let (tx2, rx2) = oneshot::channel();
let mut batch = vec![
QueueItem::new(MockTensor::new(vec![3, 4], 5), tx1),
QueueItem::new(MockTensor::new(vec![3, 4], 7), tx2),
];
let mut input = Some(MockTensor::new(vec![2, 3, 4], 5));
let output = MockTensor::new(vec![2, 3, 4], 10);
let active_count = Arc::new(Mutex::new(2));
handler.handle_outputs(&mut batch, &mut input, output, active_count.clone()).await;
assert!(batch.is_empty());
assert!(input.is_none());
assert_eq!(*active_count.lock().await, 0);
let result1 = rx1.await.unwrap();
let result2 = rx2.await.unwrap();
assert_eq!(result1.value, 10);
assert_eq!(result2.value, 10);
}
}