1use cubecl_core::prelude::KernelArg;
2use cubecl_opt::BufferVisibility;
3use rspirv::{
4 dr::Operand,
5 spirv::{
6 self, AddressingModel, Capability, Decoration, ExecutionMode, ExecutionModel, MemoryModel,
7 StorageClass, Word,
8 },
9};
10use std::{fmt::Debug, iter};
11
12use crate::{SpirvCompiler, extensions::TargetExtensions, item::Item, lookups::Buffer};
13
14pub trait SpirvTarget:
15 TargetExtensions<Self> + Debug + Clone + Default + Send + Sync + 'static
16{
17 fn set_modes(
18 &mut self,
19 b: &mut SpirvCompiler<Self>,
20 main: Word,
21 builtins: Vec<Word>,
22 cube_dims: Vec<u32>,
23 );
24 fn generate_params(
25 &mut self,
26 b: &mut SpirvCompiler<Self>,
27 bindings: &[KernelArg],
28 visibility: &[BufferVisibility],
29 ) -> Vec<Buffer>;
30 fn load_params(b: &mut SpirvCompiler<Self>);
31 fn info_storage_class(b: &mut SpirvCompiler<Self>) -> StorageClass;
32 fn params_storage_class(b: &mut SpirvCompiler<Self>, num_buffers: usize) -> StorageClass;
33
34 fn set_kernel_name(&mut self, name: impl Into<String>);
35}
36
37#[derive(Clone)]
38pub struct GLCompute {
39 kernel_name: String,
40}
41
42impl Default for GLCompute {
43 fn default() -> Self {
44 Self {
45 kernel_name: "main".into(),
46 }
47 }
48}
49
50impl Debug for GLCompute {
51 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
52 f.write_str("gl_compute")
53 }
54}
55
56impl SpirvTarget for GLCompute {
57 fn set_modes(
58 &mut self,
59 b: &mut SpirvCompiler<Self>,
60 main: Word,
61 builtins: Vec<Word>,
62 cube_dims: Vec<u32>,
63 ) {
64 let interface: Vec<u32> = builtins
65 .into_iter()
66 .chain(iter::once(b.state.params))
67 .chain(b.state.shared.values().map(|it| it.val_id))
68 .collect();
69
70 let version = b.compilation_options.vulkan.max_spirv_version;
71
72 b.capability(Capability::Shader);
73 b.capability(Capability::PhysicalStorageBufferAddresses);
74 b.capability(Capability::VulkanMemoryModel);
75 b.capability(Capability::VulkanMemoryModelDeviceScope);
76 b.capability(Capability::GroupNonUniform);
77
78 if b.compilation_options.vulkan.supports_explicit_smem {
79 b.extension("SPV_KHR_workgroup_memory_explicit_layout");
80 }
81
82 if b.compilation_options.vulkan.supports_long_vectors {
83 b.extension("SPV_EXT_long_vector");
84 b.capability(Capability::LongVectorEXT);
85 }
86
87 if b.addr_type.size_bits() == 64 {
88 b.extension("SPV_EXT_shader_64bit_indexing");
89 b.capability(Capability::Shader64BitIndexingEXT);
90 b.execution_mode(main, ExecutionMode::Shader64BitIndexingEXT, []);
91 }
92
93 let mut caps = b.capabilities.clone();
94
95 if caps.contains(&Capability::CooperativeMatrixKHR) {
96 b.extension("SPV_KHR_cooperative_matrix");
97 }
98
99 if caps.contains(&Capability::CooperativeMatrixReductionsNV)
100 || caps.contains(&Capability::CooperativeMatrixConversionsNV)
101 || caps.contains(&Capability::CooperativeMatrixPerElementOperationsNV)
102 || caps.contains(&Capability::CooperativeMatrixTensorAddressingNV)
103 || caps.contains(&Capability::CooperativeMatrixBlockLoadsNV)
104 {
105 b.extension("SPV_NV_cooperative_matrix2")
106 }
107
108 if caps.contains(&Capability::CooperativeMatrixBlockLoadsNV) {
110 b.extension("SPV_KHR_physical_storage_buffer");
111 caps.insert(Capability::PhysicalStorageBufferAddresses);
112 }
113
114 if caps.contains(&Capability::TensorAddressingNV) {
115 b.extension("SPV_NV_tensor_addressing")
116 }
117
118 if caps.contains(&Capability::AtomicFloat16AddEXT) {
119 b.extension("SPV_EXT_shader_atomic_float16_add");
120 }
121
122 if caps.contains(&Capability::AtomicFloat32AddEXT)
123 | caps.contains(&Capability::AtomicFloat64AddEXT)
124 {
125 b.extension("SPV_EXT_shader_atomic_float_add");
126 }
127
128 if caps.contains(&Capability::AtomicFloat16MinMaxEXT)
129 | caps.contains(&Capability::AtomicFloat32MinMaxEXT)
130 | caps.contains(&Capability::AtomicFloat64MinMaxEXT)
131 {
132 b.extension("SPV_EXT_shader_atomic_float_min_max");
133 }
134
135 if caps.contains(&Capability::AtomicFloat16VectorNV) {
136 b.extension("SPV_NV_shader_atomic_fp16_vector");
137 }
138
139 if caps.contains(&Capability::BFloat16TypeKHR)
140 || caps.contains(&Capability::BFloat16CooperativeMatrixKHR)
141 || caps.contains(&Capability::BFloat16DotProductKHR)
142 {
143 b.extension("SPV_KHR_bfloat16");
144 }
145
146 if caps.contains(&Capability::Float8EXT)
147 || caps.contains(&Capability::Float8CooperativeMatrixEXT)
148 {
149 b.extension("SPV_EXT_float8");
150 }
151
152 if caps.contains(&Capability::FloatControls2) {
153 b.extension("SPV_KHR_float_controls2");
154 }
155
156 if b.debug_symbols {
157 b.extension("SPV_KHR_non_semantic_info");
158 }
159
160 if version < (1, 5) {
161 b.extension("SPV_KHR_physical_storage_buffer");
162 b.extension("SPV_KHR_vulkan_memory_model");
163 if caps.contains(&Capability::StorageBuffer8BitAccess) {
164 b.extension("SPV_KHR_8bit_storage");
165 }
166 }
167
168 if version < (1, 3) && caps.contains(&Capability::StorageBuffer16BitAccess) {
169 b.extension("SPV_KHR_16bit_storage");
170 }
171
172 for cap in caps {
173 b.capability(cap);
174 }
175
176 b.memory_model(
177 AddressingModel::PhysicalStorageBuffer64,
178 MemoryModel::Vulkan,
179 );
180 b.entry_point(
181 ExecutionModel::GLCompute,
182 main,
183 &self.kernel_name,
184 interface,
185 );
186 b.execution_mode(main, spirv::ExecutionMode::LocalSize, cube_dims);
187 }
188
189 fn generate_params(
190 &mut self,
191 b: &mut SpirvCompiler<Self>,
192 bindings: &[KernelArg],
193 visibility: &[BufferVisibility],
194 ) -> Vec<Buffer> {
195 let params_class = Self::params_storage_class(b, bindings.len());
196
197 let params_struct_id = b.id();
198 let params_ptr_id = b.id();
199
200 let buffers = bindings
201 .iter()
202 .map(|binding| {
203 let buffer = self.generate_storage_buffer(b, binding);
204 b.state
205 .base_lookups
206 .values
207 .insert(binding.value.id(), buffer.id);
208 buffer
209 })
210 .collect::<Vec<_>>();
211 let info = b.info.has_info().then(|| self.generate_info_binding(b));
212
213 b.type_struct_id(
214 Some(params_struct_id),
215 buffers
216 .iter()
217 .chain(info.iter())
218 .map(|it| it.struct_ptr_ty_id),
219 );
220 b.type_pointer(Some(params_ptr_id), params_class, params_struct_id);
221
222 b.decorate(params_struct_id, Decoration::Block, []);
223 b.name(params_struct_id, "Params");
224
225 let params = b.insert_in_root(|b| b.variable(params_ptr_id, None, params_class, None));
226 b.name(params, "params");
227
228 b.state.params = params;
229
230 if !matches!(params_class, StorageClass::PushConstant) {
231 b.decorate(params, Decoration::DescriptorSet, vec![0u32.into()]);
232 b.decorate(params, Decoration::Binding, vec![0u32.into()]);
233 }
234
235 for (i, visibility) in visibility.iter().enumerate() {
236 let offset = (size_of::<u64>() * i) as u32;
237 b.member_decorate(
238 params_struct_id,
239 i as u32,
240 Decoration::Offset,
241 [offset.into()],
242 );
243 if !visibility.readable {
244 b.member_decorate(params_struct_id, i as u32, Decoration::NonReadable, []);
245 }
246 if !visibility.writable {
247 b.member_decorate(params_struct_id, i as u32, Decoration::NonWritable, []);
248 }
249 }
250
251 if let Some(info) = info {
252 let i = buffers.len();
253 let offset = (size_of::<u64>() * i) as u32;
254 b.member_decorate(
255 params_struct_id,
256 i as u32,
257 Decoration::Offset,
258 [offset.into()],
259 );
260 b.member_decorate(params_struct_id, i as u32, Decoration::NonWritable, []);
261
262 b.state.info = Some(info);
263 }
264
265 buffers
266 }
267
268 fn load_params(b: &mut SpirvCompiler<Self>) {
269 let params = b.state.params;
270 let params_class = Self::params_storage_class(b, b.state.buffers.len());
271 let zero = b.const_u32(0);
272
273 for (i, buffer) in b.state.buffers.clone().into_iter().enumerate() {
274 let field_ptr_ty = b.type_pointer(None, params_class, buffer.struct_ptr_ty_id);
276 let field_idx = b.const_u32(i as u32);
277 let ptr = b
278 .in_bounds_access_chain(field_ptr_ty, None, params, [field_idx])
279 .unwrap();
280 b.insert_in_setup(|b| {
281 let st_ptr = b
282 .load(buffer.struct_ptr_ty_id, None, ptr, None, [])
283 .unwrap();
284 b.in_bounds_access_chain(buffer.arr_ptr_ty_id, Some(buffer.id), st_ptr, [zero])
285 .unwrap()
286 });
287 b.name(buffer.id, format!("global_{i}"));
288 }
289
290 if let Some(info) = b.state.info {
291 let i = b.state.buffers.len();
292
293 let field_ptr_ty = b.type_pointer(None, params_class, info.struct_ptr_ty_id);
295 let field_idx = b.const_u32(i as u32);
296 let ptr = b
297 .in_bounds_access_chain(field_ptr_ty, None, params, [field_idx])
298 .unwrap();
299 b.insert_in_setup(|b| {
300 b.load(info.struct_ptr_ty_id, Some(info.id), ptr, None, [])
301 .unwrap()
302 });
303 b.name(info.id, "info");
304 }
305 }
306
307 fn info_storage_class(_b: &mut SpirvCompiler<Self>) -> StorageClass {
308 StorageClass::PhysicalStorageBuffer
309 }
310
311 fn params_storage_class(b: &mut SpirvCompiler<Self>, num_buffers: usize) -> StorageClass {
312 let num_addresses = match b.info.has_info() {
313 true => num_buffers + 1,
314 false => num_buffers,
315 };
316 if num_addresses > b.compilation_options.vulkan.push_constant_size / size_of::<u64>() {
317 StorageClass::Uniform
318 } else {
319 StorageClass::PushConstant
320 }
321 }
322
323 fn set_kernel_name(&mut self, name: impl Into<String>) {
324 self.kernel_name = name.into();
325 }
326}
327
328impl GLCompute {
329 fn generate_storage_buffer(
330 &mut self,
331 b: &mut SpirvCompiler<Self>,
332 binding: &KernelArg,
333 ) -> Buffer {
334 let item = b.compile_type(binding.value.ty.unwrap_ptr());
335 match item.elem().size() {
336 1 => {
337 b.capabilities.insert(Capability::StorageBuffer8BitAccess);
338 }
339 2 => {
340 b.capabilities.insert(Capability::StorageBuffer16BitAccess);
341 }
342 _ => {}
343 }
344
345 let value_size = item.value_type().size();
346
347 let arr_ty_id = item.id(b);
348 let struct_ty_id = b.id();
349 let storage_class = StorageClass::PhysicalStorageBuffer;
350
351 b.decorate(arr_ty_id, Decoration::ArrayStride, [value_size.into()]);
352
353 b.type_struct_id(Some(struct_ty_id), [arr_ty_id]);
354 b.decorate(struct_ty_id, Decoration::Block, []);
355 b.member_decorate(struct_ty_id, 0, Decoration::Offset, [0u32.into()]);
356
357 let arr_ptr_ty_id = b.type_pointer(None, storage_class, arr_ty_id);
358 let struct_ptr_ty_id = b.type_pointer(None, storage_class, struct_ty_id);
359
360 Buffer {
361 id: b.id(),
362 struct_ty_id,
363 struct_ptr_ty_id,
364 arr_ty_id,
365 arr_ptr_ty_id,
366 storage_class,
367 }
368 }
369
370 fn generate_info_binding(&mut self, b: &mut SpirvCompiler<Self>) -> Buffer {
373 let address_type = b.addr_type;
374 let struct_ty_id = b.id();
375 let storage_class = StorageClass::PhysicalStorageBuffer;
376
377 let mut fields = Vec::new();
378
379 let scalars = b.info.scalars.clone();
380
381 for scalar in scalars {
382 let scalar_ty = b.compile_storage_type(scalar.ty);
383 match scalar_ty.size() {
384 1 => {
385 b.capabilities.insert(Capability::StorageBuffer8BitAccess);
386 b.capabilities
387 .insert(Capability::UniformAndStorageBuffer8BitAccess);
388 }
389 2 => {
390 b.capabilities.insert(Capability::StorageBuffer16BitAccess);
391 b.capabilities
392 .insert(Capability::UniformAndStorageBuffer16BitAccess);
393 }
394 _ => {}
395 }
396
397 let ty_size = scalar_ty.size();
398 let scalar_ty_id = Item::Scalar(scalar_ty).id(b);
399 let arr_ty_id = b.id();
400 let len_id = b.const_u32(scalar.padded_size() as u32);
401
402 b.type_array_id(Some(arr_ty_id), scalar_ty_id, len_id);
403 b.decorate(arr_ty_id, Decoration::ArrayStride, [ty_size.into()]);
404 b.name(arr_ty_id, format!("Scalars<{}>", scalar.ty));
405
406 b.member_decorate(
407 struct_ty_id,
408 fields.len() as u32,
409 Decoration::Offset,
410 [(scalar.offset as u32).into()],
411 );
412 fields.push(arr_ty_id);
413 }
414
415 if let Some(field) = b.info.sized_meta {
416 let scalar_ty = b.compile_storage_type(field.ty);
417
418 let ty_size = scalar_ty.size();
419 let scalar_ty_id = Item::Scalar(scalar_ty).id(b);
420 let arr_ty_id = b.id();
421 let len_id = b.const_u32(field.size as u32);
422
423 b.type_array_id(Some(arr_ty_id), scalar_ty_id, len_id);
424 b.decorate(arr_ty_id, Decoration::ArrayStride, [ty_size.into()]);
425 b.name(arr_ty_id, "StaticMeta");
426
427 b.member_decorate(
428 struct_ty_id,
429 fields.len() as u32,
430 Decoration::Offset,
431 [(field.offset as u32).into()],
432 );
433 fields.push(arr_ty_id);
434 }
435
436 if b.info.has_dynamic_meta {
437 let offset = b.info.dynamic_meta_offset;
438 let scalar_ty = b.compile_storage_type(address_type);
439
440 let ty_size = scalar_ty.size();
441 let scalar_ty_id = Item::Scalar(scalar_ty).id(b);
442 let arr_ty_id = b.id();
443
444 b.type_runtime_array_id(Some(arr_ty_id), scalar_ty_id);
445 b.decorate(arr_ty_id, Decoration::ArrayStride, [ty_size.into()]);
446 b.name(arr_ty_id, "DynamicMeta");
447
448 b.member_decorate(
449 struct_ty_id,
450 fields.len() as u32,
451 Decoration::Offset,
452 [Operand::LiteralBit32(offset as u32)],
453 );
454 fields.push(arr_ty_id);
455 }
456
457 b.type_struct_id(Some(struct_ty_id), fields);
458 b.decorate(struct_ty_id, Decoration::Block, vec![]);
459 b.name(struct_ty_id, "Info");
460
461 let struct_ptr_ty_id = b.type_pointer(None, storage_class, struct_ty_id);
462
463 Buffer {
464 id: b.id(),
465 struct_ty_id,
466 struct_ptr_ty_id,
467 arr_ty_id: 0,
468 arr_ptr_ty_id: 0,
469 storage_class,
470 }
471 }
472}