Skip to main content

cubecl_cpp/shared/
item.rs

1use std::fmt::Display;
2
3use cubecl_core::{
4    ir::{BarrierLevel, Intern},
5    prelude::Visibility,
6};
7
8use crate::shared::FragmentType;
9
10use super::{Dialect, Elem};
11
12#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
13pub enum Item<D: Dialect> {
14    Scalar(Elem<D>),
15    Vector(Intern<Item<D>>, usize),
16    NativeVector(Elem<D>, usize),
17    Atomic(Intern<Item<D>>),
18    Pointer(Intern<Item<D>>, PointerClass),
19    Array(Intern<Item<D>>, usize),
20    DynamicArray(Intern<Item<D>>),
21    Fragment(FragmentType<D>),
22    Barrier(BarrierLevel),
23    BarrierToken(BarrierLevel),
24    TensorMap,
25}
26
27#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
28pub enum PointerClass {
29    Global(Visibility),
30    Shared,
31    Local,
32}
33
34impl<D: Dialect> Display for Item<D> {
35    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
36        D::compile_item(f, self)
37    }
38}
39
40impl<D: Dialect> Item<D> {
41    pub fn new(inner: Elem<D>, vectorization: usize) -> Self {
42        let scalar = Self::Scalar(inner);
43        if vectorization > 1 {
44            Self::Vector(scalar.intern(), vectorization)
45        } else {
46            scalar
47        }
48    }
49
50    pub fn intern(self) -> Intern<Self> {
51        Intern::new(self)
52    }
53
54    /// Type of the value, unwrapping pointers
55    pub fn value_ty(&self) -> &Item<D> {
56        match self {
57            Item::Pointer(inner, _) => inner.value_ty(),
58            Item::Array(inner, _) => inner.value_ty(),
59            Item::DynamicArray(inner) => inner.value_ty(),
60            other => other,
61        }
62    }
63
64    /// Type of the pointer returned when indexing
65    pub fn value_ptr(&self) -> Item<D> {
66        match self {
67            Item::Pointer(inner, class) => Item::Pointer(inner.value_ptr().intern(), *class),
68            Item::Array(inner, _) => inner.value_ptr(),
69            Item::DynamicArray(inner) => inner.value_ptr(),
70            other => *other,
71        }
72    }
73
74    pub fn elem(&self) -> &Elem<D> {
75        match self {
76            Item::Scalar(elem) | Item::NativeVector(elem, _) => elem,
77            Item::Vector(item, _)
78            | Item::Atomic(item)
79            | Item::Pointer(item, _)
80            | Item::Array(item, _)
81            | Item::DynamicArray(item) => item.elem(),
82            Item::Fragment(frag) => &frag.elem,
83            Item::BarrierToken(..) => &Elem::None,
84            Item::TensorMap => &Elem::None,
85            Item::Barrier(..) => &Elem::None,
86        }
87    }
88
89    pub fn with_elem(&self, elem: Elem<D>) -> Self {
90        match self {
91            Item::Scalar(_) => Item::Scalar(elem),
92            Item::NativeVector(_, vectorization) => Item::NativeVector(elem, *vectorization),
93            Item::Vector(inner, vectorization) => {
94                Item::Vector(inner.with_elem(elem).intern(), *vectorization)
95            }
96            Item::Atomic(inner) => Item::Atomic(inner.with_elem(elem).intern()),
97            Item::Pointer(inner, class) => Item::Pointer(inner.with_elem(elem).intern(), *class),
98            Item::Array(inner, size) => Item::Array(inner.with_elem(elem).intern(), *size),
99            Item::DynamicArray(inner) => Item::DynamicArray(inner.with_elem(elem).intern()),
100            Item::Fragment(fragment_type) => {
101                let mut frag_ty = *fragment_type;
102                frag_ty.elem = elem;
103                Item::Fragment(frag_ty)
104            }
105            Item::Barrier(..) => panic!("Can't set elem of barrier"),
106            Item::BarrierToken(..) => panic!("Can't set elem of barrier token"),
107            Item::TensorMap => panic!("Can't set elem of tensor map"),
108        }
109    }
110
111    pub fn as_scalar(&self) -> Self {
112        match self {
113            Item::Scalar(_) => *self,
114            Item::NativeVector(elem, _) => Item::Scalar(*elem),
115            Item::Vector(inner, _) => inner.as_scalar(),
116            Item::Atomic(inner) => Item::Atomic(inner.as_scalar().intern()),
117            Item::Pointer(inner, class) => Item::Pointer(inner.as_scalar().intern(), *class),
118            Item::Array(inner, size) => Item::Array(inner.as_scalar().intern(), *size),
119            Item::DynamicArray(inner) => Item::DynamicArray(inner.as_scalar().intern()),
120            Item::Fragment(fragment_type) => Item::Fragment(*fragment_type),
121            Item::Barrier(..) => panic!("Can't set elem of barrier"),
122            Item::BarrierToken(..) => panic!("Can't get elem of barrier token"),
123            Item::TensorMap => panic!("Can't get elem of tensor map"),
124        }
125    }
126
127    pub fn vectorization(&self) -> usize {
128        match self {
129            Item::Vector(_, vectorization) | Item::NativeVector(_, vectorization) => *vectorization,
130            Item::Scalar(_) => 1,
131            Item::Atomic(inner)
132            | Item::Pointer(inner, _)
133            | Item::Array(inner, _)
134            | Item::DynamicArray(inner) => inner.vectorization(),
135            Item::Fragment(_) => 1,
136            Item::Barrier(..) | Item::BarrierToken(..) => 1,
137            Item::TensorMap => 1,
138        }
139    }
140
141    pub fn size(&self) -> usize {
142        match self {
143            Item::Scalar(elem) => elem.size(),
144            Item::Vector(inner, vectorization) => inner.size() * vectorization,
145            Item::NativeVector(elem, vectorization) => elem.size() * vectorization,
146            Item::Atomic(inner) => inner.size(),
147            Item::Array(inner, size) => inner.size() * *size,
148            Item::DynamicArray(inner) => inner.size(),
149            Item::Pointer(..) => size_of::<u64>(),
150            Item::Fragment(_) => panic!("Can't read size of fragment"),
151            Item::Barrier(..) => size_of::<u64>(),
152            Item::BarrierToken(..) => size_of::<u64>(),
153            Item::TensorMap => 128,
154        }
155    }
156
157    pub fn can_be_optimized(&self) -> bool {
158        D::item_can_be_optimized()
159    }
160
161    pub fn is_optimized(&self) -> bool {
162        matches!(
163            self.elem(),
164            Elem::F16x2 | Elem::BF16x2 | Elem::FP4x2(_) | Elem::FP6x2(_) | Elem::FP8x2(_)
165        )
166    }
167
168    pub fn optimized(&self) -> Item<D> {
169        if !self.can_be_optimized() {
170            return *self;
171        }
172
173        match self {
174            Item::Scalar(elem) => Item::Scalar(*elem),
175            Item::Vector(inner, _) if !matches!(**inner, Item::Scalar(_)) => inner.optimized(),
176            Item::Vector(inner, vectorization) => match Self::optimized_elem(*inner.elem()) {
177                Some(elem) => Item::new(elem, *vectorization / elem.packing_factor()),
178                None => Item::Vector(*inner, *vectorization),
179            },
180            Item::NativeVector(elem, vectorization) => match Self::optimized_elem(*elem) {
181                Some(elem) if *vectorization > elem.packing_factor() => {
182                    Item::NativeVector(elem, *vectorization / elem.packing_factor())
183                }
184                Some(elem) => Item::Scalar(elem),
185                None => Item::NativeVector(*elem, *vectorization),
186            },
187            Item::Atomic(inner) => Item::Atomic(inner.optimized().intern()),
188            Item::Pointer(inner, pointer_class) => {
189                Item::Pointer(inner.optimized().intern(), *pointer_class)
190            }
191            Item::Array(inner, size) => Item::Array(inner.optimized().intern(), *size),
192            Item::DynamicArray(inner) => Item::DynamicArray(inner.optimized().intern()),
193            Item::Fragment(fragment_type) => Item::Fragment(*fragment_type),
194            Item::Barrier(barrier_level) => Item::Barrier(*barrier_level),
195            Item::BarrierToken(barrier_level) => Item::BarrierToken(*barrier_level),
196            Item::TensorMap => Item::TensorMap,
197        }
198    }
199
200    fn optimized_elem(elem: Elem<D>) -> Option<Elem<D>> {
201        match elem {
202            Elem::F16 => Some(Elem::F16x2),
203            Elem::BF16 => Some(Elem::BF16x2),
204            Elem::FP4(kind) => Some(Elem::FP4x2(kind)),
205            Elem::FP6(kind) => Some(Elem::FP6x2(kind)),
206            Elem::FP8(kind) => Some(Elem::FP8x2(kind)),
207            _ => None,
208        }
209    }
210
211    /// Get the number of values packed into a single storage element. (i.e. `f16x2 -> 2`)
212    pub fn packing_factor(&self) -> usize {
213        self.elem().packing_factor()
214    }
215
216    pub fn de_optimized(&self) -> Self {
217        match self {
218            Item::Scalar(elem) => Item::Scalar(*elem),
219            Item::Vector(inner, _) if !matches!(**inner, Item::Scalar(_)) => inner.de_optimized(),
220            Item::Vector(inner, vectorization) => match Self::deoptimized_elem(*inner.elem()) {
221                Some(elem) => Item::Vector(
222                    Item::Scalar(elem).intern(),
223                    *vectorization * inner.elem().packing_factor(),
224                ),
225                None => Item::Vector(*inner, *vectorization),
226            },
227            Item::NativeVector(elem, vectorization) => match Self::deoptimized_elem(*elem) {
228                Some(new_elem) => {
229                    Item::NativeVector(new_elem, *vectorization * elem.packing_factor())
230                }
231                None => Item::NativeVector(*elem, *vectorization),
232            },
233            Item::Atomic(inner) => Item::Atomic(inner.de_optimized().intern()),
234            Item::Pointer(inner, pointer_class) => {
235                Item::Pointer(inner.de_optimized().intern(), *pointer_class)
236            }
237            Item::Array(inner, size) => Item::Array(inner.de_optimized().intern(), *size),
238            Item::DynamicArray(inner) => Item::DynamicArray(inner.de_optimized().intern()),
239            Item::Fragment(fragment_type) => Item::Fragment(*fragment_type),
240            Item::Barrier(barrier_level) => Item::Barrier(*barrier_level),
241            Item::BarrierToken(barrier_level) => Item::BarrierToken(*barrier_level),
242            Item::TensorMap => Item::TensorMap,
243        }
244    }
245
246    fn deoptimized_elem(elem: Elem<D>) -> Option<Elem<D>> {
247        match elem {
248            Elem::F16x2 => Some(Elem::F16),
249            Elem::BF16x2 => Some(Elem::BF16),
250            Elem::FP4x2(kind) => Some(Elem::FP4(kind)),
251            Elem::FP6x2(kind) => Some(Elem::FP6(kind)),
252            Elem::FP8x2(kind) => Some(Elem::FP8(kind)),
253            _ => None,
254        }
255    }
256
257    pub fn is_ptr(&self) -> bool {
258        matches!(self, Item::Pointer(..))
259    }
260
261    pub fn is_array(&self) -> bool {
262        matches!(self, Item::Array(..))
263    }
264
265    pub fn is_array_like(&self) -> bool {
266        matches!(self, Item::Array(..) | Item::DynamicArray(..))
267    }
268
269    pub fn unwrap_ptr(&self) -> Item<D> {
270        match self {
271            Item::Pointer(inner, _) => **inner,
272            other => *other,
273        }
274    }
275}