1use crate::backend::ExecutableGraph;
19use rlx_driver::Device;
20use std::collections::HashSet;
21
22struct Staging {
30 prepare: Box<CompiledGraph>,
31 prepared: bool,
32 boundary: Vec<String>,
34 prepare_params: HashSet<String>,
35 main_params: HashSet<String>,
36 bound: bool,
40 values: Vec<Vec<f32>>,
42}
43
44pub struct CompiledGraph {
51 inner: Box<dyn ExecutableGraph>,
52 device: Device,
53 staging: Option<Box<Staging>>,
55}
56
57impl Clone for CompiledGraph {
58 fn clone(&self) -> Self {
61 Self {
62 inner: self.inner.clone_box(),
63 device: self.device,
64 staging: self.staging.as_ref().map(|s| {
65 Box::new(Staging {
66 prepare: s.prepare.clone(),
67 prepared: s.prepared,
68 boundary: s.boundary.clone(),
69 prepare_params: s.prepare_params.clone(),
70 main_params: s.main_params.clone(),
71 bound: s.bound,
72 values: s.values.clone(),
73 })
74 }),
75 }
76 }
77}
78
79impl CompiledGraph {
80 pub(crate) fn new(inner: Box<dyn ExecutableGraph>, device: Device) -> Self {
81 Self {
82 inner,
83 device,
84 staging: None,
85 }
86 }
87
88 pub(crate) fn with_staging(
91 mut self,
92 prepare: CompiledGraph,
93 boundary: Vec<String>,
94 prepare_params: HashSet<String>,
95 main_params: HashSet<String>,
96 ) -> Self {
97 self.staging = Some(Box::new(Staging {
98 prepare: Box::new(prepare),
99 prepared: false,
100 boundary,
101 prepare_params,
102 main_params,
103 bound: false,
104 values: Vec::new(),
105 }));
106 self
107 }
108
109 fn ensure_prepared(&mut self) {
113 let Some(st) = self.staging.as_mut() else {
114 return;
115 };
116 if st.prepared {
117 return;
118 }
119 st.prepare.finalize_params();
120 let outs = st.prepare.run(&[]);
121 let mut all_bound = !outs.is_empty();
123 for (name, out) in st.boundary.iter().zip(outs.iter()) {
124 if !self.inner.bind_handle(name, out) {
125 all_bound = false;
126 }
127 }
128 st.bound = all_bound;
129 st.values = if all_bound { Vec::new() } else { outs };
130 st.prepared = true;
131 }
132
133 pub fn device(&self) -> Device {
135 self.device
136 }
137
138 pub fn set_param(&mut self, name: &str, data: &[f32]) {
141 if let Some(st) = self.staging.as_mut() {
142 if st.prepare_params.contains(name) {
143 st.prepare.set_param(name, data);
144 st.prepared = false; }
146 if st.main_params.contains(name) {
147 self.inner.set_param(name, data);
148 }
149 } else {
150 self.inner.set_param(name, data);
151 }
152 }
153
154 pub fn run(&mut self, inputs: &[(&str, &[f32])]) -> Vec<Vec<f32>> {
157 self.ensure_prepared();
158 match self.staging.as_ref() {
159 Some(st) if !st.bound => {
160 let mut merged: Vec<(&str, &[f32])> =
161 Vec::with_capacity(inputs.len() + st.boundary.len());
162 merged.extend_from_slice(inputs);
163 for (name, vals) in st.boundary.iter().zip(st.values.iter()) {
164 merged.push((name.as_str(), vals.as_slice()));
165 }
166 self.inner.run(&merged)
167 }
168 _ => self.inner.run(inputs),
169 }
170 }
171
172 #[cfg(target_arch = "wasm32")]
175 pub async fn run_async(&mut self, inputs: &[(&str, &[f32])]) -> Result<Vec<Vec<f32>>, String> {
176 match self.device {
177 Device::WebGpu | Device::Gpu => self.inner.wgpu_run_async(inputs).await,
178 _ => Ok(self.run(inputs)),
179 }
180 }
181
182 pub fn run_read_outputs(
184 &mut self,
185 inputs: &[(&str, &[f32])],
186 read_indices: Option<&[usize]>,
187 ) -> Vec<Vec<f32>> {
188 self.ensure_prepared();
189 match self.staging.as_ref() {
190 Some(st) if !st.bound => {
191 let mut merged: Vec<(&str, &[f32])> =
192 Vec::with_capacity(inputs.len() + st.boundary.len());
193 merged.extend_from_slice(inputs);
194 for (name, vals) in st.boundary.iter().zip(st.values.iter()) {
195 merged.push((name.as_str(), vals.as_slice()));
196 }
197 self.inner.run_read_outputs(&merged, read_indices)
198 }
199 _ => self.inner.run_read_outputs(inputs, read_indices),
200 }
201 }
202
203 pub fn read_output_row(
205 &self,
206 out_idx: usize,
207 row: usize,
208 row_inner: usize,
209 ) -> Option<Vec<f32>> {
210 self.inner.read_output_row(out_idx, row, row_inner)
211 }
212
213 pub fn run_raw(&mut self, inputs: &[(&str, &[f32])]) -> Vec<(*const f32, usize)> {
220 self.ensure_prepared();
221 self.inner.run_raw(inputs)
222 }
223
224 pub fn run_slots(&mut self, inputs: &[&[f32]]) -> &[(usize, usize)] {
228 self.ensure_prepared();
229 self.inner.run_slots(inputs)
230 }
231
232 pub fn arena_ptr(&self) -> *const u8 {
234 self.inner.arena_ptr()
235 }
236
237 pub fn bind_handle(&mut self, name: &str, data: &[f32]) -> bool {
242 self.inner.bind_handle(name, data)
243 }
244
245 pub fn read_handle(&self, name: &str) -> Option<Vec<f32>> {
247 self.inner.read_handle(name)
248 }
249
250 pub fn bind_gpu_handle(&mut self, name: &str, data: &[f32]) -> bool {
252 self.inner.bind_gpu_handle(name, data)
253 }
254
255 pub fn has_gpu_handle(&self, name: &str) -> bool {
256 self.inner.has_gpu_handle(name)
257 }
258
259 pub fn set_gpu_handle_feed(&mut self, handle_name: &str, output_index: usize) -> bool {
260 self.inner.set_gpu_handle_feed(handle_name, output_index)
261 }
262
263 pub fn read_gpu_handle(&self, name: &str) -> Option<Vec<f32>> {
264 self.inner.read_gpu_handle(name)
265 }
266
267 pub fn read_gpu_handle_row(
269 &self,
270 name: &str,
271 row: usize,
272 row_inner: usize,
273 ) -> Option<Vec<f32>> {
274 self.inner.read_gpu_handle_row(name, row, row_inner)
275 }
276
277 pub fn register_kv_row_feed(&mut self, handle_name: &str, output_index: usize) -> bool {
281 self.inner.register_kv_row_feed(handle_name, output_index)
282 }
283
284 pub fn feed_kv_row(&mut self, src_row: usize, dst_row: usize, row_elems: usize) -> bool {
288 self.inner.feed_kv_row(src_row, dst_row, row_elems)
289 }
290
291 pub fn feed_kv_batch_major(
293 &mut self,
294 dst_row: usize,
295 batch: usize,
296 seq_cap: usize,
297 row_elems: usize,
298 ) -> bool {
299 self.inner
300 .feed_kv_batch_major(dst_row, batch, seq_cap, row_elems)
301 }
302
303 pub fn prepare_resident_gpu_handle(&mut self, name: &str) -> bool {
305 self.inner.prepare_resident_gpu_handle(name)
306 }
307
308 pub fn stage_bound_gpu_handles_to_arena(&mut self) -> bool {
310 self.inner.stage_bound_gpu_handles_to_arena();
311 true
312 }
313
314 pub fn seed_resident_kv_prefix_from(
316 &mut self,
317 src: &CompiledGraph,
318 prefix_tokens: usize,
319 outgoing_upper: usize,
320 kv_dim: usize,
321 n_layers: usize,
322 ) -> bool {
323 if self.device != src.device {
324 return false;
325 }
326 self.inner.seed_resident_kv_prefix_from(
327 src.inner.as_ref(),
328 prefix_tokens,
329 outgoing_upper,
330 kv_dim,
331 n_layers,
332 )
333 }
334
335 pub fn copy_resident_kv_rows_from(
337 &mut self,
338 src: &CompiledGraph,
339 from_row: usize,
340 to_row: usize,
341 outgoing_upper: usize,
342 kv_dim: usize,
343 n_layers: usize,
344 ) -> bool {
345 if self.device != src.device {
346 return false;
347 }
348 self.inner.copy_resident_kv_rows_from(
349 src.inner.as_ref(),
350 from_row,
351 to_row,
352 outgoing_upper,
353 kv_dim,
354 n_layers,
355 )
356 }
357
358 pub fn copy_params_from(&mut self, src: &CompiledGraph) -> bool {
360 if self.device != src.device {
361 return false;
362 }
363 self.inner.copy_params_from(src.inner.as_ref())
364 }
365
366 pub fn share_params_from(&mut self, src: &CompiledGraph) -> bool {
370 if self.device != src.device {
371 return false;
372 }
373 self.inner.share_params_from(src.inner.as_ref())
374 }
375
376 pub fn run_feed_gpu_handle(
378 &mut self,
379 inputs: &[(&str, &[f32])],
380 handle_name: &str,
381 output_index: usize,
382 ) -> Option<Vec<f32>> {
383 self.inner
384 .run_feed_gpu_handle(inputs, handle_name, output_index)
385 }
386
387 pub fn set_active_extent(&mut self, extent: Option<(usize, usize)>) {
394 #[cfg(feature = "cpu")]
395 if let Some((actual, _)) = extent {
396 crate::onnx_active::set_active_token_count(Some(actual))
397 }
398 self.inner.set_active_extent(extent);
399 }
400
401 pub fn set_moe_resident_experts(&mut self, mask: &[bool]) {
403 self.inner.set_moe_resident_experts(mask);
404 }
405
406 pub fn set_moe_resident_experts_per_layer(&mut self, masks: &[&[bool]]) {
408 self.inner.set_moe_resident_experts_per_layer(masks);
409 }
410
411 pub fn enable_moe_topk_capture(&mut self, num_experts: usize) -> bool {
413 self.inner.enable_moe_topk_capture(num_experts)
414 }
415
416 pub fn take_moe_topk_capture(&mut self) -> Option<Vec<Vec<u32>>> {
418 self.inner.take_moe_topk_capture()
419 }
420
421 pub fn take_moe_residency_stats(&mut self) -> Option<crate::MoeResidencyStats> {
423 self.inner.take_moe_residency_stats()
424 }
425
426 pub fn commit_no_wait(&mut self, inputs: &[(&str, &[f32])]) {
434 self.inner.commit_no_wait(inputs);
435 }
436
437 pub fn sync_pending(&mut self) {
439 self.inner.sync_pending();
440 }
441
442 pub fn run_pipelined(&mut self, input_sets: &[Vec<(&str, &[f32])>]) -> Vec<Vec<Vec<f32>>> {
448 self.inner.run_pipelined(input_sets)
449 }
450
451 pub fn set_param_typed(&mut self, name: &str, data: &[u8], dtype: rlx_ir::DType) {
456 if let Some(st) = self.staging.as_mut() {
457 if st.prepare_params.contains(name) {
458 st.prepare.set_param_typed(name, data, dtype);
459 st.prepared = false;
460 }
461 if st.main_params.contains(name) {
462 self.inner.set_param_typed(name, data, dtype);
463 }
464 } else {
465 self.inner.set_param_typed(name, data, dtype);
466 }
467 }
468
469 pub fn finalize_params(&mut self) {
471 self.ensure_prepared();
472 self.inner.finalize_params();
473 }
474
475 pub fn run_typed(
480 &mut self,
481 inputs: &[(&str, &[u8], rlx_ir::DType)],
482 ) -> Vec<(Vec<u8>, rlx_ir::DType)> {
483 self.ensure_prepared();
484 self.inner.run_typed(inputs)
485 }
486
487 pub fn set_rng(&mut self, rng: rlx_ir::RngOptions) {
489 self.inner.set_rng(rng);
490 }
491
492 pub fn rng(&self) -> rlx_ir::RngOptions {
494 self.inner.rng()
495 }
496}
497
498#[cfg(test)]
499mod tests {
500 use crate::*;
501
502 #[test]
503 #[cfg(feature = "cpu")]
504 fn end_to_end_session() {
505 let mut g = Graph::new("matmul_bias_gelu");
506 let x = g.input("x", Shape::new(&[2, 4], DType::F32));
507 let w = g.param("w", Shape::new(&[4, 3], DType::F32));
508 let b = g.param("b", Shape::new(&[3], DType::F32));
509 let mm = g.matmul(x, w, Shape::new(&[2, 3], DType::F32));
510 let add = g.binary(op::BinaryOp::Add, mm, b, Shape::new(&[2, 3], DType::F32));
511 let out = g.activation(op::Activation::Gelu, add, Shape::new(&[2, 3], DType::F32));
512 g.set_outputs(vec![out]);
513
514 let session = Session::new(Device::Cpu);
516 let mut compiled = session.compile(g);
517
518 compiled.set_param(
521 "w",
522 &[1.0, 0.0, 0.0, 0.0, 1.0, 0.0, 0.0, 0.0, 1.0, 0.0, 0.0, 0.0],
523 );
524 compiled.set_param("b", &[0.5, -0.5, 0.0]);
525
526 let x_data = vec![
528 1.0, 0.0, 0.0, 0.0, 0.0, 1.0, 0.0, 0.0, ];
531 let outputs = compiled.run(&[("x", &x_data)]);
532
533 assert_eq!(outputs.len(), 1);
534 let result = &outputs[0];
535 assert_eq!(result.len(), 6); assert!(
539 (result[0] - 1.399).abs() < 0.01,
540 "gelu(1.5) = {}",
541 result[0]
542 );
543 assert!(
544 (result[1] - -0.154).abs() < 0.01,
545 "gelu(-0.5) = {}",
546 result[1]
547 );
548 assert!((result[2]).abs() < 0.01, "gelu(0) = {}", result[2]);
549
550 assert!(
552 (result[3] - 0.346).abs() < 0.01,
553 "gelu(0.5) = {}",
554 result[3]
555 );
556 assert!(
557 (result[4] - 0.346).abs() < 0.01,
558 "gelu(0.5) = {}",
559 result[4]
560 );
561
562 let x2 = vec![0.0f32; 8];
564 let outputs2 = compiled.run(&[("x", &x2)]);
565 let r2 = &outputs2[0];
567 assert!((r2[0] - 0.346).abs() < 0.01, "gelu(0.5) = {}", r2[0]); }
569
570 #[test]
571 #[cfg(feature = "cpu")]
572 fn device_display() {
573 use crate::device_ext::is_available;
574 assert!(format!("{}", Device::Cpu).starts_with("CPU"));
575 assert!(is_available(Device::Cpu));
576 #[cfg(not(feature = "gpu"))]
579 assert!(!is_available(Device::Gpu));
580 #[cfg(not(feature = "cuda"))]
581 assert!(!is_available(Device::Cuda));
582 #[cfg(not(feature = "rocm"))]
583 assert!(!is_available(Device::Rocm));
584 #[cfg(not(feature = "tpu"))]
585 assert!(!is_available(Device::Tpu));
586 }
587}