use crate::Engine;
use cudarc::driver::CudaSlice;
use std::collections::BTreeMap;
use std::sync::{Mutex, OnceLock};
type Res<T> = Result<T, Box<dyn std::error::Error>>;
#[derive(Clone, Debug, PartialEq)]
pub struct SelRow {
pub dev: usize,
pub layer: u16,
pub sel: Vec<i32>,
pub w: Vec<f32>,
}
#[derive(Default)]
struct Ledger {
slots: BTreeMap<(usize, u16), (CudaSlice<i32>, CudaSlice<f32>)>,
host: Vec<SelRow>,
}
fn ledger() -> &'static Mutex<Ledger> {
static L: OnceLock<Mutex<Ledger>> = OnceLock::new();
L.get_or_init(|| Mutex::new(Ledger::default()))
}
pub fn armed() -> bool {
std::env::var("MEMRA_GLM5_GRAPH_SEL_LEDGER").as_deref() == Ok("1")
}
pub fn prearm(e: &Engine, layer: u16, n_used: usize) -> Res<()> {
if !armed() {
return Ok(());
}
let key = (e.ctx().ordinal(), layer);
let mut l = ledger().lock().unwrap();
if let std::collections::btree_map::Entry::Vacant(slot) = l.slots.entry(key) {
slot.insert((e.htod_i32(&vec![-1i32; n_used])?, e.zeros(n_used)?));
}
Ok(())
}
pub fn record_device(
e: &Engine,
layer: u16,
sel_d: &CudaSlice<i32>,
w_d: &CudaSlice<f32>,
) -> Res<()> {
if !armed() {
return Ok(());
}
let key = (e.ctx().ordinal(), layer);
let mut l = ledger().lock().unwrap();
let Some((sel_slot, w_slot)) = l.slots.get_mut(&key) else {
return Ok(());
};
let n = sel_slot.len().min(sel_d.len());
let mut sv = sel_slot.slice_mut(0..n);
e.stream().memcpy_dtod(&sel_d.slice(0..n), &mut sv)?;
let m = w_slot.len().min(w_d.len());
let mut wv = w_slot.slice_mut(0..m);
e.stream().memcpy_dtod(&w_d.slice(0..m), &mut wv)?;
Ok(())
}
pub fn record_host(dev: usize, layer: u16, sel: &[u32], w: &[f32]) {
if !armed() {
return;
}
ledger().lock().unwrap().host.push(SelRow {
dev,
layer,
sel: sel.iter().map(|&s| s as i32).collect(),
w: w.to_vec(),
});
}
pub fn drain_device(e: &Engine) -> Res<Vec<SelRow>> {
let dev = e.ctx().ordinal();
let l = ledger().lock().unwrap();
let mut out = Vec::new();
for ((d, layer), (sel_slot, w_slot)) in l.slots.iter() {
if *d != dev {
continue;
}
out.push(SelRow {
dev: *d,
layer: *layer,
sel: e.dtoh_i32(sel_slot)?,
w: e.dtoh(w_slot)?,
});
}
Ok(out)
}
pub fn take_host() -> Vec<SelRow> {
std::mem::take(&mut ledger().lock().unwrap().host)
}
pub fn reset_host() {
ledger().lock().unwrap().host.clear();
}