use std::process::Command;
use std::sync::Arc;
use cutile::compile_api::{CheckPlacementCounts, KernelCompiler};
use cutile::prelude::*;
use cutile_compiler::ast::Module;
use cutile_compiler::compiler::utils::CompileOptions;
use cutile_compiler::error::JITError;
pub fn compile<F: Fn() -> Module>(
module_ast_fn: F,
module_name: &str,
function_name: &str,
generics: &[&str],
strides: &[(&str, &[i32])],
) -> Result<(String, CheckPlacementCounts), JITError> {
KernelCompiler::new(module_ast_fn, module_name, function_name)
.target("sm_120")
.generics(generics.iter().map(|g| g.to_string()).collect())
.strides(strides)
.options(CompileOptions::default())
.compile()
.map(|artifacts| (artifacts.ir_text(), artifacts.check_counts()))
}
pub fn upload<T: DType>(values: Vec<T>) -> Arc<Tensor<T>> {
Arc::new(
api::copy_host_vec_to_device(&Arc::new(values))
.sync()
.expect("upload"),
)
}
pub fn host<T: DType>(tensor: &Tensor<T>) -> Vec<T> {
tensor.dup().to_host_vec().sync().expect("to_host")
}
#[derive(Debug)]
pub enum Outcome {
Ok,
Stop(String),
}
pub fn run_in_subprocess(runner_test: &str, env_var: &str, case: &str) -> Outcome {
let exe = std::env::current_exe().expect("current_exe");
let out = Command::new(exe)
.args(["--exact", runner_test, "--ignored", "--nocapture"])
.env(env_var, case)
.output()
.expect("spawn case subprocess");
let stdout = String::from_utf8_lossy(&out.stdout);
let stderr = String::from_utf8_lossy(&out.stderr);
let Some(line) = stdout.lines().find(|l| l.starts_with("OUTCOME:")) else {
if out.status.code() != Some(0) && stderr.contains("panicked") {
return Outcome::Stop("aborted during cleanup".to_string());
}
panic!("case {case} produced no OUTCOME line.\nstdout:\n{stdout}\nstderr:\n{stderr}");
};
if line == "OUTCOME:OK" {
Outcome::Ok
} else if let Some(msg) = line.strip_prefix("OUTCOME:STOP:") {
Outcome::Stop(msg.to_string())
} else {
panic!("case {case} infrastructure failure: {line}")
}
}
pub fn report_outcome(result: std::thread::Result<Result<(), String>>) {
match result {
Ok(Ok(())) => println!("OUTCOME:OK"),
Ok(Err(msg)) => println!("OUTCOME:STOP:{}", msg.replace('\n', " | ")),
Err(_) => println!("OUTCOME:STOP:panicked"),
}
}