memra_engine/
prime_graph.rs1use crate::cache::Cache;
19use crate::hybrid::HybridModel;
20use crate::Engine;
21use cudarc::driver::{CudaGraph, CudaSlice};
22
23pub struct PrimeGraph {
24 pub bucket: usize,
25 graph: CudaGraph,
26 _keeper: Vec<Box<dyn std::any::Any + Send>>,
31 #[allow(dead_code)]
37 private_scratch: Option<crate::f16_ffi::F16Scratch>,
38 scratch: Cache,
39 x_in: CudaSlice<f32>,
40 len_d: CudaSlice<i32>,
41 logits_out: CudaSlice<f32>,
42 h_seed_out: CudaSlice<f32>,
43 n_embd: usize,
44}
45
46impl PrimeGraph {
47 pub fn scratch(&self) -> &Cache {
49 &self.scratch
50 }
51}
52
53impl HybridModel {
54 pub fn prime_graph_new(&self, e: &Engine, bucket: usize)
57 -> Result<PrimeGraph, Box<dyn std::error::Error>> {
58 use cudarc::driver::sys::{CUgraphInstantiate_flags, CUstreamCaptureMode};
59 let n_embd = self.cfg.n_embd as usize;
60 let n_vocab = self.output.out_features();
61 let mut scratch = Cache::new(e, &self.cfg, bucket + 8)?;
62 let x_in = e.zeros(bucket * n_embd)?;
63 let pos_d = e.htod_i32(&(0..bucket as i32).collect::<Vec<_>>())?;
64 let len_d = e.htod_i32(&[bucket as i32])?;
65 let mut logits_out = e.uninit(n_vocab)?;
66 let mut h_seed_out = e.uninit(n_embd)?;
67
68 let n_ff_max = self.layers.iter().map(|l| match &l.ffn {
72 crate::hybrid::Ffn::Dense { ffn_gate, .. } => ffn_gate.out_features(),
73 _ => n_embd,
74 }).max().unwrap_or(n_embd).max(n_embd);
75 let private = crate::f16_ffi::F16Scratch::with_capacity(e, bucket * n_ff_max * 2)?;
76 let prev_scratch = e.f16_scratch_swap(Some(private));
77 let scratch_cell = std::cell::RefCell::new(&mut scratch);
78 let lo_cell = std::cell::RefCell::new(&mut logits_out);
79 let hs_cell = std::cell::RefCell::new(&mut h_seed_out);
80 let (graph, keeper) = e.capture_graph_retained(|e| {
81 let sc: &mut Cache = &mut scratch_cell.borrow_mut();
82 for kvl in sc.kv.iter_mut().flatten() {
83 kvl.len = 0;
84 e.stream().memset_zeros(&mut kvl.len_d)?;
85 }
86 for rl in sc.recur.iter_mut().flatten() {
87 e.stream().memset_zeros(&mut rl.conv_state)?;
88 e.stream().memset_zeros(&mut rl.ssm_state)?;
89 e.stream().memset_zeros(&mut rl.ssm_state_alt)?;
90 }
91 self.prime_chunk_captured(e, &x_in, &pos_d, bucket, sc, &len_d,
92 &mut lo_cell.borrow_mut(), &mut hs_cell.borrow_mut())
93 })?;
94 drop(scratch_cell);
95 drop(lo_cell);
96 drop(hs_cell);
97 let private_scratch = e.f16_scratch_swap(prev_scratch);
99 let _ = CUstreamCaptureMode::CU_STREAM_CAPTURE_MODE_RELAXED;
100 let _ = CUgraphInstantiate_flags::CUDA_GRAPH_INSTANTIATE_FLAG_AUTO_FREE_ON_LAUNCH;
101 Ok(PrimeGraph {
102 bucket,
103 graph,
104 _keeper: keeper,
105 private_scratch,
106 scratch,
107 x_in,
108 len_d,
109 logits_out,
110 h_seed_out,
111 n_embd,
112 })
113 }
114
115 pub fn prime_graph_run(&self, e: &Engine, pg: &mut PrimeGraph, tokens: &[u32],
118 session: &mut Cache)
119 -> Result<(Vec<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
120 let t = tokens.len();
121 assert!(t >= 2 && t <= pg.bucket, "prime_graph_run: 2 <= T <= bucket");
122 assert!(session.pos == 0, "prime_graph_run: fresh sessions only");
123 let n_embd = pg.n_embd;
124 let x = self.embed(e, tokens)?;
127 e.copy_into(&mut pg.x_in, 0, &x, t * n_embd)?;
128 if t < pg.bucket {
129 let mut tail = pg.x_in.slice_mut(t * n_embd..pg.bucket * n_embd);
130 e.stream().memset_zeros(&mut tail)?;
131 }
132 e.set_i32_one(&mut pg.len_d, t as i32)?;
133 pg.graph.launch()?;
134 for (il, kvl) in pg.scratch.kv.iter().enumerate() {
136 let (Some(src), Some(dst)) = (kvl.as_ref(), session.kv[il].as_mut()) else { continue };
137 let kb = t * src.k_tok_bytes;
138 let vb = t * src.v_tok_bytes;
139 e.stream().memcpy_dtod(&src.k.slice(0..kb), &mut dst.k.slice_mut(0..kb))?;
140 e.stream().memcpy_dtod(&src.v.slice(0..vb), &mut dst.v.slice_mut(0..vb))?;
141 dst.len = t;
142 e.set_i32_one(&mut dst.len_d, t as i32)?;
143 }
144 for (il, rl) in pg.scratch.recur.iter().enumerate() {
145 let (Some(src), Some(dst)) = (rl.as_ref(), session.recur[il].as_mut()) else { continue };
146 let cn = src.conv_state.len();
147 let sn = src.ssm_state.len();
148 e.copy_into(&mut dst.conv_state, 0, &src.conv_state, cn)?;
149 e.copy_into(&mut dst.ssm_state, 0, &src.ssm_state, sn)?;
150 }
151 session.pos = t;
152 let logits = e.dtoh(&pg.logits_out)?;
153 let mut h_seed = e.uninit(n_embd)?;
154 let hn = pg.h_seed_out.len();
155 e.copy_into(&mut h_seed, 0, &pg.h_seed_out, hn)?;
156 Ok((logits, h_seed))
157 }
158}