diff --git a/src/models/cua_s1/native/src/inputs.rs b/src/models/cua_s1/native/src/inputs.rs new file mode 100644 index 0000000..bf719ec --- /dev/null +++ b/src/models/cua_s1/native/src/inputs.rs @@ -0,0 +1,95 @@ +//! Batch-one, unpadded inputs at the adapted-vision / language-model boundary. + +use anyhow::{Result, ensure}; +use half::bf16; + +/// Image rows are already adapted to the language hidden size, in placeholder +/// order. Positions are the T/H/W slices of an int64 `[3, 1, sequence]` tensor. +/// The caller owns preprocessing, vision execution and the position calculation. +pub struct MultimodalInput<'a> { + pub token_ids: &'a [u32], + pub image_token_indices: &'a [usize], + pub image_embeddings: &'a [bf16], + pub position_ids: [&'a [i64]; 3], +} + +impl MultimodalInput<'_> { + /// Check the entire boundary before allocating buffers or launching CUDA. + pub fn validate( + &self, + hidden: usize, + vocab: usize, + image_token: u32, + max_position: usize, + ) -> Result<()> { + let t = self.token_ids.len(); + ensure!(t > 0 && t <= max_position, "empty or oversized prompt"); + ensure!( + self.token_ids.iter().all(|&id| (id as usize) < vocab), + "token id outside the vocabulary" + ); + let expected: Vec = self + .token_ids + .iter() + .enumerate() + .filter_map(|(i, &id)| (id == image_token).then_some(i)) + .collect(); + ensure!( + self.image_token_indices == expected, + "image indices must exactly match the ordered placeholders" + ); + ensure!( + Some(self.image_embeddings.len()) == expected.len().checked_mul(hidden), + "image embedding shape mismatch" + ); + ensure!( + self.image_embeddings.iter().all(|x| x.is_finite()), + "non-finite image embedding" + ); + ensure!( + self.position_ids.iter().all(|axis| axis.len() == t), + "position_ids must have shape [3, 1, sequence]" + ); + ensure!( + self.position_ids + .iter() + .flat_map(|axis| axis.iter()) + .all(|&p| p >= 0 && (p as u64) < max_position as u64), + "position outside the configured range" + ); + Ok(()) + } +} + +/// Qwen3.5's interleaved recomposition: overwrite H at 1::3 and W at 2::3 up +/// to section[axis] * 3, retaining T elsewhere. The second rotary half repeats +/// these frequencies, which the existing attention-prep kernel handles. +/// Float32 inverse frequencies/products and host float64 trig preserve the +/// original native text table rounding when all three axes are equal. +pub(crate) fn rotary_tables( + positions: [&[i64]; 3], + half: usize, + theta: f64, + sections: [usize; 3], +) -> (Vec, Vec) { + let inv: Vec = (0..half) + .map(|i| 1f32 / (theta as f32).powf((2 * i) as f32 / (2 * half) as f32)) + .collect(); + let mut cos = Vec::with_capacity(positions[0].len() * half * 2); + let mut sin = Vec::with_capacity(cos.capacity()); + for (t, _) in positions[0].iter().enumerate() { + for (i, &f) in inv.iter().enumerate() { + let axis = if i % 3 == 1 && i < sections[1] * 3 { + 1 + } else if i % 3 == 2 && i < sections[2] * 3 { + 2 + } else { + 0 + }; + let angle = (f * positions[axis][t] as f32) as f64; + cos.extend(bf16::from_f32(angle.cos() as f32).to_le_bytes()); + sin.extend(bf16::from_f32(angle.sin() as f32).to_le_bytes()); + } + } + (cos, sin) +} diff --git a/src/models/cua_s1/native/src/lib.rs b/src/models/cua_s1/native/src/lib.rs index 0d3aa5f..3d09d9c 100644 --- a/src/models/cua_s1/native/src/lib.rs +++ b/src/models/cua_s1/native/src/lib.rs @@ -5,5 +5,6 @@ pub mod contract; pub mod cuda; pub mod engine; +pub mod inputs; pub mod json; pub mod model; diff --git a/src/models/cua_s1/native/src/model.rs b/src/models/cua_s1/native/src/model.rs index ae14d55..ab4439a 100644 --- a/src/models/cua_s1/native/src/model.rs +++ b/src/models/cua_s1/native/src/model.rs @@ -6,8 +6,8 @@ //! The order of operations follows `modeling_qwen3_5.py`, and so do the points where //! it rounds to bfloat16, except inside attention and the Gated DeltaNet prefill (see //! their kernels). Text prompts use one position per -//! token, so the multimodal rotary sections all get the same position and the -//! rotary embedding is the plain one. +//! token. The explicit multimodal boundary inserts adapted image rows and supplies +//! the interleaved temporal/height/width rotary positions to the same layer loop. use std::collections::{HashMap, VecDeque}; use std::ffi::c_void; @@ -17,6 +17,7 @@ use anyhow::{Context, Result, bail, ensure}; use serde_json::Value as Json; use crate::cuda::{self, DeviceBuffer, Stream, check}; +use crate::inputs::{MultimodalInput, rotary_tables}; const ALIGN: usize = 256; const BF16: usize = 2; @@ -35,6 +36,9 @@ pub struct Config { /// Half the number of rotary dims (rotate_half pairs dim i with dim i + half). pub rotary_half: usize, pub rope_theta: f64, + pub mrope_section: [usize; 3], + pub max_positions: usize, + pub image_token_id: Option, pub lin_k_heads: usize, pub lin_v_heads: usize, pub lin_k_dim: usize, @@ -69,6 +73,31 @@ impl Config { other => bail!("unknown layer type {other:?}"), }) .collect::>>()?; + let sections = rope + .get("mrope_section") + .cloned() + .unwrap_or(serde_json::json!([11, 11, 10])); + let sections = sections + .as_array() + .context("mrope_section must be an array")? + .iter() + .map(|v| { + v.as_u64() + .and_then(|n| usize::try_from(n).ok()) + .context("mrope_section must contain integers") + }) + .collect::>>()?; + let mrope_section: [usize; 3] = sections + .try_into() + .map_err(|_| anyhow::anyhow!("mrope_section must have three entries"))?; + let image_token_id = root + .get("image_token_id") + .map(|v| { + v.as_u64() + .and_then(|id| u32::try_from(id).ok()) + .context("image_token_id must be a u32") + }) + .transpose()?; let cfg = Config { hidden: int("hidden_size")?, intermediate: int("intermediate_size")?, @@ -81,6 +110,9 @@ impl Config { .as_f64() .or(c["rope_theta"].as_f64()) .context("rope_theta")?, + mrope_section, + max_positions: int("max_position_embeddings")?, + image_token_id, lin_k_heads: int("linear_num_key_heads")?, lin_v_heads: int("linear_num_value_heads")?, lin_k_dim: int("linear_key_head_dim")?, @@ -115,6 +147,25 @@ impl Config { cfg.lin_v_dim ); ensure!(cfg.rotary_half == 32, "{} rotary dims", 2 * cfg.rotary_half); + ensure!( + cfg.rope_theta.is_finite() && cfg.rope_theta > 0.0, + "invalid rope_theta" + ); + ensure!( + rope["mrope_interleaved"].as_bool() != Some(false), + "non-interleaved mrope is unsupported" + ); + ensure!( + cfg.mrope_section + .iter() + .try_fold(0usize, |sum, &x| sum.checked_add(x)) + == Some(cfg.rotary_half), + "mrope sections must sum to rotary_half" + ); + ensure!( + cfg.max_positions > 0 && cfg.max_positions <= (1 << 24), + "unsupported max_position_embeddings" + ); ensure!( cfg.kv_heads > 0 && cfg.heads.is_multiple_of(cfg.kv_heads), "attention heads" @@ -218,7 +269,7 @@ impl Weights { .enumerate() .flat_map(|(i, st)| st.names().into_iter().map(move |n| (i, n.to_string()))) .collect(); - let prefix = ["model.language_model.", "model."] + let prefix = ["model.language_model.", "model.", ""] .into_iter() .find(|p| { names @@ -385,6 +436,8 @@ struct Scratch { act: usize, cos: usize, sin: usize, + custom_cos: usize, + custom_sin: usize, } impl Scratch { @@ -423,6 +476,8 @@ impl Scratch { take(cap * cfg.intermediate * BF16), take(cap * cfg.rotary_half * BF16), take(cap * cfg.rotary_half * BF16), + take(cap * cfg.rotary_half * BF16), + take(cap * cfg.rotary_half * BF16), ]; let buf = DeviceBuffer::new(next)?; let [ @@ -448,24 +503,16 @@ impl Scratch { act, cos, sin, + custom_cos, + custom_sin, ] = offsets; - // Rotary tables close to how Qwen3_5TextRotaryEmbedding builds them: inv_freq and - // freqs = inv_freq * position in float32, cos and sin rounded to bfloat16. Here - // cos and sin are taken in float64 on the host rather than in float32 on the - // GPU, so a few of the rounded values can differ by one bfloat16 step. - let half = cfg.rotary_half; - let inv: Vec = (0..half) - .map(|i| 1.0f32 / (cfg.rope_theta as f32).powf((2 * i) as f32 / (2 * half) as f32)) - .collect(); - let mut cos_t = Vec::with_capacity(cap * half * BF16); - let mut sin_t = Vec::with_capacity(cap * half * BF16); - for pos in 0..cap { - for &f in &inv { - let freq = (f * pos as f32) as f64; - cos_t.extend(half::bf16::from_f32(freq.cos() as f32).to_le_bytes()); - sin_t.extend(half::bf16::from_f32(freq.sin() as f32).to_le_bytes()); - } - } + let positions: Vec = (0..cap as i64).collect(); + let (cos_t, sin_t) = rotary_tables( + [&positions; 3], + cfg.rotary_half, + cfg.rope_theta, + cfg.mrope_section, + ); // SAFETY: both tables were laid out for cap * rotary_half bfloat16 values. unsafe { cuda::upload(buf.at(cos), &cos_t, stream)?; @@ -496,6 +543,8 @@ impl Scratch { act, cos, sin, + custom_cos, + custom_sin, }) } @@ -632,15 +681,7 @@ impl Model { ) } - /// The final-norm hidden state at the last position, as float32. - pub fn forward(&mut self, ids: &[u32]) -> Result> { - let t = ids.len(); - ensure!(t > 0, "empty prompt"); - let (vocab, h) = (self.embed.shape[0], self.cfg.hidden); - ensure!( - ids.iter().all(|&i| (i as usize) < vocab), - "token id outside the vocabulary" - ); + fn prepare_scratch(&mut self, t: usize) -> Result<()> { cuda::set_device(0)?; if self.scratch.as_ref().is_none_or(|s| t > s.cap) { self.graphs.clear(); @@ -651,16 +692,33 @@ impl Model { self.stream, )?); } + Ok(()) + } + + /// The final-norm hidden state at the last position, as float32. + pub fn forward(&mut self, ids: &[u32]) -> Result> { + let t = ids.len(); + ensure!( + t > 0 && t <= self.cfg.max_positions, + "empty or oversized prompt" + ); + ensure!( + ids.iter().all(|&i| (i as usize) < self.embed.shape[0]), + "token id outside the vocabulary" + ); + self.prepare_scratch(t)?; let s = self.scratch.as_ref().unwrap(); - let ids32: Vec = ids.iter().flat_map(|&i| (i as i32).to_le_bytes()).collect(); - // SAFETY: the ids buffer holds at least t int32 values. - unsafe { cuda::upload(s.at(s.ids), &ids32, self.stream)? }; + self.embed_tokens(s, ids)?; if self.graph_enabled { - if !self.graphs.iter().any(|(length, _)| *length == t) { - // Initialize every cuBLASLt plan before stream capture. - self.run(s, t)?; + if let Some((_, graph)) = self.graphs.iter().find(|(length, _)| *length == t) { + graph.launch(self.stream)?; + } else { + // Warm GEMM plans and keep this eager result for the cache miss. + // run() advances s.res in place and no longer embeds tokens, so + // launching the new graph here would advance the residual twice. + self.run(s, t, false)?; cuda::synchronize(self.stream)?; - match cuda::Graph::capture(self.stream, || self.run(s, t)) { + match cuda::Graph::capture(self.stream, || self.run(s, t, false)) { Ok(graph) => { if self.graphs.len() == 8 { self.graphs.pop_front(); @@ -669,27 +727,111 @@ impl Model { } Err(error) => { // Capture records without executing: the eager result is valid. - // Disable graphs for this worker rather than retrying failures. eprintln!("CUDA Graph capture failed; using eager execution: {error:#}"); self.graph_enabled = false; self.graphs.clear(); } } } - if self.graph_enabled { - self.graphs - .iter() - .find(|(length, _)| *length == t) - .unwrap() - .1 - .launch(self.stream)?; - } } else { - self.run(s, t)?; + self.run(s, t, false)?; } - let mut last = vec![0u8; h * BF16]; + self.last_hidden(s, t) + } + + /// Prefill one unpadded prompt with already-adapted BF16 image embeddings and + /// explicit `[3, 1, sequence]` T/H/W positions. No vision tower runs here. + /// Load a checkpoint with the matching multimodal language adapter merged. + pub fn forward_multimodal(&mut self, input: &MultimodalInput<'_>) -> Result> { + let image_token = self + .cfg + .image_token_id + .context("checkpoint has no image_token_id")?; + input.validate( + self.cfg.hidden, + self.embed.shape[0], + image_token, + self.cfg.max_positions, + )?; + let t = input.token_ids.len(); + self.prepare_scratch(t)?; + let s = self.scratch.as_ref().unwrap(); + self.upload_positions(s, input.position_ids)?; + self.embed_tokens(s, input.token_ids)?; + let bytes: Vec = input + .image_embeddings + .iter() + .flat_map(|x| x.to_le_bytes()) + .collect(); + // Coalesce adjacent placeholders. Text rows remain those of embed_tokens. + let indices = input.image_token_indices; + let mut begin = 0; + while begin < indices.len() { + let mut end = begin + 1; + while end < indices.len() && indices[end] == indices[end - 1] + 1 { + end += 1; + } + let row_bytes = self.cfg.hidden * BF16; + // SAFETY: validated indices lie in the t-row residual buffer, and + // features contain exactly one hidden-size BF16 row per placeholder. + unsafe { + cuda::upload( + s.at(s.res + indices[begin] * row_bytes), + &bytes[begin * row_bytes..end * row_bytes], + self.stream, + )?; + } + begin = end; + } + self.run(s, t, true)?; + self.last_hidden(s, t) + } + + fn upload_positions(&self, s: &Scratch, positions: [&[i64]; 3]) -> Result<()> { + let (cos, sin) = rotary_tables( + positions, + self.cfg.rotary_half, + self.cfg.rope_theta, + self.cfg.mrope_section, + ); + // SAFETY: tables contain at most s.cap rows of rotary_half BF16 values. + unsafe { + cuda::upload(s.at(s.custom_cos), &cos, self.stream)?; + cuda::upload(s.at(s.custom_sin), &sin, self.stream)?; + } + Ok(()) + } + + fn embed_tokens(&self, s: &Scratch, ids: &[u32]) -> Result<()> { + let ids32: Vec = ids.iter().flat_map(|&i| (i as i32).to_le_bytes()).collect(); + // SAFETY: IDs were checked against the vocabulary; scratch holds t rows. + unsafe { + cuda::upload(s.at(s.ids), &ids32, self.stream)?; + check( + (cuda::api().cs1_embed)( + s.at(s.ids).cast(), + self.embed.ptr, + s.at(s.res), + ids.len() as i32, + self.cfg.hidden as i32, + self.stream, + ), + "embed", + )?; + } + Ok(()) + } + + fn last_hidden(&self, s: &Scratch, t: usize) -> Result> { + let mut last = vec![0u8; self.cfg.hidden * BF16]; // SAFETY: x holds at least t rows of the hidden size. - unsafe { cuda::download(&mut last, s.at(s.x + (t - 1) * h * BF16), self.stream)? }; + unsafe { + cuda::download( + &mut last, + s.at(s.x + (t - 1) * self.cfg.hidden * BF16), + self.stream, + )?; + } let (pairs, _) = last.as_chunks::<2>(); Ok(pairs .iter() @@ -697,9 +839,10 @@ impl Model { .collect()) } - /// Queue one forward pass over the first `t` ids in `s`. The final-norm hidden - /// states end up in `s.x`. - fn run(&self, s: &Scratch, t: usize) -> Result<()> { + /// Queue language layers over prepared embeddings in `s.res`, with rotary + /// tables in immutable text buffers or separate explicit-position buffers. + /// Final-norm hidden states end up in `s.x`. + fn run(&self, s: &Scratch, t: usize, custom_positions: bool) -> Result<()> { let cfg = &self.cfg; let st = self.stream; let (ti, hi, eps) = (t as i32, cfg.hidden as i32, cfg.eps); @@ -707,13 +850,14 @@ impl Model { let (hq, hk, hd) = (cfg.heads as i32, cfg.kv_heads as i32, cfg.head_dim as i32); let w = Widths::of(cfg); let p = |off: usize| s.at(off); - // SAFETY (every kernel call below): the pointers are weights in the arena or + let (cos, sin) = if custom_positions { + (s.custom_cos, s.custom_sin) + } else { + (s.cos, s.sin) + }; + // SAFETY (every kernel call below): pointers are weights in the arena or // scratch buffers laid out for at least t tokens with the widths used here. unsafe { - check( - (cuda::api().cs1_embed)(p(s.ids).cast(), self.embed.ptr, p(s.res), ti, hi, st), - "embed", - )?; check( (cuda::api().cs1_rms_norm)( p(s.res), @@ -814,8 +958,8 @@ impl Model { ld, fa.q_norm.ptr, fa.k_norm.ptr, - p(s.cos), - p(s.sin), + p(cos), + p(sin), p(s.aq), p(s.agate), p(s.ak),