Skip to main content

cubecl_spirv/
metadata.rs

1use cubecl_core::ir::Metadata;
2use cubecl_core::ir::{self as core};
3use rspirv::spirv::{MemoryAccess, Word};
4
5use crate::{SpirvCompiler, SpirvTarget, item::Item, value::Value};
6
7impl<T: SpirvTarget> SpirvCompiler<T> {
8    pub fn compile_meta(&mut self, meta: Metadata, out: Option<core::Value>, uniform: bool) {
9        let out = out.unwrap();
10        match meta {
11            Metadata::BufferLength { list } => {
12                let list = self.compile_value(list);
13                let out = self.compile_value(out);
14                self.buffer_length(&list, Some(&out), uniform);
15            }
16            Metadata::Stride { dim, list } => {
17                let list = self.compile_value(list);
18                let dim = self.compile_value(dim);
19                let out = self.compile_value(out);
20
21                let ty_id = out.item().id(self);
22                let out_id = out.id(self);
23                self.mark_uniformity(out_id, uniform);
24
25                let pos = self.ext_pos(&list);
26
27                let offs_offset = self.info.metadata.stride_offset_index(pos);
28                let offset = self.load_const_metadata(offs_offset, None, out.item());
29                let dim_id = self.read_as(&dim, &out.item());
30
31                let index = self.i_add(ty_id, None, offset, dim_id).unwrap();
32                self.mark_uniformity(index, uniform);
33                self.load_dyn_metadata(index, Some(out_id), out.item());
34                self.write(&out, out_id);
35            }
36            Metadata::Shape { dim, list } => {
37                let list = self.compile_value(list);
38                let dim = self.compile_value(dim);
39                let out = self.compile_value(out);
40
41                let ty_id = out.item().id(self);
42                let out_id = out.id(self);
43                self.mark_uniformity(out_id, uniform);
44
45                let pos = self.ext_pos(&list);
46
47                let offs_offset = self.info.metadata.shape_offset_index(pos);
48                let offset = self.load_const_metadata(offs_offset, None, out.item());
49                let dim_id = self.read_as(&dim, &out.item());
50
51                let index = self.i_add(ty_id, None, offset, dim_id).unwrap();
52                self.load_dyn_metadata(index, Some(out_id), out.item());
53                self.write(&out, out_id);
54            }
55        }
56    }
57
58    pub fn buffer_length(&mut self, val: &Value, out: Option<&Value>, uniform: bool) -> Word {
59        let out_id = out.map(|it| self.write_id(it));
60        if let Some(out_id) = out_id {
61            self.mark_uniformity(out_id, uniform);
62        }
63        let out_ty = out
64            .map(|it| it.item())
65            .unwrap_or_else(|| self.compile_type(self.addr_type.into()));
66
67        let val_id = val.id(self);
68        let position = self.buffer_pos(val_id);
69        let offset = self.info.metadata.buffer_len_index(position);
70        let id = self.load_const_metadata(offset, out_id, out_ty);
71
72        if let Some(out) = out {
73            self.debug_name(out_id.unwrap(), format!("buffer_len({position})"));
74            self.write(out, id);
75        }
76        id
77    }
78
79    fn buffer_pos(&self, id: Word) -> u32 {
80        self.state
81            .buffers
82            .iter()
83            .position(|buf| buf.id == id)
84            .expect("Buffer should exist") as u32
85    }
86
87    pub fn load_const_metadata(&mut self, index: u32, out: Option<Word>, ty: Item) -> Word {
88        self.insert_in_setup(|b| {
89            let align = ty.size();
90            let ty_id = ty.id(b);
91            let storage_class = T::info_storage_class(b);
92            let ptr_ty = Item::Pointer(storage_class, Box::new(ty)).id(b);
93            let info = b.state.info.unwrap().id;
94            let offset = b.const_u32(b.state.scalar_bindings.len() as u32);
95            let index = b.const_u32(index);
96            let info_ptr = b
97                .in_bounds_access_chain(ptr_ty, None, info, vec![offset, index])
98                .unwrap();
99            b.load(
100                ty_id,
101                out,
102                info_ptr,
103                Some(MemoryAccess::ALIGNED),
104                [align.into()],
105            )
106            .unwrap()
107        })
108    }
109
110    pub fn load_dyn_metadata(&mut self, index: Word, out: Option<Word>, ty: Item) -> Word {
111        let align = ty.size();
112        let ty_id = ty.id(self);
113        let storage_class = T::info_storage_class(self);
114        let ptr_ty = Item::Pointer(storage_class, Box::new(ty)).id(self);
115        let info = self.state.info.unwrap().id;
116        let offset = self.const_u32(self.state.scalar_bindings.len() as u32 + 1);
117        let info_ptr = self
118            .in_bounds_access_chain(ptr_ty, None, info, vec![offset, index])
119            .unwrap();
120        self.load(
121            ty_id,
122            out,
123            info_ptr,
124            Some(MemoryAccess::ALIGNED),
125            [align.into()],
126        )
127        .unwrap()
128    }
129
130    fn ext_pos(&mut self, val: &Value) -> u32 {
131        let val_id = val.id(self);
132        let pos = self.buffer_pos(val_id);
133        self.ext_meta_pos[pos as usize]
134    }
135}