cubecl-core 0.11.0-pre.1

CubeCL core create
Documentation
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 {
    // There for consistency and in case we make it more granular like it is in SPIR-V
    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,
    }
}