use alloc::{rc::Rc, vec, vec::Vec};
use core::cell::{RefCell, RefMut};
use cubecl_runtime::kernel::Visibility;
use derive_more::{Deref, DerefMut};
use cubecl_ir::{
AddressSpace, GlobalState, Id, Instruction, Memory, Operation, Scope, Value, ValueKind,
};
use hashbrown::{HashMap, HashSet};
use crate::post_processing::{
util::AtomicCounter,
visitor::{InstructionVisitor, Visitor},
};
#[derive(Debug, Clone, Default)]
pub struct GlobalAnalyses {
ptr_source: Rc<RefCell<PointerSource>>,
used_values: Rc<RefCell<UsedValues>>,
}
impl GlobalAnalyses {
pub fn recalculate_pointer_source(&self, scope: &Scope) {
*self.ptr_source.borrow_mut() = PointerSource::new(scope);
}
pub fn recalculate_used_values(&self, scope: &Scope) {
let mut used_values = UsedValues::default();
used_values.visit_scope(scope, self, &AtomicCounter::new(0));
*self.used_values.borrow_mut() = used_values;
}
pub fn ptr_source(&self) -> RefMut<'_, PointerSource> {
self.ptr_source.borrow_mut()
}
pub fn used_values(&self) -> RefMut<'_, UsedValues> {
self.used_values.borrow_mut()
}
}
#[derive(Default, Debug, Deref, DerefMut)]
pub struct UsedValues {
used: HashSet<Value>,
}
impl InstructionVisitor for UsedValues {
fn visit_instruction(
&mut self,
mut inst: Instruction,
_global_state: &GlobalState,
analyses: &GlobalAnalyses,
_changes: &AtomicCounter,
) -> Vec<Instruction> {
let mut visitor = Visitor(self);
visitor.visit_operation(&mut inst.operation, analyses, |this, val| {
this.used.insert(*val);
});
vec![inst]
}
}
#[derive(Debug, Default, Deref, DerefMut)]
pub struct PointerSource {
sources: HashMap<ValueKind, Value>,
}
impl PointerSource {
pub fn new(scope: &Scope) -> Self {
let mut this = PointerSource::default();
this.visit_scope(scope, &GlobalAnalyses::default(), &AtomicCounter::new(0));
this
}
}
impl InstructionVisitor for PointerSource {
fn visit_instruction(
&mut self,
inst: Instruction,
_global_state: &GlobalState,
_analyses: &GlobalAnalyses,
_changes: &AtomicCounter,
) -> Vec<Instruction> {
match &inst.operation {
Operation::Copy(val) if val.ty.is_ptr() && inst.out().ty.is_ptr() => {
if let Some(source) = self.sources.get(&val.kind) {
self.sources.insert(inst.out().kind, *source);
}
}
Operation::Memory(Memory::Index(index_operands)) => {
self.sources.insert(inst.out().kind, index_operands.list);
}
_ => {}
}
vec![inst]
}
}
#[derive(Debug, Deref, Default)]
pub struct BufferVisibility {
buffers: Vec<Visibility>,
}
impl From<BufferVisibility> for Vec<Visibility> {
fn from(value: BufferVisibility) -> Self {
value.buffers
}
}
impl BufferVisibility {
pub fn new(scope: &Scope, analyses: &GlobalAnalyses) -> Self {
let mut this = BufferVisibility::default();
this.visit_scope(scope, analyses, &AtomicCounter::new(0));
this
}
}
impl InstructionVisitor for BufferVisibility {
fn visit_instruction(
&mut self,
mut inst: Instruction,
_global_state: &GlobalState,
analyses: &GlobalAnalyses,
_changes: &AtomicCounter,
) -> Vec<Instruction> {
let mut visitor = Visitor(self);
visitor.visit_instruction(
&mut inst,
analyses,
|this, val| {
if let Some(id) = global_buffer_id(val) {
this.set_readable(id as usize);
}
},
|this, val| {
if let Some(id) = global_buffer_id(val) {
this.set_writable(id as usize);
}
},
);
vec![inst]
}
}
impl BufferVisibility {
fn set_readable(&mut self, id: usize) {
if self.buffers.len() <= id {
self.buffers.resize(id + 1, Visibility::Read);
}
}
fn set_writable(&mut self, id: usize) {
if self.buffers.len() <= id {
self.buffers.resize(id + 1, Visibility::Read);
}
self.buffers[id] = Visibility::ReadWrite;
}
}
fn global_buffer_id(variable: &Value) -> Option<Id> {
match variable.address_space() {
AddressSpace::Global(id) => Some(id),
_ => None,
}
}