use crate::{Dialect, IR, OpMap};
use super::{Ready, Retired, Schedule, Selected};
pub trait ForwardSimulator {
type Dialect: Dialect;
fn select(&mut self, ready: impl Iterator<Item = Ready>)
-> impl Iterator<Item = Selected> + '_;
fn advance(&mut self) -> impl Iterator<Item = Retired>;
}
pub trait ForwardScheduler: ForwardSimulator {
fn schedule(&mut self, ir: &IR<Self::Dialect>) -> Schedule;
}
impl<T: ForwardSimulator> ForwardScheduler for T {
fn schedule(&mut self, ir: &IR<Self::Dialect>) -> Schedule {
let mut sched = Schedule::empty();
let mut tracker = Tracker::from_ir(ir);
let mut selected = Vec::new();
let mut retired = Vec::new();
while !tracker.over() {
let selected_iter = self.select(tracker.ready_iter());
selected.extend(selected_iter);
tracker.issue_selected(selected.iter().copied());
sched.issue_selected(selected.iter().copied());
let retired_iter = self.advance();
retired.extend(retired_iter);
tracker.retire(retired.iter().copied());
selected.clear();
retired.clear();
}
sched
}
}
struct Tracker<'i, D: Dialect> {
states: OpMap<State>,
ir: &'i IR<D>,
}
impl<'i, D: Dialect> Tracker<'i, D> {
fn from_ir(ir: &'i IR<D>) -> Self {
let mut states = ir.filled_opmap(State::Retired);
ir.walk_ops_linear()
.for_each(|op| match op.get_predecessors_iter().count() {
0 => {
states.insert(op.get_id(), State::Ready);
}
a => {
states.insert(op.get_id(), State::Locked(a));
}
});
Self { states, ir }
}
fn over(&self) -> bool {
self.states.iter().all(|(_, a)| *a == State::Retired)
}
fn ready_iter(&self) -> impl Iterator<Item = Ready> {
self.states
.iter()
.filter(|(_, s)| **s == State::Ready)
.map(|(opid, _)| Ready(opid))
}
fn issue_selected(&mut self, selected: impl Iterator<Item = Selected>) {
selected.for_each(|sel| {
assert_eq!(*self.states.get(&sel.0).unwrap(), State::Ready);
self.states.insert(sel.0, State::Active);
});
}
fn retire(&mut self, retired: impl Iterator<Item = Retired>) {
for Retired(ret) in retired {
assert_eq!(*self.states.get(&ret).unwrap(), State::Active);
self.states.insert(ret, State::Retired);
for dep in self.ir.get_op(ret).get_users_iter() {
let depid = dep.get_id();
let val = match self.states.get(&depid).unwrap() {
State::Locked(1) => State::Ready,
State::Locked(a) if *a > 1 => State::Locked(a - 1),
_ => unreachable!(),
};
self.states.insert(depid, val);
}
}
}
}
#[derive(Clone, Debug, PartialEq, Eq)]
enum State {
Locked(usize),
Ready,
Active,
Retired,
}