use std::num::NonZeroU32;
use std::ops::Index;
use crate::backend::treeify::Trees;
use crate::entity::EntityRef;
use crate::ir::{FunctionBody, Value, ValueDef};
#[derive(Clone, Copy, Debug, PartialEq, Eq, PartialOrd, Ord, Hash)]
pub struct CrossBlockId(NonZeroU32);
impl CrossBlockId {
pub fn new(n: u32) -> Self {
Self(NonZeroU32::new(n.checked_add(1).unwrap()).unwrap())
}
pub fn as_u32(self) -> u32 {
self.0.get() - 1
}
}
#[derive(Clone, Debug, Default)]
pub struct CrossBlockValues {
ids: Vec<Option<CrossBlockId>>,
values: Vec<Value>,
}
impl CrossBlockValues {
pub fn len(&self) -> usize {
self.values.len()
}
pub fn is_empty(&self) -> bool {
self.values.is_empty()
}
pub fn init(&mut self, total_values: usize) {
self.ids = vec![None; total_values];
}
pub fn get(&self, value: Value) -> Option<CrossBlockId> {
self.ids[value.index()]
}
pub fn value(&self, id: CrossBlockId) -> Value {
self.values[id.0.get() as usize - 1]
}
pub fn insert(&mut self, value: Value) -> CrossBlockId {
if let Some(id) = self.ids[value.index()] {
return id;
}
let id = CrossBlockId::new(self.values.len() as u32);
self.ids[value.index()] = Some(id);
self.values.push(value);
id
}
pub fn contains(&self, value: Value) -> bool {
self.ids[value.index()].is_some()
}
pub fn ids_mut(&mut self) -> &mut Vec<Option<CrossBlockId>> {
&mut self.ids
}
pub fn ids(&self) -> &Vec<Option<CrossBlockId>> {
&self.ids
}
pub fn iter_values(&self) -> impl Iterator<Item = Value> + '_ {
self.values.iter().copied()
}
pub fn build(body: &FunctionBody, trees: &Trees) -> Self {
let mut def_block = vec![u32::MAX; body.values.len()];
for (block, block_def) in body.blocks.entries() {
for &(_, param) in &block_def.params {
def_block[param.index()] = block.index() as u32;
}
for &inst in &block_def.insts {
if is_tree_or_remat(inst, trees) {
continue;
}
def_block[inst.index()] = block.index() as u32;
}
}
let mut cross_blocks = CrossBlockValues::default();
cross_blocks.init(body.values.len());
for (block, block_def) in body.blocks.entries() {
let cur = block.index() as u32;
block_def.terminator.visit_uses(|u| {
visit_use(u, body, trees, cur, &def_block, &mut cross_blocks);
});
for &inst in block_def.insts.iter().rev() {
if is_tree_or_remat(inst, trees) {
continue;
}
if let ValueDef::Operator(_, args, _) = &body.values[inst] {
for &arg in &body.arg_pool[*args] {
visit_use(arg, body, trees, cur, &def_block, &mut cross_blocks);
}
}
}
}
cross_blocks
}
}
fn is_tree_or_remat(value: Value, trees: &Trees) -> bool {
trees.owner.contains_key(&value) || trees.remat.contains(&value)
}
fn visit_use(
value: Value,
body: &FunctionBody,
trees: &Trees,
cur_block: u32,
def_block: &[u32],
cross_blocks: &mut CrossBlockValues,
) {
let value = body.resolve_alias(value);
if let ValueDef::PickOutput(value, _, _) = body.values[value] {
visit_use(value, body, trees, cur_block, def_block, cross_blocks);
return;
}
if is_tree_or_remat(value, trees) {
if let ValueDef::Operator(_, args, _) = body.values[value] {
for &arg in &body.arg_pool[args] {
visit_use(arg, body, trees, cur_block, def_block, cross_blocks);
}
}
return;
}
if def_block[value.index()] != cur_block {
cross_blocks.insert(value);
}
}
impl Index<CrossBlockId> for CrossBlockValues {
type Output = Value;
fn index(&self, id: CrossBlockId) -> &Self::Output {
&self.values[id.0.get() as usize - 1]
}
}