use std::collections::BTreeMap;
use std::os::raw::c_void;
use std::path::Path;
use cudarc::driver::{CudaSlice, CudaStream, DevicePtr, DevicePtrMut};
use memra_gguf::dsv4_forward::{
ActQuantVariant, Dsv4Model, FreqsCis, compress_topk_idxs, hc_split_sinkhorn,
precompute_freqs_cis, window_topk_idxs,
};
use crate::dsv4_ffi as k;
use crate::dsv4_ffi::ck;
type Res<T> = Result<T, String>;
fn e<E: std::fmt::Display>(what: &str) -> impl FnOnce(E) -> String + '_ {
move |err| format!("{what}: {err}")
}
#[inline]
fn sigmoid_f32(x: f32) -> f32 {
1.0 / (1.0 + (-x).exp())
}
#[inline]
fn softplus_f32(x: f32) -> f32 {
if x > 20.0 { x } else { x.exp().ln_1p() }
}
pub struct Stage {
pub dev: usize,
pub gpu: memra_runtime::Gpu,
pub layers: Vec<LayerDev>,
pub embed: Option<CudaSlice<u8>>, pub head: Option<CudaSlice<u8>>, pub trunk_norm: Option<CudaSlice<f32>>,
pub hc_head_fn: Option<CudaSlice<f32>>, pub fc_yarn: CudaSlice<f32>, pub fc_plain: CudaSlice<f32>, pub ws: CudaSlice<u8>, pub deq: [CudaSlice<u8>; 3], pub loaded_bytes: u64, pub hc_head_base_dev: Option<CudaSlice<f32>>,
pub hc_head_scale_dev: Option<CudaSlice<f32>>,
}
pub struct CmpDev {
pub ratio: usize,
pub d: usize,
pub latent: usize,
pub overlap: bool,
pub rotate: bool,
pub wkv: CudaSlice<f32>, pub wgate: CudaSlice<f32>, pub norm: CudaSlice<f32>,
pub ape: CudaSlice<f32>, }
pub struct IdxDev {
pub wq_b: DenseBf16, pub weights_proj: DenseBf16, pub wq_b_fp8: Option<Fp8Dense>,
pub weights_proj_fp8: Option<Fp8Dense>,
pub cmp: CmpDev,
pub heads: usize,
pub hd: usize,
pub topk: usize,
}
pub struct Fp8Dense {
pub codes: CudaSlice<u8>, pub scales: CudaSlice<f32>, pub sc_cols: usize, pub rows: usize,
pub cols: usize,
}
pub enum DenseBf16 {
Dev(CudaSlice<u8>),
Host(Vec<u8>),
}
impl DenseBf16 {
pub fn dev(&self) -> &CudaSlice<u8> {
match self {
DenseBf16::Dev(d) => d,
DenseBf16::Host(_) => unreachable!(
"bf16 dense slab is host-staged (fp8 dense arm): this consumer must \
ride the fp8 twins (dwsel) or the staged prefill view"
),
}
}
fn staged(&self, stream: &std::sync::Arc<CudaStream>) -> Res<DenseView<'_>> {
Ok(match self {
DenseBf16::Dev(d) => DenseView::Res(d),
DenseBf16::Host(b) => DenseView::Tmp(upload_u8(stream, b)?),
})
}
}
pub enum DenseView<'a> {
Res(&'a CudaSlice<u8>),
Tmp(CudaSlice<u8>),
}
impl DenseView<'_> {
fn slab(&self) -> &CudaSlice<u8> {
match self {
DenseView::Res(d) => d,
DenseView::Tmp(d) => d,
}
}
}
#[derive(Clone, Copy)]
pub enum DW {
Bf16(*const c_void),
Fp8 {
codes: *const c_void,
scales: *const f32,
sc_cols: i32,
},
}
impl DW {
fn offset_rows(self, rows_off: usize, cols: usize) -> DW {
match self {
DW::Bf16(p) => DW::Bf16((p as usize + rows_off * cols * 2) as *const c_void),
DW::Fp8 {
codes,
scales,
sc_cols,
} => {
assert_eq!(
rows_off % 128,
0,
"fp8 dense arm: grouped row offset {rows_off} not on the 128-row \
scale-grid boundary"
);
DW::Fp8 {
codes: (codes as usize + rows_off * cols) as *const c_void,
scales: unsafe { scales.add((rows_off / 128) * sc_cols as usize) },
sc_cols,
}
}
}
}
}
fn dwsel(
active: bool,
stream: &cudarc::driver::CudaStream,
w_bf16: &DenseBf16,
fp8: &Option<Fp8Dense>,
) -> DW {
match fp8 {
Some(f) if active => DW::Fp8 {
codes: f.codes.device_ptr(stream).0 as *const c_void,
scales: f.scales.device_ptr(stream).0 as *const f32,
sc_cols: f.sc_cols as i32,
},
_ => DW::Bf16(w_bf16.dev().device_ptr(stream).0 as *const c_void),
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum ExpertKind {
Nvfp4,
Mxfp4,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum ExpertArm {
Bf16Dequant,
Native,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum DecodePath {
Legacy,
Device { host_math: bool },
}
pub struct LayerDev {
pub il: u32,
pub ratio: usize,
pub expert_kind: ExpertKind,
pub wq_a: DenseBf16,
pub wq_b: DenseBf16,
pub wkv: DenseBf16,
pub wo_a: DenseBf16, pub wo_b: DenseBf16,
pub q_norm: CudaSlice<f32>,
pub kv_norm: CudaSlice<f32>,
pub attn_norm: CudaSlice<f32>,
pub ffn_norm: CudaSlice<f32>,
pub sink: CudaSlice<f32>,
pub cmp: Option<CmpDev>,
pub idx: Option<IdxDev>,
pub hc_attn_fn: CudaSlice<f32>,
pub hc_ffn_fn: CudaSlice<f32>,
pub hc_attn_base: Vec<f32>,
pub hc_attn_scale: Vec<f32>,
pub hc_ffn_base: Vec<f32>,
pub hc_ffn_scale: Vec<f32>,
pub hc_attn_base_dev: CudaSlice<f32>,
pub hc_attn_scale_dev: CudaSlice<f32>,
pub hc_ffn_base_dev: CudaSlice<f32>,
pub hc_ffn_scale_dev: CudaSlice<f32>,
pub gate_bias_dev: Option<CudaSlice<f32>>,
pub tid2eid_dev: Option<CudaSlice<i32>>,
pub experts_s2_dev: CudaSlice<f32>,
pub gate_w: CudaSlice<f32>, pub gate_bias: Option<Vec<f32>>,
pub tid2eid: Option<Vec<i64>>, pub experts_w: CudaSlice<u8>, pub experts_sc: CudaSlice<u8>, pub experts_s2: Vec<f32>, pub shared_w: [DenseBf16; 3], pub wq_a_fp8: Option<Fp8Dense>,
pub wq_b_fp8: Option<Fp8Dense>,
pub wkv_fp8: Option<Fp8Dense>,
pub wo_a_fp8: Option<Fp8Dense>,
pub wo_b_fp8: Option<Fp8Dense>,
pub shared_fp8: [Option<Fp8Dense>; 3],
}
#[derive(Default)]
pub struct GpuCapture {
pub embed_out: Option<Vec<f32>>,
pub layer_out: BTreeMap<u32, Vec<f32>>,
pub attn_out: BTreeMap<u32, Vec<f32>>,
pub x_dbg: BTreeMap<u32, Vec<f32>>,
pub q_dbg: BTreeMap<u32, Vec<f32>>,
pub kv_dbg: BTreeMap<u32, Vec<f32>>,
pub o_dbg: BTreeMap<u32, Vec<f32>>,
pub compressor_kv: BTreeMap<u32, (Vec<f32>, usize)>,
pub indexer_kv: BTreeMap<u32, (Vec<f32>, usize)>,
pub index_score: BTreeMap<u32, (Vec<f32>, usize)>,
pub moe_x: BTreeMap<u32, Vec<f32>>,
pub want: std::collections::BTreeSet<u32>,
}
pub struct MtpDev {
pub layer: LayerDev, pub enorm: CudaSlice<f32>,
pub hnorm: CudaSlice<f32>,
pub norm: CudaSlice<f32>,
pub e_proj: CudaSlice<u8>, pub h_proj: CudaSlice<u8>, pub hc_head_fn: CudaSlice<f32>,
pub hc_head_base: Vec<f32>,
pub hc_head_scale: Vec<f32>,
}
pub struct DsparkDev {
pub blocks: Vec<LayerDev>,
pub main_proj: CudaSlice<u8>, pub main_norm: CudaSlice<f32>,
pub norm: CudaSlice<f32>, pub markov_w1: CudaSlice<f32>, pub markov_w2: CudaSlice<f32>, pub markov_w1_host: Vec<f32>, pub conf_w: CudaSlice<f32>, pub hc_head_fn: CudaSlice<f32>, pub hc_head_base: Vec<f32>,
pub hc_head_scale: Vec<f32>,
pub block_size: usize,
pub noise_token: u32,
pub targets: Vec<usize>, pub rank: usize,
pub vocab: usize,
}
pub struct DsparkState {
pub rings: Vec<CudaSlice<f32>>,
pub taps: CudaSlice<f32>,
}
pub struct DsparkProposal {
pub out_ids: Vec<u32>,
pub confidence: Vec<f32>,
pub margins: Vec<f32>,
pub top1_logits: Vec<f32>,
pub capture: Option<DsparkCaptureOut>,
}
pub struct DsparkCaptureOut {
pub main_hidden: Vec<f32>,
pub main_x: Vec<f32>,
pub block_outs: Vec<Vec<f32>>,
pub x_collapsed: Vec<f32>,
pub logits_pre: Vec<f32>,
pub logits_post: Vec<f32>,
pub markov_embed: Vec<f32>,
}
pub struct Dsv4Gpu {
pub model: Dsv4Model,
pub stages: Vec<Stage>,
pub layer_stage: Vec<usize>, pub split_at: u32, pub max_seq: usize,
pub variant: ActQuantVariant,
pub fc_yarn_host: FreqsCis,
pub fc_plain_host: FreqsCis,
pub mtp: Option<MtpDev>,
pub dspark: Option<DsparkDev>,
pub expert_arm: ExpertArm,
pub decode_path: DecodePath,
pub dots_f32: bool,
pub chains_f32: bool,
pub dspark_head_f32: bool,
pub dense_fp8: bool,
boundary_ev: Vec<cudarc::driver::CudaEvent>,
hc_head_base: Vec<f32>,
hc_head_scale: Vec<f32>,
}
pub struct ForwardOut {
pub logits: Vec<f32>,
pub h_last: CudaSlice<f32>,
}
pub struct LayerCache {
pub kvc: CudaSlice<f32>,
pub n_blocks: usize,
pub pend_kv: Option<CudaSlice<f32>>,
pub pend_score: Option<CudaSlice<f32>>,
pub ikvc: Option<CudaSlice<f32>>,
pub i_blocks: usize,
pub ipend_kv: Option<CudaSlice<f32>>,
pub ipend_score: Option<CudaSlice<f32>>,
}
pub struct StepWs {
pub h_a: CudaSlice<f32>, pub h_b: CudaSlice<f32>, pub h_rx: CudaSlice<f32>, pub emb: CudaSlice<f32>, pub mixes: CudaSlice<f32>, pub pre: CudaSlice<f32>, pub post: CudaSlice<f32>, pub comb: CudaSlice<f32>, pub y_hc: CudaSlice<f32>, pub x: CudaSlice<f32>, pub xf: CudaSlice<f32>, pub qr: CudaSlice<f32>, pub qr_b: CudaSlice<u8>, pub q: CudaSlice<f32>, pub kv: CudaSlice<f32>, pub qi: CudaSlice<f32>, pub wproj: CudaSlice<f32>, pub score: CudaSlice<f32>, pub idx: CudaSlice<i32>, pub o: CudaSlice<f32>, pub o_b: CudaSlice<u8>, pub og: CudaSlice<f32>, pub attn_out: CudaSlice<f32>, pub gemm_xb: CudaSlice<u8>, pub raw: CudaSlice<f32>, pub sel: CudaSlice<i32>, pub selw: CudaSlice<f32>, pub order: CudaSlice<i32>, pub xq: CudaSlice<u8>, pub xs: CudaSlice<f32>, pub g1: CudaSlice<f32>, pub g3: CudaSlice<f32>, pub hbuf: CudaSlice<f32>, pub hq: CudaSlice<u8>, pub hs: CudaSlice<f32>, pub contrib: CudaSlice<f32>, pub y: CudaSlice<f32>, pub xb: CudaSlice<u8>, pub sg1: CudaSlice<f32>, pub sg3: CudaSlice<f32>,
pub shbuf: CudaSlice<f32>,
pub shb16: CudaSlice<u8>, pub sh_out: CudaSlice<f32>, pub cmp_kv_row: CudaSlice<f32>, pub cmp_sc_row: CudaSlice<f32>, pub cmp_emit: CudaSlice<f32>, pub cmp_shift: CudaSlice<f32>, pub sink_scores: CudaSlice<f32>,
pub sink_evals: CudaSlice<f32>,
pub sink_den: CudaSlice<f64>,
pub head_mixes: CudaSlice<f32>, pub head_pre: CudaSlice<f32>, pub collapsed: CudaSlice<f32>, pub logits: CudaSlice<f32>, pub argmax: CudaSlice<i32>, pub tok: CudaSlice<i32>, }
pub struct DecodeState {
pub caches: Vec<LayerCache>,
pub pos: usize,
pub cache_bytes: Vec<u64>,
pub ws: Option<Vec<StepWs>>,
}
fn sp(stream: &CudaStream) -> *mut c_void {
stream.cu_stream() as *mut c_void
}
fn upload_f32(stream: &std::sync::Arc<CudaStream>, v: &[f32]) -> Res<CudaSlice<f32>> {
let mut d = stream.alloc_zeros::<f32>(v.len()).map_err(e("alloc f32"))?;
stream.memcpy_htod(v, &mut d).map_err(e("htod f32"))?;
Ok(d)
}
fn upload_i32(stream: &std::sync::Arc<CudaStream>, v: &[i32]) -> Res<CudaSlice<i32>> {
let mut d = stream.alloc_zeros::<i32>(v.len()).map_err(e("alloc i32"))?;
stream.memcpy_htod(v, &mut d).map_err(e("htod i32"))?;
Ok(d)
}
fn upload_u8(stream: &std::sync::Arc<CudaStream>, v: &[u8]) -> Res<CudaSlice<u8>> {
let mut d = stream.alloc_zeros::<u8>(v.len()).map_err(e("alloc u8"))?;
stream.memcpy_htod(v, &mut d).map_err(e("htod u8"))?;
Ok(d)
}
fn dtoh_f32(stream: &std::sync::Arc<CudaStream>, d: &CudaSlice<f32>) -> Res<Vec<f32>> {
let mut v = vec![0f32; d.len()];
stream.memcpy_dtoh(d, &mut v[..]).map_err(e("dtoh"))?;
stream.synchronize().map_err(e("sync dtoh"))?;
Ok(v)
}
#[derive(Default, Clone)]
struct Dsv4PhaseAcc {
rows: Vec<(&'static str, u64, u64, u64)>,
stack: Vec<(usize, u64)>,
}
thread_local! {
static DSV4_PHASES: std::cell::RefCell<Dsv4PhaseAcc> =
std::cell::RefCell::new(Dsv4PhaseAcc::default());
}
fn dsv4_prof_sync() -> bool {
static V: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
*V.get_or_init(|| std::env::var("MEMRA_DSV4_ROUND_PROFILE").as_deref() == Ok("1"))
}
fn dsv4_prof_nvtx() -> bool {
static V: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
*V.get_or_init(|| std::env::var("MEMRA_DSV4_NVTX").as_deref() == Ok("1"))
}
pub fn dsv4_prof_on() -> bool {
dsv4_prof_sync() || dsv4_prof_nvtx()
}
pub struct Dsv4Phase<'a> {
stream: Option<&'a std::sync::Arc<CudaStream>>,
t0: std::time::Instant,
nvtx: bool,
}
impl<'a> Dsv4Phase<'a> {
pub fn new(name: &'static str, stream: Option<&'a std::sync::Arc<CudaStream>>) -> Option<Self> {
if !dsv4_prof_on() {
return None;
}
let nvtx = dsv4_prof_nvtx();
if nvtx {
unsafe {
k::memra_dsv4_nvtx_push(name.as_ptr() as *const std::os::raw::c_char);
}
}
let label = &name[..name.len() - 1];
DSV4_PHASES.with(|p| {
let mut p = p.borrow_mut();
let idx = match p.rows.iter().position(|r| r.0 == label) {
Some(i) => i,
None => {
p.rows.push((label, 0, 0, 0));
p.rows.len() - 1
}
};
p.stack.push((idx, 0));
});
Some(Dsv4Phase {
stream: if dsv4_prof_sync() { stream } else { None },
t0: std::time::Instant::now(),
nvtx,
})
}
}
impl Drop for Dsv4Phase<'_> {
fn drop(&mut self) {
if let Some(s) = self.stream {
let _ = s.synchronize();
}
let us = self.t0.elapsed().as_micros() as u64;
if self.nvtx {
unsafe {
k::memra_dsv4_nvtx_pop();
}
}
DSV4_PHASES.with(|p| {
let mut p = p.borrow_mut();
if let Some((idx, child)) = p.stack.pop() {
let r = &mut p.rows[idx];
r.1 += us;
r.2 += child;
r.3 += 1;
if let Some(top) = p.stack.last_mut() {
top.1 += us;
}
}
});
}
}
macro_rules! phase {
($name:literal, $stream:expr) => {
crate::dsv4_gpu::Dsv4Phase::new(concat!($name, "\0"), $stream)
};
}
pub fn dsv4_phase_report(tag: &str, rounds: u64, plain_us: f64) {
DSV4_PHASES.with(|p| {
let p = p.borrow();
if p.rows.is_empty() {
return;
}
let mode = if dsv4_prof_sync() {
"sync-bracketed (PERTURBS: compare the round total against the unbracketed A/B)"
} else {
"nvtx-only (host wall; GPU-busy comes from nsys nvtx_gpu_proj_sum)"
};
println!("\n[phase] === F ITEMISATION: {tag} ===");
println!("[phase] rounds={rounds} plain step={plain_us:.1} us mode={mode}");
println!(
"[phase] {:<26} {:>11} {:>11} {:>9} {:>12} {:>12}",
"phase", "incl_us/rd", "self_us/rd", "calls/rd", "self_plainstp", "incl_plainstp"
);
let mut rows = p.rows.clone();
rows.sort_by(|a, b| (b.1.saturating_sub(b.2)).cmp(&(a.1.saturating_sub(a.2))));
let r = rounds.max(1) as f64;
let mut leaf_sum = 0f64;
for (name, incl, child, calls) in rows {
let selfus = incl.saturating_sub(child) as f64 / r;
let inclus = incl as f64 / r;
leaf_sum += selfus;
let (sp, ip) = if plain_us > 0.0 {
(selfus / plain_us, inclus / plain_us)
} else {
(0.0, 0.0)
};
println!(
"[phase] {name:<26} {inclus:>11.1} {selfus:>11.1} {:>9.2} {sp:>12.4} {ip:>12.4}",
calls as f64 / r
);
}
println!(
"[phase] {:<26} {:>11} {:>11.1} {:>9} {:>12.4}",
"SUM of self",
"",
leaf_sum,
"",
if plain_us > 0.0 {
leaf_sum / plain_us
} else {
0.0
}
);
});
}
fn dsv4_dspark_chain_device() -> bool {
static V: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
*V.get_or_init(|| {
let on = std::env::var("MEMRA_DSV4_DSPARK_CHAIN").as_deref() == Ok("device");
if on {
println!(
"[spec] DSpark markov chain RESIDENT ON DEVICE (MEMRA_DSV4_DSPARK_CHAIN=device): \
one D2H per round instead of 2 x block_size"
);
}
on
})
}
fn dsv4_dspark_markov_rowblk() -> bool {
static V: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
*V.get_or_init(|| {
let on = std::env::var("MEMRA_DSV4_DSPARK_MARKOV").as_deref() == Ok("rowblk");
if on {
println!(
"[spec] DSpark markov bias GEMV on the ROW-BLOCKED dots twin \
(MEMRA_DSV4_DSPARK_MARKOV=rowblk; bit-identical, geometry only)"
);
}
on
})
}
pub fn dsv4_phase_reset() {
DSV4_PHASES.with(|p| *p.borrow_mut() = Dsv4PhaseAcc::default());
}
macro_rules! dp {
($slice:expr, $stream:expr) => {{ $slice.device_ptr($stream).0 as *const c_void }};
}
macro_rules! dpf {
($slice:expr, $stream:expr) => {{ $slice.device_ptr($stream).0 as *const f32 }};
}
macro_rules! dpm {
($slice:expr, $stream:expr) => {{ $slice.device_ptr_mut($stream).0 as *mut f32 }};
}
fn f32_to_bf16_exact(name: &str, v: &[f32]) -> Vec<u8> {
let mut out = Vec::with_capacity(v.len() * 2);
for (i, x) in v.iter().enumerate() {
let bits = x.to_bits();
assert!(
bits & 0xFFFF == 0,
"{name}: element {i} = {x} not exactly representable in bf16 — lane-4 rung \
exactness violated"
);
out.extend_from_slice(&((bits >> 16) as u16).to_le_bytes());
}
out
}
impl Dsv4Gpu {
fn tensor_bf16(&mut self, stage: usize, name: &str) -> Res<CudaSlice<u8>> {
let raw_name = format!("{name}.weight");
let is_bf16_raw = self
.model
.st
.raw(&raw_name)
.map(|(i, _)| i.dtype == "BF16")
.unwrap_or(false)
|| self
.model
.st
.raw(name)
.map(|(i, _)| i.dtype == "BF16")
.unwrap_or(false);
let stream = self.stages[stage].gpu.stream();
let bytes: u64;
let out = if is_bf16_raw {
let (_, raw) = self
.model
.st
.raw(&raw_name)
.or_else(|| self.model.st.raw(name))
.unwrap();
bytes = raw.len() as u64;
upload_u8(&stream, raw)?
} else {
let (_, v) = self.model.tensor_f32(name);
let b = f32_to_bf16_exact(name, &v);
bytes = b.len() as u64;
upload_u8(&stream, &b)?
};
self.stages[stage].loaded_bytes += bytes;
Ok(out)
}
fn tensor_dense(
&mut self,
stage: usize,
name: &str,
fp8_ok: bool,
) -> Res<(DenseBf16, Option<Fp8Dense>)> {
let raw_name = format!("{name}.weight");
let is_bf16_raw = self
.model
.st
.raw(&raw_name)
.map(|(i, _)| i.dtype == "BF16")
.unwrap_or(false)
|| self
.model
.st
.raw(name)
.map(|(i, _)| i.dtype == "BF16")
.unwrap_or(false);
if is_bf16_raw || !fp8_ok || !self.dense_fp8 {
return Ok((DenseBf16::Dev(self.tensor_bf16(stage, name)?), None));
}
let (wi, wraw, stem) = if let Some((i, r)) = self.model.st.raw(&raw_name) {
(i.clone(), r.to_vec(), name.to_string())
} else {
let (i, r) = self
.model
.st
.raw(name)
.unwrap_or_else(|| panic!("missing dense tensor {name}"));
let stem = name.strip_suffix(".weight").unwrap_or(name).to_string();
(i.clone(), r.to_vec(), stem)
};
if wi.dtype != "F8_E4M3" {
return Ok((DenseBf16::Dev(self.tensor_bf16(stage, name)?), None));
}
assert_eq!(wi.shape.len(), 2, "{name}: fp8 dense tensor must be 2-D");
let rows = wi.shape[0] as usize;
let cols = wi.shape[1] as usize;
assert_eq!(cols % 8, 0, "{name}: fp8 dense cols {cols} % 8 != 0");
assert_eq!(wraw.len(), rows * cols, "{name}: fp8 byte count");
let scale_name = format!("{stem}.scale");
let (si, sraw) = self
.model
.st
.raw(&scale_name)
.unwrap_or_else(|| panic!("{name}: F8_E4M3 weight without {scale_name}"));
assert_eq!(si.dtype, "F8_E8M0", "{scale_name}: dtype");
let sc_rows = rows.div_ceil(128);
let sc_cols = cols.div_ceil(128);
assert_eq!(
(si.shape[0] as usize, si.shape[1] as usize),
(sc_rows, sc_cols),
"{scale_name}: scale grid shape vs [ceil({rows}/128), ceil({cols}/128)]"
);
let sc_f32: Vec<f32> = sraw
.iter()
.map(|&b| {
assert_ne!(b, 0xFF, "{scale_name}: e8m0 NaN code");
memra_gguf::dsv4::e8m0_to_f32(b)
})
.collect();
let (_, v) = self.model.tensor_f32(name);
assert_eq!(v.len(), rows * cols, "{name}: dequant len");
let step = (v.len() / 1024).max(1);
for idx in (0..v.len()).step_by(step) {
let (r, c) = (idx / cols, idx % cols);
let got = memra_gguf::nvfp4_repack::fp8_e4m3_to_f32(wraw[idx])
* sc_f32[(r / 128) * sc_cols + c / 128];
assert_eq!(
got.to_bits(),
v[idx].to_bits(),
"{name}: fp8 arm layout check failed at [{r},{c}] ({got} vs {})",
v[idx]
);
}
let b = f32_to_bf16_exact(name, &v);
let stream = self.stages[stage].gpu.stream();
let codes = upload_u8(&stream, &wraw)?;
let scales = upload_f32(&stream, &sc_f32)?;
self.stages[stage].loaded_bytes += (wraw.len() + sc_f32.len() * 4) as u64;
Ok((
DenseBf16::Host(b),
Some(Fp8Dense {
codes,
scales,
sc_cols,
rows,
cols,
}),
))
}
fn tensor_f32_dev(&mut self, stage: usize, name: &str) -> Res<CudaSlice<f32>> {
let (_, v) = self.model.tensor_f32(name);
let stream = self.stages[stage].gpu.stream();
self.stages[stage].loaded_bytes += (v.len() * 4) as u64;
upload_f32(&stream, &v)
}
fn load_cmp(
&mut self,
stage: usize,
prefix: &str,
ratio: usize,
d: usize,
rotate: bool,
) -> Res<CmpDev> {
let (wkv_shape, _) = self.model.tensor_f32(&format!("{prefix}.wkv.weight"));
let latent = wkv_shape[0];
let overlap = ratio == 4;
assert_eq!(latent, if overlap { 2 * d } else { d }, "{prefix} latent");
Ok(CmpDev {
ratio,
d,
latent,
overlap,
rotate,
wkv: self.tensor_f32_dev(stage, &format!("{prefix}.wkv.weight"))?,
wgate: self.tensor_f32_dev(stage, &format!("{prefix}.wgate.weight"))?,
norm: self.tensor_f32_dev(stage, &format!("{prefix}.norm.weight"))?,
ape: self.tensor_f32_dev(stage, &format!("{prefix}.ape"))?,
})
}
fn load_layer(&mut self, stage: usize, il: u32, prefix: &str) -> Res<LayerDev> {
let d = self.model.cfg().clone();
let moe = self.model.mc.moe.clone().expect("moe block");
let ratio = d.compress_ratio(il) as usize;
let hd = d.head_dim as usize;
let p = prefix.to_string();
let hash = d.is_hash_layer(il);
let ne = moe.expert_count as usize;
let inter = moe.expert_ff_length as usize;
let hidden = self.model.mc.n_embd as usize;
let hc_load = |m: &Dsv4Model, fam: &str| -> (Vec<f32>, Vec<f32>, Vec<f32>) {
let fn_w = m.tensor_f32(&format!("{p}.hc_{fam}_fn")).1;
let base = m.tensor_f32(&format!("{p}.hc_{fam}_base")).1;
let scale = m.tensor_f32(&format!("{p}.hc_{fam}_scale")).1;
(fn_w, base, scale)
};
let (attn_fn, attn_base, attn_scale) = hc_load(&self.model, "attn");
let (ffn_fn, ffn_base, ffn_scale) = hc_load(&self.model, "ffn");
let stream = self.stages[stage].gpu.stream();
let hc_attn_fn = upload_f32(&stream, &attn_fn)?;
let hc_ffn_fn = upload_f32(&stream, &ffn_fn)?;
self.stages[stage].loaded_bytes += ((attn_fn.len() + ffn_fn.len()) * 4) as u64;
let (wi0, _) = self
.model
.st
.raw(&format!("{p}.ffn.experts.0.w1.weight"))
.unwrap_or_else(|| panic!("missing {p}.ffn.experts.0.w1.weight"));
let expert_kind = match wi0.dtype.as_str() {
"U8" => ExpertKind::Nvfp4,
"I8" => ExpertKind::Mxfp4,
other => panic!("{p}: unexpected expert weight dtype {other}"),
};
let wbytes = inter * hidden / 2; let sbytes = match expert_kind {
ExpertKind::Nvfp4 => inter * hidden / 16,
ExpertKind::Mxfp4 => inter * hidden / 32,
};
let mut experts_w = stream
.alloc_zeros::<u8>(ne * 3 * wbytes)
.map_err(e("alloc expert slab"))?;
let mut experts_sc = stream
.alloc_zeros::<u8>(ne * 3 * sbytes)
.map_err(e("alloc expert scale slab"))?;
let mut experts_s2 = Vec::with_capacity(ne * 3);
for ex in 0..ne {
for (pi, pname) in ["w1", "w2", "w3"].iter().enumerate() {
let base = format!("{p}.ffn.experts.{ex}.{pname}");
let (wi, wb) = self
.model
.st
.raw(&format!("{base}.weight"))
.unwrap_or_else(|| panic!("missing {base}.weight"));
assert_eq!(wb.len(), wbytes, "{base}: weight bytes");
let sb = match expert_kind {
ExpertKind::Nvfp4 => {
assert_eq!(wi.dtype, "U8", "{base}: expected NVFP4 U8 weight");
let (_, sb) = self
.model
.st
.raw(&format!("{base}.weight_scale"))
.unwrap_or_else(|| panic!("missing {base}.weight_scale"));
let (_, s2b) = self
.model
.st
.raw(&format!("{base}.weight_scale_2"))
.unwrap_or_else(|| panic!("missing {base}.weight_scale_2"));
let s2 = f32::from_le_bytes(s2b.try_into().expect("scale_2 4B"));
assert!(
s2 > 0.0 && s2.to_bits() & 0x007F_FFFF == 0,
"{base}: scale_2 {s2} not a power of two — rung exactness violated"
);
experts_s2.push(s2);
sb
}
ExpertKind::Mxfp4 => {
assert_eq!(wi.dtype, "I8", "{base}: expected MXFP4 I8 weight");
let (si, sb) = self
.model
.st
.raw(&format!("{base}.scale"))
.unwrap_or_else(|| panic!("missing {base}.scale"));
assert_eq!(si.dtype, "F8_E8M0", "{base}: expected E8M0 scale");
assert!(
!sb.contains(&0xFFu8),
"{base}: E8M0 NaN scale code — refusing"
);
experts_s2.push(1.0);
sb
}
};
assert_eq!(sb.len(), sbytes, "{base}: scale bytes");
let off = (ex * 3 + pi) * wbytes;
let mut view = experts_w.slice_mut(off..off + wbytes);
stream
.memcpy_htod(wb, &mut view)
.map_err(e("htod expert w"))?;
let soff = (ex * 3 + pi) * sbytes;
let mut sview = experts_sc.slice_mut(soff..soff + sbytes);
stream
.memcpy_htod(sb, &mut sview)
.map_err(e("htod expert sc"))?;
}
}
self.stages[stage].loaded_bytes +=
(ne * 3 * (wbytes + sbytes)) as u64 + (ne * 3 * 4) as u64;
let cmp = if ratio != 0 {
Some(self.load_cmp(stage, &format!("{p}.attn.compressor"), ratio, hd, false)?)
} else {
None
};
let fp8_ok = p.starts_with("layers.");
let idx = if d.has_indexer(il) {
let heads = d.index_n_heads as usize;
let ihd = d.index_head_dim as usize;
let (iwq_b, iwq_b_fp8) =
self.tensor_dense(stage, &format!("{p}.attn.indexer.wq_b"), fp8_ok)?;
let (iwp, iwp_fp8) = self.tensor_dense(
stage,
&format!("{p}.attn.indexer.weights_proj.weight"),
fp8_ok,
)?;
Some(IdxDev {
wq_b: iwq_b,
weights_proj: iwp,
wq_b_fp8: iwq_b_fp8,
weights_proj_fp8: iwp_fp8,
cmp: self.load_cmp(
stage,
&format!("{p}.attn.indexer.compressor"),
ratio,
ihd,
true,
)?,
heads,
hd: ihd,
topk: d.index_topk as usize,
})
} else {
None
};
let stream = self.stages[stage].gpu.stream();
let hc_attn_base_dev = upload_f32(&stream, &attn_base)?;
let hc_attn_scale_dev = upload_f32(&stream, &attn_scale)?;
let hc_ffn_base_dev = upload_f32(&stream, &ffn_base)?;
let hc_ffn_scale_dev = upload_f32(&stream, &ffn_scale)?;
let gate_bias_host: Option<Vec<f32>> = if hash {
None
} else {
Some(self.model.tensor_f32(&format!("{p}.ffn.gate.bias")).1)
};
let gate_bias_dev = match &gate_bias_host {
Some(b) => Some(upload_f32(&stream, b)?),
None => None,
};
let tid2eid_host: Option<Vec<i64>> = if hash {
Some(self.model.tensor_i64(&format!("{p}.ffn.gate.tid2eid")).1)
} else {
None
};
let tid2eid_dev = match &tid2eid_host {
Some(t) => {
let topk = moe.expert_used_count as usize;
assert_eq!(t.len() % topk, 0, "{p}: tid2eid rows");
let mut t32 = Vec::with_capacity(t.len());
for row in t.chunks(topk) {
let mut seen = std::collections::BTreeSet::new();
for &ex in row {
assert!(
(0..ne as i64).contains(&ex),
"{p}: tid2eid out of range at load"
);
assert!(seen.insert(ex), "{p}: duplicate expert id in tid2eid row");
t32.push(ex as i32);
}
}
Some(upload_i32(&stream, &t32)?)
}
None => None,
};
let experts_s2_dev = upload_f32(&stream, &experts_s2)?;
let (wq_a, wq_a_fp8) = self.tensor_dense(stage, &format!("{p}.attn.wq_a"), fp8_ok)?;
let (wq_b, wq_b_fp8) = self.tensor_dense(stage, &format!("{p}.attn.wq_b"), fp8_ok)?;
let (wkv, wkv_fp8) = self.tensor_dense(stage, &format!("{p}.attn.wkv"), fp8_ok)?;
let (wo_a, wo_a_fp8) = self.tensor_dense(stage, &format!("{p}.attn.wo_a"), fp8_ok)?;
let (wo_b, wo_b_fp8) = self.tensor_dense(stage, &format!("{p}.attn.wo_b"), fp8_ok)?;
let (sw1, sw1_fp8) =
self.tensor_dense(stage, &format!("{p}.ffn.shared_experts.w1"), fp8_ok)?;
let (sw2, sw2_fp8) =
self.tensor_dense(stage, &format!("{p}.ffn.shared_experts.w2"), fp8_ok)?;
let (sw3, sw3_fp8) =
self.tensor_dense(stage, &format!("{p}.ffn.shared_experts.w3"), fp8_ok)?;
Ok(LayerDev {
il,
ratio,
expert_kind,
hc_attn_base_dev,
hc_attn_scale_dev,
hc_ffn_base_dev,
hc_ffn_scale_dev,
gate_bias_dev,
tid2eid_dev,
experts_s2_dev,
wq_a,
wq_b,
wkv,
wo_a,
wo_b,
wq_a_fp8,
wq_b_fp8,
wkv_fp8,
wo_a_fp8,
wo_b_fp8,
q_norm: self.tensor_f32_dev(stage, &format!("{p}.attn.q_norm.weight"))?,
kv_norm: self.tensor_f32_dev(stage, &format!("{p}.attn.kv_norm.weight"))?,
attn_norm: self.tensor_f32_dev(stage, &format!("{p}.attn_norm.weight"))?,
ffn_norm: self.tensor_f32_dev(stage, &format!("{p}.ffn_norm.weight"))?,
sink: self.tensor_f32_dev(stage, &format!("{p}.attn.attn_sink"))?,
cmp,
idx,
hc_attn_fn,
hc_ffn_fn,
hc_attn_base: attn_base,
hc_attn_scale: attn_scale,
hc_ffn_base: ffn_base,
hc_ffn_scale: ffn_scale,
gate_w: self.tensor_f32_dev(stage, &format!("{p}.ffn.gate.weight"))?,
gate_bias: gate_bias_host,
tid2eid: tid2eid_host,
experts_w,
experts_sc,
experts_s2,
shared_w: [sw1, sw2, sw3],
shared_fp8: [sw1_fp8, sw2_fp8, sw3_fp8],
})
}
pub fn load(
dir: &Path,
devices: &[usize],
variant: ActQuantVariant,
max_seq: usize,
) -> Res<Self> {
assert_eq!(devices.len(), 2, "lane 4 placement is a 2-card layer split");
let model = Dsv4Model::open(dir);
let d = model.cfg().clone();
let mc = model.mc.clone();
let n_trunk = mc.n_layer - mc.nextn_predict_layers;
let rd = d.qk_rope_head_dim as usize;
let layer_bytes = |il: u32| -> u64 {
let ratio = d.compress_ratio(il);
let base = 3_875_000_000u64; match ratio {
4 => base + 66_000_000,
_ => base,
}
};
let total: u64 = (0..n_trunk).map(layer_bytes).sum();
let mut acc = 0u64;
let mut split_at = n_trunk / 2;
for il in 0..n_trunk {
acc += layer_bytes(il);
if acc * 2 >= total {
split_at = il + 1;
break;
}
}
let fc_yarn_host = precompute_freqs_cis(
rd,
max_seq,
d.rope_yarn_orig_ctx,
d.compress_rope_theta,
d.rope_yarn_factor,
d.rope_yarn_beta_fast,
d.rope_yarn_beta_slow,
);
let fc_plain_host = precompute_freqs_cis(
rd,
max_seq,
0,
mc.rope_freq_base,
d.rope_yarn_factor,
d.rope_yarn_beta_fast,
d.rope_yarn_beta_slow,
);
let flat =
|fc: &FreqsCis| -> Vec<f32> { fc.cs.iter().flat_map(|&(c, s)| [c, s]).collect() };
let inter = mc.moe.as_ref().expect("moe").expert_ff_length as usize;
let hidden = mc.n_embd as usize;
let mut stages = Vec::new();
for &dev in devices {
let gpu = memra_runtime::Gpu::new(dev).map_err(e("Gpu::new"))?;
unsafe { gpu.ctx.disable_event_tracking() };
let stream = gpu.stream();
let fc_yarn = upload_f32(&stream, &flat(&fc_yarn_host))?;
let fc_plain = upload_f32(&stream, &flat(&fc_plain_host))?;
let ws = stream.alloc_zeros::<u8>(64 << 20).map_err(e("ws alloc"))?;
let deq = [
stream
.alloc_zeros::<u8>(inter * hidden * 2)
.map_err(e("deq"))?,
stream
.alloc_zeros::<u8>(inter * hidden * 2)
.map_err(e("deq"))?,
stream
.alloc_zeros::<u8>(inter * hidden * 2)
.map_err(e("deq"))?,
];
stages.push(Stage {
dev,
gpu,
layers: Vec::new(),
embed: None,
head: None,
trunk_norm: None,
hc_head_fn: None,
fc_yarn,
fc_plain,
ws,
deq,
loaded_bytes: 0,
hc_head_base_dev: None,
hc_head_scale_dev: None,
});
}
let decode_path = match std::env::var("MEMRA_DSV4_DECODE_PATH").as_deref() {
Err(_) | Ok("") | Ok("legacy") => DecodePath::Legacy,
Ok("device-hostmath") => DecodePath::Device { host_math: true },
Ok("device") => DecodePath::Device { host_math: false },
Ok(other) => {
return Err(format!(
"MEMRA_DSV4_DECODE_PATH '{other}' unknown (legacy | device | device-hostmath)"
));
}
};
let on_device = matches!(decode_path, DecodePath::Device { .. });
let (dots_f32, chains_f32) = match std::env::var("MEMRA_DSV4_DOTS_ARM").as_deref() {
Err(_) | Ok("") => {
if on_device {
(true, true)
} else {
(false, false)
}
}
Ok(explicit @ ("f32x" | "f32")) if !on_device => {
return Err(format!(
"MEMRA_DSV4_DOTS_ARM={explicit} requires MEMRA_DSV4_DECODE_PATH=device \
(the f32 dots arms exist on the device decode path only)"
));
}
Ok("f32x") => (true, true),
Ok("f64") => (false, false),
Ok("f32") => (true, false),
Ok(other) => {
return Err(format!(
"MEMRA_DSV4_DOTS_ARM '{other}' unknown (f64 | f32 | f32x)"
));
}
};
let dspark_head_f32 = match std::env::var("MEMRA_DSV4_DSPARK_HEAD_ARM").as_deref() {
Err(_) | Ok("") | Ok("f32x") => true,
Ok("f64") => false,
Ok(other) => {
return Err(format!(
"MEMRA_DSV4_DSPARK_HEAD_ARM '{other}' unknown (f64 | f32x)"
));
}
};
let dense_fp8 = resolve_dense_arm(
std::env::var("MEMRA_DSV4_DENSE_ARM").ok().as_deref(),
on_device,
)?;
let mut me = Dsv4Gpu {
model,
stages,
layer_stage: (0..n_trunk).map(|il| usize::from(il >= split_at)).collect(),
split_at,
max_seq,
variant,
fc_yarn_host,
fc_plain_host,
mtp: None,
expert_arm: if memra_gguf::dsv4_forward::expert_arm_native() {
ExpertArm::Native
} else {
ExpertArm::Bf16Dequant
},
decode_path,
dots_f32,
chains_f32,
dspark_head_f32,
dense_fp8,
dspark: None,
boundary_ev: Vec::new(),
hc_head_base: Vec::new(),
hc_head_scale: Vec::new(),
};
eprintln!(
"[load] expert arm: {:?} | decode path: {:?} | dots arm: {}",
me.expert_arm,
me.decode_path,
if me.chains_f32 {
"f32x (dots + sink/norm/indexer chains)"
} else if me.dots_f32 {
"f32"
} else {
"f64"
}
);
eprintln!(
"[load] dspark exit-head dots arm: {} (rung-4c fork; drafts only, never the \
emitted stream)",
if me.dspark_head_f32 { "f32x" } else { "f64" }
);
eprintln!(
"[load] dense arm: {} (iteration-5; fp8 = FP8-blk linears as-stored on the \
device decode/verify paths, bit-identical twins)",
if me.dense_fp8 { "fp8" } else { "bf16" }
);
if matches!(me.decode_path, DecodePath::Device { .. }) && me.expert_arm != ExpertArm::Native
{
return Err(
"MEMRA_DSV4_DECODE_PATH=device requires MEMRA_DSV4_EXPERT_ARM=native".to_string(),
);
}
if matches!(me.decode_path, DecodePath::Device { .. }) && me.stages.len() > 1 {
use cudarc::driver::sys as cus;
for a in 0..me.stages.len() {
for b in 0..me.stages.len() {
if a == b || me.stages[a].dev == me.stages[b].dev {
continue;
}
me.stages[a]
.gpu
.ctx
.bind_to_thread()
.map_err(e("peer bind"))?;
let rc =
unsafe { cus::cuCtxEnablePeerAccess(me.stages[b].gpu.ctx.cu_ctx(), 0) };
if rc != cus::cudaError_enum::CUDA_SUCCESS
&& rc != cus::cudaError_enum::CUDA_ERROR_PEER_ACCESS_ALREADY_ENABLED
{
return Err(format!(
"cuCtxEnablePeerAccess(dev{} -> dev{}) failed: {rc:?}",
me.stages[a].dev, me.stages[b].dev
));
}
let dev = cudarc::driver::result::device::get(me.stages[a].dev as i32)
.map_err(e("device get"))?;
let mut pool: cus::CUmemoryPool = std::ptr::null_mut();
unsafe {
cus::cuDeviceGetDefaultMemPool(&mut pool, dev)
.result()
.map_err(e("default pool"))?;
}
let desc = cus::CUmemAccessDesc {
location: cus::CUmemLocation {
type_: cus::CUmemLocationType::CU_MEM_LOCATION_TYPE_DEVICE,
id: me.stages[b].dev as i32,
},
flags: cus::CUmemAccess_flags::CU_MEM_ACCESS_FLAGS_PROT_READWRITE,
};
let rc = unsafe { cus::cuMemPoolSetAccess(pool, &desc, 1) };
if rc != cus::cudaError_enum::CUDA_SUCCESS {
return Err(format!(
"cuMemPoolSetAccess(dev{} pool -> dev{}) failed: {rc:?}",
me.stages[a].dev, me.stages[b].dev
));
}
}
}
for bnd in 0..me.stages.len() - 1 {
let ev = me.stages[bnd]
.gpu
.ctx
.new_event(None)
.map_err(e("boundary event"))?;
me.boundary_ev.push(ev);
}
eprintln!(
"[load] lane-8 peer transport enabled ({} boundaries)",
me.boundary_ev.len()
);
}
me.stages[0].embed = Some({
let (_, raw) = me.model.st.raw("embed.weight").expect("embed.weight");
let stream = me.stages[0].gpu.stream();
me.stages[0].loaded_bytes += raw.len() as u64;
upload_u8(&stream, raw)?
});
let last = me.stages.len() - 1;
me.stages[last].head = Some({
let (_, raw) = me.model.st.raw("head.weight").expect("head.weight");
let stream = me.stages[last].gpu.stream();
me.stages[last].loaded_bytes += raw.len() as u64;
upload_u8(&stream, raw)?
});
me.stages[last].trunk_norm = Some(me.tensor_f32_dev(last, "norm.weight")?);
me.stages[last].hc_head_fn = Some(me.tensor_f32_dev(last, "hc_head_fn")?);
me.hc_head_base = me.model.tensor_f32("hc_head_base").1;
me.hc_head_scale = me.model.tensor_f32("hc_head_scale").1;
{
let stream = me.stages[last].gpu.stream();
let base_dev = upload_f32(&stream, &me.hc_head_base)?;
let scale_dev = upload_f32(&stream, &me.hc_head_scale)?;
me.stages[last].hc_head_base_dev = Some(base_dev);
me.stages[last].hc_head_scale_dev = Some(scale_dev);
}
let t0 = std::time::Instant::now();
for il in 0..n_trunk {
let stage = me.layer_stage[il as usize];
let l = me.load_layer(stage, il, &format!("layers.{il}"))?;
me.stages[stage].layers.push(l);
if il % 4 == 3 || il + 1 == n_trunk {
eprintln!(
"[load] layer {il} -> dev{} done t={:.0}s",
me.stages[stage].dev,
t0.elapsed().as_secs_f64()
);
}
}
let nextn = me.model.mc.nextn_predict_layers;
if nextn > 0 && me.model.has("mtp.0.e_proj.weight") {
assert_eq!(
nextn, 1,
"multi-NextN chains not wired (single MTP layer expected)"
);
let p = "mtp.0";
let layer = me.load_layer(last, n_trunk, p)?;
assert_eq!(
layer.expert_kind,
ExpertKind::Mxfp4,
"MTP experts must be MXFP4"
);
let mtp = MtpDev {
layer,
enorm: me.tensor_f32_dev(last, &format!("{p}.enorm.weight"))?,
hnorm: me.tensor_f32_dev(last, &format!("{p}.hnorm.weight"))?,
norm: me.tensor_f32_dev(last, &format!("{p}.norm.weight"))?,
e_proj: me.tensor_bf16(last, &format!("{p}.e_proj"))?,
h_proj: me.tensor_bf16(last, &format!("{p}.h_proj"))?,
hc_head_fn: me.tensor_f32_dev(last, &format!("{p}.hc_head_fn"))?,
hc_head_base: me.model.tensor_f32(&format!("{p}.hc_head_base")).1,
hc_head_scale: me.model.tensor_f32(&format!("{p}.hc_head_scale")).1,
};
me.mtp = Some(mtp);
} else if nextn > 0 {
if std::env::var("MEMRA_DSV4_DRAFTER").as_deref() == Ok("dspark") {
let cfg = memra_gguf::dsv4_dspark::DsparkConfig::load(dir, &me.model);
let hidden = me.model.mc.n_embd as usize;
let mut blocks = Vec::with_capacity(cfg.n_blocks);
for k in 0..cfg.n_blocks {
let layer = me.load_layer(last, n_trunk + k as u32, &format!("mtp.{k}"))?;
assert_eq!(layer.ratio, 0, "dspark block mtp.{k} must be ratio 0");
assert_eq!(
layer.expert_kind,
ExpertKind::Mxfp4,
"dspark experts must be MXFP4"
);
blocks.push(layer);
}
let last_p = format!("mtp.{}", cfg.n_blocks - 1);
let (mp_shape, _) = me.model.tensor_f32("mtp.0.main_proj");
assert_eq!(
mp_shape,
vec![hidden, cfg.target_layer_ids.len() * hidden],
"main_proj shape"
);
let (w1_shape, w1) = me
.model
.tensor_f32(&format!("{last_p}.markov_head.markov_w1.weight"));
let (w2_shape, w2) = me
.model
.tensor_f32(&format!("{last_p}.markov_head.markov_w2.weight"));
let vocab = w1_shape[0];
assert_eq!(w1_shape[1], cfg.markov_rank, "markov_w1 rank");
assert_eq!(w2_shape, vec![vocab, cfg.markov_rank], "markov_w2 shape");
let (cf_shape, _) = me
.model
.tensor_f32(&format!("{last_p}.confidence_head.proj.weight"));
assert_eq!(
cf_shape,
vec![1, hidden + cfg.markov_rank],
"confidence proj shape"
);
let st_stream = me.stages[last].gpu.stream();
let markov_w1 = upload_f32(&st_stream, &w1)?;
let markov_w2 = upload_f32(&st_stream, &w2)?;
let dspark = DsparkDev {
blocks,
main_proj: me.tensor_bf16(last, "mtp.0.main_proj")?,
main_norm: me.tensor_f32_dev(last, "mtp.0.main_norm.weight")?,
norm: me.tensor_f32_dev(last, &format!("{last_p}.norm.weight"))?,
markov_w1,
markov_w2,
markov_w1_host: w1,
conf_w: me
.tensor_f32_dev(last, &format!("{last_p}.confidence_head.proj.weight"))?,
hc_head_fn: me.tensor_f32_dev(last, &format!("{last_p}.hc_head_fn"))?,
hc_head_base: me.model.tensor_f32(&format!("{last_p}.hc_head_base")).1,
hc_head_scale: me.model.tensor_f32(&format!("{last_p}.hc_head_scale")).1,
block_size: cfg.block_size,
noise_token: cfg.noise_token_id,
targets: cfg.target_layer_ids.clone(),
rank: cfg.markov_rank,
vocab,
};
eprintln!(
"[load] drafter: DSpark ({} blocks, block_size {}, targets {:?}) \
resident on stage {last}",
cfg.n_blocks, cfg.block_size, cfg.target_layer_ids
);
me.dspark = Some(dspark);
} else {
eprintln!(
"[load] drafter: {nextn} DSpark block(s) (mtp.0.e_proj absent) — GPU \
drafter path off (set MEMRA_DSV4_DRAFTER=dspark); trunk-only"
);
}
}
for st in &me.stages {
st.gpu.stream().synchronize().map_err(e("load sync"))?;
}
Ok(me)
}
pub fn vram_report(&self) -> Res<Vec<(usize, u64, u64, u64)>> {
let mut out = Vec::new();
for st in &self.stages {
st.gpu.ctx.bind_to_thread().map_err(e("bind ctx"))?;
let (free, total) = st.gpu.ctx.mem_get_info().map_err(e("mem_get_info"))?;
out.push((st.dev, free as u64, total as u64, st.loaded_bytes));
}
Ok(out)
}
#[allow(clippy::too_many_arguments)]
fn gemm(
st: &Stage,
x_f32: &CudaSlice<f32>,
w_bf16: &CudaSlice<u8>,
w_off_elems: usize,
m: usize,
n: usize,
kdim: usize,
y: &mut CudaSlice<f32>,
) -> Res<()> {
let stream = st.gpu.stream();
let mut xb = stream
.alloc_zeros::<u8>(m * kdim * 2)
.map_err(e("alloc xb"))?;
unsafe {
ck(
"cvt_bf16",
k::memra_dsv4_cvt_bf16(
dpf!(x_f32, &stream),
xb.device_ptr_mut(&stream).0 as *mut c_void,
(m * kdim) as i64,
sp(&stream),
),
)?;
ck(
"gemm_bf16",
k::memra_dsv4_gemm_bf16(
(w_bf16.device_ptr(&stream).0 as usize + w_off_elems * 2) as *const c_void,
dp!(xb, &stream),
dpm!(y, &stream),
m as i32,
n as i32,
kdim as i32,
st.dev as i32,
st.ws.device_ptr(&stream).0 as *mut c_void,
st.ws.len(),
sp(&stream),
),
)?;
}
Ok(())
}
#[allow(clippy::too_many_arguments)]
fn gemm_pre(
st: &Stage,
xb: &CudaSlice<u8>,
w_bf16_ptr: *const c_void,
m: usize,
n: usize,
kdim: usize,
y: &mut CudaSlice<f32>,
) -> Res<()> {
let stream = st.gpu.stream();
unsafe {
ck(
"gemm_bf16",
k::memra_dsv4_gemm_bf16(
w_bf16_ptr,
dp!(xb, &stream),
dpm!(y, &stream),
m as i32,
n as i32,
kdim as i32,
st.dev as i32,
st.ws.device_ptr(&stream).0 as *mut c_void,
st.ws.len(),
sp(&stream),
),
)?;
}
Ok(())
}
fn dots(
st: &Stage,
x: &CudaSlice<f32>,
w_f32: &CudaSlice<f32>,
s: usize,
kdim: usize,
n: usize,
y: &mut CudaSlice<f32>,
) -> Res<()> {
let stream = st.gpu.stream();
unsafe {
ck(
"dots_f32",
k::memra_dsv4_dots_f32(
dpf!(x, &stream),
dp!(w_f32, &stream),
0,
dpm!(y, &stream),
s as i32,
kdim as i32,
n as i32,
sp(&stream),
),
)?;
}
Ok(())
}
fn dots_dev(
&self,
st: &Stage,
x: &CudaSlice<f32>,
w_f32: &CudaSlice<f32>,
s: usize,
kdim: usize,
n: usize,
y: &mut CudaSlice<f32>,
) -> Res<()> {
if !self.dots_f32 {
return Self::dots(st, x, w_f32, s, kdim, n, y);
}
let stream = st.gpu.stream();
unsafe {
ck(
"dots_f32acc",
k::memra_dsv4_dots_f32acc(
dpf!(x, &stream),
dp!(w_f32, &stream),
0,
dpm!(y, &stream),
s as i32,
kdim as i32,
n as i32,
sp(&stream),
),
)?;
}
Ok(())
}
#[allow(clippy::too_many_arguments)]
fn compressor(
&self,
st: &Stage,
cmp: &CmpDev,
x: &CudaSlice<f32>, s: usize,
hidden: usize,
fc_dev: &CudaSlice<f32>,
rd: usize,
eps: f32,
) -> Res<(
Option<(CudaSlice<f32>, usize)>,
CudaSlice<f32>,
CudaSlice<f32>,
)> {
let stream = st.gpu.stream();
let mut kv = stream
.alloc_zeros::<f32>(s * cmp.latent)
.map_err(e("cmp kv"))?;
let mut score = stream
.alloc_zeros::<f32>(s * cmp.latent)
.map_err(e("cmp score"))?;
Self::dots(st, x, &cmp.wkv, s, hidden, cmp.latent, &mut kv)?;
Self::dots(st, x, &cmp.wgate, s, hidden, cmp.latent, &mut score)?;
if s < cmp.ratio {
return Ok((None, kv, score));
}
let cutoff = s - s % cmp.ratio;
let nb = cutoff / cmp.ratio;
let mut pooled = stream
.alloc_zeros::<f32>(nb * cmp.d)
.map_err(e("cmp out"))?;
unsafe {
ck(
"compressor_pool",
k::memra_dsv4_compressor_pool(
dpf!(kv, &stream),
dpf!(score, &stream),
dpf!(cmp.ape, &stream),
dpm!(pooled, &stream),
nb as i32,
cmp.ratio as i32,
cmp.d as i32,
cmp.latent as i32,
cmp.overlap as i32,
sp(&stream),
),
)?;
ck(
"rmsnorm cmp",
k::memra_dsv4_rmsnorm(
dpf!(pooled, &stream),
dpf!(cmp.norm, &stream),
dpm!(pooled, &stream),
nb as i32,
cmp.d as i32,
eps,
sp(&stream),
),
)?;
let positions: Vec<i32> = (0..nb).map(|j| (j * cmp.ratio) as i32).collect();
let pos_dev = upload_i32(&stream, &positions)?;
ck(
"rope cmp",
k::memra_dsv4_rope(
dpm!(pooled, &stream),
nb as i32,
1,
cmp.d as i32,
rd as i32,
dpf!(fc_dev, &stream),
pos_dev.device_ptr(&stream).0 as *const i32,
0,
sp(&stream),
),
)?;
if cmp.rotate {
let scale = (cmp.d as f32).powf(-0.5);
ck(
"hadamard cmp",
k::memra_dsv4_hadamard(
dpm!(pooled, &stream),
nb as i32,
cmp.d as i32,
scale,
sp(&stream),
),
)?;
ck(
"fp4 cmp",
k::memra_dsv4_fp4_act_quant(
dpm!(pooled, &stream),
nb as i32,
cmp.d as i64,
cmp.d as i32,
sp(&stream),
),
)?;
} else {
ck(
"act_quant cmp",
k::memra_dsv4_act_quant(
dpm!(pooled, &stream),
nb as i32,
cmp.d as i64,
(cmp.d - rd) as i32,
64,
(self.variant == ActQuantVariant::ClampOnly) as i32,
sp(&stream),
),
)?;
}
}
Ok((Some((pooled, nb)), kv, score))
}
#[allow(clippy::too_many_arguments)]
fn populate_cmp_cache(
stream: &std::sync::Arc<CudaStream>,
s: usize,
cmp_ratio: usize,
latent: usize,
d: usize,
pooled: &Option<(CudaSlice<f32>, usize)>,
kv_raw: &CudaSlice<f32>,
score_raw: &CudaSlice<f32>,
store: &mut CudaSlice<f32>,
row0: usize,
blocks: &mut usize,
pend_kv: &mut CudaSlice<f32>,
pend_score: &mut CudaSlice<f32>,
overlap: bool,
) -> Res<()> {
*blocks = 0;
if let Some((buf, nb)) = pooled {
let src = buf.slice(0..nb * d);
let mut dst = store.slice_mut(row0 * d..(row0 + nb) * d);
stream.memcpy_dtod(&src, &mut dst).map_err(e("cmp store"))?;
*blocks = *nb;
}
let cutoff = s - s % cmp_ratio;
let rem = s - cutoff;
if overlap {
if cutoff >= cmp_ratio {
let a = (cutoff - cmp_ratio) * latent;
let b = cutoff * latent;
let src = kv_raw.slice(a..b);
let mut dst = pend_kv.slice_mut(0..cmp_ratio * latent);
stream
.memcpy_dtod(&src, &mut dst)
.map_err(e("pend kv prev"))?;
let src = score_raw.slice(a..b);
let mut dst = pend_score.slice_mut(0..cmp_ratio * latent);
stream
.memcpy_dtod(&src, &mut dst)
.map_err(e("pend sc prev"))?;
}
if rem > 0 {
let a = cutoff * latent;
let src = kv_raw.slice(a..s * latent);
let mut dst = pend_kv.slice_mut(cmp_ratio * latent..(cmp_ratio + rem) * latent);
stream
.memcpy_dtod(&src, &mut dst)
.map_err(e("pend kv cur"))?;
let src = score_raw.slice(a..s * latent);
let mut dst = pend_score.slice_mut(cmp_ratio * latent..(cmp_ratio + rem) * latent);
stream
.memcpy_dtod(&src, &mut dst)
.map_err(e("pend sc cur"))?;
}
} else if rem > 0 {
let a = cutoff * latent;
let src = kv_raw.slice(a..s * latent);
let mut dst = pend_kv.slice_mut(0..rem * latent);
stream.memcpy_dtod(&src, &mut dst).map_err(e("pend kv"))?;
let src = score_raw.slice(a..s * latent);
let mut dst = pend_score.slice_mut(0..rem * latent);
stream.memcpy_dtod(&src, &mut dst).map_err(e("pend sc"))?;
}
Ok(())
}
#[allow(clippy::too_many_arguments)]
fn hc_pre(
st: &Stage,
h: &CudaSlice<f32>, fn_w: &CudaSlice<f32>,
base: &[f32],
scale: &[f32],
s: usize,
hc: usize,
hidden: usize,
iters: u32,
hc_eps: f32,
) -> Res<(CudaSlice<f32>, CudaSlice<f32>, CudaSlice<f32>)> {
let stream = st.gpu.stream();
let w = hc * hidden;
let rows = (2 + hc) * hc;
let mut mixes = stream.alloc_zeros::<f32>(s * rows).map_err(e("mixes"))?;
Self::dots(st, h, fn_w, s, w, rows, &mut mixes)?;
unsafe {
ck(
"rowsq_scale",
k::memra_dsv4_rowsq_scale(
dpf!(h, &stream),
dpm!(mixes, &stream),
s as i32,
w as i32,
rows as i32,
hc_eps,
sp(&stream),
),
)?;
}
let mixes_h = dtoh_f32(&stream, &mixes)?;
let (pre, post, comb) = hc_split_sinkhorn(&mixes_h, s, hc, scale, base, iters, hc_eps);
let pre_d = upload_f32(&stream, &pre)?;
let post_d = upload_f32(&stream, &post)?;
let comb_d = upload_f32(&stream, &comb)?;
let mut y = stream.alloc_zeros::<f32>(s * hidden).map_err(e("hc y"))?;
unsafe {
ck(
"hc_collapse",
k::memra_dsv4_hc_collapse(
dpf!(h, &stream),
dpf!(pre_d, &stream),
dpm!(y, &stream),
s as i32,
hc as i32,
hidden as i32,
sp(&stream),
),
)?;
}
Ok((y, post_d, comb_d))
}
#[allow(clippy::too_many_arguments)]
fn route_host(
layer: &LayerDev,
raw_scores: &[f32], ids: &[u32],
s: usize,
ne: usize,
topk: usize,
route_scale: f32,
) -> (Vec<usize>, Vec<f32>) {
let mut scores = raw_scores.to_vec();
for v in &mut scores {
*v = softplus_f32(*v).sqrt();
}
let mut indices = vec![0usize; s * topk];
if let Some(tid2eid) = &layer.tid2eid {
for t in 0..s {
let row = &tid2eid[ids[t] as usize * topk..(ids[t] as usize + 1) * topk];
let mut seen = std::collections::BTreeSet::new();
for (kk, &ex) in row.iter().enumerate() {
assert!(
(0..ne as i64).contains(&ex),
"layer {}: tid2eid out of range",
layer.il
);
assert!(
seen.insert(ex),
"layer {}: duplicate expert id in tid2eid row {}",
layer.il,
ids[t]
);
indices[t * topk + kk] = ex as usize;
}
}
} else {
let bias = layer.gate_bias.as_ref().expect("score layer needs bias");
for t in 0..s {
let biased: Vec<f32> = (0..ne).map(|ex| scores[t * ne + ex] + bias[ex]).collect();
let mut order: Vec<usize> = (0..ne).collect();
order.sort_by(|&a, &b| {
biased[b]
.partial_cmp(&biased[a])
.unwrap_or(std::cmp::Ordering::Equal)
.then(a.cmp(&b))
});
for kk in 0..topk {
indices[t * topk + kk] = order[kk];
}
}
}
let mut weights = vec![0f32; s * topk];
for t in 0..s {
let mut sum = 0f32;
for kk in 0..topk {
let w = scores[t * ne + indices[t * topk + kk]];
weights[t * topk + kk] = w;
sum += w;
}
for kk in 0..topk {
weights[t * topk + kk] = weights[t * topk + kk] / sum * route_scale;
}
}
(indices, weights)
}
#[allow(clippy::too_many_arguments)]
fn block_forward(
&self,
st: &Stage,
layer: &LayerDev,
h: &CudaSlice<f32>,
s: usize,
ids: &[u32],
mut capture: Option<&mut GpuCapture>,
mut cache: Option<&mut LayerCache>,
) -> Res<CudaSlice<f32>> {
let d = self.model.cfg();
let mc = &self.model.mc;
let hc = d.hc_mult as usize;
let hidden = mc.n_embd as usize;
let heads = mc.n_head as usize;
let hd = d.head_dim as usize;
let rd = d.qk_rope_head_dim as usize;
let q_lora = d.q_lora_rank as usize;
let win = d.sliding_window as usize;
let o_groups = d.o_groups as usize;
let o_lora = d.o_lora_rank as usize;
let eps = mc.rms_eps;
let iters = d.hc_sinkhorn_iters;
let hc_eps = d.hc_eps;
st.gpu.ctx.bind_to_thread().map_err(e("bind ctx"))?;
let stream = st.gpu.stream();
let fc_dev = if layer.ratio != 0 {
&st.fc_yarn
} else {
&st.fc_plain
};
let clamp_only = (self.variant == ActQuantVariant::ClampOnly) as i32;
let (y, post, comb) = Self::hc_pre(
st,
h,
&layer.hc_attn_fn,
&layer.hc_attn_base,
&layer.hc_attn_scale,
s,
hc,
hidden,
iters,
hc_eps,
)?;
let mut x = stream.alloc_zeros::<f32>(s * hidden).map_err(e("x"))?;
unsafe {
ck(
"rmsnorm attn",
k::memra_dsv4_rmsnorm(
dpf!(y, &stream),
dpf!(layer.attn_norm, &stream),
dpm!(x, &stream),
s as i32,
hidden as i32,
eps,
sp(&stream),
),
)?;
}
let wq_a_v = layer.wq_a.staged(&stream)?;
let mut qr = stream.alloc_zeros::<f32>(s * q_lora).map_err(e("qr"))?;
Self::gemm(st, &x, wq_a_v.slab(), 0, s, q_lora, hidden, &mut qr)?;
unsafe {
ck(
"rmsnorm q",
k::memra_dsv4_rmsnorm(
dpf!(qr, &stream),
dpf!(layer.q_norm, &stream),
dpm!(qr, &stream),
s as i32,
q_lora as i32,
eps,
sp(&stream),
),
)?;
}
let mut qr_b = stream
.alloc_zeros::<u8>(s * q_lora * 2)
.map_err(e("qr_b"))?;
unsafe {
ck(
"cvt qr",
k::memra_dsv4_cvt_bf16(
dpf!(qr, &stream),
qr_b.device_ptr_mut(&stream).0 as *mut c_void,
(s * q_lora) as i64,
sp(&stream),
),
)?;
}
let wq_b_v = layer.wq_b.staged(&stream)?;
let mut q = stream.alloc_zeros::<f32>(s * heads * hd).map_err(e("q"))?;
Self::gemm_pre(
st,
&qr_b,
wq_b_v.slab().device_ptr(&stream).0 as *const c_void,
s,
heads * hd,
q_lora,
&mut q,
)?;
let positions: Vec<i32> = (0..s as i32).collect();
let pos_dev = upload_i32(&stream, &positions)?;
unsafe {
ck(
"headrms",
k::memra_dsv4_headrms(
dpm!(q, &stream),
(s * heads) as i32,
hd as i32,
eps,
sp(&stream),
),
)?;
ck(
"rope q",
k::memra_dsv4_rope(
dpm!(q, &stream),
s as i32,
heads as i32,
hd as i32,
rd as i32,
dpf!(fc_dev, &stream),
pos_dev.device_ptr(&stream).0 as *const i32,
0,
sp(&stream),
),
)?;
}
let wkv_v = layer.wkv.staged(&stream)?;
let mut kv = stream.alloc_zeros::<f32>(s * hd).map_err(e("kv"))?;
Self::gemm(st, &x, wkv_v.slab(), 0, s, hd, hidden, &mut kv)?;
unsafe {
ck(
"rmsnorm kv",
k::memra_dsv4_rmsnorm(
dpf!(kv, &stream),
dpf!(layer.kv_norm, &stream),
dpm!(kv, &stream),
s as i32,
hd as i32,
eps,
sp(&stream),
),
)?;
ck(
"rope kv",
k::memra_dsv4_rope(
dpm!(kv, &stream),
s as i32,
1,
hd as i32,
rd as i32,
dpf!(fc_dev, &stream),
pos_dev.device_ptr(&stream).0 as *const i32,
0,
sp(&stream),
),
)?;
ck(
"act_quant kv",
k::memra_dsv4_act_quant(
dpm!(kv, &stream),
s as i32,
hd as i64,
(hd - rd) as i32,
64,
clamp_only,
sp(&stream),
),
)?;
}
if let Some(c) = cache.as_deref_mut() {
for p in s.saturating_sub(win)..s {
let slot = p % win;
let src = kv.slice(p * hd..(p + 1) * hd);
let mut dst = c.kvc.slice_mut(slot * hd..(slot + 1) * hd);
stream.memcpy_dtod(&src, &mut dst).map_err(e("ring copy"))?;
}
}
let (widx, wslots) = window_topk_idxs(win, s);
let mut idxs: Vec<i64> = widx;
let mut slots = wslots;
let mut n_kv = s;
let mut kv_full = kv;
let mut cap_cmp: Option<(Vec<f32>, usize)> = None;
let mut cap_ikv: Option<(Vec<f32>, usize)> = None;
let mut cap_isc: Option<(Vec<f32>, usize)> = None;
let want_cap = capture
.as_ref()
.map(|c| c.want.contains(&layer.il))
.unwrap_or(false);
if layer.ratio != 0 {
let offset = s;
let (cidx, cslots) = if let Some(ix) = &layer.idx {
let mut qi = stream
.alloc_zeros::<f32>(s * ix.heads * ix.hd)
.map_err(e("qi"))?;
let iwq_b_v = ix.wq_b.staged(&stream)?;
Self::gemm_pre(
st,
&qr_b,
iwq_b_v.slab().device_ptr(&stream).0 as *const c_void,
s,
ix.heads * ix.hd,
q_lora,
&mut qi,
)?;
unsafe {
ck(
"rope qi",
k::memra_dsv4_rope(
dpm!(qi, &stream),
s as i32,
ix.heads as i32,
ix.hd as i32,
rd as i32,
dpf!(fc_dev, &stream),
pos_dev.device_ptr(&stream).0 as *const i32,
0,
sp(&stream),
),
)?;
let scale = (ix.hd as f32).powf(-0.5);
ck(
"hadamard qi",
k::memra_dsv4_hadamard(
dpm!(qi, &stream),
(s * ix.heads) as i32,
ix.hd as i32,
scale,
sp(&stream),
),
)?;
ck(
"fp4 qi",
k::memra_dsv4_fp4_act_quant(
dpm!(qi, &stream),
(s * ix.heads) as i32,
ix.hd as i64,
ix.hd as i32,
sp(&stream),
),
)?;
}
let (ckv_i, ikv_raw, isc_raw) =
self.compressor(st, &ix.cmp, &x, s, hidden, fc_dev, rd, eps)?;
if want_cap {
if let Some((buf, nb)) = &ckv_i {
cap_ikv = Some((dtoh_f32(&stream, buf)?, *nb));
}
}
if let Some(c) = cache.as_deref_mut() {
let mut i_blocks = c.i_blocks;
Self::populate_cmp_cache(
&stream,
s,
ix.cmp.ratio,
ix.cmp.latent,
ix.cmp.d,
&ckv_i,
&ikv_raw,
&isc_raw,
c.ikvc.as_mut().expect("fine layer has indexer store"),
0,
&mut i_blocks,
c.ipend_kv.as_mut().expect("ipend"),
c.ipend_score.as_mut().expect("ipend"),
ix.cmp.overlap,
)?;
c.i_blocks = i_blocks;
}
let iwp_v = ix.weights_proj.staged(&stream)?;
let mut wproj = stream.alloc_zeros::<f32>(s * ix.heads).map_err(e("wp"))?;
Self::gemm(st, &x, iwp_v.slab(), 0, s, ix.heads, hidden, &mut wproj)?;
if let Some((ckv, nb)) = &ckv_i {
let wscale = ((ix.hd as f64).powf(-0.5) * (ix.heads as f64).powf(-0.5)) as f32;
let mut score = stream.alloc_zeros::<f32>(s * nb).map_err(e("iscore"))?;
unsafe {
ck(
"indexer_score",
k::memra_dsv4_indexer_score(
dpf!(qi, &stream),
dpf!(ckv, &stream),
dpf!(wproj, &stream),
wscale,
dpm!(score, &stream),
s as i32,
ix.heads as i32,
ix.hd as i32,
*nb as i32,
layer.ratio as i32,
-1, sp(&stream),
),
)?;
}
let score_h = dtoh_f32(&stream, &score)?;
if want_cap {
cap_isc = Some((score_h.clone(), *nb));
}
let kk = ix.topk.min(*nb);
let mut cidx = vec![-1i64; s * kk];
for t in 0..s {
let lim = (t + 1) / layer.ratio;
let mut order: Vec<usize> = (0..*nb).collect();
order.sort_by(|&a, &b| {
score_h[t * nb + b]
.partial_cmp(&score_h[t * nb + a])
.unwrap_or(std::cmp::Ordering::Equal)
.then(a.cmp(&b))
});
for (slot, &j) in order.iter().take(kk).enumerate() {
cidx[t * kk + slot] = if j >= lim { -1 } else { (j + offset) as i64 };
}
}
(cidx, kk)
} else {
(Vec::new(), 0)
}
} else {
compress_topk_idxs(layer.ratio, s, offset)
};
if cslots > 0 {
let mut merged = vec![-1i64; s * (slots + cslots)];
for t in 0..s {
merged[t * (slots + cslots)..t * (slots + cslots) + slots]
.copy_from_slice(&idxs[t * slots..(t + 1) * slots]);
merged[t * (slots + cslots) + slots..(t + 1) * (slots + cslots)]
.copy_from_slice(&cidx[t * cslots..(t + 1) * cslots]);
}
idxs = merged;
slots += cslots;
}
let acmp = layer.cmp.as_ref().expect("ratio!=0 has compressor");
let (ckv, akv_raw, asc_raw) =
self.compressor(st, acmp, &x, s, hidden, fc_dev, rd, eps)?;
if want_cap {
if let Some((buf, nb)) = &ckv {
cap_cmp = Some((dtoh_f32(&stream, buf)?, *nb));
}
}
if let Some(c) = cache.as_deref_mut() {
let mut n_blocks = c.n_blocks;
Self::populate_cmp_cache(
&stream,
s,
acmp.ratio,
acmp.latent,
acmp.d,
&ckv,
&akv_raw,
&asc_raw,
&mut c.kvc,
win,
&mut n_blocks,
c.pend_kv.as_mut().expect("pend"),
c.pend_score.as_mut().expect("pend"),
acmp.overlap,
)?;
c.n_blocks = n_blocks;
}
if let Some((ckv_buf, nb)) = ckv {
let mut merged_kv = stream
.alloc_zeros::<f32>((s + nb) * hd)
.map_err(e("kv_full"))?;
{
let mut head_view = merged_kv.slice_mut(0..s * hd);
stream
.memcpy_dtod(&kv_full.slice(0..s * hd), &mut head_view)
.map_err(e("kv copy"))?;
}
{
let mut tail = merged_kv.slice_mut(s * hd..(s + nb) * hd);
stream
.memcpy_dtod(&ckv_buf.slice(0..nb * hd), &mut tail)
.map_err(e("ckv copy"))?;
}
kv_full = merged_kv;
n_kv += nb;
}
}
let _ = n_kv;
let idxs_i32: Vec<i32> = idxs.iter().map(|&v| v as i32).collect();
let idx_dev = upload_i32(&stream, &idxs_i32)?;
let mut o = stream.alloc_zeros::<f32>(s * heads * hd).map_err(e("o"))?;
let scale = (hd as f64).powf(-0.5) as f32;
unsafe {
ck(
"sink_attn",
k::memra_dsv4_sink_attn(
dpf!(q, &stream),
dpf!(kv_full, &stream),
idx_dev.device_ptr(&stream).0 as *const i32,
dpf!(layer.sink, &stream),
dpm!(o, &stream),
s as i32,
heads as i32,
hd as i32,
slots as i32,
scale,
sp(&stream),
),
)?;
ck(
"rope o inv",
k::memra_dsv4_rope(
dpm!(o, &stream),
s as i32,
heads as i32,
hd as i32,
rd as i32,
dpf!(fc_dev, &stream),
pos_dev.device_ptr(&stream).0 as *const i32,
1,
sp(&stream),
),
)?;
}
let gw = heads / o_groups * hd;
let mut og = stream
.alloc_zeros::<f32>(s * o_groups * o_lora)
.map_err(e("og"))?;
let mut o_grp = stream.alloc_zeros::<f32>(s * gw).map_err(e("o_grp"))?;
let mut y_grp = stream.alloc_zeros::<f32>(s * o_lora).map_err(e("y_grp"))?;
let wo_a_v = layer.wo_a.staged(&stream)?; for g in 0..o_groups {
unsafe {
ck(
"take_cols",
k::memra_dsv4_take_cols(
dpf!(o, &stream),
dpm!(o_grp, &stream),
s as i32,
gw as i32,
(heads * hd) as i64,
(g * gw) as i64,
sp(&stream),
),
)?;
}
Self::gemm(
st,
&o_grp,
wo_a_v.slab(),
g * o_lora * gw,
s,
o_lora,
gw,
&mut y_grp,
)?;
unsafe {
ck(
"place_cols",
k::memra_dsv4_place_cols(
dpf!(y_grp, &stream),
dpm!(og, &stream),
s as i32,
o_lora as i32,
(o_groups * o_lora) as i64,
(g * o_lora) as i64,
sp(&stream),
),
)?;
}
}
let wo_b_v = layer.wo_b.staged(&stream)?;
let mut attn_out = stream.alloc_zeros::<f32>(s * hidden).map_err(e("ao"))?;
Self::gemm(
st,
&og,
wo_b_v.slab(),
0,
s,
hidden,
o_groups * o_lora,
&mut attn_out,
)?;
let mut cap_attn: Option<Vec<f32>> = None;
if want_cap {
cap_attn = Some(dtoh_f32(&stream, &attn_out)?);
}
let mut h2 = stream
.alloc_zeros::<f32>(s * hc * hidden)
.map_err(e("h2"))?;
unsafe {
ck(
"hc_post attn",
k::memra_dsv4_hc_post(
dpf!(attn_out, &stream),
dpf!(h, &stream),
dpf!(post, &stream),
dpf!(comb, &stream),
dpm!(h2, &stream),
s as i32,
hc as i32,
hidden as i32,
sp(&stream),
),
)?;
}
let (y2, post2, comb2) = Self::hc_pre(
st,
&h2,
&layer.hc_ffn_fn,
&layer.hc_ffn_base,
&layer.hc_ffn_scale,
s,
hc,
hidden,
iters,
hc_eps,
)?;
let mut xf = stream.alloc_zeros::<f32>(s * hidden).map_err(e("xf"))?;
unsafe {
ck(
"rmsnorm ffn",
k::memra_dsv4_rmsnorm(
dpf!(y2, &stream),
dpf!(layer.ffn_norm, &stream),
dpm!(xf, &stream),
s as i32,
hidden as i32,
eps,
sp(&stream),
),
)?;
}
if let Some(c) = capture.as_deref_mut() {
if c.want.contains(&layer.il) {
c.moe_x.insert(layer.il, dtoh_f32(&stream, &xf)?);
}
}
let moe_out = self.moe_forward(st, layer, &xf, s, ids)?;
let mut h3 = stream
.alloc_zeros::<f32>(s * hc * hidden)
.map_err(e("h3"))?;
unsafe {
ck(
"hc_post ffn",
k::memra_dsv4_hc_post(
dpf!(moe_out, &stream),
dpf!(h2, &stream),
dpf!(post2, &stream),
dpf!(comb2, &stream),
dpm!(h3, &stream),
s as i32,
hc as i32,
hidden as i32,
sp(&stream),
),
)?;
}
if let Some(c) = capture {
if c.want.contains(&layer.il) {
c.layer_out.insert(layer.il, dtoh_f32(&stream, &h3)?);
c.x_dbg.insert(layer.il, dtoh_f32(&stream, &x)?);
c.q_dbg.insert(layer.il, dtoh_f32(&stream, &q)?);
{
let mut kvh = vec![0f32; s * hd];
stream
.memcpy_dtoh(&kv_full.slice(0..s * hd), &mut kvh[..])
.map_err(e("dtoh kv_dbg"))?;
stream.synchronize().map_err(e("sync kv_dbg"))?;
c.kv_dbg.insert(layer.il, kvh);
}
c.o_dbg.insert(layer.il, dtoh_f32(&stream, &o)?);
if let Some(a) = cap_attn {
c.attn_out.insert(layer.il, a);
}
if let Some(v) = cap_cmp {
c.compressor_kv.insert(layer.il, v);
}
if let Some(v) = cap_ikv {
c.indexer_kv.insert(layer.il, v);
}
if let Some(v) = cap_isc {
c.index_score.insert(layer.il, v);
}
}
}
Ok(h3)
}
fn moe_forward(
&self,
st: &Stage,
layer: &LayerDev,
x: &CudaSlice<f32>, s: usize,
ids: &[u32],
) -> Res<CudaSlice<f32>> {
let mc = &self.model.mc;
let d = self.model.cfg();
let moe = mc.moe.as_ref().expect("moe");
let hidden = mc.n_embd as usize;
let ne = moe.expert_count as usize;
let topk = moe.expert_used_count as usize;
let inter = moe.expert_ff_length as usize;
let limit = d.swiglu_limit;
let stream = st.gpu.stream();
let mut raw = stream.alloc_zeros::<f32>(s * ne).map_err(e("gate raw"))?;
Self::dots(st, x, &layer.gate_w, s, hidden, ne, &mut raw)?;
let raw_h = dtoh_f32(&stream, &raw)?;
let (indices, weights) =
Self::route_host(layer, &raw_h, ids, s, ne, topk, d.routed_scaling_factor);
let mut xb = stream
.alloc_zeros::<u8>(s * hidden * 2)
.map_err(e("xb moe"))?;
unsafe {
ck(
"cvt moe x",
k::memra_dsv4_cvt_bf16(
dpf!(x, &stream),
xb.device_ptr_mut(&stream).0 as *mut c_void,
(s * hidden) as i64,
sp(&stream),
),
)?;
}
let mut y = stream.alloc_zeros::<f32>(s * hidden).map_err(e("moe y"))?;
let wbytes = inter * hidden / 2;
let sbytes = match layer.expert_kind {
ExpertKind::Nvfp4 => inter * hidden / 16,
ExpertKind::Mxfp4 => inter * hidden / 32,
};
let mut uniq: Vec<usize> = indices.clone();
uniq.sort_unstable();
uniq.dedup();
if self.expert_arm == ExpertArm::Native {
let kind = match layer.expert_kind {
ExpertKind::Nvfp4 => 0i32,
ExpertKind::Mxfp4 => 1i32,
};
let kq_x = hidden / 128;
let kq_h = inter / 128;
let mut xq = stream.alloc_zeros::<u8>(s * hidden).map_err(e("xq"))?;
let mut xs = stream.alloc_zeros::<f32>(s * kq_x).map_err(e("xs"))?;
unsafe {
ck(
"act_quant_fp8 x",
k::memra_dsv4_act_quant_fp8(
dpf!(x, &stream),
xq.device_ptr_mut(&stream).0 as *mut c_void,
dpm!(xs, &stream),
s as i32,
hidden as i32,
sp(&stream),
),
)?;
}
let mut xgq = stream.alloc_zeros::<u8>(s * hidden).map_err(e("xgq"))?;
let mut xgs = stream.alloc_zeros::<f32>(s * kq_x).map_err(e("xgs"))?;
let mut g1 = stream.alloc_zeros::<f32>(s * inter).map_err(e("g1"))?;
let mut g3 = stream.alloc_zeros::<f32>(s * inter).map_err(e("g3"))?;
let mut hbuf = stream.alloc_zeros::<f32>(s * inter).map_err(e("hbuf"))?;
let mut hq = stream.alloc_zeros::<u8>(s * inter).map_err(e("hq"))?;
let mut hs = stream.alloc_zeros::<f32>(s * kq_h).map_err(e("hs"))?;
let mut contrib = stream
.alloc_zeros::<f32>(s * hidden)
.map_err(e("contrib"))?;
for &ex in &uniq {
let toks: Vec<(usize, usize)> = (0..s * topk)
.filter(|i| indices[*i] == ex)
.map(|i| (i / topk, i % topk))
.collect();
let g = toks.len();
let tok_rows: Vec<i32> = toks.iter().map(|&(t, _)| t as i32).collect();
let wrow: Vec<f32> = toks.iter().map(|&(t, kk)| weights[t * topk + kk]).collect();
let rows_dev = upload_i32(&stream, &tok_rows)?;
let wrow_dev = upload_f32(&stream, &wrow)?;
unsafe {
ck(
"gather xq",
k::memra_dsv4_gather_rows_u8(
dp!(xq, &stream),
rows_dev.device_ptr(&stream).0 as *const i32,
xgq.device_ptr_mut(&stream).0 as *mut c_void,
g as i32,
hidden as i64,
sp(&stream),
),
)?;
ck(
"gather xs",
k::memra_dsv4_gather_rows_u8(
xs.device_ptr(&stream).0 as *const c_void,
rows_dev.device_ptr(&stream).0 as *const i32,
xgs.device_ptr_mut(&stream).0 as *mut c_void,
g as i32,
(kq_x * 4) as i64,
sp(&stream),
),
)?;
for (pi, dst) in [(0usize, &mut g1), (2usize, &mut g3)] {
let woff = (ex * 3 + pi) * wbytes;
let soff = (ex * 3 + pi) * sbytes;
ck(
"fp4_gemm w1/w3",
k::memra_dsv4_fp4_gemm(
dp!(xgq, &stream),
dpf!(xgs, &stream),
(layer.experts_w.device_ptr(&stream).0 as usize + woff)
as *const c_void,
(layer.experts_sc.device_ptr(&stream).0 as usize + soff)
as *const c_void,
layer.experts_s2[ex * 3 + pi],
kind,
dpm!(*dst, &stream),
g as i32,
inter as i32,
hidden as i32,
sp(&stream),
),
)?;
}
ck(
"swiglu",
k::memra_dsv4_swiglu(
dpf!(g1, &stream),
dpf!(g3, &stream),
dpm!(hbuf, &stream),
g as i32,
inter as i32,
limit,
wrow_dev.device_ptr(&stream).0 as *const f32,
sp(&stream),
),
)?;
ck(
"act_quant_fp8 h",
k::memra_dsv4_act_quant_fp8(
dpf!(hbuf, &stream),
hq.device_ptr_mut(&stream).0 as *mut c_void,
dpm!(hs, &stream),
g as i32,
inter as i32,
sp(&stream),
),
)?;
let woff2 = (ex * 3 + 1) * wbytes;
let soff2 = (ex * 3 + 1) * sbytes;
ck(
"fp4_gemm w2",
k::memra_dsv4_fp4_gemm(
dp!(hq, &stream),
dpf!(hs, &stream),
(layer.experts_w.device_ptr(&stream).0 as usize + woff2)
as *const c_void,
(layer.experts_sc.device_ptr(&stream).0 as usize + soff2)
as *const c_void,
layer.experts_s2[ex * 3 + 1],
kind,
dpm!(contrib, &stream),
g as i32,
hidden as i32,
inter as i32,
sp(&stream),
),
)?;
ck(
"scatter",
k::memra_dsv4_scatter_add(
dpm!(y, &stream),
dpf!(contrib, &stream),
rows_dev.device_ptr(&stream).0 as *const i32,
g as i32,
hidden as i32,
sp(&stream),
),
)?;
}
}
return self.moe_shared_and_finish(st, layer, &xb, s, y);
}
let mut xg = stream.alloc_zeros::<u8>(s * hidden * 2).map_err(e("xg"))?;
let mut g1 = stream.alloc_zeros::<f32>(s * inter).map_err(e("g1"))?;
let mut g3 = stream.alloc_zeros::<f32>(s * inter).map_err(e("g3"))?;
let mut hbuf = stream.alloc_zeros::<f32>(s * inter).map_err(e("hbuf"))?;
let mut hb = stream.alloc_zeros::<u8>(s * inter * 2).map_err(e("hb"))?;
let mut contrib = stream
.alloc_zeros::<f32>(s * hidden)
.map_err(e("contrib"))?;
for &ex in &uniq {
let toks: Vec<(usize, usize)> = (0..s * topk)
.filter(|i| indices[*i] == ex)
.map(|i| (i / topk, i % topk))
.collect();
let g = toks.len();
let tok_rows: Vec<i32> = toks.iter().map(|&(t, _)| t as i32).collect();
let wrow: Vec<f32> = toks.iter().map(|&(t, kk)| weights[t * topk + kk]).collect();
let rows_dev = upload_i32(&stream, &tok_rows)?;
let wrow_dev = upload_f32(&stream, &wrow)?;
unsafe {
ck(
"gather",
k::memra_dsv4_gather_bf16(
dp!(xb, &stream),
rows_dev.device_ptr(&stream).0 as *const i32,
xg.device_ptr_mut(&stream).0 as *mut c_void,
g as i32,
hidden as i32,
sp(&stream),
),
)?;
for (pi, (rows, cols)) in [(inter, hidden), (hidden, inter), (inter, hidden)]
.iter()
.enumerate()
{
let woff = (ex * 3 + pi) * wbytes;
let soff = (ex * 3 + pi) * sbytes;
let wp =
(layer.experts_w.device_ptr(&stream).0 as usize + woff) as *const c_void;
let scp =
(layer.experts_sc.device_ptr(&stream).0 as usize + soff) as *const c_void;
let dst = st.deq[pi].device_ptr(&stream).0 as *mut c_void;
match layer.expert_kind {
ExpertKind::Nvfp4 => ck(
"nvfp4 deq",
k::memra_dsv4_nvfp4_deq_bf16(
wp,
scp,
layer.experts_s2[ex * 3 + pi],
*rows as i32,
*cols as i32,
dst,
sp(&stream),
),
)?,
ExpertKind::Mxfp4 => ck(
"mxfp4 deq",
k::memra_dsv4_mxfp4_deq_bf16(
wp,
scp,
*rows as i32,
*cols as i32,
dst,
sp(&stream),
),
)?,
}
}
ck(
"gemm w1",
k::memra_dsv4_gemm_bf16(
st.deq[0].device_ptr(&stream).0 as *const c_void,
dp!(xg, &stream),
dpm!(g1, &stream),
g as i32,
inter as i32,
hidden as i32,
st.dev as i32,
st.ws.device_ptr(&stream).0 as *mut c_void,
st.ws.len(),
sp(&stream),
),
)?;
ck(
"gemm w3",
k::memra_dsv4_gemm_bf16(
st.deq[2].device_ptr(&stream).0 as *const c_void,
dp!(xg, &stream),
dpm!(g3, &stream),
g as i32,
inter as i32,
hidden as i32,
st.dev as i32,
st.ws.device_ptr(&stream).0 as *mut c_void,
st.ws.len(),
sp(&stream),
),
)?;
ck(
"swiglu",
k::memra_dsv4_swiglu(
dpf!(g1, &stream),
dpf!(g3, &stream),
dpm!(hbuf, &stream),
g as i32,
inter as i32,
limit,
wrow_dev.device_ptr(&stream).0 as *const f32,
sp(&stream),
),
)?;
ck(
"cvt h",
k::memra_dsv4_cvt_bf16(
dpf!(hbuf, &stream),
hb.device_ptr_mut(&stream).0 as *mut c_void,
(g * inter) as i64,
sp(&stream),
),
)?;
ck(
"gemm w2",
k::memra_dsv4_gemm_bf16(
st.deq[1].device_ptr(&stream).0 as *const c_void,
dp!(hb, &stream),
dpm!(contrib, &stream),
g as i32,
hidden as i32,
inter as i32,
st.dev as i32,
st.ws.device_ptr(&stream).0 as *mut c_void,
st.ws.len(),
sp(&stream),
),
)?;
ck(
"scatter",
k::memra_dsv4_scatter_add(
dpm!(y, &stream),
dpf!(contrib, &stream),
rows_dev.device_ptr(&stream).0 as *const i32,
g as i32,
hidden as i32,
sp(&stream),
),
)?;
}
}
self.moe_shared_and_finish(st, layer, &xb, s, y)
}
fn moe_shared_and_finish(
&self,
st: &Stage,
layer: &LayerDev,
xb: &CudaSlice<u8>,
s: usize,
mut y: CudaSlice<f32>,
) -> Res<CudaSlice<f32>> {
let d = self.model.cfg();
let hidden = self.model.mc.n_embd as usize;
let limit = d.swiglu_limit;
let stream = st.gpu.stream();
let sh_inter = {
let (shape, _) = self
.model
.st
.raw("layers.0.ffn.shared_experts.w1.weight")
.map(|(i, _)| (i.shape.clone(), ()))
.expect("shared w1");
shape[0] as usize
};
let mut sg1 = stream.alloc_zeros::<f32>(s * sh_inter).map_err(e("sg1"))?;
let mut sg3 = stream.alloc_zeros::<f32>(s * sh_inter).map_err(e("sg3"))?;
let mut shbuf = stream.alloc_zeros::<f32>(s * sh_inter).map_err(e("shb"))?;
let mut shb16 = stream
.alloc_zeros::<u8>(s * sh_inter * 2)
.map_err(e("shb16"))?;
let mut sh_out = stream.alloc_zeros::<f32>(s * hidden).map_err(e("sh_out"))?;
let sw = [
layer.shared_w[0].staged(&stream)?,
layer.shared_w[1].staged(&stream)?,
layer.shared_w[2].staged(&stream)?,
];
Self::gemm_pre(
st,
xb,
sw[0].slab().device_ptr(&stream).0 as *const c_void,
s,
sh_inter,
hidden,
&mut sg1,
)?;
Self::gemm_pre(
st,
xb,
sw[2].slab().device_ptr(&stream).0 as *const c_void,
s,
sh_inter,
hidden,
&mut sg3,
)?;
unsafe {
ck(
"swiglu sh",
k::memra_dsv4_swiglu(
dpf!(sg1, &stream),
dpf!(sg3, &stream),
dpm!(shbuf, &stream),
s as i32,
sh_inter as i32,
limit,
std::ptr::null(),
sp(&stream),
),
)?;
ck(
"cvt sh",
k::memra_dsv4_cvt_bf16(
dpf!(shbuf, &stream),
shb16.device_ptr_mut(&stream).0 as *mut c_void,
(s * sh_inter) as i64,
sp(&stream),
),
)?;
}
Self::gemm_pre(
st,
&shb16,
sw[1].slab().device_ptr(&stream).0 as *const c_void,
s,
hidden,
sh_inter,
&mut sh_out,
)?;
unsafe {
ck(
"add shared",
k::memra_dsv4_add_inplace(
dpm!(y, &stream),
dpf!(sh_out, &stream),
(s * hidden) as i64,
sp(&stream),
),
)?;
}
Ok(y)
}
pub fn forward(
&self,
ids: &[u32],
capture: Option<&mut GpuCapture>,
early_exit_after: Option<u32>,
) -> Res<Option<ForwardOut>> {
self.forward_impl(ids, capture, early_exit_after, None)
}
pub fn prefill_with_cache(&self, ids: &[u32], state: &mut DecodeState) -> Res<ForwardOut> {
assert_eq!(state.pos, 0, "prefill_with_cache needs a fresh DecodeState");
assert!(!ids.is_empty(), "empty prompt");
let out = self
.forward_impl(ids, None, None, Some(state))?
.expect("prefill logits");
state.pos = ids.len();
Ok(out)
}
fn forward_impl(
&self,
ids: &[u32],
mut capture: Option<&mut GpuCapture>,
early_exit_after: Option<u32>,
mut state: Option<&mut DecodeState>,
) -> Res<Option<ForwardOut>> {
let mc = &self.model.mc;
let d = self.model.cfg();
let s = ids.len();
assert!(s <= self.max_seq, "seq {s} > max_seq {}", self.max_seq);
let hidden = mc.n_embd as usize;
let hc = d.hc_mult as usize;
let n_trunk = mc.n_layer - mc.nextn_predict_layers;
let st0 = &self.stages[0];
st0.gpu.ctx.bind_to_thread().map_err(e("bind ctx0"))?;
let stream0 = st0.gpu.stream();
let ids_i32: Vec<i32> = ids.iter().map(|&x| x as i32).collect();
let ids_dev = upload_i32(&stream0, &ids_i32)?;
let mut emb = stream0.alloc_zeros::<f32>(s * hidden).map_err(e("emb"))?;
unsafe {
ck(
"embed_rows",
k::memra_dsv4_embed_rows(
st0.embed
.as_ref()
.expect("embed on stage 0")
.device_ptr(&stream0)
.0 as *const c_void,
ids_dev.device_ptr(&stream0).0 as *const i32,
dpm!(emb, &stream0),
s as i32,
hidden as i32,
sp(&stream0),
),
)?;
}
if let Some(c) = capture.as_deref_mut() {
if c.embed_out.is_none() {
c.embed_out = Some(dtoh_f32(&stream0, &emb)?);
}
}
let mut h = stream0
.alloc_zeros::<f32>(s * hc * hidden)
.map_err(e("h0"))?;
unsafe {
ck(
"repeat_hc",
k::memra_dsv4_repeat_hc(
dpf!(emb, &stream0),
dpm!(h, &stream0),
s as i32,
hc as i32,
hidden as i32,
sp(&stream0),
),
)?;
}
let mut cur_stage = 0usize;
for il in 0..n_trunk {
let stage = self.layer_stage[il as usize];
if stage != cur_stage {
let src_stream = self.stages[cur_stage].gpu.stream();
let host = dtoh_f32(&src_stream, &h)?;
let dst_stream = self.stages[stage].gpu.stream();
self.stages[stage]
.gpu
.ctx
.bind_to_thread()
.map_err(e("bind"))?;
h = upload_f32(&dst_stream, &host)?;
cur_stage = stage;
}
let st = &self.stages[stage];
let lidx = st
.layers
.iter()
.position(|l| l.il == il)
.unwrap_or_else(|| panic!("layer {il} not on stage {stage}"));
let layer_cache = state.as_deref_mut().map(|ds| &mut ds.caches[il as usize]);
h = self.block_forward(
st,
&st.layers[lidx],
&h,
s,
ids,
capture.as_deref_mut(),
layer_cache,
)?;
if early_exit_after == Some(il) {
self.stages[cur_stage]
.gpu
.stream()
.synchronize()
.map_err(e("sync"))?;
return Ok(None);
}
}
let last = self.stages.len() - 1;
if cur_stage != last {
let src_stream = self.stages[cur_stage].gpu.stream();
let host = dtoh_f32(&src_stream, &h)?;
let dst_stream = self.stages[last].gpu.stream();
h = upload_f32(&dst_stream, &host)?;
}
let hc_head_fn = self.stages[last].hc_head_fn.as_ref().expect("hc_head_fn");
let trunk_norm = self.stages[last].trunk_norm.as_ref().expect("trunk norm");
let logits = self.head_logits_from(
&h,
s,
hc_head_fn,
&self.hc_head_base,
&self.hc_head_scale,
trunk_norm,
)?;
Ok(Some(ForwardOut { logits, h_last: h }))
}
fn head_logits_from(
&self,
h: &CudaSlice<f32>,
s: usize,
fn_w: &CudaSlice<f32>,
base: &[f32],
scale: &[f32],
norm: &CudaSlice<f32>,
) -> Res<Vec<f32>> {
self.head_logits_row(h, s, s - 1, fn_w, base, scale, norm)
}
#[allow(clippy::too_many_arguments)]
fn head_logits_row(
&self,
h: &CudaSlice<f32>,
s: usize,
row: usize,
fn_w: &CudaSlice<f32>,
base: &[f32],
scale: &[f32],
norm: &CudaSlice<f32>,
) -> Res<Vec<f32>> {
let d = self.model.cfg();
let mc = &self.model.mc;
let hc = d.hc_mult as usize;
let hidden = mc.n_embd as usize;
let eps = mc.rms_eps;
let last = self.stages.len() - 1;
let st = &self.stages[last];
st.gpu.ctx.bind_to_thread().map_err(e("bind ctx head"))?;
let stream = st.gpu.stream();
let w = hc * hidden;
let mut mixes = stream.alloc_zeros::<f32>(s * hc).map_err(e("hm"))?;
Self::dots(st, h, fn_w, s, w, hc, &mut mixes)?;
unsafe {
ck(
"rowsq head",
k::memra_dsv4_rowsq_scale(
dpf!(h, &stream),
dpm!(mixes, &stream),
s as i32,
w as i32,
hc as i32,
eps,
sp(&stream),
),
)?;
}
let mut mixes_h = dtoh_f32(&stream, &mixes)?;
for t in 0..s {
for c in 0..hc {
let m = mixes_h[t * hc + c];
mixes_h[t * hc + c] = sigmoid_f32(m * scale[0] + base[c]) + d.hc_eps;
}
}
let pre_d = upload_f32(&stream, &mixes_h)?;
let mut collapsed = stream.alloc_zeros::<f32>(s * hidden).map_err(e("col"))?;
unsafe {
ck(
"hc_collapse head",
k::memra_dsv4_hc_collapse(
dpf!(h, &stream),
dpf!(pre_d, &stream),
dpm!(collapsed, &stream),
s as i32,
hc as i32,
hidden as i32,
sp(&stream),
),
)?;
ck(
"rmsnorm head",
k::memra_dsv4_rmsnorm(
dpf!(collapsed, &stream),
dpf!(norm, &stream),
dpm!(collapsed, &stream),
s as i32,
hidden as i32,
eps,
sp(&stream),
),
)?;
}
assert!(row < s, "logits row {row} out of range (s = {s})");
let vocab = {
let (info, _) = self.model.st.raw("head.weight").expect("head");
info.shape[0] as usize
};
let last_row = collapsed.slice(row * hidden..(row + 1) * hidden);
let mut logits = stream.alloc_zeros::<f32>(vocab).map_err(e("logits"))?;
unsafe {
ck(
"head dots",
k::memra_dsv4_dots_f32(
last_row.device_ptr(&stream).0 as *const f32,
st.head.as_ref().expect("head").device_ptr(&stream).0 as *const c_void,
1,
dpm!(logits, &stream),
1,
hidden as i32,
vocab as i32,
sp(&stream),
),
)?;
}
dtoh_f32(&stream, &logits)
}
pub fn mtp_logits_last(&self, h_trunk: &CudaSlice<f32>, ids: &[u32]) -> Res<Vec<f32>> {
self.mtp_logits_last_cap(h_trunk, ids, None)
}
pub fn mtp_logits_last_cap(
&self,
h_trunk: &CudaSlice<f32>,
ids: &[u32],
capture: Option<&mut GpuCapture>,
) -> Res<Vec<f32>> {
let mtp = self.mtp.as_ref().expect("MTP not loaded");
let d = self.model.cfg();
let mc = &self.model.mc;
let hc = d.hc_mult as usize;
let hidden = mc.n_embd as usize;
let eps = mc.rms_eps;
let s = ids.len();
let last = self.stages.len() - 1;
let st = &self.stages[last];
st.gpu.ctx.bind_to_thread().map_err(e("bind ctx mtp"))?;
let stream = st.gpu.stream();
let e_host = self.model.embed_rows(ids);
let mut e_dev = upload_f32(&stream, &e_host)?;
unsafe {
ck(
"rmsnorm enorm",
k::memra_dsv4_rmsnorm(
dpf!(e_dev, &stream),
dpf!(mtp.enorm, &stream),
dpm!(e_dev, &stream),
s as i32,
hidden as i32,
eps,
sp(&stream),
),
)?;
}
let mut xh = stream
.alloc_zeros::<f32>(s * hc * hidden)
.map_err(e("mtp xh"))?;
unsafe {
ck(
"rmsnorm hnorm",
k::memra_dsv4_rmsnorm(
dpf!(h_trunk, &stream),
dpf!(mtp.hnorm, &stream),
dpm!(xh, &stream),
(s * hc) as i32,
hidden as i32,
eps,
sp(&stream),
),
)?;
}
let mut ep = stream.alloc_zeros::<f32>(s * hidden).map_err(e("mtp ep"))?;
Self::gemm(st, &e_dev, &mtp.e_proj, 0, s, hidden, hidden, &mut ep)?;
let mut hp = stream
.alloc_zeros::<f32>(s * hc * hidden)
.map_err(e("mtp hp"))?;
Self::gemm(st, &xh, &mtp.h_proj, 0, s * hc, hidden, hidden, &mut hp)?;
let mut xm = stream
.alloc_zeros::<f32>(s * hc * hidden)
.map_err(e("mtp xm"))?;
unsafe {
ck(
"repeat ep",
k::memra_dsv4_repeat_hc(
dpf!(ep, &stream),
dpm!(xm, &stream),
s as i32,
hc as i32,
hidden as i32,
sp(&stream),
),
)?;
ck(
"add hp",
k::memra_dsv4_add_inplace(
dpm!(xm, &stream),
dpf!(hp, &stream),
(s * hc * hidden) as i64,
sp(&stream),
),
)?;
}
let xm = self.block_forward(st, &mtp.layer, &xm, s, ids, capture, None)?;
self.head_logits_from(
&xm,
s,
&mtp.hc_head_fn,
&mtp.hc_head_base,
&mtp.hc_head_scale,
&mtp.norm,
)
}
pub fn trunk_logits_row(&self, h: &CudaSlice<f32>, s: usize, row: usize) -> Res<Vec<f32>> {
let last = self.stages.len() - 1;
let hc_head_fn = self.stages[last].hc_head_fn.as_ref().expect("hc_head_fn");
let trunk_norm = self.stages[last].trunk_norm.as_ref().expect("trunk norm");
self.head_logits_row(
h,
s,
row,
hc_head_fn,
&self.hc_head_base,
&self.hc_head_scale,
trunk_norm,
)
}
pub fn alloc_decode_state(&self) -> Res<DecodeState> {
let d = self.model.cfg();
let mc = &self.model.mc;
let win = d.sliding_window as usize;
let hd = d.head_dim as usize;
let n_trunk = mc.n_layer - mc.nextn_predict_layers;
let mut caches = Vec::with_capacity(n_trunk as usize);
let mut cache_bytes = vec![0u64; self.stages.len()];
let trans_rows = self.verify_tmax();
for il in 0..n_trunk {
let stage_i = self.layer_stage[il as usize];
let st = &self.stages[stage_i];
st.gpu.ctx.bind_to_thread().map_err(e("bind ctx cache"))?;
let stream = st.gpu.stream();
let lidx = st
.layers
.iter()
.position(|l| l.il == il)
.unwrap_or_else(|| panic!("layer {il} not on stage {stage_i}"));
let layer = &st.layers[lidx];
let ratio = layer.ratio;
let cap_blocks = if ratio != 0 { self.max_seq / ratio } else { 0 };
let kvc_rows = win + cap_blocks + trans_rows;
let mut bytes = (kvc_rows * hd * 4) as u64;
let kvc = stream
.alloc_zeros::<f32>(kvc_rows * hd)
.map_err(e("kvc alloc"))?;
let mk_pend = |latent: usize, slots: usize| -> Res<(CudaSlice<f32>, CudaSlice<f32>)> {
let kv = stream
.alloc_zeros::<f32>(slots * latent)
.map_err(e("pend kv alloc"))?;
let sc = upload_f32(&stream, &vec![f32::NEG_INFINITY; slots * latent])?;
Ok((kv, sc))
};
let (pend_kv, pend_score) = if let Some(cmp) = &layer.cmp {
let slots = if cmp.overlap {
2 * cmp.ratio
} else {
cmp.ratio
};
bytes += (2 * slots * cmp.latent * 4) as u64;
let (a, b) = mk_pend(cmp.latent, slots)?;
(Some(a), Some(b))
} else {
(None, None)
};
let (ikvc, ipend_kv, ipend_score) = if let Some(ix) = &layer.idx {
bytes += (cap_blocks * ix.cmp.d * 4) as u64;
let store = stream
.alloc_zeros::<f32>(cap_blocks * ix.cmp.d)
.map_err(e("ikvc alloc"))?;
let slots = if ix.cmp.overlap {
2 * ix.cmp.ratio
} else {
ix.cmp.ratio
};
bytes += (2 * slots * ix.cmp.latent * 4) as u64;
let (a, b) = mk_pend(ix.cmp.latent, slots)?;
(Some(store), Some(a), Some(b))
} else {
(None, None, None)
};
cache_bytes[stage_i] += bytes;
caches.push(LayerCache {
kvc,
n_blocks: 0,
pend_kv,
pend_score,
ikvc,
i_blocks: 0,
ipend_kv,
ipend_score,
});
}
let ws = if matches!(self.decode_path, DecodePath::Device { .. }) {
Some(self.alloc_step_ws()?)
} else {
None
};
for st in &self.stages {
st.gpu.stream().synchronize().map_err(e("cache sync"))?;
}
Ok(DecodeState {
caches,
pos: 0,
cache_bytes,
ws,
})
}
fn alloc_step_ws(&self) -> Res<Vec<StepWs>> {
let d = self.model.cfg();
let mc = &self.model.mc;
let moe = mc.moe.as_ref().expect("moe");
let hc = d.hc_mult as usize;
let hidden = mc.n_embd as usize;
let heads = mc.n_head as usize;
let hd = d.head_dim as usize;
let q_lora = d.q_lora_rank as usize;
let win = d.sliding_window as usize;
let o_groups = d.o_groups as usize;
let o_lora = d.o_lora_rank as usize;
let iheads = d.index_n_heads as usize;
let ihd = d.index_head_dim as usize;
let topk = moe.expert_used_count as usize;
let ne = moe.expert_count as usize;
let inter = moe.expert_ff_length as usize;
let itopk = d.index_topk as usize;
let vocab = {
let (info, _) = self.model.st.raw("head.weight").expect("head");
info.shape[0] as usize
};
let sh_inter = {
let (info, _) = self
.model
.st
.raw("layers.0.ffn.shared_experts.w1.weight")
.expect("shared w1");
info.shape[0] as usize
};
let mut max_latent = 0usize;
let mut max_d = 0usize;
let mut max_shift = 0usize;
let mut min_ratio = usize::MAX;
for st in &self.stages {
for l in &st.layers {
for cmp in l.cmp.iter().chain(l.idx.as_ref().map(|ix| &ix.cmp)) {
max_latent = max_latent.max(cmp.latent);
max_d = max_d.max(cmp.d);
if cmp.overlap {
max_shift = max_shift.max(cmp.ratio * cmp.latent);
}
min_ratio = min_ratio.min(cmp.ratio);
}
}
}
assert!(min_ratio != usize::MAX, "no compressor layers?");
let score_cap = self.max_seq / min_ratio + 1;
let idx_tail = itopk.max(self.max_seq / 128 + 1);
let max_gemm_k = (o_groups * o_lora).max(hidden).max(q_lora).max(sh_inter);
let mut out = Vec::with_capacity(self.stages.len());
for st in &self.stages {
st.gpu.ctx.bind_to_thread().map_err(e("bind ctx ws"))?;
let s = st.gpu.stream();
let f = |n: usize| s.alloc_zeros::<f32>(n).map_err(e("ws f32"));
let b = |n: usize| s.alloc_zeros::<u8>(n).map_err(e("ws u8"));
let i = |n: usize| s.alloc_zeros::<i32>(n).map_err(e("ws i32"));
out.push(StepWs {
h_a: f(hc * hidden)?,
h_b: f(hc * hidden)?,
h_rx: f(hc * hidden)?,
emb: f(hidden)?,
mixes: f((2 + hc) * hc)?,
pre: f(hc)?,
post: f(hc)?,
comb: f(hc * hc)?,
y_hc: f(hidden)?,
x: f(hidden)?,
xf: f(hidden)?,
qr: f(q_lora)?,
qr_b: b(q_lora * 2)?,
q: f(heads * hd)?,
kv: f(hd)?,
qi: f(iheads * ihd)?,
wproj: f(iheads)?,
score: f(score_cap)?,
idx: i(win + idx_tail)?,
o: f(heads * hd)?,
o_b: b(heads * hd * 2)?,
og: f(o_groups * o_lora)?,
attn_out: f(hidden)?,
gemm_xb: b(max_gemm_k * 2)?,
raw: f(ne)?,
sel: i(topk)?,
selw: f(topk)?,
order: i(topk)?,
xq: b(hidden)?,
xs: f(hidden / 128)?,
g1: f(topk * inter)?,
g3: f(topk * inter)?,
hbuf: f(topk * inter)?,
hq: b(topk * inter)?,
hs: f(topk * inter / 128)?,
contrib: f(topk * hidden)?,
y: f(hidden)?,
xb: b(hidden * 2)?,
sg1: f(sh_inter)?,
sg3: f(sh_inter)?,
shbuf: f(sh_inter)?,
shb16: b(sh_inter * 2)?,
sh_out: f(hidden)?,
cmp_kv_row: f(max_latent)?,
cmp_sc_row: f(max_latent)?,
cmp_emit: f(2 * max_d)?,
cmp_shift: f(max_shift.max(1))?,
sink_scores: f(heads * (win + idx_tail))?,
sink_evals: f(heads * (win + idx_tail))?,
sink_den: s.alloc_zeros::<f64>(heads).map_err(e("ws f64"))?,
head_mixes: f(hc)?,
head_pre: f(hc)?,
collapsed: f(hidden)?,
logits: f(vocab)?,
argmax: i(1)?,
tok: i(1)?,
});
}
Ok(out)
}
#[allow(clippy::too_many_arguments)]
fn cmp_decode(
&self,
st: &Stage,
cmp: &CmpDev,
x: &CudaSlice<f32>, pos: usize,
hidden: usize,
fc_dev: &CudaSlice<f32>,
rd: usize,
eps: f32,
pend_kv: &mut CudaSlice<f32>,
pend_score: &mut CudaSlice<f32>,
store: &mut CudaSlice<f32>,
row0: usize,
blocks: &mut usize,
) -> Res<()> {
let stream = st.gpu.stream();
let (ratio, d, latent) = (cmp.ratio, cmp.d, cmp.latent);
let mut kv_row = stream.alloc_zeros::<f32>(latent).map_err(e("dkv"))?;
let mut sc_row = stream.alloc_zeros::<f32>(latent).map_err(e("dsc"))?;
Self::dots(st, x, &cmp.wkv, 1, hidden, latent, &mut kv_row)?;
Self::dots(st, x, &cmp.wgate, 1, hidden, latent, &mut sc_row)?;
let slot = if cmp.overlap {
ratio + pos % ratio
} else {
pos % ratio
};
{
let src = kv_row.slice(0..latent);
let mut dst = pend_kv.slice_mut(slot * latent..(slot + 1) * latent);
stream.memcpy_dtod(&src, &mut dst).map_err(e("pend kv"))?;
let src = sc_row.slice(0..latent);
let mut dst = pend_score.slice_mut(slot * latent..(slot + 1) * latent);
stream.memcpy_dtod(&src, &mut dst).map_err(e("pend sc"))?;
}
if (pos + 1) % ratio != 0 {
return Ok(());
}
let j = pos / ratio;
let nb_launch = if cmp.overlap { 2usize } else { 1 };
let row_off = if cmp.overlap { d } else { 0 };
let mut out = stream
.alloc_zeros::<f32>(nb_launch * d)
.map_err(e("emit"))?;
unsafe {
ck(
"compressor_pool dec",
k::memra_dsv4_compressor_pool(
dpf!(pend_kv, &stream),
dpf!(pend_score, &stream),
dpf!(cmp.ape, &stream),
dpm!(out, &stream),
nb_launch as i32,
ratio as i32,
d as i32,
latent as i32,
cmp.overlap as i32,
sp(&stream),
),
)?;
let row_c = (out.device_ptr(&stream).0 as usize + row_off * 4) as *const f32;
let row_m = (out.device_ptr_mut(&stream).0 as usize + row_off * 4) as *mut f32;
ck(
"rmsnorm dec cmp",
k::memra_dsv4_rmsnorm(
row_c,
dpf!(cmp.norm, &stream),
row_m,
1,
d as i32,
eps,
sp(&stream),
),
)?;
let pos_dev = upload_i32(&stream, &[(j * ratio) as i32])?;
ck(
"rope dec cmp",
k::memra_dsv4_rope(
row_m,
1,
1,
d as i32,
rd as i32,
dpf!(fc_dev, &stream),
pos_dev.device_ptr(&stream).0 as *const i32,
0,
sp(&stream),
),
)?;
if cmp.rotate {
let scale = (d as f32).powf(-0.5);
ck(
"hadamard dec cmp",
k::memra_dsv4_hadamard(row_m, 1, d as i32, scale, sp(&stream)),
)?;
ck(
"fp4 dec cmp",
k::memra_dsv4_fp4_act_quant(row_m, 1, d as i64, d as i32, sp(&stream)),
)?;
} else {
ck(
"act_quant dec cmp",
k::memra_dsv4_act_quant(
row_m,
1,
d as i64,
(d - rd) as i32,
64,
(self.variant == ActQuantVariant::ClampOnly) as i32,
sp(&stream),
),
)?;
}
}
{
let src = out.slice(row_off..row_off + d);
let mut dst = store.slice_mut((row0 + j) * d..(row0 + j + 1) * d);
stream
.memcpy_dtod(&src, &mut dst)
.map_err(e("emit store"))?;
}
if cmp.overlap {
let mut tmp = stream
.alloc_zeros::<f32>(ratio * latent)
.map_err(e("shift tmp"))?;
{
let src = pend_kv.slice(ratio * latent..2 * ratio * latent);
stream.memcpy_dtod(&src, &mut tmp).map_err(e("shift1"))?;
}
{
let mut dst = pend_kv.slice_mut(0..ratio * latent);
stream
.memcpy_dtod(&tmp.slice(0..ratio * latent), &mut dst)
.map_err(e("shift2"))?;
}
{
let src = pend_score.slice(ratio * latent..2 * ratio * latent);
stream.memcpy_dtod(&src, &mut tmp).map_err(e("shift3"))?;
}
{
let mut dst = pend_score.slice_mut(0..ratio * latent);
stream
.memcpy_dtod(&tmp.slice(0..ratio * latent), &mut dst)
.map_err(e("shift4"))?;
}
}
*blocks = j + 1;
Ok(())
}
#[allow(clippy::too_many_arguments)]
fn block_decode(
&self,
st: &Stage,
layer: &LayerDev,
cache: &mut LayerCache,
h: &CudaSlice<f32>,
pos: usize,
tok: u32,
mut dump: Option<&mut Vec<(String, Vec<f32>)>>,
) -> Res<CudaSlice<f32>> {
let d = self.model.cfg();
let mc = &self.model.mc;
let hc = d.hc_mult as usize;
let hidden = mc.n_embd as usize;
let heads = mc.n_head as usize;
let hd = d.head_dim as usize;
let rd = d.qk_rope_head_dim as usize;
let q_lora = d.q_lora_rank as usize;
let win = d.sliding_window as usize;
let o_groups = d.o_groups as usize;
let o_lora = d.o_lora_rank as usize;
let eps = mc.rms_eps;
let iters = d.hc_sinkhorn_iters;
let hc_eps = d.hc_eps;
st.gpu.ctx.bind_to_thread().map_err(e("bind ctx"))?;
let stream = st.gpu.stream();
let fc_dev = if layer.ratio != 0 {
&st.fc_yarn
} else {
&st.fc_plain
};
let clamp_only = (self.variant == ActQuantVariant::ClampOnly) as i32;
let LayerCache {
kvc,
n_blocks,
pend_kv,
pend_score,
ikvc,
i_blocks,
ipend_kv,
ipend_score,
} = cache;
let (y, post, comb) = Self::hc_pre(
st,
h,
&layer.hc_attn_fn,
&layer.hc_attn_base,
&layer.hc_attn_scale,
1,
hc,
hidden,
iters,
hc_eps,
)?;
let mut x = stream.alloc_zeros::<f32>(hidden).map_err(e("x"))?;
unsafe {
ck(
"rmsnorm attn",
k::memra_dsv4_rmsnorm(
dpf!(y, &stream),
dpf!(layer.attn_norm, &stream),
dpm!(x, &stream),
1,
hidden as i32,
eps,
sp(&stream),
),
)?;
}
if let Some(dm) = dump.as_deref_mut() {
dm.push((format!("layer{}.x", layer.il), dtoh_f32(&stream, &x)?));
}
let mut qr = stream.alloc_zeros::<f32>(q_lora).map_err(e("qr"))?;
Self::gemm(st, &x, layer.wq_a.dev(), 0, 1, q_lora, hidden, &mut qr)?;
unsafe {
ck(
"rmsnorm q",
k::memra_dsv4_rmsnorm(
dpf!(qr, &stream),
dpf!(layer.q_norm, &stream),
dpm!(qr, &stream),
1,
q_lora as i32,
eps,
sp(&stream),
),
)?;
}
let mut qr_b = stream.alloc_zeros::<u8>(q_lora * 2).map_err(e("qr_b"))?;
unsafe {
ck(
"cvt qr",
k::memra_dsv4_cvt_bf16(
dpf!(qr, &stream),
qr_b.device_ptr_mut(&stream).0 as *mut c_void,
q_lora as i64,
sp(&stream),
),
)?;
}
let mut q = stream.alloc_zeros::<f32>(heads * hd).map_err(e("q"))?;
Self::gemm_pre(
st,
&qr_b,
layer.wq_b.dev().device_ptr(&stream).0 as *const c_void,
1,
heads * hd,
q_lora,
&mut q,
)?;
let pos_dev = upload_i32(&stream, &[pos as i32])?;
unsafe {
ck(
"headrms",
k::memra_dsv4_headrms(dpm!(q, &stream), heads as i32, hd as i32, eps, sp(&stream)),
)?;
ck(
"rope q",
k::memra_dsv4_rope(
dpm!(q, &stream),
1,
heads as i32,
hd as i32,
rd as i32,
dpf!(fc_dev, &stream),
pos_dev.device_ptr(&stream).0 as *const i32,
0,
sp(&stream),
),
)?;
}
if let Some(dm) = dump.as_deref_mut() {
dm.push((format!("layer{}.q", layer.il), dtoh_f32(&stream, &q)?));
}
let mut kv = stream.alloc_zeros::<f32>(hd).map_err(e("kv"))?;
Self::gemm(st, &x, layer.wkv.dev(), 0, 1, hd, hidden, &mut kv)?;
unsafe {
ck(
"rmsnorm kv",
k::memra_dsv4_rmsnorm(
dpf!(kv, &stream),
dpf!(layer.kv_norm, &stream),
dpm!(kv, &stream),
1,
hd as i32,
eps,
sp(&stream),
),
)?;
ck(
"rope kv",
k::memra_dsv4_rope(
dpm!(kv, &stream),
1,
1,
hd as i32,
rd as i32,
dpf!(fc_dev, &stream),
pos_dev.device_ptr(&stream).0 as *const i32,
0,
sp(&stream),
),
)?;
ck(
"act_quant kv",
k::memra_dsv4_act_quant(
dpm!(kv, &stream),
1,
hd as i64,
(hd - rd) as i32,
64,
clamp_only,
sp(&stream),
),
)?;
}
{
let slot = pos % win;
let src = kv.slice(0..hd);
let mut dst = kvc.slice_mut(slot * hd..(slot + 1) * hd);
stream
.memcpy_dtod(&src, &mut dst)
.map_err(e("ring write"))?;
}
if let Some(dm) = dump.as_deref_mut() {
dm.push((format!("layer{}.kv", layer.il), dtoh_f32(&stream, &kv)?));
}
let mut idxs: Vec<i64> = vec![-1; win];
if pos >= win - 1 {
let sp_ = pos % win;
let mut k_ = 0usize;
for s_ in (sp_ + 1)..win {
idxs[k_] = s_ as i64;
k_ += 1;
}
for s_ in 0..=sp_ {
idxs[k_] = s_ as i64;
k_ += 1;
}
} else {
for (p, v) in idxs.iter_mut().enumerate().take(pos + 1) {
*v = p as i64;
}
}
if layer.ratio != 0 {
let cidx: Vec<i64> = if let Some(ix) = &layer.idx {
let mut qi = stream
.alloc_zeros::<f32>(ix.heads * ix.hd)
.map_err(e("qi"))?;
Self::gemm_pre(
st,
&qr_b,
ix.wq_b.dev().device_ptr(&stream).0 as *const c_void,
1,
ix.heads * ix.hd,
q_lora,
&mut qi,
)?;
unsafe {
ck(
"rope qi",
k::memra_dsv4_rope(
dpm!(qi, &stream),
1,
ix.heads as i32,
ix.hd as i32,
rd as i32,
dpf!(fc_dev, &stream),
pos_dev.device_ptr(&stream).0 as *const i32,
0,
sp(&stream),
),
)?;
let scale = (ix.hd as f32).powf(-0.5);
ck(
"hadamard qi",
k::memra_dsv4_hadamard(
dpm!(qi, &stream),
ix.heads as i32,
ix.hd as i32,
scale,
sp(&stream),
),
)?;
ck(
"fp4 qi",
k::memra_dsv4_fp4_act_quant(
dpm!(qi, &stream),
ix.heads as i32,
ix.hd as i64,
ix.hd as i32,
sp(&stream),
),
)?;
}
self.cmp_decode(
st,
&ix.cmp,
&x,
pos,
hidden,
fc_dev,
rd,
eps,
ipend_kv.as_mut().expect("ipend"),
ipend_score.as_mut().expect("ipend"),
ikvc.as_mut().expect("ikvc"),
0,
i_blocks,
)?;
let nb = *i_blocks;
debug_assert_eq!(nb, (pos + 1) / layer.ratio, "indexer block count");
if nb > 0 {
let mut wproj = stream.alloc_zeros::<f32>(ix.heads).map_err(e("wp"))?;
Self::gemm(
st,
&x,
ix.weights_proj.dev(),
0,
1,
ix.heads,
hidden,
&mut wproj,
)?;
let wscale = ((ix.hd as f64).powf(-0.5) * (ix.heads as f64).powf(-0.5)) as f32;
let mut score = stream.alloc_zeros::<f32>(nb).map_err(e("iscore"))?;
unsafe {
ck(
"indexer_score dec",
k::memra_dsv4_indexer_score(
dpf!(qi, &stream),
dpf!(ikvc.as_ref().expect("ikvc"), &stream),
dpf!(wproj, &stream),
wscale,
dpm!(score, &stream),
1,
ix.heads as i32,
ix.hd as i32,
nb as i32,
layer.ratio as i32,
nb as i32, sp(&stream),
),
)?;
}
let score_h = dtoh_f32(&stream, &score)?;
let kk = ix.topk.min(nb);
let mut order: Vec<usize> = (0..nb).collect();
order.sort_by(|&a, &b| {
score_h[b]
.partial_cmp(&score_h[a])
.unwrap_or(std::cmp::Ordering::Equal)
.then(a.cmp(&b))
});
order
.into_iter()
.take(kk)
.map(|j| (j + win) as i64)
.collect()
} else {
Vec::new()
}
} else {
let nb = (pos + 1) / layer.ratio;
(0..nb).map(|j| (j + win) as i64).collect()
};
self.cmp_decode(
st,
layer.cmp.as_ref().expect("ratio!=0 has compressor"),
&x,
pos,
hidden,
fc_dev,
rd,
eps,
pend_kv.as_mut().expect("pend"),
pend_score.as_mut().expect("pend"),
kvc,
win,
n_blocks,
)?;
debug_assert_eq!(*n_blocks, (pos + 1) / layer.ratio, "attn block count");
idxs.extend_from_slice(&cidx);
}
let slots = idxs.len();
let idxs_i32: Vec<i32> = idxs.iter().map(|&v| v as i32).collect();
let idx_dev = upload_i32(&stream, &idxs_i32)?;
let mut o = stream.alloc_zeros::<f32>(heads * hd).map_err(e("o"))?;
let scale = (hd as f64).powf(-0.5) as f32;
unsafe {
ck(
"sink_attn dec",
k::memra_dsv4_sink_attn(
dpf!(q, &stream),
dpf!(kvc, &stream),
idx_dev.device_ptr(&stream).0 as *const i32,
dpf!(layer.sink, &stream),
dpm!(o, &stream),
1,
heads as i32,
hd as i32,
slots as i32,
scale,
sp(&stream),
),
)?;
ck(
"rope o inv",
k::memra_dsv4_rope(
dpm!(o, &stream),
1,
heads as i32,
hd as i32,
rd as i32,
dpf!(fc_dev, &stream),
pos_dev.device_ptr(&stream).0 as *const i32,
1,
sp(&stream),
),
)?;
}
if let Some(dm) = dump.as_deref_mut() {
dm.push((format!("layer{}.o", layer.il), dtoh_f32(&stream, &o)?));
}
let gw = heads / o_groups * hd;
let mut og = stream
.alloc_zeros::<f32>(o_groups * o_lora)
.map_err(e("og"))?;
let mut o_grp = stream.alloc_zeros::<f32>(gw).map_err(e("o_grp"))?;
let mut y_grp = stream.alloc_zeros::<f32>(o_lora).map_err(e("y_grp"))?;
for g in 0..o_groups {
unsafe {
ck(
"take_cols",
k::memra_dsv4_take_cols(
dpf!(o, &stream),
dpm!(o_grp, &stream),
1,
gw as i32,
(heads * hd) as i64,
(g * gw) as i64,
sp(&stream),
),
)?;
}
Self::gemm(
st,
&o_grp,
layer.wo_a.dev(),
g * o_lora * gw,
1,
o_lora,
gw,
&mut y_grp,
)?;
unsafe {
ck(
"place_cols",
k::memra_dsv4_place_cols(
dpf!(y_grp, &stream),
dpm!(og, &stream),
1,
o_lora as i32,
(o_groups * o_lora) as i64,
(g * o_lora) as i64,
sp(&stream),
),
)?;
}
}
let mut attn_out = stream.alloc_zeros::<f32>(hidden).map_err(e("ao"))?;
Self::gemm(
st,
&og,
layer.wo_b.dev(),
0,
1,
hidden,
o_groups * o_lora,
&mut attn_out,
)?;
if let Some(dm) = dump.as_deref_mut() {
dm.push((
format!("layer{}.attn_out", layer.il),
dtoh_f32(&stream, &attn_out)?,
));
}
let mut h2 = stream.alloc_zeros::<f32>(hc * hidden).map_err(e("h2"))?;
unsafe {
ck(
"hc_post attn",
k::memra_dsv4_hc_post(
dpf!(attn_out, &stream),
dpf!(h, &stream),
dpf!(post, &stream),
dpf!(comb, &stream),
dpm!(h2, &stream),
1,
hc as i32,
hidden as i32,
sp(&stream),
),
)?;
}
let (y2, post2, comb2) = Self::hc_pre(
st,
&h2,
&layer.hc_ffn_fn,
&layer.hc_ffn_base,
&layer.hc_ffn_scale,
1,
hc,
hidden,
iters,
hc_eps,
)?;
let mut xf = stream.alloc_zeros::<f32>(hidden).map_err(e("xf"))?;
unsafe {
ck(
"rmsnorm ffn",
k::memra_dsv4_rmsnorm(
dpf!(y2, &stream),
dpf!(layer.ffn_norm, &stream),
dpm!(xf, &stream),
1,
hidden as i32,
eps,
sp(&stream),
),
)?;
}
let moe_out = self.moe_forward(st, layer, &xf, 1, &[tok])?;
if let Some(dm) = dump.as_deref_mut() {
dm.push((
format!("layer{}.moe_out", layer.il),
dtoh_f32(&stream, &moe_out)?,
));
}
let mut h3 = stream.alloc_zeros::<f32>(hc * hidden).map_err(e("h3"))?;
unsafe {
ck(
"hc_post ffn",
k::memra_dsv4_hc_post(
dpf!(moe_out, &stream),
dpf!(h2, &stream),
dpf!(post2, &stream),
dpf!(comb2, &stream),
dpm!(h3, &stream),
1,
hc as i32,
hidden as i32,
sp(&stream),
),
)?;
}
if let Some(dm) = dump.as_deref_mut() {
dm.push((format!("layer{}.h3", layer.il), dtoh_f32(&stream, &h3)?));
}
Ok(h3)
}
pub fn decode_step(&self, tok: u32, state: &mut DecodeState) -> Res<Vec<f32>> {
self.decode_step_impl(tok, state, None)
}
pub fn decode_step_probe(
&self,
tok: u32,
state: &mut DecodeState,
) -> Res<(Vec<f32>, Vec<(String, Vec<f32>)>)> {
let mut dump = Vec::new();
let logits = self.decode_step_impl(tok, state, Some(&mut dump))?;
Ok((logits, dump))
}
#[allow(clippy::too_many_arguments)]
fn gemm_dev(
st: &Stage,
x_f32: *const f32,
xb: &mut CudaSlice<u8>,
w: DW,
m: usize,
n: usize,
kdim: usize,
y_ptr: *mut f32,
) -> Res<()> {
assert_eq!(m, 1, "gemm_dev is the m=1 decode path");
let stream = st.gpu.stream();
unsafe {
ck(
"cvt_bf16 dev",
k::memra_dsv4_cvt_bf16(
x_f32,
xb.device_ptr_mut(&stream).0 as *mut c_void,
kdim as i64,
sp(&stream),
),
)?;
}
let xb_ptr = xb.device_ptr(&stream).0 as *const c_void;
Self::gemv_pre_dev(st, xb_ptr, w, n, kdim, y_ptr)
}
fn gemv_pre_dev(
st: &Stage,
xb_ptr: *const c_void,
w: DW,
n: usize,
kdim: usize,
y_ptr: *mut f32,
) -> Res<()> {
let stream = st.gpu.stream();
unsafe {
match w {
DW::Bf16(w_ptr) => ck(
"gemv_bf16 pre dev",
k::memra_dsv4_gemv_bf16(
w_ptr,
xb_ptr,
y_ptr,
n as i32,
kdim as i32,
sp(&stream),
),
)?,
DW::Fp8 {
codes,
scales,
sc_cols,
} => ck(
"gemv_fp8 pre dev",
k::memra_dsv4_gemv_fp8(
codes,
scales,
sc_cols,
xb_ptr,
y_ptr,
n as i32,
kdim as i32,
sp(&stream),
),
)?,
}
}
Ok(())
}
#[allow(clippy::too_many_arguments)]
#[allow(clippy::too_many_arguments)]
unsafe fn rmsnorm_arm(
&self,
x: *const f32,
w: *const f32,
dst: *mut f32,
rows: i32,
ncols: i32,
eps: f32,
sv: *mut c_void,
) -> i32 {
unsafe {
if self.chains_f32 {
k::memra_dsv4_rmsnorm_f32acc(x, w, dst, rows, ncols, eps, sv)
} else {
k::memra_dsv4_rmsnorm(x, w, dst, rows, ncols, eps, sv)
}
}
}
unsafe fn headrms_arm(&self, x: *mut f32, rows: i32, d: i32, eps: f32, sv: *mut c_void) -> i32 {
unsafe {
if self.chains_f32 {
k::memra_dsv4_headrms_f32acc(x, rows, d, eps, sv)
} else {
k::memra_dsv4_headrms(x, rows, d, eps, sv)
}
}
}
#[allow(clippy::too_many_arguments)]
unsafe fn rowsq_scale_arm(
&self,
x: *const f32,
mixes: *mut f32,
s: i32,
w: i32,
rows: i32,
eps: f32,
sv: *mut c_void,
) -> i32 {
unsafe {
if self.chains_f32 {
k::memra_dsv4_rowsq_scale_f32acc(x, mixes, s, w, rows, eps, sv)
} else {
k::memra_dsv4_rowsq_scale(x, mixes, s, w, rows, eps, sv)
}
}
}
#[allow(clippy::too_many_arguments)]
unsafe fn indexer_score_arm(
&self,
q: *const f32,
ckv: *const f32,
w: *const f32,
wscale: f32,
score: *mut f32,
s: i32,
heads: i32,
hd: i32,
nb: i32,
ratio: i32,
lim0: i32,
sv: *mut c_void,
) -> i32 {
unsafe {
if self.chains_f32 {
k::memra_dsv4_indexer_score_f32acc(
q, ckv, w, wscale, score, s, heads, hd, nb, ratio, lim0, sv,
)
} else {
k::memra_dsv4_indexer_score(
q, ckv, w, wscale, score, s, heads, hd, nb, ratio, lim0, sv,
)
}
}
}
#[allow(clippy::too_many_arguments)]
unsafe fn sink_attn_dec_arm(
&self,
q: *const f32,
kv: *const f32,
idxs: *const i32,
sink: *const f32,
scores: *mut f32,
evals: *mut f32,
den: *mut f64,
o: *mut f32,
heads: i32,
hd: i32,
slots: i32,
scale: f32,
sv: *mut c_void,
) -> i32 {
unsafe {
if self.chains_f32 {
k::memra_dsv4_sink_attn_dec_f32acc(
q,
kv,
idxs,
sink,
scores,
evals,
den as *mut f32,
o,
heads,
hd,
slots,
scale,
sv,
)
} else {
k::memra_dsv4_sink_attn_dec(
q, kv, idxs, sink, scores, evals, den, o, heads, hd, slots, scale, sv,
)
}
}
}
fn hc_pre_dev(
&self,
st: &Stage,
h: &CudaSlice<f32>,
fn_w: &CudaSlice<f32>,
base_host: &[f32],
scale_host: &[f32],
base_dev: &CudaSlice<f32>,
scale_dev: &CudaSlice<f32>,
mixes: &mut CudaSlice<f32>,
pre: &mut CudaSlice<f32>,
post: &mut CudaSlice<f32>,
comb: &mut CudaSlice<f32>,
y_hc: &mut CudaSlice<f32>,
hc: usize,
hidden: usize,
iters: u32,
hc_eps: f32,
host_math: bool,
) -> Res<()> {
let stream = st.gpu.stream();
let w = hc * hidden;
let rows = (2 + hc) * hc;
self.dots_dev(st, h, fn_w, 1, w, rows, mixes)?;
unsafe {
ck(
"rowsq_scale dev",
self.rowsq_scale_arm(
dpf!(h, &stream),
dpm!(*mixes, &stream),
1,
w as i32,
rows as i32,
hc_eps,
sp(&stream),
),
)?;
}
if host_math {
let mixes_h = dtoh_f32(&stream, mixes)?;
let (pre_h, post_h, comb_h) =
hc_split_sinkhorn(&mixes_h, 1, hc, scale_host, base_host, iters, hc_eps);
stream.memcpy_htod(&pre_h, pre).map_err(e("htod pre"))?;
stream.memcpy_htod(&post_h, post).map_err(e("htod post"))?;
stream.memcpy_htod(&comb_h, comb).map_err(e("htod comb"))?;
} else {
unsafe {
ck(
"hc_sinkhorn",
k::memra_dsv4_hc_sinkhorn(
dpf!(*mixes, &stream),
dpf!(scale_dev, &stream),
dpf!(base_dev, &stream),
dpm!(*pre, &stream),
dpm!(*post, &stream),
dpm!(*comb, &stream),
hc as i32,
iters as i32,
hc_eps,
sp(&stream),
),
)?;
}
}
unsafe {
ck(
"hc_collapse dev",
k::memra_dsv4_hc_collapse(
dpf!(h, &stream),
dpf!(*pre, &stream),
dpm!(*y_hc, &stream),
1,
hc as i32,
hidden as i32,
sp(&stream),
),
)?;
}
Ok(())
}
#[allow(clippy::too_many_arguments)]
fn cmp_decode_dev(
&self,
st: &Stage,
cmp: &CmpDev,
x: &CudaSlice<f32>,
pos: usize,
hidden: usize,
fc_dev: &CudaSlice<f32>,
rd: usize,
eps: f32,
kv_row: &mut CudaSlice<f32>,
sc_row: &mut CudaSlice<f32>,
emit: &mut CudaSlice<f32>,
shift: &mut CudaSlice<f32>,
pend_kv: &mut CudaSlice<f32>,
pend_score: &mut CudaSlice<f32>,
store: &mut CudaSlice<f32>,
row0: usize,
blocks: &mut usize,
) -> Res<()> {
let stream = st.gpu.stream();
let (ratio, d, latent) = (cmp.ratio, cmp.d, cmp.latent);
self.dots_dev(st, x, &cmp.wkv, 1, hidden, latent, kv_row)?;
self.dots_dev(st, x, &cmp.wgate, 1, hidden, latent, sc_row)?;
let slot = if cmp.overlap {
ratio + pos % ratio
} else {
pos % ratio
};
{
let src = kv_row.slice(0..latent);
let mut dst = pend_kv.slice_mut(slot * latent..(slot + 1) * latent);
stream.memcpy_dtod(&src, &mut dst).map_err(e("pend kv"))?;
let src = sc_row.slice(0..latent);
let mut dst = pend_score.slice_mut(slot * latent..(slot + 1) * latent);
stream.memcpy_dtod(&src, &mut dst).map_err(e("pend sc"))?;
}
if (pos + 1) % ratio != 0 {
return Ok(());
}
let j = pos / ratio;
let nb_launch = if cmp.overlap { 2usize } else { 1 };
let row_off = if cmp.overlap { d } else { 0 };
unsafe {
ck(
"compressor_pool dec",
k::memra_dsv4_compressor_pool(
dpf!(*pend_kv, &stream),
dpf!(*pend_score, &stream),
dpf!(cmp.ape, &stream),
dpm!(*emit, &stream),
nb_launch as i32,
ratio as i32,
d as i32,
latent as i32,
cmp.overlap as i32,
sp(&stream),
),
)?;
let row_c = (emit.device_ptr(&stream).0 as usize + row_off * 4) as *const f32;
let row_m = (emit.device_ptr_mut(&stream).0 as usize + row_off * 4) as *mut f32;
ck(
"rmsnorm dec cmp",
self.rmsnorm_arm(
row_c,
dpf!(cmp.norm, &stream),
row_m,
1,
d as i32,
eps,
sp(&stream),
),
)?;
ck(
"rope_at dec cmp",
k::memra_dsv4_rope_at(
row_m,
1,
d as i32,
rd as i32,
dpf!(fc_dev, &stream),
(j * ratio) as i32,
0,
sp(&stream),
),
)?;
if cmp.rotate {
let scale = (d as f32).powf(-0.5);
ck(
"hadamard dec cmp",
k::memra_dsv4_hadamard(row_m, 1, d as i32, scale, sp(&stream)),
)?;
ck(
"fp4 dec cmp",
k::memra_dsv4_fp4_act_quant(row_m, 1, d as i64, d as i32, sp(&stream)),
)?;
} else {
ck(
"act_quant dec cmp",
k::memra_dsv4_act_quant(
row_m,
1,
d as i64,
(d - rd) as i32,
64,
(self.variant == ActQuantVariant::ClampOnly) as i32,
sp(&stream),
),
)?;
}
}
{
let src = emit.slice(row_off..row_off + d);
let mut dst = store.slice_mut((row0 + j) * d..(row0 + j + 1) * d);
stream
.memcpy_dtod(&src, &mut dst)
.map_err(e("emit store"))?;
}
if cmp.overlap {
{
let src = pend_kv.slice(ratio * latent..2 * ratio * latent);
let mut dst = shift.slice_mut(0..ratio * latent);
stream.memcpy_dtod(&src, &mut dst).map_err(e("shift1"))?;
}
{
let src = shift.slice(0..ratio * latent);
let mut dst = pend_kv.slice_mut(0..ratio * latent);
stream.memcpy_dtod(&src, &mut dst).map_err(e("shift2"))?;
}
{
let src = pend_score.slice(ratio * latent..2 * ratio * latent);
let mut dst = shift.slice_mut(0..ratio * latent);
stream.memcpy_dtod(&src, &mut dst).map_err(e("shift3"))?;
}
{
let src = shift.slice(0..ratio * latent);
let mut dst = pend_score.slice_mut(0..ratio * latent);
stream.memcpy_dtod(&src, &mut dst).map_err(e("shift4"))?;
}
}
*blocks = j + 1;
Ok(())
}
#[allow(clippy::too_many_arguments)]
fn block_decode_dev(
&self,
st: &Stage,
layer: &LayerDev,
cache: &mut LayerCache,
ws: &mut StepWs,
input_rx: bool,
pos: usize,
tok: u32,
host_math: bool,
) -> Res<()> {
let d = self.model.cfg();
let mc = &self.model.mc;
let hc = d.hc_mult as usize;
let hidden = mc.n_embd as usize;
let heads = mc.n_head as usize;
let hd = d.head_dim as usize;
let rd = d.qk_rope_head_dim as usize;
let q_lora = d.q_lora_rank as usize;
let win = d.sliding_window as usize;
let o_groups = d.o_groups as usize;
let o_lora = d.o_lora_rank as usize;
let eps = mc.rms_eps;
let iters = d.hc_sinkhorn_iters;
let hc_eps = d.hc_eps;
let stream = st.gpu.stream();
let fc_dev: *const f32 = if layer.ratio != 0 {
st.fc_yarn.device_ptr(&stream).0 as *const f32
} else {
st.fc_plain.device_ptr(&stream).0 as *const f32
};
let clamp_only = (self.variant == ActQuantVariant::ClampOnly) as i32;
let LayerCache {
kvc,
n_blocks,
pend_kv,
pend_score,
ikvc,
i_blocks,
ipend_kv,
ipend_score,
} = cache;
{
let StepWs {
h_a,
h_rx,
mixes,
pre,
post,
comb,
y_hc,
..
} = ws;
let h_in: &CudaSlice<f32> = if input_rx { h_rx } else { h_a };
self.hc_pre_dev(
st,
h_in,
&layer.hc_attn_fn,
&layer.hc_attn_base,
&layer.hc_attn_scale,
&layer.hc_attn_base_dev,
&layer.hc_attn_scale_dev,
mixes,
pre,
post,
comb,
y_hc,
hc,
hidden,
iters,
hc_eps,
host_math,
)?;
}
unsafe {
ck(
"rmsnorm attn dev",
self.rmsnorm_arm(
dpf!(ws.y_hc, &stream),
dpf!(layer.attn_norm, &stream),
dpm!(ws.x, &stream),
1,
hidden as i32,
eps,
sp(&stream),
),
)?;
}
Self::gemm_dev(
st,
ws.x.device_ptr(&stream).0 as *const f32,
&mut ws.gemm_xb,
dwsel(self.dense_fp8, &stream, &layer.wq_a, &layer.wq_a_fp8),
1,
q_lora,
hidden,
ws.qr.device_ptr_mut(&stream).0 as *mut f32,
)?;
unsafe {
ck(
"rmsnorm q dev",
self.rmsnorm_arm(
dpf!(ws.qr, &stream),
dpf!(layer.q_norm, &stream),
dpm!(ws.qr, &stream),
1,
q_lora as i32,
eps,
sp(&stream),
),
)?;
ck(
"cvt qr dev",
k::memra_dsv4_cvt_bf16(
dpf!(ws.qr, &stream),
ws.qr_b.device_ptr_mut(&stream).0 as *mut c_void,
q_lora as i64,
sp(&stream),
),
)?;
}
Self::gemv_pre_dev(
st,
ws.qr_b.device_ptr(&stream).0 as *const c_void,
dwsel(self.dense_fp8, &stream, &layer.wq_b, &layer.wq_b_fp8),
heads * hd,
q_lora,
ws.q.device_ptr_mut(&stream).0 as *mut f32,
)?;
unsafe {
ck(
"headrms dev",
self.headrms_arm(
dpm!(ws.q, &stream),
heads as i32,
hd as i32,
eps,
sp(&stream),
),
)?;
ck(
"rope_at q dev",
k::memra_dsv4_rope_at(
dpm!(ws.q, &stream),
heads as i32,
hd as i32,
rd as i32,
fc_dev,
pos as i32,
0,
sp(&stream),
),
)?;
}
Self::gemm_dev(
st,
ws.x.device_ptr(&stream).0 as *const f32,
&mut ws.gemm_xb,
dwsel(self.dense_fp8, &stream, &layer.wkv, &layer.wkv_fp8),
1,
hd,
hidden,
ws.kv.device_ptr_mut(&stream).0 as *mut f32,
)?;
unsafe {
ck(
"rmsnorm kv dev",
self.rmsnorm_arm(
dpf!(ws.kv, &stream),
dpf!(layer.kv_norm, &stream),
dpm!(ws.kv, &stream),
1,
hd as i32,
eps,
sp(&stream),
),
)?;
ck(
"rope_at kv dev",
k::memra_dsv4_rope_at(
dpm!(ws.kv, &stream),
1,
hd as i32,
rd as i32,
fc_dev,
pos as i32,
0,
sp(&stream),
),
)?;
ck(
"act_quant kv dev",
k::memra_dsv4_act_quant(
dpm!(ws.kv, &stream),
1,
hd as i64,
(hd - rd) as i32,
64,
clamp_only,
sp(&stream),
),
)?;
}
{
let slot = pos % win;
let src = ws.kv.slice(0..hd);
let mut dst = kvc.slice_mut(slot * hd..(slot + 1) * hd);
stream
.memcpy_dtod(&src, &mut dst)
.map_err(e("ring write"))?;
}
let mut slots = win;
if layer.ratio != 0 {
if let Some(ix) = &layer.idx {
Self::gemv_pre_dev(
st,
ws.qr_b.device_ptr(&stream).0 as *const c_void,
dwsel(self.dense_fp8, &stream, &ix.wq_b, &ix.wq_b_fp8),
ix.heads * ix.hd,
q_lora,
ws.qi.device_ptr_mut(&stream).0 as *mut f32,
)?;
unsafe {
ck(
"rope_at qi dev",
k::memra_dsv4_rope_at(
dpm!(ws.qi, &stream),
ix.heads as i32,
ix.hd as i32,
rd as i32,
fc_dev,
pos as i32,
0,
sp(&stream),
),
)?;
let scale = (ix.hd as f32).powf(-0.5);
ck(
"hadamard qi dev",
k::memra_dsv4_hadamard(
dpm!(ws.qi, &stream),
ix.heads as i32,
ix.hd as i32,
scale,
sp(&stream),
),
)?;
ck(
"fp4 qi dev",
k::memra_dsv4_fp4_act_quant(
dpm!(ws.qi, &stream),
ix.heads as i32,
ix.hd as i64,
ix.hd as i32,
sp(&stream),
),
)?;
}
{
let StepWs {
x,
cmp_kv_row,
cmp_sc_row,
cmp_emit,
cmp_shift,
..
} = ws;
self.cmp_decode_dev(
st,
&ix.cmp,
x,
pos,
hidden,
if layer.ratio != 0 {
&st.fc_yarn
} else {
&st.fc_plain
},
rd,
eps,
cmp_kv_row,
cmp_sc_row,
cmp_emit,
cmp_shift,
ipend_kv.as_mut().expect("ipend"),
ipend_score.as_mut().expect("ipend"),
ikvc.as_mut().expect("ikvc"),
0,
i_blocks,
)?;
}
let nb = *i_blocks;
debug_assert_eq!(nb, (pos + 1) / layer.ratio, "indexer block count");
unsafe {
ck(
"build_idx win",
k::memra_dsv4_build_idx(
ws.idx.device_ptr_mut(&stream).0 as *mut i32,
pos as i32,
win as i32,
-1,
win as i32,
sp(&stream),
),
)?;
}
if nb > 0 {
Self::gemm_dev(
st,
ws.x.device_ptr(&stream).0 as *const f32,
&mut ws.gemm_xb,
dwsel(
self.dense_fp8,
&stream,
&ix.weights_proj,
&ix.weights_proj_fp8,
),
1,
ix.heads,
hidden,
ws.wproj.device_ptr_mut(&stream).0 as *mut f32,
)?;
let wscale = ((ix.hd as f64).powf(-0.5) * (ix.heads as f64).powf(-0.5)) as f32;
unsafe {
ck(
"indexer_score dev",
self.indexer_score_arm(
dpf!(ws.qi, &stream),
dpf!(ikvc.as_ref().expect("ikvc"), &stream),
dpf!(ws.wproj, &stream),
wscale,
dpm!(ws.score, &stream),
1,
ix.heads as i32,
ix.hd as i32,
nb as i32,
layer.ratio as i32,
nb as i32,
sp(&stream),
),
)?;
}
let kk = ix.topk.min(nb);
if host_math {
let score_h = {
let view = ws.score.slice(0..nb);
let mut v = vec![0f32; nb];
stream
.memcpy_dtoh(&view, &mut v[..])
.map_err(e("dtoh sc"))?;
stream.synchronize().map_err(e("sync sc"))?;
v
};
let mut order: Vec<usize> = (0..nb).collect();
order.sort_by(|&a, &b| {
score_h[b]
.partial_cmp(&score_h[a])
.unwrap_or(std::cmp::Ordering::Equal)
.then(a.cmp(&b))
});
let cidx: Vec<i32> = order
.into_iter()
.take(kk)
.map(|j| (j + win) as i32)
.collect();
let mut dst = ws.idx.slice_mut(win..win + kk);
stream.memcpy_htod(&cidx, &mut dst).map_err(e("htod idx"))?;
} else {
unsafe {
let idx_tail =
(ws.idx.device_ptr_mut(&stream).0 as usize + win * 4) as *mut i32;
ck(
"topk_idx dev",
k::memra_dsv4_topk_idx(
dpf!(ws.score, &stream),
nb as i32,
kk as i32,
win as i32,
idx_tail,
sp(&stream),
),
)?;
}
}
slots = win + kk;
}
} else {
let nb = (pos + 1) / layer.ratio;
unsafe {
ck(
"build_idx coarse",
k::memra_dsv4_build_idx(
ws.idx.device_ptr_mut(&stream).0 as *mut i32,
pos as i32,
win as i32,
nb as i32,
(win + nb) as i32,
sp(&stream),
),
)?;
}
slots = win + nb;
}
{
let StepWs {
x,
cmp_kv_row,
cmp_sc_row,
cmp_emit,
cmp_shift,
..
} = ws;
self.cmp_decode_dev(
st,
layer.cmp.as_ref().expect("ratio!=0 has compressor"),
x,
pos,
hidden,
&st.fc_yarn,
rd,
eps,
cmp_kv_row,
cmp_sc_row,
cmp_emit,
cmp_shift,
pend_kv.as_mut().expect("pend"),
pend_score.as_mut().expect("pend"),
kvc,
win,
n_blocks,
)?;
}
debug_assert_eq!(*n_blocks, (pos + 1) / layer.ratio, "attn block count");
} else {
unsafe {
ck(
"build_idx window-only",
k::memra_dsv4_build_idx(
ws.idx.device_ptr_mut(&stream).0 as *mut i32,
pos as i32,
win as i32,
-1,
win as i32,
sp(&stream),
),
)?;
}
}
let scale = (hd as f64).powf(-0.5) as f32;
unsafe {
ck(
"sink_attn_dec dev",
self.sink_attn_dec_arm(
dpf!(ws.q, &stream),
dpf!(kvc, &stream),
ws.idx.device_ptr(&stream).0 as *const i32,
dpf!(layer.sink, &stream),
dpm!(ws.sink_scores, &stream),
dpm!(ws.sink_evals, &stream),
ws.sink_den.device_ptr_mut(&stream).0 as *mut f64,
dpm!(ws.o, &stream),
heads as i32,
hd as i32,
slots as i32,
scale,
sp(&stream),
),
)?;
ck(
"rope_at o inv dev",
k::memra_dsv4_rope_at(
dpm!(ws.o, &stream),
heads as i32,
hd as i32,
rd as i32,
fc_dev,
pos as i32,
1,
sp(&stream),
),
)?;
}
let gw = heads / o_groups * hd;
unsafe {
ck(
"cvt o dev",
k::memra_dsv4_cvt_bf16(
dpf!(ws.o, &stream),
ws.o_b.device_ptr_mut(&stream).0 as *mut c_void,
(heads * hd) as i64,
sp(&stream),
),
)?;
}
let wo_a_dw = dwsel(self.dense_fp8, &stream, &layer.wo_a, &layer.wo_a_fp8);
for g in 0..o_groups {
Self::gemv_pre_dev(
st,
(ws.o_b.device_ptr(&stream).0 as usize + g * gw * 2) as *const c_void,
wo_a_dw.offset_rows(g * o_lora, gw),
o_lora,
gw,
(ws.og.device_ptr_mut(&stream).0 as usize + g * o_lora * 4) as *mut f32,
)?;
}
Self::gemm_dev(
st,
ws.og.device_ptr(&stream).0 as *const f32,
&mut ws.gemm_xb,
dwsel(self.dense_fp8, &stream, &layer.wo_b, &layer.wo_b_fp8),
1,
hidden,
o_groups * o_lora,
ws.attn_out.device_ptr_mut(&stream).0 as *mut f32,
)?;
{
let StepWs {
h_a,
h_b,
h_rx,
attn_out,
post,
comb,
..
} = ws;
let h_in: &CudaSlice<f32> = if input_rx { h_rx } else { h_a };
unsafe {
ck(
"hc_post attn dev",
k::memra_dsv4_hc_post(
dpf!(attn_out, &stream),
dpf!(h_in, &stream),
dpf!(post, &stream),
dpf!(comb, &stream),
dpm!(*h_b, &stream),
1,
hc as i32,
hidden as i32,
sp(&stream),
),
)?;
}
}
{
let StepWs {
h_b,
mixes,
pre,
post,
comb,
y_hc,
..
} = ws;
self.hc_pre_dev(
st,
h_b,
&layer.hc_ffn_fn,
&layer.hc_ffn_base,
&layer.hc_ffn_scale,
&layer.hc_ffn_base_dev,
&layer.hc_ffn_scale_dev,
mixes,
pre,
post,
comb,
y_hc,
hc,
hidden,
iters,
hc_eps,
host_math,
)?;
}
unsafe {
ck(
"rmsnorm ffn dev",
self.rmsnorm_arm(
dpf!(ws.y_hc, &stream),
dpf!(layer.ffn_norm, &stream),
dpm!(ws.xf, &stream),
1,
hidden as i32,
eps,
sp(&stream),
),
)?;
}
self.moe_forward_dev(st, layer, ws, tok, host_math)?;
{
let StepWs {
h_a,
h_b,
y,
post,
comb,
..
} = ws;
unsafe {
ck(
"hc_post ffn dev",
k::memra_dsv4_hc_post(
dpf!(y, &stream),
dpf!(h_b, &stream),
dpf!(post, &stream),
dpf!(comb, &stream),
dpm!(*h_a, &stream),
1,
hc as i32,
hidden as i32,
sp(&stream),
),
)?;
}
}
Ok(())
}
fn moe_forward_dev(
&self,
st: &Stage,
layer: &LayerDev,
ws: &mut StepWs,
tok: u32,
host_math: bool,
) -> Res<()> {
let mc = &self.model.mc;
let d = self.model.cfg();
let moe = mc.moe.as_ref().expect("moe");
let hidden = mc.n_embd as usize;
let ne = moe.expert_count as usize;
let topk = moe.expert_used_count as usize;
let inter = moe.expert_ff_length as usize;
let limit = d.swiglu_limit;
let stream = st.gpu.stream();
let kind = match layer.expert_kind {
ExpertKind::Nvfp4 => 0i32,
ExpertKind::Mxfp4 => 1i32,
};
let wstride = (inter * hidden / 2) as i64;
let sstride = match layer.expert_kind {
ExpertKind::Nvfp4 => (inter * hidden / 16) as i64,
ExpertKind::Mxfp4 => (inter * hidden / 32) as i64,
};
self.dots_dev(st, &ws.xf, &layer.gate_w, 1, hidden, ne, &mut ws.raw)?;
if host_math {
let raw_h = dtoh_f32(&stream, &ws.raw)?;
let (indices, weights) =
Self::route_host(layer, &raw_h, &[tok], 1, ne, topk, d.routed_scaling_factor);
let sel: Vec<i32> = indices.iter().map(|&x| x as i32).collect();
let mut order: Vec<i32> = (0..topk as i32).collect();
order.sort_by_key(|&s| indices[s as usize]);
stream
.memcpy_htod(&sel, &mut ws.sel)
.map_err(e("htod sel"))?;
stream
.memcpy_htod(&weights, &mut ws.selw)
.map_err(e("htod selw"))?;
stream
.memcpy_htod(&order, &mut ws.order)
.map_err(e("htod order"))?;
} else {
unsafe {
ck(
"route dev",
k::memra_dsv4_route(
dpf!(ws.raw, &stream),
layer
.gate_bias_dev
.as_ref()
.map(|b| b.device_ptr(&stream).0 as *const f32)
.unwrap_or(std::ptr::null()),
layer
.tid2eid_dev
.as_ref()
.map(|t| t.device_ptr(&stream).0 as *const i32)
.unwrap_or(std::ptr::null()),
ws.tok.device_ptr(&stream).0 as *const i32,
ne as i32,
topk as i32,
d.routed_scaling_factor,
ws.sel.device_ptr_mut(&stream).0 as *mut i32,
ws.selw.device_ptr_mut(&stream).0 as *mut f32,
ws.order.device_ptr_mut(&stream).0 as *mut i32,
sp(&stream),
),
)?;
}
}
unsafe {
ck(
"act_quant_fp8 x dev",
k::memra_dsv4_act_quant_fp8(
dpf!(ws.xf, &stream),
ws.xq.device_ptr_mut(&stream).0 as *mut c_void,
dpm!(ws.xs, &stream),
1,
hidden as i32,
sp(&stream),
),
)?;
for (proj, dst) in [(0i32, &mut ws.g1), (2i32, &mut ws.g3)] {
ck(
"fp4_gemm_sel w1/w3",
k::memra_dsv4_fp4_gemm_sel(
dp!(ws.xq, &stream),
dpf!(ws.xs, &stream),
dp!(layer.experts_w, &stream),
dp!(layer.experts_sc, &stream),
dpf!(layer.experts_s2_dev, &stream),
ws.sel.device_ptr(&stream).0 as *const i32,
proj,
0,
kind,
dpm!(*dst, &stream),
topk as i32,
inter as i32,
hidden as i32,
wstride,
sstride,
sp(&stream),
),
)?;
}
ck(
"swiglu dev",
k::memra_dsv4_swiglu(
dpf!(ws.g1, &stream),
dpf!(ws.g3, &stream),
dpm!(ws.hbuf, &stream),
topk as i32,
inter as i32,
limit,
ws.selw.device_ptr(&stream).0 as *const f32,
sp(&stream),
),
)?;
ck(
"act_quant_fp8 h dev",
k::memra_dsv4_act_quant_fp8(
dpf!(ws.hbuf, &stream),
ws.hq.device_ptr_mut(&stream).0 as *mut c_void,
dpm!(ws.hs, &stream),
topk as i32,
inter as i32,
sp(&stream),
),
)?;
ck(
"fp4_gemm_sel w2",
k::memra_dsv4_fp4_gemm_sel(
dp!(ws.hq, &stream),
dpf!(ws.hs, &stream),
dp!(layer.experts_w, &stream),
dp!(layer.experts_sc, &stream),
dpf!(layer.experts_s2_dev, &stream),
ws.sel.device_ptr(&stream).0 as *const i32,
1,
1,
kind,
dpm!(ws.contrib, &stream),
topk as i32,
hidden as i32,
inter as i32,
wstride,
sstride,
sp(&stream),
),
)?;
ck(
"combine dev",
k::memra_dsv4_combine_rows(
dpf!(ws.contrib, &stream),
ws.order.device_ptr(&stream).0 as *const i32,
topk as i32,
dpm!(ws.y, &stream),
hidden as i64,
sp(&stream),
),
)?;
ck(
"cvt xb dev",
k::memra_dsv4_cvt_bf16(
dpf!(ws.xf, &stream),
ws.xb.device_ptr_mut(&stream).0 as *mut c_void,
hidden as i64,
sp(&stream),
),
)?;
}
let sh_inter = ws.sg1.len();
Self::gemv_pre_dev(
st,
ws.xb.device_ptr(&stream).0 as *const c_void,
dwsel(
self.dense_fp8,
&stream,
&layer.shared_w[0],
&layer.shared_fp8[0],
),
sh_inter,
hidden,
ws.sg1.device_ptr_mut(&stream).0 as *mut f32,
)?;
Self::gemv_pre_dev(
st,
ws.xb.device_ptr(&stream).0 as *const c_void,
dwsel(
self.dense_fp8,
&stream,
&layer.shared_w[2],
&layer.shared_fp8[2],
),
sh_inter,
hidden,
ws.sg3.device_ptr_mut(&stream).0 as *mut f32,
)?;
unsafe {
ck(
"swiglu sh dev",
k::memra_dsv4_swiglu(
dpf!(ws.sg1, &stream),
dpf!(ws.sg3, &stream),
dpm!(ws.shbuf, &stream),
1,
sh_inter as i32,
limit,
std::ptr::null(),
sp(&stream),
),
)?;
ck(
"cvt sh dev",
k::memra_dsv4_cvt_bf16(
dpf!(ws.shbuf, &stream),
ws.shb16.device_ptr_mut(&stream).0 as *mut c_void,
sh_inter as i64,
sp(&stream),
),
)?;
}
Self::gemv_pre_dev(
st,
ws.shb16.device_ptr(&stream).0 as *const c_void,
dwsel(
self.dense_fp8,
&stream,
&layer.shared_w[1],
&layer.shared_fp8[1],
),
hidden,
sh_inter,
ws.sh_out.device_ptr_mut(&stream).0 as *mut f32,
)?;
unsafe {
ck(
"add shared dev",
k::memra_dsv4_add_inplace(
dpm!(ws.y, &stream),
dpf!(ws.sh_out, &stream),
hidden as i64,
sp(&stream),
),
)?;
}
Ok(())
}
fn head_logits_dev(&self, ws: &mut StepWs, host_math: bool) -> Res<()> {
let d = self.model.cfg();
let mc = &self.model.mc;
let hc = d.hc_mult as usize;
let hidden = mc.n_embd as usize;
let eps = mc.rms_eps;
let last = self.stages.len() - 1;
let st = &self.stages[last];
let stream = st.gpu.stream();
let w = hc * hidden;
let fn_w = st.hc_head_fn.as_ref().expect("hc_head_fn");
let norm = st.trunk_norm.as_ref().expect("trunk norm");
self.dots_dev(st, &ws.h_a, fn_w, 1, w, hc, &mut ws.head_mixes)?;
unsafe {
ck(
"rowsq head dev",
self.rowsq_scale_arm(
dpf!(ws.h_a, &stream),
dpm!(ws.head_mixes, &stream),
1,
w as i32,
hc as i32,
eps,
sp(&stream),
),
)?;
}
if host_math {
let mut mixes_h = dtoh_f32(&stream, &ws.head_mixes)?;
for c in 0..hc {
let m = mixes_h[c];
mixes_h[c] =
sigmoid_f32(m * self.hc_head_scale[0] + self.hc_head_base[c]) + d.hc_eps;
}
stream
.memcpy_htod(&mixes_h, &mut ws.head_pre)
.map_err(e("htod head pre"))?;
} else {
unsafe {
ck(
"hc_head_pre dev",
k::memra_dsv4_hc_head_pre(
dpf!(ws.head_mixes, &stream),
st.hc_head_scale_dev
.as_ref()
.expect("head scale dev")
.device_ptr(&stream)
.0 as *const f32,
st.hc_head_base_dev
.as_ref()
.expect("head base dev")
.device_ptr(&stream)
.0 as *const f32,
dpm!(ws.head_pre, &stream),
hc as i32,
d.hc_eps,
sp(&stream),
),
)?;
}
}
unsafe {
ck(
"hc_collapse head dev",
k::memra_dsv4_hc_collapse(
dpf!(ws.h_a, &stream),
dpf!(ws.head_pre, &stream),
dpm!(ws.collapsed, &stream),
1,
hc as i32,
hidden as i32,
sp(&stream),
),
)?;
ck(
"rmsnorm head dev",
self.rmsnorm_arm(
dpf!(ws.collapsed, &stream),
dpf!(norm, &stream),
dpm!(ws.collapsed, &stream),
1,
hidden as i32,
eps,
sp(&stream),
),
)?;
let head_ptr = st.head.as_ref().expect("head").device_ptr(&stream).0 as *const c_void;
if self.dots_f32 {
ck(
"head dots f32acc dev",
k::memra_dsv4_dots_f32acc(
dpf!(ws.collapsed, &stream),
head_ptr,
1,
dpm!(ws.logits, &stream),
1,
hidden as i32,
ws.logits.len() as i32,
sp(&stream),
),
)?;
} else {
ck(
"head dots dev",
k::memra_dsv4_dots_f32(
dpf!(ws.collapsed, &stream),
head_ptr,
1,
dpm!(ws.logits, &stream),
1,
hidden as i32,
ws.logits.len() as i32,
sp(&stream),
),
)?;
}
}
Ok(())
}
fn decode_step_fast(
&self,
tok: u32,
state: &mut DecodeState,
want_logits: bool,
host_math: bool,
) -> Res<(Option<Vec<f32>>, u32)> {
self.decode_step_fast_tap(tok, state, want_logits, host_math, None)
}
fn decode_step_fast_tap(
&self,
tok: u32,
state: &mut DecodeState,
want_logits: bool,
host_math: bool,
mut taps: Option<(&mut CudaSlice<f32>, usize)>,
) -> Res<(Option<Vec<f32>>, u32)> {
let mc = &self.model.mc;
let d = self.model.cfg();
let pos = state.pos;
assert!(pos > 0, "decode_step needs prefill_with_cache first");
assert!(pos < self.max_seq, "pos {pos} >= max_seq {}", self.max_seq);
let hidden = mc.n_embd as usize;
let hc = d.hc_mult as usize;
let n_trunk = mc.n_layer - mc.nextn_predict_layers;
let ws_all = state.ws.as_mut().expect("device path needs StepWs");
let st0 = &self.stages[0];
st0.gpu.ctx.bind_to_thread().map_err(e("bind ctx0"))?;
let stream0 = st0.gpu.stream();
{
let ws0 = &mut ws_all[0];
stream0
.memcpy_htod(&[tok as i32], &mut ws0.tok)
.map_err(e("htod tok"))?;
unsafe {
ck(
"embed_rows dev",
k::memra_dsv4_embed_rows(
st0.embed
.as_ref()
.expect("embed on stage 0")
.device_ptr(&stream0)
.0 as *const c_void,
ws0.tok.device_ptr(&stream0).0 as *const i32,
dpm!(ws0.emb, &stream0),
1,
hidden as i32,
sp(&stream0),
),
)?;
ck(
"repeat_hc dev",
k::memra_dsv4_repeat_hc(
dpf!(ws0.emb, &stream0),
dpm!(ws0.h_a, &stream0),
1,
hc as i32,
hidden as i32,
sp(&stream0),
),
)?;
}
}
let mut cur_stage = 0usize;
let mut input_rx = false;
for il in 0..n_trunk {
let stage = self.layer_stage[il as usize];
if stage != cur_stage {
let bytes = hc * hidden * std::mem::size_of::<f32>();
let src_stream = self.stages[cur_stage].gpu.stream();
let dst_stream = self.stages[stage].gpu.stream();
let (ws_src, ws_dst) = ws_all.split_at_mut(stage);
let src_ws = &ws_src[cur_stage];
let dst_ws = &mut ws_dst[0];
self.stages[cur_stage]
.gpu
.ctx
.bind_to_thread()
.map_err(e("bind tx"))?;
let (sp_, _g0) = src_ws.h_a.device_ptr(&src_stream);
let (dp_, _g1) = dst_ws.h_rx.device_ptr_mut(&src_stream);
unsafe {
cudarc::driver::result::memcpy_peer_async(
self.stages[stage].gpu.ctx.cu_ctx(),
dp_,
self.stages[cur_stage].gpu.ctx.cu_ctx(),
sp_,
bytes,
src_stream.cu_stream(),
)
.map_err(e("peer copy h"))?;
}
let bnd = stage - 1;
self.boundary_ev[bnd]
.record(&src_stream)
.map_err(e("ev record"))?;
dst_stream
.wait(&self.boundary_ev[bnd])
.map_err(e("ev wait"))?;
self.stages[stage]
.gpu
.ctx
.bind_to_thread()
.map_err(e("bind rx"))?;
cur_stage = stage;
input_rx = true;
}
let st = &self.stages[stage];
let lidx = st
.layers
.iter()
.position(|l| l.il == il)
.unwrap_or_else(|| panic!("layer {il} not on stage {stage}"));
self.block_decode_dev(
st,
&st.layers[lidx],
&mut state.caches[il as usize],
&mut ws_all[stage],
input_rx,
pos,
tok,
host_math,
)?;
input_rx = false;
if let Some((t, base)) = taps.as_mut() {
if let Some(ds) = &self.dspark {
if let Some(k) = ds.targets.iter().position(|&tl| tl == il as usize) {
let stream = self.stages[stage].gpu.stream();
let hidden_i = hidden as i32;
unsafe {
ck(
"hc_mean tap dev",
k::memra_dsv4_hc_mean(
dpf!(ws_all[stage].h_a, &stream),
(t.device_ptr_mut(&stream).0 as usize
+ (*base + k * hidden) * 4)
as *mut f32,
1,
hc as i32,
hidden_i,
sp(&stream),
),
)?;
}
}
}
}
}
let last = self.stages.len() - 1;
assert_eq!(cur_stage, last, "device path expects the head stage last");
self.head_logits_dev(&mut ws_all[last], host_math)?;
let stream_last = self.stages[last].gpu.stream();
state.pos += 1;
if want_logits {
let logits = dtoh_f32(&stream_last, &ws_all[last].logits)?;
let mut best = 0usize;
for i in 1..logits.len() {
if logits[i] > logits[best] {
best = i;
}
}
Ok((Some(logits), best as u32))
} else {
unsafe {
ck(
"argmax dev",
k::memra_dsv4_argmax(
dpf!(ws_all[last].logits, &stream_last),
ws_all[last].logits.len() as i64,
ws_all[last].argmax.device_ptr_mut(&stream_last).0 as *mut i32,
sp(&stream_last),
),
)?;
}
let mut out = [0i32; 1];
stream_last
.memcpy_dtoh(&ws_all[last].argmax, &mut out[..])
.map_err(e("dtoh argmax"))?;
stream_last.synchronize().map_err(e("sync argmax"))?;
Ok((None, out[0] as u32))
}
}
pub fn decode_step_greedy(&self, tok: u32, state: &mut DecodeState) -> Res<u32> {
match self.decode_path {
DecodePath::Legacy => {
let logits = self.decode_step_impl(tok, state, None)?;
let mut best = 0usize;
for i in 1..logits.len() {
if logits[i] > logits[best] {
best = i;
}
}
Ok(best as u32)
}
DecodePath::Device { host_math } => {
Ok(self.decode_step_fast(tok, state, false, host_math)?.1)
}
}
}
fn decode_step_impl(
&self,
tok: u32,
state: &mut DecodeState,
mut dump: Option<&mut Vec<(String, Vec<f32>)>>,
) -> Res<Vec<f32>> {
if let DecodePath::Device { host_math } = self.decode_path {
assert!(
dump.is_none(),
"decode_step_probe is a legacy-path diagnostic (set MEMRA_DSV4_DECODE_PATH=legacy)"
);
let (logits, _) = self.decode_step_fast(tok, state, true, host_math)?;
return Ok(logits.expect("want_logits"));
}
let mc = &self.model.mc;
let d = self.model.cfg();
let pos = state.pos;
assert!(pos > 0, "decode_step needs prefill_with_cache first");
assert!(pos < self.max_seq, "pos {pos} >= max_seq {}", self.max_seq);
let hidden = mc.n_embd as usize;
let hc = d.hc_mult as usize;
let n_trunk = mc.n_layer - mc.nextn_predict_layers;
let st0 = &self.stages[0];
st0.gpu.ctx.bind_to_thread().map_err(e("bind ctx0"))?;
let stream0 = st0.gpu.stream();
let ids_dev = upload_i32(&stream0, &[tok as i32])?;
let mut emb = stream0.alloc_zeros::<f32>(hidden).map_err(e("emb"))?;
unsafe {
ck(
"embed_rows",
k::memra_dsv4_embed_rows(
st0.embed
.as_ref()
.expect("embed on stage 0")
.device_ptr(&stream0)
.0 as *const c_void,
ids_dev.device_ptr(&stream0).0 as *const i32,
dpm!(emb, &stream0),
1,
hidden as i32,
sp(&stream0),
),
)?;
}
let mut h = stream0.alloc_zeros::<f32>(hc * hidden).map_err(e("h0"))?;
unsafe {
ck(
"repeat_hc",
k::memra_dsv4_repeat_hc(
dpf!(emb, &stream0),
dpm!(h, &stream0),
1,
hc as i32,
hidden as i32,
sp(&stream0),
),
)?;
}
let mut cur_stage = 0usize;
for il in 0..n_trunk {
let stage = self.layer_stage[il as usize];
if stage != cur_stage {
let src_stream = self.stages[cur_stage].gpu.stream();
let host = dtoh_f32(&src_stream, &h)?;
let dst_stream = self.stages[stage].gpu.stream();
self.stages[stage]
.gpu
.ctx
.bind_to_thread()
.map_err(e("bind"))?;
h = upload_f32(&dst_stream, &host)?;
cur_stage = stage;
}
let st = &self.stages[stage];
let lidx = st
.layers
.iter()
.position(|l| l.il == il)
.unwrap_or_else(|| panic!("layer {il} not on stage {stage}"));
h = self.block_decode(
st,
&st.layers[lidx],
&mut state.caches[il as usize],
&h,
pos,
tok,
dump.as_deref_mut(),
)?;
}
let last = self.stages.len() - 1;
if cur_stage != last {
let src_stream = self.stages[cur_stage].gpu.stream();
let host = dtoh_f32(&src_stream, &h)?;
let dst_stream = self.stages[last].gpu.stream();
h = upload_f32(&dst_stream, &host)?;
}
let hc_head_fn = self.stages[last].hc_head_fn.as_ref().expect("hc_head_fn");
let trunk_norm = self.stages[last].trunk_norm.as_ref().expect("trunk norm");
let logits = self.head_logits_from(
&h,
1,
hc_head_fn,
&self.hc_head_base,
&self.hc_head_scale,
trunk_norm,
)?;
state.pos += 1;
Ok(logits)
}
}
impl Dsv4Gpu {
fn dspark(&self) -> &DsparkDev {
self.dspark
.as_ref()
.expect("MEMRA_DSV4_DRAFTER=dspark not loaded")
}
#[allow(clippy::too_many_arguments)]
fn dspark_head_dots(
&self,
st: &Stage,
x: *const f32,
w: *const c_void,
w_is_bf16: i32,
s: usize,
kdim: usize,
n: usize,
y: *mut f32,
) -> Res<()> {
let stream = st.gpu.stream();
unsafe {
if self.dspark_head_f32 {
ck(
"dspark head dots f32acc_mrow",
k::memra_dsv4_dots_f32acc_mrow(
x,
w,
w_is_bf16,
y,
s as i32,
kdim as i32,
n as i32,
sp(&stream),
),
)
} else {
ck(
"dspark head dots f32_mrow",
k::memra_dsv4_dots_f32_mrow(
x,
w,
w_is_bf16,
y,
s as i32,
kdim as i32,
n as i32,
sp(&stream),
),
)
}
}
}
pub fn dspark_alloc_state(&self) -> Res<DsparkState> {
let ds = self.dspark();
let d = self.model.cfg();
let hd = d.head_dim as usize;
let win = d.sliding_window as usize;
let hidden = self.model.mc.n_embd as usize;
let last = self.stages.len() - 1;
let stream = self.stages[last].gpu.stream();
let mut rings = Vec::with_capacity(ds.blocks.len());
for _ in 0..ds.blocks.len() {
rings.push(
stream
.alloc_zeros::<f32>((win + ds.block_size) * hd)
.map_err(e("dspark ring"))?,
);
}
let taps = stream
.alloc_zeros::<f32>((ds.block_size + 1) * ds.targets.len() * hidden)
.map_err(e("dspark taps"))?;
Ok(DsparkState { rings, taps })
}
fn dspark_main_x(&self, main_hidden: &CudaSlice<f32>, s: usize) -> Res<CudaSlice<f32>> {
let ds = self.dspark();
let hidden = self.model.mc.n_embd as usize;
let k = ds.targets.len() * hidden;
let last = self.stages.len() - 1;
let st = &self.stages[last];
let stream = st.gpu.stream();
let mut mx = stream.alloc_zeros::<f32>(s * hidden).map_err(e("main_x"))?;
Self::gemm(st, main_hidden, &ds.main_proj, 0, s, hidden, k, &mut mx)?;
unsafe {
ck(
"rmsnorm main_x",
k::memra_dsv4_rmsnorm(
dpf!(mx, &stream),
dpf!(ds.main_norm, &stream),
dpm!(mx, &stream),
s as i32,
hidden as i32,
self.model.mc.rms_eps,
sp(&stream),
),
)?;
}
Ok(mx)
}
fn dspark_main_kv(
&self,
blk: &LayerDev,
main_x: &CudaSlice<f32>,
s: usize,
positions: &[i32],
) -> Res<CudaSlice<f32>> {
let d = self.model.cfg();
let hd = d.head_dim as usize;
let rd = d.qk_rope_head_dim as usize;
let hidden = self.model.mc.n_embd as usize;
let eps = self.model.mc.rms_eps;
let last = self.stages.len() - 1;
let st = &self.stages[last];
let stream = st.gpu.stream();
let clamp_only = (self.variant == ActQuantVariant::ClampOnly) as i32;
let mut kv = stream.alloc_zeros::<f32>(s * hd).map_err(e("dspark kv"))?;
Self::gemm(st, main_x, blk.wkv.dev(), 0, s, hd, hidden, &mut kv)?;
let pos_dev = upload_i32(&stream, positions)?;
unsafe {
ck(
"rmsnorm dspark kv",
k::memra_dsv4_rmsnorm(
dpf!(kv, &stream),
dpf!(blk.kv_norm, &stream),
dpm!(kv, &stream),
s as i32,
hd as i32,
eps,
sp(&stream),
),
)?;
ck(
"rope dspark kv",
k::memra_dsv4_rope(
dpm!(kv, &stream),
s as i32,
1,
hd as i32,
rd as i32,
dpf!(st.fc_plain, &stream),
pos_dev.device_ptr(&stream).0 as *const i32,
0,
sp(&stream),
),
)?;
ck(
"act_quant dspark kv",
k::memra_dsv4_act_quant(
dpm!(kv, &stream),
s as i32,
hd as i64,
(hd - rd) as i32,
64,
clamp_only,
sp(&stream),
),
)?;
}
Ok(kv)
}
pub fn dspark_prime_prefill(
&self,
state: &mut DsparkState,
main_hidden: &CudaSlice<f32>,
s: usize,
) -> Res<()> {
let d = self.model.cfg();
let hd = d.head_dim as usize;
let win = d.sliding_window as usize;
let last = self.stages.len() - 1;
let stream = self.stages[last].gpu.stream();
let mx = self.dspark_main_x(main_hidden, s)?;
let positions: Vec<i32> = (0..s as i32).collect();
let n_blocks = self.dspark().blocks.len();
for bi in 0..n_blocks {
let blk = &self.dspark().blocks[bi];
let kv = self.dspark_main_kv(blk, &mx, s, &positions)?;
for p in s.saturating_sub(win)..s {
let slot = p % win;
let src = kv.slice(p * hd..(p + 1) * hd);
let mut dst = state.rings[bi].slice_mut(slot * hd..(slot + 1) * hd);
stream
.memcpy_dtod(&src, &mut dst)
.map_err(e("prime ring"))?;
}
}
Ok(())
}
pub fn dspark_prefill_prime(
&self,
ids: &[u32],
state: &mut DecodeState,
dstate: &mut DsparkState,
) -> Res<ForwardOut> {
assert_eq!(
state.pos, 0,
"dspark_prefill_prime needs a fresh DecodeState"
);
assert!(!ids.is_empty(), "empty prompt");
let hidden = self.model.mc.n_embd as usize;
let hc = self.model.cfg().hc_mult as usize;
let s = ids.len();
let targets = self.dspark().targets.clone();
let n_t = targets.len();
let mut cap = GpuCapture {
want: targets.iter().map(|&t| t as u32).collect(),
..Default::default()
};
let out = self
.forward_impl(ids, Some(&mut cap), None, Some(state))?
.expect("prefill logits");
state.pos = s;
let last = self.stages.len() - 1;
let stream = self.stages[last].gpu.stream();
self.stages[last]
.gpu
.ctx
.bind_to_thread()
.map_err(e("bind ctx prime"))?;
let mut main_hidden = stream
.alloc_zeros::<f32>(s * n_t * hidden)
.map_err(e("prefill main_hidden"))?;
let mut tmp = stream
.alloc_zeros::<f32>(s * hidden)
.map_err(e("tap tmp"))?;
for (k, &il) in targets.iter().enumerate() {
let h = cap
.layer_out
.get(&(il as u32))
.unwrap_or_else(|| panic!("prefill capture missing target layer {il}"));
assert_eq!(
h.len(),
s * hc * hidden,
"target layer {il} capture is not [s, hc, hidden]"
);
let h_dev = upload_f32(&stream, h)?;
unsafe {
ck(
"hc_mean prefill tap",
k::memra_dsv4_hc_mean(
dpf!(h_dev, &stream),
dpm!(tmp, &stream),
s as i32,
hc as i32,
hidden as i32,
sp(&stream),
),
)?;
ck(
"place_cols prefill tap",
k::memra_dsv4_place_cols(
dpf!(tmp, &stream),
dpm!(main_hidden, &stream),
s as i32,
hidden as i32,
(n_t * hidden) as i64,
(k * hidden) as i64,
sp(&stream),
),
)?;
}
}
self.dspark_prime_prefill(dstate, &main_hidden, s)?;
{
let src = main_hidden.slice((s - 1) * n_t * hidden..s * n_t * hidden);
let mut dst = dstate.taps.slice_mut(0..n_t * hidden);
stream
.memcpy_dtod(&src, &mut dst)
.map_err(e("seed tap row"))?;
}
stream.synchronize().map_err(e("prime sync"))?;
Ok(out)
}
pub fn dspark_write_rings(
&self,
state: &mut DsparkState,
tap_row: usize,
pos: usize,
) -> Res<()> {
let d = self.model.cfg();
let hd = d.head_dim as usize;
let win = d.sliding_window as usize;
let hidden = self.model.mc.n_embd as usize;
let n_t = self.dspark().targets.len();
let last = self.stages.len() - 1;
let stream = self.stages[last].gpu.stream();
let tap = {
let mut row = stream
.alloc_zeros::<f32>(n_t * hidden)
.map_err(e("tap row"))?;
let src = state
.taps
.slice(tap_row * n_t * hidden..(tap_row + 1) * n_t * hidden);
stream.memcpy_dtod(&src, &mut row).map_err(e("tap copy"))?;
row
};
let mx = self.dspark_main_x(&tap, 1)?;
let n_blocks = self.dspark().blocks.len();
for bi in 0..n_blocks {
let blk = &self.dspark().blocks[bi];
let kv = self.dspark_main_kv(blk, &mx, 1, &[pos as i32])?;
let slot = pos % win;
let src = kv.slice(0..hd);
let mut dst = state.rings[bi].slice_mut(slot * hd..(slot + 1) * hd);
stream
.memcpy_dtod(&src, &mut dst)
.map_err(e("ring write"))?;
}
Ok(())
}
#[allow(clippy::too_many_arguments)]
fn dspark_block_forward(
&self,
blk: &LayerDev,
ring: &mut CudaSlice<f32>,
h: &CudaSlice<f32>,
block: usize,
pos: usize,
) -> Res<CudaSlice<f32>> {
let d = self.model.cfg();
let mc = &self.model.mc;
let hc = d.hc_mult as usize;
let hidden = mc.n_embd as usize;
let heads = mc.n_head as usize;
let hd = d.head_dim as usize;
let rd = d.qk_rope_head_dim as usize;
let q_lora = d.q_lora_rank as usize;
let win = d.sliding_window as usize;
let o_groups = d.o_groups as usize;
let o_lora = d.o_lora_rank as usize;
let eps = mc.rms_eps;
let iters = d.hc_sinkhorn_iters;
let hc_eps = d.hc_eps;
let last = self.stages.len() - 1;
let st = &self.stages[last];
st.gpu.ctx.bind_to_thread().map_err(e("bind ctx dspark"))?;
let stream = st.gpu.stream();
let clamp_only = (self.variant == ActQuantVariant::ClampOnly) as i32;
let positions: Vec<i32> = (1..=block as i32).map(|j| pos as i32 + j).collect();
let pos_dev = upload_i32(&stream, &positions)?;
let (y, post, comb) = Self::hc_pre(
st,
h,
&blk.hc_attn_fn,
&blk.hc_attn_base,
&blk.hc_attn_scale,
block,
hc,
hidden,
iters,
hc_eps,
)?;
let mut x = stream.alloc_zeros::<f32>(block * hidden).map_err(e("x"))?;
unsafe {
ck(
"rmsnorm dspark attn",
k::memra_dsv4_rmsnorm(
dpf!(y, &stream),
dpf!(blk.attn_norm, &stream),
dpm!(x, &stream),
block as i32,
hidden as i32,
eps,
sp(&stream),
),
)?;
}
let mut qr = stream.alloc_zeros::<f32>(block * q_lora).map_err(e("qr"))?;
Self::gemm(st, &x, blk.wq_a.dev(), 0, block, q_lora, hidden, &mut qr)?;
unsafe {
ck(
"rmsnorm dspark q",
k::memra_dsv4_rmsnorm(
dpf!(qr, &stream),
dpf!(blk.q_norm, &stream),
dpm!(qr, &stream),
block as i32,
q_lora as i32,
eps,
sp(&stream),
),
)?;
}
let mut qr_b = stream
.alloc_zeros::<u8>(block * q_lora * 2)
.map_err(e("qr_b"))?;
unsafe {
ck(
"cvt dspark qr",
k::memra_dsv4_cvt_bf16(
dpf!(qr, &stream),
qr_b.device_ptr_mut(&stream).0 as *mut c_void,
(block * q_lora) as i64,
sp(&stream),
),
)?;
}
let mut q = stream
.alloc_zeros::<f32>(block * heads * hd)
.map_err(e("q"))?;
Self::gemm_pre(
st,
&qr_b,
blk.wq_b.dev().device_ptr(&stream).0 as *const c_void,
block,
heads * hd,
q_lora,
&mut q,
)?;
unsafe {
ck(
"headrms dspark",
k::memra_dsv4_headrms(
dpm!(q, &stream),
(block * heads) as i32,
hd as i32,
eps,
sp(&stream),
),
)?;
ck(
"rope dspark q",
k::memra_dsv4_rope(
dpm!(q, &stream),
block as i32,
heads as i32,
hd as i32,
rd as i32,
dpf!(st.fc_plain, &stream),
pos_dev.device_ptr(&stream).0 as *const i32,
0,
sp(&stream),
),
)?;
}
{
let mut kv = stream.alloc_zeros::<f32>(block * hd).map_err(e("dkv"))?;
Self::gemm(st, &x, blk.wkv.dev(), 0, block, hd, hidden, &mut kv)?;
unsafe {
ck(
"rmsnorm dspark dkv",
k::memra_dsv4_rmsnorm(
dpf!(kv, &stream),
dpf!(blk.kv_norm, &stream),
dpm!(kv, &stream),
block as i32,
hd as i32,
eps,
sp(&stream),
),
)?;
ck(
"rope dspark dkv",
k::memra_dsv4_rope(
dpm!(kv, &stream),
block as i32,
1,
hd as i32,
rd as i32,
dpf!(st.fc_plain, &stream),
pos_dev.device_ptr(&stream).0 as *const i32,
0,
sp(&stream),
),
)?;
ck(
"act_quant dspark dkv",
k::memra_dsv4_act_quant(
dpm!(kv, &stream),
block as i32,
hd as i64,
(hd - rd) as i32,
64,
clamp_only,
sp(&stream),
),
)?;
}
let src = kv.slice(0..block * hd);
let mut dst = ring.slice_mut(win * hd..(win + block) * hd);
stream.memcpy_dtod(&src, &mut dst).map_err(e("draft kv"))?;
}
let n_ring = win.min(pos + 1);
let mut idx_row: Vec<i32> = (0..n_ring as i32).collect();
idx_row.extend((0..block as i32).map(|j| win as i32 + j));
let slots = idx_row.len();
let mut idxs = Vec::with_capacity(block * slots);
for _ in 0..block {
idxs.extend_from_slice(&idx_row);
}
let idx_dev = upload_i32(&stream, &idxs)?;
let mut o = stream
.alloc_zeros::<f32>(block * heads * hd)
.map_err(e("o"))?;
let scale = (hd as f64).powf(-0.5) as f32;
unsafe {
ck(
"sink_attn dspark",
k::memra_dsv4_sink_attn(
dpf!(q, &stream),
dpf!(ring, &stream),
idx_dev.device_ptr(&stream).0 as *const i32,
dpf!(blk.sink, &stream),
dpm!(o, &stream),
block as i32,
heads as i32,
hd as i32,
slots as i32,
scale,
sp(&stream),
),
)?;
ck(
"rope dspark o inv",
k::memra_dsv4_rope(
dpm!(o, &stream),
block as i32,
heads as i32,
hd as i32,
rd as i32,
dpf!(st.fc_plain, &stream),
pos_dev.device_ptr(&stream).0 as *const i32,
1,
sp(&stream),
),
)?;
}
let gw = heads / o_groups * hd;
let mut og = stream
.alloc_zeros::<f32>(block * o_groups * o_lora)
.map_err(e("og"))?;
let mut o_grp = stream.alloc_zeros::<f32>(block * gw).map_err(e("o_grp"))?;
let mut y_grp = stream
.alloc_zeros::<f32>(block * o_lora)
.map_err(e("y_grp"))?;
for g in 0..o_groups {
unsafe {
ck(
"take_cols dspark",
k::memra_dsv4_take_cols(
dpf!(o, &stream),
dpm!(o_grp, &stream),
block as i32,
gw as i32,
(heads * hd) as i64,
(g * gw) as i64,
sp(&stream),
),
)?;
}
Self::gemm(
st,
&o_grp,
blk.wo_a.dev(),
g * o_lora * gw,
block,
o_lora,
gw,
&mut y_grp,
)?;
unsafe {
ck(
"place_cols dspark",
k::memra_dsv4_place_cols(
dpf!(y_grp, &stream),
dpm!(og, &stream),
block as i32,
o_lora as i32,
(o_groups * o_lora) as i64,
(g * o_lora) as i64,
sp(&stream),
),
)?;
}
}
let mut attn_out = stream.alloc_zeros::<f32>(block * hidden).map_err(e("ao"))?;
Self::gemm(
st,
&og,
blk.wo_b.dev(),
0,
block,
hidden,
o_groups * o_lora,
&mut attn_out,
)?;
let mut h2 = stream
.alloc_zeros::<f32>(block * hc * hidden)
.map_err(e("h2"))?;
unsafe {
ck(
"hc_post dspark attn",
k::memra_dsv4_hc_post(
dpf!(attn_out, &stream),
dpf!(h, &stream),
dpf!(post, &stream),
dpf!(comb, &stream),
dpm!(h2, &stream),
block as i32,
hc as i32,
hidden as i32,
sp(&stream),
),
)?;
}
let (y2, post2, comb2) = Self::hc_pre(
st,
&h2,
&blk.hc_ffn_fn,
&blk.hc_ffn_base,
&blk.hc_ffn_scale,
block,
hc,
hidden,
iters,
hc_eps,
)?;
let mut xf = stream.alloc_zeros::<f32>(block * hidden).map_err(e("xf"))?;
unsafe {
ck(
"rmsnorm dspark ffn",
k::memra_dsv4_rmsnorm(
dpf!(y2, &stream),
dpf!(blk.ffn_norm, &stream),
dpm!(xf, &stream),
block as i32,
hidden as i32,
eps,
sp(&stream),
),
)?;
}
let ids = vec![0u32; block];
let moe_out = self.moe_forward(st, blk, &xf, block, &ids)?;
let mut h3 = stream
.alloc_zeros::<f32>(block * hc * hidden)
.map_err(e("h3"))?;
unsafe {
ck(
"hc_post dspark ffn",
k::memra_dsv4_hc_post(
dpf!(moe_out, &stream),
dpf!(h2, &stream),
dpf!(post2, &stream),
dpf!(comb2, &stream),
dpm!(h3, &stream),
block as i32,
hc as i32,
hidden as i32,
sp(&stream),
),
)?;
}
Ok(h3)
}
pub fn dspark_forward_spec(
&self,
state: &mut DsparkState,
input_token: u32,
tap_row: usize,
pos: usize,
capture: bool,
) -> Res<DsparkProposal> {
let ds = self.dspark();
let mc = &self.model.mc;
let d = self.model.cfg();
let hc = d.hc_mult as usize;
let hidden = mc.n_embd as usize;
let eps = mc.rms_eps;
let block = ds.block_size;
let rank = ds.rank;
let vocab = ds.vocab;
let n_t = ds.targets.len();
let last = self.stages.len() - 1;
let st = &self.stages[last];
st.gpu.ctx.bind_to_thread().map_err(e("bind ctx spec"))?;
let stream = st.gpu.stream();
let prof = if dsv4_prof_on() {
Some(stream.clone())
} else {
None
};
let tap = {
let _p = phase!("1a.tap_copy", prof.as_ref());
let mut row = stream.alloc_zeros::<f32>(n_t * hidden).map_err(e("tapr"))?;
let src = state
.taps
.slice(tap_row * n_t * hidden..(tap_row + 1) * n_t * hidden);
stream.memcpy_dtod(&src, &mut row).map_err(e("tap cp"))?;
row
};
let mx = {
let _p = phase!("1b.main_x", prof.as_ref());
self.dspark_main_x(&tap, 1)?
};
let (cap_main_hidden, cap_main_x) = if capture {
(
Some(dtoh_f32(&stream, &tap)?),
Some(dtoh_f32(&stream, &mx)?),
)
} else {
(None, None)
};
let _p_embed = phase!("1c.embed_h2d_repeat", prof.as_ref());
let mut draft_ids = vec![ds.noise_token; block];
draft_ids[0] = input_token;
let e_rows = self.model.embed_rows(&draft_ids);
let e_dev = upload_f32(&stream, &e_rows)?;
let mut h = stream
.alloc_zeros::<f32>(block * hc * hidden)
.map_err(e("h0"))?;
unsafe {
ck(
"repeat_hc dspark",
k::memra_dsv4_repeat_hc(
dpf!(e_dev, &stream),
dpm!(h, &stream),
block as i32,
hc as i32,
hidden as i32,
sp(&stream),
),
)?;
}
drop(_p_embed);
let mut block_outs: Vec<Vec<f32>> = Vec::new();
let _p_blocks = phase!("1d.drafter_blocks", prof.as_ref());
let n_blocks = ds.blocks.len();
for bi in 0..n_blocks {
let mut ring = std::mem::replace(
&mut state.rings[bi],
stream.alloc_zeros::<f32>(0).map_err(e("swap"))?,
);
let out =
self.dspark_block_forward(&self.dspark().blocks[bi], &mut ring, &h, block, pos);
state.rings[bi] = ring;
h = out?;
if capture {
block_outs.push(dtoh_f32(&stream, &h)?);
}
}
drop(_p_blocks);
let w = hc * hidden;
let _p_mix = phase!("1e.exit_mix_dots", prof.as_ref());
let mut mixes = stream.alloc_zeros::<f32>(block * hc).map_err(e("mx"))?;
self.dspark_head_dots(
st,
h.device_ptr(&stream).0 as *const f32,
ds.hc_head_fn.device_ptr(&stream).0 as *const c_void,
0,
block,
w,
hc,
mixes.device_ptr_mut(&stream).0 as *mut f32,
)?;
unsafe {
ck(
"rowsq dspark head",
k::memra_dsv4_rowsq_scale(
dpf!(h, &stream),
dpm!(mixes, &stream),
block as i32,
w as i32,
hc as i32,
eps,
sp(&stream),
),
)?;
}
drop(_p_mix);
let _p_mixrt = phase!("1f.mix_D2H_host_H2D", prof.as_ref());
let mut mixes_h = dtoh_f32(&stream, &mixes)?;
for t in 0..block {
for c in 0..hc {
let m = mixes_h[t * hc + c];
mixes_h[t * hc + c] =
sigmoid_f32(m * ds.hc_head_scale[0] + ds.hc_head_base[c]) + d.hc_eps;
}
}
let pre_d = upload_f32(&stream, &mixes_h)?;
drop(_p_mixrt);
let _p_cn = phase!("1g.collapse_norm", prof.as_ref());
let mut xc = stream.alloc_zeros::<f32>(block * hidden).map_err(e("xc"))?;
unsafe {
ck(
"hc_collapse dspark",
k::memra_dsv4_hc_collapse(
dpf!(h, &stream),
dpf!(pre_d, &stream),
dpm!(xc, &stream),
block as i32,
hc as i32,
hidden as i32,
sp(&stream),
),
)?;
}
let mut normed = stream.alloc_zeros::<f32>(block * hidden).map_err(e("nr"))?;
unsafe {
ck(
"rmsnorm dspark head",
k::memra_dsv4_rmsnorm(
dpf!(xc, &stream),
dpf!(ds.norm, &stream),
dpm!(normed, &stream),
block as i32,
hidden as i32,
eps,
sp(&stream),
),
)?;
}
drop(_p_cn);
let _p_head = phase!("1h.exit_head_dots", prof.as_ref());
let mut logits = stream.alloc_zeros::<f32>(block * vocab).map_err(e("lg"))?;
self.dspark_head_dots(
st,
normed.device_ptr(&stream).0 as *const f32,
st.head.as_ref().expect("head").device_ptr(&stream).0 as *const c_void,
1,
block,
hidden,
vocab,
logits.device_ptr_mut(&stream).0 as *mut f32,
)?;
drop(_p_head);
let cap_logits_pre = if capture {
Some(dtoh_f32(&stream, &logits)?)
} else {
None
};
let chain_device = dsv4_dspark_chain_device();
let markov_rowblk = dsv4_dspark_markov_rowblk();
let mut w1_row = stream.alloc_zeros::<f32>(rank).map_err(e("w1r"))?;
let mut bias = stream.alloc_zeros::<f32>(vocab).map_err(e("bias"))?;
let mut am_dev = stream.alloc_zeros::<i32>(block + 1).map_err(e("am"))?;
{
let mut dst = am_dev.slice_mut(0..1);
stream
.memcpy_htod(&[input_token as i32][..], &mut dst)
.map_err(e("htod am0"))?;
}
let mut out_ids = vec![input_token];
let mut margins = Vec::with_capacity(block);
let mut top1_logits = Vec::with_capacity(block);
let mut conf_in = stream.alloc_zeros::<f32>(hidden + rank).map_err(e("cin"))?;
let mut conf_out = stream.alloc_zeros::<f32>(block).map_err(e("cout"))?;
let mut confidence = Vec::with_capacity(block);
let _p_mk = phase!("1i.markov_chain", prof.as_ref());
for i in 0..block {
{
let _p = phase!("1i1.markov_w1_gather", prof.as_ref());
if chain_device {
unsafe {
ck(
"markov w1 gather dev",
k::memra_dsv4_gather_row_by_idx(
dpf!(ds.markov_w1, &stream),
am_dev.device_ptr(&stream).0 as *const i32,
i as i32,
dpm!(w1_row, &stream),
rank as i32,
sp(&stream),
),
)?;
}
} else {
let prev = out_ids[i] as usize;
let src = ds.markov_w1.slice(prev * rank..(prev + 1) * rank);
stream.memcpy_dtod(&src, &mut w1_row).map_err(e("w1 cp"))?;
}
}
{
let _p = phase!("1i2.markov_bias_gemv", prof.as_ref());
if markov_rowblk {
unsafe {
ck(
"dots_f32 markov rowblk",
k::memra_dsv4_dots_f32_rowblk(
dpf!(w1_row, &stream),
dp!(ds.markov_w2, &stream),
0,
dpm!(bias, &stream),
1,
rank as i32,
vocab as i32,
sp(&stream),
),
)?;
}
} else {
Self::dots(st, &w1_row, &ds.markov_w2, 1, rank, vocab, &mut bias)?;
}
}
let _p_aa = phase!("1i3.markov_add_argmax", prof.as_ref());
unsafe {
ck(
"markov add dspark",
k::memra_dsv4_add_inplace(
(logits.device_ptr_mut(&stream).0 as usize + i * vocab * 4) as *mut f32,
dpf!(bias, &stream),
vocab as i64,
sp(&stream),
),
)?;
ck(
"argmax dspark",
k::memra_dsv4_argmax(
(logits.device_ptr(&stream).0 as usize + i * vocab * 4) as *const f32,
vocab as i64,
(am_dev.device_ptr_mut(&stream).0 as usize + (i + 1) * 4) as *mut i32,
sp(&stream),
),
)?;
}
drop(_p_aa);
if !chain_device {
let _p_d2h = phase!("1i4.markov_argmax_D2H_SYNC", None);
let mut am = [0i32; 1];
let view = am_dev.slice(i + 1..i + 2);
stream
.memcpy_dtoh(&view, &mut am[..])
.map_err(e("dtoh am"))?;
stream.synchronize().map_err(e("sync am"))?;
out_ids.push(am[0] as u32);
}
{
let _p = phase!("1i5.conf_in_copies", prof.as_ref());
let src = xc.slice(i * hidden..(i + 1) * hidden);
let mut dst = conf_in.slice_mut(0..hidden);
stream.memcpy_dtod(&src, &mut dst).map_err(e("cin x"))?;
let src = w1_row.slice(0..rank);
let mut dst = conf_in.slice_mut(hidden..hidden + rank);
stream.memcpy_dtod(&src, &mut dst).map_err(e("cin m"))?;
}
{
let _p = phase!("1i6.conf_dots", prof.as_ref());
unsafe {
ck(
"dots_f32 conf dspark",
k::memra_dsv4_dots_f32(
dpf!(conf_in, &stream),
dp!(ds.conf_w, &stream),
0,
(conf_out.device_ptr_mut(&stream).0 as usize + i * 4) as *mut f32,
1,
(hidden + rank) as i32,
1,
sp(&stream),
),
)?;
}
}
if !chain_device {
let _p = phase!("1i7.conf_D2H_SYNC", None);
let mut c = [0f32; 1];
let view = conf_out.slice(i..i + 1);
stream
.memcpy_dtoh(&view, &mut c[..])
.map_err(e("dtoh cf"))?;
stream.synchronize().map_err(e("sync cf"))?;
confidence.push(c[0]);
}
}
if chain_device {
let _p = phase!("1i8.chain_D2H_SYNC_once", None);
let mut ids = vec![0i32; block];
let view = am_dev.slice(1..block + 1);
stream
.memcpy_dtoh(&view, &mut ids[..])
.map_err(e("dtoh chain ids"))?;
let mut cf = vec![0f32; block];
stream
.memcpy_dtoh(&conf_out, &mut cf[..])
.map_err(e("dtoh chain conf"))?;
stream.synchronize().map_err(e("sync chain"))?;
out_ids.extend(ids.iter().map(|&x| x as u32));
confidence.extend_from_slice(&cf);
}
drop(_p_mk);
let membeds: Vec<f32> = if capture {
let mut m = Vec::with_capacity(block * rank);
for i in 0..block {
let prev = out_ids[i] as usize;
m.extend_from_slice(&ds.markov_w1_host[prev * rank..(prev + 1) * rank]);
}
m
} else {
Vec::new()
};
let cap = if capture {
let logits_post = dtoh_f32(&stream, &logits)?;
for i in 0..block {
let row = &logits_post[i * vocab..(i + 1) * vocab];
let top = out_ids[i + 1];
let mut second = f32::NEG_INFINITY;
for (vv, &val) in row.iter().enumerate() {
if vv as u32 != top && val > second {
second = val;
}
}
margins.push(row[top as usize] - second);
top1_logits.push(row[top as usize]);
}
Some(DsparkCaptureOut {
main_hidden: cap_main_hidden.unwrap(),
main_x: cap_main_x.unwrap(),
block_outs,
x_collapsed: dtoh_f32(&stream, &xc)?,
logits_pre: cap_logits_pre.unwrap(),
logits_post,
markov_embed: membeds,
})
} else {
None
};
Ok(DsparkProposal {
out_ids,
confidence,
margins,
top1_logits,
capture: cap,
})
}
pub fn decode_step_tap(
&self,
tok: u32,
state: &mut DecodeState,
dspark_state: &mut DsparkState,
tap_row: usize,
) -> Res<Vec<f32>> {
let DecodePath::Device { host_math } = self.decode_path else {
return Err("decode_step_tap requires MEMRA_DSV4_DECODE_PATH=device".into());
};
let n_t = self.dspark().targets.len();
let hidden = self.model.mc.n_embd as usize;
let (logits, _) = self.decode_step_fast_tap(
tok,
state,
true,
host_math,
Some((&mut dspark_state.taps, tap_row * n_t * hidden)),
)?;
Ok(logits.expect("want_logits"))
}
pub fn decode_step_greedy_tap(
&self,
tok: u32,
state: &mut DecodeState,
dspark_state: &mut DsparkState,
tap_row: usize,
) -> Res<u32> {
let DecodePath::Device { host_math } = self.decode_path else {
return Err("decode_step_greedy_tap requires MEMRA_DSV4_DECODE_PATH=device".into());
};
let ds = self.dspark();
let hidden = self.model.mc.n_embd as usize;
let n_t = ds.targets.len();
let (_, tok_next) = self.decode_step_fast_tap(
tok,
state,
false,
host_math,
Some((&mut dspark_state.taps, tap_row * n_t * hidden)),
)?;
Ok(tok_next)
}
}
pub struct VerifyWs {
pub tmax: usize,
h_a: CudaSlice<f32>,
h_b: CudaSlice<f32>,
h_rx: CudaSlice<f32>,
emb: CudaSlice<f32>,
mixes: CudaSlice<f32>,
pre: CudaSlice<f32>,
post: CudaSlice<f32>,
comb: CudaSlice<f32>,
y_hc: CudaSlice<f32>,
x: CudaSlice<f32>,
xf: CudaSlice<f32>,
qr: CudaSlice<f32>,
qr_b: CudaSlice<u8>,
q: CudaSlice<f32>,
kv: CudaSlice<f32>,
qi: CudaSlice<f32>,
wproj: CudaSlice<f32>,
score: CudaSlice<f32>,
idx: CudaSlice<i32>,
idx_stride: usize,
o: CudaSlice<f32>,
o_b: CudaSlice<u8>,
og: CudaSlice<f32>,
attn_out: CudaSlice<f32>,
gemm_xb: CudaSlice<u8>,
raw: CudaSlice<f32>,
sel: CudaSlice<i32>,
selw: CudaSlice<f32>,
order: CudaSlice<i32>,
xq: CudaSlice<u8>,
xs: CudaSlice<f32>,
g1: CudaSlice<f32>,
g3: CudaSlice<f32>,
hbuf: CudaSlice<f32>,
hq: CudaSlice<u8>,
hs: CudaSlice<f32>,
contrib: CudaSlice<f32>,
y: CudaSlice<f32>,
xb: CudaSlice<u8>,
sg1: CudaSlice<f32>,
sg3: CudaSlice<f32>,
shbuf: CudaSlice<f32>,
shb16: CudaSlice<u8>,
sh_out: CudaSlice<f32>,
cmp_emit: CudaSlice<f32>,
cmp_shift: CudaSlice<f32>,
sink_scores: CudaSlice<f32>,
sink_evals: CudaSlice<f32>,
sink_den: CudaSlice<f64>,
head_mixes: CudaSlice<f32>,
head_pre: CudaSlice<f32>,
collapsed: CudaSlice<f32>,
logits: CudaSlice<f32>,
tok: CudaSlice<i32>,
pos_dev: CudaSlice<i32>,
argmax: CudaSlice<i32>,
bounce: CudaSlice<f32>,
slot_rows: CudaSlice<i32>,
tap_tmp: CudaSlice<f32>,
}
struct CmpCkptDev {
kv_snap: CudaSlice<f32>,
sc_snap: CudaSlice<f32>,
rows_kv: CudaSlice<f32>,
rows_sc: CudaSlice<f32>,
latent: usize,
ratio: usize,
overlap: bool,
n_blocks0: usize,
}
struct LayerCkptDev {
cmp: Option<CmpCkptDev>,
idx: Option<CmpCkptDev>,
trans_base: usize,
}
pub struct VerifyState {
ws: Vec<VerifyWs>,
layers: Vec<LayerCkptDev>,
pub tmax: usize,
open: Option<(usize, usize)>,
pub bytes: Vec<u64>,
}
impl Dsv4Gpu {
pub fn verify_tmax(&self) -> usize {
self.dspark.as_ref().map(|d| d.block_size + 1).unwrap_or(0)
}
pub fn alloc_verify_state(&self) -> Res<VerifyState> {
let tmax = self.verify_tmax();
if tmax == 0 {
return Err("alloc_verify_state needs MEMRA_DSV4_DRAFTER=dspark".into());
}
if !matches!(self.decode_path, DecodePath::Device { .. }) {
return Err(
"batched verify is a device-path rung (MEMRA_DSV4_DECODE_PATH=device)".into(),
);
}
let d = self.model.cfg();
let mc = &self.model.mc;
let moe = mc.moe.as_ref().expect("moe");
let hc = d.hc_mult as usize;
let hidden = mc.n_embd as usize;
let heads = mc.n_head as usize;
let hd = d.head_dim as usize;
let q_lora = d.q_lora_rank as usize;
let win = d.sliding_window as usize;
let o_groups = d.o_groups as usize;
let o_lora = d.o_lora_rank as usize;
let iheads = d.index_n_heads as usize;
let ihd = d.index_head_dim as usize;
let topk = moe.expert_used_count as usize;
let ne = moe.expert_count as usize;
let inter = moe.expert_ff_length as usize;
let itopk = d.index_topk as usize;
let n_trunk = (mc.n_layer - mc.nextn_predict_layers) as usize;
let vocab = {
let (info, _) = self.model.st.raw("head.weight").expect("head");
info.shape[0] as usize
};
let sh_inter = {
let (info, _) = self
.model
.st
.raw("layers.0.ffn.shared_experts.w1.weight")
.expect("shared w1");
info.shape[0] as usize
};
let mut max_d = 0usize;
let mut max_shift = 0usize;
let mut min_ratio = usize::MAX;
for st in &self.stages {
for l in &st.layers {
for cmp in l.cmp.iter().chain(l.idx.as_ref().map(|ix| &ix.cmp)) {
max_d = max_d.max(cmp.d);
if cmp.overlap {
max_shift = max_shift.max(cmp.ratio * cmp.latent);
}
min_ratio = min_ratio.min(cmp.ratio);
}
}
}
assert!(min_ratio != usize::MAX, "no compressor layers?");
let score_cap = self.max_seq / min_ratio + 1;
let idx_tail = itopk.max(self.max_seq / 128 + 1);
let idx_stride = win + idx_tail;
let max_gemm_k = (o_groups * o_lora).max(hidden).max(q_lora).max(sh_inter);
let mut bytes = vec![0u64; self.stages.len()];
let mut ws = Vec::with_capacity(self.stages.len());
for st in &self.stages {
st.gpu.ctx.bind_to_thread().map_err(e("bind ctx vws"))?;
let s = st.gpu.stream();
let acc = std::cell::Cell::new(0u64);
let f = |n: usize| {
acc.set(acc.get() + (n * 4) as u64);
s.alloc_zeros::<f32>(n).map_err(e("vws f32"))
};
let b = |n: usize| {
acc.set(acc.get() + n as u64);
s.alloc_zeros::<u8>(n).map_err(e("vws u8"))
};
let i = |n: usize| {
acc.set(acc.get() + (n * 4) as u64);
s.alloc_zeros::<i32>(n).map_err(e("vws i32"))
};
let w = VerifyWs {
tmax,
h_a: f(tmax * hc * hidden)?,
h_b: f(tmax * hc * hidden)?,
h_rx: f(tmax * hc * hidden)?,
emb: f(tmax * hidden)?,
mixes: f(tmax * (2 + hc) * hc)?,
pre: f(tmax * hc)?,
post: f(tmax * hc)?,
comb: f(tmax * hc * hc)?,
y_hc: f(tmax * hidden)?,
x: f(tmax * hidden)?,
xf: f(tmax * hidden)?,
qr: f(tmax * q_lora)?,
qr_b: b(tmax * q_lora * 2)?,
q: f(tmax * heads * hd)?,
kv: f(tmax * hd)?,
qi: f(tmax * iheads * ihd)?,
wproj: f(tmax * iheads)?,
score: f(score_cap)?,
idx: i(tmax * idx_stride)?,
idx_stride,
o: f(tmax * heads * hd)?,
o_b: b(tmax * heads * hd * 2)?,
og: f(tmax * o_groups * o_lora)?,
attn_out: f(tmax * hidden)?,
gemm_xb: b(tmax * max_gemm_k * 2)?,
raw: f(tmax * ne)?,
sel: i(tmax * topk)?,
selw: f(tmax * topk)?,
order: i(tmax * topk)?,
xq: b(tmax * hidden)?,
xs: f(tmax * hidden / 128)?,
g1: f(tmax * topk * inter)?,
g3: f(tmax * topk * inter)?,
hbuf: f(tmax * topk * inter)?,
hq: b(tmax * topk * inter)?,
hs: f(tmax * topk * inter / 128)?,
contrib: f(tmax * topk * hidden)?,
y: f(tmax * hidden)?,
xb: b(tmax * hidden * 2)?,
sg1: f(tmax * sh_inter)?,
sg3: f(tmax * sh_inter)?,
shbuf: f(tmax * sh_inter)?,
shb16: b(tmax * sh_inter * 2)?,
sh_out: f(tmax * hidden)?,
cmp_emit: f(2 * max_d)?,
cmp_shift: f(max_shift.max(1))?,
sink_scores: f(tmax * heads * idx_stride)?,
sink_evals: f(tmax * heads * idx_stride)?,
sink_den: {
acc.set(acc.get() + (tmax * heads * 8) as u64);
s.alloc_zeros::<f64>(tmax * heads).map_err(e("vws f64"))?
},
head_mixes: f(tmax * hc)?,
head_pre: f(tmax * hc)?,
collapsed: f(tmax * hidden)?,
logits: f(tmax * vocab)?,
tok: i(tmax)?,
pos_dev: i(tmax)?,
argmax: i(tmax)?,
bounce: f(tmax * hd)?,
slot_rows: i(tmax)?,
tap_tmp: f(tmax * hidden)?,
};
bytes[st.dev] += acc.get();
ws.push(w);
}
let mut layers = Vec::with_capacity(n_trunk);
for il in 0..n_trunk {
let stage_i = self.layer_stage[il];
let st = &self.stages[stage_i];
st.gpu.ctx.bind_to_thread().map_err(e("bind ctx ckpt"))?;
let stream = st.gpu.stream();
let lidx = st
.layers
.iter()
.position(|l| l.il == il as u32)
.unwrap_or_else(|| panic!("layer {il} not on stage {stage_i}"));
let layer = &st.layers[lidx];
let cap_blocks = self.max_seq.checked_div(layer.ratio).unwrap_or(0);
let mk = |cmp: &CmpDev| -> Res<CmpCkptDev> {
let slots = if cmp.overlap {
2 * cmp.ratio
} else {
cmp.ratio
};
Ok(CmpCkptDev {
kv_snap: stream
.alloc_zeros::<f32>(slots * cmp.latent)
.map_err(e("ckpt kv snap"))?,
sc_snap: stream
.alloc_zeros::<f32>(slots * cmp.latent)
.map_err(e("ckpt sc snap"))?,
rows_kv: stream
.alloc_zeros::<f32>(tmax * cmp.latent)
.map_err(e("ckpt rows kv"))?,
rows_sc: stream
.alloc_zeros::<f32>(tmax * cmp.latent)
.map_err(e("ckpt rows sc"))?,
latent: cmp.latent,
ratio: cmp.ratio,
overlap: cmp.overlap,
n_blocks0: 0,
})
};
let cmp = match &layer.cmp {
Some(c) => Some(mk(c)?),
None => None,
};
let idxc = match &layer.idx {
Some(ix) => Some(mk(&ix.cmp)?),
None => None,
};
for c in cmp.iter().chain(idxc.iter()) {
let slots = if c.overlap { 2 * c.ratio } else { c.ratio };
bytes[st.dev] += ((2 * slots * c.latent + 2 * tmax * c.latent) * 4) as u64;
}
layers.push(LayerCkptDev {
cmp,
idx: idxc,
trans_base: d.sliding_window as usize + cap_blocks,
});
}
for st in &self.stages {
st.gpu.stream().synchronize().map_err(e("vws sync"))?;
}
Ok(VerifyState {
ws,
layers,
tmax,
open: None,
bytes,
})
}
#[allow(clippy::too_many_arguments)]
fn gemv_m_dev(
st: &Stage,
w: DW,
x_ptr: *const c_void,
y_ptr: *mut f32,
m: usize,
n: usize,
kdim: usize,
xstride: usize,
ystride: usize,
) -> Res<()> {
let stream = st.gpu.stream();
unsafe {
match w {
DW::Bf16(w_ptr) => ck(
"gemv_bf16_m dev",
k::memra_dsv4_gemv_bf16_m(
w_ptr,
x_ptr,
y_ptr,
m as i32,
n as i32,
kdim as i32,
xstride as i32,
ystride as i32,
sp(&stream),
),
),
DW::Fp8 {
codes,
scales,
sc_cols,
} => ck(
"gemv_fp8_m dev",
k::memra_dsv4_gemv_fp8_m(
codes,
scales,
sc_cols,
x_ptr,
y_ptr,
m as i32,
n as i32,
kdim as i32,
xstride as i32,
ystride as i32,
sp(&stream),
),
),
}
}
}
#[allow(clippy::too_many_arguments)]
fn gemm_m_dev(
st: &Stage,
x_f32: *const f32,
xb: &mut CudaSlice<u8>,
w: DW,
m: usize,
n: usize,
kdim: usize,
y_ptr: *mut f32,
) -> Res<()> {
let stream = st.gpu.stream();
unsafe {
ck(
"cvt_bf16 m dev",
k::memra_dsv4_cvt_bf16(
x_f32,
xb.device_ptr_mut(&stream).0 as *mut c_void,
(m * kdim) as i64,
sp(&stream),
),
)?;
}
Self::gemv_m_dev(
st,
w,
xb.device_ptr(&stream).0 as *const c_void,
y_ptr,
m,
n,
kdim,
0,
0,
)
}
#[allow(clippy::too_many_arguments)]
fn dots_m_dev(
&self,
st: &Stage,
x: *const f32,
w_f32: *const c_void,
w_is_bf16: i32,
s: usize,
kdim: usize,
n: usize,
y: *mut f32,
) -> Res<()> {
let stream = st.gpu.stream();
unsafe {
if self.dots_f32 {
ck(
"dots_f32acc_mrow",
k::memra_dsv4_dots_f32acc_mrow(
x,
w_f32,
w_is_bf16,
y,
s as i32,
kdim as i32,
n as i32,
sp(&stream),
),
)
} else {
ck(
"dots_f32_mrow",
k::memra_dsv4_dots_f32_mrow(
x,
w_f32,
w_is_bf16,
y,
s as i32,
kdim as i32,
n as i32,
sp(&stream),
),
)
}
}
}
}
impl Dsv4Gpu {
#[allow(clippy::too_many_arguments)]
fn hc_pre_batch_dev(
&self,
st: &Stage,
h_ptr: *const f32,
fn_w: &CudaSlice<f32>,
base_host: &[f32],
scale_host: &[f32],
base_dev: &CudaSlice<f32>,
scale_dev: &CudaSlice<f32>,
vws: &mut VerifyWs,
t: usize,
hc: usize,
hidden: usize,
iters: u32,
hc_eps: f32,
host_math: bool,
) -> Res<()> {
let stream = st.gpu.stream();
let w = hc * hidden;
let rows = (2 + hc) * hc;
self.dots_m_dev(
st,
h_ptr,
fn_w.device_ptr(&stream).0 as *const c_void,
0,
t,
w,
rows,
vws.mixes.device_ptr_mut(&stream).0 as *mut f32,
)?;
unsafe {
ck(
"rowsq_scale batch",
self.rowsq_scale_arm(
h_ptr,
dpm!(vws.mixes, &stream),
t as i32,
w as i32,
rows as i32,
hc_eps,
sp(&stream),
),
)?;
}
if host_math {
let mut mixes_h = vec![0f32; t * rows];
let view = vws.mixes.slice(0..t * rows);
stream
.memcpy_dtoh(&view, &mut mixes_h[..])
.map_err(e("dtoh mixes batch"))?;
stream.synchronize().map_err(e("sync mixes batch"))?;
let (pre_h, post_h, comb_h) =
hc_split_sinkhorn(&mixes_h, t, hc, scale_host, base_host, iters, hc_eps);
let mut dp = vws.pre.slice_mut(0..t * hc);
stream
.memcpy_htod(&pre_h, &mut dp)
.map_err(e("htod pre b"))?;
let mut dp = vws.post.slice_mut(0..t * hc);
stream
.memcpy_htod(&post_h, &mut dp)
.map_err(e("htod post b"))?;
let mut dp = vws.comb.slice_mut(0..t * hc * hc);
stream
.memcpy_htod(&comb_h, &mut dp)
.map_err(e("htod comb b"))?;
} else {
unsafe {
ck(
"hc_sinkhorn_m",
k::memra_dsv4_hc_sinkhorn_m(
dpf!(vws.mixes, &stream),
dpf!(scale_dev, &stream),
dpf!(base_dev, &stream),
dpm!(vws.pre, &stream),
dpm!(vws.post, &stream),
dpm!(vws.comb, &stream),
t as i32,
hc as i32,
iters as i32,
hc_eps,
sp(&stream),
),
)?;
}
}
unsafe {
ck(
"hc_collapse batch",
k::memra_dsv4_hc_collapse(
h_ptr,
dpf!(vws.pre, &stream),
dpm!(vws.y_hc, &stream),
t as i32,
hc as i32,
hidden as i32,
sp(&stream),
),
)?;
}
Ok(())
}
#[allow(clippy::too_many_arguments)]
fn cmp_decode_batch_dev(
&self,
st: &Stage,
cmp: &CmpDev,
x_ptr: *const f32,
t: usize,
pos0: usize,
hidden: usize,
fc_dev: &CudaSlice<f32>,
rd: usize,
eps: f32,
ck_dev: &mut CmpCkptDev,
emit: &mut CudaSlice<f32>,
shift: &mut CudaSlice<f32>,
pend_kv: &mut CudaSlice<f32>,
pend_score: &mut CudaSlice<f32>,
store: &mut CudaSlice<f32>,
row0: usize,
blocks: &mut usize,
) -> Res<()> {
let stream = st.gpu.stream();
let (ratio, d, latent) = (cmp.ratio, cmp.d, cmp.latent);
stream
.memcpy_dtod(pend_kv, &mut ck_dev.kv_snap)
.map_err(e("ckpt snap kv"))?;
stream
.memcpy_dtod(pend_score, &mut ck_dev.sc_snap)
.map_err(e("ckpt snap sc"))?;
ck_dev.n_blocks0 = *blocks;
self.dots_m_dev(
st,
x_ptr,
cmp.wkv.device_ptr(&stream).0 as *const c_void,
0,
t,
hidden,
latent,
ck_dev.rows_kv.device_ptr_mut(&stream).0 as *mut f32,
)?;
self.dots_m_dev(
st,
x_ptr,
cmp.wgate.device_ptr(&stream).0 as *const c_void,
0,
t,
hidden,
latent,
ck_dev.rows_sc.device_ptr_mut(&stream).0 as *mut f32,
)?;
for i in 0..t {
let pos = pos0 + i;
let slot = if cmp.overlap {
ratio + pos % ratio
} else {
pos % ratio
};
{
let src = ck_dev.rows_kv.slice(i * latent..(i + 1) * latent);
let mut dst = pend_kv.slice_mut(slot * latent..(slot + 1) * latent);
stream.memcpy_dtod(&src, &mut dst).map_err(e("pend kv b"))?;
let src = ck_dev.rows_sc.slice(i * latent..(i + 1) * latent);
let mut dst = pend_score.slice_mut(slot * latent..(slot + 1) * latent);
stream.memcpy_dtod(&src, &mut dst).map_err(e("pend sc b"))?;
}
if (pos + 1) % ratio != 0 {
continue;
}
let j = pos / ratio;
let nb_launch = if cmp.overlap { 2usize } else { 1 };
let row_off = if cmp.overlap { d } else { 0 };
unsafe {
ck(
"compressor_pool batch",
k::memra_dsv4_compressor_pool(
dpf!(*pend_kv, &stream),
dpf!(*pend_score, &stream),
dpf!(cmp.ape, &stream),
dpm!(*emit, &stream),
nb_launch as i32,
ratio as i32,
d as i32,
latent as i32,
cmp.overlap as i32,
sp(&stream),
),
)?;
let row_c = (emit.device_ptr(&stream).0 as usize + row_off * 4) as *const f32;
let row_m = (emit.device_ptr_mut(&stream).0 as usize + row_off * 4) as *mut f32;
ck(
"rmsnorm batch cmp",
self.rmsnorm_arm(
row_c,
dpf!(cmp.norm, &stream),
row_m,
1,
d as i32,
eps,
sp(&stream),
),
)?;
ck(
"rope_at batch cmp",
k::memra_dsv4_rope_at(
row_m,
1,
d as i32,
rd as i32,
dpf!(fc_dev, &stream),
(j * ratio) as i32,
0,
sp(&stream),
),
)?;
if cmp.rotate {
let scale = (d as f32).powf(-0.5);
ck(
"hadamard batch cmp",
k::memra_dsv4_hadamard(row_m, 1, d as i32, scale, sp(&stream)),
)?;
ck(
"fp4 batch cmp",
k::memra_dsv4_fp4_act_quant(row_m, 1, d as i64, d as i32, sp(&stream)),
)?;
} else {
ck(
"act_quant batch cmp",
k::memra_dsv4_act_quant(
row_m,
1,
d as i64,
(d - rd) as i32,
64,
(self.variant == ActQuantVariant::ClampOnly) as i32,
sp(&stream),
),
)?;
}
}
{
let src = emit.slice(row_off..row_off + d);
let mut dst = store.slice_mut((row0 + j) * d..(row0 + j + 1) * d);
stream
.memcpy_dtod(&src, &mut dst)
.map_err(e("emit store b"))?;
}
if cmp.overlap {
{
let src = pend_kv.slice(ratio * latent..2 * ratio * latent);
let mut dst = shift.slice_mut(0..ratio * latent);
stream.memcpy_dtod(&src, &mut dst).map_err(e("bshift1"))?;
}
{
let src = shift.slice(0..ratio * latent);
let mut dst = pend_kv.slice_mut(0..ratio * latent);
stream.memcpy_dtod(&src, &mut dst).map_err(e("bshift2"))?;
}
{
let src = pend_score.slice(ratio * latent..2 * ratio * latent);
let mut dst = shift.slice_mut(0..ratio * latent);
stream.memcpy_dtod(&src, &mut dst).map_err(e("bshift3"))?;
}
{
let src = shift.slice(0..ratio * latent);
let mut dst = pend_score.slice_mut(0..ratio * latent);
stream.memcpy_dtod(&src, &mut dst).map_err(e("bshift4"))?;
}
}
*blocks = j + 1;
}
Ok(())
}
#[allow(clippy::too_many_arguments)]
fn cmp_rollback_replay_dev(
&self,
st: &Stage,
ck_dev: &CmpCkptDev,
n_commit: usize,
t: usize,
pos0: usize,
shift: &mut CudaSlice<f32>,
pend_kv: &mut CudaSlice<f32>,
pend_score: &mut CudaSlice<f32>,
blocks: &mut usize,
) -> Res<()> {
if n_commit == t {
return Ok(()); }
let stream = st.gpu.stream();
let (ratio, latent, overlap) = (ck_dev.ratio, ck_dev.latent, ck_dev.overlap);
stream
.memcpy_dtod(&ck_dev.kv_snap, pend_kv)
.map_err(e("rb kv snap"))?;
stream
.memcpy_dtod(&ck_dev.sc_snap, pend_score)
.map_err(e("rb sc snap"))?;
*blocks = ck_dev.n_blocks0;
for i in 0..n_commit {
let pos = pos0 + i;
let slot = if overlap {
ratio + pos % ratio
} else {
pos % ratio
};
{
let src = ck_dev.rows_kv.slice(i * latent..(i + 1) * latent);
let mut dst = pend_kv.slice_mut(slot * latent..(slot + 1) * latent);
stream.memcpy_dtod(&src, &mut dst).map_err(e("rb row kv"))?;
let src = ck_dev.rows_sc.slice(i * latent..(i + 1) * latent);
let mut dst = pend_score.slice_mut(slot * latent..(slot + 1) * latent);
stream.memcpy_dtod(&src, &mut dst).map_err(e("rb row sc"))?;
}
if (pos + 1) % ratio != 0 {
continue;
}
if overlap {
{
let src = pend_kv.slice(ratio * latent..2 * ratio * latent);
let mut dst = shift.slice_mut(0..ratio * latent);
stream.memcpy_dtod(&src, &mut dst).map_err(e("rbshift1"))?;
}
{
let src = shift.slice(0..ratio * latent);
let mut dst = pend_kv.slice_mut(0..ratio * latent);
stream.memcpy_dtod(&src, &mut dst).map_err(e("rbshift2"))?;
}
{
let src = pend_score.slice(ratio * latent..2 * ratio * latent);
let mut dst = shift.slice_mut(0..ratio * latent);
stream.memcpy_dtod(&src, &mut dst).map_err(e("rbshift3"))?;
}
{
let src = shift.slice(0..ratio * latent);
let mut dst = pend_score.slice_mut(0..ratio * latent);
stream.memcpy_dtod(&src, &mut dst).map_err(e("rbshift4"))?;
}
}
*blocks += 1;
}
Ok(())
}
}
impl Dsv4Gpu {
#[allow(clippy::too_many_arguments)]
fn block_verify_dev(
&self,
st: &Stage,
layer: &LayerDev,
cache: &mut LayerCache,
lck: &mut LayerCkptDev,
vws: &mut VerifyWs,
input_rx: bool,
pos0: usize,
t: usize,
toks: &[u32],
host_math: bool,
) -> Res<()> {
let d = self.model.cfg();
let mc = &self.model.mc;
let hc = d.hc_mult as usize;
let hidden = mc.n_embd as usize;
let heads = mc.n_head as usize;
let hd = d.head_dim as usize;
let rd = d.qk_rope_head_dim as usize;
let q_lora = d.q_lora_rank as usize;
let win = d.sliding_window as usize;
let o_groups = d.o_groups as usize;
let o_lora = d.o_lora_rank as usize;
let eps = mc.rms_eps;
let iters = d.hc_sinkhorn_iters;
let hc_eps = d.hc_eps;
let stream = st.gpu.stream();
let fc_dev: *const f32 = if layer.ratio != 0 {
st.fc_yarn.device_ptr(&stream).0 as *const f32
} else {
st.fc_plain.device_ptr(&stream).0 as *const f32
};
let clamp_only = (self.variant == ActQuantVariant::ClampOnly) as i32;
let trans_base = lck.trans_base;
let LayerCache {
kvc,
n_blocks,
pend_kv,
pend_score,
ikvc,
i_blocks,
ipend_kv,
ipend_score,
} = cache;
let h_in_ptr: *const f32 = if input_rx {
vws.h_rx.device_ptr(&stream).0 as *const f32
} else {
vws.h_a.device_ptr(&stream).0 as *const f32
};
self.hc_pre_batch_dev(
st,
h_in_ptr,
&layer.hc_attn_fn,
&layer.hc_attn_base,
&layer.hc_attn_scale,
&layer.hc_attn_base_dev,
&layer.hc_attn_scale_dev,
vws,
t,
hc,
hidden,
iters,
hc_eps,
host_math,
)?;
unsafe {
ck(
"rmsnorm attn batch",
self.rmsnorm_arm(
dpf!(vws.y_hc, &stream),
dpf!(layer.attn_norm, &stream),
dpm!(vws.x, &stream),
t as i32,
hidden as i32,
eps,
sp(&stream),
),
)?;
}
Self::gemm_m_dev(
st,
vws.x.device_ptr(&stream).0 as *const f32,
&mut vws.gemm_xb,
dwsel(self.dense_fp8, &stream, &layer.wq_a, &layer.wq_a_fp8),
t,
q_lora,
hidden,
vws.qr.device_ptr_mut(&stream).0 as *mut f32,
)?;
unsafe {
ck(
"rmsnorm q batch",
self.rmsnorm_arm(
dpf!(vws.qr, &stream),
dpf!(layer.q_norm, &stream),
dpm!(vws.qr, &stream),
t as i32,
q_lora as i32,
eps,
sp(&stream),
),
)?;
ck(
"cvt qr batch",
k::memra_dsv4_cvt_bf16(
dpf!(vws.qr, &stream),
vws.qr_b.device_ptr_mut(&stream).0 as *mut c_void,
(t * q_lora) as i64,
sp(&stream),
),
)?;
}
Self::gemv_m_dev(
st,
dwsel(self.dense_fp8, &stream, &layer.wq_b, &layer.wq_b_fp8),
vws.qr_b.device_ptr(&stream).0 as *const c_void,
vws.q.device_ptr_mut(&stream).0 as *mut f32,
t,
heads * hd,
q_lora,
0,
0,
)?;
unsafe {
ck(
"headrms batch",
self.headrms_arm(
dpm!(vws.q, &stream),
(t * heads) as i32,
hd as i32,
eps,
sp(&stream),
),
)?;
ck(
"rope q batch",
k::memra_dsv4_rope(
dpm!(vws.q, &stream),
t as i32,
heads as i32,
hd as i32,
rd as i32,
fc_dev,
vws.pos_dev.device_ptr(&stream).0 as *const i32,
0,
sp(&stream),
),
)?;
}
Self::gemm_m_dev(
st,
vws.x.device_ptr(&stream).0 as *const f32,
&mut vws.gemm_xb,
dwsel(self.dense_fp8, &stream, &layer.wkv, &layer.wkv_fp8),
t,
hd,
hidden,
vws.kv.device_ptr_mut(&stream).0 as *mut f32,
)?;
unsafe {
ck(
"rmsnorm kv batch",
self.rmsnorm_arm(
dpf!(vws.kv, &stream),
dpf!(layer.kv_norm, &stream),
dpm!(vws.kv, &stream),
t as i32,
hd as i32,
eps,
sp(&stream),
),
)?;
ck(
"rope kv batch",
k::memra_dsv4_rope(
dpm!(vws.kv, &stream),
t as i32,
1,
hd as i32,
rd as i32,
fc_dev,
vws.pos_dev.device_ptr(&stream).0 as *const i32,
0,
sp(&stream),
),
)?;
ck(
"act_quant kv batch",
k::memra_dsv4_act_quant(
dpm!(vws.kv, &stream),
t as i32,
hd as i64,
(hd - rd) as i32,
64,
clamp_only,
sp(&stream),
),
)?;
}
{
let src = vws.kv.slice(0..t * hd);
let mut dst = kvc.slice_mut(trans_base * hd..(trans_base + t) * hd);
stream
.memcpy_dtod(&src, &mut dst)
.map_err(e("transient ring write"))?;
}
let mut slots = win;
if layer.ratio != 0 {
let ratio = layer.ratio;
let nbs: Vec<usize> = (0..t).map(|i| (pos0 + i + 1) / ratio).collect();
if let Some(ix) = &layer.idx {
Self::gemv_m_dev(
st,
dwsel(self.dense_fp8, &stream, &ix.wq_b, &ix.wq_b_fp8),
vws.qr_b.device_ptr(&stream).0 as *const c_void,
vws.qi.device_ptr_mut(&stream).0 as *mut f32,
t,
ix.heads * ix.hd,
q_lora,
0,
0,
)?;
unsafe {
ck(
"rope qi batch",
k::memra_dsv4_rope(
dpm!(vws.qi, &stream),
t as i32,
ix.heads as i32,
ix.hd as i32,
rd as i32,
fc_dev,
vws.pos_dev.device_ptr(&stream).0 as *const i32,
0,
sp(&stream),
),
)?;
let scale = (ix.hd as f32).powf(-0.5);
ck(
"hadamard qi batch",
k::memra_dsv4_hadamard(
dpm!(vws.qi, &stream),
(t * ix.heads) as i32,
ix.hd as i32,
scale,
sp(&stream),
),
)?;
ck(
"fp4 qi batch",
k::memra_dsv4_fp4_act_quant(
dpm!(vws.qi, &stream),
(t * ix.heads) as i32,
ix.hd as i64,
ix.hd as i32,
sp(&stream),
),
)?;
}
Self::gemm_m_dev(
st,
vws.x.device_ptr(&stream).0 as *const f32,
&mut vws.gemm_xb,
dwsel(
self.dense_fp8,
&stream,
&ix.weights_proj,
&ix.weights_proj_fp8,
),
t,
ix.heads,
hidden,
vws.wproj.device_ptr_mut(&stream).0 as *mut f32,
)?;
{
let VerifyWs {
x,
cmp_emit,
cmp_shift,
..
} = vws;
self.cmp_decode_batch_dev(
st,
&ix.cmp,
x.device_ptr(&stream).0 as *const f32,
t,
pos0,
hidden,
&st.fc_yarn,
rd,
eps,
lck.idx.as_mut().expect("idx ckpt"),
cmp_emit,
cmp_shift,
ipend_kv.as_mut().expect("ipend"),
ipend_score.as_mut().expect("ipend"),
ikvc.as_mut().expect("ikvc"),
0,
i_blocks,
)?;
}
debug_assert_eq!(*i_blocks, nbs[t - 1], "indexer block count (batch)");
let kks: Vec<usize> = nbs.iter().map(|&nb| ix.topk.min(nb)).collect();
let tail_max = kks.iter().cloned().max().unwrap_or(0);
slots = win + tail_max;
for i in 0..t {
let pos = pos0 + i;
let idx_off = i * vws.idx_stride;
unsafe {
ck(
"build_idx_redirect fine",
k::memra_dsv4_build_idx_redirect(
(vws.idx.device_ptr_mut(&stream).0 as usize + idx_off * 4)
as *mut i32,
pos as i32,
win as i32,
0, slots as i32,
pos0 as i32,
trans_base as i32,
sp(&stream),
),
)?;
}
let nb = nbs[i];
if nb == 0 {
continue;
}
let wscale = ((ix.hd as f64).powf(-0.5) * (ix.heads as f64).powf(-0.5)) as f32;
unsafe {
ck(
"indexer_score batch",
self.indexer_score_arm(
(vws.qi.device_ptr(&stream).0 as usize + i * ix.heads * ix.hd * 4)
as *const f32,
dpf!(ikvc.as_ref().expect("ikvc"), &stream),
(vws.wproj.device_ptr(&stream).0 as usize + i * ix.heads * 4)
as *const f32,
wscale,
dpm!(vws.score, &stream),
1,
ix.heads as i32,
ix.hd as i32,
nb as i32,
ratio as i32,
nb as i32,
sp(&stream),
),
)?;
}
let kk = kks[i];
if host_math {
let score_h = {
let view = vws.score.slice(0..nb);
let mut v = vec![0f32; nb];
stream
.memcpy_dtoh(&view, &mut v[..])
.map_err(e("dtoh sc b"))?;
stream.synchronize().map_err(e("sync sc b"))?;
v
};
let mut order: Vec<usize> = (0..nb).collect();
order.sort_by(|&a, &b| {
score_h[b]
.partial_cmp(&score_h[a])
.unwrap_or(std::cmp::Ordering::Equal)
.then(a.cmp(&b))
});
let cidx: Vec<i32> = order
.into_iter()
.take(kk)
.map(|j| (j + win) as i32)
.collect();
let mut dst = vws.idx.slice_mut(idx_off + win..idx_off + win + kk);
stream
.memcpy_htod(&cidx, &mut dst)
.map_err(e("htod idx b"))?;
} else {
unsafe {
let idx_tail_ptr = (vws.idx.device_ptr_mut(&stream).0 as usize
+ (idx_off + win) * 4)
as *mut i32;
ck(
"topk_idx batch",
k::memra_dsv4_topk_idx(
dpf!(vws.score, &stream),
nb as i32,
kk as i32,
win as i32,
idx_tail_ptr,
sp(&stream),
),
)?;
}
}
}
} else {
let tail_max = nbs.iter().cloned().max().unwrap_or(0);
slots = win + tail_max;
for (i, &nb_i) in nbs.iter().enumerate() {
let pos = pos0 + i;
let idx_off = i * vws.idx_stride;
unsafe {
ck(
"build_idx_redirect coarse",
k::memra_dsv4_build_idx_redirect(
(vws.idx.device_ptr_mut(&stream).0 as usize + idx_off * 4)
as *mut i32,
pos as i32,
win as i32,
nb_i as i32,
slots as i32,
pos0 as i32,
trans_base as i32,
sp(&stream),
),
)?;
}
}
}
{
let VerifyWs {
x,
cmp_emit,
cmp_shift,
..
} = vws;
self.cmp_decode_batch_dev(
st,
layer.cmp.as_ref().expect("ratio!=0 has compressor"),
x.device_ptr(&stream).0 as *const f32,
t,
pos0,
hidden,
&st.fc_yarn,
rd,
eps,
lck.cmp.as_mut().expect("cmp ckpt"),
cmp_emit,
cmp_shift,
pend_kv.as_mut().expect("pend"),
pend_score.as_mut().expect("pend"),
kvc,
win,
n_blocks,
)?;
}
debug_assert_eq!(*n_blocks, nbs[t - 1], "attn block count (batch)");
} else {
for i in 0..t {
let pos = pos0 + i;
let idx_off = i * vws.idx_stride;
unsafe {
ck(
"build_idx_redirect window-only",
k::memra_dsv4_build_idx_redirect(
(vws.idx.device_ptr_mut(&stream).0 as usize + idx_off * 4) as *mut i32,
pos as i32,
win as i32,
-1,
win as i32,
pos0 as i32,
trans_base as i32,
sp(&stream),
),
)?;
}
}
}
let scale = (hd as f64).powf(-0.5) as f32;
unsafe {
if self.chains_f32 {
ck(
"sink_attn_dec_mq_f32acc",
k::memra_dsv4_sink_attn_dec_mq_f32acc(
dpf!(vws.q, &stream),
dpf!(kvc, &stream),
vws.idx.device_ptr(&stream).0 as *const i32,
dpf!(layer.sink, &stream),
dpm!(vws.sink_scores, &stream),
dpm!(vws.sink_evals, &stream),
vws.sink_den.device_ptr_mut(&stream).0 as *mut f32,
dpm!(vws.o, &stream),
t as i32,
heads as i32,
hd as i32,
slots as i32,
vws.idx_stride as i32,
scale,
sp(&stream),
),
)?;
} else {
ck(
"sink_attn_dec_mq",
k::memra_dsv4_sink_attn_dec_mq(
dpf!(vws.q, &stream),
dpf!(kvc, &stream),
vws.idx.device_ptr(&stream).0 as *const i32,
dpf!(layer.sink, &stream),
dpm!(vws.sink_scores, &stream),
dpm!(vws.sink_evals, &stream),
vws.sink_den.device_ptr_mut(&stream).0 as *mut f64,
dpm!(vws.o, &stream),
t as i32,
heads as i32,
hd as i32,
slots as i32,
vws.idx_stride as i32,
scale,
sp(&stream),
),
)?;
}
ck(
"rope o inv batch",
k::memra_dsv4_rope(
dpm!(vws.o, &stream),
t as i32,
heads as i32,
hd as i32,
rd as i32,
fc_dev,
vws.pos_dev.device_ptr(&stream).0 as *const i32,
1,
sp(&stream),
),
)?;
}
let gw = heads / o_groups * hd;
unsafe {
ck(
"cvt o batch",
k::memra_dsv4_cvt_bf16(
dpf!(vws.o, &stream),
vws.o_b.device_ptr_mut(&stream).0 as *mut c_void,
(t * heads * hd) as i64,
sp(&stream),
),
)?;
}
let wo_a_dw = dwsel(self.dense_fp8, &stream, &layer.wo_a, &layer.wo_a_fp8);
for g in 0..o_groups {
Self::gemv_m_dev(
st,
wo_a_dw.offset_rows(g * o_lora, gw),
(vws.o_b.device_ptr(&stream).0 as usize + g * gw * 2) as *const c_void,
(vws.og.device_ptr_mut(&stream).0 as usize + g * o_lora * 4) as *mut f32,
t,
o_lora,
gw,
heads * hd,
o_groups * o_lora,
)?;
}
Self::gemm_m_dev(
st,
vws.og.device_ptr(&stream).0 as *const f32,
&mut vws.gemm_xb,
dwsel(self.dense_fp8, &stream, &layer.wo_b, &layer.wo_b_fp8),
t,
hidden,
o_groups * o_lora,
vws.attn_out.device_ptr_mut(&stream).0 as *mut f32,
)?;
unsafe {
ck(
"hc_post attn batch",
k::memra_dsv4_hc_post(
dpf!(vws.attn_out, &stream),
h_in_ptr,
dpf!(vws.post, &stream),
dpf!(vws.comb, &stream),
dpm!(vws.h_b, &stream),
t as i32,
hc as i32,
hidden as i32,
sp(&stream),
),
)?;
}
let h_b_ptr = vws.h_b.device_ptr(&stream).0 as *const f32;
self.hc_pre_batch_dev(
st,
h_b_ptr,
&layer.hc_ffn_fn,
&layer.hc_ffn_base,
&layer.hc_ffn_scale,
&layer.hc_ffn_base_dev,
&layer.hc_ffn_scale_dev,
vws,
t,
hc,
hidden,
iters,
hc_eps,
host_math,
)?;
unsafe {
ck(
"rmsnorm ffn batch",
self.rmsnorm_arm(
dpf!(vws.y_hc, &stream),
dpf!(layer.ffn_norm, &stream),
dpm!(vws.xf, &stream),
t as i32,
hidden as i32,
eps,
sp(&stream),
),
)?;
}
self.moe_verify_dev(st, layer, vws, t, toks, host_math)?;
unsafe {
ck(
"hc_post ffn batch",
k::memra_dsv4_hc_post(
dpf!(vws.y, &stream),
dpf!(vws.h_b, &stream),
dpf!(vws.post, &stream),
dpf!(vws.comb, &stream),
dpm!(vws.h_a, &stream),
t as i32,
hc as i32,
hidden as i32,
sp(&stream),
),
)?;
}
Ok(())
}
fn moe_verify_dev(
&self,
st: &Stage,
layer: &LayerDev,
vws: &mut VerifyWs,
t: usize,
toks: &[u32],
host_math: bool,
) -> Res<()> {
let mc = &self.model.mc;
let d = self.model.cfg();
let moe = mc.moe.as_ref().expect("moe");
let hidden = mc.n_embd as usize;
let ne = moe.expert_count as usize;
let topk = moe.expert_used_count as usize;
let inter = moe.expert_ff_length as usize;
let limit = d.swiglu_limit;
let stream = st.gpu.stream();
let kind = match layer.expert_kind {
ExpertKind::Nvfp4 => 0i32,
ExpertKind::Mxfp4 => 1i32,
};
let wstride = (inter * hidden / 2) as i64;
let sstride = match layer.expert_kind {
ExpertKind::Nvfp4 => (inter * hidden / 16) as i64,
ExpertKind::Mxfp4 => (inter * hidden / 32) as i64,
};
let slots = t * topk;
self.dots_m_dev(
st,
vws.xf.device_ptr(&stream).0 as *const f32,
layer.gate_w.device_ptr(&stream).0 as *const c_void,
0,
t,
hidden,
ne,
vws.raw.device_ptr_mut(&stream).0 as *mut f32,
)?;
if host_math {
let raw_h = {
let view = vws.raw.slice(0..t * ne);
let mut v = vec![0f32; t * ne];
stream
.memcpy_dtoh(&view, &mut v[..])
.map_err(e("dtoh raw b"))?;
stream.synchronize().map_err(e("sync raw b"))?;
v
};
let (indices, weights) =
Self::route_host(layer, &raw_h, toks, t, ne, topk, d.routed_scaling_factor);
let sel: Vec<i32> = indices.iter().map(|&x| x as i32).collect();
let mut order = vec![0i32; t * topk];
for p in 0..t {
let mut o: Vec<i32> = (0..topk as i32).collect();
o.sort_by_key(|&s| indices[p * topk + s as usize]);
order[p * topk..(p + 1) * topk].copy_from_slice(&o);
}
let mut dst = vws.sel.slice_mut(0..t * topk);
stream
.memcpy_htod(&sel, &mut dst)
.map_err(e("htod sel b"))?;
let mut dst = vws.selw.slice_mut(0..t * topk);
stream
.memcpy_htod(&weights, &mut dst)
.map_err(e("htod selw b"))?;
let mut dst = vws.order.slice_mut(0..t * topk);
stream
.memcpy_htod(&order, &mut dst)
.map_err(e("htod order b"))?;
} else {
unsafe {
ck(
"route_m",
k::memra_dsv4_route_m(
dpf!(vws.raw, &stream),
layer
.gate_bias_dev
.as_ref()
.map(|b| b.device_ptr(&stream).0 as *const f32)
.unwrap_or(std::ptr::null()),
layer
.tid2eid_dev
.as_ref()
.map(|x| x.device_ptr(&stream).0 as *const i32)
.unwrap_or(std::ptr::null()),
vws.tok.device_ptr(&stream).0 as *const i32,
t as i32,
ne as i32,
topk as i32,
d.routed_scaling_factor,
vws.sel.device_ptr_mut(&stream).0 as *mut i32,
vws.selw.device_ptr_mut(&stream).0 as *mut f32,
vws.order.device_ptr_mut(&stream).0 as *mut i32,
sp(&stream),
),
)?;
}
}
unsafe {
ck(
"act_quant_fp8 x batch",
k::memra_dsv4_act_quant_fp8(
dpf!(vws.xf, &stream),
vws.xq.device_ptr_mut(&stream).0 as *mut c_void,
dpm!(vws.xs, &stream),
t as i32,
hidden as i32,
sp(&stream),
),
)?;
for (proj, dst) in [(0i32, &mut vws.g1), (2i32, &mut vws.g3)] {
ck(
"fp4_gemm_sel_g w1/w3",
k::memra_dsv4_fp4_gemm_sel_g(
dp!(vws.xq, &stream),
dpf!(vws.xs, &stream),
dp!(layer.experts_w, &stream),
dp!(layer.experts_sc, &stream),
dpf!(layer.experts_s2_dev, &stream),
vws.sel.device_ptr(&stream).0 as *const i32,
proj,
0,
kind,
dpm!(*dst, &stream),
slots as i32,
inter as i32,
hidden as i32,
wstride,
sstride,
topk as i32,
sp(&stream),
),
)?;
}
ck(
"swiglu batch",
k::memra_dsv4_swiglu(
dpf!(vws.g1, &stream),
dpf!(vws.g3, &stream),
dpm!(vws.hbuf, &stream),
slots as i32,
inter as i32,
limit,
vws.selw.device_ptr(&stream).0 as *const f32,
sp(&stream),
),
)?;
ck(
"act_quant_fp8 h batch",
k::memra_dsv4_act_quant_fp8(
dpf!(vws.hbuf, &stream),
vws.hq.device_ptr_mut(&stream).0 as *mut c_void,
dpm!(vws.hs, &stream),
slots as i32,
inter as i32,
sp(&stream),
),
)?;
ck(
"fp4_gemm_sel_g w2",
k::memra_dsv4_fp4_gemm_sel_g(
dp!(vws.hq, &stream),
dpf!(vws.hs, &stream),
dp!(layer.experts_w, &stream),
dp!(layer.experts_sc, &stream),
dpf!(layer.experts_s2_dev, &stream),
vws.sel.device_ptr(&stream).0 as *const i32,
1,
1,
kind,
dpm!(vws.contrib, &stream),
slots as i32,
hidden as i32,
inter as i32,
wstride,
sstride,
0,
sp(&stream),
),
)?;
ck(
"combine_rows_m",
k::memra_dsv4_combine_rows_m(
dpf!(vws.contrib, &stream),
vws.order.device_ptr(&stream).0 as *const i32,
topk as i32,
dpm!(vws.y, &stream),
hidden as i64,
t as i32,
sp(&stream),
),
)?;
ck(
"cvt xb batch",
k::memra_dsv4_cvt_bf16(
dpf!(vws.xf, &stream),
vws.xb.device_ptr_mut(&stream).0 as *mut c_void,
(t * hidden) as i64,
sp(&stream),
),
)?;
}
let sh_inter = vws.sg1.len() / vws.tmax;
Self::gemv_m_dev(
st,
dwsel(
self.dense_fp8,
&stream,
&layer.shared_w[0],
&layer.shared_fp8[0],
),
vws.xb.device_ptr(&stream).0 as *const c_void,
vws.sg1.device_ptr_mut(&stream).0 as *mut f32,
t,
sh_inter,
hidden,
0,
0,
)?;
Self::gemv_m_dev(
st,
dwsel(
self.dense_fp8,
&stream,
&layer.shared_w[2],
&layer.shared_fp8[2],
),
vws.xb.device_ptr(&stream).0 as *const c_void,
vws.sg3.device_ptr_mut(&stream).0 as *mut f32,
t,
sh_inter,
hidden,
0,
0,
)?;
unsafe {
ck(
"swiglu sh batch",
k::memra_dsv4_swiglu(
dpf!(vws.sg1, &stream),
dpf!(vws.sg3, &stream),
dpm!(vws.shbuf, &stream),
t as i32,
sh_inter as i32,
limit,
std::ptr::null(),
sp(&stream),
),
)?;
ck(
"cvt sh batch",
k::memra_dsv4_cvt_bf16(
dpf!(vws.shbuf, &stream),
vws.shb16.device_ptr_mut(&stream).0 as *mut c_void,
(t * sh_inter) as i64,
sp(&stream),
),
)?;
}
Self::gemv_m_dev(
st,
dwsel(
self.dense_fp8,
&stream,
&layer.shared_w[1],
&layer.shared_fp8[1],
),
vws.shb16.device_ptr(&stream).0 as *const c_void,
vws.sh_out.device_ptr_mut(&stream).0 as *mut f32,
t,
hidden,
sh_inter,
0,
0,
)?;
unsafe {
ck(
"add shared batch",
k::memra_dsv4_add_inplace(
dpm!(vws.y, &stream),
dpf!(vws.sh_out, &stream),
(t * hidden) as i64,
sp(&stream),
),
)?;
}
Ok(())
}
fn head_logits_batch_dev(&self, vws: &mut VerifyWs, t: usize, host_math: bool) -> Res<()> {
let d = self.model.cfg();
let mc = &self.model.mc;
let hc = d.hc_mult as usize;
let hidden = mc.n_embd as usize;
let eps = mc.rms_eps;
let last = self.stages.len() - 1;
let st = &self.stages[last];
let stream = st.gpu.stream();
let w = hc * hidden;
let fn_w = st.hc_head_fn.as_ref().expect("hc_head_fn");
let norm = st.trunk_norm.as_ref().expect("trunk norm");
let vocab = vws.logits.len() / vws.tmax;
self.dots_m_dev(
st,
vws.h_a.device_ptr(&stream).0 as *const f32,
fn_w.device_ptr(&stream).0 as *const c_void,
0,
t,
w,
hc,
vws.head_mixes.device_ptr_mut(&stream).0 as *mut f32,
)?;
unsafe {
ck(
"rowsq head batch",
self.rowsq_scale_arm(
dpf!(vws.h_a, &stream),
dpm!(vws.head_mixes, &stream),
t as i32,
w as i32,
hc as i32,
eps,
sp(&stream),
),
)?;
}
if host_math {
let mut mixes_h = vec![0f32; t * hc];
let view = vws.head_mixes.slice(0..t * hc);
stream
.memcpy_dtoh(&view, &mut mixes_h[..])
.map_err(e("dtoh head mixes b"))?;
stream.synchronize().map_err(e("sync head mixes b"))?;
for p in 0..t {
for c in 0..hc {
let m = mixes_h[p * hc + c];
mixes_h[p * hc + c] =
sigmoid_f32(m * self.hc_head_scale[0] + self.hc_head_base[c]) + d.hc_eps;
}
}
let mut dst = vws.head_pre.slice_mut(0..t * hc);
stream
.memcpy_htod(&mixes_h, &mut dst)
.map_err(e("htod head pre b"))?;
} else {
unsafe {
ck(
"hc_head_pre_m",
k::memra_dsv4_hc_head_pre_m(
dpf!(vws.head_mixes, &stream),
st.hc_head_scale_dev
.as_ref()
.expect("head scale dev")
.device_ptr(&stream)
.0 as *const f32,
st.hc_head_base_dev
.as_ref()
.expect("head base dev")
.device_ptr(&stream)
.0 as *const f32,
dpm!(vws.head_pre, &stream),
t as i32,
hc as i32,
d.hc_eps,
sp(&stream),
),
)?;
}
}
unsafe {
ck(
"hc_collapse head batch",
k::memra_dsv4_hc_collapse(
dpf!(vws.h_a, &stream),
dpf!(vws.head_pre, &stream),
dpm!(vws.collapsed, &stream),
t as i32,
hc as i32,
hidden as i32,
sp(&stream),
),
)?;
ck(
"rmsnorm head batch",
self.rmsnorm_arm(
dpf!(vws.collapsed, &stream),
dpf!(norm, &stream),
dpm!(vws.collapsed, &stream),
t as i32,
hidden as i32,
eps,
sp(&stream),
),
)?;
}
let head_ptr = st.head.as_ref().expect("head").device_ptr(&stream).0 as *const c_void;
self.dots_m_dev(
st,
vws.collapsed.device_ptr(&stream).0 as *const f32,
head_ptr,
1,
t,
hidden,
vocab,
vws.logits.device_ptr_mut(&stream).0 as *mut f32,
)?;
Ok(())
}
}
pub struct SpecRoundGpu {
pub start_pos: usize,
pub drafts: Vec<u32>,
pub accepts: usize,
pub verified: usize,
pub t_batch: usize,
pub t_cap: usize,
pub confidence: Vec<f32>,
pub emitted: usize,
pub round_us: u64,
}
pub struct SpecRunGpu {
pub tokens: Vec<u32>,
pub rounds: Vec<SpecRoundGpu>,
}
impl Dsv4Gpu {
pub fn verify_batch_dev(
&self,
toks: &[u32],
state: &mut DecodeState,
vstate: &mut VerifyState,
taps: Option<&mut CudaSlice<f32>>,
want_logits: bool,
) -> Res<(Option<Vec<f32>>, Vec<u32>)> {
let DecodePath::Device { host_math } = self.decode_path else {
return Err("verify_batch_dev requires MEMRA_DSV4_DECODE_PATH=device".into());
};
let mc = &self.model.mc;
let d = self.model.cfg();
let t = toks.len();
assert!(
t >= 1 && t <= vstate.tmax,
"round depth {t} > tmax {}",
vstate.tmax
);
assert!(vstate.open.is_none(), "verify_batch_dev with an open round");
let pos0 = state.pos;
assert!(pos0 > 0, "batched verify needs prefill_with_cache first");
assert!(
pos0 + t <= self.max_seq,
"round [{pos0}, {}) exceeds max_seq {}",
pos0 + t,
self.max_seq
);
let hidden = mc.n_embd as usize;
let hc = d.hc_mult as usize;
let n_trunk = (mc.n_layer - mc.nextn_predict_layers) as usize;
let tok_i32: Vec<i32> = toks.iter().map(|&x| x as i32).collect();
let pos_i32: Vec<i32> = (0..t).map(|i| (pos0 + i) as i32).collect();
for (si, st) in self.stages.iter().enumerate() {
st.gpu.ctx.bind_to_thread().map_err(e("bind ctx round"))?;
let stream = st.gpu.stream();
let vws = &mut vstate.ws[si];
let mut dst = vws.tok.slice_mut(0..t);
stream
.memcpy_htod(&tok_i32, &mut dst)
.map_err(e("htod tok round"))?;
let mut dst = vws.pos_dev.slice_mut(0..t);
stream
.memcpy_htod(&pos_i32, &mut dst)
.map_err(e("htod pos round"))?;
}
{
let st0 = &self.stages[0];
st0.gpu.ctx.bind_to_thread().map_err(e("bind ctx0 round"))?;
let stream0 = st0.gpu.stream();
let vws0 = &mut vstate.ws[0];
unsafe {
ck(
"embed_rows batch",
k::memra_dsv4_embed_rows(
st0.embed
.as_ref()
.expect("embed on stage 0")
.device_ptr(&stream0)
.0 as *const c_void,
vws0.tok.device_ptr(&stream0).0 as *const i32,
dpm!(vws0.emb, &stream0),
t as i32,
hidden as i32,
sp(&stream0),
),
)?;
ck(
"repeat_hc batch",
k::memra_dsv4_repeat_hc(
dpf!(vws0.emb, &stream0),
dpm!(vws0.h_a, &stream0),
t as i32,
hc as i32,
hidden as i32,
sp(&stream0),
),
)?;
}
}
let targets = self.dspark.as_ref().map(|ds| ds.targets.clone());
let n_t = targets.as_ref().map(|x| x.len()).unwrap_or(0);
let mut taps = taps;
let mut cur_stage = 0usize;
let mut input_rx = false;
for il in 0..n_trunk {
let stage = self.layer_stage[il];
if stage != cur_stage {
let bytes = t * hc * hidden * std::mem::size_of::<f32>();
let src_stream = self.stages[cur_stage].gpu.stream();
let dst_stream = self.stages[stage].gpu.stream();
let (ws_src, ws_dst) = vstate.ws.split_at_mut(stage);
let src_ws = &ws_src[cur_stage];
let dst_ws = &mut ws_dst[0];
self.stages[cur_stage]
.gpu
.ctx
.bind_to_thread()
.map_err(e("bind tx round"))?;
let (sp_, _g0) = src_ws.h_a.device_ptr(&src_stream);
let (dp_, _g1) = dst_ws.h_rx.device_ptr_mut(&src_stream);
unsafe {
cudarc::driver::result::memcpy_peer_async(
self.stages[stage].gpu.ctx.cu_ctx(),
dp_,
self.stages[cur_stage].gpu.ctx.cu_ctx(),
sp_,
bytes,
src_stream.cu_stream(),
)
.map_err(e("peer copy h round"))?;
}
let bnd = stage - 1;
self.boundary_ev[bnd]
.record(&src_stream)
.map_err(e("ev record round"))?;
dst_stream
.wait(&self.boundary_ev[bnd])
.map_err(e("ev wait round"))?;
self.stages[stage]
.gpu
.ctx
.bind_to_thread()
.map_err(e("bind rx round"))?;
cur_stage = stage;
input_rx = true;
}
let st = &self.stages[stage];
let lidx = st
.layers
.iter()
.position(|l| l.il == il as u32)
.unwrap_or_else(|| panic!("layer {il} not on stage {stage}"));
self.block_verify_dev(
st,
&st.layers[lidx],
&mut state.caches[il],
&mut vstate.layers[il],
&mut vstate.ws[stage],
input_rx,
pos0,
t,
toks,
host_math,
)?;
input_rx = false;
if let (Some(tp), Some(tg)) = (taps.as_mut(), targets.as_ref()) {
if let Some(kk) = tg.iter().position(|&tl| tl == il) {
let stream = self.stages[stage].gpu.stream();
let vws = &mut vstate.ws[stage];
unsafe {
ck(
"hc_mean tap batch",
k::memra_dsv4_hc_mean(
dpf!(vws.h_a, &stream),
dpm!(vws.tap_tmp, &stream),
t as i32,
hc as i32,
hidden as i32,
sp(&stream),
),
)?;
ck(
"place_cols tap batch",
k::memra_dsv4_place_cols(
dpf!(vws.tap_tmp, &stream),
dpm!(**tp, &stream),
t as i32,
hidden as i32,
(n_t * hidden) as i64,
(kk * hidden) as i64,
sp(&stream),
),
)?;
}
}
}
}
let last = self.stages.len() - 1;
assert_eq!(cur_stage, last, "device path expects the head stage last");
self.head_logits_batch_dev(&mut vstate.ws[last], t, host_math)?;
let stream_last = self.stages[last].gpu.stream();
let vws = &mut vstate.ws[last];
let vocab = vws.logits.len() / vws.tmax;
let logits = if want_logits {
let mut v = vec![0f32; t * vocab];
let view = vws.logits.slice(0..t * vocab);
stream_last
.memcpy_dtoh(&view, &mut v[..])
.map_err(e("dtoh logits batch"))?;
stream_last.synchronize().map_err(e("sync logits batch"))?;
Some(v)
} else {
None
};
let mut am = vec![0i32; t];
if let Some(lg) = &logits {
for (i, slot) in am.iter_mut().enumerate() {
let row = &lg[i * vocab..(i + 1) * vocab];
let mut best = 0usize;
for j in 1..vocab {
if row[j] > row[best] {
best = j;
}
}
*slot = best as i32;
}
} else {
unsafe {
for i in 0..t {
ck(
"argmax batch",
k::memra_dsv4_argmax(
(vws.logits.device_ptr(&stream_last).0 as usize + i * vocab * 4)
as *const f32,
vocab as i64,
(vws.argmax.device_ptr_mut(&stream_last).0 as usize + i * 4)
as *mut i32,
sp(&stream_last),
),
)?;
}
}
let view = vws.argmax.slice(0..t);
stream_last
.memcpy_dtoh(&view, &mut am[..])
.map_err(e("dtoh argmax batch"))?;
stream_last.synchronize().map_err(e("sync argmax batch"))?;
}
vstate.open = Some((pos0, t));
Ok((logits, am.into_iter().map(|x| x as u32).collect()))
}
pub fn commit_verify_dev(
&self,
state: &mut DecodeState,
vstate: &mut VerifyState,
n_commit: usize,
) -> Res<()> {
let (pos0, t) = vstate
.open
.take()
.ok_or_else(|| "commit_verify_dev without an open round".to_string())?;
assert!(
n_commit >= 1 && n_commit <= t,
"commit {n_commit} outside round width {t}"
);
let d = self.model.cfg();
let mc = &self.model.mc;
let hd = d.head_dim as usize;
let win = d.sliding_window as usize;
let n_trunk = (mc.n_layer - mc.nextn_predict_layers) as usize;
let slot_rows: Vec<i32> = (0..n_commit).map(|j| ((pos0 + j) % win) as i32).collect();
for il in 0..n_trunk {
let stage = self.layer_stage[il];
let st = &self.stages[stage];
st.gpu.ctx.bind_to_thread().map_err(e("bind ctx commit"))?;
let stream = st.gpu.stream();
let lck = &mut vstate.layers[il];
let vws = &mut vstate.ws[stage];
let cache = &mut state.caches[il];
let trans_base = lck.trans_base;
{
let src = cache
.kvc
.slice(trans_base * hd..(trans_base + n_commit) * hd);
let mut dst = vws.bounce.slice_mut(0..n_commit * hd);
stream
.memcpy_dtod(&src, &mut dst)
.map_err(e("commit bounce"))?;
}
{
let mut dst = vws.slot_rows.slice_mut(0..n_commit);
stream
.memcpy_htod(&slot_rows, &mut dst)
.map_err(e("htod slot rows"))?;
}
unsafe {
ck(
"scatter_rows commit",
k::memra_dsv4_scatter_rows(
dpf!(vws.bounce, &stream),
dpm!(cache.kvc, &stream),
vws.slot_rows.device_ptr(&stream).0 as *const i32,
n_commit as i32,
hd as i32,
sp(&stream),
),
)?;
}
if let Some(ckd) = &lck.cmp {
self.cmp_rollback_replay_dev(
st,
ckd,
n_commit,
t,
pos0,
&mut vws.cmp_shift,
cache.pend_kv.as_mut().expect("pend kv"),
cache.pend_score.as_mut().expect("pend sc"),
&mut cache.n_blocks,
)?;
}
if let Some(ckd) = &lck.idx {
self.cmp_rollback_replay_dev(
st,
ckd,
n_commit,
t,
pos0,
&mut vws.cmp_shift,
cache.ipend_kv.as_mut().expect("ipend kv"),
cache.ipend_score.as_mut().expect("ipend sc"),
&mut cache.i_blocks,
)?;
}
}
for st in &self.stages {
st.gpu
.ctx
.bind_to_thread()
.map_err(e("bind ctx commit sync"))?;
st.gpu.stream().synchronize().map_err(e("commit sync"))?;
}
state.pos = pos0 + n_commit;
Ok(())
}
pub fn spec_greedy_batched_with(
&self,
prompt: &[u32],
n_new: usize,
state: &mut DecodeState,
dstate: &mut DsparkState,
vstate: &mut VerifyState,
) -> Res<SpecRunGpu> {
let depth_cap = std::env::var("MEMRA_DSV4_SPEC_DEPTH")
.ok()
.and_then(|v| v.trim().parse::<usize>().ok())
.filter(|t| *t > 0)
.unwrap_or(usize::MAX)
.max(1);
if depth_cap != usize::MAX {
println!("[spec] verify depth capped at T={depth_cap} (MEMRA_DSV4_SPEC_DEPTH)");
}
self.spec_greedy_batched_depth(prompt, n_new, state, dstate, vstate, depth_cap)
}
pub fn spec_greedy_batched_depth(
&self,
prompt: &[u32],
n_new: usize,
state: &mut DecodeState,
dstate: &mut DsparkState,
vstate: &mut VerifyState,
depth_cap: usize,
) -> Res<SpecRunGpu> {
let p0 = prompt.len();
assert!(n_new >= 1, "n_new must be positive");
let pre = self.dspark_prefill_prime(prompt, state, dstate)?;
let mut t_tok = {
let lg = &pre.logits;
let mut best = 0usize;
for i in 1..lg.len() {
if lg[i] > lg[best] {
best = i;
}
}
best as u32
};
let mut tokens: Vec<u32> = Vec::with_capacity(n_new);
let mut rounds: Vec<SpecRoundGpu> = Vec::new();
let mut mh_row = 0usize; let mut carry_pending = false;
let profile_bracket = std::env::var("MEMRA_DSV4_BENCH_PROFILE").as_deref() == Ok("1");
let depth_cap = depth_cap.max(1);
while tokens.len() < n_new {
if profile_bracket && rounds.len() == 4 {
cudarc::driver::safe::profiler_start().map_err(e("profiler_start"))?;
}
if profile_bracket && rounds.len() == 12 {
cudarc::driver::safe::profiler_stop().map_err(e("profiler_stop"))?;
}
if carry_pending {
tokens.push(t_tok);
break;
}
let round_t0 = std::time::Instant::now();
let prof_stream = if dsv4_prof_on() {
Some(self.stages[self.stages.len() - 1].gpu.stream())
} else {
None
};
let _p_round = phase!("round", prof_stream.as_ref());
let m0 = p0 + tokens.len();
let prop = {
let _p = phase!("1.drafter_forward", prof_stream.as_ref());
self.dspark_forward_spec(dstate, t_tok, mh_row, m0 - 1, false)?
};
let k_drafts = prop.out_ids.len() - 1;
tokens.push(t_tok);
if tokens.len() == n_new {
rounds.push(SpecRoundGpu {
start_pos: m0 - 1,
drafts: prop.out_ids[1..].to_vec(),
accepts: 0,
verified: 0,
t_batch: 0,
t_cap: 0,
confidence: prop.confidence.clone(),
emitted: 1,
round_us: round_t0.elapsed().as_micros() as u64,
});
break;
}
let forwards_left = n_new - tokens.len();
let t_cap = (k_drafts + 1).min(depth_cap).min(vstate.tmax);
let t_batch = t_cap.min(forwards_left);
let kv = t_batch - 1;
let mut batch_ids = Vec::with_capacity(t_batch);
batch_ids.push(t_tok);
batch_ids.extend_from_slice(&prop.out_ids[1..1 + kv]);
let (_, am) = {
let _p = phase!("2.verify_batch", prof_stream.as_ref());
self.verify_batch_dev(&batch_ids, state, vstate, Some(&mut dstate.taps), false)?
};
let mut c_d = 0usize;
let mut t_next = 0u32;
for i in 0..t_batch {
let a = am[i];
if i < kv && a == batch_ids[i + 1] {
c_d += 1;
continue;
}
t_next = a;
break;
}
let n_commit = c_d + 1;
{
let _p = phase!("3.commit_rollback", prof_stream.as_ref());
self.commit_verify_dev(state, vstate, n_commit)?;
}
{
let _p = phase!("4.ring_writes", prof_stream.as_ref());
for i in 0..n_commit {
self.dspark_write_rings(dstate, i, m0 + i)?;
}
}
{
let _p = phase!("5.round_close_sync", None);
let last = self.stages.len() - 1;
self.stages[last]
.gpu
.stream()
.synchronize()
.map_err(e("round close sync"))?;
}
mh_row = c_d;
for i in 0..c_d {
tokens.push(batch_ids[i + 1]);
}
carry_pending = c_d == kv && t_batch < t_cap;
rounds.push(SpecRoundGpu {
start_pos: m0 - 1,
drafts: prop.out_ids[1..].to_vec(),
accepts: c_d,
verified: (c_d + 1).min(kv),
t_batch,
t_cap,
confidence: prop.confidence.clone(),
emitted: 1 + c_d,
round_us: round_t0.elapsed().as_micros() as u64,
});
t_tok = t_next;
}
Ok(SpecRunGpu { tokens, rounds })
}
pub fn spec_greedy_batched(&self, prompt: &[u32], n_new: usize) -> Res<SpecRunGpu> {
let mut state = self.alloc_decode_state()?;
let mut dstate = self.dspark_alloc_state()?;
let mut vstate = self.alloc_verify_state()?;
self.spec_greedy_batched_with(prompt, n_new, &mut state, &mut dstate, &mut vstate)
}
}
impl Dsv4Gpu {
pub fn cache_classes(&self, state: &DecodeState) -> Res<Vec<(String, Vec<f32>)>> {
let d = self.model.cfg();
let hd = d.head_dim as usize;
let win = d.sliding_window as usize;
let mut out = Vec::new();
for (il, cache) in state.caches.iter().enumerate() {
let stage_i = self.layer_stage[il];
let st = &self.stages[stage_i];
st.gpu.ctx.bind_to_thread().map_err(e("bind ctx classes"))?;
let stream = st.gpu.stream();
let lidx = st
.layers
.iter()
.position(|l| l.il == il as u32)
.unwrap_or_else(|| panic!("layer {il} not on stage {stage_i}"));
let layer = &st.layers[lidx];
let read = |sl: cudarc::driver::CudaView<'_, f32>| -> Res<Vec<f32>> {
let mut v = vec![0f32; sl.len()];
stream
.memcpy_dtoh(&sl, &mut v[..])
.map_err(e("dtoh class"))?;
stream.synchronize().map_err(e("sync class"))?;
Ok(v)
};
out.push((format!("l{il}.ring"), read(cache.kvc.slice(0..win * hd))?));
if let Some(cmp) = &layer.cmp {
out.push((
format!("l{il}.cmp_store"),
read(cache.kvc.slice(win * hd..(win + cache.n_blocks) * cmp.d))?,
));
out.push((
format!("l{il}.cmp_pend_kv"),
read(cache.pend_kv.as_ref().expect("pend kv").slice(..))?,
));
out.push((
format!("l{il}.cmp_pend_score"),
read(cache.pend_score.as_ref().expect("pend sc").slice(..))?,
));
}
if let Some(ix) = &layer.idx {
let ikvc = cache.ikvc.as_ref().expect("ikvc");
out.push((
format!("l{il}.idx_store"),
read(ikvc.slice(0..cache.i_blocks * ix.cmp.d))?,
));
out.push((
format!("l{il}.idx_pend_kv"),
read(cache.ipend_kv.as_ref().expect("ipend kv").slice(..))?,
));
out.push((
format!("l{il}.idx_pend_score"),
read(cache.ipend_score.as_ref().expect("ipend sc").slice(..))?,
));
}
}
Ok(out)
}
pub fn dspark_ring_classes(&self, dstate: &DsparkState) -> Res<Vec<(String, Vec<f32>)>> {
let last = self.stages.len() - 1;
let st = &self.stages[last];
st.gpu.ctx.bind_to_thread().map_err(e("bind ctx rings"))?;
let stream = st.gpu.stream();
let d = self.model.cfg();
let hd = d.head_dim as usize;
let win = d.sliding_window as usize;
let mut out = Vec::new();
for (bi, ring) in dstate.rings.iter().enumerate() {
let view = ring.slice(0..win * hd);
let mut v = vec![0f32; view.len()];
stream
.memcpy_dtoh(&view, &mut v[..])
.map_err(e("dtoh ring class"))?;
stream.synchronize().map_err(e("sync ring class"))?;
out.push((format!("dspark.ring{bi}"), v));
}
Ok(out)
}
}
pub fn resolve_dense_arm(v: Option<&str>, on_device: bool) -> Result<bool, String> {
match v {
None | Some("") => Ok(on_device),
Some("bf16") => Ok(false),
Some("fp8") if !on_device => Err(
"MEMRA_DSV4_DENSE_ARM=fp8 requires MEMRA_DSV4_DECODE_PATH=device (the \
fp8 GEMV twins exist on the device decode/verify paths only; prefill \
and the legacy path consume the bf16 slabs)"
.to_string(),
),
Some("fp8") => Ok(true),
Some(other) => Err(format!(
"MEMRA_DSV4_DENSE_ARM '{other}' unknown (bf16 | fp8)"
)),
}
}
#[cfg(test)]
mod dense_arm_default_tests {
use super::resolve_dense_arm;
#[test]
fn ratified_default_dense_arm_is_fp8_on_device() {
assert_eq!(
resolve_dense_arm(None, true),
Ok(true),
"owner-ratified 2026-08-20: unset MEMRA_DSV4_DENSE_ARM defaults the DEVICE \
decode path to fp8 (bit-identical on four boxes, x5 A/B 41.06->47.19, \
item-3 residency green on box7)"
);
assert_eq!(resolve_dense_arm(Some(""), true), Ok(true));
assert_eq!(resolve_dense_arm(None, false), Ok(false));
assert_eq!(resolve_dense_arm(Some("bf16"), true), Ok(false));
assert_eq!(resolve_dense_arm(Some("fp8"), true), Ok(true));
assert!(
resolve_dense_arm(Some("fp8"), false).is_err(),
"legacy+fp8 stays a refusal"
);
assert!(
resolve_dense_arm(Some("q8"), true).is_err(),
"unknown values refuse"
);
}
}