memra_engine/
prime_graph.rs1use crate::Engine;
19use crate::cache::Cache;
20use crate::hybrid::HybridModel;
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(
57 &self,
58 e: &Engine,
59 bucket: usize,
60 ) -> Result<PrimeGraph, Box<dyn std::error::Error>> {
61 use cudarc::driver::sys::{CUgraphInstantiate_flags, CUstreamCaptureMode};
62 let n_embd = self.cfg.n_embd as usize;
63 let n_vocab = self.output.out_features();
64 let mut scratch = Cache::new(e, &self.cfg, bucket + 8)?;
65 let x_in = e.zeros(bucket * n_embd)?;
66 let pos_d = e.htod_i32(&(0..bucket as i32).collect::<Vec<_>>())?;
67 let len_d = e.htod_i32(&[bucket as i32])?;
68 let mut logits_out = e.uninit(n_vocab)?;
69 let mut h_seed_out = e.uninit(n_embd)?;
70
71 let n_ff_max = self
75 .layers
76 .iter()
77 .map(|l| match &l.ffn {
78 crate::hybrid::Ffn::Dense { ffn_gate, .. } => ffn_gate.out_features(),
79 _ => n_embd,
80 })
81 .max()
82 .unwrap_or(n_embd)
83 .max(n_embd);
84 let private = crate::f16_ffi::F16Scratch::with_capacity(e, bucket * n_ff_max * 2)?;
85 let prev_scratch = e.f16_scratch_swap(Some(private));
86 let scratch_cell = std::cell::RefCell::new(&mut scratch);
87 let lo_cell = std::cell::RefCell::new(&mut logits_out);
88 let hs_cell = std::cell::RefCell::new(&mut h_seed_out);
89 let (graph, keeper) = e.capture_graph_retained(|e| {
90 let sc: &mut Cache = &mut scratch_cell.borrow_mut();
91 for kvl in sc.kv.iter_mut().flatten() {
92 kvl.len = 0;
93 e.stream().memset_zeros(&mut kvl.len_d)?;
94 }
95 for rl in sc.recur.iter_mut().flatten() {
96 e.stream().memset_zeros(&mut rl.conv_state)?;
97 e.stream().memset_zeros(&mut rl.ssm_state)?;
98 e.stream().memset_zeros(&mut rl.ssm_state_alt)?;
99 }
100 self.prime_chunk_captured(
101 e,
102 &x_in,
103 &pos_d,
104 bucket,
105 sc,
106 &len_d,
107 &mut lo_cell.borrow_mut(),
108 &mut hs_cell.borrow_mut(),
109 )
110 })?;
111 drop(scratch_cell);
112 drop(lo_cell);
113 drop(hs_cell);
114 let private_scratch = e.f16_scratch_swap(prev_scratch);
116 let _ = CUstreamCaptureMode::CU_STREAM_CAPTURE_MODE_RELAXED;
117 let _ = CUgraphInstantiate_flags::CUDA_GRAPH_INSTANTIATE_FLAG_AUTO_FREE_ON_LAUNCH;
118 Ok(PrimeGraph {
119 bucket,
120 graph,
121 _keeper: keeper,
122 private_scratch,
123 scratch,
124 x_in,
125 len_d,
126 logits_out,
127 h_seed_out,
128 n_embd,
129 })
130 }
131
132 pub fn prime_graph_run(
135 &self,
136 e: &Engine,
137 pg: &mut PrimeGraph,
138 tokens: &[u32],
139 session: &mut Cache,
140 ) -> Result<(Vec<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
141 let t = tokens.len();
142 assert!(
143 t >= 2 && t <= pg.bucket,
144 "prime_graph_run: 2 <= T <= bucket"
145 );
146 assert!(session.pos == 0, "prime_graph_run: fresh sessions only");
147 let n_embd = pg.n_embd;
148 let x = self.embed(e, tokens)?;
151 e.copy_into(&mut pg.x_in, 0, &x, t * n_embd)?;
152 if t < pg.bucket {
153 let mut tail = pg.x_in.slice_mut(t * n_embd..pg.bucket * n_embd);
154 e.stream().memset_zeros(&mut tail)?;
155 }
156 e.set_i32_one(&mut pg.len_d, t as i32)?;
157 pg.graph.launch()?;
158 for (il, kvl) in pg.scratch.kv.iter().enumerate() {
160 let (Some(src), Some(dst)) = (kvl.as_ref(), session.kv[il].as_mut()) else {
161 continue;
162 };
163 let kb = t * src.k_tok_bytes;
164 let vb = t * src.v_tok_bytes;
165 e.stream()
166 .memcpy_dtod(&src.k.slice(0..kb), &mut dst.k.slice_mut(0..kb))?;
167 e.stream()
168 .memcpy_dtod(&src.v.slice(0..vb), &mut dst.v.slice_mut(0..vb))?;
169 dst.len = t;
170 e.set_i32_one(&mut dst.len_d, t as i32)?;
171 }
172 for (il, rl) in pg.scratch.recur.iter().enumerate() {
173 let (Some(src), Some(dst)) = (rl.as_ref(), session.recur[il].as_mut()) else {
174 continue;
175 };
176 let cn = src.conv_state.len();
177 let sn = src.ssm_state.len();
178 e.copy_into(&mut dst.conv_state, 0, &src.conv_state, cn)?;
179 e.copy_into(&mut dst.ssm_state, 0, &src.ssm_state, sn)?;
180 }
181 session.pos = t;
182 let logits = e.dtoh(&pg.logits_out)?;
183 let mut h_seed = e.uninit(n_embd)?;
184 let hn = pg.h_seed_out.len();
185 e.copy_into(&mut h_seed, 0, &pg.h_seed_out, hn)?;
186 Ok((logits, h_seed))
187 }
188}