custom
code
sovereign-compute
File size: 1,203 Bytes
1e4543d
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
use cudarc::driver::{CudaDevice, CudaSlice};
use half::f16;
use std::sync::Arc;

#[derive(Clone)]
pub struct AttentionMemory {
    pub device: Arc<CudaDevice>,
    pub k_cache: CudaSlice<f16>,
    pub v_cache: CudaSlice<f16>,
    pub max_seq_len: usize,
    pub num_layers: usize,
    pub num_heads: usize,
    pub head_dim: usize,
    pub current_len: usize,
}

impl AttentionMemory {
    pub fn new(device: Arc<CudaDevice>, num_layers: usize, num_heads: usize, head_dim: usize, max_seq_len: usize) -> anyhow::Result<Self> {
        let elems = num_layers * num_heads * max_seq_len * head_dim;
        let k_cache = unsafe { device.alloc::<f16>(elems)? };
        let v_cache = unsafe { device.alloc::<f16>(elems)? };
        Ok(Self { device, k_cache, v_cache, max_seq_len, num_layers, num_heads, head_dim, current_len: 0 })
    }
    pub fn append_kv(&mut self, _layer: usize, _k: &CudaSlice<f16>, _v: &CudaSlice<f16>, seq_len: usize) -> anyhow::Result<()> {
        if self.current_len + seq_len > self.max_seq_len { anyhow::bail!("KV cache overflow"); }
        self.current_len += seq_len;
        Ok(())
    }
    pub fn reset(&mut self) { self.current_len = 0; }
}