spark_storage/
predictor.rs

1// SPDX-License-Identifier: AGPL-3.0-only
2//
3// GPU-side predictor: loads the three PTX modules emitted by build.rs,
4// owns persistent device buffers for `P` (projection) and `A_g` (anchor
5// vectors per block × kv_head), and exposes three operations:
6//
7//   1. project_q       — Q @ P, per decode step
8//   2. project_kv_block — A_g[block, kv_head] = mean_token(K_block @ P), at write time
9//   3. score_blocks    — q_proj @ A_g_seq.T, max-reduced per block
10//
11// Lossless contract: scores are *advisory* (eviction priority + prefetch
12// order). Mispredictions never gate attention correctness.
13
14use anyhow::{Context, Result, bail};
15use half::bf16;
16use std::ffi::c_void;
17
18use crate::cuda_min::{
19    CudaCtx, CudaModule, DeviceBuffer, copy_d_to_h_async, copy_h_to_d_async, launch_kernel,
20    stream_sync,
21};
22use crate::projection::{PredictorShape, build_projection};
23
24include!(concat!(env!("OUT_DIR"), "/storage_ptx.rs"));
25
26#[derive(Clone, Copy, Debug)]
27pub struct PredictorDims {
28    pub num_layers: usize,
29    pub num_q_heads: usize,
30    pub num_kv_heads: usize,
31    pub head_dim: usize,
32    pub r: usize,
33    pub block_size: usize,
34    pub max_blocks: usize,
35}
36
37impl PredictorDims {
38    pub fn validate(&self) -> Result<()> {
39        if !self.num_q_heads.is_multiple_of(self.num_kv_heads) {
40            bail!(
41                "num_q_heads ({}) must divide num_kv_heads ({})",
42                self.num_q_heads,
43                self.num_kv_heads
44            );
45        }
46        Ok(())
47    }
48    pub fn gqa_ratio(&self) -> i32 {
49        (self.num_q_heads / self.num_kv_heads) as i32
50    }
51    pub fn a_g_bytes(&self) -> usize {
52        // K_lr is stored per-token: [num_layers, num_blocks, num_kv_heads, block_size, r]
53        self.num_layers * self.max_blocks * self.num_kv_heads * self.block_size * self.r * 2
54    }
55    pub fn per_layer_block_floats(&self) -> usize {
56        self.num_kv_heads * self.block_size * self.r
57    }
58    pub fn p_bytes(&self) -> usize {
59        self.head_dim * self.r * 2
60    }
61}
62
63pub struct Predictor {
64    dims: PredictorDims,
65    _modules: Vec<CudaModule>,
66    f_q_proj: u64,
67    f_kv_proj: u64,
68    f_score: u64,
69    p_dev: DeviceBuffer,   // [head_dim, r] BF16, immutable after init
70    a_g_dev: DeviceBuffer, // [num_layers, max_blocks, num_kv_heads, r] BF16
71}
72
73impl Predictor {
74    pub fn new(ctx: &CudaCtx, dims: PredictorDims, projection_seed: u64) -> Result<Self> {
75        Self::new_on_stream(ctx.stream, dims, projection_seed)
76    }
77
78    pub fn new_on_stream(stream: u64, dims: PredictorDims, projection_seed: u64) -> Result<Self> {
79        dims.validate()?;
80        // Load only the predictor PTX modules; tiled-attention kernels are
81        // owned by `TiledAttention` so each subsystem stays self-contained.
82        let mut modules: Vec<CudaModule> = Vec::new();
83        let mut f_q_proj = 0u64;
84        let mut f_kv_proj = 0u64;
85        let mut f_score = 0u64;
86        for entry in STORAGE_PTX.iter() {
87            match entry.name {
88                "q_lowrank_project" | "kv_lowrank_project" | "predictor_score" => {
89                    let m = CudaModule::from_ptx(entry.ptx)
90                        .with_context(|| format!("load PTX module {}", entry.name))?;
91                    match entry.name {
92                        "q_lowrank_project" => f_q_proj = m.function("q_lowrank_project")?,
93                        "kv_lowrank_project" => f_kv_proj = m.function("kv_lowrank_project")?,
94                        "predictor_score" => f_score = m.function("predictor_score")?,
95                        _ => unreachable!(),
96                    }
97                    modules.push(m);
98                }
99                _ => {} // tiled-attention modules handled elsewhere
100            }
101        }
102        if f_q_proj == 0 || f_kv_proj == 0 || f_score == 0 {
103            bail!("missing predictor kernel function — PTX list incomplete");
104        }
105        // Allocate and upload P.
106        let shape = PredictorShape::new(dims.head_dim, dims.r);
107        let p_host = build_projection(shape, projection_seed);
108        let p_dev = DeviceBuffer::new(dims.p_bytes())?;
109        copy_h_to_d_async(
110            p_dev.ptr,
111            p_host.as_ptr() as *const c_void,
112            dims.p_bytes(),
113            stream,
114        )?;
115        // Preflight A_g sizing against free HBM. A_g grows linearly with
116        // `r × max_blocks × num_layers × num_kv_heads × block_size`, so a
117        // greedy `r=32` default plus a max-seq-len-sized block pool can blow
118        // past free HBM on the EP head (where the MoE-transpose pass already
119        // consumed ~40 GB). The raw `cuMemAlloc_v2` error gives the user
120        // nothing to act on; this preflight names the specific knobs.
121        // PR #47 follow-up — preflight pattern only, no behavior change on
122        // the happy path.
123        let a_g_need = dims.a_g_bytes();
124        let (free_hbm, _total_hbm) = crate::cuda_min::mem_info()?;
125        // Leave a 5% safety margin for the scratch pool, tiled-attention,
126        // and the smaller predictor buffers that follow.
127        let a_g_budget = free_hbm.saturating_mul(95) / 100;
128        if a_g_need > a_g_budget {
129            bail!(
130                "HSS predictor A_g would need {:.2} GB but only {:.2} GB of HBM is free \
131                 (5% margin reserved for scratch + tiled-attention).\n\
132                 Tune one of:\n  \
133                 - --high-speed-swap-rank: current {} ; try {} (halves A_g)\n  \
134                 - --max-seq-len: max_blocks={} → currently dominates A_g; halve --max-seq-len to halve A_g\n  \
135                 - --kv-cache-dtype nvfp4: halves the KV pool, freeing room for A_g\n  \
136                 - --gpu-memory-utilization: lower so weight-side allocations leave more HBM\n\
137                 A_g sizing = num_layers ({}) × max_blocks ({}) × num_kv_heads ({}) × block_size ({}) × r ({}) × 2 bytes.",
138                a_g_need as f64 / (1u64 << 30) as f64,
139                a_g_budget as f64 / (1u64 << 30) as f64,
140                dims.r,
141                dims.r / 2,
142                dims.max_blocks,
143                dims.num_layers,
144                dims.max_blocks,
145                dims.num_kv_heads,
146                dims.block_size,
147                dims.r,
148            );
149        }
150        // Allocate A_g (zeroed initially via cuMemAlloc which doesn't zero — but
151        // unwritten slots are unused; the predictor never reads a slot that
152        // hasn't been populated by `project_kv_block`).
153        let a_g_dev = DeviceBuffer::new(a_g_need)?;
154        stream_sync(stream)?;
155        Ok(Self {
156            dims,
157            _modules: modules,
158            f_q_proj,
159            f_kv_proj,
160            f_score,
161            p_dev,
162            a_g_dev,
163        })
164    }
165
166    /// `q` device pointer to `[num_q_heads, head_dim]` BF16.
167    /// `q_proj` device pointer to `[num_q_heads, r]` BF16 output.
168    pub fn project_q(&self, ctx: &CudaCtx, q: u64, q_proj: u64) -> Result<()> {
169        self.project_q_on_stream(ctx.stream, q, q_proj)
170    }
171
172    /// Stream-only variant for production callers that already own a CUDA
173    /// context (and therefore don't need the test-side `CudaCtx` wrapper).
174    pub fn project_q_on_stream(&self, stream: u64, q: u64, q_proj: u64) -> Result<()> {
175        let mut q_v = q;
176        let mut p_v = self.p_dev.ptr;
177        let mut o_v = q_proj;
178        let mut nq = self.dims.num_q_heads as i32;
179        let mut hd = self.dims.head_dim as i32;
180        let mut r = self.dims.r as i32;
181        let mut params = [
182            &mut q_v as *mut _ as *mut c_void,
183            &mut p_v as *mut _ as *mut c_void,
184            &mut o_v as *mut _ as *mut c_void,
185            &mut nq as *mut _ as *mut c_void,
186            &mut hd as *mut _ as *mut c_void,
187            &mut r as *mut _ as *mut c_void,
188        ];
189        launch_kernel(
190            self.f_q_proj,
191            (self.dims.num_q_heads as u32, 1, 1),
192            (self.dims.r as u32, 1, 1),
193            0,
194            stream,
195            &mut params,
196        )
197    }
198
199    /// Update A_g for a single (layer, block_id) slot from the K data the
200    /// caller just wrote into the KV cache. `k_block` device ptr to
201    /// `[block_size, num_kv_heads, head_dim]` BF16.
202    pub fn project_kv_block(
203        &self,
204        ctx: &CudaCtx,
205        layer: usize,
206        block_id: usize,
207        k_block: u64,
208    ) -> Result<()> {
209        self.project_kv_block_on_stream(ctx.stream, layer, block_id, k_block)
210    }
211
212    pub fn project_kv_block_on_stream(
213        &self,
214        stream: u64,
215        layer: usize,
216        block_id: usize,
217        k_block: u64,
218    ) -> Result<()> {
219        if layer >= self.dims.num_layers || block_id >= self.dims.max_blocks {
220            bail!("project_kv_block out of range: layer {layer}, block {block_id}");
221        }
222        let slot_floats = self.dims.per_layer_block_floats();
223        let k_lr_slot = self.a_g_dev.ptr
224            + (((layer * self.dims.max_blocks + block_id) * slot_floats) * 2) as u64;
225        let mut k_v = k_block;
226        let mut p_v = self.p_dev.ptr;
227        let mut o_v = k_lr_slot;
228        let mut bs = self.dims.block_size as i32;
229        let mut nk = self.dims.num_kv_heads as i32;
230        let mut hd = self.dims.head_dim as i32;
231        let mut r = self.dims.r as i32;
232        let mut params = [
233            &mut k_v as *mut _ as *mut c_void,
234            &mut p_v as *mut _ as *mut c_void,
235            &mut o_v as *mut _ as *mut c_void,
236            &mut bs as *mut _ as *mut c_void,
237            &mut nk as *mut _ as *mut c_void,
238            &mut hd as *mut _ as *mut c_void,
239            &mut r as *mut _ as *mut c_void,
240        ];
241        launch_kernel(
242            self.f_kv_proj,
243            (
244                self.dims.num_kv_heads as u32,
245                self.dims.block_size as u32,
246                1,
247            ),
248            (self.dims.r as u32, 1, 1),
249            0,
250            stream,
251            &mut params,
252        )
253    }
254
255    /// Score `num_active_blocks` (already-laid-out) blocks for the current
256    /// layer. `q_proj` device ptr to `[num_q_heads, r]` BF16. `a_g_seq` is a
257    /// device ptr to the active sequence's per-block anchors at this layer
258    /// (`[num_active_blocks, num_kv_heads, r]` BF16). `scores_out` is a
259    /// device ptr to a `[num_active_blocks]` f32 buffer.
260    pub fn score_blocks(
261        &self,
262        ctx: &CudaCtx,
263        q_proj: u64,
264        k_lr_seq: u64,
265        scores_out: u64,
266        num_active_blocks: usize,
267    ) -> Result<()> {
268        self.score_blocks_on_stream(ctx.stream, q_proj, k_lr_seq, scores_out, num_active_blocks)
269    }
270
271    pub fn score_blocks_on_stream(
272        &self,
273        stream: u64,
274        q_proj: u64,
275        k_lr_seq: u64,
276        scores_out: u64,
277        num_active_blocks: usize,
278    ) -> Result<()> {
279        let mut q_v = q_proj;
280        let mut a_v = k_lr_seq;
281        let mut s_v = scores_out;
282        let mut nq = self.dims.num_q_heads as i32;
283        let mut nk = self.dims.num_kv_heads as i32;
284        let mut bs = self.dims.block_size as i32;
285        let mut r = self.dims.r as i32;
286        let mut gqa = self.dims.gqa_ratio();
287        let mut params = [
288            &mut q_v as *mut _ as *mut c_void,
289            &mut a_v as *mut _ as *mut c_void,
290            &mut s_v as *mut _ as *mut c_void,
291            &mut nq as *mut _ as *mut c_void,
292            &mut nk as *mut _ as *mut c_void,
293            &mut bs as *mut _ as *mut c_void,
294            &mut r as *mut _ as *mut c_void,
295            &mut gqa as *mut _ as *mut c_void,
296        ];
297        launch_kernel(
298            self.f_score,
299            (num_active_blocks as u32, 1, 1),
300            (128, 1, 1),
301            0,
302            stream,
303            &mut params,
304        )
305    }
306
307    pub fn dims(&self) -> PredictorDims {
308        self.dims
309    }
310    pub fn a_g_dev_ptr(&self) -> u64 {
311        self.a_g_dev.ptr
312    }
313}
314
315/// Helper for tests / debug: copy the predictor's K_lr slot for a given
316/// (layer, block) to host as BF16. Layout `[num_kv_heads, block_size, r]`.
317pub fn read_k_lr_slot(
318    ctx: &CudaCtx,
319    pred: &Predictor,
320    layer: usize,
321    block: usize,
322) -> Result<Vec<bf16>> {
323    let dims = pred.dims();
324    let n = dims.per_layer_block_floats();
325    let slot = pred.a_g_dev.ptr + (((layer * dims.max_blocks + block) * n) * 2) as u64;
326    let mut host = vec![bf16::from_f32(0.0); n];
327    copy_d_to_h_async(host.as_mut_ptr() as *mut c_void, slot, n * 2, ctx.stream)?;
328    stream_sync(ctx.stream)?;
329    Ok(host)
330}