1use cubecl_core::ir::Builtin;
2use rspirv::spirv::{BuiltIn, Word};
3
4use crate::{
5 SpirvCompiler, SpirvTarget,
6 item::{Elem, Item},
7};
8
9impl<T: SpirvTarget> SpirvCompiler<T> {
10 fn compile_builtin_u32(&mut self, builtin: Builtin) -> Word {
11 self.compile_builtin(builtin, &Item::builtin_u32())
12 }
13
14 pub fn compile_builtin(&mut self, builtin: Builtin, ty: &Item) -> Word {
15 match builtin {
16 Builtin::UnitPos => self.insert_global(builtin, |b| {
17 let id = b.load_builtin(BuiltIn::LocalInvocationIndex, ty);
18 b.debug_name(id, "UNIT_POS");
19 id
20 }),
21 Builtin::UnitPosX => self.insert_global(builtin, |b| {
22 let id = b.extract(BuiltIn::LocalInvocationId, 0, ty);
23 b.debug_name(id, "UNIT_POS_X");
24 id
25 }),
26 Builtin::UnitPosY => self.insert_global(builtin, |b| {
27 let id = b.extract(BuiltIn::LocalInvocationId, 1, ty);
28 b.debug_name(id, "UNIT_POS_Y");
29 id
30 }),
31 Builtin::UnitPosZ => self.insert_global(builtin, |b| {
32 let id = b.extract(BuiltIn::LocalInvocationId, 2, ty);
33 b.debug_name(id, "UNIT_POS_Z");
34 id
35 }),
36 Builtin::CubePosX => self.insert_global(builtin, |b| {
37 let id = b.extract(BuiltIn::WorkgroupId, 0, ty);
38 b.debug_name(id, "CUBE_POS_X");
39 id
40 }),
41 Builtin::CubePosY => self.insert_global(builtin, |b| {
42 let id = b.extract(BuiltIn::WorkgroupId, 1, ty);
43 b.debug_name(id, "CUBE_POS_Y");
44 id
45 }),
46 Builtin::CubePosZ => self.insert_global(builtin, |b| {
47 let id = b.extract(BuiltIn::WorkgroupId, 2, ty);
48 b.debug_name(id, "CUBE_POS_Z");
49 id
50 }),
51 Builtin::CubePosCluster
52 | Builtin::CubePosClusterX
53 | Builtin::CubePosClusterY
54 | Builtin::CubePosClusterZ => ty.const_u32(self, 0),
55 Builtin::CubeDim => self.state.cube_size,
56 Builtin::CubeDimX => self.state.cube_dims[0],
57 Builtin::CubeDimY => self.state.cube_dims[1],
58 Builtin::CubeDimZ => self.state.cube_dims[2],
59 Builtin::CubeClusterDim
60 | Builtin::CubeClusterDimX
61 | Builtin::CubeClusterDimY
62 | Builtin::CubeClusterDimZ => ty.const_u32(self, 1),
63 Builtin::CubeCount => self.insert_global(builtin, |b: &mut SpirvCompiler<T>| {
64 let ty_id = ty.id(b);
65 let x = b.compile_builtin_u32(Builtin::CubeCountX);
66 let y = b.compile_builtin_u32(Builtin::CubeCountY);
67 let z = b.compile_builtin_u32(Builtin::CubeCountZ);
68
69 let x = Item::builtin_u32().cast_to(b, None, x, ty);
70 let y = Item::builtin_u32().cast_to(b, None, y, ty);
71 let z = Item::builtin_u32().cast_to(b, None, z, ty);
72
73 let count = b.i_mul(ty_id, None, x, y).unwrap();
74 let count = b.i_mul(ty_id, None, count, z).unwrap();
75 b.debug_name(count, "CUBE_COUNT");
76 count
77 }),
78 Builtin::CubeCountX => self.insert_global(builtin, |b| {
79 let id = b.extract(BuiltIn::NumWorkgroups, 0, ty);
80 b.debug_name(id, "CUBE_COUNT_X");
81 id
82 }),
83 Builtin::CubeCountY => self.insert_global(builtin, |b| {
84 let id = b.extract(BuiltIn::NumWorkgroups, 1, ty);
85 b.debug_name(id, "CUBE_COUNT_Y");
86 id
87 }),
88 Builtin::CubeCountZ => self.insert_global(builtin, |b| {
89 let id = b.extract(BuiltIn::NumWorkgroups, 2, ty);
90 b.debug_name(id, "CUBE_COUNT_Z");
91 id
92 }),
93 Builtin::PlaneDim => self.insert_global(builtin, |b| {
94 let id = b.load_builtin(BuiltIn::SubgroupSize, ty);
95 b.debug_name(id, "PLANE_DIM");
96 id
97 }),
98 Builtin::PlanePos => self.insert_global(builtin, |b| {
99 let id = b.load_builtin(BuiltIn::SubgroupId, ty);
100 b.debug_name(id, "PLANE_POS");
101 id
102 }),
103 Builtin::UnitPosPlane => self.insert_global(builtin, |b| {
104 let id = b.load_builtin(BuiltIn::SubgroupLocalInvocationId, ty);
105 b.debug_name(id, "UNIT_POS_PLANE");
106 id
107 }),
108 Builtin::CubePos => self.insert_global(builtin, |b| {
109 let x = b.compile_builtin_u32(Builtin::CubePosX);
110 let y = b.compile_builtin_u32(Builtin::CubePosY);
111 let z = b.compile_builtin_u32(Builtin::CubePosZ);
112
113 let x = Item::builtin_u32().cast_to(b, None, x, ty);
114 let y = Item::builtin_u32().cast_to(b, None, y, ty);
115 let z = Item::builtin_u32().cast_to(b, None, z, ty);
116
117 let groups_x = b.compile_builtin_u32(Builtin::CubeCountX);
118 let groups_y = b.compile_builtin_u32(Builtin::CubeCountY);
119
120 let groups_x = Item::builtin_u32().cast_to(b, None, groups_x, ty);
121 let groups_y = Item::builtin_u32().cast_to(b, None, groups_y, ty);
122
123 let ty = ty.id(b);
124 let id = b.i_mul(ty, None, z, groups_y).unwrap();
125 let id = b.i_add(ty, None, id, y).unwrap();
126 let id = b.i_mul(ty, None, id, groups_x).unwrap();
127 let id = b.i_add(ty, None, id, x).unwrap();
128 b.debug_name(id, "CUBE_POS");
129 id
130 }),
131 Builtin::AbsolutePos => self.insert_global(builtin, |b| {
132 let x = b.compile_builtin_u32(Builtin::AbsolutePosX);
133 let y = b.compile_builtin_u32(Builtin::AbsolutePosY);
134 let z = b.compile_builtin_u32(Builtin::AbsolutePosZ);
135
136 let x = Item::builtin_u32().cast_to(b, None, x, ty);
137 let y = Item::builtin_u32().cast_to(b, None, y, ty);
138 let z = Item::builtin_u32().cast_to(b, None, z, ty);
139
140 let groups_x = b.compile_builtin_u32(Builtin::CubeCountX);
141 let groups_y = b.compile_builtin_u32(Builtin::CubeCountY);
142
143 let groups_x = Item::builtin_u32().cast_to(b, None, groups_x, ty);
144 let groups_y = Item::builtin_u32().cast_to(b, None, groups_y, ty);
145
146 let size_x = ty.const_u32(b, b.cube_dim.x);
147 let size_y = ty.const_u32(b, b.cube_dim.y);
148
149 let ty = ty.id(b);
150 let size_x = b.i_mul(ty, None, groups_x, size_x).unwrap();
151 let size_y = b.i_mul(ty, None, groups_y, size_y).unwrap();
152 let id = b.i_mul(ty, None, z, size_y).unwrap();
153 let id = b.i_add(ty, None, id, y).unwrap();
154 let id = b.i_mul(ty, None, id, size_x).unwrap();
155 let id = b.i_add(ty, None, id, x).unwrap();
156 b.debug_name(id, "ABSOLUTE_POS");
157 id
158 }),
159 Builtin::AbsolutePosX => self.insert_global(builtin, |b| {
160 let id = b.extract(BuiltIn::GlobalInvocationId, 0, ty);
161 b.debug_name(id, "ABSOLUTE_POS_X");
162 id
163 }),
164 Builtin::AbsolutePosY => self.insert_global(builtin, |b| {
165 let id = b.extract(BuiltIn::GlobalInvocationId, 1, ty);
166 b.debug_name(id, "ABSOLUTE_POS_Y");
167 id
168 }),
169 Builtin::AbsolutePosZ => self.insert_global(builtin, |b| {
170 let id = b.extract(BuiltIn::GlobalInvocationId, 2, ty);
171 b.debug_name(id, "ABSOLUTE_POS_Z");
172 id
173 }),
174 }
175 }
176
177 fn extract(&mut self, builtin: BuiltIn, idx: u32, ty: &Item) -> Word {
178 let composite_id = self.vec_global(builtin);
179 let ty = ty.id(self);
180 self.composite_extract(ty, None, composite_id, vec![idx])
181 .unwrap()
182 }
183
184 fn vec_global(&mut self, builtin: BuiltIn) -> Word {
185 let item = Item::Vector(Elem::Int(32, false), 3);
186
187 self.insert_builtin(builtin, |b| b.load_builtin(builtin, &item))
188 }
189
190 fn load_builtin(&mut self, builtin: BuiltIn, item: &Item) -> Word {
191 let item_id = item.id(self);
192 let id = self.builtin(builtin, item.clone());
193 self.load(item_id, None, id, None, vec![]).unwrap()
194 }
195}