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}