Skip to main content

cubecl_spirv/
target.rs

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        // Callback requires physical storage buffer
109        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            // uniform/push constant pointer to physical storage buffer pointer
275            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            // uniform/push constant pointer to physical storage buffer pointer
294            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    /// Generate info binding struct and variable.
371    /// SPIR-V structs have explicit offsets so unlike other targets we don't need to pad the length.
372    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}