Skip to main content

cubecl_core/codegen/
metadata.rs

1//! Metadata helpers to easily get offsets etc.
2//!
3//! Conceptually, metadata is represented like this:
4//! ```ignore
5//! struct Metadata<const NUM_BUFS: usize, const NUM_EXT: usize> {
6//!     base: BaseMeta<NUM_BUFS>,
7//!     extended: ExtendedMeta<NUM_EXT>,
8//! }
9//!
10//! struct BaseMeta<const N: usize> {
11//!     buffer_lengths: [usize; N],
12//! }
13//!
14//! struct ExtendedMeta<const N: usize> {
15//!     shape_offsets: [usize; N],
16//!     stride_offsets: [usize; N],
17//!     shapes: [usize],
18//!     strides: [usize]
19//! }
20//! ```
21//! where `Vec` isn't an actual `Vec`, just a dynamically sized series of values.
22//!
23//! Ranks and lengths have a constant offset, while shapes/strides involve loading the tensor's
24//! offset, then adding `dim` to the offset to get each shape/stride.
25
26use alloc::vec::Vec;
27use bytemuck::Pod;
28use cubecl_ir::AddressType;
29use cubecl_zspace::{Shape, Strides};
30use num_traits::NumCast;
31
32// Metadata
33const BUFFER_LEN: u32 = 0;
34const BASE_LEN: u32 = 1;
35
36// Extended Metadata
37const SHAPE_OFFSETS: u32 = 0;
38const STRIDE_OFFSETS: u32 = 1;
39const EXTENDED_LEN: u32 = 2;
40
41/// Helper to calculate metadata offsets based on buffer count and position
42#[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/// Builder for a serialized metadata struct
94///
95/// Inputs/Outputs must be added in the same order they're defined in the bind group
96#[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    /// Add an array to a builder
113    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    /// Add a tensor to a builder
125    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    /// Build the final serialized metadata struct
166    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}