Skip to main content

cubecl_spirv/
atomic.rs

1use cubecl_core::ir::{AtomicOp, ElemType, InstructionModes, IntKind, UIntKind, Value};
2use rspirv::spirv::{Capability, MemorySemantics, Scope, Word};
3
4use crate::{SpirvCompiler, SpirvTarget, item::Elem};
5
6impl<T: SpirvTarget> SpirvCompiler<T> {
7    pub fn compile_atomic(
8        &mut self,
9        atomic: AtomicOp,
10        out: Option<Value>,
11        modes: InstructionModes,
12    ) {
13        if let Some(out) = out
14            && matches!(
15                out.elem_type(),
16                ElemType::Int(IntKind::I64) | ElemType::UInt(UIntKind::U64)
17            )
18        {
19            self.capabilities.insert(Capability::Int64Atomics);
20        }
21
22        match atomic {
23            AtomicOp::Load(ptr) => {
24                let out = out.unwrap();
25
26                let ptr = self.compile_value(ptr);
27                let out = self.compile_value(out);
28                let out_ty = out.item();
29
30                let input_id = ptr.id(self);
31                let out_id = self.write_id(&out);
32
33                let ty = out_ty.id(self);
34                let memory = self.scope(&ptr);
35                let semantics = self.semantics_r(&ptr);
36
37                self.atomic_load(ty, Some(out_id), input_id, memory, semantics)
38                    .unwrap();
39                self.write(&out, out_id);
40            }
41            AtomicOp::Store(op) => {
42                if matches!(
43                    op.value.elem_type(),
44                    ElemType::Int(IntKind::I64) | ElemType::UInt(UIntKind::U64)
45                ) {
46                    self.capabilities.insert(Capability::Int64Atomics);
47                }
48
49                let value = self.compile_value(op.value);
50                let ptr = self.compile_value(op.ptr);
51
52                let value_id = self.read(&value);
53                let ptr_id = ptr.id(self);
54
55                let memory = self.scope(&ptr);
56                let semantics = self.semantics_w(&ptr);
57
58                self.atomic_store(ptr_id, memory, semantics, value_id)
59                    .unwrap();
60            }
61            AtomicOp::Swap(op) => {
62                let out = out.unwrap();
63
64                let ptr = self.compile_value(op.ptr);
65                let value = self.compile_value(op.value);
66                let out = self.compile_value(out);
67                let out_ty = out.item();
68
69                let ptr_id = ptr.id(self);
70                let value_id = self.read(&value);
71                let out_id = self.write_id(&out);
72
73                let ty = out_ty.id(self);
74                let memory = self.scope(&ptr);
75                let semantics = self.semantics_rw(&ptr);
76
77                self.atomic_exchange(ty, Some(out_id), ptr_id, memory, semantics, value_id)
78                    .unwrap();
79                self.write(&out, out_id);
80            }
81            AtomicOp::CompareAndSwap(op) => {
82                let out = out.unwrap();
83
84                let atomic = self.compile_value(op.ptr);
85                let cmp = self.compile_value(op.cmp);
86                let val = self.compile_value(op.val);
87                let out = self.compile_value(out);
88                let out_ty = out.item();
89
90                let atomic_id = atomic.id(self);
91                let cmp_id = self.read(&cmp);
92                let val_id = self.read(&val);
93                let out_id = self.write_id(&out);
94
95                let ty = out_ty.id(self);
96                let memory = self.scope(&atomic);
97                let semantics_success = self.semantics_rw(&atomic);
98                let semantics_failure = self.semantics_r(&atomic);
99
100                assert!(
101                    matches!(out_ty.elem(), Elem::Int(_, _)),
102                    "compare and swap doesn't support float atomics"
103                );
104                self.atomic_compare_exchange(
105                    ty,
106                    Some(out_id),
107                    atomic_id,
108                    memory,
109                    semantics_success,
110                    semantics_failure,
111                    val_id,
112                    cmp_id,
113                )
114                .unwrap();
115                self.write(&out, out_id);
116            }
117            AtomicOp::Add(op) => {
118                let out = out.unwrap();
119
120                let ptr = self.compile_value(op.ptr);
121                let value = self.compile_value(op.value);
122                let out = self.compile_value(out);
123                let out_ty = out.item();
124
125                let ptr_id = ptr.id(self);
126                let value_id = self.read(&value);
127                let out_id = self.write_id(&out);
128
129                let ty = out_ty.id(self);
130                let memory = self.scope(&ptr);
131                let semantics = self.semantics_rw(&ptr);
132
133                match out_ty.elem() {
134                    Elem::Int(_, _) => self
135                        .atomic_i_add(ty, Some(out_id), ptr_id, memory, semantics, value_id)
136                        .unwrap(),
137                    Elem::Float(width, None) => {
138                        match width {
139                            16 if out_ty.vectorization() == 1 => {
140                                self.capabilities.insert(Capability::AtomicFloat16AddEXT)
141                            }
142                            16 => self.capabilities.insert(Capability::AtomicFloat16VectorNV),
143                            32 => self.capabilities.insert(Capability::AtomicFloat32AddEXT),
144                            64 => self.capabilities.insert(Capability::AtomicFloat64AddEXT),
145                            _ => unreachable!(),
146                        };
147                        self.atomic_f_add_ext(ty, Some(out_id), ptr_id, memory, semantics, value_id)
148                            .unwrap()
149                    }
150                    _ => unreachable!(),
151                };
152
153                self.write(&out, out_id);
154            }
155            AtomicOp::Sub(op) => {
156                let out = out.unwrap();
157
158                let ptr = self.compile_value(op.ptr);
159                let value = self.compile_value(op.value);
160                let out = self.compile_value(out);
161                let out_ty = out.item();
162
163                let ptr_id = ptr.id(self);
164                let value_id = self.read(&value);
165                let out_id = self.write_id(&out);
166
167                let ty = out_ty.id(self);
168                let memory = self.scope(&ptr);
169                let semantics = self.semantics_rw(&ptr);
170
171                match out_ty.elem() {
172                    Elem::Int(_, _) => self
173                        .atomic_i_sub(ty, Some(out_id), ptr_id, memory, semantics, value_id)
174                        .unwrap(),
175                    Elem::Float(width, None) => {
176                        match width {
177                            16 if out_ty.vectorization() == 1 => {
178                                self.capabilities.insert(Capability::AtomicFloat16AddEXT)
179                            }
180                            16 => self.capabilities.insert(Capability::AtomicFloat16VectorNV),
181                            32 => self.capabilities.insert(Capability::AtomicFloat32AddEXT),
182                            64 => self.capabilities.insert(Capability::AtomicFloat64AddEXT),
183                            _ => unreachable!(),
184                        };
185                        let negated = self.f_negate(ty, None, value_id).unwrap();
186                        self.declare_math_mode(modes, negated);
187                        self.atomic_f_add_ext(ty, Some(out_id), ptr_id, memory, semantics, negated)
188                            .unwrap()
189                    }
190                    _ => unreachable!(),
191                };
192                self.write(&out, out_id);
193            }
194            AtomicOp::Max(op) => {
195                let out = out.unwrap();
196
197                let ptr = self.compile_value(op.ptr);
198                let value = self.compile_value(op.value);
199                let out = self.compile_value(out);
200                let out_ty = out.item();
201
202                let ptr_id = ptr.id(self);
203                let value_id = self.read(&value);
204                let out_id = self.write_id(&out);
205
206                let ty = out_ty.id(self);
207                let memory = self.scope(&ptr);
208                let semantics = self.semantics_rw(&ptr);
209
210                match out_ty.elem() {
211                    Elem::Int(_, false) => self
212                        .atomic_u_max(ty, Some(out_id), ptr_id, memory, semantics, value_id)
213                        .unwrap(),
214                    Elem::Int(_, true) => self
215                        .atomic_s_max(ty, Some(out_id), ptr_id, memory, semantics, value_id)
216                        .unwrap(),
217                    Elem::Float(width, None) => {
218                        match width {
219                            16 if out_ty.vectorization() == 1 => {
220                                self.capabilities.insert(Capability::AtomicFloat16MinMaxEXT)
221                            }
222                            16 => self.capabilities.insert(Capability::AtomicFloat16VectorNV),
223                            32 => self.capabilities.insert(Capability::AtomicFloat32MinMaxEXT),
224                            64 => self.capabilities.insert(Capability::AtomicFloat64MinMaxEXT),
225                            _ => unreachable!(),
226                        };
227                        self.atomic_f_max_ext(ty, Some(out_id), ptr_id, memory, semantics, value_id)
228                            .unwrap()
229                    }
230                    _ => unreachable!(),
231                };
232                self.write(&out, out_id);
233            }
234            AtomicOp::Min(op) => {
235                let out = out.unwrap();
236
237                let ptr = self.compile_value(op.ptr);
238                let value = self.compile_value(op.value);
239                let out = self.compile_value(out);
240                let out_ty = out.item();
241
242                let ptr_id = ptr.id(self);
243                let value_id = self.read(&value);
244                let out_id = self.write_id(&out);
245
246                let ty = out_ty.id(self);
247                let memory = self.scope(&ptr);
248                let semantics = self.semantics_rw(&ptr);
249
250                match out_ty.elem() {
251                    Elem::Int(_, false) => self
252                        .atomic_u_min(ty, Some(out_id), ptr_id, memory, semantics, value_id)
253                        .unwrap(),
254                    Elem::Int(_, true) => self
255                        .atomic_s_min(ty, Some(out_id), ptr_id, memory, semantics, value_id)
256                        .unwrap(),
257                    Elem::Float(width, None) => {
258                        match width {
259                            16 if out_ty.vectorization() == 1 => {
260                                self.capabilities.insert(Capability::AtomicFloat16MinMaxEXT)
261                            }
262                            16 => self.capabilities.insert(Capability::AtomicFloat16VectorNV),
263                            32 => self.capabilities.insert(Capability::AtomicFloat32MinMaxEXT),
264                            64 => self.capabilities.insert(Capability::AtomicFloat64MinMaxEXT),
265                            _ => unreachable!(),
266                        };
267                        self.atomic_f_min_ext(ty, Some(out_id), ptr_id, memory, semantics, value_id)
268                            .unwrap()
269                    }
270                    _ => unreachable!(),
271                };
272                self.write(&out, out_id);
273            }
274            AtomicOp::And(op) => {
275                let out = out.unwrap();
276
277                let ptr = self.compile_value(op.ptr);
278                let value = self.compile_value(op.value);
279                let out = self.compile_value(out);
280                let out_ty = out.item();
281
282                let ptr_id = ptr.id(self);
283                let value_id = self.read(&value);
284                let out_id = self.write_id(&out);
285
286                let ty = out_ty.id(self);
287                let memory = self.scope(&ptr);
288                let semantics = self.semantics_rw(&ptr);
289
290                assert!(
291                    matches!(out_ty.elem(), Elem::Int(_, _)),
292                    "and doesn't support float atomics"
293                );
294                self.atomic_and(ty, Some(out_id), ptr_id, memory, semantics, value_id)
295                    .unwrap();
296                self.write(&out, out_id);
297            }
298            AtomicOp::Or(op) => {
299                let out = out.unwrap();
300
301                let ptr = self.compile_value(op.ptr);
302                let value = self.compile_value(op.value);
303                let out = self.compile_value(out);
304                let out_ty = out.item();
305
306                let ptr_id = ptr.id(self);
307                let value_id = self.read(&value);
308                let out_id = self.write_id(&out);
309
310                let ty = out_ty.id(self);
311                let memory = self.scope(&ptr);
312                let semantics = self.semantics_rw(&ptr);
313
314                assert!(
315                    matches!(out_ty.elem(), Elem::Int(_, _)),
316                    "or doesn't support float atomics"
317                );
318                self.atomic_or(ty, Some(out_id), ptr_id, memory, semantics, value_id)
319                    .unwrap();
320                self.write(&out, out_id);
321            }
322            AtomicOp::Xor(op) => {
323                let out = out.unwrap();
324
325                let ptr = self.compile_value(op.ptr);
326                let value = self.compile_value(op.value);
327                let out = self.compile_value(out);
328                let out_ty = out.item();
329
330                let ptr_id = ptr.id(self);
331                let value_id = self.read(&value);
332                let out_id = self.write_id(&out);
333
334                let ty = out_ty.id(self);
335                let memory = self.scope(&ptr);
336                let semantics = self.semantics_rw(&ptr);
337
338                assert!(
339                    matches!(out_ty.elem(), Elem::Int(_, _)),
340                    "xor doesn't support float atomics"
341                );
342                self.atomic_xor(ty, Some(out_id), ptr_id, memory, semantics, value_id)
343                    .unwrap();
344                self.write(&out, out_id);
345            }
346        }
347    }
348
349    fn scope(&mut self, val: &crate::value::Value) -> Word {
350        let value = val.scope() as u32;
351        self.const_u32(value)
352    }
353
354    fn semantics_r(&mut self, val: &crate::value::Value) -> Word {
355        let value = self.semantics_of(val) | MemorySemantics::ACQUIRE;
356        self.const_u32(value.bits())
357    }
358
359    fn semantics_w(&mut self, val: &crate::value::Value) -> Word {
360        let value = self.semantics_of(val) | MemorySemantics::RELEASE;
361        self.const_u32(value.bits())
362    }
363
364    fn semantics_rw(&mut self, val: &crate::value::Value) -> Word {
365        let value = self.semantics_of(val) | MemorySemantics::ACQUIRE_RELEASE;
366        self.const_u32(value.bits())
367    }
368
369    fn semantics_of(&mut self, val: &crate::value::Value) -> MemorySemantics {
370        match val.scope() {
371            Scope::Device => MemorySemantics::UNIFORM_MEMORY,
372            Scope::Workgroup => MemorySemantics::WORKGROUP_MEMORY,
373            Scope::Subgroup => MemorySemantics::SUBGROUP_MEMORY,
374            other => unreachable!("Invalid scope for atomic operation, {other:?}"),
375        }
376    }
377}