use std::ops::{Deref, DerefMut, Range};
use bitvec::bitvec;
use crate::compiler::graph::{BaseOp, Graph};
use crate::{
backend::ISA,
compiler::graph::{NodeId, OpCode, Type},
};
use super::CFG;
#[derive(Debug, PartialEq, Default, Clone)]
pub struct LiveRange {
pub start: usize,
pub end: usize,
}
impl From<Range<usize>> for LiveRange {
fn from(value: Range<usize>) -> Self {
Self {
start: value.start,
end: value.end,
..Default::default()
}
}
}
#[derive(Debug, PartialEq, Eq, Clone, Copy)]
pub enum LiveIntervalValueCategory {
Int,
Float,
}
impl From<Type> for Option<LiveIntervalValueCategory> {
fn from(value: Type) -> Self {
match value {
Type::Bool | Type::I8 | Type::I16 | Type::I32 | Type::I64 => {
Some(LiveIntervalValueCategory::Int)
}
Type::F32 | Type::F64 => Some(LiveIntervalValueCategory::Float),
_ => None,
}
}
}
#[derive(Debug, PartialEq)]
pub struct LiveInterval {
rep: Option<usize>,
pub ranges: Vec<LiveRange>,
pub value_category: Option<LiveIntervalValueCategory>,
pub reg: Option<usize>,
pub mem: i32,
pub max_mem_size: usize,
pub weight: usize,
}
impl Default for LiveInterval {
fn default() -> Self {
Self {
rep: None,
ranges: vec![],
value_category: Type::Top.into(),
max_mem_size: 0,
reg: None,
mem: i32::MAX,
weight: 0,
}
}
}
impl LiveInterval {
fn new(g: &Graph, node: NodeId) -> Self {
Self {
value_category: if (g[node].op::<BaseOp>() == BaseOp::Call
|| g[node].op::<BaseOp>() == BaseOp::CallIndirect)
&& g[node].ty == Type::Void
{
Some(LiveIntervalValueCategory::Int)
} else {
g[node].ty.into()
},
max_mem_size: if g[node].ty != Type::Void && g[node].ty != Type::Top {
g[node].ty.mem_size()
} else {
0
},
reg: g[node].fixed_reg,
weight: 1,
..Default::default()
}
}
pub fn intersects(&self, other: &Self) -> bool {
self.find_first_intersection(other).is_some()
}
pub fn find_first_intersection(&self, other: &Self) -> Option<usize> {
let mut i = 0;
let mut j = 0;
while i < self.ranges.len() && j < other.ranges.len() {
let x = &self.ranges[i];
let y = &other.ranges[j];
if x.start < y.end && y.start < x.end {
return Some(usize::max(x.start, y.start));
}
if x.end < y.end {
i += 1;
} else {
j += 1;
}
}
None
}
pub fn covers(&self, index: usize) -> bool {
self.ranges
.iter()
.any(|x| x.start <= index && index < x.end)
}
fn merge(&mut self) {
self.ranges.sort_by_key(|x| x.start);
let mut index = 0;
for i in 0..self.ranges.len() {
if self.ranges[index].end >= self.ranges[i].start {
self.ranges[index].end = self.ranges[index].end.max(self.ranges[i].end);
} else {
index += 1;
self.ranges[index] = self.ranges[i].clone();
}
}
self.ranges = self.ranges[0..=index].to_vec();
}
fn add(&mut self, range: LiveRange) {
self.ranges.push(range);
self.merge();
}
pub fn union(&mut self, ranges: &[LiveRange]) {
for range in ranges {
self.ranges.push(range.clone());
}
self.merge();
}
}
#[derive(Debug, Default)]
pub struct Liveness {
pub block_ranges: Vec<LiveRange>,
pub intervals: Vec<LiveInterval>,
pub stack_map: Vec<usize>,
}
impl Deref for Liveness {
type Target = Vec<LiveInterval>;
fn deref(&self) -> &Self::Target {
&self.intervals
}
}
impl DerefMut for Liveness {
fn deref_mut(&mut self) -> &mut Self::Target {
&mut self.intervals
}
}
impl Liveness {
pub(super) fn compute<Isa: ISA>(
cfg: &mut CFG<Isa::Op>,
coalescable_values: Vec<(NodeId, NodeId)>,
) -> Self {
let mut me = Self {
intervals: vec![],
block_ranges: vec![],
stack_map: vec![],
};
me.compute_intervals::<Isa>(cfg, coalescable_values);
me
}
fn compute_intervals<Isa: ISA>(
&mut self,
cfg: &mut CFG<Isa::Op>,
coalescable_values: Vec<(NodeId, NodeId)>,
) {
if cfg!(debug_assertions) {
cfg.verify_numbering();
}
self.block_ranges = (0..cfg.blocks.len())
.map(|_| LiveRange::from(0..0))
.collect::<Vec<_>>();
let mut intervals: Vec<LiveInterval> = cfg
.nodes
.iter()
.map(|n| LiveInterval::new(&cfg.g, *n))
.collect();
for b in 0..cfg.blocks.len() {
self.block_ranges[b].start = cfg.g[cfg.blocks[b].label].cfg_id;
self.block_ranges[b].end = cfg.g[cfg.blocks[b].terminal].cfg_id + 1;
}
for b in (0..cfg.blocks.len()).rev() {
let mut live = bitvec![0; cfg.nodes.len()];
for i in 0..cfg.blocks[b].succs.len() {
let s = cfg.blocks[b].succs[i];
live |= cfg.blocks[s].live_in.clone();
}
for i in 0..cfg.blocks[b].succs.len() {
let s = cfg.blocks[b].succs[i];
let input_index = cfg.g[cfg.blocks[s].label]
.controls
.iter()
.position(|x| *x == cfg.blocks[b].terminal)
.unwrap();
for phi in &cfg.blocks[s].nodes[..cfg.blocks[s].num_phis] {
debug_assert!(input_index < cfg.g[*phi].inputs.len(), "{:?}", phi);
let input = cfg.g[*phi].inputs[input_index];
live.set(cfg.g[*phi].cfg_id, false);
live.set(cfg.g[input].cfg_id, true);
}
}
let add_range = |intervals: &mut Vec<LiveInterval>, n: usize, b: usize, end: usize| {
if self.block_ranges[b].start <= n && n < self.block_ranges[b].end {
let start = if cfg.g[cfg.nodes[n]].op::<Isa::Op>().is_phi() {
n + 1
} else {
n
};
intervals[n].add((start..end).into());
} else {
intervals[n].add((self.block_ranges[b].start..end).into());
}
};
for n in live.iter_ones() {
add_range(&mut intervals, n, b, self.block_ranges[b].end);
}
let mut process_node = |n: NodeId| {
if cfg.g[n].op::<Isa::Op>().is_phi() {
return;
}
live.set(cfg.g[n].cfg_id, false);
for input in cfg.g[n].inputs.iter() {
if !live[cfg.g[n].cfg_id] {
if !cfg.g[n].op::<Isa::Op>().is_phi() {
live.set(cfg.g[*input].cfg_id, true);
}
add_range(&mut intervals, cfg.g[*input].cfg_id, b, cfg.g[n].cfg_id);
}
}
};
process_node(cfg.blocks[b].terminal.cast());
for n in cfg.blocks[b].nodes.iter().rev() {
process_node(*n);
}
process_node(cfg.blocks[b].label.cast());
cfg.blocks[b].live_in = live;
}
assert_eq!(intervals.len(), cfg.nodes.len());
for n in &mut cfg.nodes {
cfg.g[*n].interval = cfg.g[*n].cfg_id;
}
self.intervals = intervals;
for n in &mut cfg.nodes {
if !cfg.g[*n].has_output::<Isa::Op>() {
self.intervals[cfg.g[*n].cfg_id].ranges.clear();
self.intervals[cfg.g[*n].cfg_id].rep = None;
}
}
self.coalesce_all::<Isa>(cfg, coalescable_values);
self.dump(cfg);
}
fn find_rep(&self, mut interval_id: usize) -> usize {
debug_assert_ne!(interval_id, usize::MAX);
while let Some(parent) = self.intervals[interval_id].rep {
debug_assert_ne!(parent, usize::MAX);
interval_id = parent;
}
interval_id
}
fn coalesce_all<Isa: ISA>(
&mut self,
cfg: &mut CFG<Isa::Op>,
coalescable_values: Vec<(NodeId, NodeId)>,
) {
for (x, y) in coalescable_values {
self.coalesce(&cfg.g, x, y, true)
}
for n in &cfg.nodes {
Isa::coalesce_live_intervals(&cfg.g, *n, |x, y| self.coalesce(&cfg.g, x, y, false));
}
let max_intervals = self.intervals.len();
let mut reps = vec![usize::MAX; max_intervals];
let mut count = 0;
for n in &mut cfg.nodes {
assert_ne!(cfg.g[*n].interval, usize::MAX);
let rep = self.find_rep(cfg.g[*n].interval);
if reps[rep] == usize::MAX {
reps[rep] = count;
count += 1;
}
cfg.g[*n].interval = reps[rep];
}
let mut old_intervals = std::mem::take(&mut self.intervals);
let mut intervals = vec![];
intervals.resize_with(count, Default::default);
for (old_index, new_index) in reps.into_iter().enumerate() {
if new_index != usize::MAX {
intervals[new_index] = std::mem::take(&mut old_intervals[old_index]);
}
}
self.intervals = intervals;
let max_intervals = self.intervals.len();
let mut reps = vec![usize::MAX; max_intervals];
let mut count = 0;
for n in &mut cfg.nodes {
assert_ne!(cfg.g[*n].interval, usize::MAX);
if !self.intervals[cfg.g[*n].interval].ranges.is_empty()
|| cfg.g[*n].op::<Isa::Op>().has_output()
{
if self.intervals[cfg.g[*n].interval].ranges.is_empty() {
self.intervals[cfg.g[*n].interval].ranges.push(LiveRange {
start: cfg.g[*n].cfg_id,
end: cfg.g[*n].cfg_id + 1,
})
}
if reps[cfg.g[*n].interval] == usize::MAX {
reps[cfg.g[*n].interval] = count;
count += 1;
}
cfg.g[*n].interval = reps[cfg.g[*n].interval];
} else {
cfg.g[*n].interval = usize::MAX;
}
}
let mut old_intervals = std::mem::take(&mut self.intervals);
let mut intervals = vec![];
intervals.resize_with(count, Default::default);
for (old_index, new_index) in reps.into_iter().enumerate() {
if new_index != usize::MAX {
intervals[new_index] = std::mem::take(&mut old_intervals[old_index]);
}
}
self.intervals = intervals;
}
fn conflcits_with_any_precolored_intervals(&self, i: usize, fixed_reg: usize) -> bool {
for (j, interval) in self.intervals.iter().enumerate() {
if interval.reg.is_none() {
continue;
}
let j = self.find_rep(j);
assert!(self.intervals[j].reg.is_some());
if j == i {
continue;
}
if self.intervals[j].reg.unwrap() == fixed_reg
&& self.intervals[j].intersects(&self.intervals[i])
{
return true;
}
}
false
}
fn coalesce(&mut self, g: &Graph, x: NodeId, y: NodeId, force: bool) {
assert_ne!(g[x].interval, usize::MAX);
assert_ne!(g[y].interval, usize::MAX);
let i = self.find_rep(g[x].interval);
let j = self.find_rep(g[y].interval);
if self.intervals[i].intersects(&self.intervals[j]) && !force {
return;
}
if self.intervals[i].reg.is_some() && self.intervals[j].reg.is_some() {
assert!(!force);
return;
}
assert!(self.intervals[i].value_category.is_some(), "{:?}", i);
if self.intervals[i].value_category.unwrap() != self.intervals[j].value_category.unwrap() {
assert!(!force);
return;
}
let fixed_reg_conflict = {
if let Some(fixed_reg) = self.intervals[j].reg {
self.conflcits_with_any_precolored_intervals(i, fixed_reg)
} else if let Some(fixed_reg) = self.intervals[i].reg {
self.conflcits_with_any_precolored_intervals(j, fixed_reg)
} else {
false
}
};
if fixed_reg_conflict {
assert!(!force);
return;
}
let ranges = std::mem::take(&mut self.intervals[j].ranges.clone());
self.intervals[i].union(&ranges);
self.intervals[i].max_mem_size = usize::max(
self.intervals[i].max_mem_size,
self.intervals[j].max_mem_size,
);
if let Some(fixed_reg) = self.intervals[j].reg {
self.intervals[i].reg = Some(fixed_reg);
} else if let Some(fixed_reg) = self.intervals[i].reg {
self.intervals[j].reg = Some(fixed_reg);
}
self.intervals[j].rep = Some(i);
}
pub(crate) fn dump<Op: OpCode>(&self, cfg: &CFG<Op>) {
use std::fmt::Write;
if !log_enabled!(target: "regalloc", log::Level::Trace) {
return;
}
let mut log = String::new();
let print_ranges = |log: &mut String, ranges: &[LiveRange]| {
write!(log, "[ ").unwrap();
for (i, r) in ranges.iter().enumerate() {
if i != 0 {
write!(log, ", ").unwrap();
}
write!(log, "{}..{}", r.start, r.end).unwrap();
}
write!(log, " ]").unwrap();
};
let print_interval = |log: &mut String, it: &LiveInterval| {
print_ranges(log, &it.ranges);
if let Some(r) = it.reg {
write!(log, " => R{:?}", r).unwrap();
} else if it.mem != i32::MAX {
write!(log, " => SPILL-{:?}", it.mem).unwrap();
}
};
for (n, it) in self.intervals.iter().enumerate() {
let nodes = {
let mut s = "".to_owned();
for node in &cfg.nodes {
if cfg.g[*node].interval == n {
if !s.is_empty() {
s += ", ";
}
s += &format!("{}", node.index());
}
}
s
};
write!(log, "#{:?} nodes=({})", n, nodes).unwrap();
if it.mem != i32::MAX {
write!(log, " [rsp+{}]", it.mem).unwrap();
}
write!(log, ": ").unwrap();
print_interval(&mut log, it);
if n != self.intervals.len() - 1 {
writeln!(log).unwrap();
}
}
trace!(target: "regalloc", "\n{}", log);
}
}