1use std::sync::Arc;
5
6use vyre_foundation::ir::{Ident, Node, Program};
7use vyre_foundation::memory_model::MemoryOrdering;
8
9use super::let_propagation::propagate_let_bindings;
10use super::reserve_grid_sync_vec;
11use crate::backend::BackendError;
12
13#[derive(Clone, Debug, PartialEq, Eq)]
21enum EntryWrapper {
22 Region { generator: Ident },
23 Block,
24}
25
26fn peel_entry_wrappers(program: &Program) -> (Vec<EntryWrapper>, &[Node]) {
27 let mut wrappers = Vec::new();
28 let mut entry = program.entry();
29 loop {
30 if entry.len() == 1 {
31 match &entry[0] {
32 Node::Region {
33 generator, body, ..
34 } => {
35 wrappers.push(EntryWrapper::Region {
36 generator: generator.clone(),
37 });
38 entry = body.as_slice();
39 continue;
40 }
41 Node::Block(body) => {
42 wrappers.push(EntryWrapper::Block);
43 entry = body.as_slice();
44 continue;
45 }
46 _ => {}
47 }
48 }
49 break;
50 }
51 (wrappers, entry)
52}
53
54pub(super) fn entry_sequence(program: &Program) -> &[Node] {
55 peel_entry_wrappers(program).1
56}
57
58#[must_use]
67pub fn contains_grid_sync(program: &Program) -> bool {
68 if !program.stats().has_node_barrier() {
74 return false;
75 }
76 node_slice_contains_grid_sync(entry_sequence(program))
77}
78
79fn node_slice_contains_grid_sync(nodes: &[Node]) -> bool {
80 nodes.iter().any(node_contains_grid_sync)
81}
82
83fn node_contains_grid_sync(node: &Node) -> bool {
84 match node {
85 Node::Barrier {
86 ordering: MemoryOrdering::GridSync,
87 ..
88 } => true,
89 Node::If {
90 then, otherwise, ..
91 } => node_slice_contains_grid_sync(then) || node_slice_contains_grid_sync(otherwise),
92 Node::Loop { body, .. } | Node::Block(body) => node_slice_contains_grid_sync(body),
93 Node::Region { body, .. } => node_slice_contains_grid_sync(body),
94 _ => false,
95 }
96}
97
98#[must_use]
111pub fn split_on_grid_sync(program: &Program) -> Vec<Program> {
112 try_split_on_grid_sync(program).unwrap_or_default()
113}
114
115fn hoist_grid_sync_barriers(nodes: &[Node]) -> Vec<Node> {
121 let mut new_nodes = Vec::new();
122 for node in nodes {
123 match node {
124 Node::Block(body) => {
125 let new_body = hoist_grid_sync_barriers(body);
126 let has_barrier = new_body.iter().any(|n| {
127 matches!(
128 n,
129 Node::Barrier {
130 ordering: MemoryOrdering::GridSync,
131 ..
132 }
133 )
134 });
135 if has_barrier {
136 let mut current_segment = Vec::new();
137 for b_node in new_body {
138 if matches!(
139 b_node,
140 Node::Barrier {
141 ordering: MemoryOrdering::GridSync,
142 ..
143 }
144 ) {
145 new_nodes.push(Node::Block(std::mem::take(&mut current_segment)));
146 new_nodes.push(b_node);
147 } else {
148 current_segment.push(b_node);
149 }
150 }
151 new_nodes.push(Node::Block(current_segment));
152 } else {
153 new_nodes.push(Node::Block(new_body));
154 }
155 }
156 Node::Region {
157 generator,
158 source_region,
159 body,
160 } => {
161 let new_body = hoist_grid_sync_barriers(body);
162 let has_barrier = new_body.iter().any(|n| {
163 matches!(
164 n,
165 Node::Barrier {
166 ordering: MemoryOrdering::GridSync,
167 ..
168 }
169 )
170 });
171 if has_barrier {
172 let mut current_segment = Vec::new();
173 for b_node in new_body {
174 if matches!(
175 b_node,
176 Node::Barrier {
177 ordering: MemoryOrdering::GridSync,
178 ..
179 }
180 ) {
181 new_nodes.push(Node::Region {
182 generator: generator.clone(),
183 source_region: source_region.clone(),
184 body: Arc::new(std::mem::take(&mut current_segment)),
185 });
186 new_nodes.push(b_node);
187 } else {
188 current_segment.push(b_node);
189 }
190 }
191 new_nodes.push(Node::Region {
192 generator: generator.clone(),
193 source_region: source_region.clone(),
194 body: Arc::new(current_segment),
195 });
196 } else {
197 new_nodes.push(Node::Region {
198 generator: generator.clone(),
199 source_region: source_region.clone(),
200 body: Arc::new(new_body),
201 });
202 }
203 }
204 other => {
205 new_nodes.push(other.clone());
206 }
207 }
208 }
209 new_nodes
210}
211
212pub fn try_split_on_grid_sync(program: &Program) -> Result<Vec<Program>, BackendError> {
218 let (wrappers, inner) = peel_entry_wrappers(program);
219 let hoisted_inner = hoist_grid_sync_barriers(inner);
220 let split_count = hoisted_inner
221 .iter()
222 .filter(|node| {
223 matches!(
224 node,
225 Node::Barrier {
226 ordering: MemoryOrdering::GridSync,
227 ..
228 }
229 )
230 })
231 .count();
232 if split_count == 0 {
233 let mut segments = Vec::new();
234 reserve_grid_sync_vec(&mut segments, 1, "grid-sync no-op segment")?;
235 segments.push(program.clone());
236 return Ok(segments);
237 }
238
239 let segment_count = split_count + 1;
240 let executable_nodes = hoisted_inner.len().checked_sub(split_count).ok_or_else(|| {
241 BackendError::InvalidProgram {
242 fix: format!(
243 "grid-sync split_count {split_count} exceeded entry node count {}. Fix: split_on_grid_sync must count barriers from the same entry sequence it segments.",
244 hoisted_inner.len()
245 ),
246 }
247 })?;
248 let segment_capacity = executable_nodes.div_ceil(segment_count);
249
250 let mut raw_segments = Vec::new();
251 let mut current = Vec::new();
252 reserve_grid_sync_vec(&mut current, segment_capacity, "grid-sync current segment")?;
253 for node in &hoisted_inner {
254 match node {
255 Node::Barrier {
256 ordering: MemoryOrdering::GridSync,
257 ..
258 } => {
259 let mut next = Vec::new();
260 reserve_grid_sync_vec(&mut next, segment_capacity, "grid-sync next segment")?;
261 let entry = std::mem::replace(&mut current, next);
262 raw_segments.push(entry);
263 }
264 other => {
265 current.push(other.clone());
266 }
267 }
268 }
269 raw_segments.push(current);
270
271 propagate_let_bindings(&mut raw_segments, &hoisted_inner);
272
273 let mut segments = Vec::new();
274 reserve_grid_sync_vec(
275 &mut segments,
276 raw_segments.len(),
277 "grid-sync split segments",
278 )?;
279 for entry in raw_segments {
280 segments.push(wrap_split_segment(program, &wrappers, entry));
281 }
282 Ok(segments)
283}
284
285fn wrap_split_segment(program: &Program, wrappers: &[EntryWrapper], entry: Vec<Node>) -> Program {
286 let mut wrapped_entry = entry;
290 for wrapper in wrappers.iter().rev() {
291 match wrapper {
292 EntryWrapper::Region { generator } => {
293 wrapped_entry = vec![Node::Region {
294 generator: generator.clone(),
295 source_region: None,
296 body: Arc::new(wrapped_entry),
297 }];
298 }
299 EntryWrapper::Block => {
300 wrapped_entry = vec![Node::Block(wrapped_entry)];
301 }
302 }
303 }
304 program.with_rewritten_entry(wrapped_entry)
305}
306
307#[cfg(test)]
308mod tests {
309 use super::*;
310 use crate::grid_sync::test_programs::{buffer, region};
311 use vyre_foundation::ir::Expr;
312
313 fn inner_len(program: &Program) -> usize {
315 entry_sequence(program).len()
316 }
317
318 #[test]
319 fn no_grid_sync_returns_single_segment() {
320 let program = Program::wrapped(
321 vec![buffer()],
322 [1, 1, 1],
323 vec![region(
324 "a",
325 vec![Node::store("buf", Expr::u32(0), Expr::u32(1))],
326 )],
327 );
328 assert!(!contains_grid_sync(&program));
329 let segments = split_on_grid_sync(&program);
330 assert_eq!(segments.len(), 1);
331 assert_eq!(inner_len(&segments[0]), 1);
333 }
334
335 #[test]
336 fn one_grid_sync_splits_into_two() {
337 let program = Program::wrapped(
338 vec![buffer()],
339 [1, 1, 1],
340 vec![
341 region("a", vec![Node::store("buf", Expr::u32(0), Expr::u32(1))]),
342 Node::barrier_with_ordering(MemoryOrdering::GridSync),
343 region("b", vec![Node::store("buf", Expr::u32(1), Expr::u32(2))]),
344 ],
345 );
346 assert!(contains_grid_sync(&program));
347 let segments = split_on_grid_sync(&program);
348 assert_eq!(segments.len(), 2);
349 assert_eq!(inner_len(&segments[0]), 1);
350 assert_eq!(inner_len(&segments[1]), 1);
351 }
352
353 #[test]
354 fn block_nested_grid_sync_splits_into_two() {
355 let program = Program::wrapped(
356 vec![buffer()],
357 [1, 1, 1],
358 vec![Node::Block(vec![
359 region("a", vec![Node::store("buf", Expr::u32(0), Expr::u32(1))]),
360 Node::barrier_with_ordering(MemoryOrdering::GridSync),
361 region("b", vec![Node::store("buf", Expr::u32(1), Expr::u32(2))]),
362 ])],
363 );
364 assert!(contains_grid_sync(&program));
365 let segments = split_on_grid_sync(&program);
366 assert_eq!(segments.len(), 2);
367 assert_eq!(inner_len(&segments[0]), 1);
368 assert_eq!(inner_len(&segments[1]), 1);
369 }
370
371 #[test]
372 fn three_grid_syncs_split_into_four() {
373 let program = Program::wrapped(
374 vec![buffer()],
375 [1, 1, 1],
376 vec![
377 region("a", vec![Node::Return]),
378 Node::barrier_with_ordering(MemoryOrdering::GridSync),
379 region("b", vec![Node::Return]),
380 Node::barrier_with_ordering(MemoryOrdering::GridSync),
381 region("c", vec![Node::Return]),
382 Node::barrier_with_ordering(MemoryOrdering::GridSync),
383 region("d", vec![Node::Return]),
384 ],
385 );
386 let segments = split_on_grid_sync(&program);
387 assert_eq!(segments.len(), 4);
388 }
389
390 #[test]
391 fn workgroup_barrier_does_not_split() {
392 let program = Program::wrapped(
393 vec![buffer()],
394 [1, 1, 1],
395 vec![
396 region("a", vec![Node::Return]),
397 Node::barrier_with_ordering(MemoryOrdering::SeqCst),
398 region("b", vec![Node::Return]),
399 ],
400 );
401 assert!(!contains_grid_sync(&program));
402 let segments = split_on_grid_sync(&program);
403 assert_eq!(segments.len(), 1);
404 assert_eq!(inner_len(&segments[0]), 3);
406 }
407
408 #[test]
409 fn buffers_and_workgroup_size_propagate_to_each_segment() {
410 let program = Program::wrapped(
411 vec![buffer()],
412 [256, 1, 1],
413 vec![
414 region("a", vec![Node::Return]),
415 Node::barrier_with_ordering(MemoryOrdering::GridSync),
416 region("b", vec![Node::Return]),
417 ],
418 );
419 let segments = split_on_grid_sync(&program);
420 for seg in &segments {
421 assert_eq!(seg.workgroup_size(), [256, 1, 1]);
422 assert_eq!(seg.buffers().len(), 1);
423 assert_eq!(seg.buffers()[0].name(), "buf");
424 }
425 }
426}