Skip to main content

cubecl_spirv/
branch.rs

1use cubecl_core::ir as core;
2use cubecl_opt::{ControlFlow, NodeIndex};
3use rspirv::{
4    dr::Operand,
5    spirv::{LoopControl, SelectionControl},
6};
7
8use crate::{SpirvCompiler, SpirvTarget};
9
10impl<T: SpirvTarget> SpirvCompiler<T> {
11    pub fn compile_control_flow(&mut self, control_flow: ControlFlow) {
12        match control_flow {
13            ControlFlow::IfElse {
14                cond,
15                then,
16                or_else,
17                merge,
18            } => self.compile_if_else(cond, then, or_else, merge),
19            ControlFlow::Switch {
20                value,
21                default,
22                branches,
23                merge,
24            } => self.compile_switch(value, default, branches, merge),
25            ControlFlow::Loop {
26                body,
27                continue_target,
28                merge,
29            } => self.compile_loop(body, continue_target, merge),
30            ControlFlow::LoopBreak {
31                break_cond,
32                body,
33                continue_target,
34                merge,
35            } => self.compile_loop_break(break_cond, body, continue_target, merge),
36            ControlFlow::Return { value } => {
37                if let Some(value) = value {
38                    let value = self.compile_value(value);
39                    let value_id = self.read(&value);
40                    self.ret_value(value_id).unwrap();
41                } else {
42                    self.ret().unwrap();
43                }
44                self.current_block = None;
45            }
46            ControlFlow::Unreachable => {
47                self.unreachable().unwrap();
48                self.current_block = None;
49            }
50            ControlFlow::None => {
51                let opt = self.opt.clone();
52                let func = self
53                    .current_func
54                    .map(|id| &opt.global_state.extra_functions[&id])
55                    .unwrap_or(&opt.main);
56                let children = func.successors(self.current_block.unwrap());
57                assert_eq!(
58                    children.len(),
59                    1,
60                    "None control flow should have only 1 outgoing edge"
61                );
62                let label = self.label(children[0]);
63                self.branch(label).unwrap();
64            }
65        }
66    }
67
68    fn compile_if_else(
69        &mut self,
70        cond: core::Value,
71        then: NodeIndex,
72        or_else: NodeIndex,
73        merge: Option<NodeIndex>,
74    ) {
75        let cond = self.compile_value(cond);
76        let then_label = self.label(then);
77        let else_label = self.label(or_else);
78        let cond_id = self.read(&cond);
79
80        if let Some(merge) = merge {
81            let merge_label = self.label(merge);
82            self.selection_merge(merge_label, SelectionControl::NONE)
83                .unwrap();
84        }
85        self.branch_conditional(cond_id, then_label, else_label, None)
86            .unwrap();
87    }
88
89    fn compile_switch(
90        &mut self,
91        value: core::Value,
92        default: NodeIndex,
93        branches: Vec<(u32, NodeIndex)>,
94        merge: Option<NodeIndex>,
95    ) {
96        let value = self.compile_value(value);
97        let value_id = self.read(&value);
98
99        let default_label = self.label(default);
100        let targets = branches
101            .iter()
102            .map(|(value, block)| {
103                let label = self.label(*block);
104                (Operand::LiteralBit32(*value), label)
105            })
106            .collect::<Vec<_>>();
107
108        if let Some(merge) = merge {
109            let merge_label = self.label(merge);
110            self.selection_merge(merge_label, SelectionControl::NONE)
111                .unwrap();
112        }
113
114        self.switch(value_id, default_label, targets).unwrap();
115    }
116
117    fn compile_loop(&mut self, body: NodeIndex, continue_target: NodeIndex, merge: NodeIndex) {
118        let body_label = self.label(body);
119        let continue_label = self.label(continue_target);
120        let merge_label = self.label(merge);
121
122        self.loop_merge(merge_label, continue_label, LoopControl::NONE, vec![])
123            .unwrap();
124        self.branch(body_label).unwrap();
125    }
126
127    fn compile_loop_break(
128        &mut self,
129        break_cond: core::Value,
130        body: NodeIndex,
131        continue_target: NodeIndex,
132        merge: NodeIndex,
133    ) {
134        let break_cond = self.compile_value(break_cond);
135        let cond_id = self.read(&break_cond);
136        let body_label = self.label(body);
137        let continue_label = self.label(continue_target);
138        let merge_label = self.label(merge);
139
140        self.loop_merge(merge_label, continue_label, LoopControl::NONE, [])
141            .unwrap();
142        self.branch_conditional(cond_id, body_label, merge_label, [])
143            .unwrap();
144    }
145}