use cudarc::driver::sys;
pub struct KernelNode {
pub node: sys::CUgraphNode,
pub params: sys::CUDA_KERNEL_NODE_PARAMS,
pub name: String,
}
unsafe impl Send for KernelNode {}
fn cu_try(r: sys::CUresult, what: &str) -> Result<(), Box<dyn std::error::Error>> {
if r == sys::CUresult::CUDA_SUCCESS { Ok(()) } else { Err(format!("{what}: {r:?}").into()) }
}
pub fn kernel_nodes(graph: &cudarc::driver::CudaGraph)
-> Result<Vec<KernelNode>, Box<dyn std::error::Error>>
{
let g = graph.cu_graph();
let mut n: usize = 0;
unsafe { cu_try(sys::cuGraphGetNodes(g, std::ptr::null_mut(), &mut n), "cuGraphGetNodes(count)")?; }
let mut nodes: Vec<sys::CUgraphNode> = vec![std::ptr::null_mut(); n];
unsafe { cu_try(sys::cuGraphGetNodes(g, nodes.as_mut_ptr(), &mut n), "cuGraphGetNodes")?; }
nodes.truncate(n);
let mut out = Vec::with_capacity(n);
for node in nodes {
let mut ty = sys::CUgraphNodeType::CU_GRAPH_NODE_TYPE_EMPTY;
unsafe { cu_try(sys::cuGraphNodeGetType(node, &mut ty), "cuGraphNodeGetType")?; }
if ty != sys::CUgraphNodeType::CU_GRAPH_NODE_TYPE_KERNEL { continue; }
let mut params: sys::CUDA_KERNEL_NODE_PARAMS = unsafe { std::mem::zeroed() };
unsafe { cu_try(sys::cuGraphKernelNodeGetParams_v2(node, &mut params),
"cuGraphKernelNodeGetParams_v2")?; }
let mut cname: *const std::ffi::c_char = std::ptr::null();
let name = unsafe {
if sys::cuFuncGetName(&mut cname, params.func) == sys::CUresult::CUDA_SUCCESS
&& !cname.is_null() {
std::ffi::CStr::from_ptr(cname).to_string_lossy().into_owned()
} else { String::from("<unknown>") }
};
out.push(KernelNode { node, params, name });
}
Ok(out)
}
pub fn node_census(graph: &cudarc::driver::CudaGraph)
-> Result<std::collections::BTreeMap<String, usize>, Box<dyn std::error::Error>>
{
let g = graph.cu_graph();
let mut n: usize = 0;
unsafe { cu_try(sys::cuGraphGetNodes(g, std::ptr::null_mut(), &mut n), "cuGraphGetNodes(count)")?; }
let mut nodes: Vec<sys::CUgraphNode> = vec![std::ptr::null_mut(); n];
unsafe { cu_try(sys::cuGraphGetNodes(g, nodes.as_mut_ptr(), &mut n), "cuGraphGetNodes")?; }
nodes.truncate(n);
let mut out: std::collections::BTreeMap<String, usize> = Default::default();
for node in nodes {
let mut ty = sys::CUgraphNodeType::CU_GRAPH_NODE_TYPE_EMPTY;
unsafe { cu_try(sys::cuGraphNodeGetType(node, &mut ty), "cuGraphNodeGetType")?; }
*out.entry(format!("{ty:?}")).or_insert(0) += 1;
}
Ok(out)
}
pub fn set_exec_params(graph: &cudarc::driver::CudaGraph, node: sys::CUgraphNode,
params: &sys::CUDA_KERNEL_NODE_PARAMS)
-> Result<(), Box<dyn std::error::Error>>
{
unsafe { cu_try(sys::cuGraphExecKernelNodeSetParams_v2(graph.cu_graph_exec(), node, params),
"cuGraphExecKernelNodeSetParams_v2") }
}
pub unsafe fn write_i32_arg(params: &sys::CUDA_KERNEL_NODE_PARAMS, idx: usize, val: i32) {
unsafe {
let slot = *params.kernelParams.add(idx) as *mut i32;
*slot = val;
}
}
pub unsafe fn read_i32_arg(params: &sys::CUDA_KERNEL_NODE_PARAMS, idx: usize) -> i32 {
unsafe { *(*params.kernelParams.add(idx) as *const i32) }
}
pub unsafe fn read_ptr_arg(params: &sys::CUDA_KERNEL_NODE_PARAMS, idx: usize) -> u64 {
unsafe { *(*params.kernelParams.add(idx) as *const u64) }
}
pub struct FaMain {
node: sys::CUgraphNode,
params: sys::CUDA_KERNEL_NODE_PARAMS,
nkv: u32,
bucket_splits: u32,
self_split_keys: Option<i32>,
combine: Option<(sys::CUgraphNode, sys::CUDA_KERNEL_NODE_PARAMS, usize /*nsp idx*/)>,
cur: u32,
}
unsafe impl Send for FaMain {}
const VEC_NSP_IDX: usize = 11; const VEC_PARTO_IDX: usize = 3;
const SCALAR_NSP_IDX: usize = 12; const SCALAR_SKI_IDX: usize = 13;
const COMBINE_NSP_IDX: usize = 6; const COMBINE_Q8_NSP_IDX: usize = 7; const COMBINE_PARTO_IDX: usize = 0;
pub fn fa_plan(graph: &cudarc::driver::CudaGraph)
-> Result<Vec<FaMain>, Box<dyn std::error::Error>>
{
let nodes = kernel_nodes(graph)?;
let mut mains: Vec<(usize, FaMain)> = Vec::new();
let mut combines: Vec<Option<(usize, u64, sys::CUgraphNode, sys::CUDA_KERNEL_NODE_PARAMS,
usize)>> = Vec::new();
for (i, n) in nodes.iter().enumerate() {
match n.name.as_str() {
"fa_decode_vec_q_v4_dc" | "fa_decode_vec_q_v4_deep_dc"
| "fa_decode_vec_q_v3_dc" | "fa_decode_vec_q_v2_dc"
| "fa_decode_vec_q_dc" | "fa_decode_vec_q_dpl16_dc" => {
mains.push((i, FaMain {
nkv: n.params.gridDimX, bucket_splits: n.params.gridDimY,
self_split_keys: None, combine: None, cur: n.params.gridDimY,
node: n.node, params: n.params,
}));
}
"fa_decode_f32" => {
let ski = unsafe { read_i32_arg(&n.params, SCALAR_SKI_IDX) };
mains.push((i, FaMain {
nkv: n.params.gridDimX, bucket_splits: n.params.gridDimY,
self_split_keys: Some(ski), combine: None, cur: n.params.gridDimY,
node: n.node, params: n.params,
}));
}
"fa_decode_combine_f32" | "fa_decode_combine_q8_1" => {
let po = unsafe { read_ptr_arg(&n.params, COMBINE_PARTO_IDX) };
let nsp_idx = if n.name == "fa_decode_combine_q8_1" { COMBINE_Q8_NSP_IDX }
else { COMBINE_NSP_IDX };
combines.push(Some((i, po, n.node, n.params, nsp_idx)));
}
_ => {}
}
}
let mut out = Vec::with_capacity(mains.len());
for (mi, mut m) in mains {
let po = unsafe { read_ptr_arg(&m.params, VEC_PARTO_IDX) };
let slot = combines.iter_mut()
.filter(|c| c.as_ref().is_some_and(|(ci, cpo, ..)| *ci > mi && *cpo == po))
.min_by_key(|c| c.as_ref().unwrap().0);
match slot {
Some(c) => { let (_, _, cn, cp, ci) = c.take().unwrap();
m.combine = Some((cn, cp, ci)); }
None => return Err("fa_plan: fa main has no partO-paired combine node".into()),
}
out.push(m);
}
Ok(out)
}
pub fn fa_apply(graph: &cudarc::driver::CudaGraph, plan: &mut [FaMain], t_kv: usize,
split_keys: impl Fn(usize, usize) -> usize)
-> Result<(), Box<dyn std::error::Error>>
{
for m in plan.iter_mut() {
let ns = match m.self_split_keys {
Some(ski) => (t_kv + ski as usize - 1) / (ski as usize).max(1),
None => { let sp = split_keys(t_kv, m.nkv as usize).max(1);
(t_kv + sp - 1) / sp }
}.max(1) as u32;
let ns = ns.min(m.bucket_splits);
if ns == m.cur { continue; }
m.params.gridDimY = ns;
let nsp_idx = if m.self_split_keys.is_some() { SCALAR_NSP_IDX } else { VEC_NSP_IDX };
unsafe { write_i32_arg(&m.params, nsp_idx, ns as i32); }
set_exec_params(graph, m.node, &m.params)?;
if let Some((cn, cp, ci)) = &m.combine {
unsafe { write_i32_arg(cp, *ci, ns as i32); }
set_exec_params(graph, *cn, cp)?;
}
m.cur = ns;
}
Ok(())
}