use std::cell::{Cell, RefCell};
use std::collections::HashMap;
use std::path::PathBuf;
use taconite::{Timing, bf16_to_f32, f32_to_bf16};
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};
use crate::Error;
const BF16: usize = 2;
#[derive(Debug, Clone)]
pub struct GemmSpec {
pub m: usize,
pub k: usize,
pub n: usize,
pub lda: usize,
pub c_stride: usize,
pub a_elems: Vec<usize>,
pub c_elems: Vec<usize>,
}
#[derive(Debug, Clone)]
pub struct MhaSpec {
pub rows: usize,
pub elems: [usize; 4],
pub kv_words: Vec<usize>,
pub win_words: Vec<usize>,
}
pub struct Rows {
pub buf: Buffer,
pub width: usize,
}
impl Rows {
pub fn put(&mut self, x: &[f32], cols: usize) -> Result<(), Error> {
let rows = x.len() / cols;
let w = self.width;
let dst = self.buf.as_mut_slice::<u16>();
for (r, src) in x.chunks(cols).enumerate() {
for (o, &v) in dst[r * w..r * w + cols].iter_mut().zip(src) {
*o = f32_to_bf16(v);
}
}
self.buf.sub(0, rows * w * BF16)?.sync_to_device()?;
Ok(())
}
pub fn get(&self, rows: usize, cols: usize) -> Result<Vec<f32>, Error> {
let w = self.width;
self.buf.sub(0, rows * w * BF16)?.sync_from_device()?;
let src = self.buf.as_slice::<u16>();
let mut out = Vec::with_capacity(rows * cols);
for r in 0..rows {
out.extend(src[r * w..r * w + cols].iter().map(|&v| bf16_to_f32(v)));
}
Ok(out)
}
pub fn view(&self, col: usize, elems: usize) -> Result<Buffer, Error> {
Ok(self.buf.sub(col * BF16, elems * BF16)?)
}
}
struct Src {
ctx: String,
xclbin: PathBuf,
name: String,
insts: PathBuf,
ops: u64,
}
pub struct Npu {
pub session: Session,
srcs: HashMap<String, Src>,
kernels: RefCell<HashMap<String, Kernel>>,
pub swaps: Cell<usize>,
pub gemms: HashMap<String, GemmSpec>,
pub mhas: Vec<MhaSpec>,
mha_cur: RefCell<HashMap<usize, (u32, u32)>>,
pub contexts: usize,
}
fn gemm_kernel(key: &str, n: usize) -> String {
if n == 1 { key.to_string() } else { format!("{key}_x{n}") }
}
impl Npu {
pub fn open(m: &Manifest) -> Result<Self, Error> {
let session = Session::open(0)?;
let mut srcs = HashMap::new();
let mut gemms: HashMap<String, GemmSpec> = HashMap::new();
let mut gemm_ctx: HashMap<String, String> = HashMap::new();
let mut mhas = Vec::new();
let mut ctxs: Vec<PathBuf> = Vec::new();
let words = |s: &str| -> Result<Vec<usize>, Error> {
s.split(',').map(|w| w.parse().map_err(|_| Error::Bundle(format!("bad instruction word {w}")))).collect()
};
for r in m.records() {
let (key, ctx, ops) = match r.tag.as_str() {
"gemm" => {
let g = GemmSpec {
m: r.get("M")?,
k: r.get("K")?,
n: r.get("N")?,
lda: r.get("lda")?,
c_stride: r.get("c_stride")?,
a_elems: vec![r.get("a_elems")?],
c_elems: vec![r.get("c_elems")?],
};
let ops = (2 * g.m * g.k * g.n) as u64;
let key = r.field(0)?.to_string();
gemms.insert(key.clone(), g);
gemm_ctx.insert(key.clone(), r.str("ctx")?.to_string());
(key, r.str("ctx")?.to_string(), ops)
}
"gemm_chunks" => {
let (key, n): (&str, usize) = (r.field(0)?, r.field_as(1)?);
let g = gemms
.get_mut(key)
.ok_or_else(|| Error::Bundle(format!("gemm_chunks {key} before its gemm")))?;
if g.a_elems.len() + 1 != n {
return Err(Error::Bundle(format!("gemm_chunks {key} {n} out of order")));
}
g.a_elems.push(r.get("a_elems")?);
g.c_elems.push(r.get("c_elems")?);
let ops = (2 * n * g.m * g.k * g.n) as u64;
(gemm_kernel(key, n), gemm_ctx[key].clone(), ops)
}
"mha" => {
let spec = MhaSpec {
rows: r.field_as(0)?,
elems: [r.get("q_elems")?, r.get("k_elems")?, r.get("v_elems")?, r.get("o_elems")?],
kv_words: words(r.str("kv_words")?)?,
win_words: words(r.str("win_words")?)?,
};
let key = format!("mha_{}", spec.rows);
mhas.push(spec);
(key, r.str("ctx")?.to_string(), 0)
}
_ => continue,
};
let x = m.xclbin(&ctx)?;
if !ctxs.contains(&x.path) {
ctxs.push(x.path.clone());
}
let src = Src { ctx, xclbin: x.path.clone(), name: x.kernel.clone(), insts: m.path(r.str("insts")?), ops };
srcs.insert(key, src);
}
mhas.sort_by_key(|s| s.rows);
Ok(Npu {
session,
srcs,
kernels: Default::default(),
swaps: Cell::new(0),
gemms,
mhas,
mha_cur: Default::default(),
contexts: ctxs.len(),
})
}
pub fn gemm(&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 rows(&self, rows: usize, width: usize) -> Result<Rows, Error> {
let mut buf = self.session.alloc(rows * width * BF16)?;
buf.as_mut_slice::<u16>().fill(0);
buf.sync_to_device()?;
Ok(Rows { buf, width })
}
fn with_kernel<T>(&self, key: &str, f: impl FnOnce(&Kernel) -> Result<T, Error>) -> Result<T, Error> {
if !self.kernels.borrow().contains_key(key) {
let s = self.srcs.get(key).ok_or_else(|| Error::Bundle(format!("no kernel {key} in the bundle")))?;
let load = || {
self.session
.load_kernel(&s.xclbin, &s.insts, Some(&s.name), s.ops)
.map_err(|e| Error::Npu(format!("loading {key}: {e}")))
};
let k = match load() {
Ok(k) => k,
Err(e) => {
let mut ks = self.kernels.borrow_mut();
let others: Vec<String> = ks.keys().filter(|k| self.srcs[*k].ctx != s.ctx).cloned().collect();
if others.is_empty() {
return Err(e);
}
let mut ctxs: Vec<&str> = others.iter().map(|k| self.srcs[k].ctx.as_str()).collect();
ctxs.sort();
ctxs.dedup();
self.swaps.set(self.swaps.get() + ctxs.len());
for k in &others {
ks.remove(k);
}
self.mha_cur.borrow_mut().retain(|rows, _| !others.contains(&format!("mha_{rows}")));
drop(ks);
load()?
}
};
self.kernels.borrow_mut().insert(key.to_string(), k);
}
let ks = self.kernels.borrow();
f(&ks[key])
}
#[allow(clippy::too_many_arguments)]
pub fn run_gemm(
&self,
key: &str,
a: &Rows,
col: usize,
w: &Buffer,
c: &Rows,
rows: usize,
t: &mut Timing,
) -> Result<(), Error> {
let g = self.gemm(key)?;
let n = rows.div_ceil(g.m);
let (ae, ce) = match (g.a_elems.get(n - 1), g.c_elems.get(n - 1)) {
(Some(&a), Some(&c)) => (a, c),
_ => return Err(Error::Input(format!("{key}: {rows} rows (the bundle has streams for {})", g.a_elems.len() * g.m))),
};
let (av, cv) = (a.view(col, ae)?, c.view(0, ce)?);
let d = self.with_kernel(&gemm_kernel(key, n), |k| {
k.run(&[&av, w, &cv]).map_err(|e| Error::Npu(format!("{key}: {e}")))
})?;
t.add(&format!("npu:{key}"), d);
Ok(())
}
pub fn run_mha(&self, qkv: &Rows, o: &Rows, n: usize, no_window: u32, t: &mut Timing) -> Result<(), Error> {
let spec = self
.mhas
.iter()
.find(|s| s.rows >= n)
.ok_or_else(|| Error::Input(format!("{n} patches: the largest MHA bucket is {:?}", self.mhas.last().map(|s| s.rows))))?;
let key = format!("mha_{}", spec.rows);
let want = (n as u32, no_window);
let [qe, ke, ve, oe] = spec.elems;
let views = [qkv.view(0, qe)?, qkv.view(0, ke)?, qkv.view(0, ve)?, o.view(0, oe)?];
let d = self.with_kernel(&key, |k| {
if self.mha_cur.borrow().get(&spec.rows) != Some(&want) {
let mut words: Vec<(usize, u32)> = spec.kv_words.iter().map(|&i| (i, want.0)).collect();
words.extend(spec.win_words.iter().map(|&i| (i, want.1)));
k.set_insts_words(&words).map_err(|e| Error::Npu(format!("{key}: {e}")))?;
self.mha_cur.borrow_mut().insert(spec.rows, want);
}
k.run(&[&views[0], &views[1], &views[2], &views[3]]).map_err(|e| Error::Npu(format!("{key}: {e}")))
})?;
t.add("npu:mha", d);
Ok(())
}
}