use std::fmt::Debug;
use std::sync::Mutex;
use std::sync::atomic::{AtomicU64, Ordering};
use std::sync::mpsc::{Receiver, Sender, channel};
use std::thread;
use std::time::Duration;
use crate::internal::*;
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct LaneTable {
taken: Vec<bool>,
}
impl LaneTable {
pub fn new(max_lanes: usize) -> TractResult<LaneTable> {
ensure!(max_lanes > 0, "A laned state needs at least one lane");
Ok(LaneTable { taken: vec![false; max_lanes] })
}
pub fn max_lanes(&self) -> usize {
self.taken.len()
}
pub fn taken(&self) -> usize {
self.taken.iter().filter(|t| **t).count()
}
pub fn take(&mut self) -> Option<LaneId> {
let lane = self.taken.iter().position(|t| !t)?;
self.taken[lane] = true;
Some(LaneId(lane))
}
pub fn give_back(&mut self, lane: LaneId) -> TractResult<()> {
ensure!(self.is_taken(lane), "Lane {} is not taken, so it can not be given back", lane.0);
self.taken[lane.0] = false;
Ok(())
}
pub fn is_taken(&self, lane: LaneId) -> bool {
self.taken.get(lane.0).copied().unwrap_or(false)
}
pub fn seat(&self, lanes: impl IntoIterator<Item = LaneId>) -> TractResult<Seating> {
let lanes: Vec<LaneId> = lanes.into_iter().collect();
for lane in &lanes {
ensure!(self.is_taken(*lane), "Seating lane {}, which no stream took", lane.0);
}
Seating::new(self.max_lanes(), lanes)
}
}
#[cfg(test)]
mod test {
use super::*;
#[test]
fn takes_the_lowest_free_lane() -> TractResult<()> {
let mut table = LaneTable::new(3)?;
assert_eq!(table.take(), Some(LaneId(0)));
assert_eq!(table.take(), Some(LaneId(1)));
table.give_back(LaneId(0))?;
assert_eq!(table.take(), Some(LaneId(0)));
assert_eq!(table.taken(), 2);
Ok(())
}
#[test]
fn runs_out_of_lanes() -> TractResult<()> {
let mut table = LaneTable::new(1)?;
assert_eq!(table.take(), Some(LaneId(0)));
assert_eq!(table.take(), None);
Ok(())
}
#[test]
fn gives_back_a_taken_lane_only() -> TractResult<()> {
let mut table = LaneTable::new(2)?;
assert!(table.give_back(LaneId(0)).is_err());
table.take();
table.give_back(LaneId(0))?;
assert!(table.give_back(LaneId(0)).is_err());
assert!(table.give_back(LaneId(7)).is_err());
Ok(())
}
#[test]
fn seats_taken_lanes_in_order() -> TractResult<()> {
let mut table = LaneTable::new(4)?;
table.take();
table.take();
table.take();
table.give_back(LaneId(1))?;
let seating = table.seat([LaneId(2), LaneId(0)])?;
assert_eq!(seating.max_lanes(), 4);
assert_eq!(seating.occupancy(), 2);
assert_eq!(seating.address(0), (Some(0), Some(2)));
assert_eq!(seating.address(1), (Some(1), Some(0)));
assert!(table.seat([LaneId(0), LaneId(1)]).is_err());
assert!(table.seat([LaneId(0), LaneId(0)]).is_err());
Ok(())
}
}
crate::declare_knob!(
TRACT_MAX_SEATS,
usize,
256,
"Most streams a laned runtime serves in one turn, clamped to the state's lanes."
);
crate::declare_knob!(
TRACT_TURN_LINGER_US,
usize,
0,
"How long a laned runtime waits for more streams once one is ready to run."
);
#[derive(Clone)]
pub struct LanedRunnable {
shared: Arc<Shared>,
}
struct Shared {
requests: Mutex<Sender<Request>>,
inner: Arc<dyn Runnable>,
model: Option<Arc<TypedModel>>,
plan: Option<Arc<TypedSimplePlan>>,
batch: Symbol,
max_lanes: usize,
counts: Arc<Counts>,
}
#[derive(Debug, Default)]
struct Counts {
turns: AtomicU64,
seats: AtomicU64,
}
impl LanedRunnable {
pub fn wrap(inner: Arc<dyn Runnable>, max_lanes: usize) -> TractResult<LanedRunnable> {
let model = inner.typed_model().cloned();
let plan = inner.typed_plan().cloned();
let mut symbols: Vec<Symbol> = vec![];
let mut batch_in: Vec<bool> = vec![];
for ix in 0..inner.input_count() {
let symbol = batch_symbol(inner.input_fact(ix)?);
batch_in.push(symbol.is_some());
symbols.extend(symbol);
}
let mut batch_out: Vec<bool> = vec![];
for ix in 0..inner.output_count() {
let symbol = batch_symbol(inner.output_fact(ix)?);
batch_out.push(symbol.is_some());
symbols.extend(symbol);
}
symbols.sort();
symbols.dedup();
ensure!(
symbols.len() == 1,
"A laned model carries one batch symbol on axis 0, this one carries {symbols:?}"
);
ensure!(batch_out.iter().any(|b| *b), "A laned model must batch one output at least");
let batch = symbols.remove(0);
let counts = Arc::new(Counts::default());
let max_seats = TRACT_MAX_SEATS.get().min(max_lanes);
let linger = Duration::from_micros(TRACT_TURN_LINGER_US.get() as u64);
let (requests, queue) = channel::<Request>();
let (spawned, ready) = channel::<TractResult<()>>();
let worker_counts = counts.clone();
let worker_inner = inner.clone();
thread::Builder::new().name("tract-lanes".into()).spawn(move || {
let mut state = match worker_inner.spawn().and_then(|mut state| {
let lanes: Vec<LaneId> = (0..max_lanes).map(LaneId).collect();
state.reset_lanes(&lanes).context("Preparing a laned model")?;
Ok(state)
}) {
Ok(state) => {
let _ = spawned.send(Ok(()));
state
}
Err(e) => {
let _ = spawned.send(Err(e));
return;
}
};
worker(
&mut *state,
queue,
Table { batch_in, batch_out, max_seats, linger, max_lanes, counts: worker_counts },
);
})?;
ready.recv().map_err(|_| format_err!("The laned worker died spawning the state"))??;
Ok(LanedRunnable {
shared: Arc::new(Shared {
requests: Mutex::new(requests),
inner,
model,
plan,
batch,
max_lanes,
counts,
}),
})
}
pub fn max_lanes(&self) -> usize {
self.shared.max_lanes
}
pub fn inner(&self) -> &Arc<dyn Runnable> {
&self.shared.inner
}
pub fn batch_symbol(&self) -> &Symbol {
&self.shared.batch
}
pub fn turns_and_seats(&self) -> (u64, u64) {
(
self.shared.counts.turns.load(Ordering::Relaxed),
self.shared.counts.seats.load(Ordering::Relaxed),
)
}
fn request(&self) -> TractResult<Sender<Request>> {
Ok(self.shared.requests.lock().map_err(|_| format_err!("Poisoned laned sender"))?.clone())
}
}
fn batch_symbol(fact: &TypedFact) -> Option<Symbol> {
match fact.shape.dims().first() {
Some(TDim::Sym(sym)) => Some(sym.clone()),
_ => None,
}
}
impl Debug for LanedRunnable {
fn fmt(&self, f: &mut std::fmt::Formatter) -> std::fmt::Result {
write!(f, "LanedRunnable({} lanes)", self.shared.max_lanes)
}
}
impl Runnable for LanedRunnable {
fn spawn(&self) -> TractResult<Box<dyn State>> {
let requests = self.request()?;
let (taken, lane) = channel();
requests.send(Request::Take(taken)).map_err(|_| format_err!("The laned worker is gone"))?;
let lane = lane.recv().map_err(|_| format_err!("The laned worker dropped a lane"))??;
Ok(Box::new(SessionHandle {
lease: Arc::new(Lease { lane, requests }),
runnable: self.clone(),
}))
}
fn typed_plan(&self) -> Option<&Arc<TypedSimplePlan>> {
self.shared.plan.as_ref()
}
fn typed_model(&self) -> Option<&Arc<TypedModel>> {
self.shared.model.as_ref()
}
}
#[derive(Clone, Debug)]
pub struct SessionHandle {
lease: Arc<Lease>,
runnable: LanedRunnable,
}
#[derive(Debug)]
struct Lease {
lane: LaneId,
requests: Sender<Request>,
}
impl Drop for Lease {
fn drop(&mut self) {
let _ = self.requests.send(Request::GiveBack(self.lane));
}
}
impl State for SessionHandle {
fn run(&mut self, inputs: TVec<TValue>) -> TractResult<TVec<TValue>> {
let (done, outputs) = channel();
self.lease
.requests
.send(Request::Turn(Turn { lane: self.lease.lane, inputs, done }))
.map_err(|_| format_err!("The laned worker is gone"))?;
outputs.recv().map_err(|_| format_err!("The laned worker dropped a turn"))?
}
fn runnable(&self) -> &dyn Runnable {
&self.runnable
}
}
enum Request {
Take(Sender<TractResult<LaneId>>),
GiveBack(LaneId),
Turn(Turn),
}
struct Turn {
lane: LaneId,
inputs: TVec<TValue>,
done: Sender<TractResult<TVec<TValue>>>,
}
struct Table {
batch_in: Vec<bool>,
batch_out: Vec<bool>,
max_seats: usize,
linger: Duration,
max_lanes: usize,
counts: Arc<Counts>,
}
fn worker(state: &mut dyn State, queue: Receiver<Request>, table: Table) {
let mut lanes = match LaneTable::new(table.max_lanes) {
Ok(lanes) => lanes,
Err(_) => return,
};
let mut queued: Vec<Turn> = vec![];
loop {
if queued.is_empty() {
match queue.recv() {
Ok(request) => serve(state, &mut lanes, &mut queued, request),
Err(_) => return,
}
if !table.linger.is_zero() {
thread::sleep(table.linger);
}
}
while let Ok(request) = queue.try_recv() {
serve(state, &mut lanes, &mut queued, request);
}
let mut seated: Vec<Turn> = vec![];
let mut waiting: Vec<Turn> = vec![];
for turn in queued.drain(..) {
if seated.len() < table.max_seats && !seated.iter().any(|s| s.lane == turn.lane) {
seated.push(turn);
} else {
waiting.push(turn);
}
}
queued = waiting;
if seated.is_empty() {
continue;
}
table.counts.turns.fetch_add(1, Ordering::Relaxed);
table.counts.seats.fetch_add(seated.len() as u64, Ordering::Relaxed);
match run_turn(state, &lanes, &seated, &table) {
Ok(per_seat) => {
for (turn, outputs) in seated.into_iter().zip(per_seat) {
let _ = turn.done.send(Ok(outputs));
}
}
Err(e) => {
let e = format!("{e:#}");
for turn in seated {
let _ = turn.done.send(Err(format_err!("Laned turn failed: {e}")));
}
}
}
}
}
fn serve(state: &mut dyn State, lanes: &mut LaneTable, queued: &mut Vec<Turn>, request: Request) {
match request {
Request::Take(taken) => {
let lane = lanes.take().ok_or_else(|| {
format_err!("Every one of the {} lanes is taken", lanes.max_lanes())
});
let lane = lane.and_then(|lane| {
state.reset_lanes(&[lane]).map(|_| lane).inspect_err(|_| {
let _ = lanes.give_back(lane);
})
});
let _ = taken.send(lane);
}
Request::GiveBack(lane) => {
let _ = lanes.give_back(lane);
}
Request::Turn(turn) => queued.push(turn),
}
}
fn run_turn(
state: &mut dyn State,
lanes: &LaneTable,
seated: &[Turn],
table: &Table,
) -> TractResult<Vec<TVec<TValue>>> {
let seating = lanes.seat(seated.iter().map(|turn| turn.lane))?;
let mut batched: TVec<TValue> = tvec!();
for turn in seated {
ensure!(
turn.inputs.len() == table.batch_in.len(),
"A turn feeds {} inputs, the model takes {}",
turn.inputs.len(),
table.batch_in.len()
);
}
for (ix, is_batched) in table.batch_in.iter().enumerate() {
if *is_batched {
let rows: TVec<&Tensor> = seated.iter().map(|turn| &*turn.inputs[ix]).collect();
for row in &rows {
ensure!(
row.rank() > 0 && row.shape()[0] == 1,
"A stream feeds one row per turn, input {ix} carries {:?}",
row.shape()
);
}
batched.push(Tensor::stack_tensors(0, &rows)?.into_tvalue());
} else {
let shared = &seated[0].inputs[ix];
for (seat, turn) in seated.iter().enumerate().skip(1) {
ensure!(
turn.inputs[ix] == *shared,
"Input {ix} carries no batch axis, so one value of it serves the whole \
turn, and seats 0 and {seat} feed it different ones"
);
}
batched.push(shared.clone());
}
}
state.seat(seating)?;
let outputs = state.run(batched)?;
let mut per_seat: Vec<TVec<TValue>> = seated.iter().map(|_| tvec!()).collect();
for (ix, output) in outputs.into_iter().enumerate() {
if table.batch_out.get(ix).copied().unwrap_or(false) {
ensure!(
output.shape()[0] == seated.len(),
"The turn seats {} streams, output {ix} carries {:?}",
seated.len(),
output.shape()
);
for (seat, outputs) in per_seat.iter_mut().enumerate() {
outputs.push(output.slice(0, seat, seat + 1)?.into_tvalue());
}
} else {
for outputs in per_seat.iter_mut() {
outputs.push(output.clone());
}
}
}
Ok(per_seat)
}
#[cfg(all(test, not(target_family = "wasm")))]
mod laned_test {
use super::*;
use crate::ops::math::{add, mul};
fn doubler(max_lanes: usize) -> TractResult<LanedRunnable> {
let mut model = TypedModel::default();
let batch = model.symbols.sym("B");
let input = model.add_source("input", f32::fact(dims!(batch, 3)))?;
let two = model.add_const("two", tensor2(&[[2f32]]))?;
let doubled = model.wire_node("doubled", mul(), &[input, two])?;
model.select_output_outlets(&doubled)?;
let inner = DefaultRuntime.prepare(model)?;
LanedRunnable::wrap(inner.into(), max_lanes)
}
fn turn(handle: &mut Box<dyn State>, stream: usize, turn: usize) -> TractResult<()> {
let input = tensor2(&[[stream as f32, turn as f32, 1.]]);
let output = handle.run(tvec!(input.into_tvalue()))?;
assert_eq!(&*output[0], &tensor2(&[[2. * stream as f32, 2. * turn as f32, 2.]]));
Ok(())
}
static LINGER: Mutex<()> = Mutex::new(());
fn spawn_once_free(runnable: &LanedRunnable) -> TractResult<Box<dyn State>> {
for _ in 0..100 {
if let Ok(handle) = runnable.spawn() {
return Ok(handle);
}
std::thread::sleep(Duration::from_millis(10));
}
runnable.spawn()
}
#[test]
fn one_stream_at_a_time() -> TractResult<()> {
let runnable = doubler(2)?;
let mut handle = runnable.spawn()?;
for t in 0..4 {
turn(&mut handle, 0, t)?;
}
Ok(())
}
#[test]
fn every_stream_gets_its_own_row() -> TractResult<()> {
let runnable = doubler(8)?;
let streams: Vec<_> = (0..8)
.map(|stream| {
let runnable = runnable.clone();
std::thread::spawn(move || -> TractResult<()> {
let mut handle = runnable.spawn()?;
for t in 0..32 {
turn(&mut handle, stream, t)?;
}
Ok(())
})
})
.collect();
for stream in streams {
stream.join().unwrap()?;
}
Ok(())
}
#[test]
fn a_turn_seats_the_streams_that_are_ready() -> TractResult<()> {
let _linger = LINGER.lock().unwrap_or_else(|e| e.into_inner());
TRACT_TURN_LINGER_US.set(20_000);
let runnable = doubler(8);
TRACT_TURN_LINGER_US.clear();
let runnable = runnable?;
let streams: Vec<_> = (0..8)
.map(|stream| {
let runnable = runnable.clone();
std::thread::spawn(move || -> TractResult<()> {
let mut handle = runnable.spawn()?;
for t in 0..4 {
turn(&mut handle, stream, t)?;
}
Ok(())
})
})
.collect();
for stream in streams {
stream.join().unwrap()?;
}
let (turns, seats) = runnable.turns_and_seats();
assert!(seats > turns, "{seats} seats over {turns} turns, none of them shared");
Ok(())
}
fn biased(max_lanes: usize) -> TractResult<LanedRunnable> {
let mut model = TypedModel::default();
let batch = model.symbols.sym("B");
let input = model.add_source("input", f32::fact(dims!(batch, 3)))?;
let bias = model.add_source("bias", f32::fact(dims!(1, 1)))?;
let two = model.add_const("two", tensor2(&[[2f32]]))?;
let doubled = model.wire_node("doubled", mul(), &[input, two])?;
let biased = model.wire_node("biased", add(), &[doubled[0], bias])?;
model.select_output_outlets(&biased)?;
let inner = DefaultRuntime.prepare(model)?;
LanedRunnable::wrap(inner.into(), max_lanes)
}
fn biased_turns(runnable: &LanedRunnable, biases: &[f32]) -> TractResult<Vec<TractResult<()>>> {
let handles: Vec<Box<dyn State>> =
biases.iter().map(|_| runnable.spawn()).collect::<TractResult<_>>()?;
let streams: Vec<_> = handles
.into_iter()
.zip(biases.iter().copied())
.map(|(mut handle, bias)| {
std::thread::spawn(move || -> TractResult<()> {
handle.run(tvec!(
tensor2(&[[1f32, 2., 3.]]).into_tvalue(),
tensor2(&[[bias]]).into_tvalue()
))?;
Ok(())
})
})
.collect();
Ok(streams.into_iter().map(|stream| stream.join().unwrap()).collect())
}
#[test]
fn seats_agreeing_on_a_shared_input_share_a_turn() -> TractResult<()> {
let _linger = LINGER.lock().unwrap_or_else(|e| e.into_inner());
TRACT_TURN_LINGER_US.set(100_000);
let runnable = biased(2);
TRACT_TURN_LINGER_US.clear();
let runnable = runnable?;
let served = biased_turns(&runnable, &[7., 7.])?;
assert!(served.iter().all(|s| s.is_ok()), "{served:?}");
assert_eq!(runnable.turns_and_seats(), (1, 2));
Ok(())
}
#[test]
fn seats_disagreeing_on_a_shared_input_fail_the_turn() -> TractResult<()> {
let _linger = LINGER.lock().unwrap_or_else(|e| e.into_inner());
TRACT_TURN_LINGER_US.set(100_000);
let runnable = biased(2);
TRACT_TURN_LINGER_US.clear();
let runnable = runnable?;
let served = biased_turns(&runnable, &[7., 8.])?;
assert_eq!(runnable.turns_and_seats(), (1, 2));
for stream in &served {
let error = format!("{:#}", stream.as_ref().unwrap_err());
assert!(error.contains("seats 0 and 1 feed it different ones"), "{error}");
}
Ok(())
}
#[test]
fn a_dropped_stream_gives_its_lane_back() -> TractResult<()> {
let runnable = doubler(1)?;
let mut handle = runnable.spawn()?;
turn(&mut handle, 0, 0)?;
assert!(runnable.spawn().is_err());
let clone = dyn_clone::clone_box(&*handle);
drop(handle);
assert!(runnable.spawn().is_err());
drop(clone);
let mut handle = spawn_once_free(&runnable)?;
turn(&mut handle, 1, 0)?;
Ok(())
}
}