spark_model/layers/
qsa_snapshot.rs1use anyhow::Result;
13use spark_runtime::gpu::GpuBackend;
14
15use super::{QsaIndexer, QsaSeqState};
16use crate::layers::ops;
17
18impl QsaIndexer {
19 pub fn snapshot_aux(
24 &self,
25 st: &QsaSeqState,
26 gpu: &dyn GpuBackend,
27 stream: u64,
28 ) -> Result<Vec<u8>> {
29 let hd = self.hd as usize;
30 let key_bytes = st.ingested * hd * 2;
31 let mut blob = Vec::with_capacity(16 + key_bytes);
32 blob.extend_from_slice(&(st.ingested as u64).to_le_bytes());
33 blob.extend_from_slice(&(st.pooled as u64).to_le_bytes());
34 let off = blob.len();
35 blob.resize(off + key_bytes, 0);
36 if key_bytes > 0 {
37 gpu.copy_d2h_on_stream(st.raw_keys, &mut blob[off..], stream)?;
38 }
39 Ok(blob)
40 }
41
42 pub fn restore_aux(
45 &self,
46 st: &mut QsaSeqState,
47 blob: &[u8],
48 gpu: &dyn GpuBackend,
49 stream: u64,
50 ) -> Result<()> {
51 anyhow::ensure!(blob.len() >= 16, "QSA aux blob truncated");
52 let ingested = u64::from_le_bytes(blob[..8].try_into().unwrap()) as usize;
53 let pooled = u64::from_le_bytes(blob[8..16].try_into().unwrap()) as usize;
54 let hd = self.hd as usize;
55 anyhow::ensure!(
56 blob.len() == 16 + ingested * hd * 2,
57 "QSA aux blob size mismatch"
58 );
59 anyhow::ensure!(ingested <= self.max_tokens, "QSA aux exceeds key cache");
60 if ingested > 0 {
61 gpu.copy_h2d_async(&blob[16..], st.raw_keys, stream)?;
62 }
63 st.ingested = ingested;
64 st.pooled = 0;
65 if pooled > 0 {
66 ops::qsa_block_pool(
67 gpu,
68 self.k_pool_k,
69 st.raw_keys,
70 self.k_norm_w,
71 st.block_keys,
72 0,
73 pooled as u32,
74 self.ratio,
75 self.hd,
76 self.rot,
77 self.theta,
78 self.eps,
79 stream,
80 )?;
81 st.pooled = pooled;
82 }
83 Ok(())
84 }
85}