use std::cell::RefCell;
use std::cmp::{Ordering, Reverse};
use std::collections::HashMap;
use std::ops::{Range, RangeBounds};
use std::rc::Rc;
use super::op;
#[derive(Default)]
pub struct RegAlloc(Rc<RefCell<State>>);
#[derive(Default)]
struct State {
intervals: Vec<Interval>,
event: usize,
register: usize,
}
impl State {
fn event(&mut self) -> usize {
let event = self.event;
self.event += 1;
event
}
fn registers(&mut self, n: usize) -> Range<usize> {
let start = self.register;
self.register += n;
start..start + n
}
fn alloc(&mut self) -> (usize, usize) {
let index = self.intervals.len();
let register = self.registers(1).start;
let event = self.event();
self.intervals.push(Interval {
start: event,
end: event,
entry: Entry::Register(register),
});
(index, register)
}
fn alloc_slice(&mut self, n: usize) -> (usize, Range<usize>) {
let index = self.intervals.len();
let slice = self.registers(n);
let event = self.event();
self.intervals.push(Interval {
start: event,
end: event,
entry: Entry::Slice(slice.clone()),
});
(index, slice)
}
fn access(&mut self, index: usize) {
let event = self.event();
self.intervals[index].end = event;
}
}
#[derive(Clone)]
struct Interval {
start: usize,
end: usize,
entry: Entry,
}
#[derive(Clone, PartialEq, Eq, Hash)]
enum Entry {
Register(usize),
Slice(Range<usize>),
}
impl RegAlloc {
pub fn new() -> Self {
Self::default()
}
pub fn alloc(&mut self) -> Register {
let (index, register) = self.0.borrow_mut().alloc();
Register(RegisterKind::Direct {
state: self.0.clone(),
index,
register,
})
}
pub fn alloc_slice(&mut self, n: usize) -> Slice {
let (index, slice) = self.0.borrow_mut().alloc_slice(n);
Slice {
state: self.0.clone(),
slice,
index,
}
}
pub fn finish(&self) -> (usize, Vec<usize>) {
let state = self.0.borrow();
linear_scan(&state.intervals, state.register)
}
}
#[derive(Clone)]
pub struct Register(RegisterKind);
#[derive(Clone)]
enum RegisterKind {
Direct {
state: Rc<RefCell<State>>,
index: usize,
register: usize,
},
Ref {
item: SliceItem,
},
}
impl Register {
pub fn access(&self) -> op::Register {
match &self.0 {
RegisterKind::Direct {
state,
index,
register,
} => {
state.borrow_mut().access(*index);
op::Register(*register as u32)
}
RegisterKind::Ref { item } => item.access(),
}
}
fn index(&self) -> usize {
match &self.0 {
RegisterKind::Direct {
state: _,
index,
register: _,
} => *index,
RegisterKind::Ref { item } => item.slice.index,
}
}
pub fn ensure_non_overlapping(&self, other: Register) {
fn ensure_non_overlapping_impl(state: &mut State, this: usize, other: usize) {
let temp = state.intervals[this].end;
state.intervals[this].end = state.intervals[other].start - 1;
state.intervals[other].start = temp;
}
debug_assert!(self.check_overlap(other.clone()));
let other_index = other.index();
match &self.0 {
RegisterKind::Direct {
state,
index,
register: _,
} => ensure_non_overlapping_impl(&mut state.borrow_mut(), *index, other_index),
RegisterKind::Ref { item } => ensure_non_overlapping_impl(
&mut item.slice.state.borrow_mut(),
item.slice.index,
other_index,
),
}
}
pub fn check_overlap(&self, other: Register) -> bool {
fn check_overlap_impl(state: &State, a: usize, b: usize) -> bool {
state.intervals[a].end >= state.intervals[b].start
}
let other_index = other.index();
match &self.0 {
RegisterKind::Direct {
state,
index,
register: _,
} => check_overlap_impl(&state.borrow(), *index, other_index),
RegisterKind::Ref { item } => {
check_overlap_impl(&item.slice.state.borrow(), item.slice.index, other_index)
}
}
}
}
#[derive(Clone)]
pub struct Slice {
state: Rc<RefCell<State>>,
slice: Range<usize>,
index: usize,
}
impl Slice {
pub fn access(&self, n: usize) -> op::Register {
self.state.borrow_mut().access(self.index);
assert!(self.slice.start + n < self.slice.end);
op::Register((self.slice.start + n) as u32)
}
pub fn get(&self, n: usize) -> Register {
Register(RegisterKind::Ref {
item: SliceItem {
slice: self.clone(),
n,
},
})
}
pub fn offset(&self, offset: usize) -> SliceRef {
SliceRef {
slice: self.clone(),
offset,
}
}
pub fn len(&self) -> usize {
self.slice.len()
}
}
pub struct SliceRef {
slice: Slice,
offset: usize,
}
impl SliceRef {
pub fn access(&self, n: usize) -> op::Register {
self.slice.access(self.offset + n)
}
pub fn get(&self, n: usize) -> Register {
self.slice.get(self.offset + n)
}
}
#[derive(Clone)]
struct SliceItem {
slice: Slice,
n: usize,
}
impl SliceItem {
pub fn access(&self) -> op::Register {
self.slice.access(self.n)
}
}
#[derive(Clone)]
enum Allocation {
Register(usize),
Slice(Range<usize>),
}
type Free = SortedVec<Reverse<usize>>;
type Active = HashMap<Entry, (Interval, Allocation)>;
fn linear_scan(intervals: &[Interval], total_intervals: usize) -> (usize, Vec<usize>) {
let mut mapping = Vec::new();
mapping.resize(total_intervals, 0usize);
let mut free = Free::new();
let mut active = Active::new();
let mut registers = 0usize;
for interval in intervals {
expire_old_intervals(interval, &mut free, &mut active);
match &interval.entry {
Entry::Register(index) => {
let register = allocate(&mut free, &mut registers);
active.insert(
interval.entry.clone(),
(interval.clone(), Allocation::Register(register)),
);
mapping[*index] = register;
}
Entry::Slice(indices) => {
let slice = allocate_slice(indices.len(), &mut free, &mut registers);
active.insert(
interval.entry.clone(),
(interval.clone(), Allocation::Slice(slice.clone())),
);
for (index, register) in indices.clone().zip(slice) {
mapping[index] = register;
}
}
}
}
(registers, mapping)
}
fn allocate(free: &mut Free, registers: &mut usize) -> usize {
if let Some(Reverse(reg)) = free.pop() {
reg
} else {
let reg = *registers;
*registers += 1;
reg
}
}
fn allocate_slice(n: usize, free: &mut Free, registers: &mut usize) -> Range<usize> {
assert!(n > 0);
if n == 1 {
let reg = allocate(free, registers);
reg..reg + 1
} else {
if free.is_empty() || *registers == 0 {
let start = *registers;
*registers += n;
return start..*registers;
}
if let Some(range) = find_contiguous_registers(n, free) {
let slice = &free[range.clone()];
let end = slice.first().unwrap().0 + 1;
let start = slice.last().unwrap().0;
let _ = free.drain(range);
return start..end;
}
if free[0].0 == *registers - 1 {
let mut count = 1;
for i in 1..free.len() {
let (a, b) = (free[i - 1].0, free[i].0);
if a != b + 1 {
break;
}
count += 1;
}
let slice = &free[0..count];
let mut end = slice.first().unwrap().0 + 1;
let start = slice.last().unwrap().0;
let _ = free.drain(0..count);
while (end - start) < n {
end += 1;
*registers += 1;
}
return start..end;
}
let start = *registers;
*registers += n;
start..*registers
}
}
fn find_contiguous_registers(n: usize, free: &Free) -> Option<Range<usize>> {
assert!(n >= 2);
if free.len() < n {
return None;
}
let mut start = 0;
let mut count = 1;
for i in 0..free.len() - 1 {
let (a, b) = (free[i].0, free[i + 1].0);
if a == b + 1 {
count += 1;
if count == n {
return Some(start..start + count);
}
} else {
start += count;
count = 1;
}
}
None
}
fn expire_old_intervals(i: &Interval, free: &mut Free, active: &mut Active) {
active.retain(|_, (j, allocation)| {
if j.end < i.start {
match allocation {
Allocation::Register(register) => free.insert(Reverse(*register)),
Allocation::Slice(slice) => {
for register in slice.clone() {
free.insert(Reverse(register));
}
}
}
false
} else {
true
}
});
}
#[derive(Default)]
struct SortedVec<T> {
inner: Vec<T>,
}
impl<T: Ord> SortedVec<T> {
fn new() -> Self {
SortedVec { inner: vec![] }
}
fn len(&self) -> usize {
self.inner.len()
}
fn is_empty(&self) -> bool {
self.inner.is_empty()
}
fn insert(&mut self, element: T) {
if let None | Some(Ordering::Equal | Ordering::Greater) =
self.inner.last().map(|v| element.cmp(v))
{
self.inner.push(element);
return;
}
let index = match self.inner.binary_search(&element) {
Ok(index) | Err(index) => index,
};
self.inner.insert(index, element);
}
fn pop(&mut self) -> Option<T> {
self.inner.pop()
}
fn drain<R: RangeBounds<usize>>(&mut self, range: R) -> std::vec::Drain<'_, T> {
self.inner.drain(range)
}
fn as_slice(&self) -> &[T] {
self.inner.as_slice()
}
}
impl<T, I: std::slice::SliceIndex<[T]>> std::ops::Index<I> for SortedVec<T> {
type Output = I::Output;
#[inline]
fn index(&self, index: I) -> &Self::Output {
std::ops::Index::index(&self.inner, index)
}
}
#[cfg(all(test, not(feature = "__miri")))]
mod tests;