use std::collections::HashMap;
use std::path::PathBuf;
use std::time::Duration;
use taconite_bundle::Manifest;
#[cfg(feature = "direct")]
pub use taconite::direct::{Buffer, Kernel, Session};
#[cfg(all(feature = "xrt", not(feature = "direct")))]
pub use taconite::{Buffer, Kernel, Session};
#[cfg(not(any(feature = "xrt", feature = "direct")))]
compile_error!("no NPU path: enable feature `xrt` (the default) or `direct`");
use crate::Error;
pub const BF16: usize = 2;
#[derive(Debug, Clone)]
pub struct GemmSpec {
pub key: String,
pub m: usize,
pub k: usize,
pub n: usize,
pub b_bytes: usize,
}
pub struct Npu {
pub session: Session,
kernels: HashMap<String, Kernel>,
pub gemms: HashMap<String, GemmSpec>,
pub contexts: usize,
}
pub struct Io {
pub key: String,
pub a: Buffer,
pub c: Buffer,
chunks: Vec<(Buffer, Buffer)>,
pub rows: usize,
pub k: usize,
pub n: usize,
}
impl Npu {
pub fn open(m: &Manifest) -> Result<Self, Error> {
let session = Session::open(0)?;
let mut kernels = HashMap::new();
let mut gemms = HashMap::new();
let mut ctxs: Vec<PathBuf> = Vec::new();
for r in m.records() {
let (key, xclbin, insts, name, ops) = match r.tag.as_str() {
"gemm" => {
let g = GemmSpec {
key: r.field(0)?.to_string(),
m: r.get("M")?,
k: r.get("K")?,
n: r.get("N")?,
b_bytes: r.get("b_bytes")?,
};
let x = m.xclbin(r.str("ctx")?)?;
let ops = (2 * g.m * g.k * g.n) as u64;
gemms.insert(g.key.clone(), g);
(r.field(0)?, x.path.clone(), m.path(r.str("insts")?), x.kernel.clone(), ops)
}
"op" => (r.field(0)?, m.path(r.str("xclbin")?), m.path(r.str("insts")?), r.str("name")?.to_string(), 0),
_ => continue,
};
let k = session
.load_kernel(&xclbin, &insts, Some(&name), ops)
.map_err(|e| Error::Npu(format!("loading {key}: {e}")))?;
if !ctxs.contains(&xclbin) {
ctxs.push(xclbin);
}
kernels.insert(key.to_string(), k);
}
Ok(Npu { session, kernels, gemms, contexts: ctxs.len() })
}
pub fn spec(&self, key: &str) -> Result<&GemmSpec, Error> {
self.gemms.get(key).ok_or_else(|| Error::Bundle(format!("no GEMM {key} in the bundle")))
}
pub fn upload(&self, bytes: &[u8]) -> Result<Buffer, Error> {
let mut b = self.session.alloc(bytes.len())?;
b.write(bytes)?;
Ok(b)
}
pub fn zeros(&self, elems: usize) -> Result<Buffer, Error> {
let b = self.session.alloc(elems * BF16)?;
b.sync_to_device()?;
Ok(b)
}
fn chunks(&self, g: &GemmSpec, a: &Buffer, c: &Buffer, rows: usize) -> Result<Vec<(Buffer, Buffer)>, Error> {
if rows % g.m != 0 {
return Err(Error::Bundle(format!("{}: {rows} rows is not a multiple of M = {}", g.key, g.m)));
}
(0..rows / g.m)
.map(|i| {
Ok((a.sub(i * g.m * g.k * BF16, g.m * g.k * BF16)?, c.sub(i * g.m * g.n * BF16, g.m * g.n * BF16)?))
})
.collect()
}
pub fn io(&self, key: &str, rows: usize) -> Result<Io, Error> {
let g = self.spec(key)?.clone();
let a = self.zeros(rows * g.k)?;
let c = self.zeros(rows * g.n)?;
let chunks = self.chunks(&g, &a, &c, rows)?;
Ok(Io { key: key.into(), a, c, chunks, rows, k: g.k, n: g.n })
}
pub fn io_chained(&self, key: &str, src: &Io) -> Result<Io, Error> {
let g = self.spec(key)?.clone();
if src.n != g.k {
return Err(Error::Bundle(format!("{key}: A width {} != K {}", src.n, g.k)));
}
let a = src.c.sub(0, src.rows * g.k * BF16)?;
let c = self.zeros(src.rows * g.n)?;
let chunks = self.chunks(&g, &a, &c, src.rows)?;
Ok(Io { key: key.into(), a, c, chunks, rows: src.rows, k: g.k, n: g.n })
}
fn kernel(&self, key: &str) -> Result<&Kernel, Error> {
self.kernels.get(key).ok_or_else(|| Error::Bundle(format!("no kernel {key} in the bundle")))
}
pub fn gemm(&self, io: &Io, w: &Buffer) -> Result<Duration, Error> {
let k = self.kernel(&io.key)?;
let mut t = Duration::ZERO;
for (a, c) in &io.chunks {
t += k.run(&[a, w, c]).map_err(|e| Error::Npu(format!("{}: {e}", io.key)))?;
}
Ok(t)
}
pub fn op(&self, key: &str, args: &[&Buffer]) -> Result<Duration, Error> {
self.kernel(key)?.run(args).map_err(|e| Error::Npu(format!("{key}: {e}")))
}
}