use crate::cache::Cache;
use crate::Engine;
use cudarc::driver::CudaSlice;
use std::ops::Range;
pub struct StreamBufs {
pub vtok_d: CudaSlice<u32>,
pub brk_d: CudaSlice<u32>,
pub pend_d: CudaSlice<u32>,
pub last_pred_d: CudaSlice<u32>,
pub pos_ctr: CudaSlice<i32>,
pub pos_start_d: CudaSlice<i32>,
pub ring_d: CudaSlice<u32>,
pub acc_d: CudaSlice<u32>,
pub m_rounds: usize,
pub k: usize,
}
impl StreamBufs {
pub fn new(e: &Engine, k: usize, m_rounds: usize) -> Result<Self, Box<dyn std::error::Error>> {
Ok(StreamBufs {
vtok_d: e.alloc_u32_zeroed(k + 1)?,
brk_d: e.alloc_u32_zeroed(2)?,
pend_d: e.alloc_u32_zeroed(1)?,
last_pred_d: e.alloc_u32_zeroed(1)?,
pos_ctr: e.htod_i32(&[0])?,
pos_start_d: e.htod_i32(&[0])?,
ring_d: e.alloc_u32_zeroed(m_rounds * (k + 1) + 1)?,
acc_d: e.alloc_u32_zeroed(2)?,
m_rounds,
k,
})
}
pub fn drain_ring(&self, e: &Engine) -> Result<Vec<u32>, Box<dyn std::error::Error>> {
let h = e.dtoh_u32(&self.ring_d)?;
let cnt = (h[0] as usize).min(self.ring_d.len() - 1);
Ok(h[1..1 + cnt].to_vec())
}
}
pub fn kv_len_ptr_table(
e: &Engine,
cache: &Cache,
pos_ctr: Option<&CudaSlice<i32>>,
) -> Result<CudaSlice<u64>, Box<dyn std::error::Error>> {
kv_len_ptr_table_range(e, cache, 0..cache.kv.len(), pos_ctr)
}
pub fn kv_len_ptr_table_range(
e: &Engine,
cache: &Cache,
layers: Range<usize>,
pos_ctr: Option<&CudaSlice<i32>>,
) -> Result<CudaSlice<u64>, Box<dyn std::error::Error>> {
use cudarc::driver::DevicePtr;
assert!(layers.start <= layers.end && layers.end <= cache.kv.len());
let mut ptrs: Vec<u64> = cache
.kv
.get(layers)
.expect("validated KV layer range")
.iter()
.map(|kv| match kv.as_ref() {
Some(kvl) => {
let __s_g = e.stream();
let (p, _g) = kvl.len_d.device_ptr(&__s_g);
p as u64
}
None => 0u64,
})
.collect();
if let Some(pc) = pos_ctr {
let __s_g = e.stream();
let (p, _g) = pc.device_ptr(&__s_g);
ptrs.push(p as u64);
}
Ok(e.htod_u64(&ptrs)?)
}