Skip to main content

cubecl_spirv/
globals.rs

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}