cubecl-spirv 0.11.0-pre.1

SPIR-V compiler for CubeCL
Documentation
use cubecl_core::ir::Metadata;
use cubecl_core::ir::{self as core};
use rspirv::spirv::{MemoryAccess, Word};

use crate::{SpirvCompiler, SpirvTarget, item::Item, value::Value};

impl<T: SpirvTarget> SpirvCompiler<T> {
    pub fn compile_meta(&mut self, meta: Metadata, out: Option<core::Value>, uniform: bool) {
        let out = out.unwrap();
        match meta {
            Metadata::BufferLength { list } => {
                let list = self.compile_value(list);
                let out = self.compile_value(out);
                self.buffer_length(&list, Some(&out), uniform);
            }
            Metadata::Stride { dim, list } => {
                let list = self.compile_value(list);
                let dim = self.compile_value(dim);
                let out = self.compile_value(out);

                let ty_id = out.item().id(self);
                let out_id = out.id(self);
                self.mark_uniformity(out_id, uniform);

                let pos = self.ext_pos(&list);

                let offs_offset = self.info.metadata.stride_offset_index(pos);
                let offset = self.load_const_metadata(offs_offset, None, out.item());
                let dim_id = self.read_as(&dim, &out.item());

                let index = self.i_add(ty_id, None, offset, dim_id).unwrap();
                self.mark_uniformity(index, uniform);
                self.load_dyn_metadata(index, Some(out_id), out.item());
                self.write(&out, out_id);
            }
            Metadata::Shape { dim, list } => {
                let list = self.compile_value(list);
                let dim = self.compile_value(dim);
                let out = self.compile_value(out);

                let ty_id = out.item().id(self);
                let out_id = out.id(self);
                self.mark_uniformity(out_id, uniform);

                let pos = self.ext_pos(&list);

                let offs_offset = self.info.metadata.shape_offset_index(pos);
                let offset = self.load_const_metadata(offs_offset, None, out.item());
                let dim_id = self.read_as(&dim, &out.item());

                let index = self.i_add(ty_id, None, offset, dim_id).unwrap();
                self.load_dyn_metadata(index, Some(out_id), out.item());
                self.write(&out, out_id);
            }
        }
    }

    pub fn buffer_length(&mut self, val: &Value, out: Option<&Value>, uniform: bool) -> Word {
        let out_id = out.map(|it| self.write_id(it));
        if let Some(out_id) = out_id {
            self.mark_uniformity(out_id, uniform);
        }
        let out_ty = out
            .map(|it| it.item())
            .unwrap_or_else(|| self.compile_type(self.addr_type.into()));

        let val_id = val.id(self);
        let position = self.buffer_pos(val_id);
        let offset = self.info.metadata.buffer_len_index(position);
        let id = self.load_const_metadata(offset, out_id, out_ty);

        if let Some(out) = out {
            self.debug_name(out_id.unwrap(), format!("buffer_len({position})"));
            self.write(out, id);
        }
        id
    }

    fn buffer_pos(&self, id: Word) -> u32 {
        self.state
            .buffers
            .iter()
            .position(|buf| buf.id == id)
            .expect("Buffer should exist") as u32
    }

    pub fn load_const_metadata(&mut self, index: u32, out: Option<Word>, ty: Item) -> Word {
        self.insert_in_setup(|b| {
            let align = ty.size();
            let ty_id = ty.id(b);
            let storage_class = T::info_storage_class(b);
            let ptr_ty = Item::Pointer(storage_class, Box::new(ty)).id(b);
            let info = b.state.info.unwrap().id;
            let offset = b.const_u32(b.state.scalar_bindings.len() as u32);
            let index = b.const_u32(index);
            let info_ptr = b
                .in_bounds_access_chain(ptr_ty, None, info, vec![offset, index])
                .unwrap();
            b.load(
                ty_id,
                out,
                info_ptr,
                Some(MemoryAccess::ALIGNED),
                [align.into()],
            )
            .unwrap()
        })
    }

    pub fn load_dyn_metadata(&mut self, index: Word, out: Option<Word>, ty: Item) -> Word {
        let align = ty.size();
        let ty_id = ty.id(self);
        let storage_class = T::info_storage_class(self);
        let ptr_ty = Item::Pointer(storage_class, Box::new(ty)).id(self);
        let info = self.state.info.unwrap().id;
        let offset = self.const_u32(self.state.scalar_bindings.len() as u32 + 1);
        let info_ptr = self
            .in_bounds_access_chain(ptr_ty, None, info, vec![offset, index])
            .unwrap();
        self.load(
            ty_id,
            out,
            info_ptr,
            Some(MemoryAccess::ALIGNED),
            [align.into()],
        )
        .unwrap()
    }

    fn ext_pos(&mut self, val: &Value) -> u32 {
        let val_id = val.id(self);
        let pos = self.buffer_pos(val_id);
        self.ext_meta_pos[pos as usize]
    }
}