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}