spark_model/layers/ple/
aux_state.rs1use anyhow::Result;
8use spark_runtime::gpu::{DevicePtr, GpuBackend};
9
10use super::{PleLayer, PleSeqState};
11use crate::layers::ple::ids::ple_ngram_ids;
12
13impl PleLayer {
14 pub fn snapshot_aux(
18 &self,
19 st: &PleSeqState,
20 gpu: &dyn GpuBackend,
21 stream: u64,
22 ) -> Result<Vec<u8>> {
23 let conv_bytes = self.state_len * self.hc_mult * self.hidden * 4;
24 let mut blob = Vec::with_capacity(4 + st.history.len() * 4 + conv_bytes);
25 blob.extend_from_slice(&(st.history.len() as u32).to_le_bytes());
26 for t in &st.history {
27 blob.extend_from_slice(&t.to_le_bytes());
28 }
29 let off = blob.len();
30 blob.resize(off + conv_bytes, 0);
31 gpu.copy_d2h_on_stream(st.conv, &mut blob[off..], stream)?;
32 Ok(blob)
33 }
34
35 pub fn restore_aux(
37 &self,
38 st: &mut PleSeqState,
39 blob: &[u8],
40 gpu: &dyn GpuBackend,
41 stream: u64,
42 ) -> Result<()> {
43 anyhow::ensure!(blob.len() >= 4, "PLE aux blob truncated");
44 let n = u32::from_le_bytes(blob[..4].try_into().unwrap()) as usize;
45 let conv_bytes = self.state_len * self.hc_mult * self.hidden * 4;
46 anyhow::ensure!(
47 blob.len() == 4 + n * 4 + conv_bytes,
48 "PLE aux blob size mismatch"
49 );
50 st.history = blob[4..4 + n * 4]
51 .chunks_exact(4)
52 .map(|c| u32::from_le_bytes(c.try_into().unwrap()))
53 .collect();
54 st.prestaged_va = None;
55 gpu.copy_h2d_async(&blob[4 + n * 4..], st.conv, stream)?;
56 Ok(())
57 }
58}
59
60impl PleLayer {
61 pub(super) fn reset(
63 &self,
64 st: &mut PleSeqState,
65 gpu: &dyn GpuBackend,
66 stream: u64,
67 ) -> Result<()> {
68 st.history = vec![self.dims.eos_token_id; self.dims.context_len()];
69 st.prestaged_va = None;
70 let zeros = vec![0u8; self.state_len * self.hc_mult * self.hidden * 4];
71 gpu.copy_h2d_async(&zeros, st.conv, stream)?;
72 Ok(())
73 }
74
75 pub fn prestage(
86 &self,
87 st: &mut PleSeqState,
88 tokens: &[u32],
89 gpu: &dyn GpuBackend,
90 stream: u64,
91 ) -> Result<()> {
92 if st.history.len() != self.dims.context_len() {
93 self.reset(st, gpu, stream)?;
94 }
95 let mut window = st.history.clone();
96 window.extend_from_slice(tokens);
97 let all = ple_ngram_ids(&self.dims, &window);
98 let rows = &all[all.len() - tokens.len()..];
99 let flat: Vec<u64> = rows.iter().flat_map(|r| r.iter().copied()).collect();
100 let va = self.gather_host(&flat, gpu, stream)?;
101 let keep = self.dims.context_len();
102 st.history = window[window.len() - keep..].to_vec();
103 st.prestaged_va = Some(va);
104 st.last_staged_va = va;
105 Ok(())
106 }
107
108 pub fn release_seq_state(&self, st: &mut PleSeqState, gpu: &dyn GpuBackend) -> Result<()> {
118 if st.conv.is_null() {
119 return Ok(());
120 }
121 let r = gpu.free(st.conv);
122 st.conv = DevicePtr(0);
123 st.history.clear();
124 st.prestaged_va = None;
125 st.last_staged_va = 0;
126 r
127 }
128}