1use 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 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, a_g_dev: DeviceBuffer, }
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 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 _ => {} }
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 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 let a_g_need = dims.a_g_bytes();
124 let (free_hbm, _total_hbm) = crate::cuda_min::mem_info()?;
125 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 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 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 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 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 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
315pub 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}