use std::cmp::Reverse;
use std::collections::BinaryHeap;
pub(crate) struct UnitPlan {
pub(crate) units: Vec<Vec<usize>>,
pub(crate) unit_of: Vec<usize>,
}
pub(crate) fn components(preds: &[Vec<usize>], fusible: &[bool], class: &[u64]) -> Vec<Vec<usize>> {
lumped_components(preds, fusible, class, &|_| false)
}
fn lumped_components(
preds: &[Vec<usize>],
fusible: &[bool],
class: &[u64],
lump: &dyn Fn(u64) -> bool,
) -> Vec<Vec<usize>> {
let n = preds.len();
let mut parent: Vec<usize> = (0..n).collect();
fn find(parent: &mut [usize], mut i: usize) -> usize {
while parent[i] != i {
parent[i] = parent[parent[i]];
i = parent[i];
}
i
}
for i in 0..n {
if !fusible[i] {
continue;
}
for &j in &preds[i] {
if fusible[j] && class[j] == class[i] {
let (a, b) = (find(&mut parent, i), find(&mut parent, j));
parent[a] = b;
}
}
}
let mut first_of: std::collections::HashMap<u64, usize> = Default::default();
for i in (0..n).filter(|&i| fusible[i] && lump(class[i])) {
let first = *first_of.entry(class[i]).or_insert(i);
let (a, b) = (find(&mut parent, i), find(&mut parent, first));
parent[a] = b;
}
let mut by_root: std::collections::BTreeMap<usize, Vec<usize>> = Default::default();
for i in (0..n).filter(|&i| fusible[i]) {
by_root.entry(find(&mut parent, i)).or_default().push(i);
}
by_root.into_values().collect()
}
pub(crate) fn is_convex(members: &[usize], consumers: &[Vec<usize>]) -> bool {
let n = consumers.len();
let mut is_member = vec![false; n];
for &m in members {
is_member[m] = true;
}
let mut seen = vec![false; n];
let mut stack: Vec<usize> = Vec::new();
for &m in members {
for &c in &consumers[m] {
if !is_member[c] && !seen[c] {
seen[c] = true;
stack.push(c);
}
}
}
while let Some(v) = stack.pop() {
for &c in &consumers[v] {
if is_member[c] {
return false;
}
if !seen[c] {
seen[c] = true;
stack.push(c);
}
}
}
true
}
fn convex_pieces(
members: &[usize],
preds: &[Vec<usize>],
topo: &[usize],
lumped: bool,
) -> Vec<Vec<usize>> {
let n = preds.len();
let mut is_member = vec![false; n];
for &m in members {
is_member[m] = true;
}
let mut reached = vec![false; n];
let mut stage = vec![0u64; n];
for &v in topo {
let mut s = 0;
let mut r = is_member[v];
for &p in &preds[v] {
r |= reached[p];
let returns = is_member[v] && !is_member[p] && reached[p];
s = s.max(stage[p] + returns as u64);
}
stage[v] = s;
reached[v] = r;
}
lumped_components(preds, &is_member, &stage, &|_| lumped)
}
pub(crate) fn plan_units(
preds: &[Vec<usize>],
inputs: &[Vec<usize>],
fusible: &[bool],
class: &[u64],
rank: &[usize],
lump: &dyn Fn(u64) -> bool,
) -> UnitPlan {
let n = preds.len();
let mut consumers: Vec<Vec<usize>> = vec![Vec::new(); n];
for (i, ps) in preds.iter().enumerate() {
for &p in ps {
consumers[p].push(i);
}
}
let mut unit_of = vec![usize::MAX; n];
let mut units: Vec<Vec<usize>> = Vec::new();
let mut by_rank = |mut members: Vec<usize>, units: &mut Vec<Vec<usize>>| {
members.sort_by_key(|&m| rank[m]);
for &m in &members {
unit_of[m] = units.len();
}
units.push(members);
};
let mut at_rank = vec![0usize; n];
for (i, &r) in rank.iter().enumerate() {
at_rank[r] = i;
}
for component in lumped_components(preds, fusible, class, lump) {
if component.len() == 1 || is_convex(&component, &consumers) {
by_rank(component, &mut units);
continue;
}
let lumped = lump(class[component[0]]);
for piece in convex_pieces(&component, preds, &at_rank, lumped) {
by_rank(piece, &mut units);
}
}
for i in (0..n).filter(|&i| !fusible[i]) {
by_rank(vec![i], &mut units);
}
let mut moved = false;
for m in 0..n {
if !fusible[m]
|| !preds[m].is_empty()
|| !consumers[m].is_empty()
|| inputs[m].is_empty()
|| units[unit_of[m]].len() != 1
{
continue;
}
let home = unit_of[m];
let target = units
.iter()
.enumerate()
.filter(|&(t, members)| {
t != home
&& !members.is_empty()
&& fusible[members[0]]
&& class[members[0]] == class[m]
&& members
.iter()
.any(|&x| inputs[x].iter().any(|i| inputs[m].contains(i)))
})
.map(|(t, members)| (rank[members[0]], t))
.min();
if let Some((_, t)) = target {
units[home].clear();
units[t].push(m);
units[t].sort_by_key(|&x| rank[x]);
unit_of[m] = t;
moved = true;
}
}
if moved {
units.retain(|members| !members.is_empty());
for (u, members) in units.iter().enumerate() {
for &m in members {
unit_of[m] = u;
}
}
}
let u = units.len();
let first: Vec<usize> = units.iter().map(|m| rank[m[0]]).collect();
let mut after: Vec<Vec<usize>> = vec![Vec::new(); u];
let mut waiting = vec![0usize; u];
for (to, members) in units.iter().enumerate() {
let mut from: Vec<usize> = members
.iter()
.flat_map(|&m| preds[m].iter().map(|&p| unit_of[p]))
.filter(|&f| f != to)
.collect();
from.sort_unstable();
from.dedup();
waiting[to] = from.len();
for f in from {
after[f].push(to);
}
}
let mut ready: BinaryHeap<Reverse<(usize, usize)>> = (0..u)
.filter(|&x| waiting[x] == 0)
.map(|x| Reverse((first[x], x)))
.collect();
let mut order: Vec<usize> = Vec::with_capacity(u);
while let Some(Reverse((_, x))) = ready.pop() {
order.push(x);
for &y in &after[x] {
waiting[y] -= 1;
if waiting[y] == 0 {
ready.push(Reverse((first[y], y)));
}
}
}
assert_eq!(order.len(), u, "units of a convex partition are acyclic");
let mut renumber = vec![0usize; u];
for (new, &old) in order.iter().enumerate() {
renumber[old] = new;
}
let mut ordered: Vec<Vec<usize>> = vec![Vec::new(); u];
for (old, members) in units.into_iter().enumerate() {
ordered[renumber[old]] = members;
}
for x in unit_of.iter_mut() {
*x = renumber[*x];
}
UnitPlan {
units: ordered,
unit_of,
}
}
#[cfg(test)]
mod tests {
use super::*;
fn plan(preds: Vec<Vec<usize>>, fusible: Vec<bool>) -> UnitPlan {
let n = preds.len();
plan_units(
&preds,
&vec![Vec::new(); n],
&fusible,
&vec![0; n],
&(0..n).collect::<Vec<_>>(),
&|_| false,
)
}
#[test]
fn independent_chains_are_separate_units() {
let p = plan(vec![vec![], vec![], vec![0], vec![1]], vec![true; 4]);
assert_eq!(p.units.len(), 2);
assert_eq!(p.unit_of[0], p.unit_of[2]);
assert_eq!(p.unit_of[1], p.unit_of[3]);
assert_ne!(p.unit_of[0], p.unit_of[1]);
}
#[test]
fn a_component_that_is_not_convex_splits() {
let p = plan(vec![vec![], vec![0], vec![0, 1]], vec![true, false, true]);
assert_eq!(p.units.len(), 3);
for (u, members) in p.units.iter().enumerate() {
for &m in members {
for &pr in &[vec![], vec![0], vec![0, 1]][m] {
assert!(p.unit_of[pr] <= u, "unit {u} reads a later unit");
}
}
}
}
#[test]
fn a_component_is_cut_only_where_a_path_returns() {
let preds = vec![
vec![], vec![0], vec![], vec![2], vec![0], vec![1, 4], ];
let fusible = vec![true, false, true, true, true, true];
let p = plan_units(
&preds,
&vec![Vec::new(); 6],
&fusible,
&[0; 6],
&[0, 1, 2, 3, 4, 5],
&|_| false,
);
assert_eq!(p.unit_of[0], p.unit_of[4], "0 and 4 fuse across the chain");
assert_ne!(p.unit_of[0], p.unit_of[5], "5 is where the path returns");
assert_eq!(p.unit_of[2], p.unit_of[3]);
assert_ne!(p.unit_of[2], p.unit_of[0]);
assert_eq!(p.units.len(), 4);
}
#[test]
fn a_leaf_copying_an_input_joins_a_unit_over_that_input() {
let preds = vec![vec![], vec![0], vec![], vec![], vec![]];
let inputs = vec![vec![0, 1], vec![], vec![0], vec![1], vec![9]];
let p = plan_units(
&preds,
&inputs,
&[true; 5],
&[0; 5],
&[0, 1, 2, 3, 4],
&|_| false,
);
assert_eq!(p.units.len(), 2, "{:?}", p.units);
assert_eq!(p.unit_of[2], p.unit_of[0]);
assert_eq!(p.unit_of[3], p.unit_of[0]);
assert_ne!(p.unit_of[4], p.unit_of[0]);
}
#[test]
fn classes_do_not_fuse() {
let preds = vec![vec![], vec![0]];
let p = plan_units(
&preds,
&vec![Vec::new(); 2],
&[true, true],
&[0, 1],
&[0, 1],
&|_| false,
);
assert_eq!(p.units.len(), 2);
}
}