use crate::Engine;
use crate::cache::{Cache, KvLayer};
use crate::forward::argmax;
use crate::hybrid::{FullAttnLayer, HybridModel, LinearAttnLayer, Mixer, MtpHead};
use cudarc::driver::CudaSlice;
use std::sync::atomic::{AtomicU64, Ordering};
pub fn spec_replay_env_on(value: Option<&str>) -> bool {
value == Some("1")
}
pub fn spec_replay_env_enabled() -> bool {
let value = std::env::var("MEMRA_SPEC_REPLAY").ok();
spec_replay_env_on(value.as_deref())
}
pub struct DsparkAnchorRecord {
pub position: usize,
pub hidden: Vec<f32>,
pub tokens: Vec<u32>,
pub target_top_ids: Vec<u32>,
pub target_top_logits: Vec<f32>,
pub target_top_probs: Vec<f32>,
pub target_tail_probs: Vec<f32>,
}
fn dspark_sparse_softmax_topk(
logits: &[f32],
top_k: usize,
temperature: f32,
) -> Result<(Vec<u32>, Vec<f32>, Vec<f32>, f32), Box<dyn std::error::Error>> {
if logits.is_empty() || top_k == 0 || top_k > logits.len() || temperature <= 0.0 {
return Err("invalid DSpark sparse-softmax shape or temperature".into());
}
if logits.iter().any(|value| !value.is_finite()) {
return Err("DSpark target logits contain a non-finite value".into());
}
let mut ranked: Vec<(u32, f32)> = logits
.iter()
.copied()
.enumerate()
.map(|(index, value)| (index as u32, value))
.collect();
let compare = |left: &(u32, f32), right: &(u32, f32)| {
right.1.total_cmp(&left.1).then(left.0.cmp(&right.0))
};
ranked.select_nth_unstable_by(top_k - 1, compare);
ranked[..top_k].sort_unstable_by(compare);
let max_logit = logits.iter().copied().fold(f32::NEG_INFINITY, f32::max);
let inv_temperature = 1.0f64 / temperature as f64;
let denominator: f64 = logits
.iter()
.map(|value| (((*value - max_logit) as f64) * inv_temperature).exp())
.sum();
let ids: Vec<u32> = ranked[..top_k].iter().map(|(index, _)| *index).collect();
let top_logits: Vec<f32> = ranked[..top_k].iter().map(|(_, value)| *value).collect();
let top_probs: Vec<f32> = top_logits
.iter()
.map(|value| ((((value - max_logit) as f64) * inv_temperature).exp() / denominator) as f32)
.collect();
let top_mass: f64 = top_probs.iter().map(|value| *value as f64).sum();
let tail = (1.0f64 - top_mass).clamp(0.0, 1.0) as f32;
Ok((ids, top_logits, top_probs, tail))
}
fn flatten_dspark_rows<T>(
rows: Vec<Option<Vec<T>>>,
position: usize,
label: &str,
) -> Result<Vec<T>, Box<dyn std::error::Error>> {
let mut flattened = Vec::new();
for (slot, row) in rows.into_iter().enumerate() {
flattened.extend(
row.ok_or_else(|| format!("missing DSpark {label} at {position} slot {slot}"))?,
);
}
Ok(flattened)
}
pub(crate) fn spec_hpost() -> bool {
static H: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
*H.get_or_init(|| {
std::env::var("MEMRA_SPEC_HPOST")
.map(|v| v != "0")
.unwrap_or(false)
})
}
pub(crate) fn spec_lean() -> bool {
static L: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
*L.get_or_init(|| {
std::env::var("MEMRA_SPEC_LEAN")
.map(|v| v != "0")
.unwrap_or(true)
})
}
pub(crate) fn spec_m2() -> bool {
static M: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
*M.get_or_init(|| {
std::env::var("MEMRA_SPEC_M2")
.map(|v| v != "0")
.unwrap_or(true)
})
}
pub(crate) fn spec_stream() -> bool {
static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
*ON.get_or_init(|| std::env::var("MEMRA_SPEC_STREAM").as_deref() == Ok("1"))
}
pub(crate) fn spec_stream_m() -> usize {
static M: std::sync::OnceLock<usize> = std::sync::OnceLock::new();
*M.get_or_init(|| {
std::env::var("MEMRA_SPEC_STREAM_M")
.ok()
.and_then(|v| v.parse().ok())
.unwrap_or(4)
})
}
pub(crate) fn spec_devacc() -> bool {
static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
*ON.get_or_init(|| std::env::var("MEMRA_SPEC_DEVACC").as_deref() == Ok("1"))
}
pub(crate) fn dspark_defer_readback_on() -> bool {
static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
*ON.get_or_init(|| {
std::env::var("MEMRA_DSPARK_DEFER_READBACK")
.map(|v| v != "0")
.unwrap_or(true)
})
}
pub(crate) fn state_copy_batch_on() -> bool {
static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
*ON.get_or_init(|| {
std::env::var("MEMRA_STATE_COPY_BATCH")
.map(|v| v != "0")
.unwrap_or(true)
})
}
pub(crate) fn dspark_verify_graph_on() -> bool {
static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
*ON.get_or_init(|| std::env::var("MEMRA_DSPARK_VERIFY_GRAPH").as_deref() == Ok("1"))
}
pub(crate) fn dspark_fa_rows_on() -> bool {
static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
*ON.get_or_init(|| {
std::env::var("MEMRA_DSPARK_FA_ROWS")
.map(|v| v != "0")
.unwrap_or(true)
})
}
fn debug_t_pred0(sampled: bool, base: usize, last_pred: u32, preds: &[u32]) -> String {
if base == 0 {
return last_pred.to_string();
}
match preds.get(base - 1) {
Some(p) => p.to_string(),
None => {
debug_assert!(
sampled,
"greedy spec: preds[{}] missing at base {base}",
base - 1
);
"n/a".to_string()
}
}
}
pub(crate) fn skey_probe() -> bool {
static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
*ON.get_or_init(|| std::env::var("MEMRA_SKEY_PROBE").as_deref() == Ok("1"))
}
pub trait SpecConstraint {
fn mask_logits(&mut self, logits: &mut [f32]) -> Result<(), String>;
fn mask_words(&mut self) -> Result<Vec<u32>, String>;
fn is_allowed(&mut self, tok: u32) -> Result<bool, String>;
fn consume(&mut self, tok: u32) -> Result<(), String>;
fn draft_mask_enabled(&self) -> bool {
false
}
fn draft_begin(&mut self) -> Result<(), String> {
Ok(())
}
fn draft_mask_words(&mut self) -> Result<Option<Vec<u32>>, String> {
Ok(None)
}
fn draft_advance(&mut self, _tok: u32) -> Result<bool, String> {
Ok(false)
}
}
fn upload_draft_mask(
e: &Engine,
c: &mut dyn SpecConstraint,
dst: &mut CudaSlice<u32>,
d2t: Option<&Vec<u32>>,
d_vocab: usize,
words: usize,
) -> Result<bool, Box<dyn std::error::Error>> {
let Some(tw) = c
.draft_mask_words()
.map_err(|e2| format!("constraint: {e2}"))?
else {
return Ok(false);
};
let bit = |t: usize| -> bool {
let w = t >> 5;
w < tw.len() && (tw[w] >> (t & 31)) & 1 == 1
};
let mut buf = vec![0u32; words];
match d2t {
Some(map) => {
for (i, &t) in map.iter().enumerate().take(d_vocab) {
if bit(t as usize) {
buf[i >> 5] |= 1u32 << (i & 31);
}
}
}
None => {
let n = tw.len().min(words);
buf[..n].copy_from_slice(&tw[..n]);
}
}
if buf.iter().all(|w| *w == 0) {
return Ok(false);
}
e.htod_u32_into(dst, &buf)?;
Ok(true)
}
pub(crate) fn spec_host_embd() -> bool {
static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
*ON.get_or_init(|| std::env::var("MEMRA_SPEC_HOST_EMBD").as_deref() == Ok("1"))
}
pub(crate) fn spec_fused_t() -> bool {
static F: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
*F.get_or_init(|| {
std::env::var("MEMRA_SPEC_FUSED_T")
.map(|v| v != "0")
.unwrap_or(true)
})
}
fn vbuf(e: &Engine, n: usize) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
if spec_lean() { e.uninit(n) } else { e.zeros(n) }
}
#[derive(Clone, Copy, Debug)]
pub struct SpecSampling {
pub temp: f32,
pub seed: u64,
pub top_k: i32, pub top_p: f32, pub min_p: f32, pub penalty_last_n: usize, pub penalty_repeat: f32,
pub penalty_freq: f32,
pub penalty_present: f32,
}
pub(crate) fn host_u01(seed: u64, ctr: u32) -> f32 {
let (m0, m1) = (0xD2511F53u32, 0xCD9E8D57u32);
let (mut c0, mut c1, mut c2, mut c3) = (0xFFFF_FFFEu32, ctr, 0u32, 0u32);
let (mut k0, mut k1) = ((seed & 0xFFFF_FFFF) as u32, (seed >> 32) as u32);
for _ in 0..10 {
let (h0, l0) = (((m0 as u64 * c0 as u64) >> 32) as u32, m0.wrapping_mul(c0));
let (h1, l1) = (((m1 as u64 * c2 as u64) >> 32) as u32, m1.wrapping_mul(c2));
let (n0, n1, n2, n3) = (h1 ^ c1 ^ k0, l1, h0 ^ c3 ^ k1, l0);
c0 = n0;
c1 = n1;
c2 = n2;
c3 = n3;
k0 = k0.wrapping_add(0x9E3779B9);
k1 = k1.wrapping_add(0xBB67AE85);
}
(c0 as f32 + 1.0) * (1.0 / 4294967296.0)
}
pub const SPEC_TELEM_POS: usize = 8;
#[derive(Clone, Copy, Default, Debug)]
pub struct SpecTelemetry {
pub rounds: u64,
pub drafted: u64,
pub accepted: u64,
pub pos_drafted: [u64; SPEC_TELEM_POS],
pub pos_accepted: [u64; SPEC_TELEM_POS],
}
impl SpecTelemetry {
pub fn delta_since(&self, prev: &SpecTelemetry) -> SpecTelemetry {
let mut d = SpecTelemetry {
rounds: self.rounds.saturating_sub(prev.rounds),
drafted: self.drafted.saturating_sub(prev.drafted),
accepted: self.accepted.saturating_sub(prev.accepted),
..Default::default()
};
for j in 0..SPEC_TELEM_POS {
d.pos_drafted[j] = self.pos_drafted[j].saturating_sub(prev.pos_drafted[j]);
d.pos_accepted[j] = self.pos_accepted[j].saturating_sub(prev.pos_accepted[j]);
}
d
}
pub fn merge(&mut self, d: &SpecTelemetry) {
self.rounds += d.rounds;
self.drafted += d.drafted;
self.accepted += d.accepted;
for j in 0..SPEC_TELEM_POS {
self.pos_drafted[j] += d.pos_drafted[j];
self.pos_accepted[j] += d.pos_accepted[j];
}
}
pub fn tau(&self) -> f64 {
if self.rounds > 0 {
self.accepted as f64 / self.rounds as f64
} else {
0.0
}
}
}
struct SpecTelemetryCounters {
rounds: AtomicU64,
drafted: AtomicU64,
accepted: AtomicU64,
pos_drafted: [AtomicU64; SPEC_TELEM_POS],
pos_accepted: [AtomicU64; SPEC_TELEM_POS],
}
impl Default for SpecTelemetryCounters {
fn default() -> Self {
Self {
rounds: AtomicU64::new(0),
drafted: AtomicU64::new(0),
accepted: AtomicU64::new(0),
pos_drafted: std::array::from_fn(|_| AtomicU64::new(0)),
pos_accepted: std::array::from_fn(|_| AtomicU64::new(0)),
}
}
}
impl SpecTelemetryCounters {
fn record_round(&self, drafted: usize, accepted: usize) {
debug_assert!(accepted <= drafted);
self.rounds.fetch_add(1, Ordering::Relaxed);
self.drafted.fetch_add(drafted as u64, Ordering::Relaxed);
self.accepted.fetch_add(accepted as u64, Ordering::Relaxed);
for counter in self.pos_drafted.iter().take(drafted) {
counter.fetch_add(1, Ordering::Relaxed);
}
for counter in self.pos_accepted.iter().take(accepted) {
counter.fetch_add(1, Ordering::Relaxed);
}
}
fn record_totals(&self, rounds: usize, drafted: usize, accepted: usize) {
self.rounds.fetch_add(rounds as u64, Ordering::Relaxed);
self.drafted.fetch_add(drafted as u64, Ordering::Relaxed);
self.accepted.fetch_add(accepted as u64, Ordering::Relaxed);
}
fn snapshot(&self) -> SpecTelemetry {
SpecTelemetry {
rounds: self.rounds.load(Ordering::Relaxed),
drafted: self.drafted.load(Ordering::Relaxed),
accepted: self.accepted.load(Ordering::Relaxed),
pos_drafted: std::array::from_fn(|j| self.pos_drafted[j].load(Ordering::Relaxed)),
pos_accepted: std::array::from_fn(|j| self.pos_accepted[j].load(Ordering::Relaxed)),
}
}
}
pub struct SpecSession {
pub(crate) cache: Cache,
pub(crate) scratch: MtpScratch,
pub committed: Vec<u32>,
pub(crate) last_h: Option<CudaSlice<f32>>,
pub next_pred: Option<u32>,
pub sctr: u32,
pub uctr: u32,
pub(crate) draft_ctx: Option<DraftGraphCtx>,
pub pending_tok: Option<u32>,
pub(crate) turn_ckpt: Option<SpecCheckpoint>,
telem: SpecTelemetryCounters,
pub capture_at: Option<usize>,
pub boundary_captures: Vec<SpecBoundaryCapture>,
pub ckpt_at: Option<usize>,
}
impl SpecSession {
pub fn cache_max_ctx(&self) -> usize {
self.cache.max_ctx
}
pub fn cache_ref(&self) -> &Cache {
&self.cache
}
pub fn draft_plane_ref(&self) -> Option<(&CudaSlice<u8>, &CudaSlice<u8>, usize, usize)> {
if self.scratch.kv.ring.is_some() {
return None;
}
Some((
&self.scratch.kv.k,
&self.scratch.kv.v,
self.scratch.kv.k_tok_bytes,
self.scratch.kv.v_tok_bytes,
))
}
pub fn telemetry(&self) -> SpecTelemetry {
self.telem.snapshot()
}
pub fn rewind_pos(&self) -> Option<usize> {
self.turn_ckpt.as_ref().map(|c| c.pos)
}
pub fn rewind_is_resident(&self) -> bool {
self.turn_ckpt.as_ref().is_some_and(|ckpt| {
self.cache.can_rollback(&ckpt.snap, 0) && self.scratch.can_rewind_to(ckpt.pos)
})
}
pub fn demote_ready(&self) -> bool {
self.pending_tok.is_none() && self.next_pred.is_some()
}
pub fn has_pending(&self) -> bool {
self.pending_tok.is_some()
}
pub fn committed_len(&self) -> usize {
self.committed.len()
}
pub fn into_demoted(self) -> Option<(Cache, u32)> {
if self.pending_tok.is_some() {
return None;
}
let np = self.next_pred?;
debug_assert_eq!(
self.cache.pos,
self.committed.len(),
"demotion handoff: cache rows != committed tokens"
);
Some((self.cache, np))
}
pub fn reset_graph_fallback_on_resume(&mut self) {
if let Some(line) = self
.draft_ctx
.as_mut()
.and_then(|c| c.failed.reset_on_resume())
{
eprintln!("{line}");
}
}
}
pub(crate) struct SpecCheckpoint {
snap: crate::cache::CacheSnapshot,
pos: usize,
last_h: CudaSlice<f32>,
}
pub struct SpecBoundaryCapture {
pub snap: crate::cache::CacheSnapshot,
pub pos: usize,
pub logits: Vec<f32>,
pub last_h: Vec<f32>,
}
fn capture_boundary_hidden(
e: &Engine,
h_rows: &CudaSlice<f32>,
pos: usize,
n_embd: usize,
) -> Vec<f32> {
if pos == 0 || h_rows.len() < pos * n_embd {
return Vec::new();
}
let Ok(mut row) = e.uninit(n_embd) else {
return Vec::new();
};
if e.copy_view_into(
&mut row,
0,
&h_rows.slice((pos - 1) * n_embd..pos * n_embd),
n_embd,
)
.is_err()
{
return Vec::new();
}
e.dtoh(&row).unwrap_or_default()
}
pub fn spec_sampled_boundary_on() -> bool {
static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
*ON.get_or_init(|| std::env::var("MEMRA_SPEC_SAMPLED_BOUNDARY").as_deref() != Ok("0"))
}
pub fn spec_pen_session_on() -> bool {
static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
*ON.get_or_init(|| std::env::var("MEMRA_SPEC_PEN_SESSION").as_deref() != Ok("0"))
}
pub fn spec_restore_republish_on() -> bool {
static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
*ON.get_or_init(|| std::env::var("MEMRA_SPEC_RESTORE_REPUBLISH").as_deref() != Ok("0"))
}
fn spec_boundary_trace() -> bool {
static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
*ON.get_or_init(|| std::env::var("MEMRA_SPEC_BOUNDARY_TRACE").as_deref() == Ok("1"))
}
const PEN_WINDOW_FLOOR: usize = 64;
const PEN_WINDOW_MAX: usize = 8192;
fn pen_window_seed(
session_committed: &[u32],
burst_prompt: &[u32],
penalty_last_n: usize,
) -> Vec<u32> {
let win = penalty_last_n.clamp(PEN_WINDOW_FLOOR, PEN_WINDOW_MAX);
let take_prompt = burst_prompt.len().min(win);
let take_sess = (win - take_prompt).min(session_committed.len());
let mut hist = Vec::with_capacity(take_sess + take_prompt);
hist.extend_from_slice(&session_committed[session_committed.len() - take_sess..]);
hist.extend_from_slice(&burst_prompt[burst_prompt.len() - take_prompt..]);
hist
}
#[allow(clippy::too_many_arguments)]
pub fn sample_boundary_token_dev(
e: &Engine,
logits: &CudaSlice<f32>,
n_vocab: usize,
sp: &SpecSampling,
pen_hist: &[u32],
sctr: &mut u32,
site: &str,
) -> Result<u32, Box<dyn std::error::Error>> {
debug_assert!(
sp.temp > 0.0,
"boundary sampling is the sampled regime only"
);
let mut col = e.zeros(n_vocab)?;
e.copy_into(&mut col, 0, logits, n_vocab)?;
let pen_on = sp.penalty_last_n > 0
&& (sp.penalty_repeat != 1.0 || sp.penalty_freq != 0.0 || sp.penalty_present != 0.0);
if pen_on && !pen_hist.is_empty() {
let w0 = pen_hist
.len()
.saturating_sub(sp.penalty_last_n.min(PEN_WINDOW_MAX));
let hist = &pen_hist[w0..];
let hd = e.htod_u32_v(hist)?;
e.penalize_logits(
&mut col,
&hd,
hist.len(),
sp.penalty_repeat,
sp.penalty_freq,
sp.penalty_present,
n_vocab,
)?;
}
let rows0 = e.htod_i32(&[0])?;
let (mut th_d, mut z_d, mut mx_d) = (e.zeros(1)?, e.zeros(1)?, e.zeros(1)?);
e.filter_stats(
&col, n_vocab, &rows0, &mut th_d, &mut z_d, &mut mx_d, n_vocab, 1, sp.temp, sp.top_k,
sp.top_p, sp.min_p,
)?;
let (th, mx) = (e.dtoh(&th_d)?[0], e.dtoh(&mx_d)?[0]);
let mut perturb = e.zeros(n_vocab)?;
e.gumbel_perturb_filtered(&col, &mut perturb, n_vocab, sp.seed, *sctr, sp.temp, mx, th)?;
*sctr = sctr.wrapping_add(1);
let td = e.argmax_token_device(&perturb, n_vocab)?;
let tok = e.dtoh_u32_one(&td)?;
if spec_boundary_trace() {
let raw = e.argmax_token_device(logits, n_vocab)?;
let greedy = e.dtoh_u32_one(&raw)?;
eprintln!(
"[spec-boundary] site={site} sampled={tok} argmax={greedy} \
deviates={} temp={} sctr={}",
(tok != greedy) as u8,
sp.temp,
sctr.wrapping_sub(1),
);
}
Ok(tok)
}
#[allow(clippy::too_many_arguments)]
pub fn sample_boundary_token(
e: &Engine,
logits: &[f32],
sp: &SpecSampling,
pen_hist: &[u32],
sctr: &mut u32,
site: &str,
) -> Result<u32, Box<dyn std::error::Error>> {
let n_vocab = logits.len();
let d = e.htod(logits)?;
sample_boundary_token_dev(e, &d, n_vocab, sp, pen_hist, sctr, site)
}
struct SpecPipeTraceClock {
pair: usize,
started: std::time::Instant,
}
#[derive(Clone)]
struct SpecPipeTraceCtx {
clock: std::sync::Arc<SpecPipeTraceClock>,
round: usize,
lane: usize,
}
struct SpecPipeTraceMarker {
trace: SpecPipeTraceCtx,
phase: &'static str,
edge: &'static str,
slot: Option<usize>,
}
unsafe extern "C" fn spec_pipe_trace_marker(raw: *mut std::ffi::c_void) {
let marker = unsafe { Box::from_raw(raw.cast::<SpecPipeTraceMarker>()) };
let lane = if marker.trace.lane == 0 { "A" } else { "B" };
let slot = marker
.slot
.map(|v| v.to_string())
.unwrap_or_else(|| "-".into());
let t_ms = marker.trace.clock.started.elapsed().as_secs_f64() * 1e3;
use std::io::Write as _;
let stderr = std::io::stderr();
let mut stderr = stderr.lock();
let _ = writeln!(
stderr,
"[spec-pipe-timeline] pair={} round={} lane={lane} phase={} edge={} \
slot={slot} t_ms={t_ms:.3}",
marker.trace.clock.pair, marker.trace.round, marker.phase, marker.edge,
);
}
fn enqueue_spec_pipe_trace_marker(
stream: &cudarc::driver::CudaStream,
trace: Option<&SpecPipeTraceCtx>,
phase: &'static str,
edge: &'static str,
slot: Option<usize>,
) -> Result<(), Box<dyn std::error::Error>> {
let Some(trace) = trace else {
return Ok(());
};
let marker = Box::new(SpecPipeTraceMarker {
trace: trace.clone(),
phase,
edge,
slot,
});
let raw = Box::into_raw(marker);
let result = unsafe {
cudarc::driver::result::stream::launch_host_function(
stream.cu_stream(),
spec_pipe_trace_marker,
raw.cast(),
)
};
if let Err(err) = result {
unsafe {
drop(Box::from_raw(raw));
}
return Err(err.into());
}
Ok(())
}
#[derive(Default)]
struct SpecPipeProgress {
setup_done: [bool; 2],
draft_done: [usize; 2],
stage0_done: [usize; 2],
verify_done: [usize; 2],
accept_done: [usize; 2],
finished: [bool; 2],
aborted: bool,
}
struct SpecPipeSync {
progress: std::sync::Mutex<SpecPipeProgress>,
changed: std::sync::Condvar,
primary: std::sync::Mutex<()>,
trace: Option<std::sync::Arc<SpecPipeTraceClock>>,
}
impl SpecPipeSync {
fn new() -> Self {
static TRACE_PAIR: std::sync::atomic::AtomicUsize = std::sync::atomic::AtomicUsize::new(0);
let trace = (std::env::var("MEMRA_SPEC_PIPE_TRACE").as_deref() == Ok("1")).then(|| {
std::sync::Arc::new(SpecPipeTraceClock {
pair: TRACE_PAIR.fetch_add(1, std::sync::atomic::Ordering::Relaxed) + 1,
started: std::time::Instant::now(),
})
});
Self {
progress: std::sync::Mutex::new(SpecPipeProgress::default()),
changed: std::sync::Condvar::new(),
primary: std::sync::Mutex::new(()),
trace,
}
}
}
#[derive(Clone)]
struct SpecPipeLane {
sync: std::sync::Arc<SpecPipeSync>,
lane: usize,
}
impl SpecPipeLane {
fn peer(&self) -> usize {
1 - self.lane
}
fn aborted() -> Box<dyn std::error::Error> {
"paired speculative peer aborted".into()
}
fn trace(&self, round: usize) -> Option<SpecPipeTraceCtx> {
self.sync.trace.as_ref().map(|clock| SpecPipeTraceCtx {
clock: clock.clone(),
round,
lane: self.lane,
})
}
fn setup_begin(&self) -> Result<(), Box<dyn std::error::Error>> {
let mut p = self.sync.progress.lock().unwrap();
while !p.aborted && self.lane == 1 && !p.setup_done[0] && !p.finished[0] {
p = self.sync.changed.wait(p).unwrap();
}
if p.aborted {
Err(Self::aborted())
} else {
Ok(())
}
}
fn setup_end(&self) {
let mut p = self.sync.progress.lock().unwrap();
p.setup_done[self.lane] = true;
self.sync.changed.notify_all();
}
fn draft_begin(
&self,
round: usize,
) -> Result<std::sync::MutexGuard<'_, ()>, Box<dyn std::error::Error>> {
let peer = self.peer();
let mut p = self.sync.progress.lock().unwrap();
loop {
if p.aborted {
return Err(Self::aborted());
}
let setup_ready =
(p.setup_done[0] || p.finished[0]) && (p.setup_done[1] || p.finished[1]);
let prior_ready = p.accept_done[self.lane] >= round
&& (p.accept_done[peer] >= round || p.finished[peer]);
let turn_ready = if self.lane == 0 {
true
} else {
p.draft_done[0] > round || p.finished[0]
};
if setup_ready && prior_ready && turn_ready {
break;
}
p = self.sync.changed.wait(p).unwrap();
}
drop(p);
Ok(self.sync.primary.lock().unwrap())
}
fn draft_end(&self, round: usize) {
let mut p = self.sync.progress.lock().unwrap();
p.draft_done[self.lane] = round + 1;
self.sync.changed.notify_all();
}
fn stage0_begin(&self, round: usize) -> Result<bool, Box<dyn std::error::Error>> {
let peer = self.peer();
let mut p = self.sync.progress.lock().unwrap();
loop {
if p.aborted {
return Err(Self::aborted());
}
let ready = if self.lane == 0 {
p.draft_done[0] > round && (p.draft_done[1] > round || p.finished[1])
} else {
p.draft_done[1] > round && (p.stage0_done[0] > round || p.finished[0])
};
if ready {
return Ok(self.lane == 0 || p.finished[peer]);
}
p = self.sync.changed.wait(p).unwrap();
}
}
fn stage0_end(&self, round: usize) {
let mut p = self.sync.progress.lock().unwrap();
p.stage0_done[self.lane] = round + 1;
self.sync.changed.notify_all();
}
fn stage1_begin(&self, round: usize) -> Result<(), Box<dyn std::error::Error>> {
let mut p = self.sync.progress.lock().unwrap();
while !p.aborted
&& !(p.stage0_done[self.lane] > round
&& (self.lane == 0 || p.verify_done[0] > round || p.finished[0]))
{
p = self.sync.changed.wait(p).unwrap();
}
if p.aborted {
Err(Self::aborted())
} else {
Ok(())
}
}
fn verify_end(&self, round: usize) {
let mut p = self.sync.progress.lock().unwrap();
p.verify_done[self.lane] = round + 1;
self.sync.changed.notify_all();
}
fn accept_begin(
&self,
round: usize,
) -> Result<std::sync::MutexGuard<'_, ()>, Box<dyn std::error::Error>> {
let mut p = self.sync.progress.lock().unwrap();
loop {
if p.aborted {
return Err(Self::aborted());
}
let ready = if self.lane == 0 {
p.verify_done[0] > round && (p.verify_done[1] > round || p.finished[1])
} else {
p.verify_done[1] > round && (p.accept_done[0] > round || p.finished[0])
};
if ready {
break;
}
p = self.sync.changed.wait(p).unwrap();
}
drop(p);
Ok(self.sync.primary.lock().unwrap())
}
fn accept_end(&self, round: usize) {
let mut p = self.sync.progress.lock().unwrap();
p.accept_done[self.lane] = round + 1;
self.sync.changed.notify_all();
}
fn primary(&self) -> std::sync::MutexGuard<'_, ()> {
self.sync.primary.lock().unwrap()
}
fn finish(&self, failed: bool) {
let mut p = self.sync.progress.lock().unwrap();
p.finished[self.lane] = true;
p.aborted |= failed;
self.sync.changed.notify_all();
}
}
struct SpecPipeFinish<'a> {
lane: &'a SpecPipeLane,
closed: bool,
}
impl<'a> SpecPipeFinish<'a> {
fn new(lane: &'a SpecPipeLane) -> Self {
Self {
lane,
closed: false,
}
}
fn close(&mut self, failed: bool) {
self.lane.finish(failed);
self.closed = true;
}
}
impl Drop for SpecPipeFinish<'_> {
fn drop(&mut self) {
if !self.closed {
self.lane.finish(true);
}
}
}
struct SpecPipeSessionPtr(*mut SpecSession);
unsafe impl Send for SpecPipeSessionPtr {}
impl SpecPipeSessionPtr {
unsafe fn get_mut(&mut self) -> &mut SpecSession {
unsafe { &mut *self.0 }
}
}
#[derive(Clone, Copy, PartialEq, Eq, Debug)]
pub(crate) struct SampledGraphKey {
seed: u64,
temp_bits: u32,
k: usize,
top_k: i32,
top_p_bits: u32,
min_p_bits: u32,
pen_on: bool,
}
impl SampledGraphKey {
pub(crate) fn new(
seed: u64,
temp: f32,
k: usize,
top_k: i32,
top_p: f32,
min_p: f32,
pen_on: bool,
) -> Self {
SampledGraphKey {
seed,
temp_bits: temp.to_bits(),
k,
top_k,
top_p_bits: top_p.to_bits(),
min_p_bits: min_p.to_bits(),
pen_on,
}
}
pub(crate) fn pure_temp(&self) -> bool {
self.top_k == 0
&& f32::from_bits(self.top_p_bits) >= 1.0
&& f32::from_bits(self.min_p_bits) <= 0.0
&& !self.pen_on
}
}
pub(crate) struct DraftGraphCtx {
g_tok: CudaSlice<u32>,
g_pos: CudaSlice<i32>,
g_seed: CudaSlice<f32>,
g_p: CudaSlice<f32>,
g_ctr: CudaSlice<u32>,
g_q: CudaSlice<f32>,
g_perturb: CudaSlice<f32>,
q_slots: Vec<CudaSlice<f32>>,
g_dmask: CudaSlice<u32>,
graph_masked: bool,
graph: Option<cudarc::driver::CudaGraph>,
graph_s: Option<cudarc::driver::CudaGraph>,
failed: DraftGraphFallback,
s_key: Option<SampledGraphKey>,
keeper: Vec<Box<dyn std::any::Any + Send>>,
keeper_s: Vec<Box<dyn std::any::Any + Send>>,
}
#[derive(Default)]
pub(crate) struct DraftGraphFallback {
greedy: bool,
sampled: bool,
}
impl DraftGraphFallback {
fn mark_greedy(&mut self, reason: &str) -> Option<String> {
if self.greedy {
return None;
}
self.greedy = true;
Some(format!(
"[spec] WARN: draft-graph capture failed ({reason}); eager fallback until session resume"
))
}
fn mark_sampled(&mut self, reason: &str) -> Option<String> {
if self.sampled {
return None;
}
self.sampled = true;
Some(format!(
"[spec] WARN: sampled draft-graph capture failed ({reason}); eager fallback until session resume"
))
}
fn greedy_failed(&self) -> bool {
self.greedy
}
fn sampled_failed(&self) -> bool {
self.sampled
}
fn clear_greedy(&mut self) {
self.greedy = false;
}
fn clear_sampled(&mut self) {
self.sampled = false;
}
pub(crate) fn reset_on_resume(&mut self) -> Option<String> {
if !self.greedy && !self.sampled {
return None;
}
let which = match (self.greedy, self.sampled) {
(true, true) => "greedy+sampled",
(true, false) => "greedy",
_ => "sampled",
};
self.greedy = false;
self.sampled = false;
Some(format!(
"[spec] draft-graph fallback reset on session resume ({which}); recapture eligible"
))
}
}
impl DraftGraphCtx {
fn new(e: &Engine, n_embd: usize, qlen: usize) -> Result<Self, Box<dyn std::error::Error>> {
Ok(DraftGraphCtx {
g_tok: e.alloc_u32_zeroed(1)?,
g_pos: e.htod_i32(&[0])?,
g_seed: e.zeros(n_embd)?,
g_p: e.zeros(1)?,
g_ctr: e.alloc_u32_zeroed(1)?,
g_q: e.zeros(qlen)?,
g_perturb: e.zeros(qlen)?,
q_slots: Vec::new(),
g_dmask: e.alloc_u32_zeroed(1)?,
graph_masked: false,
graph: None,
graph_s: None,
failed: DraftGraphFallback::default(),
s_key: None,
keeper: Vec::new(),
keeper_s: Vec::new(),
})
}
}
pub(crate) struct MtpScratch {
kv: KvLayer,
cap: usize,
}
fn mtp_scratch_layout(
cfg: &memra_gguf::config::ModelConfig,
geom: Option<&crate::hybrid::DraftGeom>,
) -> (usize, usize, usize, usize) {
let n_head_kv = geom.map(|g| g.n_head_kv).unwrap_or(cfg.n_head_kv as usize);
let head_dim_k = cfg.head_dim_k as usize;
let head_dim_v = cfg.head_dim_v as usize;
assert!(
head_dim_k % 32 == 0 && head_dim_v % 32 == 0,
"KVQUANT requires head_dim%32==0 (MTP scratch)"
);
let kv_dim_k = head_dim_k * n_head_kv;
let kv_dim_v = head_dim_v * n_head_kv;
let (kbb, vbb) = crate::kv_blk_bytes();
let k_tok_bytes = (kv_dim_k / 32) * kbb;
let v_tok_bytes = (kv_dim_v / 32) * vbb;
(kv_dim_k, kv_dim_v, k_tok_bytes, v_tok_bytes)
}
impl MtpScratch {
fn new(
e: &Engine,
cfg: &memra_gguf::config::ModelConfig,
cap: usize,
geom: Option<&crate::hybrid::DraftGeom>,
) -> Result<Self, Box<dyn std::error::Error>> {
let (kv_dim_k, kv_dim_v, k_tok_bytes, v_tok_bytes) = mtp_scratch_layout(cfg, geom);
let ring = if crate::cache::swa_ring_on() && cfg.arch.is_step35() {
let window = cfg.step35.as_ref().unwrap().sliding_window as usize;
Some(crate::cache::KvRing::new(
crate::cache::swa_ring_rows(window, cap),
window,
))
} else {
None
};
let alloc_rows = ring.as_ref().map(crate::cache::KvRing::rows).unwrap_or(cap);
Ok(MtpScratch {
kv: KvLayer {
k: e.alloc_u8(alloc_rows * k_tok_bytes)?,
v: e.alloc_u8(alloc_rows * v_tok_bytes)?,
kv_dim_k,
kv_dim_v,
k_tok_bytes,
v_tok_bytes,
len: 0,
ring,
len_d: e.htod_i32(&[0])?,
},
cap,
})
}
fn set_len(&mut self, e: &Engine, n: usize) -> Result<(), Box<dyn std::error::Error>> {
if self
.kv
.ring
.as_ref()
.is_some_and(|ring| !ring.can_rewind_to(n))
{
return Err("SWA ring MTP checkpoint has been lapped; full re-prime required".into());
}
self.kv.len = n;
e.set_i32_one(&mut self.kv.len_d, n as i32)
}
fn can_rewind_to(&self, n: usize) -> bool {
self.kv
.ring
.as_ref()
.is_none_or(|ring| ring.can_rewind_to(n))
}
}
struct GdnStash {
qkv_mixed: CudaSlice<f32>, q_l2: CudaSlice<f32>,
k_l2: CudaSlice<f32>,
v_g: CudaSlice<f32>, g_log: CudaSlice<f32>,
beta: CudaSlice<f32>, }
pub(crate) struct VerifyCkpt {
gdn: Vec<Option<GdnStash>>, cols: Vec<Option<Vec<(CudaSlice<f32>, CudaSlice<f32>)>>>, }
pub(crate) struct DsparkVerifyCkpt(VerifyCkpt);
pub(crate) struct DsparkVerifyGraphs {
lin: Vec<usize>,
lin_pos: std::collections::HashMap<usize, usize>,
table_all: CudaSlice<u64>,
host_table: Vec<u64>,
stash_conv: Vec<CudaSlice<f32>>,
stash_ssm: Vec<CudaSlice<f32>>,
conv_words: usize,
ssm_words: usize,
stage: std::collections::HashMap<usize, (CudaSlice<f32>, CudaSlice<f32>)>,
pub(crate) tap_bufs: std::collections::HashMap<usize, CudaSlice<f32>>,
graphs: std::collections::HashMap<(usize, usize), DsparkSegGraph>,
save_conv: CudaSlice<f32>,
save_ssm: CudaSlice<f32>,
max_run: usize,
n_embd: usize,
pub(crate) round_slab: bool,
fa: Vec<usize>,
fa_pos: std::collections::HashMap<usize, usize>,
fa_table: CudaSlice<u64>,
fa_host_table: Vec<u64>,
t_cap: usize,
pos_stage: std::collections::HashMap<usize, CudaSlice<i32>>,
full: std::collections::HashMap<(usize, usize, usize), DsparkSegGraph>,
covered: usize,
walk_uniform: bool,
}
struct DsparkSegGraph {
graph: cudarc::driver::CudaGraph,
_keeper: Vec<Box<dyn std::any::Any + Send>>,
}
pub(crate) struct FaLayerArgs<'a> {
pub pos_d: &'a CudaSlice<i32>,
pub pos_rows: &'a mut Option<Vec<CudaSlice<i32>>>,
pub pos0: usize,
pub seqs_append: bool,
pub batch_fa_on: bool,
pub graph_cap: Option<(&'a CudaSlice<u64>, usize, usize)>,
pub stream: Option<(&'a CudaSlice<u32>, &'a CudaSlice<i32>)>,
pub ckpt: Option<&'a mut VerifyCkpt>,
}
unsafe impl Send for DsparkVerifyGraphs {}
impl DsparkVerifyGraphs {
pub(crate) fn new(
e: &Engine,
cache: &Cache,
t_max: usize,
n_embd: usize,
) -> Result<Option<Self>, Box<dyn std::error::Error>> {
let lin: Vec<usize> = (0..cache.recur.len())
.filter(|&il| cache.recur[il].is_some())
.collect();
if lin.is_empty() || t_max < 2 {
return Ok(None);
}
let first = cache.recur[lin[0]].as_ref().unwrap();
let (conv_words, ssm_words) = (first.conv_state.len(), first.ssm_state.len());
for &il in &lin {
let rl = cache.recur[il].as_ref().unwrap();
if rl.conv_state.len() != conv_words || rl.ssm_state.len() != ssm_words {
return Ok(None);
}
}
let n = lin.len();
let mut lin_pos = std::collections::HashMap::with_capacity(n);
for (k, &il) in lin.iter().enumerate() {
lin_pos.insert(il, k);
}
let mut max_run = 1usize;
let mut run = 1usize;
for w in lin.windows(2) {
if w[1] == w[0] + 1 {
run += 1;
max_run = max_run.max(run);
} else {
run = 1;
}
}
let rows = t_max - 1;
let mut stash_conv = Vec::with_capacity(n);
let mut stash_ssm = Vec::with_capacity(n);
for _ in 0..n {
stash_conv.push(e.uninit(rows * conv_words)?);
stash_ssm.push(e.uninit(rows * ssm_words)?);
}
let host_table = vec![0u64; n * 6];
let table_all = e.htod_u64(&host_table)?;
let fa: Vec<usize> = (0..cache.kv.len())
.filter(|&il| cache.kv[il].is_some())
.collect();
let mut fa_pos = std::collections::HashMap::with_capacity(fa.len());
for (k, &il) in fa.iter().enumerate() {
fa_pos.insert(il, k);
}
let n_layers = cache.kv.len().max(cache.recur.len());
let walk_uniform = (0..n_layers).all(|il| {
cache.recur.get(il).is_some_and(|r| r.is_some())
!= cache.kv.get(il).is_some_and(|k| k.is_some())
});
let covered = (0..n_layers)
.take_while(|il| lin_pos.contains_key(il) || fa_pos.contains_key(il))
.count();
let t_cap = t_max;
let fa_host_table = vec![0u64; fa.len() * 2 * t_cap];
let fa_table = e.htod_u64(&fa_host_table)?;
Ok(Some(Self {
lin,
lin_pos,
table_all,
host_table,
stash_conv,
stash_ssm,
conv_words,
ssm_words,
stage: std::collections::HashMap::new(),
tap_bufs: std::collections::HashMap::new(),
graphs: std::collections::HashMap::new(),
save_conv: e.uninit(n * conv_words)?,
save_ssm: e.uninit(n * ssm_words)?,
max_run,
n_embd,
round_slab: false,
fa,
fa_pos,
fa_table,
fa_host_table,
t_cap,
pos_stage: std::collections::HashMap::new(),
full: std::collections::HashMap::new(),
covered,
walk_uniform,
}))
}
pub(crate) fn refresh_tables(
&mut self,
e: &Engine,
cache: &Cache,
) -> Result<(), Box<dyn std::error::Error>> {
use cudarc::driver::DevicePtr;
{
let s = &e.gpu.stream();
for (k, &il) in self.lin.iter().enumerate() {
let rl = cache.recur[il].as_ref().unwrap();
let (pc, _g0) = rl.conv_state.device_ptr(s);
let (p0, _g1) = rl.ssm_state.device_ptr(s);
let (p1, _g2) = rl.ssm_state_alt.device_ptr(s);
let o = k * 6;
self.host_table[o] = pc as u64;
self.host_table[o + 1] = p0 as u64;
self.host_table[o + 2] = p1 as u64;
self.host_table[o + 3] = pc as u64;
self.host_table[o + 4] = p1 as u64;
self.host_table[o + 5] = p0 as u64;
}
for (k, &il) in self.fa.iter().enumerate() {
let kvl = cache.kv[il].as_ref().unwrap();
let (pk, _g0) = kvl.k.device_ptr(s);
let (pv, _g1) = kvl.v.device_ptr(s);
let o = k * 2 * self.t_cap;
for z in 0..self.t_cap {
self.fa_host_table[o + 2 * z] = pk as u64;
self.fa_host_table[o + 2 * z + 1] = pv as u64;
}
}
}
e.htod_u64_into(&self.host_table, &mut self.table_all)?;
if !self.fa_host_table.is_empty() {
e.htod_u64_into(&self.fa_host_table, &mut self.fa_table)?;
}
Ok(())
}
#[allow(clippy::too_many_arguments)]
pub(crate) fn full_rung(
&self,
model: &crate::hybrid::HybridModel,
cache: &Cache,
lo: usize,
hi: usize,
t: usize,
seqs_arms_on: bool,
) -> Option<usize> {
if std::env::var("MEMRA_DSPARK_FULLG_DEBUG").as_deref() == Ok("1") {
static ONCE: std::sync::Once = std::sync::Once::new();
let len0 = self
.fa
.first()
.and_then(|&il| cache.kv[il].as_ref())
.map(|k| k.len);
ONCE.call_once(|| {
eprintln!(
"[fullg-debug] walk_uniform={} covered={} seqs_arms_on={} fa_rows_on={} t={} lo={} hi={} lin={} fa={} t_cap={} len0={:?}",
self.walk_uniform, self.covered, seqs_arms_on, dspark_fa_rows_on(), t, lo, hi,
self.lin.len(), self.fa.len(), self.t_cap, len0
);
});
}
if !self.walk_uniform
|| !seqs_arms_on
|| !dspark_fa_rows_on()
|| t < 2
|| lo != 0
|| hi > self.covered
|| t > self.t_cap
|| self.fa.is_empty()
{
return None;
}
let cfg = &model.cfg;
let head_dim_global = cfg.head_dim_k as usize;
let nkv = cfg.n_head_kv as usize;
let kvl0 = cache.kv[self.fa[0]].as_ref().unwrap();
let geom = cfg.full_attention_geometry_at(self.fa[0] as u32);
let kv_dim = geom.n_head_kv as usize * geom.head_dim_k as usize;
if kvl0.kv_dim_k != kv_dim || kvl0.kv_dim_v != kv_dim {
return None;
}
let len0 = kvl0.len;
let (t_kv_first, t_kv_last) = (len0 + 1, len0 + t);
if !crate::fa_seqs_eligible(t_kv_first, head_dim_global)
|| !crate::fa_seqs_eligible(t_kv_last, head_dim_global)
|| crate::fa_split_keys(t_kv_first, nkv) != crate::fa_split_keys(t_kv_last, nkv)
{
return None;
}
let rung = t_kv_last.next_power_of_two().max(256);
if crate::fa_split_keys(rung, nkv) != crate::fa_split_keys(t_kv_last, nkv) {
return None;
}
Some(rung)
}
#[allow(clippy::too_many_arguments)]
pub(crate) fn run_full(
&mut self,
model: &crate::hybrid::HybridModel,
e: &Engine,
lo: usize,
hi: usize,
x: &CudaSlice<f32>,
t: usize,
pos0: usize,
rung: usize,
cache: &mut Cache,
) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
let n_embd = self.n_embd;
if !self.stage.contains_key(&t) {
let xin = e.uninit(t * n_embd)?;
let xout = e.uninit(t * n_embd)?;
self.stage.insert(t, (xin, xout));
}
if !self.pos_stage.contains_key(&t) {
self.pos_stage.insert(t, e.htod_i32(&vec![0i32; t])?);
}
{
let pos_host: Vec<i32> = (0..t).map(|r| (pos0 + r) as i32).collect();
let pb = self.pos_stage.get_mut(&t).unwrap();
e.htod_i32_into(pb, &pos_host)?;
let (xin, _) = self.stage.get_mut(&t).unwrap();
e.copy_into(xin, 0, x, t * n_embd)?;
}
let key = (t, rung, hi);
if !self.full.contains_key(&key) {
for (k, &il) in self.lin.iter().enumerate() {
let rl = cache.recur[il].as_ref().unwrap();
e.copy_into(
&mut self.save_conv,
k * self.conv_words,
&rl.conv_state,
self.conv_words,
)?;
e.copy_into(
&mut self.save_ssm,
k * self.ssm_words,
&rl.ssm_state,
self.ssm_words,
)?;
}
let (graph, keeper) = {
let table_all = &self.table_all;
let lin_pos = &self.lin_pos;
let fa_pos = &self.fa_pos;
let fa_table = &self.fa_table;
let t_cap = self.t_cap;
let stash_conv = &mut self.stash_conv;
let stash_ssm = &mut self.stash_ssm;
let pos_d: &CudaSlice<i32> = &self.pos_stage[&t];
let (xin, xout) = self
.stage
.get_mut(&t)
.map(|(a, b)| (&*a, b))
.expect("stage bucket created above");
let cache_ref: &mut Cache = cache;
let iflag = if std::env::var("MEMRA_DSPARK_VG_AUTOFREE").as_deref() == Ok("1") {
cudarc::driver::sys::CUgraphInstantiate_flags::CUDA_GRAPH_INSTANTIATE_FLAG_AUTO_FREE_ON_LAUNCH
} else {
cudarc::driver::sys::CUgraphInstantiate_flags::CUDA_GRAPH_INSTANTIATE_FLAG_USE_NODE_PRIORITY
};
e.capture_graph_retained_flags(iflag, move |e| {
let mut xc: Option<CudaSlice<f32>> = None;
for il in lo..hi {
let xr: &CudaSlice<f32> = xc.as_ref().unwrap_or(xin);
let nx = if let Some(&k) = lin_pos.get(&il) {
model.qwen35_tparallel_linear_layer(
e,
il,
xr,
t,
cache_ref,
None,
Some((&mut stash_conv[k], &mut stash_ssm[k])),
Some((table_all, k * 6)),
)?
} else if let Some(&kf) = fa_pos.get(&il) {
let mut no_rows: Option<Vec<CudaSlice<i32>>> = None;
model.qwen35_tparallel_fa_layer(
e,
il,
xr,
t,
cache_ref,
FaLayerArgs {
pos_d,
pos_rows: &mut no_rows,
pos0,
seqs_append: true,
batch_fa_on: true,
graph_cap: Some((fa_table, kf * 2 * t_cap, rung)),
stream: None,
ckpt: None,
},
)?
} else {
return Err(format!(
"run_full: layer {il} is neither linear nor full-attention"
)
.into());
};
xc = Some(nx);
}
e.copy_into(xout, 0, xc.as_ref().unwrap(), t * n_embd)?;
Ok(())
})?
};
if t % 2 == 1 {
for &il in &self.lin {
if il < lo || il >= hi {
continue;
}
let rl = cache.recur[il].as_mut().unwrap();
std::mem::swap(&mut rl.ssm_state, &mut rl.ssm_state_alt);
}
}
for (k, &il) in self.lin.iter().enumerate() {
if il < lo || il >= hi {
continue;
}
let rl = cache.recur[il].as_mut().unwrap();
let (cw, sw) = (self.conv_words, self.ssm_words);
{
let sv = e.view(&self.save_conv, self.lin.len() * cw);
let win = sv.slice(k * cw..(k + 1) * cw);
e.copy_view_into(&mut rl.conv_state, 0, &win, cw)?;
}
{
let sv = e.view(&self.save_ssm, self.lin.len() * sw);
let win = sv.slice(k * sw..(k + 1) * sw);
e.copy_view_into(&mut rl.ssm_state, 0, &win, sw)?;
}
}
if std::env::var("MEMRA_GRAPH_CENSUS").as_deref() == Ok("1") {
if let Ok(c) = crate::graph_update::node_census(&graph) {
eprintln!("[dspark-vg-census] full vt={t} rung={rung} {c:?}");
}
}
self.full.insert(
key,
DsparkSegGraph {
graph,
_keeper: keeper,
},
);
}
self.full[&key].graph.launch()?;
if t % 2 == 1 {
for &il in &self.lin {
if il < lo || il >= hi {
continue;
}
let rl = cache.recur[il].as_mut().unwrap();
std::mem::swap(&mut rl.ssm_state, &mut rl.ssm_state_alt);
}
}
for &il in &self.fa {
if il < lo || il >= hi {
continue;
}
cache.kv[il].as_mut().unwrap().len += t;
}
let (_, xout) = self.stage.get(&t).unwrap();
let mut out = e.uninit(t * n_embd)?;
e.copy_into(&mut out, 0, xout, t * n_embd)?;
Ok(out)
}
#[allow(clippy::too_many_arguments)]
fn run_segment(
&mut self,
model: &crate::hybrid::HybridModel,
e: &Engine,
start: usize,
end: usize,
x: &CudaSlice<f32>,
t: usize,
cache: &mut Cache,
) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
let n_embd = self.n_embd;
debug_assert!(end - start <= self.max_run);
if !self.stage.contains_key(&t) {
let xin = e.uninit(t * n_embd)?;
let xout = e.uninit(t * n_embd)?;
self.stage.insert(t, (xin, xout));
}
{
let (xin, _) = self.stage.get_mut(&t).unwrap();
e.copy_into(xin, 0, x, t * n_embd)?;
}
let key = (start, t);
if !self.graphs.contains_key(&key) {
for (k, il) in (start..end).enumerate() {
let rl = cache.recur[il].as_ref().unwrap();
e.copy_into(
&mut self.save_conv,
k * self.conv_words,
&rl.conv_state,
self.conv_words,
)?;
e.copy_into(
&mut self.save_ssm,
k * self.ssm_words,
&rl.ssm_state,
self.ssm_words,
)?;
}
let (graph, keeper) = {
let table_all = &self.table_all;
let lin_pos = &self.lin_pos;
let stash_conv = &mut self.stash_conv;
let stash_ssm = &mut self.stash_ssm;
let (xin, xout) = self
.stage
.get_mut(&t)
.map(|(a, b)| (&*a, b))
.expect("stage bucket created above");
let cache_ref: &mut Cache = cache;
let iflag = if std::env::var("MEMRA_DSPARK_VG_AUTOFREE").as_deref() == Ok("1") {
cudarc::driver::sys::CUgraphInstantiate_flags::CUDA_GRAPH_INSTANTIATE_FLAG_AUTO_FREE_ON_LAUNCH
} else {
cudarc::driver::sys::CUgraphInstantiate_flags::CUDA_GRAPH_INSTANTIATE_FLAG_USE_NODE_PRIORITY
};
e.capture_graph_retained_flags(iflag, move |e| {
let mut xc: Option<CudaSlice<f32>> = None;
for il in start..end {
let k = lin_pos[&il];
let xr: &CudaSlice<f32> = xc.as_ref().unwrap_or(xin);
let nx = model.qwen35_tparallel_linear_layer(
e,
il,
xr,
t,
cache_ref,
None,
Some((&mut stash_conv[k], &mut stash_ssm[k])),
Some((table_all, k * 6)),
)?;
xc = Some(nx);
}
e.copy_into(xout, 0, xc.as_ref().unwrap(), t * n_embd)?;
Ok(())
})?
};
if t % 2 == 1 {
for il in start..end {
let rl = cache.recur[il].as_mut().unwrap();
std::mem::swap(&mut rl.ssm_state, &mut rl.ssm_state_alt);
}
}
for (k, il) in (start..end).enumerate() {
let rl = cache.recur[il].as_mut().unwrap();
let (cw, sw) = (self.conv_words, self.ssm_words);
{
let sv = e.view(&self.save_conv, self.lin.len() * cw);
let win = sv.slice(k * cw..(k + 1) * cw);
e.copy_view_into(&mut rl.conv_state, 0, &win, cw)?;
}
{
let sv = e.view(&self.save_ssm, self.lin.len() * sw);
let win = sv.slice(k * sw..(k + 1) * sw);
e.copy_view_into(&mut rl.ssm_state, 0, &win, sw)?;
}
}
if std::env::var("MEMRA_GRAPH_CENSUS").as_deref() == Ok("1") {
if let Ok(c) = crate::graph_update::node_census(&graph) {
eprintln!("[dspark-vg-census] seg={start}..{end} vt={t} {c:?}");
}
}
self.graphs.insert(
key,
DsparkSegGraph {
graph,
_keeper: keeper,
},
);
}
self.graphs[&key].graph.launch()?;
if t % 2 == 1 {
for il in start..end {
let rl = cache.recur[il].as_mut().unwrap();
std::mem::swap(&mut rl.ssm_state, &mut rl.ssm_state_alt);
}
}
let (_, xout) = self.stage.get(&t).unwrap();
let mut out = e.uninit(t * n_embd)?;
e.copy_into(&mut out, 0, xout, t * n_embd)?;
Ok(out)
}
pub(crate) fn slab_row(
&self,
e: &Engine,
il: usize,
row: usize,
) -> Option<(u64, u64, usize, usize)> {
use cudarc::driver::DevicePtr;
let k = *self.lin_pos.get(&il)?;
let s = &e.gpu.stream();
let (pc, _g0) = self.stash_conv[k].device_ptr(s);
let (ps, _g1) = self.stash_ssm[k].device_ptr(s);
Some((
pc as u64 + (row * self.conv_words * 4) as u64,
ps as u64 + (row * self.ssm_words * 4) as u64,
self.conv_words,
self.ssm_words,
))
}
}
impl VerifyCkpt {
fn new(n_layer: usize) -> Self {
VerifyCkpt {
gdn: (0..n_layer).map(|_| None).collect(),
cols: (0..n_layer).map(|_| None).collect(),
}
}
}
struct VerifyBoundaryTicket {
rt: &'static crate::pp::PpNRt,
caller_stream: std::sync::Arc<cudarc::driver::CudaStream>,
slot: usize,
pos0: usize,
t: usize,
payload: usize,
n_st: usize,
pipelined: bool,
pp_anatomy: bool,
pp_started: std::time::Instant,
reverse_ms: f64,
stage0_ms: f64,
tx_ms: f64,
trace: Option<SpecPipeTraceCtx>,
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum OptiForkGateMode {
Disabled,
Hit,
Miss,
Alternate,
Abort,
Controller,
}
static OPTI_FORK_GATE_MODE: std::sync::atomic::AtomicU8 = std::sync::atomic::AtomicU8::new(0);
static OPTI_CONTROLLER_THRESHOLD: std::sync::atomic::AtomicU32 =
std::sync::atomic::AtomicU32::new(0);
static OPTI_FORK_ATTEMPTS: std::sync::atomic::AtomicU64 = std::sync::atomic::AtomicU64::new(0);
static OPTI_FORK_HITS: std::sync::atomic::AtomicU64 = std::sync::atomic::AtomicU64::new(0);
static OPTI_FORK_MISSES: std::sync::atomic::AtomicU64 = std::sync::atomic::AtomicU64::new(0);
static OPTI_FORK_ABORT_DRAINS: std::sync::atomic::AtomicU64 = std::sync::atomic::AtomicU64::new(0);
static OPTI_FORK_REFUSALS: std::sync::atomic::AtomicU64 = std::sync::atomic::AtomicU64::new(0);
static OPTI_GATE_CHECKS: std::sync::atomic::AtomicU64 = std::sync::atomic::AtomicU64::new(0);
static OPTI_GATE_ADMITS: std::sync::atomic::AtomicU64 = std::sync::atomic::AtomicU64::new(0);
static OPTI_GATE_REJECTS: std::sync::atomic::AtomicU64 = std::sync::atomic::AtomicU64::new(0);
static OPTI_RECONCILES: std::sync::atomic::AtomicU64 = std::sync::atomic::AtomicU64::new(0);
static OPTI_WASTED_DRAFT_TOKENS: std::sync::atomic::AtomicU64 =
std::sync::atomic::AtomicU64::new(0);
static OPTI_SHADOW_DRAFT_TOKENS: std::sync::atomic::AtomicU64 =
std::sync::atomic::AtomicU64::new(0);
static OPTI_BREAKER_TRIPS: std::sync::atomic::AtomicU64 = std::sync::atomic::AtomicU64::new(0);
impl OptiForkGateMode {
fn code(self) -> u8 {
match self {
Self::Disabled => 0,
Self::Hit => 1,
Self::Miss => 2,
Self::Alternate => 3,
Self::Abort => 4,
Self::Controller => 5,
}
}
fn configured() -> Self {
match OPTI_FORK_GATE_MODE.load(std::sync::atomic::Ordering::Relaxed) {
1 => Self::Hit,
2 => Self::Miss,
3 => Self::Alternate,
4 => Self::Abort,
5 => Self::Controller,
_ => Self::Disabled,
}
}
fn action(self, generation: u64) -> OptiForkAction {
match self {
Self::Hit => OptiForkAction::Hit,
Self::Miss => OptiForkAction::Miss,
Self::Alternate if generation & 1 == 0 => OptiForkAction::Hit,
Self::Alternate => OptiForkAction::Miss,
Self::Abort => OptiForkAction::Abort,
Self::Disabled | Self::Controller => {
unreachable!("non-forced mode cannot choose a forced fork action")
}
}
}
fn is_forced(self) -> bool {
matches!(self, Self::Hit | Self::Miss | Self::Alternate | Self::Abort)
}
}
pub fn set_optipipe_gate_mode(mode: OptiForkGateMode) {
OPTI_FORK_GATE_MODE.store(mode.code(), std::sync::atomic::Ordering::Relaxed);
}
pub fn set_optipipe_controller_threshold(threshold: f32) {
assert!(threshold.is_finite() && (0.0..=1.0).contains(&threshold));
OPTI_CONTROLLER_THRESHOLD.store(threshold.to_bits(), std::sync::atomic::Ordering::Relaxed);
set_optipipe_gate_mode(OptiForkGateMode::Controller);
}
#[derive(Clone, Copy, Debug, Default, PartialEq, Eq)]
pub struct OptiForkGateStats {
pub attempts: u64,
pub hits: u64,
pub misses: u64,
pub abort_drains: u64,
pub refusals: u64,
pub gate_checks: u64,
pub gate_admits: u64,
pub gate_rejects: u64,
pub reconciles: u64,
pub wasted_draft_tokens: u64,
pub shadow_draft_tokens: u64,
pub breaker_trips: u64,
}
#[derive(Clone, Copy, Debug, Default, PartialEq, Eq)]
pub struct OptiForkStateIdentity {
pub trunk_kv_bytes: usize,
pub recurrent_bytes: usize,
pub scratch_kv_bytes: usize,
pub hidden_bytes: usize,
}
pub fn reset_optipipe_gate_stats() {
for counter in [
&OPTI_FORK_ATTEMPTS,
&OPTI_FORK_HITS,
&OPTI_FORK_MISSES,
&OPTI_FORK_ABORT_DRAINS,
&OPTI_FORK_REFUSALS,
&OPTI_GATE_CHECKS,
&OPTI_GATE_ADMITS,
&OPTI_GATE_REJECTS,
&OPTI_RECONCILES,
&OPTI_WASTED_DRAFT_TOKENS,
&OPTI_SHADOW_DRAFT_TOKENS,
&OPTI_BREAKER_TRIPS,
] {
counter.store(0, std::sync::atomic::Ordering::Relaxed);
}
}
pub fn optipipe_gate_stats() -> OptiForkGateStats {
let load = |v: &std::sync::atomic::AtomicU64| v.load(std::sync::atomic::Ordering::Relaxed);
OptiForkGateStats {
attempts: load(&OPTI_FORK_ATTEMPTS),
hits: load(&OPTI_FORK_HITS),
misses: load(&OPTI_FORK_MISSES),
abort_drains: load(&OPTI_FORK_ABORT_DRAINS),
refusals: load(&OPTI_FORK_REFUSALS),
gate_checks: load(&OPTI_GATE_CHECKS),
gate_admits: load(&OPTI_GATE_ADMITS),
gate_rejects: load(&OPTI_GATE_REJECTS),
reconciles: load(&OPTI_RECONCILES),
wasted_draft_tokens: load(&OPTI_WASTED_DRAFT_TOKENS),
shadow_draft_tokens: load(&OPTI_SHADOW_DRAFT_TOKENS),
breaker_trips: load(&OPTI_BREAKER_TRIPS),
}
}
#[derive(Clone, Copy, Debug)]
struct OptiControllerPolicy {
threshold: f32,
consecutive_misses: u8,
breaker_tripped: bool,
}
impl OptiControllerPolicy {
fn configured() -> Self {
Self {
threshold: f32::from_bits(
OPTI_CONTROLLER_THRESHOLD.load(std::sync::atomic::Ordering::Relaxed),
),
consecutive_misses: 0,
breaker_tripped: false,
}
}
fn admit(&self, q_proxy: f32) -> bool {
q_proxy.is_finite()
&& (0.0..=1.0).contains(&q_proxy)
&& (self.threshold == 0.0 || (!self.breaker_tripped && q_proxy >= self.threshold))
}
fn resolve(&mut self, hit: bool) -> bool {
if self.threshold == 0.0 {
self.consecutive_misses = 0;
return false;
}
if hit {
self.consecutive_misses = 0;
return false;
}
self.consecutive_misses = self.consecutive_misses.saturating_add(1);
if !self.breaker_tripped && self.consecutive_misses >= 3 {
self.breaker_tripped = true;
return true;
}
false
}
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
enum OptiForkAction {
Hit,
Miss,
Abort,
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
struct OptiForkGeneration {
id: u64,
slot: usize,
}
#[derive(Default)]
struct OptiForkGenerationTracker {
next: u64,
live: [Option<u64>; 2],
}
impl OptiForkGenerationTracker {
fn reserve(&mut self) -> Result<OptiForkGeneration, Box<dyn std::error::Error>> {
let generation = OptiForkGeneration {
id: self.next,
slot: (self.next & 1) as usize,
};
if let Some(live) = self.live[generation.slot] {
return Err(format!(
"optipipe snapshot slot {} still owns generation {live}; refusing to overwrite it",
generation.slot,
)
.into());
}
self.next += 1;
self.live[generation.slot] = Some(generation.id);
Ok(generation)
}
fn retire(&mut self, generation: OptiForkGeneration) -> Result<(), Box<dyn std::error::Error>> {
match self.live[generation.slot] {
Some(id) if id == generation.id => {
self.live[generation.slot] = None;
Ok(())
}
other => Err(format!(
"optipipe generation teardown mismatch: ticket={} slot={} live={other:?}",
generation.id, generation.slot,
)
.into()),
}
}
}
struct OptiForkSeedGeneration {
h_seed: CudaSlice<f32>,
fill_prev: CudaSlice<f32>,
scratch_len: usize,
}
fn opti_snapshot_stage_owned(
e: &Engine,
cache: &Cache,
rt: &'static crate::pp::PpNRt,
fence: &[usize],
) -> Result<crate::cache::CacheSnapshot, Box<dyn std::error::Error>> {
let n = cache.kv.len();
let mut snapshot = crate::cache::CacheSnapshot {
kv_len: vec![None; n],
conv: (0..n).map(|_| None).collect(),
ssm: (0..n).map(|_| None).collect(),
pos: cache.pos,
};
opti_snapshot_stage_owned_into(e, cache, rt, fence, &mut snapshot)?;
Ok(snapshot)
}
fn opti_snapshot_stage_owned_into(
e: &Engine,
cache: &Cache,
rt: &'static crate::pp::PpNRt,
fence: &[usize],
snapshot: &mut crate::cache::CacheSnapshot,
) -> Result<(), Box<dyn std::error::Error>> {
if fence.len() != rt.n_stages() + 1 || snapshot.kv_len.len() != cache.kv.len() {
return Err("optipipe stage-owned snapshot shape mismatch".into());
}
for stage in 0..rt.n_stages() {
opti_snapshot_one_stage_owned_into(e, cache, rt, fence, stage, snapshot)?;
}
snapshot.pos = cache.pos;
Ok(())
}
fn opti_snapshot_one_stage_owned_into(
e: &Engine,
cache: &Cache,
rt: &'static crate::pp::PpNRt,
fence: &[usize],
stage: usize,
snapshot: &mut crate::cache::CacheSnapshot,
) -> Result<(), Box<dyn std::error::Error>> {
if fence.len() != rt.n_stages() + 1
|| snapshot.kv_len.len() != cache.kv.len()
|| stage >= rt.n_stages()
{
return Err("optipipe single-stage snapshot shape mismatch".into());
}
let _scope = rt.enter(stage);
let owner = rt.engine(stage, e);
for il in fence[stage]..fence[stage + 1] {
snapshot.kv_len[il] = cache.kv[il].as_ref().map(|kv| kv.len);
match &cache.recur[il] {
Some(recur) => {
match snapshot.conv[il].as_mut() {
Some(dst) => {
owner.copy_into(dst, 0, &recur.conv_state, recur.conv_state.len())?
}
None => snapshot.conv[il] = Some(owner.clone_dtod(&recur.conv_state)?),
}
match snapshot.ssm[il].as_mut() {
Some(dst) => {
owner.copy_into(dst, 0, &recur.ssm_state, recur.ssm_state.len())?
}
None => snapshot.ssm[il] = Some(owner.clone_dtod(&recur.ssm_state)?),
}
}
None if snapshot.conv[il].is_some() || snapshot.ssm[il].is_some() => {
return Err(
format!("optipipe stage-owned snapshot layer {il} changed shape").into(),
);
}
None => {}
}
}
snapshot.pos = cache.pos;
Ok(())
}
struct OptiForkState {
mode: OptiForkGateMode,
controller: Option<OptiControllerPolicy>,
generations: OptiForkGenerationTracker,
active_snapshot_slot: usize,
alternate_snapshot: crate::cache::CacheSnapshot,
seeds: [OptiForkSeedGeneration; 2],
rt: &'static crate::pp::PpNRt,
fence: [usize; 3],
split: usize,
len_ptrs: CudaSlice<u64>,
saved_lens: CudaSlice<i32>,
forced_acc: CudaSlice<u32>,
valid: CudaSlice<u32>,
stage0_stream: std::sync::Arc<cudarc::driver::CudaStream>,
logical_payload_bytes: [usize; 2],
}
struct OptiForkTicket {
generation: OptiForkGeneration,
boundary: Option<VerifyBoundaryTicket>,
drain: std::sync::Arc<cudarc::driver::CudaStream>,
settled: bool,
}
struct OptiControllerTicket {
generation: OptiForkGeneration,
boundary: Option<VerifyBoundaryTicket>,
ckpt: Option<VerifyCkpt>,
verify_tokens: [u32; 2],
draft_prob: f32,
eager_seed: Option<CudaSlice<f32>>,
q_proxy: f32,
scratch_len: usize,
issued_at: std::time::Instant,
drain: std::sync::Arc<cudarc::driver::CudaStream>,
settled: bool,
}
struct OptiControllerPrepared {
verify_tokens: [u32; 2],
draft_prob: f32,
eager_seed: Option<CudaSlice<f32>>,
q_proxy: f32,
scratch_len: usize,
}
impl OptiControllerTicket {
fn take_boundary(&mut self) -> VerifyBoundaryTicket {
self.boundary
.take()
.expect("controller boundary ticket already consumed")
}
fn take_ckpt(&mut self) -> VerifyCkpt {
self.ckpt
.take()
.expect("controller verify checkpoint already consumed")
}
fn take_eager_seed(&mut self) -> Option<CudaSlice<f32>> {
self.eager_seed.take()
}
fn settle(&mut self) {
self.settled = true;
}
}
impl Drop for OptiControllerTicket {
fn drop(&mut self) {
if !self.settled {
let _ = self.drain.synchronize();
OPTI_FORK_ABORT_DRAINS.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
}
}
}
impl OptiForkTicket {
fn take_boundary(&mut self) -> VerifyBoundaryTicket {
self.boundary
.take()
.expect("fork ticket boundary already consumed")
}
fn settle(&mut self) {
self.settled = true;
}
}
impl Drop for OptiForkTicket {
fn drop(&mut self) {
if !self.settled {
let _ = self.drain.synchronize();
OPTI_FORK_ABORT_DRAINS.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
}
}
}
impl OptiForkState {
#[allow(clippy::too_many_arguments)]
fn new(
e: &Engine,
cache: &Cache,
mode: OptiForkGateMode,
alternate_snapshot: crate::cache::CacheSnapshot,
h_seed: &CudaSlice<f32>,
fill_prev: &CudaSlice<f32>,
rt: &'static crate::pp::PpNRt,
split: usize,
n_layer: usize,
) -> Result<Self, Box<dyn std::error::Error>> {
let fence = [0, split, n_layer];
let mut logical_payload_bytes = [0usize; 2];
for stage in 0..2 {
for il in fence[stage]..fence[stage + 1] {
logical_payload_bytes[stage] += alternate_snapshot.conv[il]
.as_ref()
.map_or(0, |v| v.len() * std::mem::size_of::<f32>());
logical_payload_bytes[stage] += alternate_snapshot.ssm[il]
.as_ref()
.map_or(0, |v| v.len() * std::mem::size_of::<f32>());
}
}
let seeds = [
OptiForkSeedGeneration {
h_seed: e.clone_dtod(h_seed)?,
fill_prev: e.clone_dtod(fill_prev)?,
scratch_len: 0,
},
OptiForkSeedGeneration {
h_seed: e.clone_dtod(h_seed)?,
fill_prev: e.clone_dtod(fill_prev)?,
scratch_len: 0,
},
];
let (len_ptrs, saved_lens, forced_acc, valid, stage0_stream) = {
let _stage = rt.enter(0);
let e0 = rt.engine(0, e);
(
crate::round_stream::kv_len_ptr_table_range(e0, cache, 0..split, None)?,
e0.htod_i32(&vec![0; split])?,
e0.alloc_u32_zeroed(2)?,
e0.alloc_u32_zeroed(1)?,
e0.stream(),
)
};
logical_payload_bytes[0] += seeds
.iter()
.map(|seed| (seed.h_seed.len() + seed.fill_prev.len()) * std::mem::size_of::<f32>())
.sum::<usize>();
logical_payload_bytes[0] += len_ptrs.len() * std::mem::size_of::<u64>()
+ saved_lens.len() * std::mem::size_of::<i32>()
+ forced_acc.len() * std::mem::size_of::<u32>()
+ valid.len() * std::mem::size_of::<u32>();
Ok(Self {
mode,
controller: (mode == OptiForkGateMode::Controller)
.then(OptiControllerPolicy::configured),
generations: OptiForkGenerationTracker::default(),
active_snapshot_slot: 0,
alternate_snapshot,
seeds,
rt,
fence,
split,
len_ptrs,
saved_lens,
forced_acc,
valid,
stage0_stream,
logical_payload_bytes,
})
}
fn reserve(
&mut self,
current_snapshot: &mut crate::cache::CacheSnapshot,
) -> Result<OptiForkGeneration, Box<dyn std::error::Error>> {
let generation = self.generations.reserve()?;
if generation.slot != self.active_snapshot_slot {
std::mem::swap(current_snapshot, &mut self.alternate_snapshot);
self.active_snapshot_slot = generation.slot;
}
Ok(generation)
}
fn capture_seed(
&mut self,
e: &Engine,
generation: OptiForkGeneration,
h_seed: &CudaSlice<f32>,
fill_prev: &CudaSlice<f32>,
scratch_len: usize,
) -> Result<(), Box<dyn std::error::Error>> {
let seed = &mut self.seeds[generation.slot];
e.copy_into(&mut seed.h_seed, 0, h_seed, h_seed.len())?;
e.copy_into(&mut seed.fill_prev, 0, fill_prev, fill_prev.len())?;
seed.scratch_len = scratch_len;
Ok(())
}
fn ticket(
&self,
generation: OptiForkGeneration,
boundary: VerifyBoundaryTicket,
) -> OptiForkTicket {
OptiForkTicket {
generation,
boundary: Some(boundary),
drain: self.stage0_stream.clone(),
settled: false,
}
}
#[allow(clippy::too_many_arguments)]
fn controller_ticket(
&self,
generation: OptiForkGeneration,
boundary: VerifyBoundaryTicket,
ckpt: VerifyCkpt,
verify_tokens: [u32; 2],
draft_prob: f32,
eager_seed: Option<CudaSlice<f32>>,
q_proxy: f32,
scratch_len: usize,
) -> OptiControllerTicket {
OptiControllerTicket {
generation,
boundary: Some(boundary),
ckpt: Some(ckpt),
verify_tokens,
draft_prob,
eager_seed,
q_proxy,
scratch_len,
issued_at: std::time::Instant::now(),
drain: self.stage0_stream.clone(),
settled: false,
}
}
fn reserve_successor(&mut self) -> Result<OptiForkGeneration, Box<dyn std::error::Error>> {
self.generations.reserve()
}
fn successor_snapshot_mut(&mut self) -> &mut crate::cache::CacheSnapshot {
&mut self.alternate_snapshot
}
fn promote_successor_snapshot(
&mut self,
current_snapshot: &mut crate::cache::CacheSnapshot,
generation: OptiForkGeneration,
) {
std::mem::swap(current_snapshot, &mut self.alternate_snapshot);
self.active_snapshot_slot = generation.slot;
}
fn queue_actual_reconcile(
&mut self,
e: &Engine,
snapshot: &crate::cache::CacheSnapshot,
acc: &CudaSlice<u32>,
optimistic_pending: u32,
base: usize,
) -> Result<(), Box<dyn std::error::Error>> {
let saved: Vec<i32> = (0..self.split)
.map(|il| snapshot.kv_len[il].map(|v| v as i32).unwrap_or(0))
.collect();
if self.rt.engine(0, e).ctx().ordinal() != e.ctx().ordinal() {
self.rt.fence_stages_behind(&e.stream())?;
}
let _stage = self.rt.enter(0);
let e0 = self.rt.engine(0, e);
e0.htod_i32_into(&mut self.saved_lens, &saved)?;
e0.spec_fork_valid(acc, optimistic_pending, &mut self.valid)?;
e0.spec_fork_reconcile_kv(
&self.len_ptrs,
&self.saved_lens,
acc,
&self.valid,
base,
self.split,
)
}
fn finish_actual_reconcile(
&mut self,
e: &Engine,
cache: &mut Cache,
snapshot: &crate::cache::CacheSnapshot,
n_acc: usize,
base: usize,
hit: bool,
) -> Result<(), Box<dyn std::error::Error>> {
if hit {
return Ok(());
}
let len_delta = base + n_acc;
for il in 0..self.split {
if let (Some(kv), Some(saved)) = (cache.kv[il].as_mut(), snapshot.kv_len[il]) {
kv.len = saved + len_delta;
}
}
{
let _stage = self.rt.enter(1);
let e1 = self.rt.engine(1, e);
for il in self.split..self.fence[2] {
if let (Some(kv), Some(saved)) = (cache.kv[il].as_mut(), snapshot.kv_len[il]) {
kv.len = saved + len_delta;
e1.set_i32_one(&mut kv.len_d, kv.len as i32)?;
}
}
}
self.rt.publish_to(0, &e.stream())?;
Ok(())
}
fn cancel_controller_ticket(
&mut self,
e: &Engine,
cache: &mut Cache,
scratch: &mut MtpScratch,
snapshot: &crate::cache::CacheSnapshot,
ticket: &mut OptiControllerTicket,
) -> Result<(), Box<dyn std::error::Error>> {
{
let _stage = self.rt.enter(0);
let e0 = self.rt.engine(0, e);
for il in 0..self.split {
if let (Some(kv), Some(saved)) = (cache.kv[il].as_mut(), snapshot.kv_len[il]) {
kv.len = saved;
e0.set_i32_one(&mut kv.len_d, saved as i32)?;
}
}
}
scratch.set_len(e, snapshot.pos)?;
ticket.settle();
self.generations.retire(ticket.generation)?;
OPTI_FORK_ABORT_DRAINS.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
OPTI_WASTED_DRAFT_TOKENS.fetch_add(2, std::sync::atomic::Ordering::Relaxed);
eprintln!(
"[opti-controller] tail-drain generation={} slot={}",
ticket.generation.id, ticket.generation.slot,
);
Ok(())
}
#[allow(clippy::too_many_arguments)]
fn reconcile(
&mut self,
e: &Engine,
cache: &mut Cache,
scratch: &mut MtpScratch,
snapshot: &crate::cache::CacheSnapshot,
h_seed: &mut CudaSlice<f32>,
fill_prev: &mut CudaSlice<f32>,
generation: OptiForkGeneration,
action: OptiForkAction,
optimistic_pending: u32,
) -> Result<(), Box<dyn std::error::Error>> {
debug_assert!(action != OptiForkAction::Abort);
let miss_started = std::time::Instant::now();
let keep = action == OptiForkAction::Hit;
let saved: Vec<i32> = (0..self.split)
.map(|il| snapshot.kv_len[il].map(|v| v as i32).unwrap_or(0))
.collect();
let seed = &self.seeds[generation.slot];
{
let _stage = self.rt.enter(0);
let e0 = self.rt.engine(0, e);
e0.htod_i32_into(&mut self.saved_lens, &saved)?;
let forced = if keep {
[1u32, optimistic_pending]
} else {
[0u32, optimistic_pending]
};
e0.htod_u32_into(&mut self.forced_acc, &forced)?;
e0.spec_fork_valid(&self.forced_acc, optimistic_pending, &mut self.valid)?;
e0.spec_fork_reconcile_kv(
&self.len_ptrs,
&self.saved_lens,
&self.forced_acc,
&self.valid,
0,
self.split,
)?;
for il in 0..self.split {
if let Some(recur) = cache.recur[il].as_mut() {
let conv = snapshot.conv[il]
.as_ref()
.ok_or("optipipe stage0 snapshot missing conv state")?;
let ssm = snapshot.ssm[il]
.as_ref()
.ok_or("optipipe stage0 snapshot missing ssm state")?;
e0.spec_fork_restore_f32(conv, &mut recur.conv_state, &self.valid)?;
e0.spec_fork_restore_f32(ssm, &mut recur.ssm_state, &self.valid)?;
}
}
e0.spec_fork_restore_f32(&seed.h_seed, h_seed, &self.valid)?;
e0.spec_fork_restore_f32(&seed.fill_prev, fill_prev, &self.valid)?;
}
if keep {
OPTI_FORK_HITS.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
return Ok(());
}
for il in 0..self.split {
if let (Some(kv), Some(saved)) = (cache.kv[il].as_mut(), snapshot.kv_len[il]) {
kv.len = saved;
}
}
scratch.set_len(e, seed.scratch_len)?;
let caller = e.stream();
self.rt.publish_to(0, &caller)?;
caller.synchronize()?;
let miss_ms = miss_started.elapsed().as_secs_f64() * 1e3;
eprintln!(
"[opti-fork-reconcile] generation={} slot={} miss_ms={miss_ms:.3}",
generation.id, generation.slot,
);
OPTI_FORK_MISSES.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
Ok(())
}
fn retire(&mut self, generation: OptiForkGeneration) -> Result<(), Box<dyn std::error::Error>> {
self.generations.retire(generation)
}
}
impl HybridModel {
fn opti_graph_draft_step(
&self,
e: &Engine,
mtp: &MtpHead,
dctx: &mut DraftGraphCtx,
scratch: &mut MtpScratch,
d_vocab: usize,
) -> Result<(u32, f32), Box<dyn std::error::Error>> {
dctx.graph
.as_ref()
.ok_or("optipipe controller requires the greedy draft graph")?
.launch()?;
scratch.kv.len += 1;
let idx = e.dtoh_u32_one(&dctx.g_tok)?;
if (idx as usize) >= d_vocab {
return Err(
format!("optipipe draft argmax sentinel 0x{idx:08x} >= d_vocab {d_vocab}").into(),
);
}
let probability = e.dtoh(&dctx.g_p)?[0];
if !probability.is_finite() || !(0.0..=1.0).contains(&probability) {
return Err(format!("optipipe draft probability is invalid: {probability}").into());
}
let token = match &mtp.d2t {
Some(map) => map[idx as usize],
None => idx,
};
if token != idx {
e.set_u32_one(&mut dctx.g_tok, token)?;
}
Ok((token, probability))
}
#[allow(clippy::too_many_arguments)]
fn opti_controller_draft_step(
&self,
e: &Engine,
mtp: &MtpHead,
dctx: &mut DraftGraphCtx,
scratch: &mut MtpScratch,
d_vocab: usize,
eager_state: &mut Option<(u32, CudaSlice<f32>)>,
eager_pos: usize,
embd_dev: Option<(&CudaSlice<u8>, i32, usize)>,
) -> Result<(u32, f32), Box<dyn std::error::Error>> {
if dctx.graph.is_some() {
return self.opti_graph_draft_step(e, mtp, dctx, scratch, d_vocab);
}
let (input_token, input_seed) = eager_state
.take()
.ok_or("optipipe eager continuation seed is unavailable")?;
let (logits, next_seed) = self.mtp_head_forward_dev(
e,
mtp,
input_token,
&input_seed,
scratch,
eager_pos,
embd_dev,
None,
)?;
let token_d = e.argmax_token_device(&logits, d_vocab)?;
let idx = e.dtoh_u32_one(&token_d)?;
if (idx as usize) >= d_vocab {
return Err(format!(
"optipipe eager draft argmax sentinel 0x{idx:08x} >= d_vocab {d_vocab}"
)
.into());
}
let probability_d = e.prob_of_token_device(&logits, &token_d, d_vocab)?;
let probability = e.dtoh(&probability_d)?[0];
if !probability.is_finite() || !(0.0..=1.0).contains(&probability) {
return Err(
format!("optipipe eager draft probability is invalid: {probability}").into(),
);
}
let token = match &mtp.d2t {
Some(map) => map[idx as usize],
None => idx,
};
*eager_state = Some((token, next_seed));
Ok((token, probability))
}
#[allow(clippy::too_many_arguments)]
fn mtp_head_forward_dev(
&self,
e: &Engine,
mtp: &MtpHead,
e_tok: u32,
h_seed: &CudaSlice<f32>,
scratch: &mut MtpScratch,
mtp_pos: usize,
embd_dev: Option<(&CudaSlice<u8>, i32, usize)>,
mask: Option<(&CudaSlice<u32>, usize)>,
) -> Result<(CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
use std::sync::atomic::{AtomicU64, Ordering::Relaxed};
static ANAT_NS: [AtomicU64; 5] = [
AtomicU64::new(0),
AtomicU64::new(0),
AtomicU64::new(0),
AtomicU64::new(0),
AtomicU64::new(0),
];
static ANAT_STEPS: AtomicU64 = AtomicU64::new(0);
let anat = {
static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
*ON.get_or_init(|| std::env::var("MEMRA_SPEC_ANATOMY").as_deref() == Ok("1"))
};
if anat {
e.stream().synchronize()?; }
let t_all = std::time::Instant::now();
let mut t_ph = std::time::Instant::now();
let mut anat_mark = |i: usize,
e: &Engine,
t: &mut std::time::Instant|
-> Result<(), Box<dyn std::error::Error>> {
if anat {
e.stream().synchronize()?;
ANAT_NS[i].fetch_add(t.elapsed().as_nanos() as u64, Relaxed);
*t = std::time::Instant::now();
}
Ok(())
};
let cfg = &self.cfg;
let n_embd = cfg.n_embd as usize;
let di = mtp.geom.as_ref().map(|g| g.d_inner).unwrap_or(n_embd);
let eps = cfg.rms_eps;
let pos_d = e.htod_i32(&[mtp_pos as i32])?;
let e_emb = match embd_dev {
Some((g, qt, rb)) => e.embed_gather_device_t(g, &[e_tok], n_embd, qt, rb)?,
None => e.htod(&self.embd.gather(n_embd, &[e_tok]))?,
};
let mut e_norm = e.zeros(n_embd)?;
e.rms_norm(&e_emb, mtp.enorm.float_data(), &mut e_norm, n_embd, 1, eps)?;
let mut h_norm = e.zeros(n_embd)?;
e.rms_norm(h_seed, mtp.hnorm.float_data(), &mut h_norm, n_embd, 1, eps)?;
let mut concat = e.zeros(2 * n_embd)?;
e.copy_into(&mut concat, 0, &e_norm, n_embd)?;
e.copy_into(&mut concat, n_embd, &h_norm, n_embd)?;
let inp_sa = e.matmul(&mtp.eh_proj, &concat, 1)?;
let mut a_norm = e.zeros(di)?;
e.rms_norm(&inp_sa, mtp.attn_norm.float_data(), &mut a_norm, di, 1, eps)?;
anat_mark(0, e, &mut t_ph)?;
let attn_out = match (&mtp.mixer, mtp.step35.as_ref()) {
(Mixer::Full(fa), Some(g)) => {
self.mtp_step35_attn(e, fa, g, &a_norm, &pos_d, scratch)?
}
(Mixer::Full(fa), None) => {
let out =
self.mtp_full_attn_dc(e, fa, &a_norm, &pos_d, scratch, mtp.geom.as_ref())?;
scratch.kv.len += 1;
out
}
(Mixer::Linear(_), _) => {
panic!("MTP block is full-attn in qwen35; linear MTP not supported")
}
(Mixer::Mla(_), _) => crate::hybrid::mla_forward_unimplemented(),
};
anat_mark(1, e, &mut t_ph)?;
let mut x1 = e.zeros(di)?;
e.add(&inp_sa, &attn_out, &mut x1, di)?;
let mut z = e.zeros(di)?;
e.rms_norm(&x1, mtp.post_attn_norm.float_data(), &mut z, di, 1, eps)?;
let ffn_out = match &mtp.ffn {
crate::hybrid::Ffn::Dense {
ffn_gate,
ffn_up,
ffn_down,
} => {
let n_ff = ffn_gate.out_features();
let (gate, up) = if e.uses_q8_1_fast(ffn_gate) && e.uses_q8_1_fast(ffn_up) {
let (zq, zd) = e.quantize_q8_1(&z, 1, di)?;
(
e.matmul_pre(ffn_gate, &zq, &zd, &z, 1)?,
e.matmul_pre(ffn_up, &zq, &zd, &z, 1)?,
)
} else {
(e.matmul(ffn_gate, &z, 1)?, e.matmul(ffn_up, &z, 1)?)
};
let mut act = e.zeros(n_ff)?;
Self::ffn_act_lim(
e,
&self.cfg,
&gate,
&up,
1.0,
1.0,
mtp.step35.as_ref().and_then(|s| s.clamp_shexp),
&mut act,
n_ff,
)?;
e.matmul(ffn_down, &act, 1)?
}
crate::hybrid::Ffn::Moe(m) => self.moe_ffn_il(e, m, &z, 1, u16::MAX)?,
};
anat_mark(2, e, &mut t_ph)?;
let mut h_inner = e.zeros(di)?;
e.add(&x1, &ffn_out, &mut h_inner, di)?;
let h_nextn = match mtp.geom.as_ref() {
Some(g) => e.matmul(&g.out_up, &h_inner, 1)?,
None => h_inner,
};
let final_norm = mtp.shared_head_norm.as_ref().unwrap_or(&self.output_norm);
let mut final_h = e.zeros(n_embd)?;
e.rms_norm(
&h_nextn,
final_norm.float_data(),
&mut final_h,
n_embd,
1,
eps,
)?;
let head = mtp.shared_head_head.as_ref().unwrap_or(&self.output);
let mut logits = e.matmul(head, &final_h, 1)?;
if let Some((mask_d, mw)) = mask {
let d_vocab = head.out_features();
e.mask_logits_col(&mut logits, mask_d, 0, d_vocab, mw)?;
}
anat_mark(3, e, &mut t_ph)?;
if anat {
ANAT_NS[4].fetch_add(t_all.elapsed().as_nanos() as u64, Relaxed);
let n = ANAT_STEPS.fetch_add(1, Relaxed) + 1;
if n % 128 == 0 {
let us = |i: usize| ANAT_NS[i].load(Relaxed) / n / 1000;
eprintln!(
"[spec-anatomy] steps={n} avg us/step: glue={} attn={} ffn={} head={} total={}",
us(0),
us(1),
us(2),
us(3),
us(4)
);
}
}
Ok((logits, if spec_hpost() { final_h } else { h_nextn }))
}
fn mtp_step35_attn(
&self,
e: &Engine,
fa: &FullAttnLayer,
g: &crate::hybrid::Step35MtpGeom,
h: &CudaSlice<f32>,
pos_d: &CudaSlice<i32>,
scratch: &mut MtpScratch,
) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
let (nh, nkv, hd) = (g.n_head, g.n_head_kv, self.cfg.head_dim_k as usize);
let eps = self.cfg.rms_eps;
let scale = 1.0 / (hd as f32).sqrt(); let n_embd = self.cfg.n_embd as usize;
let gw = fa
.attn_gate
.as_ref()
.ok_or("step35 MTP block is missing attn_gate.weight (head-wise attention gate)")?;
let (q0, k0, v0, gt) = if e.uses_q8_1_fast(&fa.wq)
&& e.uses_q8_1_fast(&fa.wk)
&& e.uses_q8_1_fast(&fa.wv)
&& e.uses_q8_1_fast(gw)
{
let (hq, hdq) = e.quantize_q8_1(h, 1, n_embd)?;
let (a, b, c) = match e.matmul_q8_fused3(&fa.wq, &fa.wk, &fa.wv, &hq, &hdq)? {
Some(t3) => t3,
None => (
e.matmul_pre(&fa.wq, &hq, &hdq, h, 1)?,
e.matmul_pre(&fa.wk, &hq, &hdq, h, 1)?,
e.matmul_pre(&fa.wv, &hq, &hdq, h, 1)?,
),
};
(a, b, c, e.matmul_pre(gw, &hq, &hdq, h, 1)?)
} else {
(
e.matmul(&fa.wq, h, 1)?,
e.matmul(&fa.wk, h, 1)?,
e.matmul(&fa.wv, h, 1)?,
e.matmul(gw, h, 1)?,
)
};
let mut q = e.uninit(nh * hd)?;
e.rms_norm(&q0, fa.q_norm.float_data(), &mut q, hd, nh, eps)?;
let mut k = e.uninit(nkv * hd)?;
e.rms_norm(&k0, fa.k_norm.float_data(), &mut k, hd, nkv, eps)?;
let ff = if g.swa {
None
} else {
self.step35_aux.as_ref().and_then(|a| a.rope_freqs(e))
};
#[cfg(debug_assertions)]
if let Some(ff) = ff {
crate::debug_assert_tensor_stream_device(ff, &e.stream(), "mtp_step35_attn.rope_freqs");
}
e.rope_neox2(
&mut q,
&mut k,
pos_d,
hd,
g.n_rot,
nh,
nkv,
1,
g.rope_base,
1.0,
ff,
)?;
let kv = &mut scratch.kv;
assert!(
kv.len < scratch.cap,
"step35 MTP scratch overflow ({} >= {})",
kv.len,
scratch.cap
);
let next_len = kv.len + 1;
let (off, t_kv) = if g.swa && next_len > g.window {
(next_len - g.window, g.window)
} else {
(0, next_len)
};
let write_row = e.prepare_kv_append(kv, off & !31usize, 1)?;
e.append_kv_quantized(
&k,
&v0,
&mut kv.k,
&mut kv.v,
write_row,
kv.kv_dim_k,
kv.kv_dim_v,
kv.k_tok_bytes,
kv.v_tok_bytes,
false,
)?;
kv.len = next_len;
e.set_i32_one(&mut kv.len_d, kv.len as i32)?;
let physical = kv.physical_rows(off, off + t_kv)?;
let k_view = e.view_u8_range(
&kv.k,
physical.start * kv.k_tok_bytes,
physical.end * kv.k_tok_bytes,
);
let v_view = e.view_u8_range(
&kv.v,
physical.start * kv.v_tok_bytes,
physical.end * kv.v_tok_bytes,
);
let mut attn = e.uninit(nh * hd)?;
e.fa_decode_kvmod(
&q,
&k_view,
&v_view,
&mut attn,
hd,
nh,
nkv,
t_kv,
scale,
kv.k_tok_bytes,
kv.v_tok_bytes,
false,
)?;
let mut ag = e.uninit(nh * hd)?;
e.attn_head_gate(&attn, >, &mut ag, None, hd, nh, 1)?;
Ok(e.matmul(&fa.wo, &ag, 1)?)
}
fn mtp_full_attn_dc(
&self,
e: &Engine,
fa: &FullAttnLayer,
h: &CudaSlice<f32>,
pos_d: &CudaSlice<i32>,
scratch: &mut MtpScratch,
geom: Option<&crate::hybrid::DraftGeom>,
) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
let cfg = &self.cfg;
let mtp_il = cfg.n_layer.saturating_sub(cfg.nextn_predict_layers);
let geometry = cfg.full_attention_geometry_at(mtp_il);
let n_head = geom.map(|g| g.n_head).unwrap_or(geometry.n_head as usize);
let n_head_kv = geom
.map(|g| g.n_head_kv)
.unwrap_or(geometry.n_head_kv as usize);
let head_dim = geometry.head_dim_k as usize;
let eps = cfg.rms_eps;
let scale = geometry.attention_scale();
let n_embd = geom.map(|g| g.d_inner).unwrap_or(cfg.n_embd as usize);
let bucket_max = scratch.cap;
let (qf, mut k, v) =
if e.uses_q8_1_fast(&fa.wq) && e.uses_q8_1_fast(&fa.wk) && e.uses_q8_1_fast(&fa.wv) {
let (hq, hd) = e.quantize_q8_1(h, 1, n_embd)?;
(
e.matmul_pre(&fa.wq, &hq, &hd, h, 1)?,
e.matmul_pre(&fa.wk, &hq, &hd, h, 1)?,
e.matmul_pre(&fa.wv, &hq, &hd, h, 1)?,
)
} else {
(
e.matmul(&fa.wq, h, 1)?,
e.matmul(&fa.wk, h, 1)?,
e.matmul(&fa.wv, h, 1)?,
)
};
let gated = geometry.attention_gate == memra_gguf::config::AttentionGateKind::FusedQ;
let (mut q, gate) = if gated {
let mut q = e.zeros(n_head * head_dim)?;
let mut gate = e.zeros(n_head * head_dim)?;
e.q_gate_split(&qf, &mut q, &mut gate, head_dim, n_head, 1)?;
(q, Some(gate))
} else {
(qf, None)
};
let mut qn = e.zeros(n_head * head_dim)?;
e.rms_norm(&q, fa.q_norm.float_data(), &mut qn, head_dim, n_head, eps)?;
q = qn;
let mut kn = e.zeros(n_head_kv * head_dim)?;
e.rms_norm(
&k,
fa.k_norm.float_data(),
&mut kn,
head_dim,
n_head_kv,
eps,
)?;
k = kn;
let rope_dims = geometry.n_rot as usize;
e.rope_neox(
&mut q,
pos_d,
head_dim,
rope_dims,
n_head,
1,
geometry.rope_base,
1.0,
)?;
e.rope_neox(
&mut k,
pos_d,
head_dim,
rope_dims,
n_head_kv,
1,
geometry.rope_base,
1.0,
)?;
let kv = &mut scratch.kv;
e.append_kv_quantized_dc(
&k,
&v,
&mut kv.k,
&mut kv.v,
&kv.len_d,
kv.kv_dim_k,
kv.kv_dim_v,
kv.k_tok_bytes,
kv.v_tok_bytes,
false,
)?;
e.inc_seqlen(&mut kv.len_d)?;
let k_view = e.view_u8(&kv.k, kv.k.len());
let v_view = e.view_u8(&kv.v, kv.v.len());
let (ktb, vtb) = (kv.k_tok_bytes, kv.v_tok_bytes);
let mut attn = e.zeros(n_head * head_dim)?;
e.fa_decode_dc(
&q, &k_view, &v_view, &mut attn, head_dim, n_head, n_head_kv, &kv.len_d, bucket_max,
scale, ktb, vtb, false,
)?;
let attn_g = match &gate {
Some(gate) => {
let mut gsig = e.zeros(n_head * head_dim)?;
e.sigmoid(gate, &mut gsig, n_head * head_dim)?;
let mut ag = e.zeros(n_head * head_dim)?;
e.mul(&attn, &gsig, &mut ag, n_head * head_dim)?;
ag
}
None => attn,
};
Ok(e.matmul(&fa.wo, &attn_g, 1)?)
}
#[allow(clippy::too_many_arguments)]
fn mtp_kv_fill(
&self,
e: &Engine,
mtp: &MtpHead,
tokens: &[u32],
h: &CudaSlice<f32>,
pos0: usize,
scratch: &mut MtpScratch,
embd_dev: Option<(&CudaSlice<u8>, i32, usize)>,
) -> Result<(), Box<dyn std::error::Error>> {
let cfg = &self.cfg;
let n_embd = cfg.n_embd as usize;
let eps = cfg.rms_eps;
let t = tokens.len();
assert_eq!(scratch.kv.len, pos0, "mtp_kv_fill: append slot mismatch");
assert!(pos0 + t <= scratch.cap, "mtp_kv_fill: scratch overflow");
let Mixer::Full(fa) = &mtp.mixer else {
panic!("MTP block is full-attn in qwen35; linear MTP not supported")
};
let pos_vec: Vec<i32> = (0..t).map(|i| (pos0 + i + 1) as i32).collect();
let pos_d = e.htod_i32(&pos_vec)?;
let e_emb = match embd_dev {
Some((g, qt, rb)) => e.embed_gather_device_t(g, tokens, n_embd, qt, rb)?,
None => e.htod(&self.embd.gather(n_embd, tokens))?,
};
let mut e_norm = e.zeros(t * n_embd)?;
e.rms_norm(&e_emb, mtp.enorm.float_data(), &mut e_norm, n_embd, t, eps)?;
let mut h_norm = e.zeros(t * n_embd)?;
e.rms_norm(h, mtp.hnorm.float_data(), &mut h_norm, n_embd, t, eps)?;
let mut concat = e.zeros(t * 2 * n_embd)?;
for i in 0..t {
e.copy_view_into(
&mut concat,
i * 2 * n_embd,
&e_norm.slice(i * n_embd..(i + 1) * n_embd),
n_embd,
)?;
e.copy_view_into(
&mut concat,
i * 2 * n_embd + n_embd,
&h_norm.slice(i * n_embd..(i + 1) * n_embd),
n_embd,
)?;
}
let di = mtp.geom.as_ref().map(|g| g.d_inner).unwrap_or(n_embd);
let inp_sa = e.matmul(&mtp.eh_proj, &concat, t)?;
let mut a_norm = e.zeros(t * di)?;
e.rms_norm(&inp_sa, mtp.attn_norm.float_data(), &mut a_norm, di, t, eps)?;
let n_head_kv = mtp
.geom
.as_ref()
.map(|g| g.n_head_kv)
.or(mtp.step35.as_ref().map(|s| s.n_head_kv))
.unwrap_or_else(|| {
let mtp_il = cfg.n_layer.saturating_sub(cfg.nextn_predict_layers);
cfg.full_attention_geometry_at(mtp_il).n_head_kv as usize
});
let mtp_il = cfg.n_layer.saturating_sub(cfg.nextn_predict_layers);
let geometry = cfg.full_attention_geometry_at(mtp_il);
let head_dim = geometry.head_dim_k as usize;
let mut k = e.matmul(&fa.wk, &a_norm, t)?;
let v = e.matmul(&fa.wv, &a_norm, t)?;
let mut kn = e.zeros(t * n_head_kv * head_dim)?;
e.rms_norm(
&k,
fa.k_norm.float_data(),
&mut kn,
head_dim,
n_head_kv * t,
eps,
)?;
k = kn;
let (rope_dims, rope_base, ff) = match mtp.step35.as_ref() {
Some(s) => (
s.n_rot,
s.rope_base,
if s.swa {
None
} else {
self.step35_aux.as_ref().and_then(|a| a.rope_freqs(e))
},
),
None => (geometry.n_rot as usize, geometry.rope_base, None),
};
#[cfg(debug_assertions)]
if let Some(ff) = ff {
crate::debug_assert_tensor_stream_device(ff, &e.stream(), "mtp_kv_fill.rope_freqs");
}
match ff {
Some(f) => e.rope_neox_ff(
&mut k, &pos_d, head_dim, rope_dims, n_head_kv, t, rope_base, 1.0, f,
)?,
None => e.rope_neox(
&mut k, &pos_d, head_dim, rope_dims, n_head_kv, t, rope_base, 1.0,
)?,
}
let kv = &mut scratch.kv;
let retain_from = kv
.ring
.as_ref()
.map(|ring| pos0.saturating_sub(ring.window() - 1) & !31usize)
.unwrap_or(0);
let write_row = e.prepare_kv_append(kv, retain_from, t)?;
for i in 0..t {
let k_row = k.slice(i * kv.kv_dim_k..(i + 1) * kv.kv_dim_k);
let v_row = v.slice(i * kv.kv_dim_v..(i + 1) * kv.kv_dim_v);
e.append_kv_quantized_view(
&k_row,
&v_row,
&mut kv.k,
&mut kv.v,
write_row + i,
kv.kv_dim_k,
kv.kv_dim_v,
kv.k_tok_bytes,
kv.v_tok_bytes,
false,
)?;
}
kv.len = pos0 + t;
e.set_i32_one(&mut kv.len_d, kv.len as i32)?;
Ok(())
}
#[allow(clippy::too_many_arguments)]
fn mtp_head_forward_cap(
&self,
e: &Engine,
mtp: &MtpHead,
tok_d: &mut CudaSlice<u32>,
pos_d: &mut CudaSlice<i32>,
h_seed_d: &mut CudaSlice<f32>,
p_d: &mut CudaSlice<f32>,
scratch: &mut MtpScratch,
with_prob: bool,
with_head: bool,
embd_gpu: &CudaSlice<u8>,
embd_qt: i32,
embd_rb: usize,
d_vocab: usize,
sampled_cap: Option<(
&mut CudaSlice<u32>,
&mut CudaSlice<f32>,
&mut CudaSlice<f32>,
u64,
f32,
)>,
stream_pack: Option<(&mut CudaSlice<u32>, usize, Option<&CudaSlice<u32>>)>,
mask_cap: Option<(&CudaSlice<u32>, usize)>,
) -> Result<(), Box<dyn std::error::Error>> {
let cfg = &self.cfg;
let n_embd = cfg.n_embd as usize;
if mtp.step35.is_some() {
return Err(
"step35 has no captured draft chain (fa_decode_dc cannot express the MTP \
block's SWA view offset; same root cause as the dc decode refusal) — the \
eager draft chain serves this arch"
.into(),
);
}
let di = mtp.geom.as_ref().map(|g| g.d_inner).unwrap_or(n_embd);
let eps = cfg.rms_eps;
let e_emb = e.embed_gather_device(embd_gpu, tok_d, n_embd, embd_qt, embd_rb)?;
let mut e_norm = e.zeros(n_embd)?;
e.rms_norm(&e_emb, mtp.enorm.float_data(), &mut e_norm, n_embd, 1, eps)?;
let mut h_norm = e.zeros(n_embd)?;
e.rms_norm(
&*h_seed_d,
mtp.hnorm.float_data(),
&mut h_norm,
n_embd,
1,
eps,
)?;
let mut concat = e.zeros(2 * n_embd)?;
e.copy_into(&mut concat, 0, &e_norm, n_embd)?;
e.copy_into(&mut concat, n_embd, &h_norm, n_embd)?;
let inp_sa = e.matmul(&mtp.eh_proj, &concat, 1)?;
let mut a_norm = e.zeros(di)?;
e.rms_norm(&inp_sa, mtp.attn_norm.float_data(), &mut a_norm, di, 1, eps)?;
let attn_out = match &mtp.mixer {
Mixer::Full(fa) => {
self.mtp_full_attn_dc(e, fa, &a_norm, pos_d, scratch, mtp.geom.as_ref())?
}
Mixer::Linear(_) => {
panic!("MTP block is full-attn in qwen35; linear MTP not supported")
}
Mixer::Mla(_) => crate::hybrid::mla_forward_unimplemented(),
};
let mut x1 = e.zeros(di)?;
e.add(&inp_sa, &attn_out, &mut x1, di)?;
let mut z = e.zeros(di)?;
e.rms_norm(&x1, mtp.post_attn_norm.float_data(), &mut z, di, 1, eps)?;
let ffn_out = match &mtp.ffn {
crate::hybrid::Ffn::Dense {
ffn_gate,
ffn_up,
ffn_down,
} => {
let n_ff = ffn_gate.out_features();
let (gate, up) = if e.uses_q8_1_fast(ffn_gate) && e.uses_q8_1_fast(ffn_up) {
let (zq, zd) = e.quantize_q8_1(&z, 1, di)?;
(
e.matmul_pre(ffn_gate, &zq, &zd, &z, 1)?,
e.matmul_pre(ffn_up, &zq, &zd, &z, 1)?,
)
} else {
(e.matmul(ffn_gate, &z, 1)?, e.matmul(ffn_up, &z, 1)?)
};
let mut act = e.zeros(n_ff)?;
Self::ffn_act(e, &self.cfg, &gate, &up, &mut act, n_ff)?;
e.matmul(ffn_down, &act, 1)?
}
crate::hybrid::Ffn::Moe(m) if m.dev_exps.is_some() => {
self.moe_ffn_il(e, m, &z, 1, u16::MAX)?
}
crate::hybrid::Ffn::Moe(_) => {
return Err("graph draft requires a Dense (or resident-MoE) MTP FFN".into());
}
};
let mut h_inner = e.zeros(di)?;
e.add(&x1, &ffn_out, &mut h_inner, di)?;
let h_nextn = match mtp.geom.as_ref() {
Some(g) => e.matmul(&g.out_up, &h_inner, 1)?,
None => h_inner,
};
let final_h = if with_head || spec_hpost() {
let final_norm = mtp.shared_head_norm.as_ref().unwrap_or(&self.output_norm);
let mut fh = e.zeros(n_embd)?;
e.rms_norm(&h_nextn, final_norm.float_data(), &mut fh, n_embd, 1, eps)?;
Some(fh)
} else {
None
};
if with_head {
let head = mtp.shared_head_head.as_ref().unwrap_or(&self.output);
let mut logits = e.matmul(head, final_h.as_ref().unwrap(), 1)?;
if let Some((mask_d, mw)) = mask_cap {
e.mask_logits_col(&mut logits, mask_d, 0, d_vocab, mw)?;
}
if let Some((ctr_d, perturb_d, q_out_d, seed, temp)) = sampled_cap {
e.copy_into(q_out_d, 0, &logits, d_vocab)?;
e.sctr_inc(ctr_d)?;
e.gumbel_perturb_ctr(&logits, perturb_d, d_vocab, seed, ctr_d, temp)?;
e.argmax_token_device_into(perturb_d, tok_d, d_vocab)?;
if with_prob {
e.prob_of_token_device_into(&logits, tok_d, p_d, d_vocab)?;
}
} else {
e.argmax_token_device_into(&logits, tok_d, d_vocab)?;
if with_prob {
e.prob_of_token_device_into(&logits, tok_d, p_d, d_vocab)?;
}
}
}
if let Some((out, slot, d2t)) = stream_pack {
e.pack_tok_p(tok_d, p_d, out, slot)?;
if let Some(map) = d2t {
e.tok_map_u32(tok_d, map)?;
}
}
if spec_hpost() {
e.copy_into(h_seed_d, 0, final_h.as_ref().unwrap(), n_embd)?;
} else {
e.copy_into(h_seed_d, 0, &h_nextn, n_embd)?;
}
e.inc_seqlen(pos_d)?;
Ok(())
}
pub fn decode_step_t(
&self,
e: &Engine,
tokens: &[u32],
pos0: usize,
cache: &mut Cache,
) -> Result<Vec<f32>, Box<dyn std::error::Error>> {
if self.is_gemma4_e4b() {
return Ok(self.gemma4_e4b_decode_step_t_h(e, tokens, pos0, cache)?.0);
}
if self.cfg.gemma4.is_some() {
return self.gemma4_decode_step_t(e, tokens, pos0, cache);
}
Ok(self.decode_step_t_h(e, tokens, pos0, cache)?.0)
}
pub fn decode_step_t_h(
&self,
e: &Engine,
tokens: &[u32],
pos0: usize,
cache: &mut Cache,
) -> Result<(Vec<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
self.decode_step_t_h_emb(e, tokens, pos0, cache, None)
}
pub fn decode_step_t_h_emb(
&self,
e: &Engine,
tokens: &[u32],
pos0: usize,
cache: &mut Cache,
embd_dev: Option<(&CudaSlice<u8>, i32, usize)>,
) -> Result<(Vec<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
let (logits_d, h_seed) = self.decode_step_t_h_emb_dev(e, tokens, pos0, cache, embd_dev)?;
Ok((e.dtoh(&logits_d)?, h_seed))
}
pub fn decode_step_t_h_emb_dev(
&self,
e: &Engine,
tokens: &[u32],
pos0: usize,
cache: &mut Cache,
embd_dev: Option<(&CudaSlice<u8>, i32, usize)>,
) -> Result<(CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
let n_embd = self.cfg.n_embd as usize;
let t = tokens.len();
let (logits, x) = self.decode_step_t_core(e, tokens, pos0, cache, embd_dev, None)?;
let mut hs = vbuf(e, n_embd)?; e.copy_view_into(&mut hs, 0, &x.slice((t - 1) * n_embd..t * n_embd), n_embd)?;
Ok((logits, hs))
}
fn decode_step_t_core(
&self,
e: &Engine,
tokens: &[u32],
pos0: usize,
cache: &mut Cache,
embd_dev: Option<(&CudaSlice<u8>, i32, usize)>,
mut ckpt: Option<&mut VerifyCkpt>,
) -> Result<(CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
self.decode_step_t_core_stream(
e,
tokens,
pos0,
cache,
embd_dev,
ckpt.take(),
None,
None,
None,
None,
)
}
fn decode_step_t_core_pipelined(
&self,
e: &Engine,
tokens: &[u32],
pos0: usize,
cache: &mut Cache,
embd_dev: Option<(&CudaSlice<u8>, i32, usize)>,
mut ckpt: Option<&mut VerifyCkpt>,
pipe: &SpecPipeLane,
round: usize,
) -> Result<(CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
let fence = crate::pp::pp_cuts(self.layers.len())
.ok_or("two-session speculative pipeline requires a PP stage cut")?;
if crate::pp::pp2_streams_off() || !crate::pp::spec_pp_on() {
return Err("two-session speculative pipeline requires the PP verify split".into());
}
let interval_fence = pipe.stage0_begin(round)?;
let ticket = self.verify_stage0_issue(
e,
tokens,
pos0,
cache,
embd_dev,
ckpt.as_deref_mut(),
None,
&fence,
Some(interval_fence),
pipe.trace(round),
)?;
pipe.stage0_end(round);
pipe.stage1_begin(round)?;
let result = self.verify_stage1_finish(e, ticket, cache, ckpt, None, &fence, true)?;
pipe.verify_end(round);
Ok(result)
}
#[allow(clippy::too_many_arguments)]
fn decode_step_t_core_stream(
&self,
e: &Engine,
tokens: &[u32],
pos0: usize,
cache: &mut Cache,
embd_dev: Option<(&CudaSlice<u8>, i32, usize)>,
mut ckpt: Option<&mut VerifyCkpt>,
stream: Option<(&CudaSlice<u32>, &CudaSlice<i32>)>,
pp_pipe: Option<bool>,
vtok_dev: Option<&CudaSlice<u32>>,
graphs: Option<&mut DsparkVerifyGraphs>,
) -> Result<(CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
if let Some(fence) = crate::pp::pp_cuts(self.layers.len()) {
if !crate::pp::pp2_streams_off() && crate::pp::spec_pp_on() {
if vtok_dev.is_some() {
return Err(
"device-token dspark verify (slice-2 deferred readback) has no PP \
stage-split arm; set MEMRA_DSPARK_DEFER_READBACK=0 or run the dspark \
route on one device"
.into(),
);
}
return self.decode_step_t_core_ppn(
e,
tokens,
pos0,
cache,
embd_dev,
ckpt.take(),
stream,
&fence,
pp_pipe,
);
}
}
crate::pp::refuse_unsplit_if_remote(
"decode_step_t (spec verify)",
"drop MEMRA_SPEC_PP=0 / MEMRA_PP_STREAMS=0 so the verify trunk takes its OWN stage \
split (decode_step_t_core_ppn); or run spec on one device",
)?;
let cfg = &self.cfg;
let n_embd = cfg.n_embd as usize;
let eps = cfg.rms_eps;
let t = tokens.len();
let pos_d = match stream {
Some((_, ctr)) => {
let mut p = e.alloc_uninit::<i32>(t)?;
e.pos_iota(ctr, &mut p, t)?;
p
}
None => {
let pos_vec: Vec<i32> = (0..t).map(|i| (pos0 + i) as i32).collect();
e.htod_i32(&pos_vec)?
}
};
let x = match (stream, embd_dev) {
(Some((vtok, _)), Some((g, qt, rb))) => {
e.embed_gather_device_td(g, vtok, t, n_embd, qt, rb)?
}
(None, Some((g, qt, rb))) => match vtok_dev {
Some(vt_d) => e.embed_gather_device_td(g, vt_d, t, n_embd, qt, rb)?,
None => e.embed_gather_device_t(g, tokens, n_embd, qt, rb)?,
},
_ => {
assert!(
vtok_dev.is_none(),
"device-token verify requires the resident embed table (embd_dev)"
);
e.htod(&self.embd.gather(n_embd, tokens))?
}
};
let x = self.verify_layers(
e,
x,
0,
self.layers.len(),
&pos_d,
pos0,
t,
cache,
ckpt.take(),
stream,
graphs,
)?;
let mut hn = vbuf(e, t * n_embd)?;
let serving_head = self.cfg.step35.is_some() || self.qwen35_serving_class();
let logits = if serving_head {
e.rms_norm(&x, self.output_norm.float_data(), &mut hn, n_embd, t, eps)?;
e.matmul(&self.output, &hn, t)?
} else {
e.rms_norm_decode(&x, self.output_norm.float_data(), &mut hn, n_embd, t, eps)?;
e.matmul_decode_exact(&self.output, &hn, t)?
};
if stream.is_none() {
cache.pos += t;
}
Ok((logits, if spec_hpost() { hn } else { x }))
}
#[allow(clippy::too_many_arguments)]
fn decode_step_t_core_ppn(
&self,
e: &Engine,
tokens: &[u32],
pos0: usize,
cache: &mut Cache,
embd_dev: Option<(&CudaSlice<u8>, i32, usize)>,
mut ckpt: Option<&mut VerifyCkpt>,
stream: Option<(&CudaSlice<u32>, &CudaSlice<i32>)>,
fence: &[usize],
pp_pipe: Option<bool>,
) -> Result<(CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
let ticket = self.verify_stage0_issue(
e,
tokens,
pos0,
cache,
embd_dev,
ckpt.as_deref_mut(),
stream,
fence,
pp_pipe,
None,
)?;
self.verify_stage1_finish(e, ticket, cache, ckpt, stream, fence, true)
}
#[allow(clippy::too_many_arguments)]
fn verify_stage0_issue(
&self,
e: &Engine,
tokens: &[u32],
pos0: usize,
cache: &mut Cache,
embd_dev: Option<(&CudaSlice<u8>, i32, usize)>,
mut ckpt: Option<&mut VerifyCkpt>,
stream: Option<(&CudaSlice<u32>, &CudaSlice<i32>)>,
fence: &[usize],
pp_pipe: Option<bool>,
trace: Option<SpecPipeTraceCtx>,
) -> Result<VerifyBoundaryTicket, Box<dyn std::error::Error>> {
assert!(
!self.is_gemma4_e4b() && self.cfg.gemma4.is_none(),
"decode_step_t_core_ppn covers the hybrid non-gemma4 verify trunk only \
(the gemma4 arms have their own decode_step_t twins)"
);
if crate::pp::pp_host_bounce_active() && (stream.is_some() || embd_dev.is_some()) {
return Err(
"decode_step_t_core_ppn: refused with MEMRA_PP_HOST_BOUNCE=1 — the trunk \
boundary itself is host-staged, but device-resident verify still peer-reads \
primary-device token/position/embedding buffers from stage 0. Run plain PP \
serving on this host class; spec requires local per-stage inputs first."
.into(),
);
}
let rt = crate::pp::PpNRt::get(e)?;
let n_st = fence.len() - 1;
assert_eq!(
rt.n_stages(),
n_st,
"PpNRt stage count {} != fence stages {n_st}",
rt.n_stages()
);
let n_embd = self.cfg.n_embd as usize;
let t = tokens.len();
let payload = t * n_embd;
if pp_pipe.is_some() {
assert_eq!(n_st, 2, "spec pipeline requires exactly two PP stages");
}
let pp_anatomy = n_st == 2 && std::env::var("MEMRA_SPEC_PP_ANATOMY").as_deref() == Ok("1");
let pp_started = std::time::Instant::now();
let (mut reverse_ms, mut stage0_ms, mut tx_ms) = (0.0f64, 0.0f64, 0.0f64);
let caller_stream = e.stream();
let reverse_started = std::time::Instant::now();
if pp_pipe != Some(false) {
rt.fence_stages_behind(&caller_stream)?;
}
if pp_pipe == Some(true) {
rt.prepare_overlap_slots(0, payload)?;
}
if pp_anatomy {
for s in 0..n_st {
let _st = rt.enter(s);
rt.engine(s, e).stream().synchronize()?;
}
reverse_ms = reverse_started.elapsed().as_secs_f64() * 1e3;
}
let stage_pos = |es: &Engine| -> Result<CudaSlice<i32>, Box<dyn std::error::Error>> {
match stream {
Some((_, ctr)) => {
let mut p = es.alloc_uninit::<i32>(t)?;
es.pos_iota(ctr, &mut p, t)?;
Ok(p)
}
None => {
let pos_vec: Vec<i32> = (0..t).map(|i| (pos0 + i) as i32).collect();
es.htod_i32(&pos_vec)
}
}
};
let slot = {
let _st0 = rt.enter(0);
let e0 = rt.engine(0, e);
enqueue_spec_pipe_trace_marker(&e0.stream(), trace.as_ref(), "S0", "start", None)?;
let stage0_started = std::time::Instant::now();
let pos_d = stage_pos(e0)?;
let x = match (stream, embd_dev) {
(Some((vtok, _)), Some((g, qt, rb))) => {
e0.embed_gather_device_td(g, vtok, t, n_embd, qt, rb)?
}
(None, Some((g, qt, rb))) => e0.embed_gather_device_t(g, tokens, n_embd, qt, rb)?,
_ => e0.htod(&self.embd.gather(n_embd, tokens))?,
};
let x = self.verify_layers(
e0,
x,
fence[0],
fence[1],
&pos_d,
pos0,
t,
cache,
ckpt.as_deref_mut(),
stream,
None,
)?;
if pp_anatomy {
e0.stream().synchronize()?;
stage0_ms = stage0_started.elapsed().as_secs_f64() * 1e3;
}
let tx_started = std::time::Instant::now();
let slot = if pp_pipe.is_some() {
rt.tx_pipelined(0, &x, payload)?
} else {
rt.tx(0, &x, payload)?
};
enqueue_spec_pipe_trace_marker(&e0.stream(), trace.as_ref(), "S0", "end", Some(slot))?;
if pp_anatomy {
e0.stream().synchronize()?;
tx_ms = tx_started.elapsed().as_secs_f64() * 1e3;
}
slot
};
Ok(VerifyBoundaryTicket {
rt,
caller_stream,
slot,
pos0,
t,
payload,
n_st,
pipelined: pp_pipe.is_some(),
pp_anatomy,
pp_started,
reverse_ms,
stage0_ms,
tx_ms,
trace,
})
}
#[allow(clippy::too_many_arguments)]
fn verify_stage1_finish(
&self,
e: &Engine,
ticket: VerifyBoundaryTicket,
cache: &mut Cache,
mut ckpt: Option<&mut VerifyCkpt>,
stream: Option<(&CudaSlice<u32>, &CudaSlice<i32>)>,
fence: &[usize],
publish_to_caller: bool,
) -> Result<(CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
let VerifyBoundaryTicket {
rt,
caller_stream,
slot,
pos0,
t,
payload,
n_st,
pipelined,
pp_anatomy,
pp_started,
reverse_ms,
stage0_ms,
tx_ms,
trace,
} = ticket;
let n_embd = self.cfg.n_embd as usize;
let eps = self.cfg.rms_eps;
let mut slot = slot;
let (mut rx_ms, mut stage1_ms) = (0.0f64, 0.0f64);
let stage_pos = |es: &Engine| -> Result<CudaSlice<i32>, Box<dyn std::error::Error>> {
match stream {
Some((_, ctr)) => {
let mut p = es.alloc_uninit::<i32>(t)?;
es.pos_iota(ctr, &mut p, t)?;
Ok(p)
}
None => {
let pos_vec: Vec<i32> = (0..t).map(|i| (pos0 + i) as i32).collect();
es.htod_i32(&pos_vec)
}
}
};
for s in 1..n_st - 1 {
let _st = rt.enter(s);
let es = rt.engine(s, e);
let pos_d = stage_pos(es)?;
let x = rt.rx(s - 1, slot, payload)?;
let x = self.verify_layers(
es,
x,
fence[s],
fence[s + 1],
&pos_d,
pos0,
t,
cache,
ckpt.as_deref_mut(),
stream,
None,
)?;
slot = if pipelined {
rt.tx_pipelined(s, &x, payload)?
} else {
rt.tx(s, &x, payload)?
};
}
let _stl = rt.enter(n_st - 1);
let el = rt.engine(n_st - 1, e);
let pos_d = stage_pos(el)?;
let rx_started = std::time::Instant::now();
let x = rt.rx(n_st - 2, slot, payload)?;
if pp_anatomy {
el.stream().synchronize()?;
rx_ms = rx_started.elapsed().as_secs_f64() * 1e3;
}
enqueue_spec_pipe_trace_marker(&el.stream(), trace.as_ref(), "S1", "start", Some(slot))?;
let stage1_started = std::time::Instant::now();
let x = self.verify_layers(
el,
x,
fence[n_st - 1],
fence[n_st],
&pos_d,
pos0,
t,
cache,
ckpt.as_deref_mut(),
stream,
None,
)?;
let mut hn = vbuf(el, payload)?;
let logits = if self.cfg.step35.is_some() {
el.rms_norm(&x, self.output_norm.float_data(), &mut hn, n_embd, t, eps)?;
el.matmul(&self.output, &hn, t)?
} else {
el.rms_norm_decode(&x, self.output_norm.float_data(), &mut hn, n_embd, t, eps)?;
el.matmul_decode_exact(&self.output, &hn, t)?
};
enqueue_spec_pipe_trace_marker(&el.stream(), trace.as_ref(), "S1", "end", Some(slot))?;
if pp_anatomy {
el.stream().synchronize()?;
stage1_ms = stage1_started.elapsed().as_secs_f64() * 1e3;
}
if publish_to_caller {
rt.publish_to(n_st - 1, &caller_stream)?;
}
if pp_anatomy {
if publish_to_caller {
caller_stream.synchronize()?;
}
eprintln!(
"[spec-pp-anatomy] t={t} reverse={reverse_ms:.3}ms stage0={stage0_ms:.3}ms \
tx={tx_ms:.3}ms rx={rx_ms:.3}ms stage1-head={stage1_ms:.3}ms total={:.3}ms",
pp_started.elapsed().as_secs_f64() * 1e3,
);
}
if stream.is_none() {
cache.pos += t;
}
Ok((logits, if spec_hpost() { hn } else { x }))
}
#[allow(clippy::too_many_arguments)]
fn step35_verify_batch_layers(
&self,
e: &Engine,
mut x: CudaSlice<f32>,
lo: usize,
hi: usize,
pos0: usize,
t: usize,
cache: &mut Cache,
) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
let n_embd = self.cfg.n_embd as usize;
self.cfg
.step35
.as_ref()
.ok_or("step35 verify batch requires step35 cfg")?;
let mut ph_last = std::time::Instant::now();
for il in lo..hi {
let mut next = e.uninit(t * n_embd)?;
for r in 0..t {
let mut row = e.uninit(n_embd)?;
e.dtod_copy_view(&x.slice(r * n_embd..(r + 1) * n_embd), &mut row)?;
let row_pos = e.htod_i32(&[(pos0 + r) as i32])?;
let mut one = [&mut *cache];
let out = self.step35_decode_batch_layers(
e,
row,
&mut one,
&row_pos,
il,
il + 1,
&mut ph_last,
)?;
e.dtod_copy_into(&out, &mut next, r * n_embd)?;
}
self.dflash_tap(e, cache, il, &next, t)?;
x = next;
}
Ok(x)
}
pub(crate) fn dspark_verify_t_am(
&self,
e: &Engine,
tokens: &[u32],
pos0: usize,
cache: &mut Cache,
) -> Result<Vec<u32>, Box<dyn std::error::Error>> {
let (logits, _hn) = self.decode_step_t_core_stream(
e, tokens, pos0, cache, None, None, None, None, None, None,
)?;
let t = tokens.len();
let v = self.output.out_features();
let mut am_d = e.stream().alloc_zeros::<u32>(t)?;
for r in 0..t {
e.argmax_token_device_col(&logits, r, v, &mut am_d, r)?;
}
Ok(e.dtoh_u32(&am_d)?)
}
pub(crate) fn dspark_verify_t_logits(
&self,
e: &Engine,
tokens: &[u32],
pos0: usize,
cache: &mut Cache,
) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
let (logits, _hn) = self.decode_step_t_core_stream(
e, tokens, pos0, cache, None, None, None, None, None, None,
)?;
Ok(logits)
}
pub(crate) fn dspark_verify_t_am_ckpt(
&self,
e: &Engine,
tokens: &[u32],
pos0: usize,
cache: &mut Cache,
) -> Result<(Vec<u32>, DsparkVerifyCkpt), Box<dyn std::error::Error>> {
let mut ck = VerifyCkpt::new(self.layers.len());
let (logits, _hn) = self.decode_step_t_core_stream(
e,
tokens,
pos0,
cache,
None,
Some(&mut ck),
None,
None,
None,
None,
)?;
let t = tokens.len();
let v = self.output.out_features();
let mut am_d = e.stream().alloc_zeros::<u32>(t)?;
for r in 0..t {
e.argmax_token_device_col(&logits, r, v, &mut am_d, r)?;
}
Ok((e.dtoh_u32(&am_d)?, DsparkVerifyCkpt(ck)))
}
pub(crate) fn dspark_verify_t_am_ckpt_dev(
&self,
e: &Engine,
vtok: &CudaSlice<u32>,
t: usize,
pos0: usize,
cache: &mut Cache,
embd_dev: (&CudaSlice<u8>, i32, usize),
graphs: Option<&mut DsparkVerifyGraphs>,
) -> Result<(CudaSlice<u32>, DsparkVerifyCkpt), Box<dyn std::error::Error>> {
debug_assert!(
vtok.len() >= t,
"verify window exceeds the device token buffer"
);
let mut graphs = graphs;
if let Some(g) = graphs.as_deref_mut() {
g.round_slab = false;
}
let mut ck = VerifyCkpt::new(self.layers.len());
let dummy = vec![0u32; t];
let (logits, _hn) = self.decode_step_t_core_stream(
e,
&dummy,
pos0,
cache,
Some(embd_dev),
Some(&mut ck),
None,
None,
Some(vtok),
graphs,
)?;
let v = self.output.out_features();
let mut am_d = e.stream().alloc_zeros::<u32>(t)?;
for r in 0..t {
e.argmax_token_device_col(&logits, r, v, &mut am_d, r)?;
}
Ok((am_d, DsparkVerifyCkpt(ck)))
}
pub(crate) fn dspark_verify_t_logits_ckpt(
&self,
e: &Engine,
tokens: &[u32],
pos0: usize,
cache: &mut Cache,
) -> Result<(CudaSlice<f32>, DsparkVerifyCkpt), Box<dyn std::error::Error>> {
let mut ck = VerifyCkpt::new(self.layers.len());
let (logits, _hn) = self.decode_step_t_core_stream(
e,
tokens,
pos0,
cache,
None,
Some(&mut ck),
None,
None,
None,
None,
)?;
Ok((logits, DsparkVerifyCkpt(ck)))
}
pub(crate) fn dspark_commit_prefix(
&self,
e: &Engine,
cache: &mut Cache,
snap: &crate::cache::CacheSnapshot,
ckpt: &DsparkVerifyCkpt,
keep: usize,
) -> Result<(), Box<dyn std::error::Error>> {
self.commit_verified_prefix(e, cache, snap, &ckpt.0, keep, false, None)
}
pub(crate) fn dspark_commit_prefix_slab(
&self,
e: &Engine,
cache: &mut Cache,
snap: &crate::cache::CacheSnapshot,
ctx: &DsparkVerifyGraphs,
keep: usize,
) -> Result<(), Box<dyn std::error::Error>> {
use cudarc::driver::DevicePtr;
debug_assert!(keep >= 1, "keep==0 rounds take the legacy rollback");
let mut conv_src: Vec<u64> = Vec::new();
let mut ssm_src: Vec<u64> = Vec::new();
let mut conv_dst: Vec<u64> = Vec::new();
let mut ssm_dst: Vec<u64> = Vec::new();
for il in 0..self.layers.len() {
if let (Some(kvl), Some(saved)) = (cache.kv[il].as_mut(), snap.kv_len[il]) {
kvl.len = saved + keep;
e.set_i32_one(&mut kvl.len_d, kvl.len as i32)?;
}
if let Some(rl) = cache.recur[il].as_ref() {
let (pc, ps, _cw, _sw) = ctx
.slab_row(e, il, keep - 1)
.ok_or("slab commit: linear layer missing from the graphs ctx")?;
conv_src.push(pc);
ssm_src.push(ps);
let st = &e.gpu.stream();
let (dc, _g0) = rl.conv_state.device_ptr(st);
let (ds, _g1) = rl.ssm_state.device_ptr(st);
conv_dst.push(dc as u64);
ssm_dst.push(ds as u64);
}
}
let n = conv_src.len();
if n > 0 {
if state_copy_batch_on() {
let mut tt = vec![0u64; 2 * n];
tt[..n].copy_from_slice(&conv_src);
tt[n..].copy_from_slice(&conv_dst);
let ct = e.htod_u64(&tt)?;
tt[..n].copy_from_slice(&ssm_src);
tt[n..].copy_from_slice(&ssm_dst);
let st = e.htod_u64(&tt)?;
e.copy_batch_uniform_f32(&ct, n, ctx.conv_words)?;
e.copy_batch_uniform_f32(&st, n, ctx.ssm_words)?;
} else {
let (cw, sw) = (ctx.conv_words, ctx.ssm_words);
let row = keep - 1;
for il in 0..self.layers.len() {
let Some(rl) = cache.recur[il].as_mut() else {
continue;
};
let k = ctx.lin_pos[&il];
{
let sv = e.view(&ctx.stash_conv[k], (row + 1) * cw);
let win = sv.slice(row * cw..(row + 1) * cw);
e.copy_view_into(&mut rl.conv_state, 0, &win, cw)?;
}
{
let sv = e.view(&ctx.stash_ssm[k], (row + 1) * sw);
let win = sv.slice(row * sw..(row + 1) * sw);
e.copy_view_into(&mut rl.ssm_state, 0, &win, sw)?;
}
}
}
}
cache.pos = snap.pos + keep;
Ok(())
}
#[allow(clippy::too_many_arguments)]
fn qwen35_verify_batch_layers(
&self,
e: &Engine,
x: CudaSlice<f32>,
lo: usize,
hi: usize,
pos0: usize,
t: usize,
cache: &mut Cache,
ckpt: Option<&mut VerifyCkpt>,
stream: Option<(&CudaSlice<u32>, &CudaSlice<i32>)>,
graphs: Option<&mut DsparkVerifyGraphs>,
) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
let rowwise = std::env::var("MEMRA_SPEC_VERIFY_ROWWISE").as_deref() == Ok("1")
|| !matches!(
self.cfg.arch,
memra_gguf::config::Arch::Qwen35 | memra_gguf::config::Arch::Qwen35Moe
)
|| t > 16;
if rowwise {
if stream.is_some() {
return Err("qwen35 rowwise verify has no ROUND-STREAM arm \
(t > 16 or MEMRA_SPEC_VERIFY_ROWWISE=1)"
.into());
}
self.qwen35_verify_rowwise(e, x, lo, hi, pos0, t, cache, ckpt)
} else {
self.qwen35_verify_tparallel(e, x, lo, hi, pos0, t, cache, ckpt, stream, graphs)
}
}
#[allow(clippy::too_many_arguments)]
fn qwen35_verify_rowwise(
&self,
e: &Engine,
mut x: CudaSlice<f32>,
lo: usize,
hi: usize,
pos0: usize,
t: usize,
cache: &mut Cache,
mut ckpt: Option<&mut VerifyCkpt>,
) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
let n_embd = self.cfg.n_embd as usize;
let saved_pos = cache.pos;
let mut ph_last = std::time::Instant::now();
for il in lo..hi {
let mut next = e.uninit(t * n_embd)?;
let mut col_states: Option<Vec<(CudaSlice<f32>, CudaSlice<f32>)>> =
if ckpt.is_some() && t >= 2 && matches!(self.layers[il].mixer, Mixer::Linear(_)) {
Some(Vec::with_capacity(t - 1))
} else {
None
};
for r in 0..t {
cache.pos = pos0 + r;
let mut row = e.uninit(n_embd)?;
e.dtod_copy_view(&x.slice(r * n_embd..(r + 1) * n_embd), &mut row)?;
let row_pos = e.htod_i32(&[(pos0 + r) as i32])?;
let mut one = [&mut *cache];
let ctx = self.batch_layer_ctx(e, &one, il, il + 1)?;
let out = match self.decode_batch_layers(
e,
row,
&mut one,
&ctx,
&row_pos,
&mut ph_last,
) {
Ok(out) => out,
Err(error) => {
cache.pos = saved_pos;
return Err(error);
}
};
e.dtod_copy_into(&out, &mut next, r * n_embd)?;
if r + 1 < t {
if let Some(states) = col_states.as_mut() {
let recur = cache.recur[il]
.as_ref()
.ok_or("Qwen35-MoE linear verify layer has no recurrent state")?;
states.push((
e.clone_dtod(&recur.conv_state)?,
e.clone_dtod(&recur.ssm_state)?,
));
}
}
}
if let (Some(checkpoint), Some(states)) = (ckpt.as_deref_mut(), col_states) {
checkpoint.cols[il] = Some(states);
}
x = next;
}
cache.pos = saved_pos;
Ok(x)
}
#[allow(clippy::too_many_arguments)]
fn qwen35_verify_tparallel(
&self,
e: &Engine,
mut x: CudaSlice<f32>,
lo: usize,
hi: usize,
pos0: usize,
t: usize,
cache: &mut Cache,
mut ckpt: Option<&mut VerifyCkpt>,
stream: Option<(&CudaSlice<u32>, &CudaSlice<i32>)>,
mut graphs: Option<&mut DsparkVerifyGraphs>,
) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
let seqs_append =
std::env::var("MEMRA_BATCH_APPEND").as_deref() != Ok("0") && !Engine::kv_fp8_on();
let batch_fa_on = std::env::var("MEMRA_BATCH_FA").as_deref() != Ok("0");
if stream.is_some() && graphs.is_some() {
return Err(
"qwen35 tparallel verify: ROUND-STREAM and dspark verify graphs \
cannot arm together"
.into(),
);
}
if let Some(g) = graphs.as_deref_mut() {
g.refresh_tables(e, cache)?;
g.round_slab = false;
if let Some(rung) = g.full_rung(self, cache, lo, hi, t, seqs_append && batch_fa_on) {
let out = g.run_full(self, e, lo, hi, &x, t, pos0, rung, cache)?;
g.round_slab = true;
return Ok(out);
}
}
let pos_d = match stream {
Some((_, ctr)) => {
let mut p = e.alloc_uninit::<i32>(t)?;
e.pos_iota(ctr, &mut p, t)?;
p
}
None => {
let pos_host: Vec<i32> = (0..t).map(|r| (pos0 + r) as i32).collect();
e.htod_i32(&pos_host)?
}
};
let mut pos_rows: Option<Vec<CudaSlice<i32>>> = None;
let mut il = lo;
while il < hi {
if graphs.is_some() && matches!(self.layers[il].mixer, Mixer::Linear(_)) {
let mut end = il;
while end < hi && matches!(self.layers[end].mixer, Mixer::Linear(_)) {
end += 1;
}
let g = graphs.as_deref_mut().expect("checked above");
x = g.run_segment(self, e, il, end, &x, t, cache)?;
g.round_slab = true;
il = end;
continue;
}
let layer = &self.layers[il];
if stream.is_none() && matches!(layer.mixer, Mixer::Linear(_)) {
x = self.qwen35_tparallel_linear_layer(
e,
il,
&x,
t,
cache,
ckpt.as_deref_mut(),
None,
None,
)?;
il += 1;
continue;
}
x = self.qwen35_tparallel_fa_layer(
e,
il,
&x,
t,
cache,
FaLayerArgs {
pos_d: &pos_d,
pos_rows: &mut pos_rows,
pos0,
seqs_append,
batch_fa_on,
graph_cap: None,
stream,
ckpt: ckpt.as_deref_mut(),
},
)?;
il += 1;
}
Ok(x)
}
#[allow(clippy::too_many_arguments)]
fn qwen35_tparallel_dense_ffn(
&self,
e: &Engine,
ffn_gate: &crate::model::GpuTensor,
ffn_up: &crate::model::GpuTensor,
ffn_down: &crate::model::GpuTensor,
zn: &CudaSlice<f32>,
t: usize,
n_embd: usize,
) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
let n_ff = ffn_gate.out_features();
let (zq, zd) = e.quantize_q8_1(zn, t, n_embd)?;
if Engine::tk_ffn_dual_on() {
if let Some(((g, gs), (u, us))) =
e.matmul_decode_exact_dual_pre(ffn_gate, ffn_up, &zq, &zd, t)?
{
if e.uses_q8_1_fast(ffn_down) {
let (aq, ad) = e.silu_mul_scaled_q8_1(&g, &u, gs, us, t * n_ff)?;
return e.matmul_decode_exact_pre(ffn_down, &aq, &ad, t);
}
let mut act = e.uninit(t * n_ff)?;
e.silu_mul_scaled(&g, &u, gs, us, &mut act, t * n_ff)?;
let (aq, ad) = e.quantize_q8_1(&act, t, n_ff)?;
return e.matmul_pre(ffn_down, &aq, &ad, &act, t);
}
}
let g = e.matmul_pre(ffn_gate, &zq, &zd, zn, t)?;
let u = e.matmul_pre(ffn_up, &zq, &zd, zn, t)?;
let mut act = e.uninit(t * n_ff)?;
e.silu_mul(&g, &u, &mut act, t * n_ff)?;
let (aq, ad) = e.quantize_q8_1(&act, t, n_ff)?;
e.matmul_pre(ffn_down, &aq, &ad, &act, t)
}
#[allow(clippy::too_many_arguments)]
fn qwen35_tparallel_fa_layer(
&self,
e: &Engine,
il: usize,
x: &CudaSlice<f32>,
t: usize,
cache: &mut Cache,
args: FaLayerArgs<'_>,
) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
use cudarc::driver::DevicePtr;
let cfg = &self.cfg;
let n_embd = cfg.n_embd as usize;
let eps = cfg.rms_eps;
let head_dim_global = cfg.head_dim_k as usize;
let layer = &self.layers[il];
let FaLayerArgs {
pos_d,
pos_rows,
pos0,
seqs_append,
batch_fa_on,
graph_cap,
stream,
mut ckpt,
} = args;
let anorm = layer.attn_norm.float_data();
let mut xn = e.uninit(t * n_embd)?;
e.rms_norm(x, anorm, &mut xn, n_embd, t, eps)?;
let (hq, hd) = e.quantize_q8_1(&xn, t, n_embd)?;
let mixed: CudaSlice<f32> = match &layer.mixer {
Mixer::Mla(_) => crate::hybrid::mla_forward_unimplemented(),
Mixer::Linear(la) if stream.is_some() => {
if !(t >= 3 || (t == 2 && spec_m2()))
|| !self.mixer_in_q8_1_fast(e, &layer.mixer)
|| !e.uses_q8_1_fast(&la.ssm_out)
{
return Err("qwen35 stream verify: GDN batched arm requires t>=3 \
(or MEMRA_SPEC_M2 at t=2) and q8_1-fast projections"
.into());
}
let want = ckpt.is_some();
let (out, stash) =
self.linear_attn_verify_t(e, la, &xn, Some((&hq, &hd)), t, cache, il, want)?;
if let (Some(ck), Some(st)) = (ckpt.as_deref_mut(), stash) {
ck.gdn[il] = Some(st);
}
out
}
Mixer::Linear(_) => {
unreachable!("linear layers ride qwen35_tparallel_linear_layer")
}
Mixer::Full(fa) => {
let geometry = cfg.full_attention_geometry_at(il as u32);
let n_head = geometry.n_head as usize;
let n_head_kv = geometry.n_head_kv as usize;
let head_dim = geometry.head_dim_k as usize;
let rope_dims = geometry.n_rot as usize;
let rope_base = geometry.rope_base;
let scale = geometry.attention_scale();
let (qf, mut k, v) = match e.matmul_decode_exact_group3_pre(
[&fa.wq, &fa.wk, &fa.wv],
&hq,
&hd,
t,
)? {
Some(mut g3) => {
let v = g3.pop().unwrap();
let k = g3.pop().unwrap();
let qf = g3.pop().unwrap();
(qf, k, v)
}
None => (
e.matmul_pre(&fa.wq, &hq, &hd, &xn, t)?,
e.matmul_pre(&fa.wk, &hq, &hd, &xn, t)?,
e.matmul_pre(&fa.wv, &hq, &hd, &xn, t)?,
),
};
let gated =
geometry.attention_gate == memra_gguf::config::AttentionGateKind::FusedQ;
let (mut q, gate) = if gated {
let mut qs = e.uninit(t * n_head * head_dim)?;
let mut gs = e.uninit(t * n_head * head_dim)?;
e.q_gate_split(&qf, &mut qs, &mut gs, head_dim, n_head, t)?;
(qs, Some(gs))
} else {
(qf, None)
};
let mut qn = e.uninit(t * n_head * head_dim)?;
e.rms_norm(
&q,
fa.q_norm.float_data(),
&mut qn,
head_dim,
t * n_head,
eps,
)?;
q = qn;
let mut kn = e.uninit(t * n_head_kv * head_dim)?;
e.rms_norm(
&k,
fa.k_norm.float_data(),
&mut kn,
head_dim,
t * n_head_kv,
eps,
)?;
k = kn;
e.rope_neox(
&mut q, pos_d, head_dim, rope_dims, n_head, t, rope_base, 1.0,
)?;
e.rope_neox(
&mut k, pos_d, head_dim, rope_dims, n_head_kv, t, rope_base, 1.0,
)?;
let q_dim = n_head * head_dim;
let kv_dim = n_head_kv * head_dim;
let mut attn = e.uninit(t * q_dim)?;
let (kdk, kdv, ktb, vtb, len0, kv_local) = {
let kvl = cache.kv[il].as_ref().unwrap();
let local: Option<CudaSlice<u64>> = match graph_cap {
Some(_) => None,
None => {
let s = &e.gpu.stream();
let (pk, _g) = kvl.k.device_ptr(s);
let (pv, _g2) = kvl.v.device_ptr(s);
let mut tbl = Vec::with_capacity(2 * t);
for _ in 0..t {
tbl.push(pk as u64);
tbl.push(pv as u64);
}
Some(e.htod_u64(&tbl)?)
}
};
(
kvl.kv_dim_k,
kvl.kv_dim_v,
kvl.k_tok_bytes,
kvl.v_tok_bytes,
kvl.len,
local,
)
};
let (kv_tbl, kv_off): (&CudaSlice<u64>, usize) = match graph_cap {
Some((tb, off, _)) => (tb, off),
None => (kv_local.as_ref().expect("built above"), 0),
};
let t_kv_first = len0 + 1;
let t_kv_last = len0 + t;
let rows_batched = t >= 2
&& seqs_append
&& batch_fa_on
&& dspark_fa_rows_on()
&& kdk == kv_dim
&& kdv == kv_dim
&& crate::fa_seqs_eligible(t_kv_first, head_dim_global)
&& crate::fa_seqs_eligible(t_kv_last, head_dim_global)
&& crate::fa_split_keys(t_kv_first, cfg.n_head_kv as usize)
== crate::fa_split_keys(t_kv_last, cfg.n_head_kv as usize);
let (size_kv_max, sp) = match graph_cap {
Some((_, _, rung)) => {
if !rows_batched {
return Err(format!(
"fa graph capture: layer {il} round is not batchable \
(t_kv {t_kv_first}..{t_kv_last}) — the per-row fallback \
must never be captured"
)
.into());
}
let sp_r = crate::fa_split_keys(rung, cfg.n_head_kv as usize);
if t_kv_last > rung
|| sp_r != crate::fa_split_keys(t_kv_last, cfg.n_head_kv as usize)
{
return Err(format!(
"fa graph capture: rung {rung} does not cover round \
t_kv {t_kv_first}..{t_kv_last} on one split ladder step"
)
.into());
}
(rung, sp_r)
}
None => (
t_kv_last,
crate::fa_split_keys(t_kv_last, cfg.n_head_kv as usize),
),
};
if let Some((_, ctr)) = stream {
let kvl = cache.kv[il].as_mut().unwrap();
e.append_kv_quantized_rows_dc(
&k,
&v,
&mut kvl.k,
&mut kvl.v,
ctr,
t,
kdk,
kdv,
ktb,
vtb,
Engine::kv_fp8_on(),
)?;
let upper = (kvl.len + t + 64).min(cache.max_ctx);
let k_view = e.view_u8(&kvl.k, upper * ktb);
let v_view = e.view_u8(&kvl.v, upper * vtb);
e.fa_decode_rows_dc(
&q, &k_view, &v_view, &mut attn, head_dim, n_head, n_head_kv, ctr, upper,
t, scale, ktb, vtb, 0, false,
)?;
} else if rows_batched {
e.append_kv_quantized_seqs(
&k,
&v,
&kv_tbl.slice(kv_off..kv_off + 2 * t),
pos_d,
t,
kdk,
kdv,
ktb,
vtb,
)?;
if graph_cap.is_none() {
cache.kv[il].as_mut().unwrap().len += t;
}
e.fa_decode_batch_seqs_v4(
&q,
&kv_tbl.slice(kv_off..kv_off + 2 * t),
pos_d,
&mut attn,
head_dim,
n_head,
n_head_kv,
t,
size_kv_max,
scale,
sp,
ktb,
vtb,
)?;
} else {
if pos_rows.is_none() {
*pos_rows = Some(match stream {
Some((_, ctr)) => (0..t)
.map(|r| {
let mut b = e.alloc_uninit::<i32>(1)?;
e.i32_copy_add(ctr, &mut b, r as i32)?;
Ok(b)
})
.collect::<Result<_, Box<dyn std::error::Error>>>()?,
None => (0..t)
.map(|r| e.htod_i32(&[(pos0 + r) as i32]))
.collect::<Result<_, _>>()?,
});
}
let pos_rows = pos_rows.as_ref().unwrap();
for r in 0..t {
let mut k_row = e.uninit(kv_dim)?;
e.dtod_copy_view(&k.slice(r * kv_dim..(r + 1) * kv_dim), &mut k_row)?;
let mut v_row = e.uninit(kv_dim)?;
e.dtod_copy_view(&v.slice(r * kv_dim..(r + 1) * kv_dim), &mut v_row)?;
let pos_row = &pos_rows[r];
let kvl = cache.kv[il].as_mut().unwrap();
if seqs_append {
e.append_kv_quantized_seqs(
&k_row,
&v_row,
&kv_tbl.slice(kv_off..kv_off + 2),
pos_row,
1,
kdk,
kdv,
ktb,
vtb,
)?;
kvl.len += 1;
} else {
e.append_kv_quantized_view(
&k_row.slice(0..kv_dim),
&v_row.slice(0..kv_dim),
&mut kvl.k,
&mut kvl.v,
kvl.len,
kvl.kv_dim_k,
kvl.kv_dim_v,
kvl.k_tok_bytes,
kvl.v_tok_bytes,
Engine::kv_fp8_on(),
)?;
kvl.len += 1;
}
let t_kv = kvl.len;
let mut q_row = e.uninit(q_dim)?;
e.dtod_copy_view(&q.slice(r * q_dim..(r + 1) * q_dim), &mut q_row)?;
let mut a_row = e.uninit(q_dim)?;
if batch_fa_on && crate::fa_seqs_eligible(t_kv, head_dim_global) {
let sp0_r = crate::fa_split_keys(t_kv, cfg.n_head_kv as usize);
e.fa_decode_batch_seqs_v4(
&q_row,
&kv_tbl.slice(kv_off..kv_off + 2),
pos_row,
&mut a_row,
head_dim,
n_head,
n_head_kv,
1,
t_kv,
scale,
sp0_r,
ktb,
vtb,
)?;
} else {
let k_view = e.view_u8(&kvl.k, t_kv * kvl.k_tok_bytes);
let v_view = e.view_u8(&kvl.v, t_kv * kvl.v_tok_bytes);
let mut a_view = a_row.slice_mut(0..q_dim);
e.fa_decode_kvmod_view(
&q_row.slice(0..q_dim),
&k_view,
&v_view,
&mut a_view,
head_dim,
n_head,
n_head_kv,
t_kv,
scale,
kvl.k_tok_bytes,
kvl.v_tok_bytes,
Engine::kv_fp8_on(),
)?;
}
e.dtod_copy_into(&a_row, &mut attn, r * q_dim)?;
}
}
let attn_g = match &gate {
Some(g) => {
let n = t * q_dim;
let mut gsig = e.uninit(n)?;
e.sigmoid(g, &mut gsig, n)?;
let mut ag = e.uninit(n)?;
e.mul(&attn, &gsig, &mut ag, n)?;
ag
}
None => attn,
};
e.matmul(&fa.wo, &attn_g, t)?
}
};
let pnorm = layer.post_attn_norm.float_data();
let mut x1 = e.uninit(t * n_embd)?;
let mut zn = e.uninit(t * n_embd)?;
e.add_rms_norm(x, &mixed, pnorm, &mut x1, &mut zn, n_embd, t, eps)?;
let ffn_out = match &layer.ffn {
crate::hybrid::Ffn::Dense {
ffn_gate,
ffn_up,
ffn_down,
} => {
assert!(
self.cfg.m3.is_none(),
"qwen35 t-parallel verify: M3 swigluoai FFN not yet batched"
);
self.qwen35_tparallel_dense_ffn(e, ffn_gate, ffn_up, ffn_down, &zn, t, n_embd)?
}
crate::hybrid::Ffn::Moe(m) => self.moe_ffn_il_zq8(e, m, &zn, None, t, il as u16)?,
};
let mut x2 = e.uninit(t * n_embd)?;
e.add(&x1, &ffn_out, &mut x2, t * n_embd)?;
self.dflash_tap(e, cache, il, &x2, t)?;
Ok(x2)
}
#[allow(clippy::too_many_arguments)]
fn qwen35_tparallel_linear_layer(
&self,
e: &Engine,
il: usize,
x: &CudaSlice<f32>,
t: usize,
cache: &mut Cache,
mut ckpt: Option<&mut VerifyCkpt>,
stash: Option<(&mut CudaSlice<f32>, &mut CudaSlice<f32>)>,
table_src: Option<(&CudaSlice<u64>, usize)>,
) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
use cudarc::driver::DevicePtr;
let cfg = &self.cfg;
let n_embd = cfg.n_embd as usize;
let eps = cfg.rms_eps;
let layer = &self.layers[il];
let Mixer::Linear(la) = &layer.mixer else {
return Err("qwen35_tparallel_linear_layer on a non-linear layer".into());
};
let anorm = layer.attn_norm.float_data();
let mut xn = e.uninit(t * n_embd)?;
e.rms_norm(x, anorm, &mut xn, n_embd, t, eps)?;
let (hq, hd) = e.quantize_q8_1(&xn, t, n_embd)?;
let ssm = cfg.ssm.as_ref().expect("linear mixer requires ssm cfg");
let d_state = ssm.state_size as usize;
let num_k = ssm.group_count as usize;
let num_v = ssm.time_step_rank as usize;
let d_conv = ssm.conv_kernel as usize;
let key_dim = d_state * num_k;
let value_dim = d_state * num_v;
let conv_dim = key_dim * 2 + value_dim;
let gdn_scale = 1.0 / (d_state as f32).sqrt();
let (qkv_mixed, z, beta_raw, alpha) = match e.matmul_decode_exact_group4_pre(
[&la.wqkv, &la.wqkv_gate, &la.ssm_beta, &la.ssm_alpha],
&hq,
&hd,
t,
)? {
Some(mut g4) => {
let alpha = g4.pop().unwrap();
let beta_raw = g4.pop().unwrap();
let z = g4.pop().unwrap();
let qkv_mixed = g4.pop().unwrap();
(qkv_mixed, z, beta_raw, alpha)
}
None => (
e.matmul_pre(&la.wqkv, &hq, &hd, &xn, t)?,
e.matmul_pre(&la.wqkv_gate, &hq, &hd, &xn, t)?,
e.matmul_pre(&la.ssm_beta, &hq, &hd, &xn, t)?,
e.matmul_pre(&la.ssm_alpha, &hq, &hd, &xn, t)?,
),
};
let beta_w = la.ssm_beta.out_features();
let alpha_w = la.ssm_alpha.out_features();
let qkv_w = la.wqkv.out_features();
let table_local: Option<CudaSlice<u64>> = match table_src {
Some(_) => None,
None => {
let rl = cache.recur[il].as_ref().unwrap();
let s = &e.gpu.stream();
let (pc, _g0) = rl.conv_state.device_ptr(s);
let (p0, _g1) = rl.ssm_state.device_ptr(s);
let (p1, _g2) = rl.ssm_state_alt.device_ptr(s);
Some(e.htod_u64(&[
pc as u64, p0 as u64, p1 as u64, pc as u64, p1 as u64, p0 as u64,
])?)
}
};
let (table, toff): (&CudaSlice<u64>, usize) = match table_src {
Some((tb, off)) => (tb, off),
None => (table_local.as_ref().unwrap(), 0),
};
let mut o_all = e.uninit(t * value_dim)?;
let mut col_states: Option<Vec<(CudaSlice<f32>, CudaSlice<f32>)>> =
if ckpt.is_some() && stash.is_none() && t >= 2 {
Some(Vec::with_capacity(t - 1))
} else {
None
};
let mut stash = stash;
let mut conv_out = e.uninit(conv_dim)?;
let mut q_l2 = e.uninit(value_dim)?;
let mut k_l2 = e.uninit(value_dim)?;
let mut v_gd = e.uninit(value_dim)?;
let mut beta_b = e.uninit(num_v)?;
let mut g_log = e.uninit(num_v)?;
for r in 0..t {
let base = toff + if r % 2 == 0 { 0 } else { 3 };
let conv_view = table.slice(base..base + 1);
let in_view = table.slice(base + 1..base + 2);
let out_view = table.slice(base + 2..base + 3);
e.ssm_conv1d_fused_decode_b_view(
&qkv_mixed.slice(r * qkv_w..(r + 1) * qkv_w),
&conv_view,
la.ssm_conv1d.float_data(),
&mut conv_out,
conv_dim,
d_conv,
1,
)?;
e.gdn_prep_decode_b_view(
&conv_out,
&beta_raw.slice(r * beta_w..(r + 1) * beta_w),
&alpha.slice(r * alpha_w..(r + 1) * alpha_w),
la.ssm_dt.float_data(),
la.ssm_a.float_data(),
&mut q_l2,
&mut k_l2,
&mut v_gd,
&mut beta_b,
&mut g_log,
d_state,
num_v,
num_k,
key_dim,
eps,
conv_dim,
1,
)?;
let mut o_row = o_all.slice_mut(r * value_dim..(r + 1) * value_dim);
e.gdn_scan_s128_batched_view(
&q_l2, &k_l2, &v_gd, &g_log, &beta_b, &in_view, &out_view, &mut o_row, num_v, 1,
gdn_scale,
)?;
if r + 1 < t {
let rl = cache.recur[il]
.as_ref()
.ok_or("qwen35 linear verify layer has no recurrent state")?;
let ssm_src = if r % 2 == 0 {
&rl.ssm_state_alt
} else {
&rl.ssm_state
};
match stash.as_mut() {
Some((conv_slab, ssm_slab)) => {
e.copy_indirect_src_f32(
&conv_view,
conv_slab,
r * conv_dim * (d_conv - 1),
conv_dim * (d_conv - 1),
)?;
e.copy_indirect_src_f32(
&out_view,
ssm_slab,
r * d_state * d_state * num_v,
d_state * d_state * num_v,
)?;
}
None => {
if let Some(states) = col_states.as_mut() {
states.push((e.clone_dtod(&rl.conv_state)?, e.clone_dtod(ssm_src)?));
}
}
}
}
}
if t % 2 == 1 {
let rl = cache.recur[il].as_mut().unwrap();
std::mem::swap(&mut rl.ssm_state, &mut rl.ssm_state_alt);
}
if let (Some(checkpoint), Some(states)) = (ckpt.as_deref_mut(), col_states) {
checkpoint.cols[il] = Some(states);
}
let mixed = if e.uses_q8_1_fast(&la.ssm_out) {
let (gq, gd) = e.gated_rmsnorm_q8_1(
&o_all,
la.ssm_norm.float_data(),
&z,
d_state,
t * num_v,
eps,
)?;
let g0 = e.zeros(0)?;
e.matmul_pre(&la.ssm_out, &gq, &gd, &g0, t)?
} else {
let mut gn = e.uninit(t * value_dim)?;
e.gated_rmsnorm(
&o_all,
la.ssm_norm.float_data(),
&z,
&mut gn,
d_state,
t * num_v,
eps,
)?;
e.matmul(&la.ssm_out, &gn, t)?
};
let pnorm = layer.post_attn_norm.float_data();
let mut x1 = e.uninit(t * n_embd)?;
let mut zn = e.uninit(t * n_embd)?;
e.add_rms_norm(x, &mixed, pnorm, &mut x1, &mut zn, n_embd, t, eps)?;
let ffn_out = match &layer.ffn {
crate::hybrid::Ffn::Dense {
ffn_gate,
ffn_up,
ffn_down,
} => {
assert!(
self.cfg.m3.is_none(),
"qwen35 t-parallel verify: M3 swigluoai FFN not yet batched"
);
self.qwen35_tparallel_dense_ffn(e, ffn_gate, ffn_up, ffn_down, &zn, t, n_embd)?
}
crate::hybrid::Ffn::Moe(m) => self.moe_ffn_il_zq8(e, m, &zn, None, t, il as u16)?,
};
let mut x2 = e.uninit(t * n_embd)?;
e.add(&x1, &ffn_out, &mut x2, t * n_embd)?;
self.dflash_tap(e, cache, il, &x2, t)?;
Ok(x2)
}
#[allow(clippy::too_many_arguments)]
fn verify_layers(
&self,
e: &Engine,
mut x: CudaSlice<f32>,
lo: usize,
hi: usize,
pos_d: &CudaSlice<i32>,
pos0: usize,
t: usize,
cache: &mut Cache,
mut ckpt: Option<&mut VerifyCkpt>,
stream: Option<(&CudaSlice<u32>, &CudaSlice<i32>)>,
graphs: Option<&mut DsparkVerifyGraphs>,
) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
if self.cfg.step35.is_some() {
if stream.is_some() {
return Err(
"step35 has no ROUND-STREAM verify arm (the device-counter _dc twins \
cannot express the SWA offset KV view)"
.into(),
);
}
return self.step35_verify_batch_layers(e, x, lo, hi, pos0, t, cache);
}
if self.qwen35_serving_class() {
return self.qwen35_verify_batch_layers(
e,
x,
lo,
hi,
pos0,
t,
cache,
ckpt.take(),
stream,
graphs,
);
}
let n_embd = self.cfg.n_embd as usize;
let eps = self.cfg.rms_eps;
let mut pending: Option<(CudaSlice<f32>, CudaSlice<f32>)> = None;
for il in lo..hi {
let layer = &self.layers[il];
let mixer_fast = self.mixer_in_q8_1_fast(e, &layer.mixer);
let norm_fused = std::env::var("MEMRA_NO_FUSE_NORMQ").is_err() && mixer_fast;
let lin_q8_only = match &layer.mixer {
Mixer::Linear(la) => {
(t >= 3 || (t == 2 && spec_m2())) && e.uses_q8_1_fast(&la.ssm_out)
}
Mixer::Full(_) if self.cfg.step35.is_some() => false,
_ => true,
};
let taken = pending.take();
let (h, h_q8) = if norm_fused && lin_q8_only {
let pair = match taken {
Some((x1p, f1p)) => {
let mut x2 = vbuf(e, t * n_embd)?; let p = e.add_rms_norm_q8_1(
&x1p,
&f1p,
layer.attn_norm.float_data(),
&mut x2,
n_embd,
t,
eps,
)?;
x = x2;
p
}
None => e.rms_norm_q8_1(&x, layer.attn_norm.float_data(), n_embd, t, eps)?,
};
(e.zeros(0)?, Some(pair)) } else {
if let Some((x1p, f1p)) = taken {
let mut x2 = vbuf(e, t * n_embd)?; e.add(&x1p, &f1p, &mut x2, t * n_embd)?;
x = x2;
}
let mut h = vbuf(e, t * n_embd)?; if norm_fused {
e.rms_norm_decode(&x, layer.attn_norm.float_data(), &mut h, n_embd, t, eps)?;
} else {
e.rms_norm(&x, layer.attn_norm.float_data(), &mut h, n_embd, t, eps)?;
}
(h, None)
};
let h_q8_ref = h_q8.as_ref().map(|(q, d)| (q, d));
let mixed = match &layer.mixer {
Mixer::Full(fa) => self.full_attn_verify(
e,
fa,
&h,
h_q8_ref,
pos_d,
t,
cache,
il,
stream.map(|(_, c)| c),
)?,
Mixer::Mla(_) => crate::hybrid::mla_forward_unimplemented(),
Mixer::Linear(la) => {
if (t >= 3 || (t == 2 && spec_m2()))
&& mixer_fast
&& e.uses_q8_1_fast(&la.ssm_out)
{
let want = ckpt.is_some();
let (out, stash) =
self.linear_attn_verify_t(e, la, &h, h_q8_ref, t, cache, il, want)?;
if let (Some(ck), Some(st)) = (ckpt.as_deref_mut(), stash) {
ck.gdn[il] = Some(st);
}
out
} else {
let mut out = vbuf(e, t * n_embd)?; let mut col_states: Option<Vec<(CudaSlice<f32>, CudaSlice<f32>)>> =
if ckpt.is_some() && t >= 2 {
Some(Vec::with_capacity(t - 1))
} else {
None
};
for col in 0..t {
let mut h_col = vbuf(e, n_embd)?; let src = h.slice(col * n_embd..(col + 1) * n_embd);
e.copy_view_into(&mut h_col, 0, &src, n_embd)?;
let m_col = self.linear_attn_decode(e, la, &h_col, cache, il)?;
e.copy_into(&mut out, col * n_embd, &m_col, n_embd)?;
if let Some(cs) = col_states.as_mut() {
if col + 1 < t {
let rl = cache.recur[il].as_ref().unwrap();
cs.push((
e.clone_dtod(&rl.conv_state)?,
e.clone_dtod(&rl.ssm_state)?,
));
}
}
}
if let (Some(ck), Some(cs)) = (ckpt.as_deref_mut(), col_states) {
if std::env::var("MEMRA_SPEC_STATS").as_deref() == Ok("1") {
static ONCE: std::sync::Once = std::sync::Once::new();
let bytes: usize =
cs.iter().map(|(c, s)| (c.len() + s.len()) * 4).sum();
ONCE.call_once(|| eprintln!(
"[verify-ckpt] per-column layer il={il}: {} clones, {:.2} MB/layer/round",
cs.len(), bytes as f64 / 1e6));
}
ck.cols[il] = Some(cs);
}
out
}
}
};
let ffn_fuse = match &layer.ffn {
crate::hybrid::Ffn::Dense {
ffn_gate, ffn_up, ..
} => {
std::env::var("MEMRA_NO_FUSE_NORMQ").is_err()
&& e.uses_q8_1_fast(ffn_gate)
&& e.uses_q8_1_fast(ffn_up)
}
crate::hybrid::Ffn::Moe(_) => false,
};
let dense_lim = self.cfg.clamp_shexp_at(il as u32);
let fuse_q8 = ffn_fuse && self.cfg.m3.is_none() && dense_lim.is_none();
let mut x1 = vbuf(e, t * n_embd)?; let mut z = e.zeros(0)?; let z_q8 = if fuse_q8 {
Some(e.add_rms_norm_q8_1(
&x,
&mixed,
layer.post_attn_norm.float_data(),
&mut x1,
n_embd,
t,
eps,
)?)
} else {
let mut zf = vbuf(e, t * n_embd)?; if ffn_fuse {
e.add(&x, &mixed, &mut x1, t * n_embd)?;
e.rms_norm_decode(
&x1,
layer.post_attn_norm.float_data(),
&mut zf,
n_embd,
t,
eps,
)?;
} else {
e.add_rms_norm(
&x,
&mixed,
layer.post_attn_norm.float_data(),
&mut x1,
&mut zf,
n_embd,
t,
eps,
)?;
}
z = zf;
None
};
let ffn_out = match &layer.ffn {
crate::hybrid::Ffn::Dense {
ffn_gate,
ffn_up,
ffn_down,
} => {
let n_ff = ffn_gate.out_features();
if let Some((zq, zd)) = z_q8.as_ref() {
let pair =
match e.matmul_decode_exact_dual_pre(ffn_gate, ffn_up, zq, zd, t)? {
Some(((g, gs), (u, us))) => Some((g, gs, u, us)),
None => None,
};
let (gate, gs, up, us) = match pair {
Some(x4) => x4,
None => (
e.matmul_decode_exact_pre(ffn_gate, zq, zd, t)?,
1.0, e.matmul_decode_exact_pre(ffn_up, zq, zd, t)?,
1.0,
),
};
if e.uses_q8_1_fast(ffn_down) {
let (aq, ad) = e.silu_mul_scaled_q8_1(&gate, &up, gs, us, t * n_ff)?;
e.matmul_decode_exact_pre(ffn_down, &aq, &ad, t)?
} else {
let mut act = vbuf(e, t * n_ff)?;
e.silu_mul_scaled(&gate, &up, gs, us, &mut act, t * n_ff)?;
e.matmul_decode_exact(ffn_down, &act, t)?
}
} else {
let (gate, up) =
match e.matmul_decode_exact_dual(ffn_gate, ffn_up, &z, t)? {
Some(pair) => pair,
None => (
e.matmul_decode_exact(ffn_gate, &z, t)?,
e.matmul_decode_exact(ffn_up, &z, t)?,
),
};
let mut act = vbuf(e, t * n_ff)?; Self::ffn_act_lim(
e,
&self.cfg,
&gate,
&up,
1.0,
1.0,
dense_lim,
&mut act,
t * n_ff,
)?;
e.matmul_decode_exact(ffn_down, &act, t)?
}
}
crate::hybrid::Ffn::Moe(m) => self.moe_ffn_il(e, m, &z, t, il as u16)?,
};
pending = Some((x1, ffn_out));
}
if let Some((x1p, f1p)) = pending.take() {
let mut x2 = vbuf(e, t * n_embd)?; e.add(&x1p, &f1p, &mut x2, t * n_embd)?;
x = x2;
}
Ok(x)
}
#[allow(clippy::too_many_arguments)]
fn linear_attn_verify_t(
&self,
e: &Engine,
la: &LinearAttnLayer,
h: &CudaSlice<f32>,
h_q8: Option<(&CudaSlice<i8>, &CudaSlice<f32>)>,
t: usize,
cache: &mut Cache,
il: usize,
want_stash: bool,
) -> Result<(CudaSlice<f32>, Option<GdnStash>), Box<dyn std::error::Error>> {
let cfg = &self.cfg;
let ssm = cfg.ssm.as_ref().unwrap();
let d_state = ssm.state_size as usize;
let num_k = ssm.group_count as usize;
let num_v = ssm.time_step_rank as usize;
let d_conv = ssm.conv_kernel as usize;
let key_dim = d_state * num_k;
let conv_dim = key_dim * 2 + d_state * num_v;
let eps = cfg.rms_eps;
let scale = 1.0 / (d_state as f32).sqrt();
let h_q8_t = if h_q8.is_none()
&& spec_fused_t()
&& (2..=4).contains(&t)
&& ((e.uses_q8_1_fast(&la.wqkv) && e.uses_q8_1_fast(&la.wqkv_gate))
|| (e.uses_q8_1_fast(&la.ssm_beta) && e.uses_q8_1_fast(&la.ssm_alpha)))
{
Some(e.quantize_q8_1(h, t, cfg.n_embd as usize)?)
} else {
None
};
let hq8_any: Option<(&CudaSlice<i8>, &CudaSlice<f32>)> =
h_q8.or(h_q8_t.as_ref().map(|(q, d)| (q, d)));
let (qkv_mixed, z) = {
let mut fused = None;
if t == 1 && e.uses_q8_1_fast(&la.wqkv) && e.uses_q8_1_fast(&la.wqkv_gate) {
let (hq, hd) = e.quantize_q8_1(h, 1, cfg.n_embd as usize)?;
fused = e.matmul_q8_fused2(&la.wqkv, &la.wqkv_gate, &hq, &hd)?;
} else if let Some((hq, hd)) = hq8_any {
if spec_fused_t() && (2..=4).contains(&t) {
fused = e.matmul_q8_fused2_t(&la.wqkv, &la.wqkv_gate, hq, hd, t)?;
}
}
match (fused, hq8_any) {
(Some(pair), _) => pair,
(None, Some((hq, hd))) if h_q8.is_some() => (
e.matmul_decode_exact_pre(&la.wqkv, hq, hd, t)?,
e.matmul_decode_exact_pre(&la.wqkv_gate, hq, hd, t)?,
),
(None, _) => (
e.matmul_decode_exact(&la.wqkv, h, t)?,
e.matmul_decode_exact(&la.wqkv_gate, h, t)?,
),
}
};
let (beta_raw, alpha) = if t == 1 {
let (hq, hd) = e.quantize_q8_1(h, 1, cfg.n_embd as usize)?;
match e.matmul_pre_dual_noscale(&la.ssm_beta, &la.ssm_alpha, &hq, &hd, 1)? {
Some(((mut b, bs), (mut a, as_))) => {
if bs != 1.0 {
e.scale_inplace(&mut b, bs, la.ssm_beta.out_features())?;
}
if as_ != 1.0 {
e.scale_inplace(&mut a, as_, la.ssm_alpha.out_features())?;
}
(b, a)
}
None => match e.matmul_q8_fused2(&la.ssm_beta, &la.ssm_alpha, &hq, &hd)? {
Some((b, a)) => (b, a),
None => (
e.matmul_decode_exact(&la.ssm_beta, h, 1)?,
e.matmul_decode_exact(&la.ssm_alpha, h, 1)?,
),
},
}
} else {
let mut nvfp4_fused = None;
let mut q8_fused = None;
if let Some((hq, hd)) = hq8_any {
if t == 3 && std::env::var("MEMRA_NVFP4_AUX_DUAL").as_deref() != Ok("0") {
nvfp4_fused =
e.matmul_decode_exact_dual_pre(&la.ssm_beta, &la.ssm_alpha, hq, hd, t)?;
if nvfp4_fused.is_some() && std::env::var("MEMRA_DEBUG").is_ok() {
static ONCE: std::sync::Once = std::sync::Once::new();
ONCE.call_once(|| {
eprintln!("[memra] NVFP4 beta+alpha batched aux dual ENGAGED (t={t})")
});
}
}
if nvfp4_fused.is_none() && spec_fused_t() && (2..=4).contains(&t) {
q8_fused = e.matmul_q8_fused2_t(&la.ssm_beta, &la.ssm_alpha, hq, hd, t)?;
}
}
if let Some(((mut b, bs), (mut a, as_))) = nvfp4_fused {
if bs != 1.0 {
e.scale_inplace(&mut b, bs, t * la.ssm_beta.out_features())?;
}
if as_ != 1.0 {
e.scale_inplace(&mut a, as_, t * la.ssm_alpha.out_features())?;
}
(b, a)
} else if let Some(pair) = q8_fused {
pair
} else {
match hq8_any {
Some((hq, hd)) if h_q8.is_some() => (
e.matmul_decode_exact_pre(&la.ssm_beta, hq, hd, t)?,
e.matmul_decode_exact_pre(&la.ssm_alpha, hq, hd, t)?,
),
_ => (
e.matmul_decode_exact(&la.ssm_beta, h, t)?,
e.matmul_decode_exact(&la.ssm_alpha, h, t)?,
),
}
}
};
let rl = cache.recur[il].as_mut().unwrap();
let mut conv_out = e.uninit(conv_dim * t)?;
e.ssm_conv1d_tm_state(
&qkv_mixed,
&mut rl.conv_state,
la.ssm_conv1d.float_data(),
&mut conv_out,
conv_dim,
t,
d_conv,
)?;
let mut q_g = e.uninit(d_state * num_v * t)?;
let mut k_g = e.uninit(d_state * num_v * t)?;
let mut v_g = e.uninit(d_state * num_v * t)?;
e.qkv_to_gdn_repack(
&conv_out, &mut q_g, &mut k_g, &mut v_g, d_state, num_v, num_k, key_dim, t,
)?;
let mut q_l2 = e.uninit(d_state * num_v * t)?;
e.l2_norm_decode(&q_g, &mut q_l2, d_state, num_v * t, eps)?;
let mut k_l2 = e.uninit(d_state * num_v * t)?;
e.l2_norm_decode(&k_g, &mut k_l2, d_state, num_v * t, eps)?;
let mut beta = e.uninit(t * num_v)?;
e.sigmoid(&beta_raw, &mut beta, t * num_v)?;
let mut g_log = e.uninit(t * num_v)?;
e.gdn_glog(
&alpha,
la.ssm_dt.float_data(),
la.ssm_a.float_data(),
&mut g_log,
num_v,
t,
)?;
let mut o = e.uninit(d_state * num_v * t)?;
{
let crate::cache::RecurLayer {
ssm_state,
ssm_state_alt,
..
} = rl;
e.gdn_scan_s128(
&q_l2,
&k_l2,
&v_g,
&g_log,
&beta,
ssm_state,
ssm_state_alt,
&mut o,
num_v,
t,
scale,
)?;
}
std::mem::swap(&mut rl.ssm_state, &mut rl.ssm_state_alt);
let out = if e.uses_q8_1_fast(&la.ssm_out) {
let (gq, gd) =
e.gated_rmsnorm_q8_1(&o, la.ssm_norm.float_data(), &z, d_state, num_v * t, eps)?;
e.matmul_decode_exact_pre(&la.ssm_out, &gq, &gd, t)?
} else {
let mut gn = e.uninit(d_state * num_v * t)?;
e.gated_rmsnorm(
&o,
la.ssm_norm.float_data(),
&z,
&mut gn,
d_state,
num_v * t,
eps,
)?;
e.matmul_decode_exact(&la.ssm_out, &gn, t)?
};
let stash = if want_stash {
Some(GdnStash {
qkv_mixed,
q_l2,
k_l2,
v_g,
g_log,
beta,
})
} else {
None
};
Ok((out, stash))
}
fn commit_verified_prefix(
&self,
e: &Engine,
cache: &mut Cache,
snap: &crate::cache::CacheSnapshot,
ckpt: &VerifyCkpt,
j: usize,
kv_lens_done: bool,
dev_j: Option<(&CudaSlice<u32>, usize, usize)>,
) -> Result<(), Box<dyn std::error::Error>> {
let cfg = &self.cfg;
let ssm = cfg.ssm.as_ref().unwrap();
let d_state = ssm.state_size as usize;
let num_k = ssm.group_count as usize;
let num_v = ssm.time_step_rank as usize;
let d_conv = ssm.conv_kernel as usize;
let conv_dim = d_state * num_k * 2 + d_state * num_v;
let scale = 1.0 / (d_state as f32).sqrt();
let mut batched_cols = false;
if state_copy_batch_on() && dev_j.is_none() {
use cudarc::driver::DevicePtr;
let s = &e.gpu.stream();
let mut conv_pairs: Vec<(u64, u64)> = Vec::new();
let mut ssm_pairs: Vec<(u64, u64)> = Vec::new();
let (mut conv_words, mut ssm_words) = (0usize, 0usize);
let mut uniform = true;
for il in 0..self.layers.len() {
let Some(rl) = cache.recur[il].as_ref() else {
continue;
};
if ckpt.gdn[il].is_some() {
continue; }
let Some(cols) = &ckpt.cols[il] else {
continue; };
let (c, st) = &cols[j - 1];
if conv_pairs.is_empty() {
conv_words = c.len();
ssm_words = st.len();
} else if c.len() != conv_words || st.len() != ssm_words {
uniform = false;
break;
}
let (pc, _g0) = c.device_ptr(s);
let (dc, _g1) = rl.conv_state.device_ptr(s);
let (ps, _g2) = st.device_ptr(s);
let (ds, _g3) = rl.ssm_state.device_ptr(s);
conv_pairs.push((pc as u64, dc as u64));
ssm_pairs.push((ps as u64, ds as u64));
}
if uniform && !conv_pairs.is_empty() {
let n = conv_pairs.len();
let mut t = vec![0u64; 2 * n];
for (k, &(src, dst)) in conv_pairs.iter().enumerate() {
t[k] = src;
t[n + k] = dst;
}
let conv_t = e.htod_u64(&t)?;
for (k, &(src, dst)) in ssm_pairs.iter().enumerate() {
t[k] = src;
t[n + k] = dst;
}
let ssm_t = e.htod_u64(&t)?;
e.copy_batch_uniform_f32(&conv_t, n, conv_words)?;
e.copy_batch_uniform_f32(&ssm_t, n, ssm_words)?;
batched_cols = true;
}
}
for il in 0..self.layers.len() {
if let (Some(kvl), Some(saved)) = (cache.kv[il].as_mut(), snap.kv_len[il]) {
kvl.len = saved + j;
if !kv_lens_done {
e.set_i32_one(&mut kvl.len_d, kvl.len as i32)?;
}
}
if let Some(rl) = cache.recur[il].as_mut() {
if let Some(st) = &ckpt.gdn[il] {
let ring_old = snap.conv[il].as_ref().expect("snapshot missing conv");
let state_in = snap.ssm[il].as_ref().expect("snapshot missing ssm");
if let Some((acc, base, t_v)) = dev_j {
e.ssm_conv_ring_rebuild_dc(
&st.qkv_mixed,
ring_old,
&mut rl.conv_state,
conv_dim,
acc,
base,
t_v,
d_conv,
)?;
let mut o = e.uninit(d_state * num_v * j.max(1))?;
e.gdn_scan_s128_dc(
&st.q_l2,
&st.k_l2,
&st.v_g,
&st.g_log,
&st.beta,
state_in,
&mut rl.ssm_state,
&mut o,
num_v,
acc,
base,
t_v,
scale,
)?;
} else {
e.ssm_conv_ring_rebuild(
&st.qkv_mixed,
ring_old,
&mut rl.conv_state,
conv_dim,
j,
d_conv,
)?;
let mut o = e.uninit(d_state * num_v * j)?; e.gdn_scan_s128(
&st.q_l2,
&st.k_l2,
&st.v_g,
&st.g_log,
&st.beta,
state_in,
&mut rl.ssm_state,
&mut o,
num_v,
j,
scale,
)?;
}
} else if let Some(cols) = &ckpt.cols[il] {
if !batched_cols {
let (c, s) = &cols[j - 1];
e.copy_into(&mut rl.conv_state, 0, c, c.len())?;
e.copy_into(&mut rl.ssm_state, 0, s, s.len())?;
}
} else {
return Err(
"commit_verified_prefix: verify ckpt missing for linear layer".into(),
);
}
}
}
cache.pos = snap.pos + j;
Ok(())
}
fn commit_verified_prefix_stream(
&self,
e: &Engine,
cache: &mut Cache,
snap: &crate::cache::CacheSnapshot,
ckpt: &VerifyCkpt,
acc: &CudaSlice<u32>,
base: usize,
t_v: usize,
) -> Result<(), Box<dyn std::error::Error>> {
let cfg = &self.cfg;
let ssm = cfg.ssm.as_ref().unwrap();
let d_state = ssm.state_size as usize;
let num_k = ssm.group_count as usize;
let num_v = ssm.time_step_rank as usize;
let d_conv = ssm.conv_kernel as usize;
let conv_dim = d_state * num_k * 2 + d_state * num_v;
let scale = 1.0 / (d_state as f32).sqrt();
for il in 0..self.layers.len() {
if let Some(rl) = cache.recur[il].as_mut() {
let st = ckpt.gdn[il]
.as_ref()
.ok_or("stream restore: batched-linear stash missing")?;
let ring_old = snap.conv[il].as_ref().expect("snapshot missing conv");
let state_in = snap.ssm[il].as_ref().expect("snapshot missing ssm");
e.ssm_conv_ring_rebuild_dc(
&st.qkv_mixed,
ring_old,
&mut rl.conv_state,
conv_dim,
acc,
base,
t_v,
d_conv,
)?;
let mut o = e.uninit(d_state * num_v * t_v)?;
e.gdn_scan_s128_dc(
&st.q_l2,
&st.k_l2,
&st.v_g,
&st.g_log,
&st.beta,
state_in,
&mut rl.ssm_state,
&mut o,
num_v,
acc,
base,
t_v,
scale,
)?;
}
}
Ok(())
}
pub fn decode_step_t_aux2(
&self,
e: &Engine,
tokens: &[u32],
pos0: usize,
cache: &mut Cache,
aux_layers: &[usize],
pred_col: Option<usize>,
) -> Result<
(Vec<f32>, Vec<CudaSlice<f32>>, Option<Vec<CudaSlice<f32>>>),
Box<dyn std::error::Error>,
> {
let cfg = &self.cfg;
let n_embd = cfg.n_embd as usize;
let eps = cfg.rms_eps;
let t = tokens.len();
let pos_vec: Vec<i32> = (0..t).map(|i| (pos0 + i) as i32).collect();
let pos_d = e.htod_i32(&pos_vec)?;
let mut x = e.htod(&self.embd.gather(n_embd, tokens))?;
let mut aux_last: Vec<CudaSlice<f32>> = Vec::with_capacity(aux_layers.len());
let mut aux_pred: Vec<CudaSlice<f32>> = Vec::new();
let want_pred = pred_col.is_some();
for (il, layer) in self.layers.iter().enumerate() {
let mixer_fast = self.mixer_in_q8_1_fast(e, &layer.mixer);
let norm_fused = std::env::var("MEMRA_NO_FUSE_NORMQ").is_err() && mixer_fast;
let mut h = vbuf(e, t * n_embd)?; if norm_fused {
e.rms_norm_decode(&x, layer.attn_norm.float_data(), &mut h, n_embd, t, eps)?;
} else {
e.rms_norm(&x, layer.attn_norm.float_data(), &mut h, n_embd, t, eps)?;
}
let mixed = match &layer.mixer {
Mixer::Full(fa) => {
self.full_attn_verify(e, fa, &h, None, &pos_d, t, cache, il, None)?
}
Mixer::Mla(_) => crate::hybrid::mla_forward_unimplemented(),
Mixer::Linear(la) => {
let mut out = e.zeros(t * n_embd)?;
for col in 0..t {
let mut h_col = e.zeros(n_embd)?;
let src = h.slice(col * n_embd..(col + 1) * n_embd);
e.copy_view_into(&mut h_col, 0, &src, n_embd)?;
let m_col = self.linear_attn_decode(e, la, &h_col, cache, il)?;
e.copy_into(&mut out, col * n_embd, &m_col, n_embd)?;
}
out
}
};
let ffn_fuse = match &layer.ffn {
crate::hybrid::Ffn::Dense {
ffn_gate, ffn_up, ..
} => {
std::env::var("MEMRA_NO_FUSE_NORMQ").is_err()
&& e.uses_q8_1_fast(ffn_gate)
&& e.uses_q8_1_fast(ffn_up)
}
crate::hybrid::Ffn::Moe(_) => false,
};
let mut x1 = vbuf(e, t * n_embd)?; let mut z = vbuf(e, t * n_embd)?; if ffn_fuse {
e.add(&x, &mixed, &mut x1, t * n_embd)?;
e.rms_norm_decode(
&x1,
layer.post_attn_norm.float_data(),
&mut z,
n_embd,
t,
eps,
)?;
} else {
e.add_rms_norm(
&x,
&mixed,
layer.post_attn_norm.float_data(),
&mut x1,
&mut z,
n_embd,
t,
eps,
)?;
}
let ffn_out = match &layer.ffn {
crate::hybrid::Ffn::Dense {
ffn_gate,
ffn_up,
ffn_down,
} => {
let n_ff = ffn_gate.out_features();
let gate = e.matmul_decode_exact(ffn_gate, &z, t)?;
let up = e.matmul_decode_exact(ffn_up, &z, t)?;
let mut act = vbuf(e, t * n_ff)?; Self::ffn_act_lim(
e,
&self.cfg,
&gate,
&up,
1.0,
1.0,
self.cfg.clamp_shexp_at(il as u32),
&mut act,
t * n_ff,
)?;
e.matmul_decode_exact(ffn_down, &act, t)?
}
crate::hybrid::Ffn::Moe(m) => self.moe_ffn_il(e, m, &z, t, il as u16)?,
};
let mut x2 = vbuf(e, t * n_embd)?; e.add(&x1, &ffn_out, &mut x2, t * n_embd)?;
if aux_layers.contains(&il) {
let mut a = e.zeros(n_embd)?;
e.copy_view_into(&mut a, 0, &x2.slice((t - 1) * n_embd..t * n_embd), n_embd)?;
aux_last.push(a);
if let Some(pc) = pred_col {
let mut ap = e.zeros(n_embd)?;
e.copy_view_into(
&mut ap,
0,
&x2.slice(pc * n_embd..(pc + 1) * n_embd),
n_embd,
)?;
aux_pred.push(ap);
}
}
x = x2;
}
let mut hn = vbuf(e, t * n_embd)?; e.rms_norm_decode(&x, self.output_norm.float_data(), &mut hn, n_embd, t, eps)?;
let logits = e.matmul_decode_exact(&self.output, &hn, t)?;
let host = e.dtoh(&logits)?;
cache.pos += t;
Ok((
host,
aux_last,
if want_pred { Some(aux_pred) } else { None },
))
}
#[allow(clippy::too_many_arguments)]
fn step35_verify(
&self,
e: &Engine,
fa: &FullAttnLayer,
h: &CudaSlice<f32>,
h_q8: Option<(&CudaSlice<i8>, &CudaSlice<f32>)>,
t: usize,
cache: &mut Cache,
il: usize,
) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
let n_embd = self.cfg.n_embd as usize;
assert_eq!(
h.len(),
t * n_embd,
"step35_verify needs the f32 attn-normed rows ([t*n_embd]); the caller took the \
fused q8-only norm arm (h_q8={}) — step35 must stay on the unfused arm",
h_q8.is_some()
);
let mut out = vbuf(e, t * n_embd)?; for r in 0..t {
let pos_d = e.htod_i32(&[(cache.pos + r) as i32])?;
let mut h_row = vbuf(e, n_embd)?; e.copy_view_into(
&mut h_row,
0,
&h.slice(r * n_embd..(r + 1) * n_embd),
n_embd,
)?;
let o = self.step35_decode_attn(e, fa, il, &h_row, None, &pos_d, cache)?;
debug_assert_eq!(
o.len(),
n_embd,
"step35_decode_attn returns post-wo [n_embd]"
);
e.copy_into(&mut out, r * n_embd, &o, n_embd)?;
}
Ok(out)
}
#[allow(clippy::too_many_arguments)]
fn full_attn_verify(
&self,
e: &Engine,
fa: &FullAttnLayer,
h: &CudaSlice<f32>,
h_q8: Option<(&CudaSlice<i8>, &CudaSlice<f32>)>,
pos_d: &CudaSlice<i32>,
t: usize,
cache: &mut Cache,
il: usize,
stream_ctr: Option<&CudaSlice<i32>>,
) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
if self.cfg.step35.is_some() {
if stream_ctr.is_some() {
return Err(
"step35 has no ROUND-STREAM verify arm (the device-counter _dc twins \
cannot express the SWA offset KV view; same root cause as the dc \
decode refusal) — run spec without the stream arm"
.into(),
);
}
return self.step35_verify(e, fa, h, h_q8, t, cache, il);
}
let cfg = &self.cfg;
let geometry = cfg.full_attention_geometry_at(il as u32);
let n_head = geometry.n_head as usize;
let n_head_kv = geometry.n_head_kv as usize;
let head_dim = geometry.head_dim_k as usize;
let eps = cfg.rms_eps;
let scale = geometry.attention_scale();
let n_embd = cfg.n_embd as usize;
let (qf, mut k, v) = {
let mut fused = None;
let qkv_fast =
e.uses_q8_1_fast(&fa.wq) && e.uses_q8_1_fast(&fa.wk) && e.uses_q8_1_fast(&fa.wv);
if t == 1 && qkv_fast {
let (hq_o, hd_o);
let (hq, hd): (&CudaSlice<i8>, &CudaSlice<f32>) = match h_q8 {
Some(p) => p,
None => {
(hq_o, hd_o) = e.quantize_q8_1(h, 1, n_embd)?;
(&hq_o, &hd_o)
}
};
fused = e.matmul_q8_fused3(&fa.wq, &fa.wk, &fa.wv, hq, hd)?;
} else if spec_fused_t() && (2..=4).contains(&t) && qkv_fast {
let (hq_o, hd_o);
let (hq, hd): (&CudaSlice<i8>, &CudaSlice<f32>) = match h_q8 {
Some(p) => p,
None => {
(hq_o, hd_o) = e.quantize_q8_1(h, t, n_embd)?;
(&hq_o, &hd_o)
}
};
fused = e.matmul_q8_fused3_t(&fa.wq, &fa.wk, &fa.wv, hq, hd, t)?;
}
match (fused, h_q8) {
(Some(triple), _) => triple,
(None, Some((hq, hd))) if qkv_fast => (
e.matmul_decode_exact_pre(&fa.wq, hq, hd, t)?,
e.matmul_decode_exact_pre(&fa.wk, hq, hd, t)?,
e.matmul_decode_exact_pre(&fa.wv, hq, hd, t)?,
),
(None, _) => (
e.matmul_decode_exact(&fa.wq, h, t)?,
e.matmul_decode_exact(&fa.wk, h, t)?,
e.matmul_decode_exact(&fa.wv, h, t)?,
),
}
};
let gated = geometry.attention_gate == memra_gguf::config::AttentionGateKind::FusedQ;
let (mut q, gate) = if gated {
let mut q = vbuf(e, t * n_head * head_dim)?; let mut gate = vbuf(e, t * n_head * head_dim)?; e.q_gate_split(&qf, &mut q, &mut gate, head_dim, n_head, t)?;
(q, Some(gate))
} else {
(qf, None)
};
let mut qn = vbuf(e, t * n_head * head_dim)?; e.rms_norm(
&q,
fa.q_norm.float_data(),
&mut qn,
head_dim,
n_head * t,
eps,
)?;
q = qn;
let mut kn = vbuf(e, t * n_head_kv * head_dim)?; e.rms_norm(
&k,
fa.k_norm.float_data(),
&mut kn,
head_dim,
n_head_kv * t,
eps,
)?;
k = kn;
let rope_dims = geometry.n_rot as usize;
e.rope_neox(
&mut q,
pos_d,
head_dim,
rope_dims,
n_head,
t,
geometry.rope_base,
1.0,
)?;
e.rope_neox(
&mut k,
pos_d,
head_dim,
rope_dims,
n_head_kv,
t,
geometry.rope_base,
1.0,
)?;
let kvl = cache.kv[il].as_mut().unwrap();
let (kv_dim_k, kv_dim_v, ktb, vtb) =
(kvl.kv_dim_k, kvl.kv_dim_v, kvl.k_tok_bytes, kvl.v_tok_bytes);
if let Some(ctr) = stream_ctr {
e.append_kv_quantized_rows_dc(
&k,
&v,
&mut kvl.k,
&mut kvl.v,
ctr,
t,
kv_dim_k,
kv_dim_v,
ktb,
vtb,
crate::Engine::kv_fp8_on(),
)?;
} else {
for i in 0..t {
let k_row = k.slice(i * kv_dim_k..(i + 1) * kv_dim_k);
let v_row = v.slice(i * kv_dim_v..(i + 1) * kv_dim_v);
e.append_kv_quantized_view(
&k_row,
&v_row,
&mut kvl.k,
&mut kvl.v,
kvl.len + i,
kv_dim_k,
kv_dim_v,
ktb,
vtb,
crate::Engine::kv_fp8_on(),
)?;
}
kvl.len += t;
}
let mut attn = vbuf(e, t * n_head * head_dim)?; let base_len = kvl.len - t; if let Some(ctr) = stream_ctr {
let upper = kvl.len + t + 64;
let k_view = e.view_u8(&kvl.k, (upper.min(cache.max_ctx)) * ktb);
let v_view = e.view_u8(&kvl.v, (upper.min(cache.max_ctx)) * vtb);
e.fa_decode_rows_dc(
&q,
&k_view,
&v_view,
&mut attn,
head_dim,
n_head,
n_head_kv,
ctr,
upper.min(cache.max_ctx),
t,
scale,
ktb,
vtb,
0,
false,
)?;
} else if spec_lean() && t == 1 {
let t_kv = base_len + 1;
let k_view = e.view_u8(&kvl.k, t_kv * ktb);
let v_view = e.view_u8(&kvl.v, t_kv * vtb);
e.fa_decode_kvmod(
&q,
&k_view,
&v_view,
&mut attn,
head_dim,
n_head,
n_head_kv,
t_kv,
scale,
ktb,
vtb,
crate::Engine::kv_fp8_on(),
)?;
} else if e.fa_rows_eligible(base_len, head_dim) {
let k_view = e.view_u8(&kvl.k, (base_len + t) * ktb);
let v_view = e.view_u8(&kvl.v, (base_len + t) * vtb);
e.fa_decode_rows(
&q,
&k_view,
&v_view,
&mut attn,
head_dim,
n_head,
n_head_kv,
base_len,
t,
scale,
ktb,
vtb,
None,
false,
crate::Engine::kv_fp8_on(),
None,
)?;
} else {
for r in 0..t {
let t_kv_r = base_len + r + 1; let k_view_r = e.view_u8(&kvl.k, t_kv_r * ktb);
let v_view_r = e.view_u8(&kvl.v, t_kv_r * vtb);
let mut q_row = vbuf(e, n_head * head_dim)?; let q_src = q.slice(r * n_head * head_dim..(r + 1) * n_head * head_dim);
e.copy_view_into(&mut q_row, 0, &q_src, n_head * head_dim)?;
let mut attn_row = vbuf(e, n_head * head_dim)?; e.fa_decode_kvmod(
&q_row,
&k_view_r,
&v_view_r,
&mut attn_row,
head_dim,
n_head,
n_head_kv,
t_kv_r,
scale,
ktb,
vtb,
crate::Engine::kv_fp8_on(),
)?;
e.copy_into(
&mut attn,
r * n_head * head_dim,
&attn_row,
n_head * head_dim,
)?;
}
}
let attn_g = match &gate {
Some(gate) => {
let mut gsig = vbuf(e, t * n_head * head_dim)?; e.sigmoid(gate, &mut gsig, t * n_head * head_dim)?;
let mut ag = vbuf(e, t * n_head * head_dim)?; e.mul(&attn, &gsig, &mut ag, t * n_head * head_dim)?;
ag
}
None => attn,
};
Ok(e.matmul_decode_exact(&fa.wo, &attn_g, t)?)
}
pub fn plain_session_kv_bytes_per_token(&self) -> usize {
crate::cache::cache_bytes_per_token(&self.cfg)
}
pub fn plain_session_kv_shape(&self) -> (usize, usize, usize) {
(
self.plain_session_kv_bytes_per_token(),
crate::cache::cache_ring_bytes_per_token(&self.cfg),
crate::cache::cache_ring_row_cap(&self.cfg),
)
}
pub fn spec_session_kv_bytes_per_token(&self) -> usize {
let scratch = self
.mtp
.as_ref()
.map(|mtp| {
let (_, _, k, v) = mtp_scratch_layout(&self.cfg, mtp.geom.as_ref());
k + v
})
.unwrap_or(0);
self.plain_session_kv_bytes_per_token()
.saturating_add(scratch)
}
pub fn spec_session_kv_shape(&self) -> (usize, usize, usize) {
let total = self.spec_session_kv_bytes_per_token();
let (_, mut ring, rows) = self.plain_session_kv_shape();
if rows > 0 {
ring = ring.saturating_add(
self.mtp
.as_ref()
.map(|mtp| {
let (_, _, k, v) = mtp_scratch_layout(&self.cfg, mtp.geom.as_ref());
k + v
})
.unwrap_or(0),
);
}
(total, ring, rows)
}
pub fn new_session(
&self,
e: &Engine,
max_ctx: usize,
) -> Result<SpecSession, Box<dyn std::error::Error>> {
Ok(SpecSession {
cache: crate::pp::new_cache(e, &self.cfg, max_ctx)?,
scratch: MtpScratch::new(
e,
&self.cfg,
max_ctx,
self.mtp.as_ref().and_then(|m| m.geom.as_ref()),
)?,
committed: Vec::new(),
last_h: None,
next_pred: None,
sctr: 0,
uctr: 0,
draft_ctx: None,
pending_tok: None,
turn_ckpt: None,
telem: SpecTelemetryCounters::default(),
capture_at: None,
boundary_captures: Vec::new(),
ckpt_at: None,
})
}
#[allow(clippy::too_many_arguments)]
pub fn spec_session_from_restored(
&self,
e: &Engine,
mut cache: Cache,
prefix: Vec<u32>,
suffix: &[u32],
draft_k: &CudaSlice<u8>,
draft_v: &CudaSlice<u8>,
draft_k_tok_bytes: usize,
draft_v_tok_bytes: usize,
draft_len: usize,
last_h: &[f32],
boundary_logits: &[f32],
sampling: Option<SpecSampling>,
require_anchor: bool,
max_ctx: usize,
republish_at: Option<usize>,
) -> Result<SpecSession, (Option<Cache>, String)> {
let pos = prefix.len();
let fail = |cache: Cache, msg: String| -> Result<SpecSession, (Option<Cache>, String)> {
Err((Some(cache), msg))
};
if self.mtp.is_none() {
return fail(cache, "no MTP head attached (nothing to draft with)".into());
}
if pos == 0 {
return fail(cache, "empty committed prefix".into());
}
if cache.pos != pos {
let msg = format!(
"restored cache pos {} != restored prefix len {pos}",
cache.pos
);
return fail(cache, msg);
}
if draft_len != pos {
return fail(
cache,
format!("draft plane len {draft_len} != restored prefix len {pos}"),
);
}
if pos + suffix.len() >= max_ctx {
return fail(
cache,
format!(
"prompt {} + suffix would not leave generation room in ctx {max_ctx}",
pos + suffix.len(),
),
);
}
let mut scratch = match MtpScratch::new(
e,
&self.cfg,
max_ctx,
self.mtp.as_ref().and_then(|m| m.geom.as_ref()),
) {
Ok(s) => s,
Err(err) => return fail(cache, format!("draft scratch alloc failed: {err}")),
};
if scratch.kv.ring.is_some() {
return fail(
cache,
"ring-backed draft scratch (Step35 SWA) cannot take a flat prefix restore".into(),
);
}
if scratch.kv.k_tok_bytes != draft_k_tok_bytes
|| scratch.kv.v_tok_bytes != draft_v_tok_bytes
{
return fail(
cache,
format!(
"draft plane layout {draft_k_tok_bytes}/{draft_v_tok_bytes} != scratch \
{}/{} bytes/token (stale entry across a format change)",
scratch.kv.k_tok_bytes, scratch.kv.v_tok_bytes,
),
);
}
if pos > scratch.cap {
return fail(
cache,
format!(
"draft plane rows {pos} exceed scratch capacity {}",
scratch.cap
),
);
}
let kb = pos * draft_k_tok_bytes;
let vb = pos * draft_v_tok_bytes;
if draft_k.len() < kb || draft_v.len() < vb {
return fail(
cache,
format!(
"truncated draft plane: K {} < {kb} or V {} < {vb} bytes",
draft_k.len(),
draft_v.len(),
),
);
}
if kb > 0 {
if let Err(err) = e.copy_u8_into(&mut scratch.kv.k, 0, draft_k, kb) {
return fail(cache, format!("draft K restore copy failed: {err}"));
}
}
if vb > 0 {
if let Err(err) = e.copy_u8_into(&mut scratch.kv.v, 0, draft_v, vb) {
return fail(cache, format!("draft V restore copy failed: {err}"));
}
}
if let Err(err) = scratch.set_len(e, pos) {
return fail(cache, format!("draft scratch len set failed: {err}"));
}
let mut last_h_dev = if last_h.len() == self.cfg.n_embd as usize {
e.htod(last_h).ok()
} else {
None
};
if require_anchor && last_h_dev.is_none() {
return fail(
cache,
"empty-suffix continuation requires the entry's boundary hidden anchor".into(),
);
}
let mut committed = prefix;
let next_pred;
let mut sctr = 0u32;
let sampled = sampling.is_some_and(|s| s.temp > 0.0) && spec_sampled_boundary_on();
let mut boundary_captures: Vec<SpecBoundaryCapture> = Vec::new();
let mut restored_turn_ckpt: Option<SpecCheckpoint> = None;
if !suffix.is_empty() {
let dirty =
|msg: String| -> Result<SpecSession, (Option<Cache>, String)> { Err((None, msg)) };
let n_embd = self.cfg.n_embd as usize;
let t = suffix.len();
let mut h_rows = match e.uninit(t * n_embd) {
Ok(b) => b,
Err(err) => return fail(cache, format!("suffix hidden buffer alloc: {err}")),
};
let b_rel = republish_at
.and_then(|abs| abs.checked_sub(pos))
.filter(|&r| r > 0 && r < t);
let mut feed_logits = Vec::new();
let tokenwise_env = std::env::var("MEMRA_PRIME_TOKENWISE").is_ok()
|| e.frozen_cpu_experts_prefer_tokenwise_prime();
let mut fed = 0usize;
for seg_end in [b_rel, Some(t)].into_iter().flatten() {
if seg_end <= fed {
continue;
}
let seg = &suffix[fed..seg_end];
let batched = seg.len() >= crate::hybrid_forward::PRIME_MIN_T && !tokenwise_env;
if batched {
match self.prime_cache(e, seg, &mut cache, t - seg_end) {
Ok((l, _h_seed, hiddens)) => {
if let Err(err) =
e.copy_into(&mut h_rows, fed * n_embd, &hiddens, seg.len() * n_embd)
{
return dirty(format!("suffix hidden copy: {err}"));
}
feed_logits = l;
}
Err(err) => return dirty(format!("suffix prime failed: {err}")),
}
} else {
for (i, &tok) in seg.iter().enumerate() {
match self.decode_step_h(e, tok, &mut cache) {
Ok((l, h)) => {
if let Err(err) =
e.copy_into(&mut h_rows, (fed + i) * n_embd, &h, n_embd)
{
return dirty(format!("suffix hidden copy: {err}"));
}
feed_logits = l;
}
Err(err) => return dirty(format!("suffix decode_step failed: {err}")),
}
}
}
fed = seg_end;
if Some(seg_end) == b_rel {
debug_assert_eq!(
cache.pos,
pos + seg_end,
"stable-boundary capture off the feed split"
);
if spec_restore_republish_on() {
if let Ok(snap) = cache.snapshot(e) {
boundary_captures.push(SpecBoundaryCapture {
snap,
pos: pos + seg_end,
logits: feed_logits.clone(),
last_h: capture_boundary_hidden(e, &h_rows, seg_end, n_embd),
});
}
}
let anchor: Result<CudaSlice<f32>, Box<dyn std::error::Error>> =
e.uninit(n_embd).and_then(|mut a| {
e.copy_view_into(
&mut a,
0,
&h_rows.slice((seg_end - 1) * n_embd..seg_end * n_embd),
n_embd,
)?;
Ok(a)
});
if let (Ok(snap), Ok(last_h)) = (cache.snapshot(e), anchor) {
restored_turn_ckpt = Some(SpecCheckpoint {
snap,
pos: pos + seg_end,
last_h,
});
}
}
}
let mtp = self.mtp.as_ref().expect("mtp checked above");
let (embd_qt, embd_rb) = self.embd.qt_and_row_bytes(n_embd);
let embd_gpu = if spec_host_embd() {
None
} else {
Some(
self.embd_gpu
.get_or_init(|| e.upload_u8(&self.embd.raw).expect("embed table upload")),
)
};
let embd_dev = embd_gpu.map(|g| (g, embd_qt, embd_rb));
let fill_chunk = 4096usize;
let mut filled = true;
let mut start = 0usize;
'fill: while start < t {
let end = (start + fill_chunk).min(t);
let tc = end - start;
let Ok(mut phs) = e.zeros(tc * n_embd) else {
filled = false;
break 'fill;
};
let (src_lo, dst_off, n_copy) = if start == 0 {
(0, n_embd, (tc - 1) * n_embd)
} else {
((start - 1) * n_embd, 0, tc * n_embd)
};
if start == 0 {
if let Some(lh) = last_h_dev.as_ref() {
if e.copy_into(&mut phs, 0, lh, n_embd).is_err() {
filled = false;
break 'fill;
}
}
}
if n_copy > 0
&& e.copy_view_into(
&mut phs,
dst_off,
&h_rows.slice(src_lo..src_lo + n_copy),
n_copy,
)
.is_err()
{
filled = false;
break 'fill;
}
if self
.mtp_kv_fill(
e,
mtp,
&suffix[start..end],
&phs,
pos + start,
&mut scratch,
embd_dev,
)
.is_err()
{
filled = false;
break 'fill;
}
start = end;
}
if !filled {
if let Err(err) = scratch.set_len(e, pos) {
return dirty(format!("scratch truncation after failed fill: {err}"));
}
}
if spec_restore_republish_on() && boundary_captures.is_empty() {
debug_assert_eq!(
cache.pos,
pos + t,
"extended-entry capture must sit at the restored session's prompt end",
);
if let Ok(snap) = cache.snapshot(e) {
boundary_captures.push(SpecBoundaryCapture {
snap,
pos: pos + t,
logits: feed_logits.clone(),
last_h: capture_boundary_hidden(e, &h_rows, t, n_embd),
});
}
}
next_pred = Some(if sampled {
let sp = sampling.expect("sampled implies a sampler");
let hist = pen_window_seed(&committed, suffix, sp.penalty_last_n);
match sample_boundary_token(
e,
&feed_logits,
&sp,
&hist,
&mut sctr,
"restore-suffix-feed",
) {
Ok(t) => t,
Err(err) => {
return dirty(format!("boundary token draw failed: {err}"));
}
}
} else {
argmax(&feed_logits) as u32
});
let mut lh = match e.uninit(n_embd) {
Ok(b) => b,
Err(err) => return dirty(format!("boundary hidden alloc: {err}")),
};
if let Err(err) = e.copy_view_into(
&mut lh,
0,
&h_rows.slice((t - 1) * n_embd..t * n_embd),
n_embd,
) {
return dirty(format!("boundary hidden copy: {err}"));
}
last_h_dev = Some(lh);
committed.extend_from_slice(suffix);
} else {
if boundary_logits.is_empty() {
return fail(
cache,
"full-cover restore without the entry's boundary logits".into(),
);
}
next_pred = Some(if sampled {
let sp = sampling.expect("sampled implies a sampler");
let hist = pen_window_seed(&committed, &[], sp.penalty_last_n);
match sample_boundary_token(
e,
boundary_logits,
&sp,
&hist,
&mut sctr,
"restore-full-cover",
) {
Ok(t) => t,
Err(err) => {
return fail(cache, format!("boundary token draw failed: {err}"));
}
}
} else {
argmax(boundary_logits) as u32
});
}
Ok(SpecSession {
cache,
scratch,
committed,
last_h: last_h_dev,
next_pred,
sctr,
uctr: 0,
draft_ctx: None,
pending_tok: None,
turn_ckpt: restored_turn_ckpt,
telem: SpecTelemetryCounters::default(),
capture_at: None,
boundary_captures,
ckpt_at: None,
})
}
pub fn optipipe_compare_session_state(
&self,
e: &Engine,
reference: &SpecSession,
candidate: &SpecSession,
) -> Result<OptiForkStateIdentity, Box<dyn std::error::Error>> {
fn fail(what: &str) -> Box<dyn std::error::Error> {
format!("optipipe state mismatch: {what}").into()
}
fn same_f32(a: &[f32], b: &[f32]) -> bool {
a.len() == b.len() && a.iter().zip(b).all(|(x, y)| x.to_bits() == y.to_bits())
}
fn compare_layers(
es: &Engine,
range: std::ops::Range<usize>,
reference: &SpecSession,
candidate: &SpecSession,
report: &mut OptiForkStateIdentity,
) -> Result<(), Box<dyn std::error::Error>> {
for il in range {
match (&reference.cache.kv[il], &candidate.cache.kv[il]) {
(Some(a), Some(b)) => {
if a.len != b.len {
return Err(fail(&format!(
"layer {il} host KV len {} != {}",
a.len, b.len
)));
}
let ad = es.dtoh_i32(&a.len_d)?;
let bd = es.dtoh_i32(&b.len_d)?;
if ad != bd || ad.first().copied() != Some(a.len as i32) {
return Err(fail(&format!(
"layer {il} device KV len {ad:?} != {bd:?} (host={})",
a.len,
)));
}
let kb = a.len * a.k_tok_bytes;
let vb = a.len * a.v_tok_bytes;
if kb > 0 {
let ak = es.dtoh_u8_view(&a.k.slice(0..kb))?;
let bk = es.dtoh_u8_view(&b.k.slice(0..kb))?;
if ak != bk {
let at = ak.iter().zip(&bk).position(|(x, y)| x != y).unwrap();
return Err(fail(&format!(
"layer {il} K bytes at byte {at} row {} offset {}: {} != {}",
at / a.k_tok_bytes,
at % a.k_tok_bytes,
ak[at],
bk[at],
)));
}
}
if vb > 0 {
let av = es.dtoh_u8_view(&a.v.slice(0..vb))?;
let bv = es.dtoh_u8_view(&b.v.slice(0..vb))?;
if av != bv {
let at = av.iter().zip(&bv).position(|(x, y)| x != y).unwrap();
return Err(fail(&format!(
"layer {il} V bytes at byte {at} row {} offset {}: {} != {}",
at / a.v_tok_bytes,
at % a.v_tok_bytes,
av[at],
bv[at],
)));
}
}
report.trunk_kv_bytes += kb + vb;
}
(None, None) => {}
_ => return Err(fail(&format!("layer {il} KV presence"))),
}
match (&reference.cache.recur[il], &candidate.cache.recur[il]) {
(Some(a), Some(b)) => {
let ac = es.dtoh(&a.conv_state)?;
let bc = es.dtoh(&b.conv_state)?;
if !same_f32(&ac, &bc) {
return Err(fail(&format!("layer {il} conv state")));
}
let as_ = es.dtoh(&a.ssm_state)?;
let bs = es.dtoh(&b.ssm_state)?;
if !same_f32(&as_, &bs) {
return Err(fail(&format!("layer {il} SSM state")));
}
report.recurrent_bytes += (ac.len() + as_.len()) * 4;
}
(None, None) => {}
_ => return Err(fail(&format!("layer {il} recurrent presence"))),
}
}
Ok(())
}
if reference.committed != candidate.committed {
return Err(fail("committed token ids"));
}
if reference.cache.pos != candidate.cache.pos
|| reference.cache.max_ctx != candidate.cache.max_ctx
{
return Err(fail("cache pos/capacity"));
}
if reference.pending_tok != candidate.pending_tok
|| reference.next_pred != candidate.next_pred
|| reference.sctr != candidate.sctr
|| reference.uctr != candidate.uctr
{
return Err(fail("pending/prediction/counter tail"));
}
let mut report = OptiForkStateIdentity::default();
if let Some(fence) = crate::pp::pp_cuts(self.layers.len()) {
let rt = crate::pp::PpNRt::get(e)?;
for stage in 0..rt.n_stages() {
let _scope = rt.enter(stage);
compare_layers(
rt.engine(stage, e),
fence[stage]..fence[stage + 1],
reference,
candidate,
&mut report,
)?;
}
} else {
compare_layers(e, 0..self.layers.len(), reference, candidate, &mut report)?;
}
let (a, b) = (&reference.scratch.kv, &candidate.scratch.kv);
if a.len != b.len || e.dtoh_i32(&a.len_d)? != e.dtoh_i32(&b.len_d)? {
return Err(fail("draft scratch length"));
}
let kb = a.len * a.k_tok_bytes;
let vb = a.len * a.v_tok_bytes;
if kb > 0 && e.dtoh_u8_view(&a.k.slice(0..kb))? != e.dtoh_u8_view(&b.k.slice(0..kb))? {
return Err(fail("draft scratch K bytes"));
}
if vb > 0 && e.dtoh_u8_view(&a.v.slice(0..vb))? != e.dtoh_u8_view(&b.v.slice(0..vb))? {
return Err(fail("draft scratch V bytes"));
}
report.scratch_kv_bytes = kb + vb;
match (&reference.last_h, &candidate.last_h) {
(Some(a), Some(b)) => {
let ah = e.dtoh(a)?;
let bh = e.dtoh(b)?;
if !same_f32(&ah, &bh) {
return Err(fail("last hidden/seed bytes"));
}
report.hidden_bytes = ah.len() * 4;
}
(None, None) => {}
_ => return Err(fail("last hidden/seed presence")),
}
Ok(report)
}
pub fn spec_rewind_to_checkpoint(
&self,
e: &Engine,
sess: &mut SpecSession,
) -> Result<Option<usize>, Box<dyn std::error::Error>> {
if sess.turn_ckpt.as_ref().is_some_and(|ckpt| {
!sess.cache.can_rollback(&ckpt.snap, 0) || !sess.scratch.can_rewind_to(ckpt.pos)
}) {
return Err(
"SWA ring rewind checkpoint has been lapped; full re-prime required".into(),
);
}
let Some(ckpt) = sess.turn_ckpt.take() else {
return Ok(None);
};
assert!(
ckpt.pos <= sess.committed.len(),
"checkpoint past committed ({} > {})",
ckpt.pos,
sess.committed.len()
);
crate::pp::restore_cache_checkpoint(e, &self.cfg, None, &mut sess.cache, &ckpt.snap)?;
debug_assert_eq!(
sess.cache.pos, ckpt.pos,
"rollback landed off the checkpoint"
);
sess.scratch.set_len(e, ckpt.pos)?;
sess.committed.truncate(ckpt.pos);
sess.last_h = Some(ckpt.last_h);
sess.next_pred = None;
sess.pending_tok = None;
Ok(Some(ckpt.pos))
}
pub fn spec_grow_and_rewind_to_checkpoint(
&self,
e: &Engine,
sess: &mut SpecSession,
target_cap: usize,
) -> Result<Option<usize>, Box<dyn std::error::Error>> {
if target_cap <= sess.cache.max_ctx {
return self.spec_rewind_to_checkpoint(e, sess);
}
let Some(ckpt) = sess.turn_ckpt.as_ref() else {
return Ok(None);
};
if ckpt.pos == 0 || ckpt.pos > sess.committed.len() {
return Err(format!(
"checkpoint pos {} outside committed length {}",
ckpt.pos,
sess.committed.len(),
)
.into());
}
if ckpt.pos > target_cap {
return Err(format!(
"checkpoint pos {} exceeds grown capacity {target_cap}",
ckpt.pos,
)
.into());
}
let mut grown_cache = crate::pp::new_cache(e, &self.cfg, target_cap)?;
let mut grown_scratch = MtpScratch::new(
e,
&self.cfg,
target_cap,
self.mtp.as_ref().and_then(|m| m.geom.as_ref()),
)?;
crate::pp::restore_cache_checkpoint(
e,
&self.cfg,
Some(&sess.cache),
&mut grown_cache,
&ckpt.snap,
)?;
let src = &sess.scratch.kv;
let dst = &mut grown_scratch.kv;
if ckpt.pos > src.len
|| src.kv_dim_k != dst.kv_dim_k
|| src.kv_dim_v != dst.kv_dim_v
|| src.k_tok_bytes != dst.k_tok_bytes
|| src.v_tok_bytes != dst.v_tok_bytes
{
return Err(format!(
"checkpoint draft layout mismatch (pos {}, source len {})",
ckpt.pos, src.len,
)
.into());
}
let kb = ckpt.pos * src.k_tok_bytes;
let vb = ckpt.pos * src.v_tok_bytes;
if kb > 0 {
e.copy_u8_into(&mut dst.k, 0, &src.k, kb)?;
}
if vb > 0 {
e.copy_u8_into(&mut dst.v, 0, &src.v, vb)?;
}
grown_scratch.set_len(e, ckpt.pos)?;
e.stream().synchronize()?;
let ckpt = sess
.turn_ckpt
.take()
.expect("checkpoint remained present through transactional grow");
let pos = ckpt.pos;
sess.cache = grown_cache;
sess.scratch = grown_scratch;
sess.committed.truncate(pos);
sess.last_h = Some(ckpt.last_h);
sess.next_pred = None;
sess.pending_tok = None;
sess.draft_ctx = None;
debug_assert_eq!(sess.cache.pos, pos, "grown rewind landed off checkpoint");
debug_assert_eq!(
sess.scratch.kv.len, pos,
"grown draft rewind landed off checkpoint"
);
Ok(Some(pos))
}
pub fn spec_flush_pending(
&self,
e: &Engine,
sess: &mut SpecSession,
sampling: Option<SpecSampling>,
) -> Result<(), Box<dyn std::error::Error>> {
let Some(b) = sess.pending_tok.take() else {
return Ok(());
};
let mtp = self
.mtp
.as_ref()
.expect("pending carry requires an MTP head");
let n_embd = self.cfg.n_embd as usize;
let (embd_qt, embd_rb) = self.embd.qt_and_row_bytes(n_embd);
let embd_gpu = if spec_host_embd() {
None
} else {
Some(
self.embd_gpu
.get_or_init(|| e.upload_u8(&self.embd.raw).expect("embed table upload")),
)
};
let embd_dev = embd_gpu.map(|g| (g, embd_qt, embd_rb));
let pos_b = sess.cache.pos;
sess.scratch.set_len(e, pos_b)?;
let (lg_b, hb) = self.spec_target_step_h(e, b, &mut sess.cache)?;
sess.next_pred = Some(match sampling {
Some(sp) if sp.temp > 0.0 && spec_sampled_boundary_on() => {
let hist = pen_window_seed(&sess.committed, &[b], sp.penalty_last_n);
sample_boundary_token(e, &lg_b, &sp, &hist, &mut sess.sctr, "flush-pending")?
}
_ => argmax(&lg_b) as u32,
});
let anchor = sess
.last_h
.as_ref()
.expect("pending carry requires last_h (the predecessor-row anchor)");
self.mtp_kv_fill(e, mtp, &[b], anchor, pos_b, &mut sess.scratch, embd_dev)?;
sess.last_h = Some(hb);
sess.committed.push(b);
Ok(())
}
fn spec_target_step_h(
&self,
e: &Engine,
token: u32,
cache: &mut Cache,
) -> Result<(Vec<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
if self.cfg.step35.is_none() && !self.qwen35_serving_class() {
return self.decode_step_h(e, token, cache);
}
let pos0 = cache.pos;
let (logits, hidden) = self.decode_step_t_core(e, &[token], pos0, cache, None, None)?;
Ok((e.dtoh(&logits)?, hidden))
}
fn mtp_graph_capturable(&self) -> bool {
self.mtp
.as_ref()
.map(|m| match &m.ffn {
crate::hybrid::Ffn::Dense { .. } => true,
crate::hybrid::Ffn::Moe(mo) => mo.dev_exps.is_some(),
})
.unwrap_or(false)
}
fn qwen35_serving_class(&self) -> bool {
matches!(
self.cfg.arch,
memra_gguf::config::Arch::Qwen35 | memra_gguf::config::Arch::Qwen35Moe
)
}
pub fn spec_pipe_available(&self, e: &Engine) -> bool {
if std::env::var("MEMRA_SPEC_PIPE").as_deref() != Ok("1")
|| !spec_devacc()
|| spec_replay_env_enabled()
|| spec_stream()
|| std::env::var("MEMRA_SPEC_ADAPT").as_deref() == Ok("1")
|| std::env::var("MEMRA_SPEC_PMIN0").as_deref() == Ok("1")
|| std::env::var("MEMRA_SPEC_PP_ANATOMY").as_deref() == Ok("1")
|| std::env::var("MEMRA_SPEC_PMIN")
.ok()
.and_then(|v| v.parse::<f32>().ok())
.unwrap_or(0.0)
> 0.0
|| self.is_gemma4_e4b()
|| self.cfg.gemma4.is_some()
|| self.mtp.is_none()
{
return false;
}
let Some(cuts) = crate::pp::pp_cuts(self.layers.len()) else {
return false;
};
if cuts.len() != 3 || crate::pp::pp2_streams_off() || !crate::pp::spec_pp_on() {
return false;
}
crate::pp::PpNRt::get(e)
.map(|rt| rt.n_stages() == 2 && rt.cross_device())
.unwrap_or(false)
}
#[allow(clippy::too_many_arguments)]
pub fn generate_spec_session_pair(
&self,
e: &Engine,
sess_a: &mut SpecSession,
max_new_a: usize,
k_a: usize,
sess_b: &mut SpecSession,
max_new_b: usize,
k_b: usize,
) -> Result<((Vec<u32>, usize, usize), (Vec<u32>, usize, usize)), Box<dyn std::error::Error>>
{
if !self.spec_pipe_available(e) {
return Err("two-session speculative pipeline is outside its reduced matrix".into());
}
if max_new_a == 0 || max_new_b == 0 || k_a == 0 || k_b == 0 {
return Err(
"two-session speculative pipeline requires non-empty positive-K bursts".into(),
);
}
for sess in [&*sess_a, &*sess_b] {
if sess.committed.is_empty()
|| sess.last_h.is_none()
|| (sess.next_pred.is_none() && sess.pending_tok.is_none())
{
return Err("two-session speculative pipeline requires warm continuations".into());
}
}
let graph_ok = std::env::var("MEMRA_SPEC_NOGRAPH").is_err()
&& !spec_host_embd()
&& self.mtp_graph_capturable()
&& !crate::model::full_prec_enabled();
let graph_a = graph_ok && k_a + 2 < 96;
let graph_b = graph_ok && k_b + 2 < 96;
let was_tracking = e.ctx().is_event_tracking();
if (graph_a || graph_b) && was_tracking {
unsafe {
e.ctx().disable_event_tracking();
}
}
static LOGGED: std::sync::Once = std::sync::Once::new();
LOGGED.call_once(|| {
eprintln!("[spec-pipe] two-session PP-2 continuation pipeline engaged");
});
let sync = std::sync::Arc::new(SpecPipeSync::new());
let lane_a = SpecPipeLane {
sync: sync.clone(),
lane: 0,
};
let lane_b = SpecPipeLane { sync, lane: 1 };
let mut sess_b_ptr = SpecPipeSessionPtr(sess_b as *mut SpecSession);
let (result_a, result_b) = std::thread::scope(|scope| {
let b = scope.spawn(move || {
let mut finish = SpecPipeFinish::new(&lane_b);
let sess_b = unsafe { sess_b_ptr.get_mut() };
let result = e
.ctx()
.bind_to_thread()
.map_err(|err| err.to_string())
.and_then(|_| {
self.generate_spec_inner2(
e,
&[],
max_new_b,
k_b,
graph_b,
Some(sess_b),
None,
None,
None,
None,
Some(&lane_b),
)
.map_err(|err| err.to_string())
});
finish.close(result.is_err());
result
});
let mut finish = SpecPipeFinish::new(&lane_a);
let result_a = self.generate_spec_inner2(
e,
&[],
max_new_a,
k_a,
graph_a,
Some(sess_a),
None,
None,
None,
None,
Some(&lane_a),
);
finish.close(result_a.is_err());
let result_b = b
.join()
.map_err(|_| "paired speculative session B panicked".to_string())
.and_then(|r| r);
(result_a, result_b)
});
if (graph_a || graph_b) && was_tracking {
unsafe {
e.ctx().enable_event_tracking();
}
}
let result_a = result_a?;
let result_b = result_b.map_err(|err| -> Box<dyn std::error::Error> { err.into() })?;
Ok((result_a, result_b))
}
pub fn generate_spec_session(
&self,
e: &Engine,
sess: &mut SpecSession,
suffix: &[u32],
max_new: usize,
k: usize,
) -> Result<(Vec<u32>, usize, usize), Box<dyn std::error::Error>> {
self.generate_spec_session_sampled(e, sess, suffix, max_new, k, None, None)
}
#[allow(clippy::too_many_arguments)]
pub fn generate_spec_session_sampled(
&self,
e: &Engine,
sess: &mut SpecSession,
suffix: &[u32],
max_new: usize,
k: usize,
sampling: Option<SpecSampling>,
on_commit: Option<&mut dyn FnMut(&[u32]) -> bool>,
) -> Result<(Vec<u32>, usize, usize), Box<dyn std::error::Error>> {
self.generate_spec_session_sampled_prime_split(
e, sess, suffix, max_new, k, sampling, None, on_commit,
)
}
#[allow(clippy::too_many_arguments)]
pub fn generate_spec_session_sampled_prime_split(
&self,
e: &Engine,
sess: &mut SpecSession,
suffix: &[u32],
max_new: usize,
k: usize,
sampling: Option<SpecSampling>,
prime_split: Option<usize>,
on_commit: Option<&mut dyn FnMut(&[u32]) -> bool>,
) -> Result<(Vec<u32>, usize, usize), Box<dyn std::error::Error>> {
self.generate_spec_session_constrained_prime_split(
e,
sess,
suffix,
max_new,
k,
sampling,
None,
prime_split,
on_commit,
)
}
#[allow(clippy::too_many_arguments)]
pub fn generate_spec_session_constrained(
&self,
e: &Engine,
sess: &mut SpecSession,
suffix: &[u32],
max_new: usize,
k: usize,
sampling: Option<SpecSampling>,
constraint: Option<&mut dyn SpecConstraint>,
on_commit: Option<&mut dyn FnMut(&[u32]) -> bool>,
) -> Result<(Vec<u32>, usize, usize), Box<dyn std::error::Error>> {
self.generate_spec_session_constrained_prime_split(
e, sess, suffix, max_new, k, sampling, constraint, None, on_commit,
)
}
#[allow(clippy::too_many_arguments)]
pub fn generate_spec_session_constrained_prime_split(
&self,
e: &Engine,
sess: &mut SpecSession,
suffix: &[u32],
max_new: usize,
k: usize,
sampling: Option<SpecSampling>,
constraint: Option<&mut dyn SpecConstraint>,
prime_split: Option<usize>,
on_commit: Option<&mut dyn FnMut(&[u32]) -> bool>,
) -> Result<(Vec<u32>, usize, usize), Box<dyn std::error::Error>> {
if constraint.is_some() && sampling.is_some_and(|s| s.temp > 0.0) {
return Err(
"constrained spec decode is greedy-only (worker routes sampled \
constrained to plain decode)"
.into(),
);
}
if sess.pending_tok.is_some()
&& (!suffix.is_empty() || sampling.map_or(false, |s| s.temp > 0.0))
{
self.spec_flush_pending(e, sess, sampling)?;
}
let graph_draft = std::env::var("MEMRA_SPEC_NOGRAPH").is_err()
&& !spec_host_embd()
&& self.mtp_graph_capturable()
&& k + 2 < 96
&& !crate::model::full_prec_enabled();
let was_tracking = e.ctx().is_event_tracking();
if graph_draft && was_tracking {
unsafe {
e.ctx().disable_event_tracking();
}
}
let r = self.generate_spec_inner2(
e,
suffix,
max_new,
k,
graph_draft,
Some(sess),
sampling,
constraint,
on_commit,
prime_split,
None,
);
if graph_draft && was_tracking {
unsafe {
e.ctx().enable_event_tracking();
}
}
let (out, d, a) = r?;
Ok((out, d, a))
}
pub fn generate_spec(
&self,
e: &Engine,
prompt: &[u32],
max_new: usize,
k: usize,
) -> Result<(Vec<u32>, usize, usize), Box<dyn std::error::Error>> {
let graph_draft = std::env::var("MEMRA_SPEC_NOGRAPH").is_err()
&& !spec_host_embd()
&& self.mtp_graph_capturable()
&& k + 2 < 96
&& !crate::model::full_prec_enabled();
if !graph_draft {
return self.generate_spec_inner2(
e, prompt, max_new, k, false, None, None, None, None, None, None,
);
}
let was_tracking = e.ctx().is_event_tracking();
if was_tracking {
unsafe {
e.ctx().disable_event_tracking();
}
}
let r = self.generate_spec_inner2(
e, prompt, max_new, k, true, None, None, None, None, None, None,
);
if was_tracking {
unsafe {
e.ctx().enable_event_tracking();
}
}
r
}
fn generate_spec_inner2(
&self,
e: &Engine,
prompt: &[u32],
max_new: usize,
k: usize,
graph_draft: bool,
mut sess: Option<&mut SpecSession>,
sampling: Option<SpecSampling>,
mut constraint: Option<&mut dyn SpecConstraint>,
mut on_commit: Option<&mut dyn FnMut(&[u32]) -> bool>,
prime_split: Option<usize>,
pipe: Option<&SpecPipeLane>,
) -> Result<(Vec<u32>, usize, usize), Box<dyn std::error::Error>> {
assert!(k >= 1, "k must be >= 1");
if let Some(p) = pipe {
p.setup_begin()?;
}
let mut flushed = 0usize;
let mut keep_going;
let mtp = self
.mtp
.as_ref()
.expect("generate_spec requires an MTP head (nextn_predict_layers>0)");
let n_vocab = self.output.out_features();
let d_vocab = mtp
.shared_head_head
.as_ref()
.unwrap_or(&self.output)
.out_features();
let n_embd = self.cfg.n_embd as usize;
let session_mode = sess.is_some();
let max_ctx = match sess.as_ref() {
Some(s) => s.cache.max_ctx,
None => prompt.len() + max_new + k + 8,
};
let mut own_cache;
let mut own_scratch;
let mut sess_capture: Option<(Option<usize>, &mut Vec<SpecBoundaryCapture>)> = None;
let mut ckpt_req: Option<usize> = None;
let (
cache,
scratch,
mut sess_tail,
mut sess_draft_slot,
mut sess_pending_slot,
sess_ckpt_slot,
sess_telem,
): (
&mut Cache,
&mut MtpScratch,
Option<(
&mut Vec<u32>,
&mut Option<CudaSlice<f32>>,
&mut Option<u32>,
&mut u32,
&mut u32,
)>,
Option<&mut Option<DraftGraphCtx>>,
Option<&mut Option<u32>>,
Option<&mut Option<SpecCheckpoint>>,
Option<&SpecTelemetryCounters>,
) = match sess.take() {
Some(sr) => {
let SpecSession {
cache,
scratch,
committed,
last_h,
next_pred,
sctr: s_sctr,
uctr: s_uctr,
draft_ctx,
pending_tok,
turn_ckpt,
telem,
capture_at,
boundary_captures,
ckpt_at,
} = sr;
sess_capture = Some((capture_at.take(), boundary_captures));
ckpt_req = ckpt_at.take();
(
cache,
scratch,
Some((committed, last_h, next_pred, s_sctr, s_uctr)),
Some(draft_ctx),
Some(pending_tok),
Some(turn_ckpt),
Some(telem),
)
}
None => {
own_cache = crate::pp::new_cache(e, &self.cfg, max_ctx)?;
own_scratch = MtpScratch::new(
e,
&self.cfg,
max_ctx,
self.mtp.as_ref().and_then(|m| m.geom.as_ref()),
)?;
(
&mut own_cache,
&mut own_scratch,
None,
None,
None,
None,
None,
)
}
};
let base = cache.pos;
let carried_pending: Option<u32> = sess_pending_slot.as_mut().and_then(|s| s.take());
let spec_replay = spec_replay_env_enabled();
if constraint.is_some() && spec_replay {
return Err(
"constrained spec decode does not support MEMRA_SPEC_REPLAY=1 \
(legacy replay commits an unmasked bonus)"
.into(),
);
}
let refresh = std::env::var("MEMRA_SPEC_NOREFRESH").is_err();
let continuation = prompt.is_empty();
if continuation {
assert!(session_mode, "empty prompt requires a session");
assert!(
sess_tail
.as_ref()
.map_or(false, |(c, lh, np, _, _)| !c.is_empty()
&& lh.is_some()
&& (np.is_some() || carried_pending.is_some())),
"empty-suffix continuation needs a primed session (committed + last_h + next_pred|pending)"
);
}
let mut prime_logits;
let mut prompt_h: Option<CudaSlice<f32>> = None;
let t_prime = std::time::Instant::now();
let batched_prime = !continuation
&& prompt.len() >= crate::hybrid_forward::PRIME_MIN_T
&& std::env::var("MEMRA_PRIME_TOKENWISE").is_err()
&& !e.frozen_cpu_experts_prefer_tokenwise_prime();
let prime_split = prime_split.filter(|&split| split > 0 && split < prompt.len());
if prime_split.is_some() && continuation {
return Err("spec prime split requires a non-empty prime".into());
}
let ckpt_rel = if continuation {
None
} else {
ckpt_req
.and_then(|abs| abs.checked_sub(base))
.filter(|&r| r > 0 && r < prompt.len())
};
let mut stops: Vec<usize> = Vec::new();
for b in [prime_split, ckpt_rel].into_iter().flatten() {
if !stops.contains(&b) {
stops.push(b);
}
}
stops.sort_unstable();
let mut ckpt_early: Option<Option<SpecCheckpoint>> = None;
if continuation {
prime_logits = Vec::new();
} else if !stops.is_empty() {
if let Some(&first) = stops.first() {
if prime_split == Some(first) && first < crate::hybrid_forward::PRIME_MIN_T {
return Err(format!(
"spec prime split {first} is below PRIME_MIN_T {}",
crate::hybrid_forward::PRIME_MIN_T,
)
.into());
}
}
let mut h_all = e.uninit(prompt.len() * n_embd)?;
prime_logits = Vec::new();
let mut prev = 0usize;
for seg_end in stops.iter().copied().chain(std::iter::once(prompt.len())) {
if seg_end <= prev {
continue;
}
let seg = &prompt[prev..seg_end];
let is_final = seg_end == prompt.len();
let batched_seg = seg.len() >= crate::hybrid_forward::PRIME_MIN_T
&& (!is_final
|| (std::env::var("MEMRA_PRIME_TOKENWISE").is_err()
&& !e.frozen_cpu_experts_prefer_tokenwise_prime()));
if batched_seg {
let (l, _, h_seg) =
self.prime_cache(e, seg, &mut *cache, prompt.len() - seg_end)?;
e.copy_into(&mut h_all, prev * n_embd, &h_seg, seg.len() * n_embd)?;
prime_logits = l;
} else {
for (i, &tok) in seg.iter().enumerate() {
let (l, h) = self.decode_step_h(e, tok, &mut *cache)?;
e.copy_into(&mut h_all, (prev + i) * n_embd, &h, n_embd)?;
prime_logits = l;
}
}
prev = seg_end;
if is_final {
break;
}
debug_assert_eq!(cache.pos, base + seg_end, "prime stop landed off boundary");
if base == 0 {
if let Some((requested, slot)) = sess_capture.as_mut() {
if *requested == Some(seg_end) || ckpt_rel == Some(seg_end) {
if let Ok(snap) = cache.snapshot(e) {
slot.push(SpecBoundaryCapture {
snap,
pos: seg_end,
logits: prime_logits.clone(),
last_h: capture_boundary_hidden(e, &h_all, seg_end, n_embd),
});
}
}
}
}
if ckpt_rel == Some(seg_end) {
let anchor: Result<CudaSlice<f32>, Box<dyn std::error::Error>> =
e.uninit(n_embd).and_then(|mut a| {
e.copy_view_into(
&mut a,
0,
&h_all.slice((seg_end - 1) * n_embd..seg_end * n_embd),
n_embd,
)?;
Ok(a)
});
ckpt_early = Some(match (cache.snapshot(e), anchor) {
(Ok(snap), Ok(last_h)) => Some(SpecCheckpoint {
snap,
pos: base + seg_end,
last_h,
}),
_ => None,
});
}
}
if std::env::var("MEMRA_SPEC_STATS").as_deref() == Ok("1") {
eprintln!(
"[spec-prime] stops={stops:?} tail={}",
prompt.len() - stops.last().copied().unwrap_or(0)
);
}
prompt_h = Some(h_all);
} else if batched_prime {
let (l, _h_seed, hiddens) = self.prime_cache(e, prompt, &mut *cache, 0)?;
prime_logits = l;
prompt_h = Some(hiddens);
} else {
prime_logits = Vec::new();
prompt_h = Some(e.uninit(prompt.len() * n_embd)?);
for (i, &tok) in prompt.iter().enumerate() {
let (l, h) = self.spec_target_step_h(e, tok, &mut *cache)?;
if let Some(ph) = prompt_h.as_mut() {
e.copy_into(ph, i * n_embd, &h, n_embd)?;
}
prime_logits = l;
}
}
e.stream().synchronize()?;
if !continuation && base == 0 {
if let Some((requested, slot)) = sess_capture.as_mut() {
if *requested == Some(prompt.len()) && slot.is_empty() {
debug_assert_eq!(cache.pos, prompt.len(), "seed capture off prompt end");
if let Ok(snap) = cache.snapshot(e) {
slot.push(SpecBoundaryCapture {
snap,
pos: prompt.len(),
logits: prime_logits.clone(),
last_h: prompt_h
.as_ref()
.map(|ph| capture_boundary_hidden(e, ph, prompt.len(), n_embd))
.unwrap_or_default(),
});
}
}
}
}
crate::PRIME_NANOS.store(
t_prime.elapsed().as_nanos() as u64,
std::sync::atomic::Ordering::Relaxed,
);
let (embd_qt, embd_rb) = self.embd.qt_and_row_bytes(n_embd);
let host_embd = spec_host_embd();
let embd_gpu = if host_embd {
None
} else {
Some(
self.embd_gpu
.get_or_init(|| e.upload_u8(&self.embd.raw).expect("embed table upload")),
)
};
let embd_dev = embd_gpu.map(|g| (g, embd_qt, embd_rb));
if host_embd {
eprintln!(
"[spec] host-row embedding: {} bytes kept off HBM",
self.embd.raw.len()
);
}
let mut out: Vec<u32> = Vec::with_capacity(max_new);
let mut total_drafted = 0usize;
let mut total_accepted = 0usize;
let sp = sampling.unwrap_or_else(|| SpecSampling {
temp: std::env::var("MEMRA_SPEC_TEMP")
.ok()
.and_then(|v| v.parse().ok())
.unwrap_or(0.0),
seed: std::env::var("MEMRA_SEED")
.ok()
.and_then(|v| v.parse().ok())
.unwrap_or(42),
top_k: std::env::var("MEMRA_TOP_K")
.ok()
.and_then(|v| v.parse().ok())
.unwrap_or(0),
top_p: std::env::var("MEMRA_TOP_P")
.ok()
.and_then(|v| v.parse().ok())
.unwrap_or(1.0),
min_p: std::env::var("MEMRA_MIN_P")
.ok()
.and_then(|v| v.parse().ok())
.unwrap_or(0.0),
penalty_last_n: std::env::var("MEMRA_PENALTY_LAST_N")
.ok()
.and_then(|v| v.parse().ok())
.unwrap_or(0),
penalty_repeat: std::env::var("MEMRA_PENALTY_REPEAT")
.ok()
.and_then(|v| v.parse().ok())
.unwrap_or(1.0),
penalty_freq: std::env::var("MEMRA_PENALTY_FREQ")
.ok()
.and_then(|v| v.parse().ok())
.unwrap_or(0.0),
penalty_present: std::env::var("MEMRA_PENALTY_PRESENT")
.ok()
.and_then(|v| v.parse().ok())
.unwrap_or(0.0),
});
let (sp_temp, sp_seed) = (sp.temp, sp.seed);
let sampled = sp_temp > 0.0;
let mut sctr: u32 = sess_tail.as_ref().map(|(_, _, _, s, _)| **s).unwrap_or(0);
let mut uctr: u32 = sess_tail.as_ref().map(|(_, _, _, _, u)| **u).unwrap_or(0);
let pen_on = sampled
&& sp.penalty_last_n > 0
&& (sp.penalty_repeat != 1.0 || sp.penalty_freq != 0.0 || sp.penalty_present != 0.0);
let mut pen_hist: Vec<u32> = if pen_on {
let sess_hist: &[u32] = if spec_pen_session_on() {
sess_tail
.as_ref()
.map(|(c, ..)| c.as_slice())
.unwrap_or(&[])
} else {
&[] };
pen_window_seed(sess_hist, prompt, sp.penalty_last_n)
} else {
Vec::new()
};
if let Some(c) = constraint.as_deref_mut() {
if continuation && carried_pending.is_none() {
return Err("constrained spec continuation requires a carried pending \
(pool resume is unconstrained-only)"
.into());
}
if !continuation {
c.mask_logits(&mut prime_logits)
.map_err(|e2| format!("constraint: {e2}"))?;
}
}
let mut last_token = if let Some(b) = carried_pending {
b
} else if continuation {
sess_tail.as_ref().unwrap().2.unwrap()
} else if sampled && constraint.is_none() && spec_sampled_boundary_on() {
sample_boundary_token(e, &prime_logits, &sp, &pen_hist, &mut sctr, "cold-prime")?
} else {
argmax(&prime_logits) as u32
};
if pen_on {
pen_hist.push(last_token);
}
if carried_pending.is_none() {
out.push(last_token);
if let Some(c) = constraint.as_deref_mut() {
c.consume(last_token)
.map_err(|e2| format!("constraint: {e2}"))?;
}
}
if continuation {
scratch.set_len(e, base)?;
}
fn flush_commit(
cb: &mut Option<&mut dyn FnMut(&[u32]) -> bool>,
out: &[u32],
flushed: &mut usize,
) -> bool {
if let Some(f) = cb.as_mut() {
let keep = f(&out[*flushed..]);
*flushed = out.len();
keep
} else {
true
}
}
keep_going = flush_commit(&mut on_commit, &out, &mut flushed);
let d2t_dev: Option<CudaSlice<u32>> = if sampled || crate::spec::spec_stream() {
match &mtp.d2t {
Some(map) => Some(e.htod_u32_v(map)?),
None => None,
}
} else {
None
};
let mut q_full_buf: Option<CudaSlice<f32>> = None;
let mut draft_logits: Vec<CudaSlice<f32>> = Vec::new(); let mut draft_stats: Vec<(f32, f32, f32)> = Vec::new(); let mut perturb_buf: Option<CudaSlice<f32>> = None; let mut sample_tok = e.alloc_u32_zeroed(1)?; let mut col_buf: Option<CudaSlice<f32>> = None; let mut pen_hist_d: Option<CudaSlice<u32>> = None;
let mut pcol_buf: Option<CudaSlice<f32>> = None; let setup_trace = std::env::var("MEMRA_SPEC_SETUP_TRACE").as_deref() == Ok("1");
let t_ent = std::time::Instant::now();
if let Some(slot) = sess_ckpt_slot {
if let Some(early) = ckpt_early {
if early.is_none() && std::env::var("MEMRA_DEBUG_SPEC").is_ok() {
eprintln!(
"[spec] stable-boundary turn checkpoint skipped; \
next turn re-primes in full"
);
}
*slot = early;
} else if !continuation {
let pos = cache.pos;
debug_assert_eq!(
pos,
base + prompt.len(),
"turn checkpoint must sit at the prompt end, before the init feed"
);
let anchor: Result<CudaSlice<f32>, Box<dyn std::error::Error>> =
if let Some(ph) = &prompt_h {
let np = prompt.len();
e.uninit(n_embd).and_then(|mut a| {
e.copy_view_into(
&mut a,
0,
&ph.slice((np - 1) * n_embd..np * n_embd),
n_embd,
)?;
Ok(a)
})
} else {
Err("no prompt hiddens".into())
};
match (cache.snapshot(e), anchor) {
(Ok(snap), Ok(last_h)) => {
*slot = Some(SpecCheckpoint { snap, pos, last_h });
}
(s, a) => {
*slot = None; if std::env::var("MEMRA_DEBUG_SPEC").is_ok() {
let err = s
.err()
.map(|e| e.to_string())
.or_else(|| a.err().map(|e| e.to_string()))
.unwrap_or_default();
eprintln!(
"[spec] turn checkpoint skipped ({err}); \
next turn re-primes in full"
);
}
}
}
}
}
let mut last_pred = 0u32;
let mut last_col_logits: Option<CudaSlice<f32>> = None;
let mut init_logits_host: Option<Vec<f32>> = None;
let h_seed0: CudaSlice<f32> = if carried_pending.is_none() {
let (init_logits, h) = self.spec_target_step_h(e, last_token, &mut *cache)?;
last_pred = argmax(&init_logits) as u32;
if constraint.is_some() {
init_logits_host = Some(init_logits.clone());
}
if sampled {
last_col_logits = Some(e.htod(&init_logits)?);
}
h
} else {
let lh = sess_tail
.as_ref()
.unwrap()
.1
.as_ref()
.expect("pending carry requires last_h");
e.clone_dtod(lh)?
};
let t_init = t_ent.elapsed();
let mut last_col_stats: Option<(f32, f32, f32)> = None;
let mut h_seed_buf = e.clone_dtod(&h_seed0)?;
let mut fill_prev = e.clone_dtod(&h_seed0)?;
{
if let Some(ph) = &prompt_h {
let np = prompt.len();
e.copy_view_into(
&mut h_seed_buf,
0,
&ph.slice((np - 1) * n_embd..np * n_embd),
n_embd,
)?;
} else if continuation {
if let Some((_, lh, _, _, _)) = sess_tail.as_ref() {
if let Some(lh) = lh.as_ref() {
e.copy_into(&mut h_seed_buf, 0, lh, n_embd)?;
}
}
}
}
let mut preds_d = e.alloc_u32_zeroed(k + 2)?;
let debug_spec = std::env::var("MEMRA_DEBUG_SPEC").is_ok();
let fork_mode = OptiForkGateMode::configured();
let spec_stats = std::env::var("MEMRA_SPEC_STATS").is_ok();
let mut st_drafted = vec![0usize; k];
let mut st_accepted = vec![0usize; k];
let mut st_len_hist = vec![0usize; k + 1];
let mut st_full = 0usize;
static PMIN: std::sync::OnceLock<f32> = std::sync::OnceLock::new();
let p_min = *PMIN.get_or_init(|| {
std::env::var("MEMRA_SPEC_PMIN")
.ok()
.and_then(|v| v.parse().ok())
.unwrap_or(0.0)
});
let pmin0 = std::env::var("MEMRA_SPEC_PMIN0")
.map(|v| v == "1")
.unwrap_or(false);
let mut dctx: DraftGraphCtx = match sess_draft_slot.as_mut().and_then(|s| s.take()) {
Some(c) => c,
None => DraftGraphCtx::new(e, n_embd, if sampled { d_vocab } else { 1 })?,
};
if sampled && dctx.g_q.len() < d_vocab {
dctx.g_q = e.zeros(d_vocab)?;
dctx.g_perturb = e.zeros(d_vocab)?;
}
let dmask_on = constraint
.as_deref()
.is_some_and(|c| c.draft_mask_enabled());
let dmask_words = if dmask_on { d_vocab.div_ceil(32) } else { 0 };
if dmask_on && dctx.g_dmask.len() < dmask_words {
dctx.g_dmask = e.alloc_u32_zeroed(dmask_words)?;
dctx.graph = None; dctx.failed.clear_greedy();
dctx.keeper.clear();
}
if dctx.graph.is_some() && dctx.graph_masked != dmask_on {
dctx.graph = None;
dctx.failed.clear_greedy();
dctx.keeper.clear();
}
if graph_draft && !sampled && dctx.graph.is_none() && !dctx.failed.greedy_failed() {
let DraftGraphCtx {
g_tok,
g_pos,
g_seed,
g_p,
g_dmask,
..
} = &mut dctx;
if dmask_on {
e.htod_u32_into(g_dmask, &vec![u32::MAX; dmask_words])?;
}
let g_dmask_ro: &CudaSlice<u32> = &*g_dmask;
let cap_res = e.capture_graph_retained(|e| {
self.mtp_head_forward_cap(
e,
mtp,
g_tok,
g_pos,
g_seed,
g_p,
&mut *scratch,
p_min > 0.0 || fork_mode == OptiForkGateMode::Controller,
true,
embd_gpu.expect("graph draft requires resident embedding"),
embd_qt,
embd_rb,
d_vocab,
None,
None,
if dmask_on {
Some((g_dmask_ro, dmask_words))
} else {
None
},
)
});
match cap_res {
Ok((g, keep)) => {
scratch.set_len(e, base)?;
dctx.graph = Some(g);
dctx.graph_masked = dmask_on;
dctx.keeper = keep;
}
Err(err) => {
scratch.set_len(e, base)?;
if let Some(line) = dctx.failed.mark_greedy(&err.to_string()) {
eprintln!("{line}");
}
}
}
}
let s_key = SampledGraphKey::new(sp_seed, sp_temp, k, sp.top_k, sp.top_p, sp.min_p, pen_on);
let pure_temp = s_key.pure_temp();
if sampled && dctx.s_key.is_some_and(|old| old != s_key) {
dctx.graph_s = None;
dctx.failed.clear_sampled();
dctx.s_key = None;
dctx.q_slots.clear();
dctx.keeper_s.clear();
}
if graph_draft
&& sampled
&& pure_temp
&& dctx.graph_s.is_none()
&& !dctx.failed.sampled_failed()
{
let DraftGraphCtx {
g_tok,
g_pos,
g_seed,
g_p,
g_ctr,
g_perturb,
g_q,
..
} = &mut dctx;
let cap_res = e.capture_graph_retained(|e| {
self.mtp_head_forward_cap(
e,
mtp,
g_tok,
g_pos,
g_seed,
g_p,
&mut *scratch,
p_min > 0.0,
true,
embd_gpu.expect("graph draft requires resident embedding"),
embd_qt,
embd_rb,
d_vocab,
Some((g_ctr, g_perturb, g_q, sp_seed, sp_temp)),
None,
None, )
});
match cap_res {
Ok((g, keep)) => {
scratch.set_len(e, base)?;
for _ in 0..k {
dctx.q_slots.push(e.zeros(d_vocab)?);
}
dctx.graph_s = Some(g);
dctx.s_key = Some(s_key);
dctx.keeper_s = keep;
}
Err(err) => {
scratch.set_len(e, base)?;
if let Some(line) = dctx.failed.mark_sampled(&err.to_string()) {
eprintln!("{line}");
}
}
}
}
if sampled && !pure_temp && dctx.graph_s.is_some() {
debug_assert!(
false,
"sampled draft graph parked under {:?} survived into a FILTERED request \
(top_k={} top_p={} min_p={} pen_on={}): the in-graph chain draws from the RAW \
softmax, so the verify's filtered q would test a distribution the draft was \
never sampled from",
dctx.s_key, sp.top_k, sp.top_p, sp.min_p, pen_on,
);
eprintln!(
"[spec] BUG: dropping a parked sampled draft graph that outlived its capture \
regime (s_key={:?}, request top_k={} top_p={} min_p={} pen_on={}); drafting \
EAGER — the key must carry every field that shapes q",
dctx.s_key, sp.top_k, sp.top_p, sp.min_p, pen_on,
);
dctx.graph_s = None;
dctx.s_key = None;
dctx.q_slots.clear();
dctx.keeper_s.clear();
}
if skey_probe() {
eprintln!(
"[skey] burst sampled={} pure_temp={} temp={} top_k={} top_p={} min_p={} \
pen_on={} k={} graph_draft={} graph_s_parked={} s_key_parked={:?}",
sampled as u8,
pure_temp as u8,
sp_temp,
sp.top_k,
sp.top_p,
sp.min_p,
pen_on as u8,
k,
graph_draft as u8,
dctx.graph_s.is_some() as u8,
dctx.s_key,
);
}
let t_cap = t_ent.elapsed();
if let Some(ph) = &prompt_h {
scratch.set_len(e, base)?;
let tp = prompt.len();
let fill_chunk: usize = if crate::cache::swa_ring_on() {
crate::hybrid_forward::prime_chunk_tokens(tp, self.layers.len())
} else {
std::env::var("MEMRA_PRIME_CHUNK")
.ok()
.and_then(|v| v.parse().ok())
.unwrap_or(4096)
};
let fill_chunk = if fill_chunk == 0 { tp } else { fill_chunk };
let mut start = 0usize;
while start < tp {
let end = (start + fill_chunk).min(tp);
let tc = end - start;
{
let mut phs = e.zeros(tc * n_embd)?;
let (src_lo, dst_off) = if start == 0 {
(0, n_embd)
} else {
((start - 1) * n_embd, 0)
};
let n_copy = if start == 0 {
(tc - 1) * n_embd
} else {
tc * n_embd
};
if start == 0 {
if let Some((_, lh, _, _, _)) = sess_tail.as_ref() {
if let Some(lh) = lh.as_ref() {
e.copy_into(&mut phs, 0, lh, n_embd)?;
}
}
}
if n_copy > 0 {
e.copy_view_into(
&mut phs,
dst_off,
&ph.slice(src_lo..src_lo + n_copy),
n_copy,
)?;
}
self.mtp_kv_fill(
e,
mtp,
&prompt[start..end],
&phs,
base + start,
&mut *scratch,
embd_dev,
)?;
}
start = end;
}
}
if std::env::var("MEMRA_PROFILE_SPEC").as_deref() == Ok("2") {
unsafe extern "C" {
fn cudaProfilerStart() -> i32;
}
unsafe {
cudaProfilerStart();
}
}
let stream_on = crate::spec::spec_stream()
&& !sampled
&& !spec_replay
&& constraint.is_none()
&& !session_mode
&& embd_gpu.is_some()
&& !crate::model::full_prec_enabled()
&& k + 2 < 96;
let mut stream_graph: Option<cudarc::driver::CudaGraph> = None;
let mut g_tokp2k = e.alloc_u32_zeroed(2 * k.max(1))?;
if stream_on {
let cap = e.capture_graph(|e| {
for j in 0..k.max(1) {
self.mtp_head_forward_cap(
e,
mtp,
&mut dctx.g_tok,
&mut dctx.g_pos,
&mut dctx.g_seed,
&mut dctx.g_p,
&mut *scratch,
true,
true,
embd_gpu.expect("round stream requires resident embedding"),
embd_qt,
embd_rb,
d_vocab,
None,
Some((&mut g_tokp2k, j, d2t_dev.as_ref())),
None, )?;
}
Ok(())
});
match cap {
Ok(g) => {
scratch.set_len(e, 0)?;
stream_graph = Some(g);
}
Err(err) => {
scratch.set_len(e, 0)?;
if debug_spec {
eprintln!("[spec] stream-graph capture failed ({err}); stream off");
}
}
}
}
let stream_active = stream_on && stream_graph.is_some();
if debug_spec {
eprintln!(
"[spec] stream_on={stream_on} env={} samp={sampled} dg={} captured={} active={stream_active} session={session_mode} replay={spec_replay}",
crate::spec::spec_stream(),
dctx.graph.is_some(),
stream_graph.is_some()
);
}
let t_v_s = k + 1;
let sb = crate::round_stream::StreamBufs::new(e, k, crate::spec::spec_stream_m())?;
let crate::round_stream::StreamBufs {
mut vtok_d,
mut brk_d,
mut pend_d,
last_pred_d,
mut pos_ctr,
mut pos_start_d,
mut ring_d,
acc_d: mut stream_acc,
m_rounds,
k: _,
} = sb;
let stream_ptrs: Option<CudaSlice<u64>> = if stream_active {
Some(crate::round_stream::kv_len_ptr_table(
e,
cache,
Some(&pos_ctr),
)?)
} else {
None
};
let t_fill = t_ent.elapsed();
let mut round = 0usize;
let adapt = std::env::var("MEMRA_SPEC_ADAPT").as_deref() == Ok("1");
let adapt_floor_env: Option<usize> = std::env::var("MEMRA_SPEC_ADAPT_FLOOR")
.ok()
.and_then(|v| v.parse().ok());
let adapt_floor_default: usize = if self.cfg.n_embd as usize >= 3500 {
4
} else if self.cfg.n_embd as usize >= 2500 {
2
} else {
1
};
let adapt_floor: usize = adapt_floor_env.unwrap_or(adapt_floor_default);
let floor_ctx: usize = std::env::var("MEMRA_SPEC_FLOOR_CTX")
.ok()
.and_then(|v| v.parse().ok())
.unwrap_or(1024);
let floor_at = |pos: usize| -> usize {
if adapt_floor_env.is_some() || pos < floor_ctx {
adapt_floor
} else if adapt_floor >= 4 {
1
} else {
adapt_floor
}
};
let cap_max: usize = std::env::var("MEMRA_SPEC_CAPMAX")
.ok()
.and_then(|v| v.parse().ok())
.unwrap_or(7);
let k_cap = k.min(cap_max).max(1);
let mut kc = k_cap;
let mut opti_fork: Option<OptiForkState> = None;
let mut fork_snapshot: Option<crate::cache::CacheSnapshot> = None;
if fork_mode != OptiForkGateMode::Disabled {
let fence = crate::pp::pp_cuts(self.layers.len());
let refusal = if !session_mode {
Some("not-session")
} else if k != 1 || adapt {
Some("requires-fixed-k1")
} else if sampled || constraint.is_some() || spec_replay {
Some("sampled-constrained-or-replay")
} else if pipe.is_some() {
Some("two-session-pipeline")
} else if !spec_devacc() {
Some("requires-device-accept")
} else if stream_active || crate::spec::spec_stream() {
Some("round-stream")
} else if crate::cache::swa_ring_on() || cache.has_swa_ring() {
Some("swa-ring")
} else if crate::pp::pp_host_bounce_active() {
Some("host-bounce")
} else if fork_mode == OptiForkGateMode::Controller
&& cache.recur.iter().any(Option::is_some)
{
Some("controller-requires-zero-recurrent-state")
} else if fence.as_ref().is_none_or(|f| f.len() != 3) {
Some("requires-pp2")
} else {
None
};
if let Some(reason) = refusal {
OPTI_FORK_REFUSALS.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
eprintln!("[opti-fork] refused reason={reason}");
} else {
let fence = fence.expect("validated PP-2 fence");
let rt = crate::pp::PpNRt::get(e)?;
let primary_stage0 = rt.engine(0, e).ctx().ordinal() == e.ctx().ordinal();
let primary_stage1 = rt.engine(1, e).ctx().ordinal() == e.ctx().ordinal();
let primary_supported =
primary_stage0 || (fork_mode == OptiForkGateMode::Controller && primary_stage1);
if !rt.cross_device() || !primary_supported {
OPTI_FORK_REFUSALS.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
eprintln!("[opti-fork] refused reason=requires-supported-primary-cross-device");
} else {
let current_snapshot = opti_snapshot_stage_owned(e, cache, rt, &fence)?;
let alternate_snapshot = opti_snapshot_stage_owned(e, cache, rt, &fence)?;
let fork = OptiForkState::new(
e,
cache,
fork_mode,
alternate_snapshot,
&h_seed_buf,
&fill_prev,
rt,
fence[1],
self.layers.len(),
)?;
eprintln!(
"[opti-fork] armed mode={fork_mode:?} snapshots=2 seeds=2 split={} \
payload_dev0={} payload_dev1={} q_threshold={:.3}",
fence[1],
fork.logical_payload_bytes[0],
fork.logical_payload_bytes[1],
fork.controller.map_or(0.0, |policy| policy.threshold),
);
fork_snapshot = Some(current_snapshot);
opti_fork = Some(fork);
}
}
}
let mut snap = match fork_snapshot {
Some(snapshot) => snapshot,
None => cache.snapshot(e)?,
};
let mut carried_opti: Option<OptiControllerTicket> = None;
let kv_len_ptrs: Option<CudaSlice<u64>> = if spec_devacc() && !spec_replay {
Some(crate::round_stream::kv_len_ptr_table(e, cache, None)?)
} else {
None
};
let mut pending: Option<u32> = carried_pending;
let anatomy_on = std::env::var("MEMRA_SPEC_PP_ANATOMY").as_deref() == Ok("1");
let phase_on = anatomy_on || std::env::var("MEMRA_SPEC_PHASE").as_deref() == Ok("1");
let (mut dm_clone_ns, mut dm_rounds) = (0u128, 0usize);
let (mut dm_cuts, mut dm_cut_tokens) = (0usize, 0usize);
let (mut ph_draft, mut ph_verify, mut ph_rest) = (0f64, 0f64, 0f64);
let mut ph_wait = 0f64;
let mut ph_commit = 0f64;
let mut ph_t = std::time::Instant::now();
let mut ph_mark = |acc: &mut f64, on: bool| {
if on {
let now = std::time::Instant::now();
*acc += (now - ph_t).as_secs_f64();
ph_t = now;
}
};
if let Some(p) = pipe {
p.setup_end();
}
while keep_going && out.len() < max_new {
if let (true, Some(sg), Some(ptrs)) = (
stream_active && round >= 1 && pending.is_some(),
&stream_graph,
&stream_ptrs,
) {
if debug_spec {
static ONCE: std::sync::Once = std::sync::Once::new();
ONCE.call_once(|| {
eprintln!("[memra] ROUND-STREAM burst engaged (M={m_rounds} k={k})")
});
}
e.set_i32_one(&mut pos_ctr, cache.pos as i32)?;
e.set_u32_one(&mut pend_d, pending.unwrap())?;
e.set_u32_one(&mut ring_d, 0)?; for _mi in 0..m_rounds {
e.i32_copy_add(&pos_ctr, &mut pos_start_d, 0)?;
cache.snapshot_into(e, &mut snap)?; e.i32_copy_add(&pos_ctr, &mut scratch.kv.len_d, 0)?; e.i32_copy_add(&pos_ctr, &mut dctx.g_pos, 1)?; e.u32_copy(&pend_d, &mut dctx.g_tok)?;
e.copy_into(&mut dctx.g_seed, 0, &h_seed_buf, n_embd)?;
sg.launch()?;
e.spec_assemble_verify(
&g_tokp2k,
&pend_d,
d2t_dev.as_ref(),
&mut vtok_d,
&mut brk_d,
p_min,
k,
pmin0,
)?;
let mut ck = VerifyCkpt::new(self.layers.len());
let dummy = vec![0u32; t_v_s];
let (tl_d, vx) = self.decode_step_t_core_stream(
e,
&dummy,
0,
&mut *cache,
embd_dev,
Some(&mut ck),
Some((&vtok_d, &pos_ctr)),
None,
None,
None,
)?;
for j in 0..t_v_s {
e.argmax_token_device_col(&tl_d, j, n_vocab, &mut preds_d, j)?;
}
e.spec_accept_greedy_dc(
&preds_d,
&vtok_d,
&last_pred_d,
&brk_d,
&mut stream_acc,
)?;
e.spec_seed_gather(&vx, &fill_prev, &stream_acc, &mut h_seed_buf, 1, n_embd)?;
e.copy_into(&mut fill_prev, 0, &h_seed_buf, n_embd)?;
self.commit_verified_prefix_stream(
e,
&mut *cache,
&snap,
&ck,
&stream_acc,
1,
t_v_s,
)?;
e.spec_rollback_stream(
ptrs,
&pos_start_d,
&stream_acc,
1,
self.layers.len() + 1,
)?;
e.spec_ring_commit(&vtok_d, &stream_acc, &brk_d, &mut ring_d, &mut pend_d)?;
}
e.stream().synchronize()?;
let ring_h = e.dtoh_u32(&ring_d)?;
let cnt = ring_h[0] as usize;
for i in 0..cnt {
if out.len() < max_new {
out.push(ring_h[1 + i]);
}
}
let pos_h = e.dtoh_i32(&pos_ctr)?[0] as usize;
for il in 0..self.layers.len() {
if let Some(kvl) = cache.kv[il].as_mut() {
kvl.len = pos_h;
}
}
cache.pos = pos_h;
scratch.kv.len = pos_h;
pending = Some(ring_h[cnt]); last_token = ring_h[cnt];
total_drafted += k * m_rounds; total_accepted += cnt.saturating_sub(m_rounds);
if let Some(t) = sess_telem {
t.record_totals(m_rounds, k * m_rounds, cnt.saturating_sub(m_rounds));
}
round += m_rounds;
keep_going = flush_commit(&mut on_commit, &out, &mut flushed);
continue;
}
let pipe_draft = match pipe {
Some(p) => Some(p.draft_begin(round)?),
None => None,
};
let pos = cache.pos; let mut current_opti = carried_opti.take();
let mut fork_generation = if current_opti.is_none() && pending.is_some() {
match opti_fork.as_mut() {
Some(fork) if fork.mode.is_forced() => Some(fork.reserve(&mut snap)?),
None => None,
Some(_) => None,
}
} else {
None
};
if current_opti.is_none() {
if let Some(fork) = opti_fork.as_ref() {
opti_snapshot_stage_owned_into(e, cache, fork.rt, &fork.fence, &mut snap)?;
} else {
cache.snapshot_into(e, &mut snap)?;
}
} else if snap.pos != pos {
return Err(format!(
"optipipe carried snapshot pos {} != current pos {pos}",
snap.pos
)
.into());
} ph_mark(&mut ph_rest, phase_on);
let base0 = if pending.is_some() { 1usize } else { 0usize };
let k_this = if adapt { kc } else { k };
let mut draft: Vec<u32> = Vec::with_capacity(k);
let mut draft_idx: Vec<u32> = Vec::with_capacity(k); let mut controller_draft_prob: Option<f32> = None;
let mut controller_eager_state: Option<(u32, CudaSlice<f32>)> = None;
if let Some(ticket) = current_opti.as_mut() {
let carried_pending = pending.ok_or("optipipe carried successor lost pending")?;
if ticket.verify_tokens[0] != carried_pending {
return Err(format!(
"optipipe carried pending mismatch: ticket={} live={carried_pending}",
ticket.verify_tokens[0],
)
.into());
}
draft.push(ticket.verify_tokens[1]);
controller_draft_prob = Some(ticket.draft_prob);
controller_eager_state = ticket
.take_eager_seed()
.map(|seed| (ticket.verify_tokens[1], seed));
} else {
scratch.set_len(e, pos + base0 - 1)?;
if pen_on {
let win = sp.penalty_last_n.min(PEN_WINDOW_MAX);
let w0 = pen_hist.len().saturating_sub(win);
pen_hist_d = Some(e.htod_u32_v(&pen_hist[w0..])?);
}
if sampled {
draft_logits.clear();
draft_stats.clear();
}
let mut dmask_live = dmask_on;
if dmask_live {
let t_c = std::time::Instant::now();
constraint
.as_deref_mut()
.unwrap()
.draft_begin()
.map_err(|e2| format!("constraint: {e2}"))?;
dm_clone_ns += t_c.elapsed().as_nanos();
dm_rounds += 1;
}
if let (false, Some(gr)) = (sampled || pen_on, &dctx.graph) {
e.set_i32_one(&mut dctx.g_pos, (pos + base0) as i32)?;
e.set_u32_one(&mut dctx.g_tok, last_token)?;
e.copy_into(&mut dctx.g_seed, 0, &h_seed_buf, n_embd)?;
for j in 0..k_this {
if dmask_live
&& !upload_draft_mask(
e,
constraint.as_deref_mut().unwrap(),
&mut dctx.g_dmask,
mtp.d2t.as_ref(),
d_vocab,
dmask_words,
)?
{
e.htod_u32_into(&mut dctx.g_dmask, &vec![u32::MAX; dmask_words])?;
dmask_live = false;
}
gr.launch()?;
scratch.kv.len += 1; let idx = e.dtoh_u32_one(&dctx.g_tok)?;
if (idx as usize) >= d_vocab {
let seed_h = e.dtoh(&dctx.g_seed)?;
let seed_nan = seed_h.iter().filter(|v| v.is_nan()).count();
let in_h = e.dtoh(&h_seed_buf)?;
let in_nan = in_h.iter().filter(|v| v.is_nan()).count();
return Err(format!(
"draft(graph) argmax sentinel 0x{idx:08x} >= d_vocab {d_vocab} at \
round {round} j={j} pos={pos}: head-out NaN {seed_nan}/{n_embd}, \
round-input-seed NaN {in_nan}/{n_embd} — refusing to dereference \
the embed row (#87 trap)"
)
.into());
}
let d = match &mtp.d2t {
Some(map) => map[idx as usize],
None => idx,
};
let draft_p = if p_min > 0.0
|| opti_fork
.as_ref()
.is_some_and(|fork| fork.controller.is_some())
{
Some(e.dtoh(&dctx.g_p)?[0])
} else {
None
};
if j == 0 {
controller_draft_prob = draft_p;
}
if let Some(p) = draft_p.filter(|_| p_min > 0.0) {
if p < p_min && (j > 0 || (pmin0 && base0 == 1)) {
break;
}
}
draft.push(d);
if d != idx {
e.set_u32_one(&mut dctx.g_tok, d)?;
}
if dmask_live
&& !constraint
.as_deref_mut()
.unwrap()
.draft_advance(d)
.map_err(|e2| format!("constraint: {e2}"))?
{
e.htod_u32_into(&mut dctx.g_dmask, &vec![u32::MAX; dmask_words])?;
break;
}
}
} else if let (true, Some(gr)) = (sampled && pure_temp, &dctx.graph_s) {
if skey_probe() {
eprintln!(
"[skey] chain=graph_s round={round} pure_temp={} top_k={} \
top_p={} min_p={} s_key_parked={:?}",
pure_temp as u8, sp.top_k, sp.top_p, sp.min_p, dctx.s_key,
);
}
e.set_i32_one(&mut dctx.g_pos, (pos + base0) as i32)?;
e.set_u32_one(&mut dctx.g_tok, last_token)?;
e.copy_into(&mut dctx.g_seed, 0, &h_seed_buf, n_embd)?;
e.set_u32_one(&mut dctx.g_ctr, sctr.wrapping_sub(1))?;
for j in 0..k_this {
gr.launch()?;
scratch.kv.len += 1; sctr += 1; e.copy_into(&mut dctx.q_slots[j], 0, &dctx.g_q, d_vocab)?;
let idx = e.dtoh_u32_one(&dctx.g_tok)?;
if (idx as usize) >= d_vocab {
let seed_h = e.dtoh(&dctx.g_seed)?;
let seed_nan = seed_h.iter().filter(|v| v.is_nan()).count();
return Err(format!(
"draft(graph-sampled) argmax sentinel 0x{idx:08x} >= d_vocab \
{d_vocab} at round {round} j={j} pos={pos}: round-seed NaN \
{seed_nan}/{n_embd} — refusing to dereference the embed row \
(#87 trap)"
)
.into());
}
let d = match &mtp.d2t {
Some(map) => map[idx as usize],
None => idx,
};
draft_idx.push(idx);
if p_min > 0.0 {
let p = e.dtoh(&dctx.g_p)?[0];
if p < p_min && (j > 0 || (pmin0 && base0 == 1)) {
break;
}
}
draft.push(d);
if d != idx {
e.set_u32_one(&mut dctx.g_tok, d)?;
}
}
for j in 0..draft.len().max(draft_idx.len()) {
let rows0 = e.htod_i32(&[0])?;
let (mut th_d, mut z_d, mut mx_d) = (e.zeros(1)?, e.zeros(1)?, e.zeros(1)?);
e.filter_stats(
&dctx.q_slots[j],
d_vocab,
&rows0,
&mut th_d,
&mut z_d,
&mut mx_d,
d_vocab,
1,
sp_temp,
sp.top_k,
sp.top_p,
sp.min_p,
)?;
draft_stats.push((e.dtoh(&mx_d)?[0], e.dtoh(&th_d)?[0], e.dtoh(&z_d)?[0]));
}
} else {
if skey_probe() && sampled {
eprintln!(
"[skey] chain=eager round={round} pure_temp={} top_k={} \
top_p={} min_p={} s_key_parked={:?}",
pure_temp as u8, sp.top_k, sp.top_p, sp.min_p, dctx.s_key,
);
}
let mut e_tok = last_token;
let mut d_seed = e.clone_dtod(&h_seed_buf)?;
for j in 0..k_this {
let mtp_pos = pos + base0 + j;
if dmask_live {
dmask_live = upload_draft_mask(
e,
constraint.as_deref_mut().unwrap(),
&mut dctx.g_dmask,
mtp.d2t.as_ref(),
d_vocab,
dmask_words,
)?;
}
let (dl_d, h_nextn) = self.mtp_head_forward_dev(
e,
mtp,
e_tok,
&d_seed,
&mut *scratch,
mtp_pos,
embd_dev,
if dmask_live {
Some((&dctx.g_dmask, dmask_words))
} else {
None
},
)?;
let tok_d = if sampled {
if perturb_buf.is_none() {
perturb_buf = Some(e.zeros(d_vocab.max(n_vocab))?);
}
let mut q_row = e.clone_dtod(&dl_d)?; if pen_on {
let h = pen_hist_d.as_ref().unwrap();
let nh = h.len();
e.penalize_logits(
&mut q_row,
h,
nh,
sp.penalty_repeat,
sp.penalty_freq,
sp.penalty_present,
d_vocab,
)?;
}
let rows0 = e.htod_i32(&[0])?;
let (mut th_d, mut z_d, mut mx_d) =
(e.zeros(1)?, e.zeros(1)?, e.zeros(1)?);
e.filter_stats(
&q_row, d_vocab, &rows0, &mut th_d, &mut z_d, &mut mx_d, d_vocab,
1, sp_temp, sp.top_k, sp.top_p, sp.min_p,
)?;
let (th, z, mx) =
(e.dtoh(&th_d)?[0], e.dtoh(&z_d)?[0], e.dtoh(&mx_d)?[0]);
let pb = perturb_buf.as_mut().unwrap();
e.gumbel_perturb_filtered(
&q_row, pb, d_vocab, sp_seed, sctr, sp_temp, mx, th,
)?;
sctr += 1;
draft_logits.push(q_row);
draft_stats.push((mx, th, z));
e.argmax_token_device(pb, d_vocab)?
} else {
e.argmax_token_device(&dl_d, d_vocab)?
};
let idx = e.dtoh_u32_one(&tok_d)?;
if (idx as usize) >= d_vocab {
let dl_h = e.dtoh(&dl_d)?;
let dl_nan = dl_h.iter().filter(|v| v.is_nan()).count();
let seed_h = e.dtoh(&d_seed)?;
let seed_nan = seed_h.iter().filter(|v| v.is_nan()).count();
return Err(format!(
"draft(eager) argmax sentinel 0x{idx:08x} >= d_vocab {d_vocab} at \
round {round} j={j} pos={pos}: head-logits NaN {dl_nan}/{d_vocab}, \
step-seed NaN {seed_nan}/{n_embd} — refusing to dereference the \
embed row (#87 trap)"
)
.into());
}
let d = match &mtp.d2t {
Some(map) => map[idx as usize],
None => idx,
};
if sampled {
draft_idx.push(idx);
}
let draft_p = if p_min > 0.0
|| opti_fork
.as_ref()
.is_some_and(|fork| fork.controller.is_some())
{
let p_d = e.prob_of_token_device(&dl_d, &tok_d, d_vocab)?;
Some(e.dtoh(&p_d)?[0])
} else {
None
};
if j == 0 {
controller_draft_prob = draft_p;
}
if let Some(p) = draft_p.filter(|_| p_min > 0.0) {
if p < p_min && (j > 0 || (pmin0 && base0 == 1)) {
break;
}
}
draft.push(d);
e_tok = d;
d_seed = h_nextn;
if dmask_live
&& !constraint
.as_deref_mut()
.unwrap()
.draft_advance(d)
.map_err(|e2| format!("constraint: {e2}"))?
{
break;
}
}
if opti_fork
.as_ref()
.is_some_and(|fork| fork.controller.is_some())
{
controller_eager_state = Some((e_tok, d_seed));
}
}
}
let k_round = draft.len();
if let Some(p) = pipe {
p.draft_end(round);
}
drop(pipe_draft);
ph_mark(&mut ph_draft, phase_on);
let verify_tokens: Vec<u32> = match pending {
Some(b) => {
let mut v = Vec::with_capacity(k_round + 1);
v.push(b);
v.extend_from_slice(&draft);
v
}
None => draft.clone(),
};
let base = if pending.is_some() { 1 } else { 0 };
let mut ckpt = if let Some(ticket) = current_opti.as_mut() {
Some(ticket.take_ckpt())
} else if spec_replay {
None
} else {
Some(VerifyCkpt::new(self.layers.len()))
};
let controller_can_probe = base == 1
&& k_round == 1
&& out.len().saturating_add(2) < max_new
&& controller_draft_prob.is_some()
&& opti_fork
.as_ref()
.and_then(|fork| fork.controller.as_ref())
.is_some_and(|policy| !policy.breaker_tripped);
let mut successor_attempt: Option<OptiControllerTicket> = None;
let mut rejected_probe: Option<(f32, u32)> = None;
let mut controller_prepared: Option<OptiControllerPrepared> = None;
if controller_can_probe {
let eager_pos = scratch.kv.len + 1;
let (optimistic_pending, pending_probability) = self.opti_controller_draft_step(
e,
mtp,
&mut dctx,
&mut *scratch,
d_vocab,
&mut controller_eager_state,
eager_pos,
embd_dev,
)?;
let first_probability = controller_draft_prob
.ok_or("optipipe controller probe lost first-token probability")?;
let q_proxy = first_probability * pending_probability;
OPTI_GATE_CHECKS.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
OPTI_SHADOW_DRAFT_TOKENS.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
let admitted = opti_fork
.as_ref()
.and_then(|fork| fork.controller.as_ref())
.ok_or("optipipe controller policy disappeared")?
.admit(q_proxy);
if admitted {
OPTI_GATE_ADMITS.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
let eager_pos = scratch.kv.len + 1;
let (optimistic_draft, optimistic_draft_probability) = self
.opti_controller_draft_step(
e,
mtp,
&mut dctx,
&mut *scratch,
d_vocab,
&mut controller_eager_state,
eager_pos,
embd_dev,
)?;
OPTI_SHADOW_DRAFT_TOKENS.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
let eager_seed = controller_eager_state.take().map(|(token, seed)| {
debug_assert_eq!(token, optimistic_draft);
seed
});
controller_prepared = Some(OptiControllerPrepared {
verify_tokens: [optimistic_pending, optimistic_draft],
draft_prob: optimistic_draft_probability,
eager_seed,
q_proxy,
scratch_len: scratch.kv.len,
});
} else {
OPTI_GATE_REJECTS.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
OPTI_WASTED_DRAFT_TOKENS.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
rejected_probe = Some((q_proxy, optimistic_pending));
eprintln!(
"[opti-controller] reject q={q_proxy:.6} threshold={:.3}",
opti_fork
.as_ref()
.and_then(|fork| fork.controller.as_ref())
.expect("controller policy")
.threshold,
);
}
}
let fork_attempt = match fork_generation.take() {
Some(generation) if base == 1 && k_round == 1 => Some(generation),
Some(generation) => {
opti_fork
.as_mut()
.expect("fork generation without fork state")
.retire(generation)?;
None
}
None => None,
};
let (tlogits_d, vx) = if let Some(p) = pipe {
self.decode_step_t_core_pipelined(
e,
&verify_tokens,
pos,
&mut *cache,
embd_dev,
ckpt.as_mut(),
p,
round,
)?
} else if controller_can_probe {
let fence = opti_fork
.as_ref()
.ok_or("optipipe controller probe lost fork state")?
.fence;
let boundary = match current_opti.as_mut() {
Some(ticket) => ticket.take_boundary(),
None => self.verify_stage0_issue(
e,
&verify_tokens,
pos,
&mut *cache,
embd_dev,
ckpt.as_mut(),
None,
&fence,
Some(true),
None,
)?,
};
if let Some(prepared) = controller_prepared.take() {
let generation = {
let fork = opti_fork
.as_mut()
.ok_or("optipipe controller admission lost fork state")?;
let generation = fork.reserve_successor()?;
let rt = fork.rt;
let snapshot_fence = fork.fence;
opti_snapshot_one_stage_owned_into(
e,
cache,
rt,
&snapshot_fence,
0,
fork.successor_snapshot_mut(),
)?;
generation
};
let mut successor_ckpt = VerifyCkpt::new(self.layers.len());
let successor_boundary = self.verify_stage0_issue(
e,
&prepared.verify_tokens,
pos + verify_tokens.len(),
&mut *cache,
embd_dev,
Some(&mut successor_ckpt),
None,
&fence,
Some(false),
None,
)?;
OPTI_FORK_ATTEMPTS.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
let fork = opti_fork
.as_ref()
.ok_or("optipipe controller ticket lost fork state")?;
successor_attempt = Some(fork.controller_ticket(
generation,
successor_boundary,
successor_ckpt,
prepared.verify_tokens,
prepared.draft_prob,
prepared.eager_seed,
prepared.q_proxy,
prepared.scratch_len,
));
eprintln!(
"[opti-controller] issue generation={} q={:.6} threshold={:.3} \
verify={:?}",
generation.id,
prepared.q_proxy,
fork.controller.expect("controller policy").threshold,
prepared.verify_tokens,
);
}
let result = self.verify_stage1_finish(
e,
boundary,
&mut *cache,
ckpt.as_mut(),
None,
&fence,
successor_attempt.is_none(),
)?;
if let Some(ticket) = current_opti.as_mut() {
ticket.settle();
}
if successor_attempt.is_some() {
let fork = opti_fork
.as_mut()
.ok_or("optipipe successor snapshot lost fork state")?;
let rt = fork.rt;
let snapshot_fence = fork.fence;
opti_snapshot_one_stage_owned_into(
e,
cache,
rt,
&snapshot_fence,
1,
fork.successor_snapshot_mut(),
)?;
fork.rt.publish_to(1, &e.stream())?;
}
result
} else if let Some(ticket) = current_opti.as_mut() {
let fork = opti_fork
.as_mut()
.ok_or("optipipe carried controller ticket lost fork state")?;
let boundary = ticket.take_boundary();
let result = self.verify_stage1_finish(
e,
boundary,
&mut *cache,
ckpt.as_mut(),
None,
&fork.fence,
true,
)?;
ticket.settle();
result
} else if let Some(generation) = fork_attempt {
let fork = opti_fork
.as_mut()
.expect("fork generation without fork state");
fork.capture_seed(e, generation, &h_seed_buf, &fill_prev, scratch.kv.len)?;
let action = fork.mode.action(generation.id);
let boundary = self.verify_stage0_issue(
e,
&verify_tokens,
pos,
&mut *cache,
embd_dev,
ckpt.as_mut(),
None,
&fork.fence,
Some(true),
None,
)?;
OPTI_FORK_ATTEMPTS.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
let mut ticket = fork.ticket(generation, boundary);
if action == OptiForkAction::Abort {
return Err(format!(
"optipipe forced abort with generation {} stage0 in flight",
generation.id,
)
.into());
}
fork.reconcile(
e,
&mut *cache,
&mut *scratch,
&snap,
&mut h_seed_buf,
&mut fill_prev,
generation,
action,
verify_tokens[0],
)?;
let result = if action == OptiForkAction::Hit {
let boundary = ticket.take_boundary();
self.verify_stage1_finish(
e,
boundary,
&mut *cache,
ckpt.as_mut(),
None,
&fork.fence,
true,
)?
} else {
self.decode_step_t_core(
e,
&verify_tokens,
pos,
&mut *cache,
embd_dev,
ckpt.as_mut(),
)?
};
ticket.settle();
debug_assert_eq!(ticket.generation, generation);
fork.retire(generation)?;
result
} else {
self.decode_step_t_core(
e,
&verify_tokens,
pos,
&mut *cache,
embd_dev,
ckpt.as_mut(),
)?
};
let pipe_accept = match pipe {
Some(p) => Some(p.accept_begin(round)?),
None => None,
};
ph_mark(&mut ph_verify, phase_on);
let t_v = verify_tokens.len();
let mut preds: Vec<u32> = Vec::new();
if !sampled {
for j in 0..t_v {
e.argmax_token_device_col(&tlogits_d, j, n_vocab, &mut preds_d, j)?;
}
preds = e.dtoh_u32(&preds_d)?; if let Some(bad) = preds[..t_v].iter().position(|&p| (p as usize) >= n_vocab) {
let col = &tlogits_d.slice(bad * n_vocab..(bad + 1) * n_vocab);
let mut probe = e.zeros(n_vocab)?;
e.copy_view_into(&mut probe, 0, col, n_vocab)?;
let col_h = e.dtoh(&probe)?;
let col_nan = col_h.iter().filter(|v| v.is_nan()).count();
return Err(format!(
"verify argmax sentinel 0x{:08x} >= n_vocab {n_vocab} at round {round} \
col {bad}/{t_v} pos={pos}: verify-logits col NaN {col_nan}/{n_vocab} \
— the stage-split verify produced a poisoned column (#87 trap)",
preds[bad]
)
.into());
}
}
ph_mark(&mut ph_wait, phase_on);
let t_pred = |j: usize| -> u32 {
if j == 0 && base == 0 {
last_pred
} else {
debug_assert!(
!sampled,
"t_pred is greedy-only: `preds` is empty in the sampled arm"
);
preds[base + j - 1]
}
};
let mut devacc_seeded = false;
let mut devacc_acc: Option<CudaSlice<u32>> = None;
let (n_acc, bonus) = if !sampled {
if crate::spec::spec_devacc() && k_round > 0 && !spec_replay && constraint.is_none()
{
let draft_d = e.htod_u32_v(&draft)?;
let mut acc_out = e.alloc_u32_zeroed(2)?;
e.spec_accept_greedy(
&preds_d,
&draft_d,
last_pred,
base,
k_round,
&mut acc_out,
)?;
devacc_acc = Some(acc_out.clone());
e.spec_seed_gather(&vx, &fill_prev, &acc_out, &mut h_seed_buf, base, n_embd)?;
if let Some(successor) = successor_attempt.as_ref() {
opti_fork
.as_mut()
.ok_or("optipipe successor reconcile lost fork state")?
.queue_actual_reconcile(
e,
&snap,
&acc_out,
successor.verify_tokens[0],
base,
)?;
} else if let Some(ptrs) = &kv_len_ptrs {
let saved: Vec<i32> = (0..self.layers.len())
.map(|il| snap.kv_len[il].map(|v| v as i32).unwrap_or(0))
.collect();
let saved_d = e.htod_i32(&saved)?;
e.spec_rollback_kv(ptrs, &saved_d, &acc_out, base, self.layers.len())?;
}
devacc_seeded = true;
let ab = e.dtoh_u32(&acc_out)?;
(ab[0] as usize, ab[1])
} else {
let mut n_acc = 0usize;
for j in 0..k_round {
if t_pred(j) == draft[j] {
n_acc += 1;
} else {
break;
}
}
(n_acc, t_pred(n_acc))
}
} else {
if col_buf.is_none() {
col_buf = Some(e.zeros(n_vocab)?);
}
let mut pj = vec![0f32; k_round.max(1)];
let mut col_stats: Vec<(f32, f32, f32)> = Vec::new(); if k_round > 0 {
let mut ids: Vec<u32> = Vec::new();
let mut rows: Vec<i32> = Vec::new();
for j in 0..k_round {
if j > 0 || base == 1 {
ids.push(draft[j]);
rows.push((base + j) as i32 - 1);
}
}
if !ids.is_empty() {
let nr = rows.len();
let p_rows: Vec<i32> = if pen_on {
(0..nr as i32).collect()
} else {
rows.clone()
};
if pen_on {
if pcol_buf.as_ref().map(|b| b.len()).unwrap_or(0) < nr * n_vocab {
pcol_buf = Some(e.zeros(nr * n_vocab)?);
}
let pc = pcol_buf.as_mut().unwrap();
for (i2, &r) in rows.iter().enumerate() {
let c = r as usize;
e.copy_view_into(
pc,
i2 * n_vocab,
&tlogits_d.slice(c * n_vocab..(c + 1) * n_vocab),
n_vocab,
)?;
}
let h = pen_hist_d.as_ref().unwrap();
let nh = h.len();
e.penalize_logits_rows(
pc,
h,
nh,
sp.penalty_repeat,
sp.penalty_freq,
sp.penalty_present,
n_vocab,
nr,
)?;
}
let p_src: &CudaSlice<f32> = if pen_on {
pcol_buf.as_ref().unwrap()
} else {
&tlogits_d
};
let rowsd = e.htod_i32(&p_rows)?;
let (mut th_d, mut z_d, mut mx_d) =
(e.zeros(nr)?, e.zeros(nr)?, e.zeros(nr)?);
e.filter_stats(
p_src, n_vocab, &rowsd, &mut th_d, &mut z_d, &mut mx_d, n_vocab, nr,
sp_temp, sp.top_k, sp.top_p, sp.min_p,
)?;
let idsd = e.htod_u32_v(&ids)?;
let mut outd = e.zeros(nr)?;
e.softmax_gather_filtered(
p_src, n_vocab, &idsd, &rowsd, &th_d, &z_d, &mut outd, n_vocab, nr,
sp_temp,
)?;
let outv = e.dtoh(&outd)?;
let (thv, zv, mxv) = (e.dtoh(&th_d)?, e.dtoh(&z_d)?, e.dtoh(&mx_d)?);
let mut oi = 0usize;
for j in 0..k_round {
if j > 0 || base == 1 {
pj[j] = outv[oi];
oi += 1;
}
}
col_stats = (0..nr).map(|i| (mxv[i], thv[i], zv[i])).collect();
}
if base == 0 {
let lc: &CudaSlice<f32> = if pen_on {
if col_buf.is_none() {
col_buf = Some(e.zeros(n_vocab)?);
}
let cb = col_buf.as_mut().unwrap();
e.copy_into(
cb,
0,
last_col_logits
.as_ref()
.expect("sampled: last_col_logits unset"),
n_vocab,
)?;
let h = pen_hist_d.as_ref().unwrap();
let nh = h.len();
e.penalize_logits(
cb,
h,
nh,
sp.penalty_repeat,
sp.penalty_freq,
sp.penalty_present,
n_vocab,
)?;
col_buf.as_ref().unwrap()
} else {
last_col_logits
.as_ref()
.expect("sampled: last_col_logits unset")
};
let rows0 = e.htod_i32(&[0])?;
let (mut th_d, mut z_d, mut mx_d) = (e.zeros(1)?, e.zeros(1)?, e.zeros(1)?);
e.filter_stats(
lc, n_vocab, &rows0, &mut th_d, &mut z_d, &mut mx_d, n_vocab, 1,
sp_temp, sp.top_k, sp.top_p, sp.min_p,
)?;
let idsd = e.htod_u32_v(&[draft[0]])?;
let mut outd = e.zeros(1)?;
e.softmax_gather_filtered(
lc, n_vocab, &idsd, &rows0, &th_d, &z_d, &mut outd, n_vocab, 1, sp_temp,
)?;
pj[0] = e.dtoh(&outd)?[0];
last_col_stats =
Some((e.dtoh(&mx_d)?[0], e.dtoh(&th_d)?[0], e.dtoh(&z_d)?[0]));
}
}
let q_bufs: &[CudaSlice<f32>] = if dctx.graph_s.is_some() {
&dctx.q_slots
} else {
&draft_logits
};
let mut n_acc = 0usize;
for j in 0..k_round {
let (qmx, qth, qz) = draft_stats[j];
let idsd = e.htod_u32_v(&[draft_idx[j]])?;
let rowsd = e.htod_i32(&[0])?;
let thd = e.htod(&[qth])?;
let zd = e.htod(&[qz])?;
let _ = qmx;
let mut outd = e.zeros(1)?;
e.softmax_gather_filtered(
&q_bufs[j], d_vocab, &idsd, &rowsd, &thd, &zd, &mut outd, d_vocab, 1,
sp_temp,
)?;
let qj = e.dtoh(&outd)?[0];
let u = host_u01(sp_seed, uctr);
uctr += 1;
let accept = (u as f64) * (qj as f64) < pj[j] as f64;
if skey_probe() && qj == 0.0 {
eprintln!(
"[skey] EXACTNESS q=0 round={round} j={j} draft_tok={} \
draft_idx={} p={:e} u={u} accepted={} th_z={:?}",
draft[j], draft_idx[j], pj[j], accept as u8, draft_stats[j],
);
}
if accept {
n_acc += 1;
} else {
break;
}
}
let bonus = if n_acc == k_round {
let col = base + k_round - 1;
let cb = col_buf.as_mut().unwrap();
e.copy_view_into(
cb,
0,
&tlogits_d.slice(col * n_vocab..(col + 1) * n_vocab),
n_vocab,
)?;
if pen_on {
let h = pen_hist_d.as_ref().unwrap();
let nh = h.len();
e.penalize_logits(
cb,
h,
nh,
sp.penalty_repeat,
sp.penalty_freq,
sp.penalty_present,
n_vocab,
)?;
}
if perturb_buf.is_none() {
perturb_buf = Some(e.zeros(d_vocab.max(n_vocab))?);
}
let (mx, th) = {
let rows0 = e.htod_i32(&[0])?;
let (mut th_d, mut z_d, mut mx_d) = (e.zeros(1)?, e.zeros(1)?, e.zeros(1)?);
let cb0 = col_buf.as_ref().unwrap();
e.filter_stats(
cb0, n_vocab, &rows0, &mut th_d, &mut z_d, &mut mx_d, n_vocab, 1,
sp_temp, sp.top_k, sp.top_p, sp.min_p,
)?;
(e.dtoh(&mx_d)?[0], e.dtoh(&th_d)?[0])
};
let pb = perturb_buf.as_mut().unwrap();
let cb2 = col_buf.as_ref().unwrap();
e.gumbel_perturb_filtered(cb2, pb, n_vocab, sp_seed, sctr, sp_temp, mx, th)?;
sctr += 1;
let td = e.argmax_token_device(pb, n_vocab)?;
e.dtoh_u32_one(&td)?
} else {
let cb = col_buf.as_mut().unwrap();
if n_acc > 0 || base == 1 {
let col = base + n_acc - 1;
e.copy_view_into(
cb,
0,
&tlogits_d.slice(col * n_vocab..(col + 1) * n_vocab),
n_vocab,
)?;
} else {
let lc = last_col_logits.as_ref().unwrap();
e.copy_into(cb, 0, lc, n_vocab)?;
}
if pen_on {
let h = pen_hist_d.as_ref().unwrap();
let nh = h.len();
e.penalize_logits(
cb,
h,
nh,
sp.penalty_repeat,
sp.penalty_freq,
sp.penalty_present,
n_vocab,
)?;
}
let cb2 = col_buf.as_ref().unwrap();
let sc = sctr;
sctr += 1;
let p_stats = if n_acc > 0 || base == 1 {
let gi = if base == 1 { n_acc } else { n_acc - 1 };
col_stats.get(gi).copied().unwrap_or_else(|| {
(0.0, 0.0, 1.0) })
} else {
last_col_stats.expect("sampled: last_col_stats unset at reject")
};
let q_stats = draft_stats[n_acc];
if let Some(map) = &d2t_dev {
if q_full_buf.is_none() {
q_full_buf = Some(e.zeros(n_vocab)?);
}
let qf = q_full_buf.as_mut().unwrap();
e.scatter_trim_logits(&q_bufs[n_acc], map, qf, d_vocab, n_vocab)?;
let qf2 = q_full_buf.as_ref().unwrap();
e.residual_sample_filtered(
cb2,
Some(qf2),
n_vocab,
sp_temp,
sp_seed,
sc,
p_stats,
q_stats,
&mut sample_tok,
)?;
} else {
e.residual_sample_filtered(
cb2,
Some(&q_bufs[n_acc]),
n_vocab,
sp_temp,
sp_seed,
sc,
p_stats,
q_stats,
&mut sample_tok,
)?;
}
e.dtoh_u32(&sample_tok)?[0]
};
(n_acc, bonus)
};
let (n_acc, bonus) = match constraint.as_deref_mut() {
None => (n_acc, bonus),
Some(c) => {
fn ce(e2: String) -> Box<dyn std::error::Error> {
format!("constraint: {e2}").into()
}
let mut na = n_acc;
let mut cut = false;
for (j, &d) in draft.iter().enumerate().take(n_acc) {
if c.is_allowed(d).map_err(ce)? {
c.consume(d).map_err(ce)?;
} else {
na = j;
cut = true;
dm_cut_tokens += n_acc - j;
break;
}
}
if cut {
dm_cuts += 1;
}
let mut bo = bonus;
if cut || !c.is_allowed(bo).map_err(ce)? {
let mut row = if na == 0 && base == 0 {
init_logits_host
.clone()
.ok_or("constraint: init logits missing (round-0 cut)")?
} else {
e.dtoh_view(
&tlogits_d.slice((base + na - 1) * n_vocab..(base + na) * n_vocab),
)?
};
c.mask_logits(&mut row).map_err(ce)?;
bo = argmax(&row) as u32;
}
c.consume(bo).map_err(ce)?;
(na, bo)
}
};
let mut successor_valid = false;
if let Some((q_proxy, expected_d2)) = rejected_probe {
let v_n = n_acc == 1 && bonus == expected_d2;
eprintln!(
"[opti-controller] shadow q={q_proxy:.6} admitted=false v_n={v_n} \
expected_d2={expected_d2} n_acc={n_acc} bonus={bonus}",
);
}
if let Some(successor) = successor_attempt.as_ref() {
successor_valid = n_acc == 1 && bonus == successor.verify_tokens[0];
let generation = successor.generation;
let q_proxy = successor.q_proxy;
let expected_pending = successor.verify_tokens[0];
let resolution_ms = successor.issued_at.elapsed().as_secs_f64() * 1e3;
let fork = opti_fork
.as_mut()
.ok_or("optipipe successor resolution lost fork state")?;
fork.finish_actual_reconcile(e, &mut *cache, &snap, n_acc, base, successor_valid)?;
if successor_valid {
OPTI_FORK_HITS.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
} else {
OPTI_FORK_MISSES.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
OPTI_RECONCILES.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
OPTI_WASTED_DRAFT_TOKENS.fetch_add(2, std::sync::atomic::Ordering::Relaxed);
}
let breaker_tripped = fork
.controller
.as_mut()
.expect("controller policy")
.resolve(successor_valid);
if breaker_tripped {
OPTI_BREAKER_TRIPS.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
}
eprintln!(
"[opti-controller] resolve generation={} hit={} q={q_proxy:.6} \
expected_pending={expected_pending} n_acc={n_acc} bonus={bonus} \
resolution_ms={resolution_ms:.3} reconcile={} breaker={}",
generation.id, successor_valid, !successor_valid, breaker_tripped,
);
if !successor_valid {
let mut successor = successor_attempt
.take()
.expect("controller successor disappeared on miss");
successor.settle();
fork.retire(generation)?;
}
}
total_drafted += k_round;
total_accepted += n_acc;
if let Some(t) = sess_telem {
t.record_round(k_round, n_acc);
}
if spec_stats {
st_len_hist[k_round] += 1;
for j in 0..k_round {
st_drafted[j] += 1;
}
for j in 0..n_acc {
st_accepted[j] += 1;
}
if n_acc == k_round {
st_full += 1;
}
}
if debug_spec {
eprintln!(
"[R{round}] pos={pos} out_len={} last_tok={last_token} draft={draft:?} n_acc={n_acc} bonus={bonus} t_pred0={}",
out.len(),
debug_t_pred0(sampled, base, last_pred, &preds)
);
}
let commit_started = std::time::Instant::now();
for j in 0..n_acc {
if !session_mode && out.len() >= max_new {
break;
}
out.push(draft[j]);
}
if pen_on {
pen_hist.extend_from_slice(&draft[0..n_acc]);
pen_hist.push(bonus);
}
let bonus_emitted = session_mode || out.len() < max_new;
if bonus_emitted {
out.push(bonus);
}
last_token = bonus;
if n_acc == k_round && !spec_replay {
let mut vh_seed = e.zeros(n_embd)?;
e.copy_view_into(
&mut vh_seed,
0,
&vx.slice((t_v - 1) * n_embd..t_v * n_embd),
n_embd,
)?;
if refresh {
scratch.set_len(e, pos)?;
let mut vxs = e.zeros(t_v * n_embd)?;
e.copy_into(&mut vxs, 0, &fill_prev, n_embd)?;
if t_v > 1 {
e.copy_view_into(
&mut vxs,
n_embd,
&vx.slice(0..(t_v - 1) * n_embd),
(t_v - 1) * n_embd,
)?;
}
self.mtp_kv_fill(e, mtp, &verify_tokens, &vxs, pos, &mut *scratch, embd_dev)?;
} else {
scratch.set_len(e, pos + base + k_round - 1)?;
let mut hp = e.zeros(n_embd)?;
if t_v >= 2 {
e.copy_view_into(
&mut hp,
0,
&vx.slice((t_v - 2) * n_embd..(t_v - 1) * n_embd),
n_embd,
)?;
} else {
e.copy_into(&mut hp, 0, &fill_prev, n_embd)?;
}
self.mtp_kv_fill(
e,
mtp,
&[draft[k_round - 1]],
&hp,
pos + base + k_round - 1,
&mut *scratch,
embd_dev,
)?;
}
if !devacc_seeded {
e.copy_into(&mut h_seed_buf, 0, &vh_seed, n_embd)?;
e.copy_into(&mut fill_prev, 0, &vh_seed, n_embd)?;
}
pending = Some(bonus);
if debug_spec {
eprintln!(" -> FULL ACCEPT (bonus pending, prev-h seed)");
}
} else if !spec_replay && base + n_acc >= 1 {
let j = base + n_acc;
self.commit_verified_prefix(
e,
&mut *cache,
&snap,
ckpt.as_ref().unwrap(),
j,
devacc_seeded,
if devacc_seeded {
devacc_acc.as_ref().map(|a| (a, base, t_v))
} else {
None
},
)?;
let mut seed = e.zeros(n_embd)?;
e.copy_view_into(
&mut seed,
0,
&vx.slice((j - 1) * n_embd..j * n_embd),
n_embd,
)?;
if refresh {
scratch.set_len(e, pos)?;
let mut vxs = e.zeros(j * n_embd)?;
e.copy_into(&mut vxs, 0, &fill_prev, n_embd)?;
if j > 1 {
e.copy_view_into(
&mut vxs,
n_embd,
&vx.slice(0..(j - 1) * n_embd),
(j - 1) * n_embd,
)?;
}
self.mtp_kv_fill(
e,
mtp,
&verify_tokens[0..j],
&vxs,
pos,
&mut *scratch,
embd_dev,
)?;
} else {
scratch.set_len(e, pos + j)?;
}
if !devacc_seeded {
e.copy_into(&mut h_seed_buf, 0, &seed, n_embd)?;
e.copy_into(&mut fill_prev, 0, &seed, n_embd)?;
}
pending = Some(bonus);
if debug_spec {
eprintln!(" -> PARTIAL(replay-free j={j}, bonus pending, prev-h seed)");
}
} else if !spec_replay {
cache.rollback(e, &snap, 0)?;
scratch.set_len(e, pos)?;
e.copy_into(&mut h_seed_buf, 0, &fill_prev, n_embd)?;
pending = Some(bonus);
if debug_spec {
eprintln!(" -> ZERO-ROUND FOLD (bonus pending, fill_prev seed)");
}
} else {
cache.rollback(e, &snap, 0)?; let mut replay: Vec<u32> = Vec::with_capacity(base + n_acc + 1);
if let Some(b) = pending.take() {
replay.push(b);
}
replay.extend_from_slice(&draft[0..n_acc]);
replay.push(bonus);
let (rl_d, rx) = if self.qwen35_serving_class() {
let mut logits = Vec::with_capacity(replay.len() * n_vocab);
let mut hidden = e.uninit(replay.len() * n_embd)?;
for (row, &token) in replay.iter().enumerate() {
let (row_logits, row_hidden) =
self.spec_target_step_h(e, token, &mut *cache)?;
logits.extend_from_slice(&row_logits);
e.dtod_copy_into(&row_hidden, &mut hidden, row * n_embd)?;
}
(e.htod(&logits)?, hidden)
} else {
self.decode_step_t_core(e, &replay, pos, &mut *cache, embd_dev, None)?
};
e.argmax_token_device_col(&rl_d, replay.len() - 1, n_vocab, &mut preds_d, 0)?;
last_pred = e.dtoh_u32(&preds_d)?[0];
if sampled {
let lr0 = replay.len();
let lc = last_col_logits
.as_mut()
.expect("sampled: last_col_logits unset");
e.copy_view_into(
lc,
0,
&rl_d.slice((lr0 - 1) * n_vocab..lr0 * n_vocab),
n_vocab,
)?;
}
let lr = replay.len();
if lr >= 2 {
e.copy_view_into(
&mut h_seed_buf,
0,
&rx.slice((lr - 2) * n_embd..(lr - 1) * n_embd),
n_embd,
)?;
} else {
e.copy_into(&mut h_seed_buf, 0, &fill_prev, n_embd)?;
}
let mut rh_last = e.zeros(n_embd)?;
e.copy_view_into(
&mut rh_last,
0,
&rx.slice((lr - 1) * n_embd..lr * n_embd),
n_embd,
)?;
e.copy_into(&mut fill_prev, 0, &rh_last, n_embd)?;
if debug_spec {
eprintln!(" -> PARTIAL(replay={replay:?}), next_pred={last_pred}");
}
}
if devacc_seeded {
e.copy_into(&mut fill_prev, 0, &h_seed_buf, n_embd)?;
}
if successor_valid {
let optimistic_scratch_len = successor_attempt
.as_ref()
.expect("valid controller successor disappeared")
.scratch_len;
scratch.set_len(e, optimistic_scratch_len)?;
}
if let Some(current) = current_opti.take() {
opti_fork
.as_mut()
.ok_or("optipipe current retirement lost fork state")?
.retire(current.generation)?;
}
if successor_valid {
let successor = successor_attempt
.take()
.expect("valid controller successor disappeared before promotion");
let generation = successor.generation;
opti_fork
.as_mut()
.ok_or("optipipe successor promotion lost fork state")?
.promote_successor_snapshot(&mut snap, generation);
carried_opti = Some(successor);
}
if anatomy_on {
e.stream().synchronize()?;
ph_commit += commit_started.elapsed().as_secs_f64();
}
if adapt {
let fl_now = floor_at(cache.pos);
kc = (n_acc + 1).clamp(fl_now.min(k_cap), k_cap);
}
ph_mark(&mut ph_rest, phase_on);
if let Some(p) = pipe {
p.accept_end(round);
}
drop(pipe_accept);
round += 1;
keep_going = flush_commit(&mut on_commit, &out, &mut flushed);
}
if let Some(mut ticket) = carried_opti.take() {
opti_fork
.as_mut()
.ok_or("optipipe tail drain lost fork state")?
.cancel_controller_ticket(e, &mut *cache, &mut *scratch, &snap, &mut ticket)?;
}
let _ = flush_commit(&mut on_commit, &out, &mut flushed);
if spec_stats {
let per_slot: Vec<String> = (0..k)
.map(|j| {
if st_drafted[j] > 0 {
format!(
"{}/{}={:.3}",
st_accepted[j],
st_drafted[j],
st_accepted[j] as f64 / st_drafted[j] as f64
)
} else {
"0/0".into()
}
})
.collect();
let acc = if total_drafted > 0 {
total_accepted as f64 / total_drafted as f64
} else {
0.0
};
eprintln!(
"[spec-stats] rounds={round} full_accept={st_full} len_hist={st_len_hist:?} \
per_slot=[{}] total={total_accepted}/{total_drafted}={acc:.3} \
tok_per_round={:.3}",
per_slot.join(" "),
(total_accepted + round) as f64 / round.max(1) as f64
);
}
if constraint.is_some() {
eprintln!(
"[draft-mask] mask_rounds={dm_rounds} clone_total={:.3}ms \
clone_per_round={:.4}ms gram_cuts={dm_cuts}/{round} cut_tokens={dm_cut_tokens}",
dm_clone_ns as f64 / 1e6,
dm_clone_ns as f64 / 1e6 / dm_rounds.max(1) as f64
);
}
if phase_on {
let tot = ph_draft + ph_verify + ph_wait + ph_rest;
eprintln!(
"[spec-phase] draft={:.1}ms ({:.1}%) verify-issue={:.1}ms ({:.1}%) verify-wait={:.1}ms ({:.1}%) commit-host={:.1}ms ({:.1}%) rounds={round}",
ph_draft * 1e3,
ph_draft / tot * 100.0,
ph_verify * 1e3,
ph_verify / tot * 100.0,
ph_wait * 1e3,
ph_wait / tot * 100.0,
ph_rest * 1e3,
ph_rest / tot * 100.0
);
}
if anatomy_on {
let rounds_f = round.max(1) as f64;
let other = (ph_rest - ph_commit).max(0.0);
eprintln!(
"[spec-anatomy] per-round draft={:.3}ms pp-verify={:.3}ms \
verify-accept={:.3}ms commit-rollback={:.3}ms other={:.3}ms rounds={round}",
ph_draft * 1e3 / rounds_f,
ph_verify * 1e3 / rounds_f,
ph_wait * 1e3 / rounds_f,
ph_commit * 1e3 / rounds_f,
other * 1e3 / rounds_f,
);
}
let _pipe_tail = pipe.map(|p| p.primary());
if let Some(slot) = sess_draft_slot.take() {
*slot = Some(dctx);
}
let t_rounds = t_ent.elapsed();
if let Some((committed, last_h, next_pred_slot, sctr_slot, uctr_slot)) = sess_tail.take() {
*next_pred_slot = Some(last_pred);
let sample_boundary = sampled && constraint.is_none() && spec_sampled_boundary_on();
let mut stashed_pending = false;
if let Some(b) = pending.take() {
if !sampled {
debug_assert_eq!(out.last(), Some(&b), "pending must be the last emitted");
if let Some(slot) = sess_pending_slot.take() {
*slot = Some(b);
}
*next_pred_slot = None;
*last_h = Some(e.clone_dtod(&fill_prev)?);
stashed_pending = true;
} else {
let pos_b = cache.pos;
scratch.set_len(e, pos_b)?;
let (lg_b, hb) = self.spec_target_step_h(e, b, &mut *cache)?;
*next_pred_slot = Some(if sample_boundary {
sample_boundary_token(
e,
&lg_b,
&sp,
&pen_hist,
&mut sctr,
"burst-tail-commit",
)?
} else {
argmax(&lg_b) as u32
});
self.mtp_kv_fill(e, mtp, &[b], &fill_prev, pos_b, &mut *scratch, embd_dev)?;
*last_h = Some(hb);
}
} else {
*last_h = Some(e.clone_dtod(&fill_prev)?);
if sample_boundary {
match last_col_logits.as_ref() {
Some(lc) => {
*next_pred_slot = Some(sample_boundary_token_dev(
e,
lc,
n_vocab,
&sp,
&pen_hist,
&mut sctr,
"burst-tail-nopending",
)?);
}
None => eprintln!(
"[spec-boundary] sampled tail kept the ARGMAX boundary token \
(reason: no retained boundary logits row)"
),
}
}
}
*sctr_slot = sctr;
*uctr_slot = uctr;
committed.extend_from_slice(prompt);
if let Some(cb) = carried_pending {
committed.push(cb);
}
if stashed_pending {
let emitted = out.len().saturating_sub(1);
committed.extend_from_slice(&out[..emitted]);
} else {
committed.extend_from_slice(&out); }
debug_assert_eq!(
cache.pos,
committed.len(),
"session invariant: cache rows == committed tokens"
);
if setup_trace {
e.stream().synchronize()?; let t_tail = t_ent.elapsed();
eprintln!(
"[spec-setup] init={:.2}ms cap={:.2}ms fill={:.2}ms rounds={:.2}ms tail={:.2}ms total={:.2}ms out={} cont={}",
t_init.as_secs_f64() * 1e3,
(t_cap - t_init).as_secs_f64() * 1e3,
(t_fill - t_cap).as_secs_f64() * 1e3,
(t_rounds - t_fill).as_secs_f64() * 1e3,
(t_tail - t_rounds).as_secs_f64() * 1e3,
t_tail.as_secs_f64() * 1e3,
out.len(),
continuation
);
}
return Ok((out, total_drafted, total_accepted));
}
out.truncate(max_new);
Ok((out, total_drafted, total_accepted))
}
pub fn extract_dspark_anchors(
&self,
e: &Engine,
tokens: &[u32],
anchor_positions: &[usize],
gamma: usize,
top_k: usize,
chunk: usize,
temperature: f32,
) -> Result<Vec<DsparkAnchorRecord>, Box<dyn std::error::Error>> {
if tokens.len() < gamma + 2 || gamma == 0 || chunk < 2 {
return Err("DSpark extraction token tape/gamma/chunk is invalid".into());
}
if anchor_positions.windows(2).any(|pair| pair[0] >= pair[1]) {
return Err("DSpark anchor positions must be sorted and unique".into());
}
for &position in anchor_positions {
if position == 0 || position + gamma >= tokens.len() {
return Err(format!(
"DSpark anchor {position} has no predecessor or cannot cover gamma={gamma} in {} tokens",
tokens.len()
)
.into());
}
}
let n_vocab = self.output.out_features();
let n_embd = self.cfg.n_embd as usize;
let mut cache = crate::pp::new_cache(e, &self.cfg, tokens.len() + gamma + 8)?;
let (embd_qt, embd_rb) = self.embd.qt_and_row_bytes(n_embd);
let embd_gpu = if spec_host_embd() {
None
} else {
Some(
self.embd_gpu
.get_or_init(|| e.upload_u8(&self.embd.raw).expect("embed table upload")),
)
};
let embd_dev = embd_gpu.map(|gpu| (gpu, embd_qt, embd_rb));
struct PendingRecord {
position: usize,
hidden: Option<Vec<f32>>,
tokens: Vec<u32>,
target_top_ids: Vec<Option<Vec<u32>>>,
target_top_logits: Vec<Option<Vec<f32>>>,
target_top_probs: Vec<Option<Vec<f32>>>,
target_tail_probs: Vec<Option<f32>>,
}
let mut pending: Vec<PendingRecord> = anchor_positions
.iter()
.map(|&position| PendingRecord {
position,
hidden: None,
tokens: tokens[position..=position + gamma].to_vec(),
target_top_ids: vec![None; gamma],
target_top_logits: vec![None; gamma],
target_top_probs: vec![None; gamma],
target_tail_probs: vec![None; gamma],
})
.collect();
let mut start = 0usize;
while start < tokens.len() {
let end = (start + chunk).min(tokens.len());
let chunk_tokens = &tokens[start..end];
let (target_logits, hidden_rows) =
self.decode_step_t_core(e, chunk_tokens, start, &mut cache, embd_dev, None)?;
for record in &mut pending {
let hidden_position = record.position - 1;
if hidden_position >= start && hidden_position < end {
let local = hidden_position - start;
record.hidden = Some(
e.dtoh_view(&hidden_rows.slice(local * n_embd..(local + 1) * n_embd))?,
);
}
for slot in 0..gamma {
let target_row = record.position + slot;
if target_row < start || target_row >= end {
continue;
}
let local = target_row - start;
let logits =
e.dtoh_view(&target_logits.slice(local * n_vocab..(local + 1) * n_vocab))?;
let (ids, top_logits, probs, tail) =
dspark_sparse_softmax_topk(&logits, top_k, temperature)?;
record.target_top_ids[slot] = Some(ids);
record.target_top_logits[slot] = Some(top_logits);
record.target_top_probs[slot] = Some(probs);
record.target_tail_probs[slot] = Some(tail);
}
}
start = end;
}
pending
.into_iter()
.map(|record| {
let hidden = record
.hidden
.ok_or_else(|| format!("missing DSpark hidden at {}", record.position))?;
let target_top_ids =
flatten_dspark_rows(record.target_top_ids, record.position, "target ids")?;
let target_top_logits = flatten_dspark_rows(
record.target_top_logits,
record.position,
"target logits",
)?;
let target_top_probs =
flatten_dspark_rows(record.target_top_probs, record.position, "target probs")?;
let target_tail_probs = record
.target_tail_probs
.into_iter()
.enumerate()
.map(|(slot, value)| {
value.ok_or_else(|| {
format!("missing DSpark tail at {} slot {slot}", record.position)
})
})
.collect::<Result<Vec<_>, _>>()?;
Ok(DsparkAnchorRecord {
position: record.position,
hidden,
tokens: record.tokens,
target_top_ids,
target_top_logits,
target_top_probs,
target_tail_probs,
})
})
.collect()
}
pub fn replay_acceptance(
&self,
e: &Engine,
tokens: &[u32],
k: usize,
stride: usize,
chunk: usize,
mut hdump: Option<&mut std::fs::File>,
) -> Result<(Vec<(usize, Vec<u32>, Vec<u32>)>, Vec<u32>), Box<dyn std::error::Error>> {
assert!(k >= 1 && stride >= 1 && chunk >= 2);
let mtp = self
.mtp
.as_ref()
.expect("replay_acceptance requires an MTP head");
let n_vocab = self.output.out_features();
let d_vocab = mtp
.shared_head_head
.as_ref()
.unwrap_or(&self.output)
.out_features();
let n_embd = self.cfg.n_embd as usize;
let t_total = tokens.len();
assert!(t_total >= 8, "corpus too short ({t_total} tokens)");
let mut cache = crate::pp::new_cache(e, &self.cfg, t_total + k + 8)?;
let mut scratch = MtpScratch::new(
e,
&self.cfg,
t_total + k + 8,
self.mtp.as_ref().and_then(|m| m.geom.as_ref()),
)?;
let (embd_qt, embd_rb) = self.embd.qt_and_row_bytes(n_embd);
let embd_gpu = if spec_host_embd() {
None
} else {
Some(
self.embd_gpu
.get_or_init(|| e.upload_u8(&self.embd.raw).expect("embed table upload")),
)
};
let embd_dev = embd_gpu.map(|g| (g, embd_qt, embd_rb));
let mut bg: Vec<u32> = vec![0; t_total + 1];
let mut rows: Vec<(usize, Vec<u32>, Vec<u32>)> = Vec::new();
let mut prev_last_h = e.zeros(n_embd)?; let mut seed_buf = e.zeros(n_embd)?;
let mut preds_d = e.alloc_u32_zeroed(chunk)?;
let nll_on = std::env::var("MEMRA_REPLAY_NLL").as_deref() == Ok("1");
let (mut nll_sum, mut nll_cnt) = (0f64, 0u64);
let mut s = 0usize;
while s < t_total {
let cend = (s + chunk).min(t_total);
let tc = cend - s;
let ch = &tokens[s..cend];
let (tl_d, vx) = self.decode_step_t_core(e, ch, s, &mut cache, embd_dev, None)?;
for j in 0..tc {
e.argmax_token_device_col(&tl_d, j, n_vocab, &mut preds_d, j)?;
}
let preds = e.dtoh_u32(&preds_d)?;
for j in 0..tc {
bg[s + j + 1] = preds[j];
}
if nll_on {
let jmax = if cend < t_total { tc } else { tc - 1 }; if jmax > 0 {
let ids: Vec<u32> = (0..jmax).map(|j| tokens[s + j + 1]).collect();
let rows: Vec<i32> = (0..jmax as i32).collect();
let idsd = e.htod_u32_v(&ids)?;
let rowsd = e.htod_i32(&rows)?;
let mut outd = e.zeros(jmax)?;
e.softmax_gather(&tl_d, n_vocab, &idsd, &rowsd, &mut outd, n_vocab, jmax, 1.0)?;
for pr in e.dtoh(&outd)? {
nll_sum += -((pr.max(1e-30)) as f64).ln();
nll_cnt += 1;
}
}
}
if let Some(f) = hdump.as_deref_mut() {
use std::io::Write;
let host: Vec<f32> = e.dtoh(&vx)?;
let mut bytes = Vec::with_capacity(tc * n_embd * 2);
for v in &host[..tc * n_embd] {
let b = v.to_bits();
let r = b.wrapping_add(0x7FFF + ((b >> 16) & 1));
bytes.extend_from_slice(&((r >> 16) as u16).to_le_bytes());
}
f.write_all(&bytes)?;
}
let chainless = stride > t_total;
if chainless {
e.copy_view_into(
&mut prev_last_h,
0,
&vx.slice((tc - 1) * n_embd..tc * n_embd),
n_embd,
)?;
s = cend;
continue;
}
let mut vxs = e.zeros(tc * n_embd)?;
e.copy_into(&mut vxs, 0, &prev_last_h, n_embd)?;
if tc > 1 {
e.copy_view_into(
&mut vxs,
n_embd,
&vx.slice(0..(tc - 1) * n_embd),
(tc - 1) * n_embd,
)?;
}
scratch.set_len(e, s)?;
self.mtp_kv_fill(e, mtp, ch, &vxs, s, &mut scratch, embd_dev)?;
let ps: Vec<usize> = (s..cend)
.filter(|p| *p >= 1 && *p % stride == 0 && *p + k <= t_total)
.collect();
for &p in ps.iter().rev() {
scratch.set_len(e, p)?;
if p == s {
e.copy_into(&mut seed_buf, 0, &prev_last_h, n_embd)?;
} else {
e.copy_view_into(
&mut seed_buf,
0,
&vx.slice((p - 1 - s) * n_embd..(p - s) * n_embd),
n_embd,
)?;
}
let mut e_tok = tokens[p];
let mut d_seed = e.clone_dtod(&seed_buf)?;
let mut drafts: Vec<u32> = Vec::with_capacity(k);
for j in 0..k {
let (dl_d, h_nextn) = self.mtp_head_forward_dev(
e,
mtp,
e_tok,
&d_seed,
&mut scratch,
p + 1 + j,
embd_dev,
None, )?;
let tok_d = e.argmax_token_device(&dl_d, d_vocab)?;
let idx = e.dtoh_u32_one(&tok_d)?;
let d = match &mtp.d2t {
Some(map) => map[idx as usize],
None => idx,
};
drafts.push(d);
e_tok = d;
d_seed = h_nextn;
}
rows.push((p, drafts, Vec::new()));
}
scratch.set_len(e, s)?;
self.mtp_kv_fill(e, mtp, ch, &vxs, s, &mut scratch, embd_dev)?;
e.copy_view_into(
&mut prev_last_h,
0,
&vx.slice((tc - 1) * n_embd..tc * n_embd),
n_embd,
)?;
s = cend;
}
for (p, drafts, targets) in rows.iter_mut() {
for j in 0..drafts.len() {
targets.push(bg[*p + 1 + j]);
}
}
rows.sort_by_key(|r| r.0);
if nll_cnt > 0 {
let mean = nll_sum / nll_cnt as f64;
println!(
"[replay-nll] tokens={nll_cnt} nll/token={mean:.5} ppl={:.4}",
mean.exp()
);
}
Ok((rows, bg))
}
}
#[cfg(test)]
mod dspark_sparse_tests {
use super::dspark_sparse_softmax_topk;
#[test]
fn topk_keeps_full_softmax_mass_and_stable_ties() {
let logits = [1.0f32, 3.0, 3.0, -2.0];
let (ids, top_logits, probs, tail) = dspark_sparse_softmax_topk(&logits, 2, 1.0).unwrap();
assert_eq!(ids, vec![1, 2]);
assert_eq!(top_logits, vec![3.0, 3.0]);
let denominator = logits.iter().map(|value| (value - 3.0).exp()).sum::<f32>();
let expected = 1.0 / denominator;
assert!((probs[0] - expected).abs() < 1.0e-6);
assert!((probs[1] - expected).abs() < 1.0e-6);
assert!((tail - (1.0 - 2.0 * expected)).abs() < 1.0e-6);
assert!((probs.iter().sum::<f32>() + tail - 1.0).abs() < 1.0e-6);
}
}
#[cfg(test)]
mod spec_replay_env_tests {
use super::spec_replay_env_on;
#[test]
fn replay_requires_literal_one() {
assert!(!spec_replay_env_on(None));
assert!(!spec_replay_env_on(Some("")));
assert!(!spec_replay_env_on(Some("0")));
assert!(!spec_replay_env_on(Some("true")));
assert!(!spec_replay_env_on(Some("2")));
assert!(spec_replay_env_on(Some("1")));
}
}
#[cfg(test)]
mod telem_tests {
use super::{SPEC_TELEM_POS, SpecTelemetry, SpecTelemetryCounters};
#[test]
fn synthetic_accept_masks_produce_tau_and_position_histogram() {
let counters = SpecTelemetryCounters::default();
for mask in [
[true, true, true],
[true, true, false],
[true, false, false],
[false, false, false],
] {
let accepted = mask.iter().take_while(|&&value| value).count();
counters.record_round(mask.len(), accepted);
}
let snapshot = counters.snapshot();
assert_eq!(
(snapshot.rounds, snapshot.drafted, snapshot.accepted),
(4, 12, 6)
);
assert_eq!(&snapshot.pos_drafted[..3], &[4, 4, 4]);
assert_eq!(&snapshot.pos_accepted[..3], &[3, 2, 1]);
assert_eq!(snapshot.tau(), 1.5);
assert_eq!(snapshot.pos_drafted[3..], [0; SPEC_TELEM_POS - 3]);
assert_eq!(snapshot.pos_accepted[3..], [0; SPEC_TELEM_POS - 3]);
}
#[test]
fn delta_isolates_burst_contribution() {
let mut t = SpecTelemetry::default();
for (kr, na) in [(3usize, 3usize), (3, 1)] {
t.rounds += 1;
t.drafted += kr as u64;
t.accepted += na as u64;
for j in 0..kr {
t.pos_drafted[j] += 1;
}
for j in 0..na {
t.pos_accepted[j] += 1;
}
}
let before = t;
t.rounds += 1;
t.drafted += 3;
t.accepted += 2;
for j in 0..3 {
t.pos_drafted[j] += 1;
}
for j in 0..2 {
t.pos_accepted[j] += 1;
}
let d = t.delta_since(&before);
assert_eq!((d.rounds, d.drafted, d.accepted), (1, 3, 2));
assert_eq!(&d.pos_drafted[..3], &[1, 1, 1]);
assert_eq!(&d.pos_accepted[..3], &[1, 1, 0]);
assert_eq!(d.pos_drafted[3..], [0; SPEC_TELEM_POS - 3]);
}
#[test]
fn merge_accumulates_fieldwise() {
let mut agg = SpecTelemetry::default();
let mut d1 = SpecTelemetry {
rounds: 2,
drafted: 6,
accepted: 4,
..Default::default()
};
d1.pos_drafted[0] = 2;
d1.pos_accepted[0] = 2;
let mut d2 = SpecTelemetry {
rounds: 1,
drafted: 3,
accepted: 1,
..Default::default()
};
d2.pos_drafted[0] = 1;
d2.pos_accepted[0] = 1;
d2.pos_drafted[1] = 1;
agg.merge(&d1);
agg.merge(&d2);
assert_eq!((agg.rounds, agg.drafted, agg.accepted), (3, 9, 5));
assert_eq!(agg.pos_drafted[0], 3);
assert_eq!(agg.pos_accepted[0], 3);
assert_eq!(agg.pos_drafted[1], 1);
assert_eq!(agg.pos_accepted[1], 0);
}
#[test]
fn delta_saturates_never_wraps() {
let small = SpecTelemetry {
rounds: 1,
drafted: 2,
accepted: 1,
..Default::default()
};
let big = SpecTelemetry {
rounds: 5,
drafted: 15,
accepted: 9,
..Default::default()
};
let d = small.delta_since(&big);
assert_eq!((d.rounds, d.drafted, d.accepted), (0, 0, 0));
}
}
#[cfg(test)]
mod opti_fork_tests {
use super::{
OptiControllerPolicy, OptiForkAction, OptiForkGateMode, OptiForkGenerationTracker,
};
#[test]
fn controller_threshold_and_three_miss_breaker_are_exact() {
let mut policy = OptiControllerPolicy {
threshold: 0.7,
consecutive_misses: 0,
breaker_tripped: false,
};
assert!(!policy.admit(0.699_999));
assert!(policy.admit(0.7));
assert!(!policy.resolve(false));
assert!(!policy.resolve(false));
assert!(policy.resolve(false));
assert!(policy.breaker_tripped);
assert!(!policy.admit(1.0));
assert!(
!policy.resolve(true),
"a resolved hit cannot re-arm a tripped request"
);
assert!(policy.breaker_tripped);
}
#[test]
fn zero_threshold_is_the_true_unconditional_measurement_arm() {
let mut policy = OptiControllerPolicy {
threshold: 0.0,
consecutive_misses: 0,
breaker_tripped: false,
};
for _ in 0..16 {
assert!(policy.admit(0.0));
assert!(!policy.resolve(false));
}
for invalid in [f32::NAN, f32::INFINITY, -0.01, 1.01] {
assert!(
!policy.admit(invalid),
"invalid q proxy must fail closed: {invalid}"
);
}
assert!(!policy.breaker_tripped);
assert_eq!(policy.consecutive_misses, 0);
}
#[test]
fn alternating_mode_flips_by_generation_not_round_parity() {
assert_eq!(OptiForkGateMode::Alternate.action(0), OptiForkAction::Hit);
assert_eq!(OptiForkGateMode::Alternate.action(1), OptiForkAction::Miss);
assert_eq!(OptiForkGateMode::Alternate.action(8), OptiForkAction::Hit);
assert_eq!(OptiForkGateMode::Alternate.action(9), OptiForkAction::Miss);
}
#[test]
fn live_generation_cannot_be_overwritten() {
let mut tracker = OptiForkGenerationTracker::default();
let g0 = tracker.reserve().unwrap();
let g1 = tracker.reserve().unwrap();
let err = tracker.reserve().unwrap_err().to_string();
assert!(
err.contains("still owns generation 0"),
"unexpected error: {err}"
);
tracker.retire(g0).unwrap();
let g2 = tracker.reserve().unwrap();
assert_eq!((g2.id, g2.slot), (2, 0));
tracker.retire(g1).unwrap();
tracker.retire(g2).unwrap();
}
#[test]
fn teardown_rejects_a_stale_generation_tag() {
let mut tracker = OptiForkGenerationTracker::default();
let g0 = tracker.reserve().unwrap();
tracker.retire(g0).unwrap();
let err = tracker.retire(g0).unwrap_err().to_string();
assert!(err.contains("teardown mismatch"), "unexpected error: {err}");
}
}
#[cfg(test)]
mod draft_graph_fallback_tests {
use super::DraftGraphFallback;
#[test]
fn flip_is_loud_once_and_memoized_after() {
let mut f = DraftGraphFallback::default();
let line = f
.mark_greedy("out of memory")
.expect("first flip must return the warn line");
assert!(
line.contains("WARN"),
"flip line must be warn-level: {line}"
);
assert!(
line.contains("out of memory"),
"flip line must carry the reason: {line}"
);
assert!(f.greedy_failed());
assert!(f.mark_greedy("out of memory").is_none());
assert!(f.greedy_failed());
assert!(!f.sampled_failed());
let line_s = f
.mark_sampled("capture unsupported")
.expect("sampled flip is its own flip");
assert!(
line_s.contains("sampled"),
"sampled flip names itself: {line_s}"
);
assert!(f.mark_sampled("capture unsupported").is_none());
}
#[test]
fn reset_on_resume_clears_flags_and_logs_once() {
let mut f = DraftGraphFallback::default();
assert!(f.reset_on_resume().is_none());
f.mark_greedy("oom").unwrap();
f.mark_sampled("oom").unwrap();
let note = f
.reset_on_resume()
.expect("a set flag must produce the reset note");
assert!(
note.contains("greedy+sampled"),
"note names what was reset: {note}"
);
assert!(
!f.greedy_failed() && !f.sampled_failed(),
"both flags cleared"
);
assert!(f.mark_greedy("oom again").is_some());
let note2 = f.reset_on_resume().expect("greedy-only reset");
assert!(note2.contains("(greedy)"), "single-flag note: {note2}");
}
#[test]
fn shape_change_clears_are_silent() {
let mut f = DraftGraphFallback::default();
f.mark_greedy("oom").unwrap();
f.clear_greedy();
assert!(!f.greedy_failed());
f.mark_sampled("oom").unwrap();
f.clear_sampled();
assert!(!f.sampled_failed());
assert!(f.reset_on_resume().is_none());
}
}
#[cfg(test)]
mod sampled_graph_key_tests {
use super::{SampledGraphKey, debug_t_pred0};
fn legacy_key(k: &SampledGraphKey) -> (u64, u32, usize) {
(k.seed, k.temp_bits, k.k)
}
fn pure_temp_key() -> SampledGraphKey {
SampledGraphKey::new(12345, 1.0, 3, 0, 1.0, 0.0, false)
}
#[test]
fn vendor_filters_change_the_key() {
let parked = pure_temp_key();
let vendor = SampledGraphKey::new(12345, 1.0, 3, 20, 0.95, 0.0, false);
assert_eq!(
legacy_key(&parked),
legacy_key(&vendor),
"pre-fix key collided: this is the bug, and the reason a test asserts on it",
);
assert_ne!(parked, vendor, "post-fix key must separate the two regimes");
assert!(parked.pure_temp());
assert!(!vendor.pure_temp());
}
#[test]
fn every_filter_field_is_keyed() {
let base = pure_temp_key();
for (what, other) in [
(
"top_k",
SampledGraphKey::new(12345, 1.0, 3, 20, 1.0, 0.0, false),
),
(
"top_p",
SampledGraphKey::new(12345, 1.0, 3, 0, 0.95, 0.0, false),
),
(
"min_p",
SampledGraphKey::new(12345, 1.0, 3, 0, 1.0, 0.05, false),
),
(
"penalties",
SampledGraphKey::new(12345, 1.0, 3, 0, 1.0, 0.0, true),
),
] {
assert_ne!(base, other, "{what} must be part of the key");
assert!(!other.pure_temp(), "{what} leaves the pure-temp regime");
assert_eq!(
legacy_key(&base),
legacy_key(&other),
"{what} was invisible to the pre-fix key",
);
}
}
#[test]
fn baked_constants_stay_keyed() {
let base = pure_temp_key();
assert_ne!(
base,
SampledGraphKey::new(999, 1.0, 3, 0, 1.0, 0.0, false),
"seed"
);
assert_ne!(
base,
SampledGraphKey::new(12345, 0.7, 3, 0, 1.0, 0.0, false),
"temp"
);
assert_ne!(
base,
SampledGraphKey::new(12345, 1.0, 4, 0, 1.0, 0.0, false),
"k"
);
assert_eq!(
SampledGraphKey::new(1, 0.7, 3, 0, 1.0, 0.0, false),
SampledGraphKey::new(1, 7.0 / 10.0, 3, 0, 1.0, 0.0, false),
);
}
#[test]
fn seed_alone_still_rekeys_the_draft_graph() {
let parked = pure_temp_key();
let reseeded = SampledGraphKey::new(999, 1.0, 3, 0, 1.0, 0.0, false);
assert_ne!(
parked, reseeded,
"a seed-only change MUST drop the parked sampled graph — the resume predicate's \
decision not to compare seed rests on exactly this",
);
assert!(parked.pure_temp() && reseeded.pure_temp());
}
#[test]
fn equal_keys_agree_on_the_regime() {
let a = SampledGraphKey::new(7, 0.8, 3, 20, 0.95, 0.0, false);
let b = SampledGraphKey::new(7, 0.8, 3, 20, 0.95, 0.0, false);
assert_eq!(a, b);
assert_eq!(a.pure_temp(), b.pure_temp());
assert!(SampledGraphKey::new(7, 0.8, 3, 0, 1.0, 0.0, false).pure_temp());
assert!(SampledGraphKey::new(7, 0.8, 3, 0, 1.5, -1.0, false).pure_temp());
}
#[test]
fn debug_print_survives_the_sampled_arm() {
assert_eq!(debug_t_pred0(true, 1, 4242, &[]), "n/a");
assert_eq!(debug_t_pred0(true, 2, 4242, &[]), "n/a");
assert_eq!(debug_t_pred0(true, 0, 4242, &[]), "4242");
assert_eq!(debug_t_pred0(false, 0, 4242, &[7, 8]), "4242");
assert_eq!(debug_t_pred0(false, 1, 4242, &[7, 8]), "7");
assert_eq!(debug_t_pred0(false, 2, 4242, &[7, 8]), "8");
}
}