use crate::parsing::c::lex::tokens::*;
use crate::parsing::composition::child_phase;
use crate::region::wrap_anonymous;
use vyre::ir::{BufferAccess, BufferDecl, DataType, Expr, Node, Program};
use vyre_foundation::memory_model::MemoryOrdering;
const STMT_TILE_TOKENS: u32 = 256;
const STMT_TILE_WORDS: u32 = 8;
const TILE_PAREN_CLOSERS: u32 = 0;
const TILE_PAREN_OPENERS: u32 = 1;
const TILE_BRACKET_CLOSERS: u32 = 2;
const TILE_BRACKET_OPENERS: u32 = 3;
const TILE_PAREN_ENTRY: u32 = 4;
const TILE_BRACKET_ENTRY: u32 = 5;
const TILE_FIRST_BOUNDARY: u32 = 6;
const TILE_NEXT_BOUNDARY: u32 = 7;
const STMT_NO_BOUNDARY: u32 = u32::MAX;
pub const C11_STATEMENT_BOUNDS_SCRATCH: &str = "c11_stmt_boundary_scratch";
#[must_use]
pub const fn c11_statement_bounds_scratch_words(num_tokens: u32) -> u32 {
let tokens = if num_tokens == 0 { 1 } else { num_tokens };
let tiles = (tokens - 1) / STMT_TILE_TOKENS + 1;
tokens
.saturating_add(tiles.saturating_mul(STMT_TILE_WORDS))
.saturating_add(STMT_TILE_WORDS)
}
struct BoundsCtx<'a> {
tok_types: &'a str,
tile_base: u32,
t: Expr,
active: Expr,
tile_count: Expr,
}
impl BoundsCtx<'_> {
fn tile_slot(&self, tile: Expr, offset: u32) -> Expr {
Expr::add(
Expr::add(
Expr::u32(self.tile_base),
Expr::mul(tile, Expr::u32(STMT_TILE_WORDS)),
),
Expr::u32(offset),
)
}
fn own_slot(&self, offset: u32) -> Expr {
self.tile_slot(self.t.clone(), offset)
}
fn own_tile_lo(&self) -> Expr {
Expr::mul(self.t.clone(), Expr::u32(STMT_TILE_TOKENS))
}
fn own_tile_hi(&self) -> Expr {
Expr::min(
Expr::add(self.own_tile_lo(), Expr::u32(STMT_TILE_TOKENS)),
self.active.clone(),
)
}
fn clamped_dec(name: &str) -> Expr {
Expr::sub(Expr::max(Expr::var(name), Expr::u32(1)), Expr::u32(1))
}
fn pass_tile_reduce(&self) -> Vec<Node> {
vec![
Node::let_bind("reduce_lo", self.own_tile_lo()),
Node::let_bind("reduce_hi", self.own_tile_hi()),
Node::let_bind("paren_closers", Expr::u32(0)),
Node::let_bind("paren_openers", Expr::u32(0)),
Node::let_bind("bracket_closers", Expr::u32(0)),
Node::let_bind("bracket_openers", Expr::u32(0)),
Node::loop_for(
"reduce_tok",
Expr::var("reduce_lo"),
Expr::var("reduce_hi"),
vec![
Node::let_bind(
"reduce_token",
Expr::load(self.tok_types, Expr::var("reduce_tok")),
),
Node::if_then(
Expr::eq(Expr::var("reduce_token"), Expr::u32(TOK_LPAREN)),
vec![Node::assign(
"paren_openers",
Expr::add(Expr::var("paren_openers"), Expr::u32(1)),
)],
),
Node::if_then(
Expr::eq(Expr::var("reduce_token"), Expr::u32(TOK_RPAREN)),
vec![Node::if_then_else(
Expr::gt(Expr::var("paren_openers"), Expr::u32(0)),
vec![Node::assign(
"paren_openers",
Expr::sub(Expr::var("paren_openers"), Expr::u32(1)),
)],
vec![Node::assign(
"paren_closers",
Expr::add(Expr::var("paren_closers"), Expr::u32(1)),
)],
)],
),
Node::if_then(
Expr::eq(Expr::var("reduce_token"), Expr::u32(TOK_LBRACKET)),
vec![Node::assign(
"bracket_openers",
Expr::add(Expr::var("bracket_openers"), Expr::u32(1)),
)],
),
Node::if_then(
Expr::eq(Expr::var("reduce_token"), Expr::u32(TOK_RBRACKET)),
vec![Node::if_then_else(
Expr::gt(Expr::var("bracket_openers"), Expr::u32(0)),
vec![Node::assign(
"bracket_openers",
Expr::sub(Expr::var("bracket_openers"), Expr::u32(1)),
)],
vec![Node::assign(
"bracket_closers",
Expr::add(Expr::var("bracket_closers"), Expr::u32(1)),
)],
)],
),
],
),
Node::store(
C11_STATEMENT_BOUNDS_SCRATCH,
self.own_slot(TILE_PAREN_CLOSERS),
Expr::var("paren_closers"),
),
Node::store(
C11_STATEMENT_BOUNDS_SCRATCH,
self.own_slot(TILE_PAREN_OPENERS),
Expr::var("paren_openers"),
),
Node::store(
C11_STATEMENT_BOUNDS_SCRATCH,
self.own_slot(TILE_BRACKET_CLOSERS),
Expr::var("bracket_closers"),
),
Node::store(
C11_STATEMENT_BOUNDS_SCRATCH,
self.own_slot(TILE_BRACKET_OPENERS),
Expr::var("bracket_openers"),
),
]
}
fn pass_tile_compose(&self) -> Vec<Node> {
vec![
Node::let_bind("entry_paren", Expr::u32(0)),
Node::let_bind("entry_bracket", Expr::u32(0)),
Node::store(
C11_STATEMENT_BOUNDS_SCRATCH,
self.tile_slot(self.tile_count.clone(), TILE_NEXT_BOUNDARY),
Expr::u32(STMT_NO_BOUNDARY),
),
Node::loop_for(
"compose_tile",
Expr::u32(0),
self.tile_count.clone(),
vec![
Node::store(
C11_STATEMENT_BOUNDS_SCRATCH,
self.tile_slot(Expr::var("compose_tile"), TILE_PAREN_ENTRY),
Expr::var("entry_paren"),
),
Node::store(
C11_STATEMENT_BOUNDS_SCRATCH,
self.tile_slot(Expr::var("compose_tile"), TILE_BRACKET_ENTRY),
Expr::var("entry_bracket"),
),
Node::let_bind(
"compose_pc",
Expr::load(
C11_STATEMENT_BOUNDS_SCRATCH,
self.tile_slot(Expr::var("compose_tile"), TILE_PAREN_CLOSERS),
),
),
Node::let_bind(
"compose_po",
Expr::load(
C11_STATEMENT_BOUNDS_SCRATCH,
self.tile_slot(Expr::var("compose_tile"), TILE_PAREN_OPENERS),
),
),
Node::let_bind(
"compose_bc",
Expr::load(
C11_STATEMENT_BOUNDS_SCRATCH,
self.tile_slot(Expr::var("compose_tile"), TILE_BRACKET_CLOSERS),
),
),
Node::let_bind(
"compose_bo",
Expr::load(
C11_STATEMENT_BOUNDS_SCRATCH,
self.tile_slot(Expr::var("compose_tile"), TILE_BRACKET_OPENERS),
),
),
Node::assign(
"entry_paren",
Expr::add(
Expr::sub(
Expr::max(Expr::var("entry_paren"), Expr::var("compose_pc")),
Expr::var("compose_pc"),
),
Expr::var("compose_po"),
),
),
Node::assign(
"entry_bracket",
Expr::add(
Expr::sub(
Expr::max(Expr::var("entry_bracket"), Expr::var("compose_bc")),
Expr::var("compose_bc"),
),
Expr::var("compose_bo"),
),
),
],
),
]
}
fn pass_mark_boundaries(&self) -> Vec<Node> {
vec![
Node::let_bind("mark_lo", self.own_tile_lo()),
Node::let_bind("mark_hi", self.own_tile_hi()),
Node::let_bind(
"paren_depth",
Expr::load(
C11_STATEMENT_BOUNDS_SCRATCH,
self.own_slot(TILE_PAREN_ENTRY),
),
),
Node::let_bind(
"bracket_depth",
Expr::load(
C11_STATEMENT_BOUNDS_SCRATCH,
self.own_slot(TILE_BRACKET_ENTRY),
),
),
Node::let_bind("first_boundary", Expr::u32(STMT_NO_BOUNDARY)),
Node::loop_for(
"mark_tok",
Expr::var("mark_lo"),
Expr::var("mark_hi"),
vec![
Node::let_bind("token", Expr::load(self.tok_types, Expr::var("mark_tok"))),
Node::if_then(
Expr::eq(Expr::var("token"), Expr::u32(TOK_LPAREN)),
vec![Node::assign(
"paren_depth",
Expr::add(Expr::var("paren_depth"), Expr::u32(1)),
)],
),
Node::if_then(
Expr::eq(Expr::var("token"), Expr::u32(TOK_RPAREN)),
vec![Node::assign(
"paren_depth",
Self::clamped_dec("paren_depth"),
)],
),
Node::if_then(
Expr::eq(Expr::var("token"), Expr::u32(TOK_LBRACKET)),
vec![Node::assign(
"bracket_depth",
Expr::add(Expr::var("bracket_depth"), Expr::u32(1)),
)],
),
Node::if_then(
Expr::eq(Expr::var("token"), Expr::u32(TOK_RBRACKET)),
vec![Node::assign(
"bracket_depth",
Self::clamped_dec("bracket_depth"),
)],
),
Node::let_bind(
"at_top_level_expr",
Expr::and(
Expr::eq(Expr::var("paren_depth"), Expr::u32(0)),
Expr::eq(Expr::var("bracket_depth"), Expr::u32(0)),
),
),
Node::let_bind(
"is_brace_boundary",
Expr::and(
Expr::var("at_top_level_expr"),
Expr::or(
Expr::eq(Expr::var("token"), Expr::u32(TOK_LBRACE)),
Expr::eq(Expr::var("token"), Expr::u32(TOK_RBRACE)),
),
),
),
Node::let_bind(
"is_statement_boundary",
Expr::or(
Expr::eq(Expr::var("token"), Expr::u32(TOK_SEMICOLON)),
Expr::var("is_brace_boundary"),
),
),
Node::store(
C11_STATEMENT_BOUNDS_SCRATCH,
Expr::var("mark_tok"),
Expr::select(
Expr::var("is_statement_boundary"),
Expr::var("mark_tok"),
Expr::u32(STMT_NO_BOUNDARY),
),
),
Node::if_then(
Expr::and(
Expr::var("is_statement_boundary"),
Expr::eq(Expr::var("first_boundary"), Expr::u32(STMT_NO_BOUNDARY)),
),
vec![Node::assign("first_boundary", Expr::var("mark_tok"))],
),
],
),
Node::store(
C11_STATEMENT_BOUNDS_SCRATCH,
self.own_slot(TILE_FIRST_BOUNDARY),
Expr::var("first_boundary"),
),
]
}
fn pass_resolve_tiles(&self) -> Vec<Node> {
vec![
Node::let_bind("tile_running", Expr::u32(STMT_NO_BOUNDARY)),
Node::loop_for(
"resolve_step",
Expr::u32(0),
self.tile_count.clone(),
vec![
Node::let_bind(
"resolve_tile",
Expr::sub(
Expr::sub(self.tile_count.clone(), Expr::u32(1)),
Expr::var("resolve_step"),
),
),
Node::let_bind(
"resolve_first",
Expr::load(
C11_STATEMENT_BOUNDS_SCRATCH,
self.tile_slot(Expr::var("resolve_tile"), TILE_FIRST_BOUNDARY),
),
),
Node::if_then(
Expr::ne(Expr::var("resolve_first"), Expr::u32(STMT_NO_BOUNDARY)),
vec![Node::assign("tile_running", Expr::var("resolve_first"))],
),
Node::store(
C11_STATEMENT_BOUNDS_SCRATCH,
self.tile_slot(Expr::var("resolve_tile"), TILE_NEXT_BOUNDARY),
Expr::var("tile_running"),
),
],
),
]
}
fn pass_emit_spans(&self, out_statements: &str, out_counts: &str) -> Vec<Node> {
vec![
Node::let_bind("emit_lo", self.own_tile_lo()),
Node::let_bind("emit_hi", self.own_tile_hi()),
Node::let_bind(
"emit_running",
Expr::load(
C11_STATEMENT_BOUNDS_SCRATCH,
self.tile_slot(Expr::add(self.t.clone(), Expr::u32(1)), TILE_NEXT_BOUNDARY),
),
),
Node::loop_for(
"emit_step",
Expr::var("emit_lo"),
Expr::var("emit_hi"),
vec![
Node::let_bind(
"emit_pos",
Expr::sub(
Expr::sub(Expr::var("emit_hi"), Expr::u32(1)),
Expr::sub(Expr::var("emit_step"), Expr::var("emit_lo")),
),
),
Node::let_bind(
"emit_mark",
Expr::load(C11_STATEMENT_BOUNDS_SCRATCH, Expr::var("emit_pos")),
),
Node::if_then(
Expr::ne(Expr::var("emit_mark"), Expr::u32(STMT_NO_BOUNDARY)),
vec![Node::assign("emit_running", Expr::var("emit_mark"))],
),
Node::let_bind(
"stmt_bound_end",
Expr::select(
Expr::eq(Expr::var("emit_running"), Expr::u32(STMT_NO_BOUNDARY)),
Expr::var("emit_pos"),
Expr::add(Expr::var("emit_running"), Expr::u32(1)),
),
),
Node::let_bind(
"stmt_idx",
Expr::atomic_add(out_counts, Expr::u32(0), Expr::u32(2)),
),
Node::store(out_statements, Expr::var("stmt_idx"), Expr::var("emit_pos")),
Node::store(
out_statements,
Expr::add(Expr::var("stmt_idx"), Expr::u32(1)),
Expr::var("stmt_bound_end"),
),
],
),
]
}
}
#[must_use]
pub fn c11_statement_bounds(
tok_types: &str,
num_tokens: Expr,
out_statements: &str,
out_counts: &str,
) -> Program {
let t = Expr::InvocationId { axis: 0 };
let tok_count = match &num_tokens {
Expr::LitU32(0) => 1,
Expr::LitU32(n) => *n,
other => panic!(
"c11_statement_bounds requires a literal token-window count for build-time output \
buffer sizing, got a non-literal expression {other:?}. Fix: pass Expr::u32(N)."
),
};
let active = Expr::min(Expr::buf_len(tok_types), Expr::u32(tok_count));
let tile_count = Expr::div(
Expr::add(active.clone(), Expr::u32(STMT_TILE_TOKENS - 1)),
Expr::u32(STMT_TILE_TOKENS),
);
let ctx = BoundsCtx {
tok_types,
tile_base: tok_count,
t: t.clone(),
active,
tile_count: tile_count.clone(),
};
let owns_tile = Expr::lt(t.clone(), tile_count.clone());
let is_lead_lane = Expr::eq(t.clone(), Expr::u32(0));
let has_tiles = Expr::gt(tile_count, Expr::u32(0));
Program::wrapped(
vec![
BufferDecl::storage(tok_types, 0, BufferAccess::ReadOnly, DataType::U32),
BufferDecl::storage(out_statements, 1, BufferAccess::ReadWrite, DataType::U32)
.with_count(tok_count.saturating_mul(2)),
BufferDecl::storage(out_counts, 2, BufferAccess::ReadWrite, DataType::U32)
.with_count(1),
BufferDecl::storage(
C11_STATEMENT_BOUNDS_SCRATCH,
3,
BufferAccess::ReadWrite,
DataType::U32,
)
.with_count(c11_statement_bounds_scratch_words(tok_count)),
],
[256, 1, 1],
vec![wrap_anonymous(
"vyre-libs::parsing::c11_statement_bounds",
vec![
Node::if_then(
is_lead_lane.clone(),
vec![Node::store(out_counts, Expr::u32(0), Expr::u32(0))],
),
Node::Barrier {
ordering: MemoryOrdering::GridSync,
},
Node::if_then(owns_tile.clone(), ctx.pass_tile_reduce()),
Node::Barrier {
ordering: MemoryOrdering::GridSync,
},
Node::if_then(is_lead_lane.clone(), ctx.pass_tile_compose()),
Node::Barrier {
ordering: MemoryOrdering::GridSync,
},
Node::if_then(owns_tile.clone(), ctx.pass_mark_boundaries()),
Node::Barrier {
ordering: MemoryOrdering::GridSync,
},
Node::if_then(Expr::and(is_lead_lane, has_tiles), ctx.pass_resolve_tiles()),
Node::Barrier {
ordering: MemoryOrdering::GridSync,
},
child_phase(
"vyre-libs::parsing::c11_statement_bounds",
vyre_primitives::bitset::select::OP_ID,
vec![Node::if_then(
owns_tile,
ctx.pass_emit_spans(out_statements, out_counts),
)],
),
],
)],
)
.with_entry_op_id("vyre-libs::parsing::c11_statement_bounds")
.with_non_composable_with_self(true)
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn statement_bounds_sizes_outputs_to_full_literal_window_without_fixed_clamp() {
let token_window = crate::parsing::c::pipeline::stages::C11_AST_MAX_TOK_SCAN + 1;
let program = c11_statement_bounds(
"tok_types",
Expr::u32(token_window),
"out_statements",
"out_counts",
);
let out_statements = program
.buffers
.iter()
.find(|buffer| buffer.name() == "out_statements")
.expect("Fix: out_statements buffer must exist");
assert_eq!(out_statements.count, token_window.saturating_mul(2));
}
#[test]
fn statement_bounds_rejects_non_literal_token_count_for_buffer_sizing() {
let panic = std::panic::catch_unwind(|| {
let _ = c11_statement_bounds(
"tok_types",
Expr::var("dynamic_tokens"),
"out_statements",
"out_counts",
);
})
.expect_err("non-literal statement-bound count must fail");
let message = panic
.downcast_ref::<String>()
.map(String::as_str)
.or_else(|| panic.downcast_ref::<&'static str>().copied())
.unwrap_or("<non-string panic>");
assert!(
message.contains("requires a literal token-window count"),
"{message}"
);
}
#[test]
fn statement_bounds_initializes_atomic_count_inside_kernel() {
let source = include_str!("structure_statement.rs");
assert!(
source.contains("Node::store(out_counts, Expr::u32(0), Expr::u32(0))"),
"Fix: statement bounds must zero out_counts in-kernel before atomic_add."
);
assert!(
source.contains("MemoryOrdering::GridSync"),
"Fix: statement bounds must synchronize after zeroing out_counts before worker lanes \
append records. The fence must be GridSync, not a workgroup-scope ordering: the \
append lanes span the whole grid, so a workgroup-only fence lets another workgroup \
read out_counts before lane 0's zeroing is visible."
);
}
#[test]
fn scratch_words_covers_marks_tiles_and_the_successor_slot() {
assert_eq!(
c11_statement_bounds_scratch_words(256),
256 + STMT_TILE_WORDS + STMT_TILE_WORDS
);
assert_eq!(
c11_statement_bounds_scratch_words(257),
257 + 2 * STMT_TILE_WORDS + STMT_TILE_WORDS
);
assert_eq!(
c11_statement_bounds_scratch_words(0),
c11_statement_bounds_scratch_words(1)
);
}
#[test]
fn launch_geometry_stays_within_cooperative_residency_at_the_pipeline_cap() {
let geometry = |n: u32| -> (u32, [u32; 3]) {
let program =
c11_statement_bounds("tok_types", Expr::u32(n), "out_statements", "out_counts");
let plan = vyre_driver::binding::BindingPlan::build(&program)
.expect("binding plan must build from declared buffer counts");
let count = vyre_driver::dispatch_element_count_for_program(&program, &plan.bindings);
let grid = vyre_driver::infer_dispatch_grid_for_count(count, program.workgroup_size())
.expect("grid must infer for a 1D workgroup");
(count, grid)
};
assert_eq!(
geometry(256),
(512, [2, 1, 1]),
"Fix: element count must be 2 * num_tokens (out_statements), so 256 tokens is 512 lanes."
);
let cap = crate::parsing::c::pipeline::stages::C11_AST_MAX_TOK_SCAN;
assert_eq!(
geometry(cap),
(cap * 2, [cap / 128, 1, 1]),
"Fix: at the pipeline scan cap the grid must stay at ceil(2n/256) blocks."
);
assert_eq!(
geometry(cap).1[0],
512,
"Fix: the pipeline cap must launch 512 blocks, half the 1020 block cooperative \
residency limit on a 170 SM Blackwell part. More than this spends the headroom."
);
assert_eq!(
geometry(130_560).1[0],
1_020,
"Fix: 130560 tokens must land exactly on 1020 blocks, the documented point where the \
launch route changes from one cooperative launch to a kernel split."
);
assert_eq!(
geometry(130_688).1[0],
1_021,
"Fix: one tile past the transition must exceed 1020 blocks, which is what moves the \
dispatch onto the split route. The split is correct at any width, so this is a cost \
boundary and not a correctness cliff."
);
}
#[test]
fn autotuner_does_not_widen_the_declared_workgroup() {
let limits = vyre_driver::validation::LaunchGeometryLimits {
backend: "cuda",
max_threads_per_block: 1024,
max_block_dim: [1024, 1024, 64],
max_grid_dim: [u32::MAX, u32::MAX, u32::MAX],
max_threads_per_sm: 0,
};
let config = vyre_driver::DispatchConfig::default();
for n in [65_536_u32, 87_040, 130_560] {
let program =
c11_statement_bounds("tok_types", Expr::u32(n), "out_statements", "out_counts");
assert!(
program.non_composable_with_self,
"Fix: this kernel must stay non-composable-with-self; that flag is the only thing \
keeping the autotuner from widening it to 1024 and moving the split crossing \
point from 130560 tokens down to 87040."
);
let plan = vyre_driver::binding::BindingPlan::build(&program)
.expect("binding plan must build from declared buffer counts");
let count = vyre_driver::dispatch_element_count_for_program(&program, &plan.bindings);
for mode in [
vyre_driver::tuner::Mode::production_default(),
vyre_driver::tuner::Mode::OffUseDefault,
] {
let effective = vyre_driver::launch::resolve_launch_workgroup_for_mode(
&program, &config, limits, count, mode,
);
assert_eq!(
effective,
[256, 1, 1],
"Fix: effective workgroup must stay 256 under {mode:?} at {n} tokens. A 1024 \
wide launch would cut the cooperative lane ceiling from 261120 to 174080."
);
}
}
}
#[test]
fn scratch_buffer_is_declared_after_both_reported_outputs() {
let program =
c11_statement_bounds("tok_types", Expr::u32(64), "out_statements", "out_counts");
let names: Vec<&str> = program.buffers.iter().map(|b| b.name()).collect();
assert_eq!(
names,
vec![
"tok_types",
"out_statements",
"out_counts",
C11_STATEMENT_BOUNDS_SCRATCH
]
);
}
}