use std::collections::{HashMap, HashSet};
use vyre_foundation::ir::{BufferAccess, BufferDecl, Expr, Ident, MemoryKind, Node, Program};
use super::barrier_split::{entry_sequence, try_split_on_grid_sync};
use super::{reserve_grid_sync_hash_map, reserve_grid_sync_hash_set, reserve_grid_sync_vec};
use crate::backend::BackendError;
pub(super) struct PlannedGridSyncSegment {
pub(super) program: Program,
pub(super) input_names: Vec<Ident>,
pub(super) output_names: Vec<Ident>,
}
pub fn plan_host_grid_sync_segment_programs(
program: &Program,
) -> Result<Vec<Program>, BackendError> {
Ok(plan_host_grid_sync_segments(program)?
.into_iter()
.map(|segment| segment.program)
.collect())
}
pub(super) fn plan_host_grid_sync_segments(
program: &Program,
) -> Result<Vec<PlannedGridSyncSegment>, BackendError> {
let split = try_split_on_grid_sync(program)?;
let first_writer = first_writer_segment_per_buffer(&split, program)?;
let mut planned = Vec::new();
reserve_grid_sync_vec(&mut planned, split.len(), "grid-sync planned host segments")?;
for (segment_idx, segment) in split.into_iter().enumerate() {
let rewritten =
rewrite_segment_buffers_for_host_split(program, &segment, segment_idx, &first_writer)?;
let input_names = segment_input_names(&rewritten)?;
let output_names = segment_output_names(&rewritten)?;
planned.push(PlannedGridSyncSegment {
program: rewritten,
input_names,
output_names,
});
}
Ok(planned)
}
fn first_writer_segment_per_buffer(
split: &[Program],
program: &Program,
) -> Result<HashMap<Ident, usize>, BackendError> {
let mut first_writer: HashMap<Ident, usize> = HashMap::new();
reserve_grid_sync_hash_map(
&mut first_writer,
program.buffers().len(),
"grid-sync first-writer map",
)?;
for (segment_idx, segment) in split.iter().enumerate() {
let mut reads = HashSet::new();
let mut writes = HashSet::new();
reserve_grid_sync_hash_set(
&mut reads,
program.buffers().len(),
"grid-sync first-writer read scan",
)?;
reserve_grid_sync_hash_set(
&mut writes,
program.buffers().len(),
"grid-sync first-writer write scan",
)?;
for node in entry_sequence(segment) {
collect_segment_buffer_targets(node, &mut reads, &mut writes);
}
for name in writes {
first_writer.entry(name).or_insert(segment_idx);
}
}
Ok(first_writer)
}
fn rewrite_segment_buffers_for_host_split(
source: &Program,
segment: &Program,
segment_idx: usize,
first_writer: &HashMap<Ident, usize>,
) -> Result<Program, BackendError> {
let mut reads = HashSet::new();
let mut writes = HashSet::new();
reserve_grid_sync_hash_set(
&mut reads,
source.buffers().len(),
"grid-sync segment read set",
)?;
reserve_grid_sync_hash_set(
&mut writes,
source.buffers().len(),
"grid-sync segment write set",
)?;
for node in entry_sequence(segment) {
collect_segment_buffer_targets(node, &mut reads, &mut writes);
}
let mut buffers = Vec::new();
reserve_grid_sync_vec(
&mut buffers,
source.buffers().len(),
"grid-sync segment buffers",
)?;
for buffer in source.buffers() {
let name = Ident::from(buffer.name());
let reads_this = reads.contains(&name);
let writes_this = writes.contains(&name);
let readwrite_passthrough = matches!(buffer.access(), BufferAccess::ReadWrite)
&& !buffer.is_output()
&& !buffer.is_pipeline_live_out()
&& !reads_this
&& !writes_this;
if !reads_this && !writes_this && !readwrite_passthrough {
continue;
}
let mut rewritten = buffer.clone();
if matches!(rewritten.access(), BufferAccess::Workgroup) {
buffers.push(rewritten);
continue;
}
let is_source_output = buffer.is_output() || buffer.is_pipeline_live_out();
let earlier_segment_wrote_output = is_source_output
&& first_writer
.get(&name)
.is_some_and(|&first| first < segment_idx);
let access = if readwrite_passthrough {
BufferAccess::ReadWrite
} else if earlier_segment_wrote_output && writes_this {
BufferAccess::ReadWrite
} else {
match (reads_this, writes_this) {
(true, true) => BufferAccess::ReadWrite,
(true, false) => BufferAccess::ReadOnly,
(false, true) => BufferAccess::WriteOnly,
(false, false) => BufferAccess::ReadWrite,
}
};
rewrite_segment_buffer_access(&mut rewritten, access);
rewritten.is_output = false;
rewritten.pipeline_live_out = false;
buffers.push(rewritten);
}
Ok(segment.with_rewritten_buffers(buffers))
}
fn rewrite_segment_buffer_access(buffer: &mut BufferDecl, access: BufferAccess) {
buffer.kind = match &access {
BufferAccess::ReadOnly => MemoryKind::Readonly,
BufferAccess::Uniform => MemoryKind::Uniform,
BufferAccess::Workgroup => MemoryKind::Shared,
_ => MemoryKind::Global,
};
buffer.access = access;
}
pub(super) fn segment_input_names(segment: &Program) -> Result<Vec<Ident>, BackendError> {
let mut names = Vec::new();
reserve_grid_sync_vec(
&mut names,
segment.buffers().len(),
"grid-sync segment input names",
)?;
for buffer in segment.buffers() {
if matches!(buffer.access(), BufferAccess::Workgroup) {
continue;
}
if segment_buffer_consumes_input(buffer) {
names.push(Ident::from(buffer.name()));
}
}
Ok(names)
}
pub(super) fn segment_output_names(segment: &Program) -> Result<Vec<Ident>, BackendError> {
let mut names = Vec::new();
reserve_grid_sync_vec(
&mut names,
segment.buffers().len(),
"grid-sync segment output names",
)?;
for buffer in segment.buffers() {
if matches!(buffer.access(), BufferAccess::Workgroup) {
continue;
}
if segment_buffer_produces_output(buffer) {
names.push(Ident::from(buffer.name()));
}
}
Ok(names)
}
pub(super) fn original_input_names(program: &Program) -> Result<Vec<Ident>, BackendError> {
segment_input_names(program)
}
pub(super) fn original_output_names(program: &Program) -> Result<Vec<Ident>, BackendError> {
segment_output_names(program)
}
pub(super) fn segment_buffer_consumes_input(buffer: &BufferDecl) -> bool {
if buffer.is_output() || buffer.is_pipeline_live_out() {
return false;
}
matches!(
buffer.access(),
BufferAccess::ReadOnly | BufferAccess::ReadWrite | BufferAccess::Uniform
)
}
pub(super) fn segment_buffer_produces_output(buffer: &BufferDecl) -> bool {
buffer.is_output()
|| buffer.is_pipeline_live_out()
|| matches!(
buffer.access(),
BufferAccess::ReadWrite | BufferAccess::WriteOnly
)
}
fn collect_segment_buffer_targets(
node: &Node,
reads: &mut HashSet<Ident>,
writes: &mut HashSet<Ident>,
) {
match node {
Node::Let { value, .. } | Node::Assign { value, .. } => {
collect_segment_expr_targets(value, reads, writes);
}
Node::Store {
buffer,
index,
value,
} => {
writes.insert(Ident::from(buffer));
collect_segment_expr_targets(index, reads, writes);
collect_segment_expr_targets(value, reads, writes);
}
Node::If {
cond,
then,
otherwise,
} => {
collect_segment_expr_targets(cond, reads, writes);
for child in then.iter().chain(otherwise.iter()) {
collect_segment_buffer_targets(child, reads, writes);
}
}
Node::Loop { from, to, body, .. } => {
collect_segment_expr_targets(from, reads, writes);
collect_segment_expr_targets(to, reads, writes);
for child in body {
collect_segment_buffer_targets(child, reads, writes);
}
}
Node::Block(body) => {
for child in body {
collect_segment_buffer_targets(child, reads, writes);
}
}
Node::Region { body, .. } => {
for child in body.iter() {
collect_segment_buffer_targets(child, reads, writes);
}
}
Node::AllReduce { buffer, .. } | Node::Broadcast { buffer, .. } => {
reads.insert(buffer.clone());
writes.insert(buffer.clone());
}
Node::AllGather { input, output, .. } | Node::ReduceScatter { input, output, .. } => {
reads.insert(input.clone());
writes.insert(output.clone());
}
Node::IndirectDispatch { .. }
| Node::Return
| Node::Barrier { .. }
| Node::AsyncLoad { .. }
| Node::AsyncStore { .. }
| Node::AsyncWait { .. }
| Node::Trap { .. }
| Node::Resume { .. }
| Node::Opaque(_) => {}
_ => {}
}
}
fn collect_segment_expr_targets(
expr: &Expr,
reads: &mut HashSet<Ident>,
writes: &mut HashSet<Ident>,
) {
vyre_foundation::visit::visit_expr_buffer_accesses(expr, |access, buffer| {
reads.insert(buffer.clone());
if access == vyre_foundation::visit::ExprBufferAccess::Atomic {
writes.insert(buffer.clone());
}
});
}
#[cfg(test)]
mod tests {
use super::*;
use crate::grid_sync::test_programs::region;
use vyre_foundation::ir::DataType;
use vyre_foundation::memory_model::MemoryOrdering;
#[test]
fn split_keeps_multi_segment_output_as_readwrite_accumulator() {
let out = BufferDecl::output("out", 0, DataType::U32).with_count(4);
let program = Program::wrapped(
vec![out],
[1, 1, 1],
vec![
region("a", vec![Node::store("out", Expr::u32(0), Expr::u32(0xAA))]),
Node::barrier_with_ordering(MemoryOrdering::GridSync),
region("b", vec![Node::store("out", Expr::u32(2), Expr::u32(0xBB))]),
],
);
let segments =
plan_host_grid_sync_segment_programs(&program).expect("plan host grid-sync segments");
assert_eq!(segments.len(), 2, "one GridSync barrier -> two segments");
let seg0_out = segments[0]
.buffers()
.iter()
.find(|b| b.name() == "out")
.expect("segment 0 must declare the output it writes");
assert_eq!(
seg0_out.access(),
BufferAccess::WriteOnly,
"the first writer establishes the accumulator as write-only"
);
assert!(
!seg0_out.is_output() && !seg0_out.is_pipeline_live_out(),
"split segment buffers must never be marked program-output; final values are reassembled by name"
);
let seg1_out = segments[1]
.buffers()
.iter()
.find(|b| b.name() == "out")
.expect("segment 1 must declare the output it writes");
assert_eq!(
seg1_out.access(),
BufferAccess::ReadWrite,
"a later writer of a multi-segment output must read+merge the accumulated value, not overwrite it"
);
assert!(
!seg1_out.is_output() && !seg1_out.is_pipeline_live_out(),
"the later writer must consume its forwarded prior value, which `segment_buffer_consumes_input` refuses for is_output buffers"
);
assert!(
segment_input_names(&segments[1])
.expect("segment 1 input names")
.iter()
.any(|n| n.as_str() == "out"),
"the accumulated output must be forwarded as an input to the later writing segment"
);
}
}