cubecl_opt/analyses/
uniformity.rs1use alloc::vec::Vec;
2use cubecl_ir::{
3 Builtin, Operation, OperationReflect, Operator, Plane, Synchronization, Value, ValueKind,
4};
5use hashbrown::{HashMap, HashSet};
6use petgraph::{graph::EdgeIndex, visit::EdgeRef};
7
8use crate::{ControlFlow, Function, GlobalState, NodeIndex};
9
10use super::Analysis;
11
12#[derive(Default, Clone)]
13pub struct Uniformity {
14 block_uniformity: HashMap<NodeIndex, bool>,
15 value_uniformity: HashMap<Value, bool>,
16 visited: HashSet<EdgeIndex>,
17}
18
19impl Analysis for Uniformity {
20 fn init(func: &mut Function, _: &GlobalState) -> Self {
21 let mut this = Self::default();
22 this.run(func);
23 this
24 }
25}
26
27impl Uniformity {
28 fn run(&mut self, func: &Function) {
29 let root = func.root;
30 self.block_uniformity.insert(root, true);
31 while self.analyze_block(func, root).is_none() {}
32 }
33
34 fn analyze_block(&mut self, func: &Function, block_id: NodeIndex) -> Option<()> {
35 let block = func.block(block_id);
36 let mut block_uniform = self.block_uniformity[&block_id];
37
38 for phi in block.phi_nodes.borrow().iter() {
39 let uniform = phi.entries.iter().all(|entry| {
40 let block_uniform = self.is_block_uniform(entry.block);
41 let value_uniform = self.is_val_uniform(entry.value);
42 block_uniform && value_uniform
43 }) && block_uniform;
44 self.mark_uniformity(phi.out, uniform && block_uniform)?;
45 }
46
47 for inst in block.ops.borrow().values() {
48 if inst.out.is_none() {
49 continue;
50 }
51 let out = inst.out.unwrap();
52 match &inst.operation {
53 Operation::Plane(plane) => match plane {
54 Plane::Elect
57 | Plane::ExclusiveSum(_)
58 | Plane::InclusiveSum(_)
59 | Plane::ExclusiveProd(_)
60 | Plane::InclusiveProd(_) => self.mark_uniformity(out, false)?,
61 Plane::Sum(_)
63 | Plane::Prod(_)
64 | Plane::Min(_)
65 | Plane::Max(_)
66 | Plane::All(_)
67 | Plane::Any(_)
68 | Plane::Ballot(_) => self.mark_uniformity(out, block_uniform)?,
69 Plane::Broadcast(op) => {
72 let input_uniform =
73 self.is_val_uniform(op.lhs) || self.is_val_uniform(op.rhs);
74 self.mark_uniformity(out, input_uniform && block_uniform)?;
75 }
76 Plane::Shuffle(op)
79 | Plane::ShuffleXor(op)
80 | Plane::ShuffleUp(op)
81 | Plane::ShuffleDown(op) => {
82 let input_uniform = self.is_val_uniform(op.lhs);
83 self.mark_uniformity(out, input_uniform && block_uniform)?;
84 }
85 },
86 Operation::Synchronization(sync) => match sync {
87 Synchronization::SyncCube | Synchronization::SyncStorage => {
88 block_uniform = true;
89 }
90 Synchronization::SyncAsyncProxyShared => {}
91 Synchronization::SyncPlane => {
92 }
94 },
95 Operation::Operator(Operator::ReadBuiltin(builtin)) => {
96 self.mark_uniformity(out, is_builtin_uniform(builtin) && block_uniform)?;
97 }
98 Operation::Operator(Operator::ReadScalar(_)) => {
99 self.mark_uniformity(out, block_uniform);
100 }
101 op => {
102 let is_uniform =
103 op.is_pure() && self.is_all_uniform(op.args()) && block_uniform;
104 self.mark_uniformity(out, is_uniform)?;
105 }
106 }
107 }
108
109 match &*block.control_flow.borrow() {
110 ControlFlow::IfElse {
111 cond,
112 then,
113 or_else,
114 merge,
115 } => {
116 let is_uniform = self.is_val_uniform(*cond);
117 self.block_uniformity
118 .insert(*then, is_uniform && block_uniform);
119 self.block_uniformity
120 .insert(*or_else, is_uniform && block_uniform);
121 if let Some(merge) = merge {
122 self.block_uniformity.insert(*merge, block_uniform);
123 }
124 }
125 ControlFlow::Switch {
126 value,
127 default,
128 branches,
129 merge,
130 } => {
131 let is_uniform = self.is_val_uniform(*value);
132 self.block_uniformity
133 .insert(*default, is_uniform && block_uniform);
134 for branch in branches {
135 self.block_uniformity
136 .insert(branch.1, is_uniform && block_uniform);
137 }
138 if let Some(merge) = merge {
139 self.block_uniformity.insert(*merge, block_uniform);
140 }
141 }
142 ControlFlow::Loop {
143 body,
144 continue_target,
145 merge,
146 } => {
147 self.block_uniformity.insert(block_id, false);
149 self.block_uniformity.insert(*body, false);
150 self.block_uniformity.insert(*continue_target, false);
151 self.block_uniformity.insert(*merge, false);
152 }
153 ControlFlow::LoopBreak {
154 break_cond,
155 body,
156 continue_target,
157 merge,
158 } => {
159 let is_uniform = self.is_val_uniform(*break_cond);
160 self.block_uniformity
161 .insert(block_id, is_uniform && block_uniform);
162 self.block_uniformity
163 .insert(*body, is_uniform && block_uniform);
164 self.block_uniformity
165 .insert(*continue_target, is_uniform && block_uniform);
166 self.block_uniformity
167 .insert(*merge, is_uniform && block_uniform);
168 }
169 ControlFlow::Return { .. } | ControlFlow::Unreachable => {}
170 ControlFlow::None => {
171 let successor = func.successors(block_id)[0];
172 self.block_uniformity
173 .entry(successor)
174 .and_modify(|it| {
175 *it |= block_uniform;
176 })
177 .or_insert(block_uniform);
178 }
179 }
180
181 for edge in func.edges(block_id) {
182 if !self.visited.contains(&edge.id()) {
183 self.visited.insert(edge.id());
184 self.analyze_block(func, edge.target())?;
185 }
186 }
187
188 Some(())
189 }
190
191 fn mark_uniformity(&mut self, val: Value, new_value: bool) -> Option<()> {
192 if let Some(prev_value) = self.value_uniformity.get_mut(&val) {
193 let invalidate = !new_value && *prev_value;
197 *prev_value = *prev_value && new_value;
198 if invalidate {
199 self.visited.clear();
200 return None;
201 }
202 } else {
203 self.value_uniformity.insert(val, new_value);
204 }
205 Some(())
206 }
207
208 fn is_all_uniform(&self, args: Option<Vec<Value>>) -> bool {
209 args.map(|it| it.iter().all(|it| self.is_val_uniform(*it)))
210 .unwrap_or(false)
211 }
212
213 pub fn is_val_uniform(&self, val: Value) -> bool {
215 match val.kind {
216 ValueKind::Constant(_) => true,
217 ValueKind::Value { .. } => self.value_uniformity.get(&val).copied().unwrap_or(true),
218 }
219 }
220
221 pub fn is_block_uniform(&self, block: NodeIndex) -> bool {
222 self.block_uniformity.get(&block).copied().unwrap_or(true)
223 }
224}
225
226fn is_builtin_uniform(builtin: &Builtin) -> bool {
227 match builtin {
228 Builtin::UnitPosPlane
229 | Builtin::PlanePos
230 | Builtin::AbsolutePos
231 | Builtin::AbsolutePosX
232 | Builtin::AbsolutePosY
233 | Builtin::AbsolutePosZ
234 | Builtin::UnitPos
235 | Builtin::UnitPosX
236 | Builtin::UnitPosY
237 | Builtin::UnitPosZ => false,
238 Builtin::CubePos
239 | Builtin::CubePosX
240 | Builtin::CubePosY
241 | Builtin::CubePosZ
242 | Builtin::CubePosCluster
243 | Builtin::CubePosClusterX
244 | Builtin::CubePosClusterY
245 | Builtin::CubePosClusterZ
246 | Builtin::CubeDim
247 | Builtin::CubeDimX
248 | Builtin::CubeDimY
249 | Builtin::CubeDimZ
250 | Builtin::CubeClusterDim
251 | Builtin::CubeClusterDimX
252 | Builtin::CubeClusterDimY
253 | Builtin::CubeClusterDimZ
254 | Builtin::CubeCount
255 | Builtin::CubeCountX
256 | Builtin::CubeCountY
257 | Builtin::CubeCountZ
258 | Builtin::PlaneDim => true,
259 }
260}