use std::{
fmt::{self, Debug},
marker::PhantomData,
ops::ControlFlow,
sync::Mutex,
};
use bitvec::{bitvec, vec::BitVec};
use crate::{
backend::ISA,
compiler::graph::{BaseOp, CtrlOp, Graph, NodeId, OpCode, Signature, Top},
};
use super::liveness::Liveness;
pub struct CFG<Op: OpCode = BaseOp> {
pub g: Box<Graph>,
pub signature: Signature,
pub blocks: Vec<Box<Block>>,
pub nodes: Vec<NodeId>,
pub entry_block: usize,
pub liveness: Liveness,
pub used_registers: BitVec,
pub stack_map: Vec<NodeId>,
pub stack_size: usize,
_p: PhantomData<Op>,
}
impl<Op: OpCode> CFG<Op> {
pub(super) fn new(g: Box<Graph>) -> Self {
Self {
signature: g.signature.clone(),
g,
blocks: vec![],
nodes: vec![],
entry_block: 0,
liveness: Liveness::default(),
used_registers: bitvec![0; 0],
stack_map: vec![],
stack_size: 0,
_p: PhantomData,
}
}
fn visit_blocks_postorder(
blocks: &mut Vec<Box<Block>>,
b: usize,
prologue: &mut impl FnMut(&mut Vec<Box<Block>>, usize) -> ControlFlow<(), ()>,
visit: &mut impl FnMut(&mut Vec<Box<Block>>, usize),
) {
match prologue(blocks, b) {
ControlFlow::Break(_) => return,
_ => {}
}
for s in blocks[b].succs.clone() {
Self::visit_blocks_postorder(blocks, s, prologue, visit);
}
visit(blocks, b);
assert!(blocks[b].id != usize::MAX - 1);
}
pub(super) fn sort_blocks_rpo(&mut self) {
let mut id = 0usize;
Self::visit_blocks_postorder(
&mut self.blocks,
self.entry_block,
&mut |blocks, b| {
if blocks[b].id != usize::MAX {
ControlFlow::Break(())
} else {
blocks[b].id = usize::MAX - 1;
ControlFlow::Continue(())
}
},
&mut |blocks, b| {
debug_assert!(id < usize::MAX - 1);
blocks[b].id = id;
id += 1;
},
);
let num_blocks = self.blocks.len();
for b in &mut self.blocks {
b.id = num_blocks - 1 - b.id;
}
for b in 0..num_blocks {
for i in 0..self.blocks[b].preds.len() {
self.blocks[b].preds[i] = self.blocks[self.blocks[b].preds[i]].id;
}
for i in 0..self.blocks[b].succs.len() {
self.blocks[b].succs[i] = self.blocks[self.blocks[b].succs[i]].id;
}
}
for n in 0..self.g.nodes.len() {
let n = NodeId::<Top>::from(n);
if let Some(block) = self.g[n].block {
self.g[n].block = Some(self.blocks[block].id);
}
}
self.blocks.sort_by_key(|b| b.id);
assert_eq!(
self.g[self.blocks[0].label].op::<Op>().ctrl_op(),
Some(CtrlOp::Start)
);
self.entry_block = 0;
}
fn compute_dom_depth(&mut self, b: usize) -> usize {
if let Some(depth) = self.blocks[b].dom_depth {
return depth;
}
if b == self.entry_block {
self.blocks[b].dom_depth = Some(0);
return 0;
}
debug_assert_ne!(self.blocks[b].idom.unwrap(), b);
let depth = self.compute_dom_depth(self.blocks[b].idom.unwrap()) + 1;
self.blocks[b].dom_depth = Some(depth);
depth
}
pub(super) fn build_dom_tree(&mut self) {
let num_blocks = self.blocks.len();
let mut doms = vec![usize::MAX; num_blocks];
doms[self.entry_block] = self.entry_block;
let intersect = |mut b1: usize, mut b2: usize, doms: &[usize]| {
while b1 != b2 {
while b1 > b2 {
debug_assert_ne!(doms[b1], usize::MAX);
b1 = doms[b1];
}
while b2 > b1 {
debug_assert_ne!(doms[b2], usize::MAX);
b2 = doms[b2];
}
}
b1
};
let mut changed = true;
while changed {
changed = false;
for i in (0..self.blocks.len()).rev() {
if i == self.entry_block {
continue;
}
let Some(mut new_idom) = self.blocks[i]
.preds
.iter()
.find(|p| doms[**p] != usize::MAX)
.cloned()
else {
continue;
};
for p in 0..self.blocks[i].preds.len() {
if doms[self.blocks[i].preds[p]] != usize::MAX {
new_idom = intersect(self.blocks[i].preds[p], new_idom, &doms);
}
}
if doms[i] != new_idom {
doms[i] = new_idom;
changed = true;
}
}
}
for i in 0..self.blocks.len() {
if i == self.entry_block {
continue;
}
self.blocks[i].idom = Some(doms[i]);
}
for i in 0..self.blocks.len() {
self.blocks[i].dom = DOM::new(self.blocks.len());
if i == self.entry_block {
continue;
}
let mut parent = self.blocks[i].idom;
while let Some(b) = parent {
self.blocks[i].dom.insert(b);
parent = self.blocks[b].idom;
}
}
for i in 0..self.blocks.len() {
self.compute_dom_depth(i);
}
}
pub fn compute_loops(&mut self) {
let mut loops = vec![];
for i in 0..self.blocks.len() {
self.blocks[i].loop_depth = Some(0);
if self.blocks[i].preds.len() < 2 {
continue;
}
for j in 0..self.blocks[i].preds.len() {
let p = self.blocks[i].preds[j];
if !self.blocks[p].dom.contains(i) {
continue;
}
loops.push(i..p);
break;
}
}
let mut visited = vec![false; self.blocks.len()];
for loop_ in loops {
let mut queue = vec![loop_.end];
while let Some(b) = queue.pop() {
if visited[b] {
continue;
}
visited[b] = true;
self.blocks[b].loop_depth = Some(self.blocks[b].loop_depth.unwrap() + 1);
if b == loop_.start {
continue;
}
for p in &self.blocks[b].preds {
queue.push(*p);
}
}
}
for b in 0..self.blocks.len() {
for c in 0..self.blocks[b].succs.len() {
let c = self.blocks[b].succs[c];
if self.blocks[b].dom.contains(c) {
self.blocks[c].loop_info = Some(LoopInfo::Header { end: b });
self.blocks[b].loop_info = Some(LoopInfo::End { header: c });
}
}
}
}
pub(crate) fn verify_numbering(&self) {
let mut entry_block = None;
for (i, b) in self.blocks.iter().enumerate() {
assert_eq!(b.id, i);
if self.g[b.label].op::<Op>().ctrl_op() == Some(CtrlOp::Start) {
assert!(entry_block.is_none());
entry_block = Some(i);
}
}
assert!(entry_block.is_some());
assert_eq!(self.entry_block, entry_block.unwrap());
for (i, n) in self.nodes.iter().enumerate() {
assert_eq!(self.g[*n].cfg_id, i);
}
}
pub(crate) fn renumber_nodes(&mut self) {
let mut id = 0;
macro_rules! assign_id {
($n:expr) => {{
let n = $n;
self.g[n].cfg_id = id;
id += 1;
self.nodes.push(n);
}};
}
for b in 0..self.blocks.len() {
assign_id!(self.blocks[b].label.cast());
let mut num_phis = 0;
for n in &self.blocks[b].nodes {
if self.g[*n].op::<BaseOp>().is_phi() {
num_phis += 1;
}
assign_id!(*n);
}
self.blocks[b].num_phis = num_phis;
assign_id!(self.blocks[b].terminal.cast());
}
}
pub fn compute_liveness<Isa: ISA<Op = Op>>(
&mut self,
coalescable_values: Vec<(NodeId, NodeId)>,
) {
self.liveness = Liveness::compute::<Isa>(self, coalescable_values);
}
}
impl<Op: OpCode> fmt::Debug for CFG<Op> {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
writeln!(f, "CFG {:?} {{", self.signature)?;
for b in &self.blocks {
*b.g.lock().unwrap() = Some(&*self.g);
write!(f, "{:?}", b)?;
}
write!(f, "}}")
}
}
#[derive(Clone)]
pub struct DOM {
table: BitVec,
}
impl DOM {
fn new(size: usize) -> Self {
let mut bit_vec = BitVec::new();
bit_vec.resize(size, false);
Self { table: bit_vec }
}
fn uninit() -> Self {
Self {
table: BitVec::new(),
}
}
pub fn contains(&self, block: usize) -> bool {
self.table[block]
}
pub fn equals(&self, other: &Self) -> bool {
assert_eq!(self.table.len(), other.table.len());
self.table == other.table
}
fn insert(&mut self, block: usize) {
self.table.set(block, true);
}
pub fn iter(&self) -> impl Iterator<Item = usize> + '_ {
self.table
.iter()
.enumerate()
.filter(|(_, x)| **x)
.map(|(i, _)| i)
}
}
impl Debug for DOM {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
write!(f, "DOM(")?;
let mut first = true;
for i in self.iter() {
if first {
first = false;
write!(f, "#{}", i)?;
} else {
write!(f, ", #{}", i)?;
}
}
write!(f, ")")
}
}
pub enum LoopInfo {
Header { end: usize },
End { header: usize },
}
pub struct Block {
pub id: usize,
pub label: NodeId<()>,
pub nodes: Vec<NodeId>,
pub terminal: NodeId<()>,
pub preds: Vec<usize>,
pub succs: Vec<usize>,
pub dom: DOM,
pub idom: Option<usize>,
pub dom_depth: Option<usize>,
pub loop_depth: Option<usize>,
pub loop_info: Option<LoopInfo>,
pub live_in: BitVec,
pub live_out: BitVec,
pub num_phis: usize,
g: Mutex<Option<*const Graph>>,
}
impl Block {
pub(super) fn new(label: NodeId<()>, terminal: NodeId<()>) -> Self {
Self {
id: usize::MAX,
label,
nodes: vec![],
terminal,
preds: vec![],
succs: vec![],
dom: DOM::uninit(),
idom: None,
dom_depth: None,
loop_depth: None,
loop_info: None,
live_in: BitVec::new(),
live_out: BitVec::new(),
num_phis: 0,
g: Default::default(),
}
}
fn traverse_nodes_rpo_impl(
&self,
g: &mut Graph,
node: NodeId,
nodes: &mut Vec<NodeId>,
mark: usize,
) {
if g[node].block != Some(self.id) || g[node].mark == mark {
return;
}
g[node].mark = mark;
for i in 0..g[node].inputs.len() {
let i = g[node].inputs[i];
if g[i].block == Some(self.id) && g[i].mark != mark {
self.traverse_nodes_rpo_impl(g, i, nodes, mark)
}
}
for i in 0..g[node].controls.len() {
let i = g[node].controls[i];
if g[i].block == Some(self.id)
&& !g[i].op::<BaseOp>().is_label_or_terminal()
&& g[i].mark != mark
{
self.traverse_nodes_rpo_impl(g, i.cast(), nodes, mark)
}
}
for i in 0..g[node].effects.len() {
let i = g[node].effects[i];
if g[i].block == Some(self.id)
&& !g[i].op::<BaseOp>().is_label_or_terminal()
&& g[i].mark != mark
{
self.traverse_nodes_rpo_impl(g, i, nodes, mark)
}
}
nodes.push(node);
}
pub(super) fn sort_nodes_rpo<Op: OpCode>(&mut self, g: &mut Graph, mark: usize) {
let mut nodes = vec![];
for i in &g[self.terminal].inputs.clone() {
self.traverse_nodes_rpo_impl(g, *i, &mut nodes, mark);
}
for i in self.nodes.iter().rev() {
self.traverse_nodes_rpo_impl(g, *i, &mut nodes, mark);
}
assert_eq!(nodes.len(), self.nodes.len());
self.nodes.clear();
nodes.retain(|n| {
let op = g[*n].op::<BaseOp>();
if op.is_phi() || op.is_param() {
self.nodes.push(*n);
false
} else {
true
}
});
for n in nodes {
self.nodes.push(n)
}
let mut moves = vec![];
self.nodes.retain(|n| {
let is_constrained_move = g[*n].op::<Op>().is_move()
&& (g[*n].move_node_after.is_some() || g[*n].move_node_before.is_some());
if is_constrained_move {
moves.push(*n);
}
!is_constrained_move
});
for mov in moves {
if let Some(n) = g[mov].move_node_after {
debug_assert_eq!(g[n].block, g[mov].block);
let i = self.nodes.iter().position(|x| *x == n).unwrap();
self.nodes.insert(i + 1, mov);
} else if let Some(n) = g[mov].move_node_before {
debug_assert_eq!(g[n].block, g[mov].block);
if let Some(i) = self.nodes.iter().position(|x| *x == n) {
self.nodes.insert(i, mov);
} else {
debug_assert_eq!(n, self.terminal.cast());
self.nodes.push(mov);
}
}
}
}
}
impl fmt::Debug for Block {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
let pred = self
.preds
.iter()
.map(|x| format!("{}", x))
.collect::<Vec<_>>()
.join(",");
let succ = self
.succs
.iter()
.map(|x| format!("{}", x))
.collect::<Vec<_>>()
.join(",");
writeln!(
f,
"[block#{}] pred=({}) succ=({}) idom={:?} dom_depth={:?} loop_depth={:?}",
self.id, pred, succ, self.idom, self.dom_depth, self.loop_depth
)?;
let g = unsafe { &*(self.g.lock().unwrap().unwrap()) };
for n in &self.nodes {
write!(f, " {:?}", g[*n])?;
if g[*n].fixed_reg.is_some() {
write!(f, " (Fixed-Reg: {:?})", g[*n].fixed_reg.unwrap())?;
}
writeln!(f, "")?;
}
writeln!(f, " {:?}", g[self.terminal])?;
Ok(())
}
}