cubecl_core/compute/
builder.rs1use 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#[derive(Deref)]
18pub struct KernelBuilder {
19 #[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 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 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 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 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 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 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}