Skip to content
Open
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
79 changes: 63 additions & 16 deletions src/models/qwen3_5/native/src/model.rs
Original file line number Diff line number Diff line change
Expand Up @@ -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<Scratch>,
/// 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<usize>, cuda::Graph)>,
mm_graphs: VecDeque<(Vec<usize>, cuda::Graph)>,
}

// SAFETY: the raw pointers are device addresses and a cuBLASLt handle owned by the
Expand All @@ -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) };
}
Expand Down Expand Up @@ -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)
}
Expand Down Expand Up @@ -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,
Expand Down Expand Up @@ -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()
Expand Down Expand Up @@ -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<cuda::Graph>) {
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<()> {
Expand Down
Loading