cubecl_core/codegen/
metadata.rs1use alloc::vec::Vec;
27use bytemuck::Pod;
28use cubecl_ir::AddressType;
29use cubecl_zspace::{Shape, Strides};
30use num_traits::NumCast;
31
32const BUFFER_LEN: u32 = 0;
34const BASE_LEN: u32 = 1;
35
36const SHAPE_OFFSETS: u32 = 0;
38const STRIDE_OFFSETS: u32 = 1;
39const EXTENDED_LEN: u32 = 2;
40
41#[derive(Clone, Copy, Debug, Default)]
43pub struct Metadata {
44 num_meta: u32,
45 num_extended_meta: u32,
46}
47
48impl Metadata {
49 pub fn new(num_meta: u32, num_extended_meta: u32) -> Self {
50 Self {
51 num_meta,
52 num_extended_meta,
53 }
54 }
55
56 fn offset_of(&self, id: u32) -> u32 {
57 self.num_meta * id
58 }
59
60 fn base_len(&self) -> u32 {
61 self.num_meta * BASE_LEN
62 }
63
64 pub fn static_len(&self) -> u32 {
65 self.num_meta * BASE_LEN + self.num_extended_meta * EXTENDED_LEN
66 }
67
68 pub fn num_meta(&self) -> u32 {
69 self.num_meta
70 }
71
72 pub fn num_extended_meta(&self) -> u32 {
73 self.num_extended_meta
74 }
75
76 fn offset_of_extended(&self, id: u32) -> u32 {
77 self.base_len() + self.num_extended_meta * id
78 }
79
80 pub fn buffer_len_index(&self, buffer_idx: u32) -> u32 {
81 self.offset_of(BUFFER_LEN) + buffer_idx
82 }
83
84 pub fn shape_offset_index(&self, buffer_idx: u32) -> u32 {
85 self.offset_of_extended(SHAPE_OFFSETS) + buffer_idx
86 }
87
88 pub fn stride_offset_index(&self, buffer_idx: u32) -> u32 {
89 self.offset_of_extended(STRIDE_OFFSETS) + buffer_idx
90 }
91}
92
93#[derive(Default)]
97pub struct MetadataBuilder {
98 state_32: State<u32>,
99 state_64: State<u64>,
100}
101
102#[derive(Default)]
103struct State<T: Pod> {
104 buffer_lens: Vec<T>,
105 shapes: Vec<T>,
106 strides: Vec<T>,
107
108 offsets: Vec<usize>,
109}
110
111impl MetadataBuilder {
112 pub fn register_buffer(&mut self, buffer_len: u64, address_type: AddressType) {
114 match address_type {
115 AddressType::U64 => {
116 self.state_64.buffer_lens.push(buffer_len);
117 }
118 AddressType::U32 => {
119 self.state_32.buffer_lens.push(buffer_len as u32);
120 }
121 }
122 }
123
124 pub fn register_tensor(
126 &mut self,
127 buffer_len: u64,
128 shape: Shape,
129 strides: Strides,
130 address_type: AddressType,
131 ) {
132 match address_type {
133 AddressType::U64 => {
134 let state = &mut self.state_64;
135 state.buffer_lens.push(buffer_len);
136 state.offsets.push(state.shapes.len());
137 state.shapes.extend(shape.iter().map(|s| *s as u64));
138 state.strides.extend(strides.iter().map(|s| *s as u64));
139 }
140 AddressType::U32 => {
141 let state = &mut self.state_32;
142 state.buffer_lens.push(buffer_len as u32);
143 state.offsets.push(state.shapes.len());
144 state.shapes.extend(shape.iter().map(|s| *s as u32));
145 state.strides.extend(strides.iter().map(|s| *s as u32));
146 }
147 }
148 }
149
150 pub fn static_len(&self, address_type: AddressType) -> usize {
151 let (base, ext) = match address_type {
152 AddressType::U32 => (self.state_32.buffer_lens.len(), self.state_32.offsets.len()),
153 AddressType::U64 => (self.state_64.buffer_lens.len(), self.state_64.offsets.len()),
154 };
155 base * BASE_LEN as usize + ext * EXTENDED_LEN as usize
156 }
157
158 pub fn dynamic_len(&self, address_type: AddressType) -> usize {
159 match address_type {
160 AddressType::U32 => self.state_32.shapes.len() + self.state_32.strides.len(),
161 AddressType::U64 => self.state_64.shapes.len() + self.state_64.strides.len(),
162 }
163 }
164
165 pub fn finish(&mut self, address_type: AddressType, out: (&mut [u64], &mut [u64])) {
167 fn finish_inner<T: Pod + NumCast>(state: &mut State<T>, out: (&mut [u64], &mut [u64])) {
168 let mut sized = bytemuck::cast_slice_mut::<u64, u8>(out.0);
169 let mut dynamic = bytemuck::cast_slice_mut::<u64, u8>(out.1);
170
171 {
172 let buffer_lens = bytemuck::cast_slice::<T, u8>(&state.buffer_lens);
173
174 sized[..buffer_lens.len()].copy_from_slice(buffer_lens);
175 sized = &mut sized[buffer_lens.len()..];
176 }
177
178 state.buffer_lens.clear();
179
180 let strides_offset_base = state.shapes.len();
181
182 for offs in state.offsets.iter() {
183 let offset = [T::from(*offs).unwrap()];
184 let bytes = bytemuck::cast_slice(&offset);
185 sized[..bytes.len()].copy_from_slice(bytes);
186 sized = &mut sized[size_of::<T>()..];
187 }
188
189 for offs in state.offsets.drain(..) {
190 let offset = [T::from(strides_offset_base + offs).unwrap()];
191 let bytes = bytemuck::cast_slice(&offset);
192 sized[..bytes.len()].copy_from_slice(bytes);
193 sized = &mut sized[size_of::<T>()..];
194 }
195
196 {
197 let shapes = bytemuck::cast_slice::<T, u8>(&state.shapes);
198 let strides = bytemuck::cast_slice::<T, u8>(&state.strides);
199
200 dynamic[..shapes.len()].copy_from_slice(shapes);
201 dynamic = &mut dynamic[shapes.len()..];
202
203 dynamic[..strides.len()].copy_from_slice(strides);
204 }
205
206 state.shapes.clear();
207 state.strides.clear();
208 }
209
210 match address_type {
211 AddressType::U32 => finish_inner(&mut self.state_32, out),
212 AddressType::U64 => finish_inner(&mut self.state_64, out),
213 }
214 }
215}