Skip to main content

cubecl_opt/analyses/
uniformity.rs

1use 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                    // Elect returns true on only one unit, so it's always non-uniform
55                    // Inclusive/exclusive scans are non-uniform by definition
56                    Plane::Elect
57                    | Plane::ExclusiveSum(_)
58                    | Plane::InclusiveSum(_)
59                    | Plane::ExclusiveProd(_)
60                    | Plane::InclusiveProd(_) => self.mark_uniformity(out, false)?,
61                    // Reductions are always uniform if executed in uniform control flow
62                    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                    // Broadcast maps to shuffle or broadcast, if id or value is uniform, so will
70                    // the output, otherwise not.
71                    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                    // Shuffle operations: if offset/mask/delta is uniform, output is non-uniform
77                    // (each thread gets a different value). If value is uniform, output is uniform.
78                    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                        // TODO: not sure
93                    }
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                // If we don't know the break condition, we can't detect whether it's uniform
148                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            // If the value was already set before and has been invalidated, we need to revisit
194            // all edges. This only happens for loopback edges, where an uninitialized value
195            // was assumed to be uniform but actually isn't
196            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    /// Whether a value is plane uniform
214    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}