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}