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 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 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 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}