use alloc::collections::vec_deque::VecDeque;
use cubecl_ir::{AddressSpace, Id, Value, ValueKind};
use hashbrown::{HashMap, HashSet};
use petgraph::graph::NodeIndex;
use crate::{Function, GlobalState, analyses::post_order::PostOrder};
use super::Analysis;
type LivePredicate = fn(&Function, &Value) -> Option<Id>;
pub struct Liveness {
live_vars: HashMap<NodeIndex, HashSet<Id>>,
}
#[derive(Debug, Clone)]
struct BlockSets {
generated: HashSet<Id>,
kill: HashSet<Id>,
}
struct State {
worklist: VecDeque<NodeIndex>,
block_sets: HashMap<NodeIndex, BlockSets>,
}
impl Analysis for Liveness {
fn init(func: &mut Function, state: &GlobalState) -> Self {
Self {
live_vars: compute_liveness(func, state, Function::local_variable_id),
}
}
}
impl Liveness {
pub fn empty(func: &Function) -> Self {
let live_vars = func
.node_ids()
.iter()
.map(|it| (*it, HashSet::new()))
.collect();
Self { live_vars }
}
pub fn at_block(&self, block: NodeIndex) -> &HashSet<Id> {
&self.live_vars[&block]
}
pub fn is_dead(&self, node: NodeIndex, var: Id) -> bool {
!self.at_block(node).contains(&var)
}
}
pub struct MemoryLiveness {
live_vars: HashMap<NodeIndex, HashSet<Id>>,
}
impl Analysis for MemoryLiveness {
fn init(func: &mut Function, state: &GlobalState) -> Self {
Self {
live_vars: compute_liveness(func, state, Function::local_memory_id),
}
}
}
impl MemoryLiveness {
pub fn empty(func: &Function) -> Self {
let live_vars = func
.node_ids()
.iter()
.map(|it| (*it, HashSet::new()))
.collect();
Self { live_vars }
}
pub fn at_block(&self, block: NodeIndex) -> &HashSet<Id> {
&self.live_vars[&block]
}
pub fn is_dead(&self, node: NodeIndex, var: Id) -> bool {
!self.at_block(node).contains(&var)
}
}
fn compute_liveness(
func: &mut Function,
global_state: &GlobalState,
pred: LivePredicate,
) -> HashMap<NodeIndex, HashSet<Id>> {
let mut live_vars: HashMap<NodeIndex, HashSet<Id>> = func
.node_ids()
.iter()
.map(|it| (*it, HashSet::new()))
.collect();
let mut state = State {
worklist: VecDeque::from(func.analysis::<PostOrder>(global_state).forward()),
block_sets: HashMap::new(),
};
while let Some(block) = state.worklist.pop_front() {
analyze_block(func, global_state, block, &mut state, &mut live_vars, pred);
}
live_vars
}
fn analyze_block(
func: &mut Function,
global_state: &GlobalState,
block: NodeIndex,
state: &mut State,
live_vars: &mut HashMap<NodeIndex, HashSet<Id>>,
pred: LivePredicate,
) {
let BlockSets { generated, kill } = block_sets(func, global_state, block, state, pred);
let mut block_live = generated.clone();
for successor in func.successors(block) {
let successor = &live_vars[&successor];
block_live.extend(successor.difference(kill));
}
if block_live != live_vars[&block] {
state.worklist.extend(func.predecessors(block));
live_vars.insert(block, block_live);
}
}
fn block_sets<'a>(
func: &mut Function,
global_state: &GlobalState,
block: NodeIndex,
state: &'a mut State,
pred: LivePredicate,
) -> &'a BlockSets {
let block_sets = state.block_sets.entry(block);
block_sets.or_insert_with(|| calculate_block_sets(func, global_state, block, pred))
}
fn calculate_block_sets(
func: &mut Function,
state: &GlobalState,
block: NodeIndex,
pred: LivePredicate,
) -> BlockSets {
let mut generated = HashSet::new();
let mut kill = HashSet::new();
let ops = func[block].ops.clone();
let control_flow = func[block].control_flow.clone();
func.visit_control_flow(&mut control_flow.borrow_mut(), |func, val| {
if let Some(id) = pred(func, val) {
generated.insert(id);
}
});
let mut ops = ops.borrow().clone();
for op in ops.values_mut().rev() {
func.visit_out(&mut op.out, |func, val| {
if let Some(id) = pred(func, val) {
kill.insert(id);
generated.remove(&id);
}
});
func.visit_operation(state, &mut op.operation, |func, val| {
if let Some(id) = pred(func, val) {
generated.insert(id);
}
});
}
BlockSets { generated, kill }
}
impl Function {
pub fn local_variable_id(&self, value: &Value) -> Option<Id> {
match value.kind {
ValueKind::Value { id } if self.destructurable_local_memories().contains_key(&id) => {
Some(id)
}
_ => None,
}
}
pub fn local_memory_id(&self, value: &Value) -> Option<Id> {
match value.kind {
ValueKind::Value { id } => match self.memories.get(&id) {
Some(mem)
if matches!(mem.address_space, AddressSpace::Local)
&& !mem.value_ty.is_atomic() =>
{
Some(id)
}
_ => None,
},
_ => None,
}
}
}
pub mod shared {
use alloc::vec::Vec;
use cubecl_ir::{AddressSpace, Marker, Operation, Type, Value, ValueKind};
use crate::{MemoryBlock, Uniformity};
use super::*;
#[derive(Debug, Clone, Copy, PartialEq, Eq, PartialOrd, Ord, Hash)]
pub struct SmemAllocation {
pub id: Id,
pub smem: MemoryBlock,
pub offset: usize,
}
#[derive(Default, Clone)]
pub struct SharedLiveness {
live_vars: HashMap<NodeIndex, HashSet<Id>>,
pub shared_memories: HashMap<Id, MemoryBlock>,
pub allocations: HashMap<Id, SmemAllocation>,
}
impl Analysis for SharedLiveness {
fn init(func: &mut Function, state: &GlobalState) -> Self {
let mut this = Self::empty(func);
this.analyze_liveness(func, state);
this.uniformize_liveness(func, state);
this.allocate_slices(func);
this
}
}
impl SharedLiveness {
pub fn empty(func: &Function) -> Self {
let live_vars = func
.node_ids()
.iter()
.map(|it| (*it, HashSet::new()))
.collect();
Self {
live_vars,
shared_memories: Default::default(),
allocations: Default::default(),
}
}
pub fn at_block(&self, block: NodeIndex) -> &HashSet<Id> {
&self.live_vars[&block]
}
fn is_live(&self, node: NodeIndex, var: Id) -> bool {
self.at_block(node).contains(&var)
}
fn analyze_liveness(&mut self, func: &mut Function, global_state: &GlobalState) {
let mut state = State {
worklist: VecDeque::from(func.analysis::<PostOrder>(global_state).reverse()),
block_sets: HashMap::new(),
};
while let Some(block) = state.worklist.pop_front() {
self.analyze_block(func, global_state, block, &mut state);
}
}
fn uniformize_liveness(&mut self, func: &mut Function, global_state: &GlobalState) {
let mut state = State {
worklist: VecDeque::from(func.analysis::<PostOrder>(global_state).forward()),
block_sets: HashMap::new(),
};
while let Some(block) = state.worklist.pop_front() {
self.uniformize_block(func, global_state, block, &mut state);
}
}
fn allocate_slices(&mut self, func: &mut Function) {
for block in func.node_ids() {
for live_smem in self.at_block(block).clone() {
if !self.allocations.contains_key(&live_smem) {
let smem = self.shared_memories[&live_smem];
let offset = self.allocate_slice(block, smem.size(), smem.alignment);
self.allocations.insert(
live_smem,
SmemAllocation {
id: live_smem,
smem,
offset,
},
);
}
}
}
}
fn allocate_slice(&mut self, block: NodeIndex, size: usize, align: usize) -> usize {
let live_slices = self.live_slices(block);
if live_slices.is_empty() {
return 0;
}
for i in 0..live_slices.len() - 1 {
let slice_0 = &live_slices[i];
let slice_1 = &live_slices[i + 1];
let end_0 = (slice_0.offset + slice_0.smem.size()).next_multiple_of(align);
let gap = slice_1.offset.saturating_sub(end_0);
if gap >= size {
return end_0;
}
}
let last_slice = &live_slices[live_slices.len() - 1];
(last_slice.offset + last_slice.smem.size()).next_multiple_of(align)
}
fn live_slices(&mut self, block: NodeIndex) -> Vec<SmemAllocation> {
let mut live_slices = self
.allocations
.iter()
.filter(|(k, _)| self.is_live(block, **k))
.map(|it| *it.1)
.collect::<Vec<_>>();
live_slices.sort_by_key(|it| it.offset);
live_slices
}
fn analyze_block(
&mut self,
func: &mut Function,
global_state: &GlobalState,
block: NodeIndex,
state: &mut State,
) {
let BlockSets { generated, kill } = self.block_sets(func, global_state, block, state);
let mut live_vars = generated.clone();
for predecessor in func.predecessors(block) {
let predecessor = &self.live_vars[&predecessor];
live_vars.extend(predecessor.difference(kill));
}
if live_vars != self.live_vars[&block] {
state.worklist.extend(func.successors(block));
self.live_vars.insert(block, live_vars);
}
}
fn uniformize_block(
&mut self,
func: &mut Function,
global_state: &GlobalState,
block: NodeIndex,
state: &mut State,
) {
let mut live_vars = self.live_vars[&block].clone();
let uniformity = func.analysis::<Uniformity>(global_state);
for successor in func.successors(block) {
if !uniformity.is_block_uniform(successor) {
let successor = &self.live_vars[&successor];
live_vars.extend(successor);
}
}
if live_vars != self.live_vars[&block] {
state.worklist.extend(func.predecessors(block));
self.live_vars.insert(block, live_vars);
}
}
fn block_sets<'a>(
&mut self,
func: &mut Function,
global_state: &GlobalState,
block: NodeIndex,
state: &'a mut State,
) -> &'a BlockSets {
let block_sets = state.block_sets.entry(block);
block_sets.or_insert_with(|| self.calculate_block_sets(func, global_state, block))
}
fn calculate_block_sets(
&mut self,
func: &mut Function,
state: &GlobalState,
block: NodeIndex,
) -> BlockSets {
let mut generated = HashSet::new();
let mut kill = HashSet::new();
let ops = func[block].ops.clone();
for op in ops.borrow_mut().values_mut() {
func.visit_out(&mut op.out, |func, var| {
if let Some((id, smem)) = shared_memory(func, var) {
generated.insert(id);
self.shared_memories.insert(id, smem);
}
});
func.visit_operation(state, &mut op.operation, |func, var| {
if let Some((id, smem)) = shared_memory(func, var) {
generated.insert(id);
self.shared_memories.insert(id, smem);
}
});
if let Operation::Marker(Marker::Free(Value {
ty: Type::Pointer(_, AddressSpace::Shared),
kind: ValueKind::Value { id, .. },
..
})) = &op.operation
{
kill.insert(*id);
generated.remove(id);
}
}
BlockSets { generated, kill }
}
}
fn shared_memory(func: &Function, var: &Value) -> Option<(Id, MemoryBlock)> {
match var.kind {
ValueKind::Value { id } => {
if let Some(mem) = func.memories.get(&id)
&& matches!(mem.address_space, AddressSpace::Shared)
{
Some((id, *mem))
} else {
None
}
}
_ => None,
}
}
}
mod captures {
use cubecl_ir::Value;
use super::*;
pub struct Captures {
live_vars: HashMap<NodeIndex, HashSet<Value>>,
}
#[derive(Clone)]
struct BlockSets {
generated: HashSet<Value>,
kill: HashSet<Value>,
}
struct State {
worklist: VecDeque<NodeIndex>,
block_sets: HashMap<NodeIndex, BlockSets>,
}
impl Analysis for Captures {
fn init(func: &mut Function, state: &GlobalState) -> Self {
let mut this = Self::empty(func);
this.analyze_liveness(func, state);
this
}
}
impl Captures {
pub fn empty(func: &Function) -> Self {
let live_vars = func
.node_ids()
.iter()
.map(|it| (*it, HashSet::new()))
.collect();
Self { live_vars }
}
pub fn at_block(&self, block: NodeIndex) -> &HashSet<Value> {
&self.live_vars[&block]
}
pub fn analyze_liveness(&mut self, func: &mut Function, global_state: &GlobalState) {
let mut state = State {
worklist: VecDeque::from(func.analysis::<PostOrder>(global_state).forward()),
block_sets: HashMap::new(),
};
while let Some(block) = state.worklist.pop_front() {
self.analyze_block(func, global_state, block, &mut state);
}
}
fn analyze_block(
&mut self,
func: &mut Function,
global_state: &GlobalState,
block: NodeIndex,
state: &mut State,
) {
let BlockSets { generated, kill } = block_sets(func, global_state, block, state);
let mut live_vars = generated.clone();
for successor in func.successors(block) {
let successor = &self.live_vars[&successor];
live_vars.extend(successor.difference(kill));
}
if live_vars != self.live_vars[&block] {
state.worklist.extend(func.predecessors(block));
self.live_vars.insert(block, live_vars);
}
}
}
fn block_sets<'a>(
func: &mut Function,
global_state: &GlobalState,
block: NodeIndex,
state: &'a mut State,
) -> &'a BlockSets {
let block_sets = state.block_sets.entry(block);
block_sets.or_insert_with(|| calculate_block_sets(func, global_state, block))
}
fn calculate_block_sets(
func: &mut Function,
state: &GlobalState,
block: NodeIndex,
) -> BlockSets {
let mut generated = HashSet::new();
let mut kill = HashSet::new();
let ops = func[block].ops.clone();
let control_flow = func[block].control_flow.clone();
func.visit_control_flow(&mut control_flow.borrow_mut(), |_, var| {
generated.insert(*var);
});
for inst in ops.borrow_mut().values_mut().rev() {
func.visit_out(&mut inst.out, |_, var| {
kill.insert(*var);
generated.remove(var);
});
func.visit_operation(state, &mut inst.operation, |_, var| {
generated.insert(*var);
});
}
BlockSets { generated, kill }
}
}
pub use captures::Captures;