use rucc_mir::{Func, Reg};
use rucc_target::RegClass;
use crate::assign::{self, Env};
use crate::live::{Area, Live};
use crate::order::{Order, Point};
#[derive(Debug, Clone, Default)]
pub struct Pressure {
wanted: Vec<Vec<u32>>,
room: Vec<Vec<u32>>,
}
impl Pressure {
#[must_use]
pub fn of(func: &Func, order: &Order, live: &Live, env: &Env) -> Self {
let points = usize::try_from(order.points()).expect("a point count");
let mut forced = vec![false; func.vregs()];
for reg in assign::forced(func) {
forced[index(reg)] = true;
}
let mut steps: Vec<Vec<i64>> = Vec::new();
for (number, &forced) in forced.iter().enumerate() {
let reg = Reg::virtual_reg(u32::try_from(number).expect("a register number"));
let (Some(area), Some(class)) = (live.area(reg), func.class_of(reg)) else {
continue;
};
if forced {
continue;
}
let class = usize::from(class.number());
if steps.len() <= class {
steps.resize_with(class + 1, Vec::new);
}
let row = &mut steps[class];
if row.is_empty() {
*row = vec![0; points + 1];
}
for piece in area.pieces() {
row[at(piece.start)] += 1;
row[at(piece.end) + 1] -= 1;
}
}
let wanted = steps
.into_iter()
.map(|row| {
let mut count = 0i64;
let mut wanted: Vec<u32> = row
.iter()
.map(|step| {
count += step;
u32::try_from(count).expect("a count of values")
})
.collect();
wanted.truncate(points);
wanted
})
.collect();
let mut room: Vec<Vec<u32>> = env
.offered()
.map(|offered| {
let count = u32::try_from(offered.len()).expect("a count of registers");
if count == 0 { Vec::new() } else { vec![count; points] }
})
.collect();
let mut taken: Vec<_> = assign::blocked(func, order).taken().collect();
taken.dedup();
for (class, reg, point) in taken {
let Some(row) = room.get_mut(usize::from(class.number())) else { continue };
if row.is_empty() || !env.order(class).contains(®) {
continue;
}
row[at(point)] = row[at(point)].saturating_sub(1);
}
Self { wanted, room }
}
#[must_use]
pub fn wanted(&self, class: RegClass, point: Point) -> u32 {
read(&self.wanted, class, point)
}
#[must_use]
pub fn room(&self, class: RegClass, point: Point) -> u32 {
read(&self.room, class, point)
}
#[must_use]
pub fn excess(&self, class: RegClass, point: Point) -> u32 {
self.wanted(class, point).saturating_sub(self.room(class, point))
}
#[must_use]
pub fn peak(&self, class: RegClass) -> u32 {
row(&self.wanted, class).iter().copied().max().unwrap_or(0)
}
pub fn over(&self, class: RegClass) -> impl Iterator<Item = Point> + '_ {
let wanted = row(&self.wanted, class);
(0..wanted.len())
.filter_map(|point| u32::try_from(point).ok())
.filter(move |&point| self.excess(class, point) > 0)
}
pub(crate) fn lift(&mut self, class: RegClass, area: Area<'_>) {
let Some(row) = self.wanted.get_mut(usize::from(class.number())) else { return };
for piece in area.pieces() {
for point in piece.start..=piece.end {
if let Some(count) = row.get_mut(at(point)) {
*count = count.saturating_sub(1);
}
}
}
}
}
fn row(table: &[Vec<u32>], class: RegClass) -> &[u32] {
table.get(usize::from(class.number())).map_or(&[], Vec::as_slice)
}
fn read(table: &[Vec<u32>], class: RegClass, point: Point) -> u32 {
row(table, class).get(at(point)).copied().unwrap_or(0)
}
fn at(point: Point) -> usize {
usize::try_from(point).expect("a point")
}
fn index(reg: Reg) -> usize {
usize::try_from(reg.number().expect("a virtual register")).expect("a register number")
}
#[cfg(test)]
mod tests {
use rucc_base::Interner;
use rucc_mir::{Opcode, Operand};
use rucc_target::x86_64::{GPR, RAX, SYSV};
use super::*;
fn narrow(count: usize) -> Env {
Env::new().with(GPR, &SYSV.int_order[..count], &SYSV.int_order[count..count + 1])
}
fn of(func: &Func, env: &Env) -> (Order, Pressure) {
let order = Order::of(func);
let live = Live::of(func, &order);
let pressure = Pressure::of(func, &order, &live, env);
(order, pressure)
}
#[test]
fn every_value_live_at_a_point_is_counted_there() {
let mut names = Interner::new();
let mut func = Func::new(names.intern("f"));
let opcode = Opcode::new(names.intern("x64.nop"));
let block = func.create_block();
let values: Vec<Reg> = (0..3).map(|_| func.new_vreg(GPR)).collect();
let defs: Vec<_> = values
.iter()
.map(|&value| func.build(block, opcode).def(value, GPR).finish())
.collect();
let reads = func.build(block, opcode);
let reads = values.iter().fold(reads, |build, &value| build.uses(value, GPR));
let last = reads.finish();
let (order, pressure) = of(&func, &narrow(2));
assert_eq!(pressure.wanted(GPR, order.early(last)), 3);
assert_eq!(pressure.room(GPR, order.early(last)), 2);
assert_eq!(pressure.excess(GPR, order.early(last)), 1);
assert_eq!(pressure.peak(GPR), 3);
let over = [order.late(defs[2]), order.early(last)];
assert_eq!(pressure.over(GPR).collect::<Vec<_>>(), over);
}
#[test]
fn a_register_an_instruction_writes_outright_is_not_room_where_it_writes_it() {
let mut names = Interner::new();
let mut func = Func::new(names.intern("f"));
let opcode = Opcode::new(names.intern("x64.nop"));
let block = func.create_block();
let value = func.new_vreg(GPR);
func.build(block, opcode).def(value, GPR).finish();
let call =
func.build(block, opcode).operand(Operand::write(Reg::physical(RAX), GPR)).finish();
func.build(block, opcode).uses(value, GPR).finish();
let (order, pressure) = of(&func, &narrow(2));
assert_eq!(pressure.room(GPR, order.early(call)), 2);
assert_eq!(pressure.room(GPR, order.late(call)), 1);
assert_eq!(pressure.wanted(GPR, order.late(call)), 1);
assert_eq!(pressure.over(GPR).count(), 0);
}
#[test]
fn a_value_sent_to_memory_is_taken_off_every_point_it_was_live_at() {
let mut names = Interner::new();
let mut func = Func::new(names.intern("f"));
let opcode = Opcode::new(names.intern("x64.nop"));
let block = func.create_block();
let first = func.new_vreg(GPR);
let second = func.new_vreg(GPR);
func.build(block, opcode).def(first, GPR).finish();
func.build(block, opcode).def(second, GPR).finish();
let last = func.build(block, opcode).uses(first, GPR).uses(second, GPR).finish();
let order = Order::of(&func);
let live = Live::of(&func, &order);
let mut pressure = Pressure::of(&func, &order, &live, &narrow(1));
assert_eq!(pressure.excess(GPR, order.early(last)), 1);
pressure.lift(GPR, live.area(first).expect("live"));
assert_eq!(pressure.over(GPR).count(), 0);
assert_eq!(pressure.wanted(GPR, order.early(last)), 1);
}
}