use std::{
cmp::Reverse,
collections::{BTreeMap, BTreeSet, BinaryHeap},
ops::RangeInclusive
};
use crate::{
Add, AddressingMode, CanVisitInstructions as _, DependencyAnalyzer, Div,
DropHighest, DropLowest, Exp, Function, Instruction, InstructionVisitor,
Max, Mod, Mul, Neg, Optimizer, ProgramCounter, RegisterIndex, Return,
RollCustomDice, RollRange, RollStandardDice, Sub, SumRollingRecord
};
type LivenessMap = BTreeMap<RegisterIndex, RangeInclusive<ProgramCounter>>;
type Coloring = BTreeMap<RegisterIndex, RegisterIndex>;
#[derive(Debug, Clone, Default, PartialEq, Eq)]
pub struct RegisterCoalescer
{
function: Option<Function>,
coloring: Coloring,
instructions: Vec<Instruction>
}
impl Optimizer<()> for RegisterCoalescer
{
fn optimize(mut self, mut function: Function) -> Result<Function, ()>
{
self.function = Some(function.clone());
let analyzer =
DependencyAnalyzer::analyze(&self.function().instructions);
let liveness = self.compute_liveness(analyzer);
self.coloring = Self::colorize(&liveness);
for instruction in &function.instructions
{
instruction.visit(&mut self).unwrap();
}
function.register_count = self
.coloring
.values()
.max()
.map(|count| count.0 + 1)
.unwrap_or(0);
function.instructions = self.instructions;
Ok(function)
}
}
impl RegisterCoalescer
{
fn function(&self) -> &Function { self.function.as_ref().unwrap() }
fn compute_liveness(&self, analyzer: DependencyAnalyzer<'_>)
-> LivenessMap
{
let mut liveness = LivenessMap::new();
for r in 0..self.function().register_count
{
let r: RegisterIndex = r.into();
let start = match analyzer.writers().get(&r.into())
{
Some(writers) =>
{
(writers.first().unwrap().0 + 1).into()
},
None => 0.into()
};
let end = match analyzer.readers().get(&r.into())
{
Some(readers) => *readers.last().unwrap(),
None => start
};
let range = start..=end;
liveness.insert(r, range);
}
liveness
}
fn colorize(liveness: &LivenessMap) -> Coloring
{
let mut order = liveness
.iter()
.map(|(r, range)| (*range.start(), *r))
.collect::<Vec<_>>();
order.sort();
let mut colors = Coloring::new();
let mut free = BTreeSet::<RegisterIndex>::new();
let mut live =
BinaryHeap::<Reverse<(ProgramCounter, RegisterIndex)>>::new();
let mut palette = 0;
for (start, r) in order
{
while let Some(&Reverse((end, color))) = live.peek()
&& end < start
{
live.pop();
free.insert(color);
}
let color = free.pop_first().unwrap_or_else(|| {
palette += 1;
RegisterIndex(palette - 1)
});
colors.insert(r, color);
live.push(Reverse((*liveness[&r].end(), color)));
}
colors
}
#[inline]
fn color(&self, op: impl Into<AddressingMode>) -> AddressingMode
{
match op.into()
{
AddressingMode::Register(r) => self.coloring[&r].into(),
op => op
}
}
#[inline]
fn emit(&mut self, inst: impl Into<Instruction>) -> Result<(), ()>
{
self.instructions.push(inst.into());
Ok(())
}
}
impl InstructionVisitor<()> for RegisterCoalescer
{
fn visit_roll_range(&mut self, inst: &RollRange) -> Result<(), ()>
{
self.emit(RollRange {
dest: inst.dest,
start: self.color(inst.start),
end: self.color(inst.end)
})
}
fn visit_roll_standard_dice(
&mut self,
inst: &RollStandardDice
) -> Result<(), ()>
{
self.emit(RollStandardDice {
dest: inst.dest,
count: self.color(inst.count),
faces: self.color(inst.faces)
})
}
fn visit_roll_custom_dice(
&mut self,
inst: &RollCustomDice
) -> Result<(), ()>
{
self.emit(RollCustomDice {
dest: inst.dest,
count: self.color(inst.count),
faces: inst.faces.clone()
})
}
fn visit_drop_lowest(&mut self, inst: &DropLowest) -> Result<(), ()>
{
self.emit(DropLowest {
dest: inst.dest,
count: self.color(inst.count)
})
}
fn visit_drop_highest(&mut self, inst: &DropHighest) -> Result<(), ()>
{
self.emit(DropHighest {
dest: inst.dest,
count: self.color(inst.count)
})
}
fn visit_sum_rolling_record(
&mut self,
inst: &SumRollingRecord
) -> Result<(), ()>
{
self.emit(SumRollingRecord {
dest: self.color(inst.dest).try_into().unwrap(),
src: inst.src
})
}
fn visit_add(&mut self, inst: &Add) -> Result<(), ()>
{
self.emit(Add {
dest: self.color(inst.dest).try_into().unwrap(),
op1: self.color(inst.op1),
op2: self.color(inst.op2)
})
}
fn visit_sub(&mut self, inst: &Sub) -> Result<(), ()>
{
self.emit(Sub {
dest: self.color(inst.dest).try_into().unwrap(),
op1: self.color(inst.op1),
op2: self.color(inst.op2)
})
}
fn visit_mul(&mut self, inst: &Mul) -> Result<(), ()>
{
self.emit(Mul {
dest: self.color(inst.dest).try_into().unwrap(),
op1: self.color(inst.op1),
op2: self.color(inst.op2)
})
}
fn visit_div(&mut self, inst: &Div) -> Result<(), ()>
{
self.emit(Div {
dest: self.color(inst.dest).try_into().unwrap(),
op1: self.color(inst.op1),
op2: self.color(inst.op2)
})
}
fn visit_mod(&mut self, inst: &Mod) -> Result<(), ()>
{
self.emit(Mod {
dest: self.color(inst.dest).try_into().unwrap(),
op1: self.color(inst.op1),
op2: self.color(inst.op2)
})
}
fn visit_exp(&mut self, inst: &Exp) -> Result<(), ()>
{
self.emit(Exp {
dest: self.color(inst.dest).try_into().unwrap(),
op1: self.color(inst.op1),
op2: self.color(inst.op2)
})
}
fn visit_max(&mut self, inst: &Max) -> Result<(), ()>
{
self.emit(Max {
dest: self.color(inst.dest).try_into().unwrap(),
op1: self.color(inst.op1),
op2: self.color(inst.op2)
})
}
fn visit_neg(&mut self, inst: &Neg) -> Result<(), ()>
{
self.emit(Neg {
dest: self.color(inst.dest).try_into().unwrap(),
op: self.color(inst.op)
})
}
fn visit_return(&mut self, inst: &Return) -> Result<(), ()>
{
self.emit(Return {
src: self.color(inst.src)
})
}
}