use std::marker::PhantomData;
use std::sync::Arc;
use tokio::sync::{oneshot, Mutex};
use crate::backend::{Backend, Unsqueezable};
use super::queue_item::QueueItem;
use async_trait::async_trait;
use oneshot::channel;
use crate::communication::Pill;
use crate::core::batch::batch_inference_loop;
use crate::core::worker::BatchWorkerHandle;
use crate::feedforward::core_trait::{Feedforward, FeedforwardBatcher};
use crate::feedforward::handler::FeedForwardHandler;
use crate::feedforward::item::Item;
pub struct FeedforwardBatchInference<B, O, const S: usize>
{
waiting_requests: Arc<Mutex<Vec<QueueItem<B, O>>>>,
handle: BatchWorkerHandle
}
impl <B, O, const S: usize> FeedforwardBatchInference<B, O, S>
where B: Backend + Unsqueezable, O: Backend
{
pub fn new<M>(
model: M,
) -> Self
where M: Feedforward<B, O> + Send + Sync + 'static,
{
let waiting_requests = Arc::new(Mutex::new(vec![]));
let pill = Pill::new();
let worker_handle = BatchWorkerHandle::new( {
let waiting_requests = waiting_requests.clone();
move |running, notifier| {
tokio::spawn(async move {
#[allow(unused_variables)]
let moved_pill = pill;
let inference_handler = FeedForwardHandler{
_marker: PhantomData,
model
};
batch_inference_loop::<FeedForwardHandler<M, B, O>, S>(
&inference_handler,
running,
notifier,
waiting_requests,
)
.await;
})
}
});
Self {
waiting_requests,
handle: worker_handle,
}
}
}
#[async_trait]
impl <B, O, const S: usize> FeedforwardBatcher<B, O> for FeedforwardBatchInference<B, O, S>
where B: Backend, O: Backend
{
async fn run(&self, item: B) -> Item<O> {
let (tx, rx) = channel();
let queue_item = QueueItem::new(
item,
tx,
);
{
let mut senders = self.waiting_requests.lock().await;
senders.push(queue_item);
}
self.handle.notify();
Item::new(rx)
}
}
#[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;
use crate::core::handler::BatchHandler;
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)
}
}
#[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);
}
}