vyre_foundation/execution_plan/fusion/
fuse.rs1use rustc_hash::{FxHashMap, FxHashSet};
4
5use crate::execution_plan::SchedulingPolicy;
6use crate::ir::{BufferAccess, BufferDecl, Ident, Node, Program};
7
8use super::alpha_rename::{multiply_declared_names, push_alpha_renamed_arm_entry_node, ArmRenamer};
9use super::collectors::collect_buffer_targets;
10use super::divergence::{
11 has_divergent_invocation_gated_store, has_launch_geometry_dependent_write,
12};
13use super::{
14 FusionError, FusionOverDispatchError, FusionSelfAliasingError, FusionWorkgroupGeometryError,
15};
16
17pub fn fuse_programs(programs: &[Program]) -> Result<Program, FusionError> {
25 match programs.len() {
26 0 => Ok(Program::empty()),
27 1 => Ok(programs[0].clone()),
28 _ => fuse_programs_multi(programs),
29 }
30}
31
32#[inline]
41#[must_use]
42pub fn fuse_programs_vec(mut programs: Vec<Program>) -> Result<Program, FusionError> {
43 match programs.len() {
44 0 => Ok(Program::empty()),
45 1 => {
46 let Some(program) = programs.pop() else {
47 return Ok(Program::empty());
48 };
49 Ok(program)
50 }
51 _ => fuse_programs_multi(programs.as_slice()),
52 }
53}
54
55#[derive(Clone, Copy, PartialEq, Eq)]
57pub(crate) enum ArmNamespace {
58 Isolated,
65 Shared,
76}
77
78pub fn merge_programs_shared(programs: &[Program]) -> Result<Program, FusionError> {
92 match programs.len() {
93 0 => Ok(Program::empty()),
94 1 => Ok(programs[0].clone()),
95 _ => fuse_programs_multi_with(programs, ArmNamespace::Shared),
96 }
97}
98
99fn fuse_programs_multi(programs: &[Program]) -> Result<Program, FusionError> {
100 fuse_programs_multi_with(programs, ArmNamespace::Isolated)
101}
102
103fn fuse_programs_multi_with(
104 programs: &[Program],
105 namespace: ArmNamespace,
106) -> Result<Program, FusionError> {
107 reject_non_composable_self_fusion(programs)?;
108
109 let mut merged_buffers: Vec<BufferDecl> = Vec::new();
114 let mut name_to_index: FxHashMap<Ident, usize> = FxHashMap::default();
115 let mut next_binding = 0_u32;
116
117 let mut read_arms_per_buffer: FxHashMap<Ident, Vec<usize>> = FxHashMap::default();
118 let mut write_arms_per_buffer: FxHashMap<Ident, Vec<usize>> = FxHashMap::default();
125 let mut barrier_after_arm: FxHashSet<usize> = FxHashSet::default();
126 let mut grid_sync_writer_arms: FxHashSet<usize> = FxHashSet::default();
131
132 let mut fused_workgroup = [1u32, 1, 1];
133 let mut max_arm_threads: u64 = 1;
134
135 let mut arm_entries: Vec<Vec<Node>> = Vec::with_capacity(programs.len());
136
137 let multiply_declared: FxHashSet<Ident> = match namespace {
144 ArmNamespace::Isolated => FxHashSet::default(),
145 ArmNamespace::Shared => {
146 let entries: Vec<&[Node]> = programs.iter().map(Program::entry).collect();
147 multiply_declared_names(&entries)
148 }
149 };
150
151 for (arm_idx, prog) in programs.iter().enumerate() {
152 let entry = prog.entry();
162 let mut segment = Vec::with_capacity(entry.len());
163 let mut atomic_targets: FxHashSet<Ident> = FxHashSet::default();
164 let mut load_targets: FxHashSet<Ident> = FxHashSet::default();
165 let mut store_targets: FxHashSet<Ident> = FxHashSet::default();
166 let mut divergent_store_seen = false;
167 for node in entry {
168 match namespace {
169 ArmNamespace::Isolated => {
170 push_alpha_renamed_arm_entry_node(&mut segment, node, arm_idx);
171 }
172 ArmNamespace::Shared => {
173 ArmRenamer::shared(arm_idx, &multiply_declared)
174 .push_entry_node(&mut segment, node);
175 }
176 }
177 collect_buffer_targets(
178 node,
179 &mut load_targets,
180 &mut store_targets,
181 &mut atomic_targets,
182 );
183 if has_divergent_invocation_gated_store(node, false) {
184 divergent_store_seen = true;
185 }
186 }
187 if divergent_store_seen || has_launch_geometry_dependent_write(prog.entry()) {
188 grid_sync_writer_arms.insert(arm_idx);
189 }
190 arm_entries.push(segment);
191
192 let mut arm_reads: FxHashSet<Ident> = FxHashSet::default();
193 let mut arm_explicit_writes: FxHashSet<Ident> = FxHashSet::default();
194 classify_and_merge_arm_buffers(
195 prog,
196 &mut arm_reads,
197 &mut arm_explicit_writes,
198 &mut merged_buffers,
199 &mut name_to_index,
200 &mut next_binding,
201 );
202
203 for target in &load_targets {
208 arm_reads.insert(target.clone());
209 }
210 for target in &store_targets {
212 arm_explicit_writes.insert(target.clone());
213 }
214
215 let mut arm_writes = arm_explicit_writes.clone();
217 for target in &atomic_targets {
218 if !arm_reads.contains(target) && !arm_explicit_writes.contains(target) {
219 arm_writes.insert(target.clone());
220 }
221 }
222
223 for write_buf in &arm_writes {
227 if let Some(read_arms) = read_arms_per_buffer.get(write_buf) {
228 for &read_arm in read_arms {
229 barrier_after_arm.insert(read_arm);
230 }
231 }
232 }
233
234 for read_buf in &arm_reads {
245 if let Some(write_arms) = write_arms_per_buffer.get(read_buf) {
246 for &write_arm in write_arms {
247 barrier_after_arm.insert(write_arm);
248 }
249 }
250 }
251
252 for read_buf in &arm_reads {
254 read_arms_per_buffer
255 .entry(read_buf.clone())
256 .or_default()
257 .push(arm_idx);
258 }
259 for write_buf in &arm_writes {
261 write_arms_per_buffer
262 .entry(write_buf.clone())
263 .or_default()
264 .push(arm_idx);
265 }
266
267 let wg = prog.workgroup_size();
269 fused_workgroup[0] = fused_workgroup[0].max(wg[0]);
270 fused_workgroup[1] = fused_workgroup[1].max(wg[1]);
271 fused_workgroup[2] = fused_workgroup[2].max(wg[2]);
272 let arm_threads = u64::from(wg[0]) * u64::from(wg[1]) * u64::from(wg[2]);
273 max_arm_threads = max_arm_threads.max(arm_threads);
274 }
275
276 reject_workgroup_geometry_change(programs, fused_workgroup)?;
277
278 let combined_entry = flatten_arm_entries(
279 arm_entries,
280 &barrier_after_arm,
281 &grid_sync_writer_arms,
282 programs.len(),
283 namespace,
284 );
285 reject_overdispatch(fused_workgroup, max_arm_threads)?;
286
287 let non_composable = programs.iter().any(Program::is_non_composable_with_self);
298 Ok(
299 Program::wrapped(merged_buffers, fused_workgroup, combined_entry)
300 .with_non_composable_with_self(non_composable),
301 )
302}
303
304fn classify_and_merge_arm_buffers(
305 prog: &Program,
306 arm_reads: &mut FxHashSet<Ident>,
307 arm_explicit_writes: &mut FxHashSet<Ident>,
308 merged_buffers: &mut Vec<BufferDecl>,
309 name_to_index: &mut FxHashMap<Ident, usize>,
310 next_binding: &mut u32,
311) {
312 for buf in prog.buffers() {
313 let name = Ident::from(buf.name());
314 match buf.access() {
315 BufferAccess::ReadOnly | BufferAccess::Uniform => {
316 arm_reads.insert(name.clone());
317 }
318 BufferAccess::ReadWrite => {
319 arm_explicit_writes.insert(name.clone());
320 }
321 _ => {}
322 }
323 if let Some(&idx) = name_to_index.get(&name) {
324 let existing = &mut merged_buffers[idx];
325 let access = buf.access();
326 upgrade_buffer_access(existing, &access);
327 if buf.count > existing.count {
328 existing.count = buf.count;
329 }
330 if buf.is_output() {
331 existing.is_output = true;
332 existing.pipeline_live_out = true;
333 }
334 } else {
335 let mut merged = buf.clone();
336 if merged.access() != BufferAccess::Workgroup {
337 merged.binding = *next_binding;
338 *next_binding += 1;
339 }
340 name_to_index.insert(Ident::from(merged.name()), merged_buffers.len());
341 merged_buffers.push(merged);
342 }
343 }
344}
345
346fn reject_non_composable_self_fusion(programs: &[Program]) -> Result<(), FusionError> {
347 let mut seen_op_ids: FxHashMap<String, bool> = FxHashMap::default();
348 for prog in programs {
349 let key = prog
350 .entry_op_id()
351 .map_or_else(|| fallback_composition_key(prog), ToString::to_string);
352 let is_non_comp = prog.is_non_composable_with_self();
353 match seen_op_ids.get_mut(&key) {
354 Some(has_non_comp) if *has_non_comp || is_non_comp => {
355 return Err(FusionError::SelfAliasing(FusionSelfAliasingError {
356 op_id: key,
357 fix: "rename the second parser's workgroup buffer or split into two separate dispatches",
358 }));
359 }
360 Some(_) => {}
361 None => {
362 seen_op_ids.insert(key, is_non_comp);
363 }
364 }
365 }
366 Ok(())
367}
368
369fn reject_workgroup_geometry_change(
384 programs: &[Program],
385 fused_workgroup: [u32; 3],
386) -> Result<(), FusionError> {
387 for (arm, prog) in programs.iter().enumerate() {
388 let arm_workgroup = prog.workgroup_size();
389 if arm_workgroup == fused_workgroup {
390 continue;
391 }
392 let uses_workgroup_memory = prog
393 .buffers()
394 .iter()
395 .any(|buf| buf.access() == BufferAccess::Workgroup);
396 let synchronizes = has_barrier(prog.entry());
397 let reason = match (uses_workgroup_memory, synchronizes) {
398 (true, true) => "keeps state in workgroup memory and synchronizes its workgroup",
399 (true, false) => "keeps state in workgroup memory sized for its own workgroup",
400 (false, true) => "synchronizes its workgroup with a barrier",
401 (false, false) => continue,
402 };
403 return Err(FusionError::WorkgroupGeometry(
404 FusionWorkgroupGeometryError {
405 arm,
406 arm_workgroup,
407 fused_workgroup,
408 reason,
409 fix: "dispatch this arm separately, or rebuild it for the wider workgroup before fusing",
410 },
411 ));
412 }
413 Ok(())
414}
415
416fn has_barrier(nodes: &[Node]) -> bool {
418 nodes.iter().any(|node| match node {
419 Node::Barrier { .. } => true,
420 Node::Region { body, .. } => has_barrier(body),
421 Node::Block(body) | Node::Loop { body, .. } => has_barrier(body),
422 Node::If {
423 then, otherwise, ..
424 } => has_barrier(then) || has_barrier(otherwise),
425 _ => false,
426 })
427}
428
429fn flatten_arm_entries(
430 arm_entries: Vec<Vec<Node>>,
431 barrier_after_arm: &FxHashSet<usize>,
432 grid_sync_writer_arms: &FxHashSet<usize>,
433 program_count: usize,
434 namespace: ArmNamespace,
435) -> Vec<Node> {
436 let total_nodes: usize = arm_entries.iter().map(Vec::len).sum();
437 let mut combined_entry = Vec::with_capacity(total_nodes + program_count);
438 for (arm_idx, segment) in arm_entries.into_iter().enumerate() {
439 match namespace {
440 ArmNamespace::Isolated => combined_entry.push(Node::Block(segment)),
443 ArmNamespace::Shared => combined_entry.extend(segment),
446 }
447 if barrier_after_arm.contains(&arm_idx) {
448 let ordering = if grid_sync_writer_arms.contains(&arm_idx) {
454 crate::memory_model::MemoryOrdering::GridSync
455 } else {
456 crate::memory_model::MemoryOrdering::SeqCst
457 };
458 combined_entry.push(Node::barrier_with_ordering(ordering));
459 }
460 }
461 combined_entry
462}
463
464fn reject_overdispatch(fused_workgroup: [u32; 3], max_arm_threads: u64) -> Result<(), FusionError> {
465 let fused_threads = u64::from(fused_workgroup[0])
466 * u64::from(fused_workgroup[1])
467 * u64::from(fused_workgroup[2]);
468 let policy = SchedulingPolicy::standard();
469 if policy.allow_fused_threads(fused_threads, max_arm_threads) {
470 return Ok(());
471 }
472 Err(FusionError::OverDispatch(FusionOverDispatchError {
473 max_arm_threads,
474 fused_threads,
475 fix: "split the batch or use per-arm dispatch; axis-wise max exceeds the shared over-dispatch policy",
476 }))
477}
478
479pub(super) fn fallback_composition_key(prog: &Program) -> String {
480 let mut hasher = blake3::Hasher::new();
481 for buf in prog.buffers() {
482 hasher.update(buf.name().as_bytes());
483 hasher.update(&[0]);
484 }
485 for dim in prog.workgroup_size() {
486 hasher.update(&dim.to_le_bytes());
487 }
488 hasher.update(&(prog.entry().len() as u64).to_le_bytes());
489 format!("{}", hasher.finalize().to_hex())
490}
491
492pub(super) fn upgrade_buffer_access(buffer: &mut BufferDecl, other: &BufferAccess) {
494 let current = buffer.access();
495 buffer.access = match (¤t, &other) {
496 (BufferAccess::ReadWrite, _)
497 | (_, BufferAccess::ReadWrite)
498 | (BufferAccess::WriteOnly, BufferAccess::ReadOnly | BufferAccess::Uniform)
499 | (BufferAccess::ReadOnly | BufferAccess::Uniform, BufferAccess::WriteOnly) => {
500 BufferAccess::ReadWrite
501 }
502 (BufferAccess::WriteOnly, BufferAccess::WriteOnly) => BufferAccess::WriteOnly,
503 (BufferAccess::Uniform, _) | (_, BufferAccess::Uniform) => BufferAccess::Uniform,
504 (BufferAccess::Workgroup, _) | (_, BufferAccess::Workgroup) => BufferAccess::Workgroup,
505 _ => BufferAccess::ReadOnly,
506 };
507 buffer.kind = match buffer.access {
509 BufferAccess::ReadOnly => crate::ir::MemoryKind::Readonly,
510 BufferAccess::Uniform => crate::ir::MemoryKind::Uniform,
511 BufferAccess::Workgroup => crate::ir::MemoryKind::Shared,
512 _ => crate::ir::MemoryKind::Global,
513 };
514}