Skip to main content

cubecl_core/compute/
builder.rs

1use alloc::vec::Vec;
2use core::sync::atomic::{AtomicI8, Ordering};
3use derive_more::Deref;
4
5use crate::{
6    BufferInfo, KernelExpansion, KernelIntegrator, KernelSettings, ScalarInfo,
7    ir::{Id, Type},
8    prelude::KernelDefinition,
9};
10use alloc::collections::BTreeMap;
11use cubecl_ir::{DeviceProperties, Scope, StorageType, TargetProperties, Value};
12use cubecl_runtime::config::{
13    CubeClRuntimeConfig, RuntimeConfig, compilation::CompilationLogLevel,
14};
15
16/// Prepare a kernel to create a [`KernelDefinition`].
17#[derive(Deref)]
18pub struct KernelBuilder {
19    /// Cube [scope](Scope).
20    #[deref]
21    pub scope: Scope,
22    buffers: Vec<BufferInfo>,
23    scalars: BTreeMap<StorageType, usize>,
24    tensor_maps: Vec<BufferInfo>,
25}
26
27static DEBUG: AtomicI8 = AtomicI8::new(-1);
28
29impl KernelBuilder {
30    /// Register a scalar and return the [element](Value) to be used for kernel expansion.
31    pub fn scalar(&mut self, storage: StorageType) -> Id {
32        let current_id = self.scalars.entry(storage).or_default();
33        let id = *current_id;
34        *current_id += 1;
35        id as Id
36    }
37
38    fn buffer_id(&self) -> Id {
39        self.buffers.len() as Id + self.tensor_maps.len() as Id
40    }
41
42    /// Register a buffer and return the [element](Value) to be used for kernel expansion.
43    pub fn buffer(&mut self, value_ty: Type) -> Value {
44        let id = self.buffer_id();
45        let value = self.scope.global(id, value_ty);
46        self.buffers.push(BufferInfo {
47            id,
48            value,
49            has_extended_meta: false,
50        });
51        value
52    }
53
54    /// Register a tensor and return the [element](Value) to be used for kernel expansion.
55    pub fn tensor(&mut self, value_ty: Type) -> Value {
56        let id = self.buffer_id();
57        let value = self.scope.global(id, value_ty);
58        self.buffers.push(BufferInfo {
59            id,
60            value,
61            has_extended_meta: true,
62        });
63        value
64    }
65
66    /// Register a tensor map and return the [element](Value) to be used for kernel expansion.
67    pub fn tensor_map(&mut self) -> Value {
68        let id = self.buffer_id();
69        let value = self.scope.tensor_map(id);
70        self.tensor_maps.push(BufferInfo {
71            id,
72            value,
73            has_extended_meta: true,
74        });
75        value
76    }
77
78    /// Register an output that uses the same resource as the input as the given position.
79    pub fn inplace(&mut self, position: Id) -> Value {
80        let input = self.buffers.get_mut(position as usize);
81        input.expect("Position valid").value
82    }
83
84    pub fn runtime_properties(&mut self, properties: TargetProperties) {
85        self.scope.state_mut().target_properties = properties;
86    }
87
88    pub fn device_properties(&mut self, properties: &DeviceProperties) {
89        self.scope.device_properties(properties);
90    }
91
92    /// Build the [kernel definition](KernelDefinition).
93    pub fn build(self, settings: KernelSettings) -> KernelDefinition {
94        let scalars = self
95            .scalars
96            .into_iter()
97            .map(|(ty, count)| ScalarInfo { ty, count })
98            .collect();
99        KernelIntegrator::new(KernelExpansion {
100            scope: self.scope,
101            buffers: self.buffers,
102            scalars,
103            tensor_maps: self.tensor_maps,
104        })
105        .integrate(settings)
106    }
107
108    pub fn new() -> Self {
109        let debug = DEBUG.load(Ordering::Relaxed);
110        let debug = if debug == -1 {
111            let val = match CubeClRuntimeConfig::get().compilation.logger.level {
112                CompilationLogLevel::Full => 1,
113                _ => 0,
114            };
115
116            DEBUG.store(val, Ordering::Relaxed);
117            val == 1
118        } else {
119            debug == 1
120        };
121
122        Self {
123            scope: Scope::root(debug),
124            buffers: Default::default(),
125            scalars: Default::default(),
126            tensor_maps: Default::default(),
127        }
128    }
129}
130
131impl Default for KernelBuilder {
132    fn default() -> Self {
133        Self::new()
134    }
135}