spark_storage/
scratch_pool.rs1use anyhow::{Result, bail};
27use std::collections::{HashMap, VecDeque};
28
29use crate::cuda_min::DeviceBuffer;
30
31#[derive(Clone, Copy, Debug, PartialEq, Eq, Hash)]
32pub struct ResidentKey {
33 pub layer: u32,
34 pub block: u32,
35}
36
37#[derive(Clone, Copy, Debug)]
38pub struct ScratchDims {
39 pub num_slots: u32,
40 pub num_kv_heads: u16,
41 pub group_stride: u64, }
43
44impl ScratchDims {
45 pub fn slot_bytes(&self) -> usize {
46 (2 * self.num_kv_heads as u64 * self.group_stride) as usize
47 }
48 pub fn pool_bytes(&self) -> usize {
49 self.num_slots as usize * self.slot_bytes()
50 }
51}
52
53pub struct ScratchPool {
54 dims: ScratchDims,
55 pool: DeviceBuffer,
56 residents: Vec<Option<ResidentKey>>, lookup: HashMap<ResidentKey, u32>,
58 free_list: VecDeque<u32>,
59}
60
61impl ScratchPool {
62 pub fn new(dims: ScratchDims) -> Result<Self> {
63 if dims.num_slots == 0 {
64 bail!("ScratchPool requires at least one slot");
65 }
66 let pool = DeviceBuffer::new(dims.pool_bytes())?;
67 let residents = vec![None; dims.num_slots as usize];
68 let free_list = (0..dims.num_slots).collect();
69 Ok(Self {
70 dims,
71 pool,
72 residents,
73 lookup: HashMap::new(),
74 free_list,
75 })
76 }
77
78 pub fn dims(&self) -> ScratchDims {
79 self.dims
80 }
81 pub fn pool_dev_ptr(&self) -> u64 {
82 self.pool.ptr
83 }
84 pub fn slot_dev_ptr(&self, slot: u32) -> u64 {
85 self.pool.ptr + (slot as u64) * (self.dims.slot_bytes() as u64)
86 }
87 pub fn slot_k_ptr(&self, slot: u32, kv_head: u16) -> u64 {
89 self.slot_dev_ptr(slot) + (kv_head as u64) * self.dims.group_stride
90 }
91 pub fn slot_v_ptr(&self, slot: u32, kv_head: u16) -> u64 {
93 self.slot_dev_ptr(slot)
94 + (self.dims.num_kv_heads as u64 + kv_head as u64) * self.dims.group_stride
95 }
96
97 pub fn lookup(&self, key: ResidentKey) -> Option<u32> {
98 self.lookup.get(&key).copied()
99 }
100
101 pub fn invalidate(&mut self, key: ResidentKey) {
110 if let Some(slot) = self.lookup.remove(&key) {
111 self.residents[slot as usize] = None;
112 self.free_list.push_back(slot);
113 }
114 }
115
116 pub fn capacity(&self) -> u32 {
117 self.dims.num_slots
118 }
119 pub fn free_count(&self) -> u32 {
120 self.free_list.len() as u32
121 }
122
123 pub fn assign(&mut self, key: ResidentKey, evict_candidates: &[u32]) -> Result<u32> {
128 if let Some(&slot) = self.lookup.get(&key) {
129 return Ok(slot); }
131 let slot = match self.free_list.pop_front() {
132 Some(s) => s,
133 None => {
134 let mut chosen = None;
137 for &c in evict_candidates {
138 if self
139 .residents
140 .get(c as usize)
141 .and_then(|r| r.as_ref())
142 .is_some()
143 {
144 chosen = Some(c);
145 break;
146 }
147 }
148 let s = chosen.ok_or_else(|| {
149 anyhow::anyhow!("no slot available and no eviction candidate is resident")
150 })?;
151 if let Some(prev) = self.residents[s as usize].take() {
152 self.lookup.remove(&prev);
153 }
154 s
155 }
156 };
157 self.residents[slot as usize] = Some(key);
158 self.lookup.insert(key, slot);
159 Ok(slot)
160 }
161
162 pub fn residents(&self) -> Vec<(u32, ResidentKey)> {
165 self.residents
166 .iter()
167 .enumerate()
168 .filter_map(|(i, r)| r.map(|k| (i as u32, k)))
169 .collect()
170 }
171
172 pub fn clear(&mut self) {
174 self.lookup.clear();
175 for r in self.residents.iter_mut() {
176 *r = None;
177 }
178 self.free_list.clear();
179 for s in 0..self.dims.num_slots {
180 self.free_list.push_back(s);
181 }
182 }
183}
184
185#[cfg(test)]
186mod tests {
187 use super::*;
188
189 fn ctx() -> crate::cuda_min::CudaCtx {
190 crate::cuda_min::CudaCtx::new(0).expect("cuda init")
191 }
192
193 #[test]
194 #[ignore = "requires GPU"]
195 fn assign_and_lookup() {
196 let _ctx = ctx();
197 let mut pool = ScratchPool::new(ScratchDims {
198 num_slots: 4,
199 num_kv_heads: 2,
200 group_stride: 4096,
201 })
202 .unwrap();
203 let k0 = ResidentKey { layer: 0, block: 7 };
204 let s0 = pool.assign(k0, &[]).unwrap();
205 assert_eq!(pool.lookup(k0), Some(s0));
206 let s0_again = pool.assign(k0, &[]).unwrap();
208 assert_eq!(s0, s0_again);
209 for b in 8..11 {
211 pool.assign(ResidentKey { layer: 0, block: b }, &[])
212 .unwrap();
213 }
214 assert_eq!(pool.free_count(), 0);
215 let evicted = pool
217 .assign(
218 ResidentKey {
219 layer: 0,
220 block: 99,
221 },
222 &[s0],
223 )
224 .unwrap();
225 assert_eq!(evicted, s0); assert_eq!(pool.lookup(k0), None); }
228
229 #[test]
230 #[ignore = "requires GPU"]
231 fn slot_pointer_layout() {
232 let _ctx = ctx();
233 let pool = ScratchPool::new(ScratchDims {
234 num_slots: 2,
235 num_kv_heads: 4,
236 group_stride: 4096,
237 })
238 .unwrap();
239 let base = pool.pool_dev_ptr();
240 assert_eq!(pool.slot_dev_ptr(0), base);
241 assert_eq!(pool.slot_dev_ptr(1), base + 8 * 4096);
242 assert_eq!(pool.slot_k_ptr(0, 2), base + 2 * 4096);
243 assert_eq!(pool.slot_v_ptr(0, 2), base + (4 + 2) * 4096);
244 }
245}