mod autotune;
mod emit;
mod ir;
mod model;
mod toolchain;
use anyhow::{Context, Result};
use ir::{plan_chunks, UnhandledOperand};
use model::{ExpressionsInfo, StarkInfo};
use rayon::prelude::*;
use std::path::{Path, PathBuf};
use toolchain::Toolchain;
const DEFAULT_CAP: usize = 40000; const DEFAULT_CHUNK: usize = 512; const SLOTS_CAP: u64 = 1000;
#[derive(Debug, Clone)]
pub struct GenConfig {
pub cap: usize,
pub chunk: Option<usize>,
pub archspec: String,
pub stark_src: Option<PathBuf>,
pub keep_dir: Option<PathBuf>,
pub dry_run: bool,
}
impl Default for GenConfig {
fn default() -> Self {
GenConfig {
cap: DEFAULT_CAP,
chunk: None,
archspec: "auto".into(),
stark_src: None,
keep_dir: None,
dry_run: false,
}
}
}
#[derive(Debug, Clone)]
pub struct GeneratedAir {
pub name: String,
pub base: String,
pub sym: String,
pub nbits: u64,
pub cexp: i64,
pub n_ops: usize,
pub slots: u64,
}
#[derive(Debug, Default)]
pub struct GenSummary {
pub generated: Vec<GeneratedAir>,
pub skipped: Vec<(String, String)>,
pub placed: usize,
pub max_scratch_bytes: u64,
}
struct Candidate {
stark_info: StarkInfo,
expr_info: ExpressionsInfo,
sym: String,
nbits: u64,
cexp: i64,
name: String,
n_ops: usize,
base: String,
}
struct Placement {
name: String,
base: String,
sym: String,
air_dir: PathBuf,
}
fn find_starkinfos(root: &Path) -> Vec<PathBuf> {
let mut out = Vec::new();
let mut stack = vec![root.to_path_buf()];
while let Some(dir) = stack.pop() {
let Ok(entries) = std::fs::read_dir(&dir) else { continue };
for e in entries.flatten() {
let p = e.path();
if p.is_dir() {
stack.push(p);
} else if p.file_name().and_then(|s| s.to_str()).is_some_and(|s| s.ends_with(".starkinfo.json")) {
out.push(p);
}
}
}
out.sort();
out
}
fn proof_phase(air_dir: &Path) -> String {
let comp = air_dir.file_name().and_then(|s| s.to_str()).unwrap_or("");
match comp {
"air" => "basic",
"compressor" => "compressor",
"recursive1" | "recursive2" => "recursive",
c if c.starts_with("vadcop_final") => "final",
other => other, }
.to_string()
}
fn make_sym(si: &StarkInfo, phase: &str) -> String {
format!("{phase}_a{}_{}_b{}_e{}", si.airgroup_id, si.air_id, si.stark_struct.n_bits, si.c_exp_id)
}
fn load_candidate(
stark_info_path: &Path,
root: &Path,
cap: usize,
) -> Result<Option<Candidate>, Option<(String, String)>> {
let air_dir = stark_info_path.parent().unwrap();
let fname = stark_info_path.file_name().unwrap().to_string_lossy();
let base = fname.strip_suffix(".starkinfo.json").unwrap().to_string();
let expr_info_path = air_dir.join(format!("{base}.expressionsinfo.json"));
if !expr_info_path.exists() {
return Err(None);
}
let name = air_dir.strip_prefix(root).unwrap_or(air_dir).to_string_lossy().to_string();
let stark_info: StarkInfo = match std::fs::read(stark_info_path).ok().and_then(|b| serde_json::from_slice(&b).ok())
{
Some(si) => si,
None => return Err(None), };
let expr_info: ExpressionsInfo =
match std::fs::read(&expr_info_path).ok().and_then(|b| serde_json::from_slice(&b).ok()) {
Some(ei) => ei,
None => return Err(None),
};
let cexp = stark_info.c_exp_id;
let Some(code) = expr_info.expressions_code.iter().find(|e| e.exp_id == cexp) else {
return Err(None); };
let n_ops = code.code.len();
let nbits = stark_info.stark_struct.n_bits;
if n_ops > cap {
return Err(Some((name, format!("{n_ops} ops > CAP"))));
}
let sym = make_sym(&stark_info, &proof_phase(air_dir));
Ok(Some(Candidate { stark_info, expr_info, sym, nbits, cexp, name, n_ops, base }))
}
pub fn generate_all(proving_key: &Path, cfg: &GenConfig) -> Result<GenSummary> {
if cfg.dry_run && cfg.keep_dir.is_none() {
anyhow::bail!("dry_run requires keep_dir (nowhere to write the .cu otherwise)");
}
let work = WorkDir::new(cfg.keep_dir.clone())?;
std::fs::write(work.path().join("gen_common.cuh"), emit::COMMON_CUH)?;
let tc = Toolchain::new(cfg.stark_src.clone(), &cfg.archspec, work.path())?;
eprintln!("[exps-codegen] generating kernels for {} (archs: {})", proving_key.display(), tc.arch_summary());
let mut candidates: Vec<Candidate> = Vec::new();
let mut placements: Vec<Placement> = Vec::new();
let mut skipped: Vec<(String, String)> = Vec::new();
let mut seen: std::collections::HashSet<String> = std::collections::HashSet::new();
for si_path in find_starkinfos(proving_key) {
match load_candidate(&si_path, proving_key, cfg.cap) {
Ok(Some(c)) => {
placements.push(Placement {
name: c.name.clone(),
base: c.base.clone(),
sym: c.sym.clone(),
air_dir: si_path.parent().unwrap().to_path_buf(),
});
if seen.insert(c.sym.clone()) {
candidates.push(c);
}
}
Ok(None) => {}
Err(Some(skip)) => skipped.push(skip),
Err(None) => {}
}
}
let mut summary = run_pipeline(&tc, work.path(), &candidates, &placements, cfg)?;
summary.skipped.extend(skipped);
summary.skipped.sort();
print_summary(&summary, cfg);
Ok(summary)
}
pub fn generate_air(air_dir: &Path, cfg: &GenConfig) -> Result<PathBuf> {
let si_path = find_starkinfos(air_dir)
.into_iter()
.find(|p| p.parent() == Some(air_dir))
.with_context(|| format!("no *.starkinfo.json in {}", air_dir.display()))?;
let work = WorkDir::new(cfg.keep_dir.clone())?;
std::fs::write(work.path().join("gen_common.cuh"), emit::COMMON_CUH)?;
let tc = Toolchain::new(cfg.stark_src.clone(), &cfg.archspec, work.path())?;
let root = air_dir;
let candidate = match load_candidate(&si_path, root, cfg.cap) {
Ok(Some(c)) => c,
Ok(None) => anyhow::bail!("{} is not a codegen target", air_dir.display()),
Err(Some((_, why))) => anyhow::bail!("skipped: {why}"),
Err(None) => anyhow::bail!("{} missing/invalid expressionsinfo", air_dir.display()),
};
let placement = Placement {
name: candidate.name.clone(),
base: candidate.base.clone(),
sym: candidate.sym.clone(),
air_dir: air_dir.to_path_buf(),
};
let dest = air_dir.join(format!("{}.exps.so", candidate.base));
let summary =
run_pipeline(&tc, work.path(), std::slice::from_ref(&candidate), std::slice::from_ref(&placement), cfg)?;
if summary.placed == 1 {
Ok(dest)
} else {
let why = summary.skipped.first().map(|(_, w)| w.clone()).unwrap_or_else(|| "unknown".into());
anyhow::bail!("{}: not generated ({why})", candidate.name)
}
}
fn run_pipeline(
tc: &Toolchain,
work: &Path,
candidates: &[Candidate],
placements: &[Placement],
cfg: &GenConfig,
) -> Result<GenSummary> {
let autotune = cfg.chunk.is_none();
let built: Vec<(&Candidate, std::result::Result<ir::Ir, String>)> = candidates
.iter()
.map(|c| {
let r = match ir::build_ir(&c.stark_info, &c.expr_info) {
Ok(ir) => Ok(ir),
Err(e) if e.downcast_ref::<UnhandledOperand>().is_some() => Err("unhandled operand".to_string()),
Err(e) => Err(format!("build_ir error: {e}")),
};
(c, r)
})
.collect();
let chunk_map: std::collections::HashMap<String, Option<usize>> = if autotune {
built
.par_iter()
.filter_map(|(c, r)| {
r.as_ref()
.ok()
.map(|ir| autotune::tune_chunk(tc, ir, &c.sym, c.n_ops, work).map(|ck| (c.sym.clone(), ck)))
})
.collect::<Result<std::collections::HashMap<_, _>>>()?
} else {
Default::default()
};
let mut slots_by_sym: std::collections::HashMap<String, u64> = std::collections::HashMap::new();
let mut exprs_by_sym: std::collections::HashMap<String, usize> = std::collections::HashMap::new();
let mut generated: Vec<GeneratedAir> = Vec::new();
let mut skipped: Vec<(String, String)> = Vec::new();
let mut max_scratch: u64 = 0;
for (c, r) in &built {
let ir = match r {
Ok(ir) => ir,
Err(why) => {
skipped.push((c.name.clone(), why.clone()));
continue;
}
};
let chunk = if autotune {
match chunk_map.get(&c.sym).copied().flatten() {
Some(ck) => ck,
None => {
skipped.push((c.name.clone(), format!("{} ops: still spills at CHUNK_MIN", c.n_ops)));
continue;
}
}
} else {
cfg.chunk.unwrap_or(DEFAULT_CHUNK)
};
let plan = plan_chunks(ir, chunk, &c.sym)?;
if plan.total_slots > SLOTS_CAP {
skipped.push((c.name.clone(), format!("slots {} > SLOTS_CAP (wide cut)", plan.total_slots)));
continue;
}
for (fname, text) in emit::emit_air(ir, &plan, &c.sym) {
std::fs::write(work.join(&fname), text)?;
}
{
const EXPR_CAP: usize = 512;
let mut items: Vec<(i64, ir::Ir, u64)> = Vec::new();
for ec in &c.expr_info.expressions_code {
if ec.exp_id == c.cexp || ec.code.is_empty() || ec.code.len() > EXPR_CAP {
continue;
}
let Ok(eir) = ir::build_ir_expr(&c.stark_info, &c.expr_info, ec.exp_id, false) else {
continue;
};
if eir.uses_zi() {
continue;
}
let Some(od) = eir.out_dim() else { continue };
if od != 1 && od != 3 {
continue;
}
items.push((ec.exp_id, eir, od));
}
if !items.is_empty() {
let n_exprs = items.len();
std::fs::write(work.join(format!("gen_{}_cexprs.cu", c.sym)), emit::emit_exprs_tu(&c.sym, &items))?;
exprs_by_sym.insert(c.sym.clone(), n_exprs);
}
}
let n_ext = 1u64 << c.stark_info.stark_struct.n_bits_ext;
max_scratch = max_scratch.max(plan.total_slots * n_ext);
slots_by_sym.insert(c.sym.clone(), plan.total_slots);
generated.push(GeneratedAir {
name: c.name.clone(),
base: c.base.clone(),
sym: c.sym.clone(),
nbits: c.nbits,
cexp: c.cexp,
n_ops: c.n_ops,
slots: plan.total_slots,
});
}
{
let mut jobs: Vec<(PathBuf, PathBuf)> = Vec::new();
for sym in slots_by_sym.keys() {
for cu in collect_artifacts(work, sym, "cu") {
let obj = cu.with_extension("o");
if !obj.exists() {
jobs.push((cu, obj));
}
}
}
let par = std::thread::available_parallelism().map(|n| n.get()).unwrap_or(8);
for batch in jobs.chunks(par) {
let errs: Vec<String> = std::thread::scope(|scope| {
let handles: Vec<_> = batch
.iter()
.map(|(cu, obj)| {
scope.spawn(move || -> Option<String> {
match tc.compile_tu(cu, obj, Some(work)) {
Ok((true, _)) => None,
Ok((false, log)) => Some(format!("nvcc failed for {}: {log}", cu.display())),
Err(e) => Some(format!("nvcc spawn failed for {}: {e}", cu.display())),
}
})
})
.collect();
handles.into_iter().filter_map(|h| h.join().ok().flatten()).collect()
});
if let Some(e) = errs.into_iter().next() {
anyhow::bail!(e);
}
}
}
let placed: Vec<&Placement> = placements.iter().filter(|p| slots_by_sym.contains_key(&p.sym)).collect();
write_gen_log(work, &placed, &slots_by_sym)?;
if !cfg.dry_run {
placed.par_iter().try_for_each(|p| -> Result<()> {
let dest = p.air_dir.join(format!("{}.exps.so", p.base));
link_one(tc, work, &p.sym, &dest)
})?;
}
Ok(GenSummary { placed: placed.len(), generated, skipped, max_scratch_bytes: max_scratch * 8 })
}
fn link_one(tc: &Toolchain, work: &Path, sym: &str, dest: &Path) -> Result<()> {
let objs = collect_artifacts(work, sym, "o");
if !objs.is_empty() {
tc.link_objs(&objs, dest)
} else {
let cus = collect_artifacts(work, sym, "cu");
tc.compile_link_cus(&cus, dest)
}
}
fn collect_artifacts(work: &Path, sym: &str, ext: &str) -> Vec<PathBuf> {
let mut out = Vec::new();
let main = work.join(format!("gen_{sym}.{ext}"));
if main.exists() {
out.push(main);
}
let prefix = format!("gen_{sym}_c");
if let Ok(entries) = std::fs::read_dir(work) {
let mut chunks: Vec<PathBuf> = entries
.flatten()
.map(|e| e.path())
.filter(|p| {
p.file_name()
.and_then(|s| s.to_str())
.is_some_and(|s| s.starts_with(&prefix) && s.ends_with(&format!(".{ext}")))
})
.collect();
chunks.sort();
out.extend(chunks);
}
out
}
fn write_gen_log(
work: &Path,
placed: &[&Placement],
slots_by_sym: &std::collections::HashMap<String, u64>,
) -> Result<()> {
let mut log = String::new();
for p in placed {
log.push_str(&format!("{}\t{}\t{}\t{}\n", p.name, p.base, p.sym, slots_by_sym[&p.sym]));
}
std::fs::write(work.join("gen.log"), log)?;
Ok(())
}
fn print_summary(s: &GenSummary, cfg: &GenConfig) {
let chunk_info = if let Some(chunk) = cfg.chunk { format!(", chunk={}", chunk) } else { String::new() };
eprintln!(
"generated {} kernels -> {} per-AIR .exps.so (CAP={}{}, max scratch {:.0}MB):",
s.generated.len(),
s.placed,
cfg.cap,
chunk_info,
s.max_scratch_bytes as f64 / 1e6
);
for g in &s.generated {
let chunked = if g.slots > 0 { format!("CHUNKED slots={}", g.slots) } else { "single".into() };
eprintln!(" {:40} {}.exps.so nBits={} cExp={} ops={} {}", g.name, g.base, g.nbits, g.cexp, g.n_ops, chunked);
}
for (name, why) in &s.skipped {
eprintln!(" SKIP {name:38} {why}");
}
}
struct WorkDir {
path: PathBuf,
temp: bool,
}
impl WorkDir {
fn new(keep_dir: Option<PathBuf>) -> Result<Self> {
match keep_dir {
Some(p) => {
std::fs::create_dir_all(&p)?;
eprintln!("[exps-codegen] keeping generated code in {}", p.display());
Ok(WorkDir { path: p, temp: false })
}
None => {
use std::sync::atomic::{AtomicU64, Ordering};
static SEQ: AtomicU64 = AtomicU64::new(0);
let seq = SEQ.fetch_add(1, Ordering::Relaxed);
let p = std::env::temp_dir().join(format!("genexps_{}_{}", std::process::id(), seq));
let _ = std::fs::remove_dir_all(&p);
std::fs::create_dir_all(&p)?;
Ok(WorkDir { path: p, temp: true })
}
}
}
fn path(&self) -> &Path {
&self.path
}
}
impl Drop for WorkDir {
fn drop(&mut self) {
if self.temp {
let _ = std::fs::remove_dir_all(&self.path);
}
}
}