diff --git a/src/models/qwen3_5/native/src/model.rs b/src/models/qwen3_5/native/src/model.rs index 6d37cd2f..b243f229 100644 --- a/src/models/qwen3_5/native/src/model.rs +++ b/src/models/qwen3_5/native/src/model.rs @@ -563,9 +563,10 @@ pub struct Model { gemm: *mut c_void, /// Buffers for the largest packed token count so far; grows as needed. scratch: Option, - /// Opt-in replay with at most 64 captures keyed by ordered sequence lengths. + /// Opt-in replay with at most 64 captures per mode, keyed by sequence lengths. graph_enabled: bool, graphs: VecDeque<(Vec, cuda::Graph)>, + mm_graphs: VecDeque<(Vec, cuda::Graph)>, } // SAFETY: the raw pointers are device addresses and a cuBLASLt handle owned by the @@ -575,6 +576,7 @@ unsafe impl Send for Model {} impl Drop for Model { fn drop(&mut self) { self.graphs.clear(); + self.mm_graphs.clear(); // SAFETY: created by cs1_gemm_create and not destroyed before. unsafe { (cuda::api().cs1_gemm_destroy)(self.gemm) }; } @@ -656,6 +658,7 @@ impl Model { scratch: None, graph_enabled: std::env::var("CUA_S1_GRAPH").as_deref() == Ok("1"), graphs: VecDeque::new(), + mm_graphs: VecDeque::new(), }; Ok(model) } @@ -709,6 +712,7 @@ impl Model { cuda::set_device(0)?; if self.scratch.as_ref().is_none_or(|s| t > s.cap) { self.graphs.clear(); + self.mm_graphs.clear(); self.scratch = None; self.scratch = Some(Scratch::new( &self.cfg, @@ -757,28 +761,21 @@ impl Model { if self.graph_enabled { if let Some((_, graph)) = self.graphs.iter().find(|(shape, _)| *shape == lengths) { graph.launch(self.stream)?; + if std::env::var("CUA_S1_GRAPH_TRACE").as_deref() == Ok("1") { + eprintln!("CUA_S1_GRAPH mode=text event=replay tokens={t}"); + } } else { // Warm plans and keep the eager result: run() advances the residual // in place, so replaying on this cache miss would advance it twice. self.run(s, &lengths, false)?; cuda::synchronize(self.stream)?; - match cuda::Graph::capture(self.stream, || self.run(s, &lengths, false)) { - Ok(graph) => { - if self.graphs.len() == 64 { - self.graphs.pop_front(); - } - self.graphs.push_back((lengths.clone(), graph)); - } - Err(error) => { - eprintln!("CUDA Graph capture failed; using eager execution: {error:#}"); - self.graph_enabled = false; - self.graphs.clear(); - } - } + let captured = cuda::Graph::capture(self.stream, || self.run(s, &lengths, false)); + self.cache_graph(&lengths, false, captured); } } else { self.run(s, &lengths, false)?; } + let s = self.scratch.as_ref().unwrap(); let mut end = 0; lengths .iter() @@ -833,8 +830,58 @@ impl Model { } begin = end; } - self.run(s, &[t], true)?; - self.last_hidden(s, t) + if self.graph_enabled { + if let Some((_, graph)) = self + .mm_graphs + .iter() + .find(|(shape, _)| shape.as_slice() == [t]) + { + graph.launch(self.stream)?; + if std::env::var("CUA_S1_GRAPH_TRACE").as_deref() == Ok("1") { + eprintln!("CUA_S1_GRAPH mode=multimodal event=replay tokens={t}"); + } + } else { + // Uploads and embedding are outside capture. Keep the warmed eager + // result; launching now would advance the residual a second time. + self.run(s, &[t], true)?; + cuda::synchronize(self.stream)?; + let captured = cuda::Graph::capture(self.stream, || self.run(s, &[t], true)); + self.cache_graph(&[t], true, captured); + } + } else { + self.run(s, &[t], true)?; + } + self.last_hidden(self.scratch.as_ref().unwrap(), t) + } + + /// Both modes retain the eager miss result; capture records without executing. + fn cache_graph(&mut self, lengths: &[usize], multimodal: bool, captured: Result) { + let mode = if multimodal { "multimodal" } else { "text" }; + match captured { + Ok(graph) => { + let cache = if multimodal { + &mut self.mm_graphs + } else { + &mut self.graphs + }; + if cache.len() == 64 { + cache.pop_front(); + } + cache.push_back((lengths.to_vec(), graph)); + if std::env::var("CUA_S1_GRAPH_TRACE").as_deref() == Ok("1") { + let t: usize = lengths.iter().sum(); + eprintln!("CUA_S1_GRAPH mode={mode} event=capture tokens={t}"); + } + } + Err(error) => { + eprintln!( + "CUDA Graph capture failed in {mode} mode; using eager execution: {error:#}" + ); + self.graph_enabled = false; + self.graphs.clear(); + self.mm_graphs.clear(); + } + } } fn upload_positions(&self, s: &Scratch, positions: [&[i64]; 3]) -> Result<()> {