Skip to main content

vyre_driver/grid_sync/
barrier_split.rs

1//! Grid-sync barrier detection and the split of a program's entry sequence
2//! into one segment per barrier.
3
4use 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/// Walk past `Program::wrapped`'s synthetic outer Region. Real
14/// programs are constructed via `wrapped`, which inserts a single
15/// outer Region around the user's entry sequence; the split logic
16/// must operate on the inner sequence so a `GridSync` barrier inside
17/// the wrapper actually splits the program. Programs constructed
18/// via `Program::new` use the entry sequence directly  -  in that
19/// case we just return it unchanged.
20#[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/// Whether `program` contains any `Node::Barrier { ordering: GridSync }`
59/// in its dispatch-level entry sequence (peeled past any synthetic
60/// outer Region).
61///
62/// The check is intentionally shallow: nested grid-sync barriers
63/// inside `Node::Loop` or inner `Node::Region` bodies are a contract
64/// violation (`validate::barrier` rejects them) and never reach this
65/// path. The split operates at the dispatch-level granularity.
66#[must_use]
67pub fn contains_grid_sync(program: &Program) -> bool {
68    // O(1) negative gate: if the cached ProgramStats bitset records no
69    // Barrier of any kind in the entire tree, there is definitely no
70    // top-level GridSync barrier either. Skip the entry-sequence walk
71    // (which itself is shallow but still pays a buffers/buffer_index
72    // dispatch on every backend dispatch path).
73    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/// Split `program` at every top-level `Node::Barrier { GridSync }`.
99///
100/// Returns a vector of segments in execution order. The barrier nodes
101/// themselves are dropped from the segments  -  the kernel-launch
102/// boundary between segments takes their place.
103///
104/// Each returned segment is a complete `Program` that shares the
105/// original's buffer table, workgroup size, and metadata; only the
106/// entry sequence changes. Segments without any executable nodes are
107/// preserved (an empty segment between two adjacent barriers becomes
108/// a no-op kernel that completes with byte-identical inputs and
109/// outputs).
110#[must_use]
111pub fn split_on_grid_sync(program: &Program) -> Vec<Program> {
112    try_split_on_grid_sync(program).unwrap_or_default()
113}
114
115/// Fallible variant of [`split_on_grid_sync`] for production dispatch paths.
116///
117/// # Errors
118/// Returns an actionable [`BackendError`] if segment storage cannot be
119/// reserved or if split accounting overflows.
120fn 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
212/// Fallible variant of [`split_on_grid_sync`] for production dispatch paths.
213///
214/// # Errors
215/// Returns an actionable [`BackendError`] if segment storage cannot be
216/// reserved or if split accounting overflows.
217pub 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    // Re-wrap each segment in the same wrapper stack the source had,
287    // so tagged/fused programs keep provenance and structure while the
288    // executable body is split at launch boundaries.
289    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    /// Get the inner-segment node count for a wrapped or unwrapped Program.
314    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        // Original entry was [Region("a", ...)] so the inner sequence is 1.
332        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        // Region("a"), Barrier(SeqCst), Region("b") = 3 inner nodes.
405        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}