cubecl_std/
reinterpret_slice.rs1use core::marker::PhantomData;
2
3use cubecl::prelude::*;
4use cubecl_core::{self as cubecl, ir::VectorSize, unexpanded};
5
6#[derive(CubeType)]
14pub struct ReinterpretSlice<'a, S: CubePrimitive, T: CubePrimitive> {
15 slice: &'a [S],
17
18 #[cube(comptime)]
19 vector_size: VectorSize,
20
21 #[cube(comptime)]
22 load_many: Option<usize>,
23
24 #[cube(comptime)]
25 _phantom: PhantomData<T>,
26}
27
28#[cube]
29impl<'a, S: CubePrimitive, T: CubePrimitive> ReinterpretSlice<'a, S, T> {
30 pub fn new(slice: &'a [S]) -> ReinterpretSlice<'a, S, T> {
31 let in_vector_size = slice.vector_size();
32 let source_size = S::Scalar::size();
33 let target_size = T::Scalar::size();
34 let (optimized_vector_size, load_many) = comptime!(optimize_vector_size(
35 source_size,
36 in_vector_size,
37 target_size
38 ));
39 match comptime!(optimized_vector_size) {
40 Some(vector_size) => {
41 let size!(N2) = vector_size;
42 let slice = slice.as_vectorized().with_vector_size::<N2>();
43
44 ReinterpretSlice::<'a, S, T> {
45 slice: unsafe { slice.downcast_unchecked() },
46 vector_size,
47 load_many,
48 _phantom: PhantomData,
49 }
50 }
51 None => ReinterpretSlice::<'a, S, T> {
52 slice: unsafe { slice.downcast_unchecked() },
53 vector_size: in_vector_size,
54 load_many,
55 _phantom: PhantomData,
56 },
57 }
58 }
59
60 pub fn read(&self, index: usize) -> T {
61 let size!(N) = self.vector_size;
62 let slice = self.slice.as_vectorized().with_vector_size::<N>();
63 match comptime!(self.load_many) {
64 Some(amount) => {
65 let first = index * amount;
66 let size!(N2) = comptime!(amount * self.vector_size);
67 let mut vector = Vector::<S::Scalar, N2>::empty();
68 #[unroll]
69 for k in 0..amount {
70 let elem = slice[first + k];
71 #[unroll]
72 for j in 0..self.vector_size {
73 vector.insert(comptime!(k * self.vector_size + j), elem.extract(j));
74 }
75 }
76 T::reinterpret(vector)
77 }
78 None => T::reinterpret(slice[index]),
79 }
80 }
81}
82
83#[derive(CubeType)]
91pub struct ReinterpretSliceMut<'a, S: CubePrimitive, T: CubePrimitive> {
92 slice: &'a mut [S],
93
94 #[cube(comptime)]
95 vector_size: VectorSize,
96
97 #[cube(comptime)]
98 load_many: Option<usize>,
99
100 #[cube(comptime)]
101 _phantom: PhantomData<T>,
102}
103
104#[cube]
105impl<'a, S: CubePrimitive, T: CubePrimitive> ReinterpretSliceMut<'a, S, T> {
106 pub fn new(slice: &'a mut [S]) -> ReinterpretSliceMut<'a, S, T> {
107 let in_vector_size = slice.vector_size();
108 let source_size = S::Scalar::size();
109 let target_size = T::Scalar::size();
110 let (optimized_vector_size, load_many) = comptime!(optimize_vector_size(
111 source_size,
112 in_vector_size,
113 target_size
114 ));
115 match comptime!(optimized_vector_size) {
116 Some(vector_size) => {
117 let size!(N2) = vector_size;
118 let slice = slice.as_vectorized_mut().with_vector_size_mut::<N2>();
119
120 ReinterpretSliceMut::<'a, S, T> {
121 slice: unsafe { slice.downcast_mut_unchecked() },
122 vector_size,
123 load_many,
124 _phantom: PhantomData,
125 }
126 }
127 None => ReinterpretSliceMut::<'a, S, T> {
128 slice: unsafe { slice.downcast_mut_unchecked() },
129 vector_size: in_vector_size,
130 load_many,
131 _phantom: PhantomData,
132 },
133 }
134 }
135
136 pub fn read(&self, index: usize) -> T {
137 let size!(N) = self.vector_size;
138 let slice = self.slice.as_vectorized().with_vector_size::<N>();
139 match comptime!(self.load_many) {
140 Some(amount) => {
141 let first = index * amount;
142 let size!(N2) = comptime!(amount * self.vector_size);
143 let mut vector = Vector::<S::Scalar, N2>::empty();
144 #[unroll]
145 for k in 0..amount {
146 let elem = slice[first + k];
147 #[unroll]
148 for j in 0..self.vector_size {
149 vector.insert(comptime!(k * self.vector_size + j), elem.extract(j));
150 }
151 }
152 T::reinterpret(vector)
153 }
154 None => T::reinterpret(slice[index]),
155 }
156 }
157
158 pub fn write(&mut self, index: usize, value: T) {
159 let size!(N) = self.vector_size;
160 let slice = self.slice.as_vectorized_mut().with_vector_size_mut::<N>();
161 let size!(N1) = S::reinterpret_vectorization::<T>();
162 let reinterpreted = Vector::<S::Scalar, N1>::reinterpret(value);
163 match comptime!(self.load_many) {
164 Some(amount) => {
165 let first = index * amount;
166 let reinterpreted_vec = reinterpreted.vector_size();
167 let vector_size = comptime!(reinterpreted_vec / amount);
168
169 #[unroll]
170 for k in 0..amount {
171 let mut vector = Vector::empty();
172 #[unroll]
173 for j in 0..vector_size {
174 vector.insert(j, reinterpreted.extract(k * vector_size + j));
175 }
176 slice[first + k] = vector;
177 }
178 }
179 None => slice[index] = Vector::cast_from(reinterpreted),
180 }
181 }
182}
183
184fn optimize_vector_size(
185 source_size: usize,
186 vector_size: VectorSize,
187 target_size: usize,
188) -> (Option<usize>, Option<usize>) {
189 let vector_source_size = source_size * vector_size;
190 match vector_source_size.cmp(&target_size) {
191 core::cmp::Ordering::Less => {
192 if !target_size.is_multiple_of(vector_source_size) {
193 panic!("incompatible number of bytes");
194 }
195
196 let ratio = target_size / vector_source_size;
197
198 (None, Some(ratio))
199 }
200 core::cmp::Ordering::Greater => {
201 if !vector_source_size.is_multiple_of(target_size) {
202 panic!("incompatible number of bytes");
203 }
204 let ratio = vector_source_size / target_size;
205
206 (Some(vector_size / ratio), None)
207 }
208 core::cmp::Ordering::Equal => (None, None),
209 }
210}
211
212pub fn size_of<S: CubePrimitive>() -> u32 {
213 unexpanded!()
214}
215
216pub mod size_of {
217 use super::*;
218 #[allow(unused, clippy::all)]
219 pub fn expand<S: CubePrimitive>(context: &Scope) -> u32 {
220 S::__expand_size(context) as u32
221 }
222}