#![allow(
clippy::manual_is_multiple_of,
clippy::collapsible_if,
clippy::needless_range_loop,
clippy::too_many_arguments,
clippy::unnecessary_unwrap,
clippy::needless_question_mark,
clippy::needless_option_as_deref,
clippy::extend_with_drain,
clippy::type_complexity,
clippy::large_enum_variant
)]
use std::os::raw::c_void;
use cudarc::driver::{CudaSlice, CudaView, DevicePtr, DevicePtrMut, LaunchConfig, PushKernelArg};
use memra_gguf::model_plan::{
AttentionPlan, FullAttentionPlan, GatedDeltaNetPlan, GdnGateActivation, MicroBlockIndexPlan,
MlpPlan, ModelPlan, MoeMlpPlan, PleEmbeddingPlan, ResidualTopology, RopeFactors, RopePlan,
RouterPlan, TensorPresence, yarn_attention_factor, yarn_frequency_divisors,
};
use memra_gguf::tensor_contract::{LayerTensor, TensorId};
use memra_reference::{ReferenceTensor, ReferenceWeights};
use crate::Engine;
type Res<T> = Result<T, Box<dyn std::error::Error>>;
struct GateW {
norm: Vec<CudaSlice<f32>>,
norm_stack: CudaSlice<f32>,
down: Vec<CudaSlice<f32>>,
up: Vec<CudaSlice<f32>>,
inject: Option<CudaSlice<f32>>,
down_b16: Option<CudaSlice<u8>>,
up_b16: Option<CudaSlice<u8>>,
inject_b16: Option<CudaSlice<u8>>,
}
fn gate_rank(gate: &GateW, hidden: usize, streams: usize) -> Res<usize> {
if gate.down[0].len() >= hidden {
return Ok(gate.down[0].len() / hidden);
}
match gate.down_b16.as_ref() {
Some(w) => Ok(w.len() / (2 * streams * hidden)),
None => Err("qwen4exp_gpu: gate rank underivable (f32 dropped and no bf16 twin)".into()),
}
}
struct QsaW {
attn: FullAttentionPlan,
overlay: MicroBlockIndexPlan,
wq: CudaSlice<f32>, wk: CudaSlice<f32>, wv: CudaSlice<f32>, wo: CudaSlice<f32>, q_norm: Option<CudaSlice<f32>>,
k_norm: Option<CudaSlice<f32>>,
idx_proj: CudaSlice<f32>, idx_q_norm: Vec<f32>,
idx_k_norm: Vec<f32>,
proj_b16: Option<CudaSlice<u8>>,
wo_b16: Option<CudaSlice<u8>>,
yarn: Option<YarnRopeW>,
}
struct YarnRopeW {
ff: CudaSlice<f32>,
ff_host: Vec<f32>,
mscale: f32,
}
fn build_yarn(
e: &Engine,
rope: &RopePlan,
overlay: Option<&MicroBlockIndexPlan>,
layer: u32,
) -> Res<Option<YarnRopeW>> {
match rope.factors {
RopeFactors::None => Ok(None),
RopeFactors::Yarn {
factor,
original_context,
beta_fast,
beta_slow,
} => {
if let Some(overlay) = overlay
&& overlay.rope_dimensions != rope.dimensions
{
return Err(format!(
"qwen4exp_gpu: layer {layer} indexer rope width {} != attention rope \
width {} — the shared yarn table would be wrong",
overlay.rope_dimensions, rope.dimensions
)
.into());
}
let ff_host = yarn_frequency_divisors(
rope.dimensions,
rope.base,
factor,
original_context,
beta_fast,
beta_slow,
);
Ok(Some(YarnRopeW {
ff: e.htod(&ff_host)?,
ff_host,
mscale: yarn_attention_factor(factor),
}))
}
_ => Err(format!(
"qwen4exp_gpu: layer {layer}: only plain or yarn rope factors are supported"
)
.into()),
}
}
struct GdnW {
plan: GatedDeltaNetPlan,
qkv: CudaSlice<f32>, z: CudaSlice<f32>, beta: CudaSlice<f32>, alpha: CudaSlice<f32>, conv_w: CudaSlice<f32>, a: CudaSlice<f32>, dt: CudaSlice<f32>, norm: CudaSlice<f32>, out: CudaSlice<f32>, proj_b16: Option<CudaSlice<u8>>,
out_b16: Option<CudaSlice<u8>>,
}
enum MixerW {
Qsa(QsaW),
Gdn(GdnW),
}
enum BankHalf {
F32(CudaSlice<f32>),
Nvfp4 {
codes: CudaSlice<u8>,
scales: CudaSlice<u8>,
macros: Vec<f32>,
macros_dev: CudaSlice<f32>,
},
HostBf16(Vec<u8>),
DeviceBf16(CudaSlice<u8>),
}
struct ExpertBank {
gate: BankHalf, up: BankHalf, down: BankHalf, }
struct MoeW {
plan: MoeMlpPlan,
router: CudaSlice<f32>, router_b16: Option<CudaSlice<u8>>,
bank: ExpertBank,
shared_gate: CudaSlice<f32>,
shared_up: CudaSlice<f32>,
shared_down: CudaSlice<f32>,
shared_input_gate: Option<CudaSlice<f32>>, shared_gu_b16: Option<CudaSlice<u8>>,
shared_down_b16: Option<CudaSlice<u8>>,
}
enum NgramTable {
F32(Vec<f32>),
Bf16(Vec<u8>),
}
impl NgramTable {
fn rows(&self, head_dim: usize) -> usize {
match self {
Self::F32(data) => data.len() / head_dim,
Self::Bf16(bytes) => bytes.len() / 2 / head_dim,
}
}
fn gather_into(&self, row: usize, head_dim: usize, dst: &mut [f32]) {
match self {
Self::F32(data) => {
dst.copy_from_slice(&data[row * head_dim..(row + 1) * head_dim]);
}
Self::Bf16(bytes) => {
let start = row * head_dim * 2;
for (i, out) in dst.iter_mut().enumerate() {
let b = u16::from_le_bytes([bytes[start + 2 * i], bytes[start + 2 * i + 1]]);
*out = f32::from_bits(u32::from(b) << 16);
}
}
}
}
}
struct PleW {
plan: PleEmbeddingPlan,
key_proj: Vec<CudaSlice<f32>>, value_proj: CudaSlice<f32>, norm_key: Vec<CudaSlice<f32>>, norm_query: Vec<CudaSlice<f32>>,
norm_conv: Vec<CudaSlice<f32>>,
conv_w: Vec<CudaSlice<f32>>, multipliers: Vec<i64>,
sizes: Vec<i64>,
offsets: Vec<i64>,
table: NgramTable,
}
struct LayerW {
index: u32,
eps_attn: f32,
eps_mlp: f32,
attn_gate: GateW,
mlp_gate: GateW,
mixer: MixerW,
moe: MoeW,
ple: Option<PleW>,
}
struct MtpW {
eps_embed: f32,
eps_hidden: f32,
pre_norm_embed: CudaSlice<f32>,
pre_norm_hidden: CudaSlice<f32>,
fc_embed: CudaSlice<f32>, fc_embed_b16: Option<CudaSlice<u8>>,
fc_hidden: CudaSlice<f32>, fc_hidden_b16: Option<CudaSlice<u8>>,
layer: LayerW,
mixer: GateW,
}
struct MtpDev1 {
dev: usize,
output: CudaSlice<f32>,
output_b16: Option<CudaSlice<u8>>,
}
struct DraftTrim {
n: usize,
d2t: Vec<u32>,
head_b16: Option<CudaSlice<u8>>,
head: Option<CudaSlice<f32>>,
}
fn linear_trim_into(
e: &Engine,
trim: &DraftTrim,
x: &CudaSlice<f32>,
y: &mut CudaSlice<f32>,
t: usize,
in_f: usize,
) -> Res<()> {
if trunk_bf16_on() {
if let Some(w) = trim.head_b16.as_ref() {
if (2..=12).contains(&t) && verify_mt_on() {
return launch_qmatvec_bf16w_mt(e, w, 0, x, y, in_f, trim.n, t);
}
return launch_qmatvec_bf16w(e, w, x, y, in_f, trim.n, t, 1, 0, 0, in_f, 0);
}
}
let w = trim.head.as_ref().ok_or(
"qwen4exp_gpu: the draft trim was gathered bf16-only — the f32 head arm needs \
trunk_bf16 on (or a checkpoint without a bf16 lm-head twin)",
)?;
e.linear_device_into(x, w, y, t, in_f, trim.n)
}
pub struct Qwen4ExpGpu {
pub plan: ModelPlan,
hidden: usize,
streams: usize,
vocab: usize,
embed_host: Vec<f32>, output: CudaSlice<f32>, output_b16: Option<CudaSlice<u8>>,
layers: Vec<LayerW>,
exit_mixer: GateW,
exit_eps: f32,
mtp: Option<MtpW>,
mtp_dev1: Option<MtpDev1>,
draft_trim: Option<DraftTrim>,
draft_trim_parked: Option<DraftTrim>,
chain_embed: Option<ChainEmbed>,
}
struct ChainEmbed {
table: CudaSlice<u8>,
qt: i32,
row_bytes: usize,
rows: usize,
for_trim: bool,
dev: usize,
}
struct PleState {
conv_hist: Vec<CudaSlice<f32>>,
ngram_ids: Vec<i64>,
ngram_history: Vec<i64>,
ngram_last_eos: i64,
}
fn q8_row_bytes(dim: usize) -> usize {
dim.div_ceil(32) * 34
}
fn q5_row_bytes(dim: usize) -> usize {
dim.div_ceil(32) * 24
}
enum QsaKvStore {
F32 {
k: CudaSlice<f32>, v: CudaSlice<f32>, },
Q8Q5 {
k: CudaSlice<u8>, v: CudaSlice<u8>, },
}
impl QsaKvStore {
fn is_quant(&self) -> bool {
matches!(self, QsaKvStore::Q8Q5 { .. })
}
fn capacity_rows(&self, kv_dim: usize) -> usize {
match self {
QsaKvStore::F32 { k, .. } => k.len() / kv_dim,
QsaKvStore::Q8Q5 { k, .. } => k.len() / q8_row_bytes(kv_dim),
}
}
}
fn host_quant_q8_row(row: &[f32], dim: usize, out: &mut Vec<u8>) {
for b in 0..dim.div_ceil(32) {
let mut amax = 0.0f32;
for l in 0..32 {
let e = b * 32 + l;
let x = if e < dim { row[e] } else { 0.0 };
amax = amax.max(x.abs());
}
let d = amax / 127.0f32;
let mut id = if d != 0.0 { 1.0f32 / d } else { 0.0 };
if !id.is_finite() {
id = 0.0;
}
out.extend_from_slice(&memra_gguf::nvfp4_repack::f32_to_f16_bits(d).to_le_bytes());
for l in 0..32 {
let e = b * 32 + l;
let x = if e < dim { row[e] } else { 0.0 };
let q = ((x * id).round_ties_even() as i32).clamp(-127, 127);
out.push(q as i8 as u8);
}
}
}
fn host_deq_q8_rows(bytes: &[u8], row0: usize, rows: usize, dim: usize, out: &mut Vec<f32>) {
let rb = q8_row_bytes(dim);
for r in row0..row0 + rows {
let row = &bytes[r * rb..(r + 1) * rb];
for e in 0..dim {
let blk = &row[(e >> 5) * 34..];
let d = memra_gguf::dequant::fp16_to_f32(u16::from_le_bytes([blk[0], blk[1]]));
let q = blk[2 + (e & 31)] as i8 as f32;
out.push(d * q);
}
}
}
fn host_quant_q5_row(row: &[f32], dim: usize, out: &mut Vec<u8>) {
for b in 0..dim.div_ceil(32) {
let lane = |l: usize| -> f32 {
let e = b * 32 + l;
if e < dim { row[e] } else { 0.0 }
};
let mut mn = f32::INFINITY;
let mut mx = f32::NEG_INFINITY;
for l in 0..32 {
mn = mn.min(lane(l));
mx = mx.max(lane(l));
}
let d = (mx - mn) / 31.0f32;
let mut id = if d != 0.0 { 1.0f32 / d } else { 0.0 };
if !id.is_finite() {
id = 0.0;
}
let q5 = |l: usize| -> u32 {
(((lane(l) - mn) * id).round_ties_even() as i32).clamp(0, 31) as u32
};
let mut qh = 0u32;
for l in 0..32 {
qh |= ((q5(l) >> 4) & 1) << l;
}
out.extend_from_slice(&memra_gguf::nvfp4_repack::f32_to_f16_bits(d).to_le_bytes());
out.extend_from_slice(&memra_gguf::nvfp4_repack::f32_to_f16_bits(mn).to_le_bytes());
out.extend_from_slice(&qh.to_le_bytes());
for l in 0..16 {
out.push(((q5(l) & 0x0F) | ((q5(l + 16) & 0x0F) << 4)) as u8);
}
}
}
fn host_deq_q5_rows(bytes: &[u8], row0: usize, rows: usize, dim: usize, out: &mut Vec<f32>) {
let rb = q5_row_bytes(dim);
for r in row0..row0 + rows {
let row = &bytes[r * rb..(r + 1) * rb];
for e in 0..dim {
let blk = &row[(e >> 5) * 24..];
let d = memra_gguf::dequant::fp16_to_f32(u16::from_le_bytes([blk[0], blk[1]]));
let m = memra_gguf::dequant::fp16_to_f32(u16::from_le_bytes([blk[2], blk[3]]));
let qh = u32::from_le_bytes([blk[4], blk[5], blk[6], blk[7]]);
let lane = e & 31;
let lo = if lane < 16 {
blk[8 + lane] & 0x0F
} else {
blk[8 + lane - 16] >> 4
};
let q5 = (lo as u32) | (((qh >> lane) & 1) << 4);
out.push(d.mul_add(q5 as f32, m));
}
}
}
fn f32_to_bf16_rne(x: f32) -> u16 {
let bits = x.to_bits();
let rounding_bias = 0x7fff + ((bits >> 16) & 1);
(bits.wrapping_add(rounding_bias) >> 16) as u16
}
enum IdxRawCache {
F32(Vec<f32>),
Q8(Vec<u8>),
Bf16(Vec<u16>),
}
impl IdxRawCache {
fn new(mode: IdxQMode) -> Self {
match mode {
IdxQMode::F32 => IdxRawCache::F32(Vec::new()),
IdxQMode::Q8 => IdxRawCache::Q8(Vec::new()),
IdxQMode::Bf16 => IdxRawCache::Bf16(Vec::new()),
}
}
fn rows(&self, idx_dim: usize) -> usize {
match self {
IdxRawCache::F32(v) => v.len() / idx_dim,
IdxRawCache::Q8(v) => v.len() / q8_row_bytes(idx_dim),
IdxRawCache::Bf16(v) => v.len() / idx_dim,
}
}
fn truncate_rows(&mut self, rows: usize, idx_dim: usize) {
match self {
IdxRawCache::F32(v) => v.truncate(rows * idx_dim),
IdxRawCache::Q8(v) => v.truncate(rows * q8_row_bytes(idx_dim)),
IdxRawCache::Bf16(v) => v.truncate(rows * idx_dim),
}
}
fn append_rows_f32(&mut self, rows: &[f32], n: usize, idx_dim: usize) {
match self {
IdxRawCache::F32(v) => v.extend_from_slice(&rows[..n * idx_dim]),
IdxRawCache::Q8(v) => {
for r in 0..n {
host_quant_q8_row(&rows[r * idx_dim..(r + 1) * idx_dim], idx_dim, v);
}
}
IdxRawCache::Bf16(v) => {
v.extend(rows[..n * idx_dim].iter().map(|&x| f32_to_bf16_rne(x)));
}
}
}
fn rows_f32(&self, row0: usize, n: usize, idx_dim: usize, out: &mut Vec<f32>) {
out.clear();
match self {
IdxRawCache::F32(v) => out.extend_from_slice(&v[row0 * idx_dim..(row0 + n) * idx_dim]),
IdxRawCache::Q8(v) => host_deq_q8_rows(v, row0, n, idx_dim, out),
IdxRawCache::Bf16(v) => out.extend(
v[row0 * idx_dim..(row0 + n) * idx_dim]
.iter()
.map(|&b| memra_gguf::dequant::bf16_to_f32(b)),
),
}
}
}
enum IdxRawDev {
F32(CudaSlice<f32>),
Q8(CudaSlice<u8>),
Bf16(CudaSlice<u16>),
}
fn idx_materialize_host(
e: &Engine,
raw_keys: &mut IdxRawCache,
raw_dev: &Option<IdxRawDev>,
raw_dev_rows: usize,
idx_dim: usize,
) -> Res<()> {
let host_rows = raw_keys.rows(idx_dim);
if raw_dev_rows <= host_rows {
return Ok(());
}
let m = raw_dev
.as_ref()
.ok_or("idxcache: rows counted without a cache")?;
match (m, raw_keys) {
(IdxRawDev::F32(d), IdxRawCache::F32(h)) => {
let delta = e.dtoh_view(&d.slice(host_rows * idx_dim..raw_dev_rows * idx_dim))?;
h.extend_from_slice(&delta);
}
(IdxRawDev::Q8(d), IdxRawCache::Q8(h)) => {
let rb = q8_row_bytes(idx_dim);
let delta = e.dtoh_u8_view(&d.slice(host_rows * rb..raw_dev_rows * rb))?;
h.extend_from_slice(&delta);
}
(IdxRawDev::Bf16(d), IdxRawCache::Bf16(h)) => {
let delta = e
.gpu
.stream()
.clone_dtoh(&d.slice(host_rows * idx_dim..raw_dev_rows * idx_dim))?;
e.gpu.stream().synchronize()?;
h.extend_from_slice(&delta);
}
_ => return Err("idxcache: device/host raw-key formats disagree".into()),
}
Ok(())
}
struct IdxAudit {
raw_f32: IdxRawCache, pooled_f32: Vec<f32>,
}
enum MixerState {
Qsa {
kv: QsaKvStore,
raw_keys: IdxRawCache,
pooled_keys: Vec<f32>,
pooled_dev: Option<CudaSlice<f32>>,
pooled_dev_rows: usize,
raw_dev: Option<IdxRawDev>,
raw_dev_rows: usize,
idx_audit: Option<Box<IdxAudit>>,
},
Gdn {
conv: CudaSlice<f32>, state: CudaSlice<f32>, },
}
struct LayerState {
mixer: MixerState,
ple: Option<PleState>,
}
struct GdnStash {
states: CudaSlice<f32>,
conv_pre: CudaSlice<f32>,
qkv_rows: CudaSlice<f32>,
scan_graph: Option<(usize, GraphEntry)>,
scan_warm: Option<usize>,
}
struct PleStash {
hist_pre: Vec<CudaSlice<f32>>, normed_rows: Vec<CudaSlice<f32>>, }
pub struct VerifyStash {
k_cap: usize,
chunk: Option<(usize, usize)>,
fused_chunk: Option<(usize, usize)>,
gdn: Vec<Option<GdnStash>>,
ple: Vec<Option<PleStash>>,
wide: CudaSlice<f32>,
ring_rows: usize,
wide_dev1: Option<CudaSlice<f32>>,
argmax: Vec<u32>,
toks: CudaSlice<u32>,
want_argmax: bool,
want_argmax_t1: bool,
last_row_only: bool,
}
pub struct Qwen4ExpState {
pos: usize,
capacity: usize,
reserve: usize,
tokens: Vec<u32>,
layers: Vec<LayerState>,
ws: StepPool,
graphs: StepGraphs,
tp2: Option<Tp2State>,
verify: Option<VerifyStash>,
}
#[derive(Default)]
struct StepPool {
f32s: std::collections::BTreeMap<&'static str, CudaSlice<f32>>,
i32s: std::collections::BTreeMap<&'static str, CudaSlice<i32>>,
u8s: std::collections::BTreeMap<&'static str, CudaSlice<u8>>,
u64s: std::collections::BTreeMap<&'static str, CudaSlice<u64>>,
}
const EXIT_PLANE_SLOTS: [&str; 8] = [
"exit.p0", "exit.p1", "exit.p2", "exit.p3", "exit.p4", "exit.p5", "exit.p6", "exit.p7",
];
const PLANE_SLOTS: [&str; 8] = [
"plane.0", "plane.1", "plane.2", "plane.3", "plane.4", "plane.5", "plane.6", "plane.7",
];
const INJECT_SLOTS: [&str; 8] = [
"hc.inj.0", "hc.inj.1", "hc.inj.2", "hc.inj.3", "hc.inj.4", "hc.inj.5", "hc.inj.6", "hc.inj.7",
];
impl StepPool {
fn take_f32(
&mut self,
e: &Engine,
name: &'static str,
len: usize,
reserve: usize,
) -> Res<CudaSlice<f32>> {
if step_ws_on() {
if let Some(buf) = self.f32s.remove(name) {
if buf.len() >= len {
return Ok(buf);
}
}
e.uninit(len.max(reserve))
} else {
e.uninit(len)
}
}
fn put_f32(&mut self, name: &'static str, buf: CudaSlice<f32>) {
if step_ws_on() {
self.f32s.insert(name, buf);
}
}
fn shed(&mut self) -> usize {
let bytes = self.f32s.values().map(|b| b.len() * 4).sum::<usize>()
+ self.i32s.values().map(|b| b.len() * 4).sum::<usize>()
+ self.u8s.values().map(|b| b.len()).sum::<usize>()
+ self.u64s.values().map(|b| b.len() * 8).sum::<usize>();
self.f32s.clear();
self.i32s.clear();
self.u8s.clear();
self.u64s.clear();
bytes
}
fn take_i32(
&mut self,
e: &Engine,
name: &'static str,
host: &[i32],
reserve: usize,
) -> Res<CudaSlice<i32>> {
if step_ws_on() {
let mut buf = match self.i32s.remove(name) {
Some(buf) if buf.len() >= host.len() => buf,
_ => e.alloc_uninit::<i32>(host.len().max(reserve))?,
};
let mut view = buf.slice_mut(0..host.len());
e.gpu.stream().memcpy_htod(host, &mut view)?;
Ok(buf)
} else {
e.htod_i32(host)
}
}
fn put_i32(&mut self, name: &'static str, buf: CudaSlice<i32>) {
if step_ws_on() {
self.i32s.insert(name, buf);
}
}
fn take_i32_slot(
&mut self,
e: &Engine,
name: &'static str,
len: usize,
reserve: usize,
) -> Res<CudaSlice<i32>> {
if step_ws_on() {
if let Some(buf) = self.i32s.remove(name) {
if buf.len() >= len {
return Ok(buf);
}
}
}
e.alloc_uninit::<i32>(len.max(reserve))
}
fn take_f32_h2d(
&mut self,
e: &Engine,
name: &'static str,
host: &[f32],
reserve: usize,
) -> Res<CudaSlice<f32>> {
if step_ws_on() {
let mut buf = match self.f32s.remove(name) {
Some(buf) if buf.len() >= host.len() => buf,
_ => e.uninit(host.len().max(reserve))?,
};
let mut view = buf.slice_mut(0..host.len());
e.gpu.stream().memcpy_htod(host, &mut view)?;
Ok(buf)
} else {
e.htod(host)
}
}
fn take_u8_h2d(
&mut self,
e: &Engine,
name: &'static str,
host: &[u8],
reserve: usize,
) -> Res<CudaSlice<u8>> {
if step_ws_on() {
let mut buf = match self.u8s.remove(name) {
Some(buf) if buf.len() >= host.len() => buf,
_ => e.alloc_u8_uninit(host.len().max(reserve))?,
};
let mut view = buf.slice_mut(0..host.len());
e.gpu.stream().memcpy_htod(host, &mut view)?;
Ok(buf)
} else {
e.htod_bytes(host)
}
}
fn put_u8(&mut self, name: &'static str, buf: CudaSlice<u8>) {
if step_ws_on() {
self.u8s.insert(name, buf);
}
}
fn take_u8(
&mut self,
e: &Engine,
name: &'static str,
len: usize,
reserve: usize,
) -> Res<CudaSlice<u8>> {
if step_ws_on() {
if let Some(buf) = self.u8s.remove(name) {
if buf.len() >= len {
return Ok(buf);
}
}
}
e.alloc_u8_uninit(len.max(reserve))
}
fn upsert_u8(
&mut self,
e: &Engine,
name: &'static str,
host: &[u8],
reserve: usize,
) -> Res<()> {
if !self.u8s.contains_key(name) {
let buf = self.take_u8(e, name, host.len(), reserve)?;
self.put_u8(name, buf);
}
let buf = self
.u8s
.get_mut(name)
.ok_or_else(|| format!("step workspace: slot {name} is not parked"))?;
if buf.len() < host.len() {
return Err(format!("step workspace: slot {name} is too small").into());
}
let mut view = buf.slice_mut(0..host.len());
e.gpu.stream().memcpy_htod(host, &mut view)?;
Ok(())
}
fn peek_u8(&self, name: &'static str) -> Res<&CudaSlice<u8>> {
self.u8s
.get(name)
.ok_or_else(|| format!("step workspace: slot {name} is not parked").into())
}
fn take_u64_h2d(
&mut self,
e: &Engine,
name: &'static str,
host: &[u64],
reserve: usize,
) -> Res<CudaSlice<u64>> {
if step_ws_on() {
let mut buf = match self.u64s.remove(name) {
Some(buf) if buf.len() >= host.len() => buf,
_ => e.alloc_uninit::<u64>(host.len().max(reserve))?,
};
let mut view = buf.slice_mut(0..host.len());
e.gpu.stream().memcpy_htod(host, &mut view)?;
Ok(buf)
} else {
e.htod_u64(host)
}
}
fn put_u64(&mut self, name: &'static str, buf: CudaSlice<u64>) {
if step_ws_on() {
self.u64s.insert(name, buf);
}
}
fn peek_f32(&self, name: &'static str) -> Res<&CudaSlice<f32>> {
self.f32s
.get(name)
.ok_or_else(|| format!("step workspace: slot {name} is not parked").into())
}
fn write_i32(&mut self, e: &Engine, name: &'static str, host: &[i32]) -> Res<()> {
let buf = self
.i32s
.get_mut(name)
.ok_or_else(|| format!("step workspace: slot {name} is not parked"))?;
if buf.len() < host.len() {
return Err(format!("step workspace: slot {name} is too small").into());
}
let mut view = buf.slice_mut(0..host.len());
e.gpu.stream().memcpy_htod(host, &mut view)?;
Ok(())
}
fn write_f32(&mut self, e: &Engine, name: &'static str, host: &[f32]) -> Res<()> {
let buf = self
.f32s
.get_mut(name)
.ok_or_else(|| format!("step workspace: slot {name} is not parked"))?;
if buf.len() < host.len() {
return Err(format!("step workspace: slot {name} is too small").into());
}
let mut view = buf.slice_mut(0..host.len());
e.gpu.stream().memcpy_htod(host, &mut view)?;
Ok(())
}
}
#[derive(Default)]
struct StepGraphs {
warm: bool,
a: Vec<Option<GraphEntry>>,
b: Vec<Option<GraphEntry>>,
exit: Option<GraphEntry>,
}
type GraphEntry = (
cudarc::driver::CudaGraph,
Vec<Box<dyn std::any::Any + Send>>,
);
impl Qwen4ExpState {
pub fn position(&self) -> usize {
self.pos
}
}
pub struct PrefillCapture {
pub layer_wide: Vec<Vec<f32>>,
pub exit_mixed: Vec<f32>,
}
pub mod prof {
use std::cell::RefCell;
use std::collections::BTreeMap;
thread_local! {
static STATE: RefCell<Option<BTreeMap<&'static str, (f64, u64)>>> =
const { RefCell::new(None) };
}
thread_local! {
static PREFILL: RefCell<Option<BTreeMap<&'static str, (f64, u64)>>> =
const { RefCell::new(None) };
static ROUNDS_ONLY: RefCell<bool> = const { RefCell::new(false) };
}
pub fn enable() {
STATE.with(|s| *s.borrow_mut() = Some(BTreeMap::new()));
PREFILL.with(|s| *s.borrow_mut() = None);
}
pub fn set_rounds_only(on: bool) {
ROUNDS_ONLY.with(|s| *s.borrow_mut() = on);
}
pub fn rounds_only() -> bool {
ROUNDS_ONLY.with(|s| *s.borrow())
}
pub fn split_prefill() {
if !on() || !rounds_only() {
return;
}
let pre = STATE.with(|s| s.borrow_mut().replace(BTreeMap::new()));
PREFILL.with(|s| *s.borrow_mut() = pre);
}
pub fn take_prefill() -> Vec<(&'static str, f64, u64)> {
PREFILL.with(|s| {
s.borrow_mut()
.take()
.map(|map| map.into_iter().map(|(k, (t, c))| (k, t, c)).collect())
.unwrap_or_default()
})
}
pub fn on() -> bool {
STATE.with(|s| s.borrow().is_some())
}
pub fn take() -> Vec<(&'static str, f64, u64)> {
STATE.with(|s| {
s.borrow_mut()
.take()
.map(|map| map.into_iter().map(|(k, (t, c))| (k, t, c)).collect())
.unwrap_or_default()
})
}
pub(super) fn add(name: &'static str, seconds: f64) {
STATE.with(|s| {
if let Some(map) = s.borrow_mut().as_mut() {
let entry = map.entry(name).or_insert((0.0, 0));
entry.0 += seconds;
entry.1 += 1;
}
});
}
}
static MOE_SEL_PATH: std::sync::atomic::AtomicBool = std::sync::atomic::AtomicBool::new(true);
pub fn set_moe_sel_path(on: bool) {
MOE_SEL_PATH.store(on, std::sync::atomic::Ordering::Relaxed);
}
fn moe_sel_path_on() -> bool {
MOE_SEL_PATH.load(std::sync::atomic::Ordering::Relaxed)
}
static HC_FUSED_GATE: std::sync::atomic::AtomicBool = std::sync::atomic::AtomicBool::new(true);
pub fn set_hc_fused_gate(on: bool) {
HC_FUSED_GATE.store(on, std::sync::atomic::Ordering::Relaxed);
}
fn hc_fused_gate_on() -> bool {
HC_FUSED_GATE.load(std::sync::atomic::Ordering::Relaxed)
}
static TRUNK_BF16: std::sync::atomic::AtomicBool = std::sync::atomic::AtomicBool::new(true);
pub fn set_trunk_bf16(on: bool) {
TRUNK_BF16.store(on, std::sync::atomic::Ordering::Relaxed);
}
fn trunk_bf16_on() -> bool {
TRUNK_BF16.load(std::sync::atomic::Ordering::Relaxed)
}
static PREFILL_GROUPED_ALL: std::sync::atomic::AtomicBool =
std::sync::atomic::AtomicBool::new(false);
pub fn set_prefill_grouped_all(on: bool) {
PREFILL_GROUPED_ALL.store(on, std::sync::atomic::Ordering::Relaxed);
}
fn prefill_grouped_all_on() -> bool {
PREFILL_GROUPED_ALL.load(std::sync::atomic::Ordering::Relaxed)
}
static STEP_WS: std::sync::atomic::AtomicBool = std::sync::atomic::AtomicBool::new(true);
pub fn set_step_ws(on: bool) {
STEP_WS.store(on, std::sync::atomic::Ordering::Relaxed);
}
fn step_ws_on() -> bool {
STEP_WS.load(std::sync::atomic::Ordering::Relaxed)
}
static DECODE_GRAPHS: std::sync::atomic::AtomicBool = std::sync::atomic::AtomicBool::new(true);
pub fn set_decode_graphs(on: bool) {
DECODE_GRAPHS.store(on, std::sync::atomic::Ordering::Relaxed);
}
fn decode_graphs_on() -> bool {
DECODE_GRAPHS.load(std::sync::atomic::Ordering::Relaxed)
}
static VERIFY_GRAPHS: std::sync::atomic::AtomicBool = std::sync::atomic::AtomicBool::new(false);
pub fn set_verify_graphs(on: bool) {
VERIFY_GRAPHS.store(on, std::sync::atomic::Ordering::Relaxed);
}
fn verify_graphs_on() -> bool {
VERIFY_GRAPHS.load(std::sync::atomic::Ordering::Relaxed)
}
static SEL_V2: std::sync::atomic::AtomicBool = std::sync::atomic::AtomicBool::new(true);
pub fn set_sel_v2(on: bool) {
SEL_V2.store(on, std::sync::atomic::Ordering::Relaxed);
}
fn sel_v2_on() -> bool {
SEL_V2.load(std::sync::atomic::Ordering::Relaxed)
}
pub const SEL_V3_DEFAULT: bool = true;
static SEL_V3: std::sync::atomic::AtomicBool = std::sync::atomic::AtomicBool::new(SEL_V3_DEFAULT);
pub fn set_sel_v3(on: bool) {
SEL_V3.store(on, std::sync::atomic::Ordering::Relaxed);
}
fn sel_v3_on() -> bool {
SEL_V3.load(std::sync::atomic::Ordering::Relaxed)
}
const SEL_GROUP_OFF: u32 = 0;
const SEL_GROUP_AUTO: u32 = 1;
static SEL_GROUP_DN: std::sync::atomic::AtomicU32 =
std::sync::atomic::AtomicU32::new(SEL_GROUP_AUTO);
static SEL_GROUP_GU: std::sync::atomic::AtomicU32 =
std::sync::atomic::AtomicU32::new(SEL_GROUP_AUTO);
fn sel_group_dn() -> u32 {
SEL_GROUP_DN.load(std::sync::atomic::Ordering::Relaxed)
}
fn sel_group_gu() -> u32 {
SEL_GROUP_GU.load(std::sync::atomic::Ordering::Relaxed)
}
pub fn sel_group_spec() -> String {
let one = |c: u32| -> String {
match c {
SEL_GROUP_OFF => "off".to_string(),
SEL_GROUP_AUTO => "auto".to_string(),
v => format!("{}:{}", (v >> 8) & 0xff, v & 0xff),
}
};
format!("dn:{}+gu:{}", one(sel_group_dn()), one(sel_group_gu()))
}
pub fn set_sel_group(spec: &str) -> bool {
let parse_one = |s: &str| -> Option<u32> {
match s {
"off" | "0" => Some(SEL_GROUP_OFF),
"auto" | "1" | "" => Some(SEL_GROUP_AUTO),
other => {
let (g, rows) = other.split_once(':')?;
let g: u32 = g.parse().ok()?;
let rows: u32 = rows.parse().ok()?;
if !matches!(g, 1 | 2 | 4 | 8 | 16 | 32) || !matches!(rows, 1 | 2 | 4) {
return None;
}
Some((g << 8) | rows)
}
}
};
if let Some(both) = parse_one(spec) {
SEL_GROUP_DN.store(both, std::sync::atomic::Ordering::Relaxed);
SEL_GROUP_GU.store(both, std::sync::atomic::Ordering::Relaxed);
return true;
}
let mut dn = None;
let mut gu = None;
for part in spec.split('+').filter(|p| !p.is_empty()) {
let Some((family, rest)) = part.split_once(':') else {
return false;
};
let Some(code) = parse_one(rest) else {
return false;
};
match family {
"dn" | "down" => dn = Some(code),
"gu" | "gateup" => gu = Some(code),
_ => return false,
}
}
if dn.is_none() && gu.is_none() {
return false;
}
if let Some(c) = dn {
SEL_GROUP_DN.store(c, std::sync::atomic::Ordering::Relaxed);
}
if let Some(c) = gu {
SEL_GROUP_GU.store(c, std::sync::atomic::Ordering::Relaxed);
}
true
}
fn sel_group_resolve(code: u32, in_f: usize, out_f: usize) -> Option<(usize, usize)> {
if code == SEL_GROUP_OFF || in_f % 32 != 0 {
return None;
}
let pairs = in_f / 32;
if code != SEL_GROUP_AUTO {
let (g, rows) = (((code >> 8) & 0xff) as usize, (code & 0xff) as usize);
if !matches!(g, 1 | 2 | 4 | 8 | 16 | 32) || !matches!(rows, 1 | 2 | 4) {
return None;
}
if out_f % ((32 / g) * rows) != 0 {
return None;
}
return Some((g, rows));
}
let mut g = 1usize;
for cand in [2usize, 4, 8, 16, 32] {
if pairs % cand != 0 {
break;
}
g = cand;
}
let mut rows = 4usize;
while rows > 1 && out_f % ((32 / g) * rows) != 0 {
rows /= 2;
}
let rows_per_warp = (32 / g) * rows;
if out_f % rows_per_warp != 0 {
return None;
}
Some((g, rows))
}
static HC_MICRO: std::sync::atomic::AtomicBool = std::sync::atomic::AtomicBool::new(true);
pub fn set_hc_micro(on: bool) {
HC_MICRO.store(on, std::sync::atomic::Ordering::Relaxed);
}
fn hc_micro_on() -> bool {
HC_MICRO.load(std::sync::atomic::Ordering::Relaxed)
}
pub const GDN_STEP_DEFAULT: bool = true;
static GDN_STEP: std::sync::atomic::AtomicBool =
std::sync::atomic::AtomicBool::new(GDN_STEP_DEFAULT);
pub fn set_gdn_step(on: bool) {
GDN_STEP.store(on, std::sync::atomic::Ordering::Relaxed);
}
fn gdn_step_on() -> bool {
GDN_STEP.load(std::sync::atomic::Ordering::Relaxed)
}
pub const GDN_FUSE_DEFAULT: bool = true;
static GDN_FUSE: std::sync::atomic::AtomicBool =
std::sync::atomic::AtomicBool::new(GDN_FUSE_DEFAULT);
pub fn set_gdn_fuse(on: bool) {
GDN_FUSE.store(on, std::sync::atomic::Ordering::Relaxed);
}
fn gdn_fuse_on() -> bool {
GDN_FUSE.load(std::sync::atomic::Ordering::Relaxed)
}
pub const PROJ_STACK_DEFAULT: bool = true;
static PROJ_STACK: std::sync::atomic::AtomicBool =
std::sync::atomic::AtomicBool::new(PROJ_STACK_DEFAULT);
pub fn set_proj_stack(on: bool) {
PROJ_STACK.store(on, std::sync::atomic::Ordering::Relaxed);
}
fn proj_stack_on() -> bool {
PROJ_STACK.load(std::sync::atomic::Ordering::Relaxed)
}
pub const HC_DIET_DEFAULT: bool = true;
static HC_DIET: std::sync::atomic::AtomicBool = std::sync::atomic::AtomicBool::new(HC_DIET_DEFAULT);
pub fn set_hc_diet(on: bool) {
HC_DIET.store(on, std::sync::atomic::Ordering::Relaxed);
}
fn hc_diet_on() -> bool {
HC_DIET.load(std::sync::atomic::Ordering::Relaxed)
}
pub const SEL_GUFUSE_DEFAULT: bool = true;
static SEL_GUFUSE: std::sync::atomic::AtomicBool =
std::sync::atomic::AtomicBool::new(SEL_GUFUSE_DEFAULT);
pub fn set_sel_gufuse(on: bool) {
SEL_GUFUSE.store(on, std::sync::atomic::Ordering::Relaxed);
}
fn sel_gufuse_on() -> bool {
SEL_GUFUSE.load(std::sync::atomic::Ordering::Relaxed)
}
pub const VERIFY_MT_DEFAULT: bool = true;
static VERIFY_MT: std::sync::atomic::AtomicBool =
std::sync::atomic::AtomicBool::new(VERIFY_MT_DEFAULT);
pub fn set_verify_mt(on: bool) {
VERIFY_MT.store(on, std::sync::atomic::Ordering::Relaxed);
}
fn verify_mt_on() -> bool {
VERIFY_MT.load(std::sync::atomic::Ordering::Relaxed)
}
pub const VERIFY_FUSED_DEFAULT: bool = false;
static VERIFY_FUSED: std::sync::atomic::AtomicBool =
std::sync::atomic::AtomicBool::new(VERIFY_FUSED_DEFAULT);
pub fn set_verify_fused(on: bool) {
VERIFY_FUSED.store(on, std::sync::atomic::Ordering::Relaxed);
}
pub fn verify_fused_on() -> bool {
VERIFY_FUSED.load(std::sync::atomic::Ordering::Relaxed)
}
pub const ROUTER_B16_DEFAULT: bool = true;
static ROUTER_B16: std::sync::atomic::AtomicBool =
std::sync::atomic::AtomicBool::new(ROUTER_B16_DEFAULT);
pub fn set_router_bf16(on: bool) {
ROUTER_B16.store(on, std::sync::atomic::Ordering::Relaxed);
}
fn router_bf16_on() -> bool {
ROUTER_B16.load(std::sync::atomic::Ordering::Relaxed)
}
const SDPA_MASK_TKV_BOUND: usize = 12288;
static IDX_DEV: std::sync::atomic::AtomicBool = std::sync::atomic::AtomicBool::new(true);
fn idx_dev_on() -> bool {
IDX_DEV.load(std::sync::atomic::Ordering::Relaxed)
}
pub fn set_idx_dev(on: bool) {
IDX_DEV.store(on, std::sync::atomic::Ordering::Relaxed);
}
pub const IDX_SEL_DEFAULT: bool = true;
static IDX_SEL: std::sync::atomic::AtomicBool = std::sync::atomic::AtomicBool::new(IDX_SEL_DEFAULT);
fn idx_sel_on() -> bool {
IDX_SEL.load(std::sync::atomic::Ordering::Relaxed)
}
pub fn set_idx_sel(on: bool) {
IDX_SEL.store(on, std::sync::atomic::Ordering::Relaxed);
}
pub const PLE_CACHE_DEFAULT: bool = true;
static PLE_CACHE: std::sync::atomic::AtomicBool =
std::sync::atomic::AtomicBool::new(PLE_CACHE_DEFAULT);
fn ple_cache_on() -> bool {
PLE_CACHE.load(std::sync::atomic::Ordering::Relaxed)
}
pub fn set_ple_cache(on: bool) {
PLE_CACHE.store(on, std::sync::atomic::Ordering::Relaxed);
}
fn ple_cache_audit_on() -> bool {
static C: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
*C.get_or_init(|| std::env::var("MEMRA_Q4E_PLECACHE_AUDIT").as_deref() == Ok("1"))
}
static PLE_CACHE_AUDIT_ROWS: std::sync::atomic::AtomicU64 = std::sync::atomic::AtomicU64::new(0);
static PLE_CACHE_AUDIT_MISMATCH: std::sync::atomic::AtomicU64 =
std::sync::atomic::AtomicU64::new(0);
static PLE_CACHE_AUDIT_MAX_FILL: std::sync::atomic::AtomicU64 =
std::sync::atomic::AtomicU64::new(0);
pub fn ple_cache_audit_stats() -> (u64, u64, u64) {
(
PLE_CACHE_AUDIT_ROWS.load(std::sync::atomic::Ordering::Relaxed),
PLE_CACHE_AUDIT_MISMATCH.load(std::sync::atomic::Ordering::Relaxed),
PLE_CACHE_AUDIT_MAX_FILL.load(std::sync::atomic::Ordering::Relaxed),
)
}
fn idx_sel_audit_on() -> bool {
static C: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
*C.get_or_init(|| std::env::var("MEMRA_Q4E_IDXSEL_AUDIT").as_deref() == Ok("1"))
}
static IDX_SEL_AUDIT_ROWS: std::sync::atomic::AtomicU64 = std::sync::atomic::AtomicU64::new(0);
static IDX_SEL_AUDIT_MISMATCH: std::sync::atomic::AtomicU64 = std::sync::atomic::AtomicU64::new(0);
static IDX_SEL_AUDIT_MAX_BLOCKS: std::sync::atomic::AtomicU64 =
std::sync::atomic::AtomicU64::new(0);
pub fn idx_sel_audit_stats() -> (u64, u64, u64) {
(
IDX_SEL_AUDIT_ROWS.load(std::sync::atomic::Ordering::Relaxed),
IDX_SEL_AUDIT_MISMATCH.load(std::sync::atomic::Ordering::Relaxed),
IDX_SEL_AUDIT_MAX_BLOCKS.load(std::sync::atomic::Ordering::Relaxed),
)
}
pub const ROUTER_DEV_DEFAULT: bool = true;
static ROUTER_DEV: std::sync::atomic::AtomicBool =
std::sync::atomic::AtomicBool::new(ROUTER_DEV_DEFAULT);
fn router_dev_on() -> bool {
ROUTER_DEV.load(std::sync::atomic::Ordering::Relaxed)
}
pub fn set_router_dev(on: bool) {
ROUTER_DEV.store(on, std::sync::atomic::Ordering::Relaxed);
}
fn route_dev_geometry(experts: usize, selected: usize) -> bool {
selected > 0
&& selected <= 32
&& selected <= experts
&& experts % 2 == 0
&& experts * 12 <= 48 * 1024
}
fn router_audit_on() -> bool {
static C: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
*C.get_or_init(|| std::env::var("MEMRA_Q4E_ROUTER_AUDIT").as_deref() == Ok("1"))
}
fn route_sync_diag() -> bool {
static C: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
*C.get_or_init(|| std::env::var("MEMRA_Q4E_ROUTE_SYNC").as_deref() == Ok("1"))
}
pub fn peer_kv_max_cap() -> usize {
static C: std::sync::OnceLock<usize> = std::sync::OnceLock::new();
*C.get_or_init(|| {
std::env::var("MEMRA_Q4E_PEER_KV_MAX_CAP")
.ok()
.and_then(|v| v.trim().parse::<usize>().ok())
.unwrap_or(8192)
})
}
const ROUTE_AUDIT_ULP_BOUND: u32 = 8;
static ROUTE_AUDIT_ROWS: std::sync::atomic::AtomicU64 = std::sync::atomic::AtomicU64::new(0);
static ROUTE_AUDIT_MAX_ULP: std::sync::atomic::AtomicU32 = std::sync::atomic::AtomicU32::new(0);
pub fn route_audit_stats() -> (u64, u32) {
(
ROUTE_AUDIT_ROWS.load(std::sync::atomic::Ordering::Relaxed),
ROUTE_AUDIT_MAX_ULP.load(std::sync::atomic::Ordering::Relaxed),
)
}
pub const IDX_CACHE_DEFAULT: bool = true;
static IDX_CACHE: std::sync::atomic::AtomicBool =
std::sync::atomic::AtomicBool::new(IDX_CACHE_DEFAULT);
fn idx_cache_on() -> bool {
IDX_CACHE.load(std::sync::atomic::Ordering::Relaxed)
}
pub fn set_idx_cache(on: bool) {
IDX_CACHE.store(on, std::sync::atomic::Ordering::Relaxed);
}
pub const KV_QUANT_DEFAULT: bool = true;
static KV_QUANT: std::sync::atomic::AtomicBool =
std::sync::atomic::AtomicBool::new(KV_QUANT_DEFAULT);
fn kv_quant_on() -> bool {
KV_QUANT.load(std::sync::atomic::Ordering::Relaxed)
}
pub fn set_kv_quant(on: bool) {
KV_QUANT.store(on, std::sync::atomic::Ordering::Relaxed);
}
pub const KV_HOIST_DEFAULT: bool = false;
static KV_HOIST: std::sync::atomic::AtomicBool =
std::sync::atomic::AtomicBool::new(KV_HOIST_DEFAULT);
fn kv_hoist_on() -> bool {
KV_HOIST.load(std::sync::atomic::Ordering::Relaxed)
}
pub fn set_kv_hoist(on: bool) {
KV_HOIST.store(on, std::sync::atomic::Ordering::Relaxed);
}
pub fn kv_hoist_is_on() -> bool {
kv_hoist_on()
}
pub const POOL_T_DEFAULT: bool = false;
static POOL_T: std::sync::atomic::AtomicBool = std::sync::atomic::AtomicBool::new(POOL_T_DEFAULT);
fn pool_t_on() -> bool {
POOL_T.load(std::sync::atomic::Ordering::Relaxed)
}
pub fn set_pool_t(on: bool) {
POOL_T.store(on, std::sync::atomic::Ordering::Relaxed);
}
pub fn pool_t_is_on() -> bool {
pool_t_on()
}
pub fn kv_quant_is_on() -> bool {
kv_quant_on()
}
#[derive(Clone, Copy, PartialEq, Eq, Debug)]
pub enum IdxQMode {
F32,
Q8,
Bf16,
}
static IDXQ_MODE: std::sync::atomic::AtomicU8 = std::sync::atomic::AtomicU8::new(1);
fn idxq_mode() -> IdxQMode {
match IDXQ_MODE.load(std::sync::atomic::Ordering::Relaxed) {
1 => IdxQMode::Q8,
2 => IdxQMode::Bf16,
_ => IdxQMode::F32,
}
}
pub fn set_idxq(mode: &str) {
let v = match mode {
"q8" | "1" => 1,
"bf16" => 2,
_ => 0,
};
IDXQ_MODE.store(v, std::sync::atomic::Ordering::Relaxed);
}
pub fn idxq_mode_name() -> &'static str {
match idxq_mode() {
IdxQMode::F32 => "f32",
IdxQMode::Q8 => "q8",
IdxQMode::Bf16 => "bf16",
}
}
fn idxq_audit_on() -> bool {
static C: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
*C.get_or_init(|| std::env::var("MEMRA_Q4E_IDXQ_AUDIT").as_deref() == Ok("1"))
}
static IDXQ_AUDIT_ROWS: std::sync::atomic::AtomicU64 = std::sync::atomic::AtomicU64::new(0);
static IDXQ_AUDIT_FLIPPED: std::sync::atomic::AtomicU64 = std::sync::atomic::AtomicU64::new(0);
static IDXQ_AUDIT_BLOCKS: std::sync::atomic::AtomicU64 = std::sync::atomic::AtomicU64::new(0);
pub fn idxq_audit_stats() -> (u64, u64, u64) {
(
IDXQ_AUDIT_ROWS.load(std::sync::atomic::Ordering::Relaxed),
IDXQ_AUDIT_FLIPPED.load(std::sync::atomic::Ordering::Relaxed),
IDXQ_AUDIT_BLOCKS.load(std::sync::atomic::Ordering::Relaxed),
)
}
#[derive(Clone, Copy, PartialEq, Eq)]
enum LongAttMode {
Auto,
Force,
Off,
}
static LONGATT_MODE: std::sync::atomic::AtomicU8 = std::sync::atomic::AtomicU8::new(0);
fn longatt_mode() -> LongAttMode {
match LONGATT_MODE.load(std::sync::atomic::Ordering::Relaxed) {
1 => LongAttMode::Force,
2 => LongAttMode::Off,
_ => LongAttMode::Auto,
}
}
pub fn set_longatt(mode: &str) {
let v = match mode {
"force" | "1" => 1,
"off" | "0" => 2,
_ => 0,
};
LONGATT_MODE.store(v, std::sync::atomic::Ordering::Relaxed);
}
#[derive(Debug, Clone)]
pub struct Tp2Placement {
by_layer: std::collections::BTreeMap<u32, Vec<u8>>,
expert_count: usize,
entry_rank: u8,
strategy: String,
source: String,
}
#[derive(Debug, Clone)]
pub struct LayerPlacement {
pub card1: Vec<u32>,
local_of: Vec<u32>,
rank_of: Vec<u8>,
}
impl LayerPlacement {
#[inline]
pub fn rank(&self, expert: usize) -> u8 {
self.rank_of[expert]
}
#[inline]
pub fn local(&self, expert: usize) -> usize {
self.local_of[expert] as usize
}
pub fn is_even(&self) -> bool {
let half = self.rank_of.len() / 2;
self.card1.len() == half
&& self
.card1
.iter()
.enumerate()
.all(|(i, &e)| e as usize == half + i)
}
}
impl Tp2Placement {
pub fn even(expert_count: usize) -> Self {
Self {
by_layer: std::collections::BTreeMap::new(),
expert_count,
entry_rank: 0,
strategy: "even".to_string(),
source: "built-in (MEMRA_Q4E_EP_MAP unset)".to_string(),
}
}
pub fn strategy(&self) -> &str {
&self.strategy
}
pub fn source(&self) -> &str {
&self.source
}
pub fn entry_rank(&self) -> u8 {
self.entry_rank
}
pub fn from_env(expert_count: usize) -> Res<Option<Self>> {
let Ok(path) = std::env::var("MEMRA_Q4E_EP_MAP") else {
return Ok(None);
};
if path.is_empty() || path == "0" {
return Ok(None);
}
Some(Self::load(std::path::Path::new(&path), expert_count)).transpose()
}
pub fn load(path: &std::path::Path, expert_count: usize) -> Res<Self> {
let text = std::fs::read_to_string(path)
.map_err(|e| format!("MEMRA_Q4E_EP_MAP {}: {e}", path.display()))?;
let v = memra_tokenizer::json::parse(&text)
.map_err(|e| format!("MEMRA_Q4E_EP_MAP {}: {e}", path.display()))?;
let want = |k: &str| -> Res<Self> {
Err(format!("MEMRA_Q4E_EP_MAP {}: {k}", path.display()).into())
};
match v.get("format").and_then(|f| f.as_str()) {
Some("memra-ep-map-v1") => {}
other => {
return want(&format!(
"format is {other:?}, expected \"memra-ep-map-v1\" (mint it with \
tools/build_expert_placement_map.py)"
));
}
}
let ranks = v.get("ranks").and_then(|r| r.as_u64()).unwrap_or(0);
if ranks != 2 {
return want(&format!(
"ranks={ranks}, but the TP2 route is exactly two cards"
));
}
let map_experts = v.get("expert_count").and_then(|r| r.as_u64()).unwrap_or(0) as usize;
if map_experts != expert_count {
return want(&format!(
"expert_count={map_experts} but this plan has {expert_count} experts"
));
}
if expert_count % 2 != 0 {
return want(&format!(
"this plan has {expert_count} routed experts, which is ODD: the TP2 route \
splits the bank into two EQUAL-size device allocations, so no two-card \
placement exists for it"
));
}
let entry_rank = v.get("entry_rank").and_then(|r| r.as_u64()).unwrap_or(0) as u8;
if entry_rank > 1 {
return want(&format!("entry_rank={entry_rank} outside {{0,1}}"));
}
let strategy = v
.get("strategy")
.and_then(|s| s.as_str())
.unwrap_or("unnamed")
.to_string();
let Some(layers) = v.get("layers").and_then(|l| l.as_arr()) else {
return want("no `layers` array");
};
let half = expert_count / 2;
let mut by_layer = std::collections::BTreeMap::new();
for row in layers {
let Some(index) = row.get("layer").and_then(|l| l.as_u64()) else {
return want("a layer row without an integer `layer`");
};
let Some(assign) = row.get("assignment").and_then(|a| a.as_arr()) else {
return want(&format!("layer {index}: no `assignment` array"));
};
if assign.len() != expert_count {
return want(&format!(
"layer {index}: assignment has {} entries, expected {expert_count}",
assign.len()
));
}
let mut ranks_vec = Vec::with_capacity(expert_count);
for (eid, a) in assign.iter().enumerate() {
match a.as_u64() {
Some(r) if r <= 1 => ranks_vec.push(r as u8),
other => {
return want(&format!(
"layer {index} expert {eid}: rank {other:?} outside {{0,1}}"
));
}
}
}
let on1 = ranks_vec.iter().filter(|&&r| r == 1).count();
if on1 != half {
return want(&format!(
"layer {index}: card 1 owns {on1} experts but the bank halves are \
equal-size allocations, so it must own exactly {half} — rebalance \
the map (build_expert_placement_map.py --balance-tolerance)"
));
}
by_layer.insert(index as u32, ranks_vec);
}
if by_layer.is_empty() {
return want("`layers` is empty");
}
Ok(Self {
by_layer,
expert_count,
entry_rank,
strategy,
source: path.display().to_string(),
})
}
pub fn layer(&self, index: u32, expert_count: usize) -> Res<LayerPlacement> {
if expert_count != self.expert_count {
return Err(format!(
"qwen4exp_gpu tp2 placement: layer {index} has {expert_count} experts, \
map is for {}",
self.expert_count
)
.into());
}
if expert_count % 2 != 0 {
return Err(format!(
"qwen4exp_gpu tp2 placement: layer {index} has {expert_count} routed \
experts, which is ODD: the TP2 route splits the bank into two EQUAL-size \
device allocations, so no two-card placement exists for it"
)
.into());
}
let half = expert_count / 2;
let rank_of: Vec<u8> = if self.by_layer.is_empty() {
(0..expert_count).map(|e| u8::from(e >= half)).collect()
} else {
self.by_layer
.get(&index)
.ok_or_else(|| {
format!(
"qwen4exp_gpu tp2 placement: map {} does not cover MoE layer \
{index} (fail-closed; a partly-applied map is not a placement)",
self.source
)
})?
.clone()
};
let card1: Vec<u32> = (0..expert_count)
.filter(|&e| rank_of[e] == 1)
.map(|e| e as u32)
.collect();
let mut local_of = vec![0u32; expert_count];
for (slot, &eid) in card1.iter().enumerate() {
local_of[eid as usize] = slot as u32;
}
for e in 0..expert_count {
if rank_of[e] == 0 {
local_of[e] = e as u32;
}
}
Ok(LayerPlacement {
card1,
local_of,
rank_of,
})
}
}
static TP2_PEER_EXPERT_SLOTS: std::sync::atomic::AtomicU64 = std::sync::atomic::AtomicU64::new(0);
static TP2_HOME_EXPERT_SLOTS: std::sync::atomic::AtomicU64 = std::sync::atomic::AtomicU64::new(0);
static TP2_BOTH_TOUCH_ROWS: std::sync::atomic::AtomicU64 = std::sync::atomic::AtomicU64::new(0);
static TP2_TOUCH_ROWS: std::sync::atomic::AtomicU64 = std::sync::atomic::AtomicU64::new(0);
pub fn tp2_expert_split_stats() -> (u64, u64, u64, u64) {
use std::sync::atomic::Ordering::Relaxed;
(
TP2_PEER_EXPERT_SLOTS.load(Relaxed),
TP2_HOME_EXPERT_SLOTS.load(Relaxed),
TP2_BOTH_TOUCH_ROWS.load(Relaxed),
TP2_TOUCH_ROWS.load(Relaxed),
)
}
fn tp2_count_split(routes0: &[Vec<(usize, f32)>], routes1: &[Vec<(usize, f32)>]) {
use std::sync::atomic::Ordering::Relaxed;
let (mut peer, mut home, mut both) = (0u64, 0u64, 0u64);
for (r0, r1) in routes0.iter().zip(routes1.iter()) {
home += r0.len() as u64;
peer += r1.len() as u64;
if !r0.is_empty() && !r1.is_empty() {
both += 1;
}
}
TP2_HOME_EXPERT_SLOTS.fetch_add(home, Relaxed);
TP2_PEER_EXPERT_SLOTS.fetch_add(peer, Relaxed);
TP2_BOTH_TOUCH_ROWS.fetch_add(both, Relaxed);
TP2_TOUCH_ROWS.fetch_add(routes0.len() as u64, Relaxed);
}
#[derive(Clone, Copy, PartialEq, Eq, Debug)]
pub enum Tp2GateRed {
None,
SkipPeerMoe,
PeerLocalIds,
ReverseePeerWeights,
}
fn tp2_gate_red() -> Res<Tp2GateRed> {
static C: std::sync::OnceLock<Result<Tp2GateRed, String>> = std::sync::OnceLock::new();
C.get_or_init(
|| match std::env::var("MEMRA_Q4E_TP2_GATE_RED").as_deref() {
Err(_) | Ok("") | Ok("0") | Ok("none") => Ok(Tp2GateRed::None),
Ok("skip-peer-moe") => Ok(Tp2GateRed::SkipPeerMoe),
Ok("peer-local-ids") => Ok(Tp2GateRed::PeerLocalIds),
Ok("reverse-peer-weights") => Ok(Tp2GateRed::ReverseePeerWeights),
Ok(other) => Err(format!(
"MEMRA_Q4E_TP2_GATE_RED={other:?}: want skip-peer-moe|peer-local-ids|\
reverse-peer-weights|none"
)),
},
)
.clone()
.map_err(Into::into)
}
fn trace_moe_routes(layer: u32, t: usize, routes: &[Vec<(usize, f32)>]) {
use std::io::Write as _;
static IDS: std::sync::OnceLock<Option<String>> = std::sync::OnceLock::new();
static WEIGHTS: std::sync::OnceLock<Option<String>> = std::sync::OnceLock::new();
let ids = IDS.get_or_init(|| {
std::env::var("MEMRA_MOE_TRACE")
.ok()
.filter(|p| !p.is_empty())
});
let weights = WEIGHTS.get_or_init(|| {
std::env::var("MEMRA_MOE_WEIGHT_TRACE")
.ok()
.filter(|p| !p.is_empty())
});
if ids.is_none() && weights.is_none() {
return;
}
let flat: Vec<&(usize, f32)> = routes.iter().flatten().collect();
let append = |path: &str, body: String| {
if let Ok(mut f) = std::fs::OpenOptions::new()
.create(true)
.append(true)
.open(path)
{
let _ = writeln!(f, "{layer} {t} {body}");
}
};
if let Some(path) = ids {
let body: Vec<String> = flat.iter().map(|(e, _)| e.to_string()).collect();
append(path, body.join(","));
}
if let Some(path) = weights {
let body: Vec<String> = flat.iter().map(|(e, w)| format!("{e}:{w:.9}")).collect();
append(path, body.join(","));
}
}
pub fn set_seam(name: &str, on: bool, value: Option<&str>) -> bool {
seam_dispatch(name, on, value, true)
}
pub fn seam_state(name: &str) -> Option<bool> {
Some(match name {
"moe" => moe_sel_path_on(),
"hc" => hc_fused_gate_on(),
"trunk" => trunk_bf16_on(),
"ws" => step_ws_on(),
"graph" => decode_graphs_on(),
"selv2" => sel_v2_on(),
"hcmicro" => hc_micro_on(),
"selv3" => sel_v3_on(),
"gdnstep" => gdn_step_on(),
"gdnfuse" => gdn_fuse_on(),
"projstack" => proj_stack_on(),
"hcdiet" => hc_diet_on(),
"gufuse" => sel_gufuse_on(),
"routerb16" => router_bf16_on(),
"vgraph" => verify_graphs_on(),
"vfuse" => verify_fused_on(),
"idxdev" => idx_dev_on(),
"idxsel" => idx_sel_on(),
"plecache" => ple_cache_on(),
"routerdev" => router_dev_on(),
"idxcache" => idx_cache_on(),
"kvq" => kv_quant_on(),
"kvhoist" => kv_hoist_on(),
"poolT" => pool_t_on(),
"selgroup" => sel_group_dn() != SEL_GROUP_OFF || sel_group_gu() != SEL_GROUP_OFF,
_ => return None,
})
}
pub fn seam_exists(name: &str) -> bool {
seam_dispatch(name, false, None, false)
}
pub fn seam_names() -> &'static [&'static str] {
&[
"moe",
"hc",
"trunk",
"ws",
"graph",
"selv2",
"hcmicro",
"selv3",
"gdnstep",
"gdnfuse",
"projstack",
"hcdiet",
"gufuse",
"routerb16",
"vgraph",
"vfuse",
"longatt",
"idxdev",
"idxsel",
"plecache",
"routerdev",
"idxcache",
"kvq",
"idxq",
"kvhoist",
"poolT",
"selgroup",
]
}
fn seam_dispatch(name: &str, on: bool, value: Option<&str>, apply: bool) -> bool {
macro_rules! seam {
($call:expr) => {{
if apply {
$call;
}
true
}};
}
match name {
"moe" => seam!(set_moe_sel_path(on)),
"hc" => seam!(set_hc_fused_gate(on)),
"trunk" => seam!(set_trunk_bf16(on)),
"ws" => seam!(set_step_ws(on)),
"graph" => seam!(set_decode_graphs(on)),
"selv2" => seam!(set_sel_v2(on)),
"hcmicro" => seam!(set_hc_micro(on)),
"selv3" => seam!(set_sel_v3(on)),
"gdnstep" => seam!(set_gdn_step(on)),
"gdnfuse" => seam!(set_gdn_fuse(on)),
"projstack" => seam!(set_proj_stack(on)),
"hcdiet" => seam!(set_hc_diet(on)),
"gufuse" => seam!(set_sel_gufuse(on)),
"routerb16" => seam!(set_router_bf16(on)),
"vgraph" => seam!(set_verify_graphs(on)),
"vfuse" => seam!(set_verify_fused(on)),
"longatt" => seam!(set_longatt(if on { "force" } else { "off" })),
"idxdev" => seam!(set_idx_dev(on)),
"idxsel" => seam!(set_idx_sel(on)),
"plecache" => seam!(set_ple_cache(on)),
"routerdev" => seam!(set_router_dev(on)),
"idxcache" => seam!(set_idx_cache(on)),
"kvq" => seam!(set_kv_quant(on)),
"kvhoist" => seam!(set_kv_hoist(on)),
"poolT" => seam!(set_pool_t(on)),
"idxq" => seam!(set_idxq(value.unwrap_or("q8"))),
"selgroup" => {
if apply {
set_sel_group(if on { value.unwrap_or("auto") } else { "off" })
} else {
true
}
}
_ => {
debug_assert!(
!seam_names().contains(&name),
"seam_names() lists {name:?} but seam_dispatch has no arm for it"
);
false
}
}
}
pub fn apply_env_seams() {
let Ok(spec) = std::env::var("MEMRA_Q4E_SEAMS") else {
return;
};
for part in spec.split(',').filter(|p| !p.is_empty()) {
let (name, on) = match part.split_once('=') {
Some((n, v)) => (n, v != "0"),
None => (part, true),
};
if !set_seam(name, on, part.split_once('=').map(|(_, v)| v)) {
eprintln!("MEMRA_Q4E_SEAMS: unknown seam {name:?} ignored");
}
}
}
fn micro_env_on(name: &'static str, cell: &'static std::sync::OnceLock<bool>) -> bool {
*cell.get_or_init(|| std::env::var(name).as_deref() != Ok("0"))
}
fn micro_norm_on() -> bool {
static C: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
hc_micro_on() && micro_env_on("MEMRA_Q4E_MICRO_NORM", &C)
}
fn micro_inj_on() -> bool {
static C: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
hc_micro_on() && micro_env_on("MEMRA_Q4E_MICRO_INJ", &C)
}
fn micro_shexp_on() -> bool {
static C: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
hc_micro_on() && micro_env_on("MEMRA_Q4E_MICRO_SHEXP", &C)
}
fn prof_section<T>(e: &Engine, name: &'static str, f: impl FnOnce() -> Res<T>) -> Res<T> {
if !prof::on() {
return f();
}
e.gpu.stream().synchronize()?;
let t0 = std::time::Instant::now();
let out = f()?;
e.gpu.stream().synchronize()?;
prof::add(name, t0.elapsed().as_secs_f64());
Ok(out)
}
fn host_sigmoid(x: f32) -> f32 {
1.0 / (1.0 + (-x).exp())
}
fn host_softmax(values: &mut [f32]) {
let max = values.iter().copied().fold(f32::NEG_INFINITY, f32::max);
let mut sum = 0.0;
for value in values.iter_mut() {
*value = (*value - max).exp();
sum += *value;
}
for value in values {
*value /= sum;
}
}
const ROUTE_DENOM_FLOOR: f32 = 6.103_515_6e-5;
fn host_route_softmax_topk(logits: &[f32], selected: usize) -> Vec<(usize, f32)> {
let mut weights = logits.to_vec();
host_softmax(&mut weights);
let mut indices: Vec<usize> = (0..logits.len()).collect();
indices.sort_by(|&left, &right| {
weights[right]
.total_cmp(&weights[left])
.then(left.cmp(&right))
});
indices.truncate(selected);
let denominator = indices
.iter()
.map(|&index| weights[index])
.sum::<f32>()
.max(ROUTE_DENOM_FLOOR);
indices
.into_iter()
.map(|index| (index, weights[index] / denominator))
.collect()
}
fn host_rms_norm(x: &mut [f32], width: usize, weight: &[f32], epsilon: f32) {
for row in x.chunks_exact_mut(width) {
let mean_square = row.iter().map(|v| v * v).sum::<f32>() / width as f32;
let inverse = 1.0 / (mean_square + epsilon).sqrt();
for (value, w) in row.iter_mut().zip(weight) {
*value = *value * inverse * w;
}
}
}
fn host_rope_at(
values: &mut [f32],
head_dim: usize,
dimensions: usize,
base: f32,
yarn: Option<(&[f32], f32)>,
position: usize,
) {
let dimensions = dimensions.min(head_dim) / 2 * 2;
let half = dimensions / 2;
for head in values.chunks_exact_mut(head_dim) {
for index in 0..half {
let frequency = base.powf(-2.0 * index as f32 / dimensions as f32);
let frequency = match yarn {
Some((ff, _)) => frequency / ff[index],
None => frequency,
};
let angle = position as f32 * frequency;
let (sin, cos) = angle.sin_cos();
let (sin, cos) = match yarn {
Some((_, mscale)) => (sin * mscale, cos * mscale),
None => (sin, cos),
};
let first = head[index];
let second = head[index + half];
head[index] = first * cos - second * sin;
head[index + half] = first * sin + second * cos;
}
}
}
#[derive(Clone, Copy, PartialEq, Eq)]
pub enum HeadMode {
All,
LastRow,
Skip,
}
struct RowSel {
full: bool,
blocks: Vec<u32>,
visible: usize,
}
#[allow(clippy::too_many_arguments)]
fn extend_pooled_keys(
pooled_keys: &mut Vec<f32>,
raw_keys: &IdxRawCache,
head_dim: usize,
block_size: usize,
idx_k_norm: &[f32],
epsilon: f32,
rope_dims: usize,
rope_base: f32,
yarn: Option<(&[f32], f32)>,
pos_off: usize,
) {
let complete_total = raw_keys.rows(head_dim) / block_size;
let cached = pooled_keys.len() / head_dim;
let mut block_rows: Vec<f32> = Vec::new();
for block in cached..complete_total {
let start = block * block_size;
raw_keys.rows_f32(start, block_size, head_dim, &mut block_rows);
let mut pooled = vec![0.0f32; head_dim];
for offset in 0..block_size {
for dim in 0..head_dim {
pooled[dim] += block_rows[offset * head_dim + dim];
}
}
for value in &mut pooled {
*value /= block_size as f32;
}
host_rms_norm(&mut pooled, head_dim, idx_k_norm, epsilon);
host_rope_at(
&mut pooled,
head_dim,
rope_dims,
rope_base,
yarn,
start + pos_off,
);
pooled_keys.extend_from_slice(&pooled);
}
}
#[inline]
fn sel_cmp(scores: &[f32], a: u32, b: u32) -> std::cmp::Ordering {
scores[b as usize]
.total_cmp(&scores[a as usize])
.then(a.cmp(&b))
}
fn top_blocks_ascending(scores: &[f32], budget: usize, threads: usize) -> Vec<u32> {
fn cut(scores: &[f32], idx: &mut Vec<u32>, budget: usize) {
let k = budget.min(idx.len());
if k < idx.len() {
idx.select_nth_unstable_by(k - 1, |&a, &b| sel_cmp(scores, a, b));
idx.truncate(k);
}
}
let complete = scores.len();
debug_assert!(budget < complete);
const PAR_MIN: usize = 1 << 15;
let mut candidates: Vec<u32> = if threads > 1 && complete >= PAR_MIN {
let ranges: Vec<(u32, u32)> = {
let per = complete.div_ceil(threads);
(0..threads)
.map(|i| ((i * per) as u32, ((i + 1) * per).min(complete) as u32))
.filter(|(a, b)| a < b)
.collect()
};
std::thread::scope(|scope| {
let handles: Vec<_> = ranges
.iter()
.map(|&(a, b)| {
scope.spawn(move || {
let mut idx: Vec<u32> = (a..b).collect();
cut(scores, &mut idx, budget);
idx
})
})
.collect();
handles
.into_iter()
.flat_map(|h| h.join().unwrap())
.collect()
})
} else {
(0..complete as u32).collect()
};
cut(scores, &mut candidates, budget);
candidates.sort_unstable();
candidates
}
fn score_blocks(
query: &[f32],
pooled_keys: &[f32],
heads: usize,
head_dim: usize,
complete: usize,
scale: f32,
threads: usize,
) -> Vec<f32> {
let mut scores = vec![0.0f32; complete];
let run = |scores: &mut [f32], block0: usize| {
for (i, slot) in scores.iter_mut().enumerate() {
let block = block0 + i;
let pooled = &pooled_keys[block * head_dim..(block + 1) * head_dim];
let mut score = 0.0f32;
for head in 0..heads {
let mut dot = 0.0f32;
for dim in 0..head_dim {
dot += query[head * head_dim + dim] * pooled[dim];
}
score += dot.max(0.0);
}
*slot = score / scale;
}
};
const PAR_MIN: usize = 1 << 14;
if threads > 1 && complete >= PAR_MIN {
let per = complete.div_ceil(threads);
let run = &run;
std::thread::scope(|scope| {
for (i, chunk) in scores.chunks_mut(per).enumerate() {
scope.spawn(move || run(chunk, i * per));
}
});
} else {
run(&mut scores, 0);
}
scores
}
#[allow(clippy::too_many_arguments)]
fn indexer_select_rows(
overlay: &MicroBlockIndexPlan,
rope_base: f32,
yarn: Option<(&[f32], f32)>,
epsilon: f32,
idx_q_norm: &[f32],
idx_k_norm: &[f32],
proj_rows: &[f32], raw_keys: &IdxRawCache, pooled_keys: &mut Vec<f32>,
mut dev: Option<(&Engine, &mut Option<CudaSlice<f32>>, &mut usize)>,
base_pos: usize,
t: usize,
t_kv: usize,
pos_off: usize,
) -> Res<Vec<RowSel>> {
let heads = overlay.query_heads as usize;
let head_dim = overlay.head_dim as usize;
let block_size = overlay.block_size as usize;
let budget_blocks = overlay.budget_blocks as usize;
let rope_dims = overlay.rope_dimensions as usize;
let qk_width = (heads + overlay.kv_heads as usize) * head_dim;
let scale = (head_dim as f32).sqrt();
debug_assert_eq!(raw_keys.rows(head_dim), t_kv);
extend_pooled_keys(
pooled_keys,
raw_keys,
head_dim,
block_size,
idx_k_norm,
epsilon,
rope_dims,
rope_base,
yarn,
pos_off,
);
let threads = std::thread::available_parallelism()
.map(|n| n.get())
.unwrap_or(1);
if let Some((e, mirror, mirrored)) = dev.as_mut() {
let rows_needed: Vec<usize> = (0..t)
.map(|qt| (base_pos + qt + 1) / block_size)
.filter(|&c| c > budget_blocks)
.collect();
if let Some(&max_blocks) = rows_needed.iter().max() {
let pooled_rows = pooled_keys.len() / head_dim;
let want = pooled_rows.max(max_blocks);
if mirror
.as_ref()
.is_none_or(|m| m.len() < want * head_dim * POOL_PLANES)
{
let cap_rows = want.next_power_of_two().max(1024);
let fresh = e.zeros(cap_rows * head_dim * POOL_PLANES)?;
**mirror = Some(fresh);
**mirrored = 0;
}
let m = mirror.as_mut().expect("allocated above");
if pooled_rows > **mirrored {
let delta = &pooled_keys[**mirrored * head_dim..pooled_rows * head_dim];
let mut view = m.slice_mut(**mirrored * head_dim..pooled_rows * head_dim);
e.gpu.stream().memcpy_htod(delta, &mut view)?;
let cap_rows = m.len() / (head_dim * POOL_PLANES);
launch_qsa_pooled_transpose(
e,
m,
**mirrored,
pooled_rows - **mirrored,
head_dim,
cap_rows,
)?;
**mirrored = pooled_rows;
}
let mut sels: Vec<RowSel> = Vec::with_capacity(t);
let mut queries: Vec<f32> = Vec::new();
let mut scored_rows: Vec<usize> = Vec::new();
for qt in 0..t {
let row = base_pos + qt;
let visible = row + 1;
let complete = visible / block_size;
if complete <= budget_blocks {
sels.push(RowSel {
full: true,
blocks: Vec::new(),
visible,
});
continue;
}
let mut query = proj_rows[qt * qk_width..qt * qk_width + heads * head_dim].to_vec();
host_rms_norm(&mut query, head_dim, idx_q_norm, epsilon);
host_rope_at(
&mut query,
head_dim,
rope_dims,
rope_base,
yarn,
row + pos_off,
);
queries.extend_from_slice(&query);
scored_rows.push(qt);
sels.push(RowSel {
full: false,
blocks: Vec::new(),
visible,
});
}
let score_cap_mf: usize = std::env::var("MEMRA_Q4E_IDX_SCORE_CAP_MF")
.ok()
.and_then(|v| v.parse::<usize>().ok())
.filter(|v| *v > 0)
.unwrap_or(32);
let score_cap: usize = score_cap_mf << 20;
#[allow(non_snake_case)]
let SCORE_CAP = score_cap;
let mut done = 0usize;
while done < scored_rows.len() {
let batch_max = scored_rows[done..]
.iter()
.map(|&qt| (base_pos + qt + 1) / block_size)
.max()
.unwrap_or(0);
let per = (SCORE_CAP / batch_max.max(1)).max(1);
let n = per.min(scored_rows.len() - done);
let qslab = &queries[done * heads * head_dim..(done + n) * heads * head_dim];
let q_dev = e.htod(qslab)?;
let mut scores_dev = e.uninit(n * batch_max)?;
launch_qsa_index_score(
e,
&q_dev,
m,
&mut scores_dev,
heads,
head_dim,
batch_max,
n,
scale,
)?;
if idx_sel_on() {
let counts: Vec<usize> = (0..n)
.map(|i| (base_pos + scored_rows[done + i] + 1) / block_size)
.collect();
let picked =
launch_qsa_index_topk(e, &scores_dev, &counts, batch_max, budget_blocks)?;
if idx_sel_audit_on() {
let host = e.dtoh(&scores_dev)?;
let mut mismatched = 0u64;
let mut deepest = 0u64;
for i in 0..n {
let complete = counts[i];
let row_scores = &host[i * batch_max..i * batch_max + complete];
let twin = top_blocks_ascending(row_scores, budget_blocks, threads);
if twin != picked[i] {
mismatched += 1;
}
deepest = deepest.max(complete as u64);
}
IDX_SEL_AUDIT_ROWS
.fetch_add(n as u64, std::sync::atomic::Ordering::Relaxed);
IDX_SEL_AUDIT_MISMATCH
.fetch_add(mismatched, std::sync::atomic::Ordering::Relaxed);
IDX_SEL_AUDIT_MAX_BLOCKS
.fetch_max(deepest, std::sync::atomic::Ordering::Relaxed);
if mismatched > 0 {
return Err(format!(
"idxsel audit: {mismatched} of {n} device selections differ \
from the host twin (ids or order) at fill {t_kv}"
)
.into());
}
}
for (i, blocks) in picked.into_iter().enumerate() {
sels[scored_rows[done + i]].blocks = blocks;
}
} else {
let host = e.dtoh(&scores_dev)?;
for i in 0..n {
let qt = scored_rows[done + i];
let complete = (base_pos + qt + 1) / block_size;
let row_scores = &host[i * batch_max..i * batch_max + complete];
sels[qt].blocks = top_blocks_ascending(row_scores, budget_blocks, threads);
}
}
done += n;
}
for sel in &sels {
if sel.visible == 0
|| (!sel.full && sel.blocks.is_empty() && sel.visible % block_size == 0)
{
return Err("indexer selection left a query with no visible source".into());
}
}
return Ok(sels);
}
}
let pooled_ref: &[f32] = pooled_keys;
let select_row = |qt: usize, threads_in_row: usize| -> RowSel {
let row = base_pos + qt;
let position = row + pos_off;
let visible = row + 1;
let complete = visible / block_size;
if complete <= budget_blocks {
return RowSel {
full: true,
blocks: Vec::new(),
visible,
};
}
let mut query = proj_rows[qt * qk_width..qt * qk_width + heads * head_dim].to_vec();
host_rms_norm(&mut query, head_dim, idx_q_norm, epsilon);
host_rope_at(&mut query, head_dim, rope_dims, rope_base, yarn, position);
let scores = score_blocks(
&query,
pooled_ref,
heads,
head_dim,
complete,
scale,
threads_in_row,
);
let blocks = top_blocks_ascending(&scores, budget_blocks, threads_in_row);
RowSel {
full: false,
blocks,
visible,
}
};
const ROW_PAR_MIN_WORK: usize = 1 << 16;
let total_scored_blocks: usize = (0..t)
.map(|qt| {
let complete = (base_pos + qt + 1) / block_size;
if complete <= budget_blocks {
0
} else {
complete
}
})
.sum();
let sels: Vec<RowSel> = if t > 1 && threads > 1 && total_scored_blocks >= ROW_PAR_MIN_WORK {
let cursor = std::sync::atomic::AtomicUsize::new(0);
let mut out: Vec<Option<RowSel>> = (0..t).map(|_| None).collect();
let slots = std::sync::Mutex::new(&mut out);
std::thread::scope(|scope| {
for _ in 0..threads.min(t) {
scope.spawn(|| {
loop {
let qt = cursor.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
if qt >= t {
break;
}
let sel = select_row(qt, 1);
slots.lock().unwrap()[qt] = Some(sel);
}
});
}
});
out.into_iter().map(|s| s.unwrap()).collect()
} else {
(0..t).map(|qt| select_row(qt, threads)).collect()
};
for sel in &sels {
if sel.visible == 0 || (!sel.full && sel.blocks.is_empty() && sel.visible % block_size == 0)
{
return Err("indexer selection left a query with no visible source".into());
}
}
Ok(sels)
}
fn rowsel_to_mask(sels: &[RowSel], block_size: usize, t_kv: usize) -> Vec<u8> {
let t = sels.len();
let mut mask = vec![0u8; t * t_kv];
for (qt, sel) in sels.iter().enumerate() {
let row = &mut mask[qt * t_kv..(qt + 1) * t_kv];
if sel.full {
for slot in row.iter_mut().take(sel.visible) {
*slot = 1;
}
continue;
}
for &block in &sel.blocks {
for offset in 0..block_size {
row[block as usize * block_size + offset] = 1;
}
}
let complete = sel.visible / block_size;
for slot in row.iter_mut().take(sel.visible).skip(complete * block_size) {
*slot = 1;
}
}
mask
}
fn rowsel_positions(sels: &[RowSel], block_size: usize) -> (Vec<i32>, Vec<i32>, usize) {
let mut flat: Vec<i32> = Vec::new();
let mut meta: Vec<i32> = Vec::with_capacity(sels.len() * 2);
let mut max_count = 0usize;
for sel in sels {
let start = flat.len();
if sel.full {
flat.extend(0..sel.visible as i32);
} else {
for &block in &sel.blocks {
let first = block as usize * block_size;
flat.extend(first as i32..(first + block_size) as i32);
}
let complete = sel.visible / block_size;
flat.extend((complete * block_size) as i32..sel.visible as i32);
}
let count = flat.len() - start;
max_count = max_count.max(count);
meta.push(start as i32);
meta.push(count as i32);
}
(flat, meta, max_count)
}
#[allow(clippy::too_many_arguments)]
fn launch_qsa_index_score(
e: &Engine,
q: &CudaSlice<f32>,
pooled: &CudaSlice<f32>,
out: &mut CudaSlice<f32>,
heads: usize,
head_dim: usize,
n_blocks: usize,
rows: usize,
scale: f32,
) -> Res<()> {
if rows == 0 || n_blocks == 0 {
return Ok(());
}
if out.len() < rows * n_blocks {
return Err("qsa_index_score_f32: score slab too short".into());
}
if rows > 65535 {
return Err("qsa_index_score_f32: rows exceed grid.y (caller sub-batches)".into());
}
let cap_rows = pooled.len() / (head_dim * POOL_PLANES);
let pool_t = pool_t_on();
if pool_t && cap_rows < n_blocks {
return Err("qsa_index_score_f32_t: pooled plane capacity below n_blocks".into());
}
let f = e.func(if pool_t {
"qsa_index_score_f32_t"
} else {
"qsa_index_score_f32"
});
const TPB: usize = 128;
let cfg = LaunchConfig {
grid_dim: (n_blocks.div_ceil(TPB) as u32, rows as u32, 1),
block_dim: (TPB as u32, 1, 1),
shared_mem_bytes: 0,
};
let (h, hd, nb, r) = (heads as i32, head_dim as i32, n_blocks as i32, rows as i32);
let pitch = cap_rows as i64;
let stream = e.gpu.stream();
if pool_t {
let plane = pooled.slice(cap_rows * head_dim..cap_rows * head_dim * POOL_PLANES);
let mut b = stream.launch_builder(&f);
b.arg(q)
.arg(&plane)
.arg(&mut *out)
.arg(&h)
.arg(&hd)
.arg(&nb)
.arg(&r)
.arg(&scale)
.arg(&pitch);
unsafe {
b.launch(cfg)?;
}
return Ok(());
}
let mut b = stream.launch_builder(&f);
b.arg(q)
.arg(pooled)
.arg(&mut *out)
.arg(&h)
.arg(&hd)
.arg(&nb)
.arg(&r)
.arg(&scale);
unsafe {
b.launch(cfg)?;
}
Ok(())
}
const POOL_PLANES: usize = 2;
fn launch_qsa_pooled_transpose(
e: &Engine,
buf: &mut CudaSlice<f32>,
r0: usize,
rows: usize,
head_dim: usize,
cap_rows: usize,
) -> Res<()> {
if rows == 0 {
return Ok(());
}
if r0 + rows > cap_rows {
return Err("qsa_pooled_transpose_f32: delta exceeds the plane capacity".into());
}
let f = e.func("qsa_pooled_transpose_f32");
const TPB: usize = 128;
let cfg = LaunchConfig {
grid_dim: (rows.div_ceil(TPB) as u32, head_dim as u32, 1),
block_dim: (TPB as u32, 1, 1),
shared_mem_bytes: 0,
};
let (r, hd, r0i) = (rows as i32, head_dim as i32, r0 as i32);
let cap = cap_rows as i64;
let stream = e.gpu.stream();
let mut b = stream.launch_builder(&f);
b.arg(buf).arg(&r).arg(&hd).arg(&r0i).arg(&cap);
unsafe {
b.launch(cfg)?;
}
Ok(())
}
fn launch_qsa_index_topk(
e: &Engine,
scores: &CudaSlice<f32>,
counts: &[usize],
stride: usize,
budget: usize,
) -> Res<Vec<Vec<u32>>> {
let rows = counts.len();
if rows == 0 || budget == 0 {
return Ok(Vec::new());
}
if rows > 65535 {
return Err("qsa_index_topk_u32: rows exceed grid.x (caller sub-batches)".into());
}
if scores.len() < rows * stride {
return Err("qsa_index_topk_u32: score slab too short".into());
}
for (r, &c) in counts.iter().enumerate() {
if c <= budget || c > stride {
return Err(format!(
"qsa_index_topk_u32: row {r} block count {c} outside (budget {budget}, \
stride {stride}] — the caller only routes scored rows here"
)
.into());
}
}
let counts_i32: Vec<i32> = counts.iter().map(|&c| c as i32).collect();
let counts_dev = e.htod_i32(&counts_i32)?;
let mut out = e.htod_i32(&vec![-1i32; rows * budget])?;
let f = e.func("qsa_index_topk_u32");
let cfg = LaunchConfig {
grid_dim: (rows as u32, 1, 1),
block_dim: (256, 1, 1),
shared_mem_bytes: 0,
};
let (st, bu, ro) = (stride as i32, budget as i32, rows as i32);
let stream = e.gpu.stream();
let mut b = stream.launch_builder(&f);
b.arg(scores)
.arg(&counts_dev)
.arg(&mut out)
.arg(&st)
.arg(&bu)
.arg(&ro);
unsafe {
b.launch(cfg)?;
}
let host = e.gpu.stream().clone_dtoh(&out)?;
e.gpu.stream().synchronize()?;
let mut out_rows: Vec<Vec<u32>> = Vec::with_capacity(rows);
for r in 0..rows {
let row = &host[r * budget..(r + 1) * budget];
let mut blocks: Vec<u32> = Vec::with_capacity(budget);
let mut prev: i64 = -1;
for (j, &v) in row.iter().enumerate() {
if v < 0 || (v as usize) >= counts[r] || (v as i64) <= prev {
return Err(format!(
"qsa_index_topk_u32: row {r} slot {j} = {v} is not a strictly ascending \
in-range block id (blocks {}, budget {budget})",
counts[r]
)
.into());
}
prev = v as i64;
blocks.push(v as u32);
}
out_rows.push(blocks);
}
Ok(out_rows)
}
#[allow(clippy::too_many_arguments)]
fn launch_copy_rows_col(
e: &Engine,
src: &CudaSlice<f32>,
dst: &mut CudaSlice<f32>,
rows: usize,
width: usize,
src_stride: usize,
src_col: usize,
dst_row: usize,
) -> Res<()> {
if rows == 0 {
return Ok(());
}
if src.len() < (rows - 1) * src_stride + src_col + width || dst.len() < (dst_row + rows) * width
{
return Err("copy_rows_col_f32: window out of range".into());
}
let f = e.func("copy_rows_col_f32");
let total = rows * width;
let cfg = LaunchConfig::for_num_elems(total as u32);
let (r, w) = (rows as i32, width as i32);
let (ss, sc, dr) = (src_stride as i64, src_col as i64, dst_row as i64);
let stream = e.gpu.stream();
let mut b = stream.launch_builder(&f);
b.arg(src)
.arg(&mut *dst)
.arg(&r)
.arg(&w)
.arg(&ss)
.arg(&sc)
.arg(&dr);
unsafe {
b.launch(cfg)?;
}
Ok(())
}
#[allow(clippy::too_many_arguments)]
fn launch_route_topk(
e: &Engine,
logits: &CudaSlice<f32>,
sel: &mut CudaSlice<i32>,
w: &mut CudaSlice<f32>,
tok: Option<(&mut CudaSlice<i32>, usize)>,
experts: usize,
selected: usize,
rows: usize,
) -> Res<()> {
if rows == 0 {
return Ok(());
}
if selected == 0 || selected > 32 || selected > experts {
return Err("qwen4exp_route_topk_f32: selected out of range (caller guards)".into());
}
if experts % 2 != 0 {
return Err("qwen4exp_route_topk_f32: odd expert count".into());
}
if logits.len() < rows * experts || sel.len() < rows * selected || w.len() < rows * selected {
return Err("qwen4exp_route_topk_f32: buffer too short".into());
}
let smem = experts * 12; if smem > 48 * 1024 {
return Err("qwen4exp_route_topk_f32: experts exceed the smem bound".into());
}
let stream = e.gpu.stream();
let (tok_raw, tok_base) = match tok {
Some((buf, base)) => {
if buf.len() < rows * selected {
return Err("qwen4exp_route_topk_f32: tok map too short".into());
}
(buf.device_ptr(&stream).0, base)
}
None => (0u64, 0usize),
};
let f = e.func("qwen4exp_route_topk_f32");
let cfg = LaunchConfig {
grid_dim: (rows as u32, 1, 1),
block_dim: (128, 1, 1),
shared_mem_bytes: smem as u32,
};
let (ex, se, ro, tb) = (
experts as i32,
selected as i32,
rows as i32,
tok_base as i32,
);
let floor = ROUTE_DENOM_FLOOR;
let mut b = stream.launch_builder(&f);
b.arg(logits)
.arg(&mut *sel)
.arg(&mut *w)
.arg(&tok_raw)
.arg(&ex)
.arg(&se)
.arg(&ro)
.arg(&tb)
.arg(&floor);
unsafe {
b.launch(cfg)?;
}
Ok(())
}
#[allow(clippy::too_many_arguments)]
fn route_topk_device(
e: &Engine,
logits: &CudaSlice<f32>,
sel: &mut CudaSlice<i32>,
w: &mut CudaSlice<f32>,
tok: Option<(&mut CudaSlice<i32>, usize)>,
experts: usize,
selected: usize,
rows: usize,
layer: u32,
) -> Res<()> {
launch_route_topk(e, logits, sel, w, tok, experts, selected, rows)?;
if !router_audit_on() {
return Ok(());
}
let k = selected.min(experts);
let lg = e.dtoh_view(&logits.slice(0..rows * experts))?;
let sel_h = e.gpu.stream().clone_dtoh(&sel.slice(0..rows * selected))?;
let w_h = e.gpu.stream().clone_dtoh(&w.slice(0..rows * selected))?;
{
let routes: Vec<Vec<(usize, f32)>> = (0..rows)
.map(|row| {
(0..selected)
.map(|j| {
(
sel_h[row * selected + j].max(0) as usize,
w_h[row * selected + j],
)
})
.collect()
})
.collect();
trace_moe_routes(layer, rows, &routes);
}
let mut worst: u32 = 0;
for row in 0..rows {
let twin = host_route_softmax_topk(&lg[row * experts..(row + 1) * experts], selected);
if twin.len() != k {
return Err("router audit: host twin emitted an unexpected selection width".into());
}
for (j, &(idx, wt)) in twin.iter().enumerate() {
let ds = sel_h[row * selected + j];
let dw = w_h[row * selected + j];
if ds != idx as i32 {
return Err(format!(
"router audit: selection mismatch at row {row} slot {j}: \
device {ds} vs host {idx} (host w {wt:e})"
)
.into());
}
let ulp = (dw.to_bits() as i64 - wt.to_bits() as i64).unsigned_abs();
let ulp = u32::try_from(ulp).unwrap_or(u32::MAX);
worst = worst.max(ulp);
if ulp > ROUTE_AUDIT_ULP_BOUND {
return Err(format!(
"router audit: weight ULP {ulp} > bound {ROUTE_AUDIT_ULP_BOUND} at \
row {row} slot {j}: device {dw:e} vs host {wt:e}"
)
.into());
}
}
}
ROUTE_AUDIT_ROWS.fetch_add(rows as u64, std::sync::atomic::Ordering::Relaxed);
ROUTE_AUDIT_MAX_ULP.fetch_max(worst, std::sync::atomic::Ordering::Relaxed);
Ok(())
}
#[allow(dead_code)]
#[allow(clippy::too_many_arguments)]
fn indexer_mask_rows(
overlay: &MicroBlockIndexPlan,
rope_base: f32,
yarn: Option<(&[f32], f32)>,
epsilon: f32,
idx_q_norm: &[f32],
idx_k_norm: &[f32],
proj_rows: &[f32],
raw_keys: &IdxRawCache,
pooled_keys: &mut Vec<f32>,
base_pos: usize,
t: usize,
t_kv: usize,
pos_off: usize,
) -> Res<Vec<u8>> {
let sels = indexer_select_rows(
overlay,
rope_base,
yarn,
epsilon,
idx_q_norm,
idx_k_norm,
proj_rows,
raw_keys,
pooled_keys,
None,
base_pos,
t,
t_kv,
pos_off,
)?;
Ok(rowsel_to_mask(&sels, overlay.block_size as usize, t_kv))
}
fn shift_right_ignore_eos(history: &[i64], shift: usize, eos: i64) -> Vec<i64> {
if shift == 0 {
return history.to_vec();
}
let mut last_eos_inclusive: i64 = -1;
let mut output = Vec::with_capacity(history.len());
for (position, &token) in history.iter().enumerate() {
let previous_eos = last_eos_inclusive;
if token == eos {
last_eos_inclusive = position as i64;
}
let segment_start = previous_eos + 1;
let position_in_segment = position as i64 - segment_start;
let source = position as i64 - shift as i64;
let valid = position_in_segment >= shift as i64 && source >= 0;
output.push(if valid { history[source as usize] } else { eos });
}
output
}
#[allow(clippy::too_many_arguments)]
fn host_ngram_ids_cached(
cache_ids: &mut Vec<i64>,
cache_history: &mut Vec<i64>,
cache_last_eos: &mut i64,
token_ids: &[u32],
multipliers: &[i64],
sizes: &[i64],
offsets: &[i64],
max_ngram: usize,
heads_per_ngram: usize,
eos_token_id: u32,
) {
let context = max_ngram - 1;
let eos = eos_token_id as i64;
let total_heads = (max_ngram - 1) * heads_per_ngram;
if cache_history.is_empty() {
cache_history.extend(std::iter::repeat_n(eos, context));
*cache_last_eos = context as i64 - 1; cache_ids.clear();
}
let cached_tokens = (cache_history.len() - context).min(cache_ids.len() / total_heads);
let mut keep = cached_tokens.min(token_ids.len());
for i in 0..keep {
if cache_history[context + i] != token_ids[i] as i64 {
keep = i;
break;
}
}
if keep < cached_tokens {
cache_history.truncate(context + keep);
cache_ids.truncate(keep * total_heads);
*cache_last_eos = cache_history
.iter()
.rposition(|&v| v == eos)
.map(|p| p as i64)
.unwrap_or(-1);
}
for &token in &token_ids[keep..] {
let position = cache_history.len();
let value = token as i64;
cache_history.push(value);
let previous_eos = *cache_last_eos;
if value == eos {
*cache_last_eos = position as i64;
}
let segment_start = previous_eos + 1;
let position_in_segment = position as i64 - segment_start;
let shifted_at = |shift: usize| -> i64 {
if shift == 0 {
return cache_history[position];
}
let source = position as i64 - shift as i64;
if position_in_segment >= shift as i64 && source >= 0 {
cache_history[source as usize]
} else {
eos
}
};
let mut row = vec![0i64; total_heads];
for ngram in 2..=max_ngram {
let head_start = (ngram - 2) * heads_per_ngram;
let mut mixed = shifted_at(0).wrapping_mul(multipliers[0]);
for shift in 1..ngram {
mixed ^= shifted_at(shift).wrapping_mul(multipliers[shift]);
}
for head in 0..heads_per_ngram {
let index = head_start + head;
row[index] = mixed.rem_euclid(sizes[index]) + offsets[index];
}
}
cache_ids.extend_from_slice(&row);
}
debug_assert_eq!(cache_ids.len(), token_ids.len() * total_heads);
}
fn host_ngram_ids(
token_ids: &[u32],
multipliers: &[i64],
sizes: &[i64],
offsets: &[i64],
max_ngram: usize,
heads_per_ngram: usize,
eos_token_id: u32,
) -> Vec<i64> {
let context = max_ngram - 1;
let eos = eos_token_id as i64;
let total_heads = (max_ngram - 1) * heads_per_ngram;
let mut history = Vec::with_capacity(context + token_ids.len());
history.extend(std::iter::repeat_n(eos, context));
history.extend(token_ids.iter().map(|&token| token as i64));
let shifted: Vec<Vec<i64>> = (0..max_ngram)
.map(|shift| shift_right_ignore_eos(&history, shift, eos))
.collect();
let tokens = token_ids.len();
let mut ids = vec![0i64; tokens * total_heads];
for ngram in 2..=max_ngram {
let head_start = (ngram - 2) * heads_per_ngram;
for token in 0..tokens {
let position = context + token;
let mut mixed = shifted[0][position].wrapping_mul(multipliers[0]);
for (shift, row) in shifted.iter().enumerate().take(ngram).skip(1) {
mixed ^= row[position].wrapping_mul(multipliers[shift]);
}
for head in 0..heads_per_ngram {
let index = head_start + head;
ids[token * total_heads + index] = mixed.rem_euclid(sizes[index]) + offsets[index];
}
}
}
ids
}
#[allow(clippy::too_many_arguments)]
fn launch_sdpa_mask(
e: &Engine,
q: &CudaSlice<f32>,
k: &CudaView<'_, f32>,
v: &CudaView<'_, f32>,
o: &mut CudaSlice<f32>,
mask: &CudaSlice<u8>,
head_dim: usize,
n_head: usize,
n_head_kv: usize,
t: usize,
t_kv: usize,
scale: f32,
) -> Res<()> {
if t_kv * 4 > 48 * 1024 {
return Err(
"sdpa_naive_mask_f32: T_kv exceeds the smem bound; the gmem twin is perf-lane work"
.into(),
);
}
let f = e.func("sdpa_naive_mask_f32");
let cfg = LaunchConfig {
grid_dim: (n_head as u32, t as u32, 1),
block_dim: (128, 1, 1),
shared_mem_bytes: (t_kv * 4) as u32,
};
let (hd, nh, nkv, ti, tkvi) = (
head_dim as i32,
n_head as i32,
n_head_kv as i32,
t as i32,
t_kv as i32,
);
let stream = e.gpu.stream();
let mut b = stream.launch_builder(&f);
b.arg(q)
.arg(k)
.arg(v)
.arg(o)
.arg(mask)
.arg(&hd)
.arg(&nh)
.arg(&nkv)
.arg(&ti)
.arg(&tkvi)
.arg(&scale);
unsafe {
b.launch(cfg)?;
}
Ok(())
}
#[allow(clippy::too_many_arguments)]
fn launch_sdpa_blocklist(
e: &Engine,
q: &CudaSlice<f32>,
k: &CudaView<'_, f32>,
v: &CudaView<'_, f32>,
o: &mut CudaSlice<f32>,
pos: &CudaSlice<i32>,
meta: &CudaSlice<i32>,
head_dim: usize,
n_head: usize,
n_head_kv: usize,
t: usize,
max_count: usize,
scale: f32,
) -> Res<()> {
let smem = (max_count * 8) as u32;
if smem > 48 * 1024 {
return Err("sdpa_blocklist_f32: selection exceeds the smem budget".into());
}
let f = e.func("sdpa_blocklist_f32");
let cfg = LaunchConfig {
grid_dim: (n_head as u32, t as u32, 1),
block_dim: (128, 1, 1),
shared_mem_bytes: smem,
};
let (hd, nh, nkv, ti, mc) = (
head_dim as i32,
n_head as i32,
n_head_kv as i32,
t as i32,
max_count as i32,
);
let stream = e.gpu.stream();
let mut b = stream.launch_builder(&f);
b.arg(q)
.arg(k)
.arg(v)
.arg(o)
.arg(pos)
.arg(meta)
.arg(&hd)
.arg(&nh)
.arg(&nkv)
.arg(&ti)
.arg(&mc)
.arg(&scale);
unsafe {
b.launch(cfg)?;
}
Ok(())
}
#[allow(clippy::too_many_arguments)]
fn launch_q4e_kv_append(
e: &Engine,
k_rows: &CudaSlice<f32>,
v_rows: &CudaSlice<f32>,
k: &mut CudaSlice<u8>,
v: &mut CudaSlice<u8>,
base_pos: usize,
t: usize,
kv_dim: usize,
) -> Res<()> {
let f = e.func("q4e_kv_append_q8q5_rows");
let blocks = kv_dim.div_ceil(32);
let cfg = LaunchConfig {
grid_dim: (blocks as u32, t as u32, 1),
block_dim: (32, 1, 1),
shared_mem_bytes: 0,
};
let (t0, dk, dv) = (base_pos as i32, kv_dim as i32, kv_dim as i32);
let (ktb, vtb) = (q8_row_bytes(kv_dim) as i64, q5_row_bytes(kv_dim) as i64);
let stream = e.gpu.stream();
let mut b = stream.launch_builder(&f);
b.arg(k_rows)
.arg(v_rows)
.arg(k)
.arg(v)
.arg(&t0)
.arg(&dk)
.arg(&dv)
.arg(&ktb)
.arg(&vtb);
unsafe {
b.launch(cfg)?;
}
Ok(())
}
#[allow(clippy::too_many_arguments)]
fn launch_q4e_kv_dequant_rows(
e: &Engine,
k: &CudaSlice<u8>,
v: &CudaSlice<u8>,
k_out: &mut CudaSlice<f32>,
v_out: &mut CudaSlice<f32>,
r0: usize,
rows: usize,
kv_dim: usize,
) -> Res<()> {
let f = e.func("q4e_kv_dequant_rows");
let blocks = kv_dim.div_ceil(32);
let cfg = LaunchConfig {
grid_dim: (blocks as u32, rows as u32, 1),
block_dim: (32, 1, 1),
shared_mem_bytes: 0,
};
let (r0i, dk, dv) = (r0 as i32, kv_dim as i32, kv_dim as i32);
let (ktb, vtb) = (q8_row_bytes(kv_dim) as i64, q5_row_bytes(kv_dim) as i64);
let stream = e.gpu.stream();
let mut b = stream.launch_builder(&f);
b.arg(k)
.arg(v)
.arg(k_out)
.arg(v_out)
.arg(&r0i)
.arg(&dk)
.arg(&dv)
.arg(&ktb)
.arg(&vtb);
unsafe {
b.launch(cfg)?;
}
Ok(())
}
#[allow(clippy::too_many_arguments)]
fn launch_q4e_sdpa_blocklist_q8q5(
e: &Engine,
q: &CudaSlice<f32>,
k: &CudaSlice<u8>,
v: &CudaSlice<u8>,
o: &mut CudaSlice<f32>,
pos: &CudaSlice<i32>,
meta: &CudaSlice<i32>,
head_dim: usize,
n_head: usize,
n_head_kv: usize,
t: usize,
max_count: usize,
scale: f32,
) -> Res<()> {
let smem = (max_count * 8) as u32;
if smem > 48 * 1024 {
return Err("q4e_sdpa_blocklist_q8q5: selection exceeds the smem budget".into());
}
let f = e.func(if kv_hoist_on() {
"q4e_sdpa_blocklist_q8q5_hoist"
} else {
"q4e_sdpa_blocklist_q8q5"
});
let cfg = LaunchConfig {
grid_dim: (n_head as u32, t as u32, 1),
block_dim: (128, 1, 1),
shared_mem_bytes: smem,
};
let kv_dim = n_head_kv * head_dim;
let (hd, nh, nkv, ti, mc) = (
head_dim as i32,
n_head as i32,
n_head_kv as i32,
t as i32,
max_count as i32,
);
let (ktb, vtb) = (q8_row_bytes(kv_dim) as i64, q5_row_bytes(kv_dim) as i64);
let stream = e.gpu.stream();
let mut b = stream.launch_builder(&f);
b.arg(q)
.arg(k)
.arg(v)
.arg(o)
.arg(pos)
.arg(meta)
.arg(&hd)
.arg(&nh)
.arg(&nkv)
.arg(&ti)
.arg(&mc)
.arg(&scale)
.arg(&ktb)
.arg(&vtb);
unsafe {
b.launch(cfg)?;
}
Ok(())
}
#[allow(clippy::too_many_arguments)]
fn launch_q4e_idx_append_q8(
e: &Engine,
src: &CudaSlice<f32>,
dst: &mut CudaSlice<u8>,
rows: usize,
width: usize,
src_stride: usize,
src_col: usize,
dst_row: usize,
) -> Res<()> {
let f = e.func("q4e_idx_append_q8");
let cfg = LaunchConfig {
grid_dim: (width.div_ceil(32) as u32, rows as u32, 1),
block_dim: (32, 1, 1),
shared_mem_bytes: 0,
};
let (r, w) = (rows as i32, width as i32);
let (ss, sc, dr) = (src_stride as i64, src_col as i64, dst_row as i64);
let stream = e.gpu.stream();
let mut b = stream.launch_builder(&f);
b.arg(src)
.arg(dst)
.arg(&r)
.arg(&w)
.arg(&ss)
.arg(&sc)
.arg(&dr);
unsafe {
b.launch(cfg)?;
}
Ok(())
}
#[allow(clippy::too_many_arguments)]
fn launch_q4e_idx_append_bf16(
e: &Engine,
src: &CudaSlice<f32>,
dst: &mut CudaSlice<u16>,
rows: usize,
width: usize,
src_stride: usize,
src_col: usize,
dst_row: usize,
) -> Res<()> {
let f = e.func("q4e_idx_append_bf16");
let total = rows * width;
let cfg = LaunchConfig {
grid_dim: (total.div_ceil(256) as u32, 1, 1),
block_dim: (256, 1, 1),
shared_mem_bytes: 0,
};
let (r, w) = (rows as i32, width as i32);
let (ss, sc, dr) = (src_stride as i64, src_col as i64, dst_row as i64);
let stream = e.gpu.stream();
let mut b = stream.launch_builder(&f);
b.arg(src)
.arg(dst)
.arg(&r)
.arg(&w)
.arg(&ss)
.arg(&sc)
.arg(&dr);
unsafe {
b.launch(cfg)?;
}
Ok(())
}
#[allow(clippy::too_many_arguments)]
fn launch_gdn_scan(
e: &Engine,
qkv: &CudaSlice<f32>,
g_log: &CudaSlice<f32>,
beta_raw: &CudaSlice<f32>,
state: &mut CudaSlice<f32>,
o: &mut CudaSlice<f32>,
nk: usize,
nv: usize,
hk: usize,
hv: usize,
t: usize,
scale: f32,
eps: f32,
) -> Res<()> {
if hk > 128 {
return Err("gdn_scan_naive_f32: hk > 128".into());
}
let f = e.func("gdn_scan_naive_f32");
let cfg = LaunchConfig {
grid_dim: (nv as u32, 1, 1),
block_dim: (hv as u32, 1, 1),
shared_mem_bytes: ((2 * hk + 2) * 4) as u32,
};
let (nki, nvi, hki, hvi, ti) = (nk as i32, nv as i32, hk as i32, hv as i32, t as i32);
let stream = e.gpu.stream();
let mut b = stream.launch_builder(&f);
b.arg(qkv)
.arg(g_log)
.arg(beta_raw)
.arg(state)
.arg(o)
.arg(&nki)
.arg(&nvi)
.arg(&hki)
.arg(&hvi)
.arg(&ti)
.arg(&scale)
.arg(&eps);
unsafe {
b.launch(cfg)?;
}
Ok(())
}
#[allow(clippy::too_many_arguments)]
fn launch_gdn_scan_step(
e: &Engine,
qkv: &CudaSlice<f32>,
g_log: &CudaSlice<f32>,
beta_raw: &CudaSlice<f32>,
state: &mut CudaSlice<f32>,
o: &mut CudaSlice<f32>,
nk: usize,
nv: usize,
hk: usize,
hv: usize,
scale: f32,
eps: f32,
) -> Res<()> {
if hk % 32 != 0 || hk > 1024 {
return Err("gdn_scan_step_f32: hk must be a multiple of 32 and <= 1024".into());
}
let f = e.func("gdn_scan_step_f32");
let cfg = LaunchConfig {
grid_dim: (nv as u32, hv as u32, 1),
block_dim: (hk as u32, 1, 1),
shared_mem_bytes: 0,
};
let (nki, nvi, hki, hvi) = (nk as i32, nv as i32, hk as i32, hv as i32);
let stream = e.gpu.stream();
let mut b = stream.launch_builder(&f);
b.arg(qkv)
.arg(g_log)
.arg(beta_raw)
.arg(state)
.arg(o)
.arg(&nki)
.arg(&nvi)
.arg(&hki)
.arg(&hvi)
.arg(&scale)
.arg(&eps);
unsafe {
b.launch(cfg)?;
}
Ok(())
}
#[allow(clippy::too_many_arguments)]
fn launch_gdn_scan_step_at(
e: &Engine,
conv_out: &CudaSlice<f32>,
g_log: &CudaSlice<f32>,
beta_raw: &CudaSlice<f32>,
state: &mut CudaSlice<f32>,
o: &mut CudaSlice<f32>,
tok: usize,
nk: usize,
nv: usize,
hk: usize,
hv: usize,
scale: f32,
eps: f32,
) -> Res<()> {
if hk % 32 != 0 || hk > 1024 {
return Err("gdn_scan_step_f32: hk must be a multiple of 32 and <= 1024".into());
}
let conv_dim = 2 * nk * hk + nv * hv;
let qv = conv_out.slice(tok * conv_dim..(tok + 1) * conv_dim);
let gv = g_log.slice(tok * nv..(tok + 1) * nv);
let bv = beta_raw.slice(tok * nv..(tok + 1) * nv);
let mut ov = o.slice_mut(tok * nv * hv..(tok + 1) * nv * hv);
let f = e.func("gdn_scan_step_f32");
let cfg = LaunchConfig {
grid_dim: (nv as u32, hv as u32, 1),
block_dim: (hk as u32, 1, 1),
shared_mem_bytes: 0,
};
let (nki, nvi, hki, hvi) = (nk as i32, nv as i32, hk as i32, hv as i32);
let stream = e.gpu.stream();
let mut b = stream.launch_builder(&f);
b.arg(&qv)
.arg(&gv)
.arg(&bv)
.arg(&mut *state)
.arg(&mut ov)
.arg(&nki)
.arg(&nvi)
.arg(&hki)
.arg(&hvi)
.arg(&scale)
.arg(&eps);
unsafe {
b.launch(cfg)?;
}
Ok(())
}
#[allow(clippy::too_many_arguments)]
fn launch_gdn_scan_at(
e: &Engine,
conv_out: &CudaSlice<f32>,
g_log: &CudaSlice<f32>,
beta_raw: &CudaSlice<f32>,
state: &mut CudaSlice<f32>,
o: &mut CudaSlice<f32>,
tok: usize,
nk: usize,
nv: usize,
hk: usize,
hv: usize,
scale: f32,
eps: f32,
) -> Res<()> {
if hk > 128 {
return Err("gdn_scan_naive_f32: hk > 128".into());
}
let conv_dim = 2 * nk * hk + nv * hv;
let qv = conv_out.slice(tok * conv_dim..(tok + 1) * conv_dim);
let gv = g_log.slice(tok * nv..(tok + 1) * nv);
let bv = beta_raw.slice(tok * nv..(tok + 1) * nv);
let mut ov = o.slice_mut(tok * nv * hv..(tok + 1) * nv * hv);
let f = e.func("gdn_scan_naive_f32");
let cfg = LaunchConfig {
grid_dim: (nv as u32, 1, 1),
block_dim: (hv as u32, 1, 1),
shared_mem_bytes: ((2 * hk + 2) * 4) as u32,
};
let (nki, nvi, hki, hvi, ti) = (nk as i32, nv as i32, hk as i32, hv as i32, 1i32);
let stream = e.gpu.stream();
let mut b = stream.launch_builder(&f);
b.arg(&qv)
.arg(&gv)
.arg(&bv)
.arg(&mut *state)
.arg(&mut ov)
.arg(&nki)
.arg(&nvi)
.arg(&hki)
.arg(&hvi)
.arg(&ti)
.arg(&scale)
.arg(&eps);
unsafe {
b.launch(cfg)?;
}
Ok(())
}
#[allow(clippy::too_many_arguments)]
fn launch_rms_sigmul(
e: &Engine,
x: &CudaSlice<f32>,
w: &CudaSlice<f32>,
z: &CudaSlice<f32>,
dst: &mut CudaSlice<f32>,
ncols: usize,
nrows: usize,
eps: f32,
) -> Res<()> {
let f = e.func("rms_sigmul_f32");
let cfg = LaunchConfig {
grid_dim: (nrows as u32, 1, 1),
block_dim: (crate::rms_block(), 1, 1),
shared_mem_bytes: 0,
};
let (nc, ep) = (ncols as i32, eps);
let stream = e.gpu.stream();
let mut b = stream.launch_builder(&f);
b.arg(x).arg(w).arg(z).arg(dst).arg(&nc).arg(&ep);
unsafe {
b.launch(cfg)?;
}
Ok(())
}
#[allow(clippy::too_many_arguments)]
fn launch_dwconv(
e: &Engine,
x: &CudaSlice<f32>,
hist: &CudaSlice<f32>,
w: &CudaSlice<f32>,
y: &mut CudaSlice<f32>,
t: usize,
th: usize,
c: usize,
k: usize,
dilation: usize,
mode: i32,
) -> Res<()> {
let f = e.func("dwconv_causal_f32");
let cfg = LaunchConfig::for_num_elems((t * c) as u32);
let (ti, thi, ci, ki, di) = (t as i32, th as i32, c as i32, k as i32, dilation as i32);
let stream = e.gpu.stream();
let mut b = stream.launch_builder(&f);
b.arg(x)
.arg(hist)
.arg(w)
.arg(y)
.arg(&ti)
.arg(&thi)
.arg(&ci)
.arg(&ki)
.arg(&di)
.arg(&mode);
unsafe {
b.launch(cfg)?;
}
Ok(())
}
#[allow(clippy::too_many_arguments)]
fn run_routed_expert(
e: &Engine,
xg: &CudaSlice<f32>,
gate: &CudaView<'_, f32>,
up: &CudaView<'_, f32>,
down: &CudaView<'_, f32>,
m_e: usize,
hidden: usize,
ff: usize,
) -> Res<CudaSlice<f32>> {
let xg_view = xg.slice(0..m_e * hidden);
let mut gate_out = e.uninit(m_e * ff)?;
e.linear_device_into(&xg_view, gate, &mut gate_out, m_e, hidden, ff)?;
let mut up_out = e.uninit(m_e * ff)?;
e.linear_device_into(&xg_view, up, &mut up_out, m_e, hidden, ff)?;
let mut act = e.uninit(m_e * ff)?;
e.silu_mul(&gate_out, &up_out, &mut act, m_e * ff)?;
let mut down_out = e.uninit(m_e * hidden)?;
e.linear_device_into(
&act.slice(0..m_e * ff),
down,
&mut down_out,
m_e,
ff,
hidden,
)?;
Ok(down_out)
}
fn launch_rms_norm_into_view(
e: &Engine,
x: &CudaSlice<f32>,
w: &CudaSlice<f32>,
dst: &mut cudarc::driver::CudaViewMut<'_, f32>,
ncols: usize,
nrows: usize,
eps: f32,
) -> Res<()> {
let kname = if Engine::norm_ilp_on() {
"rms_norm_f32_v2"
} else {
"rms_norm_f32"
};
let f = e.func(kname);
let cfg = LaunchConfig {
grid_dim: (nrows as u32, 1, 1),
block_dim: (crate::rms_block(), 1, 1),
shared_mem_bytes: 0,
};
let (nc, ep) = (ncols as i32, eps);
let stream = e.gpu.stream();
let mut b = stream.launch_builder(&f);
b.arg(x).arg(w).arg(dst).arg(&nc).arg(&ep);
unsafe {
b.launch(cfg)?;
}
Ok(())
}
fn launch_hc_lowrank_reduce(
e: &Engine,
parts: &CudaSlice<f32>,
low_act: &mut CudaSlice<f32>,
streams: usize,
t: usize,
rank: usize,
) -> Res<()> {
let f = e.func("hc_lowrank_reduce_f32");
let cfg = LaunchConfig::for_num_elems((t * rank) as u32);
let (si, ti, ri) = (streams as i32, t as i32, rank as i32);
let inv = 1.0f32 / streams as f32;
let stream = e.gpu.stream();
let mut b = stream.launch_builder(&f);
b.arg(parts)
.arg(low_act)
.arg(&si)
.arg(&ti)
.arg(&ri)
.arg(&inv);
unsafe {
b.launch(cfg)?;
}
Ok(())
}
fn launch_hc_mix_epilogue(
e: &Engine,
gates: &CudaSlice<f32>,
normed: &CudaSlice<f32>,
mixed: &mut CudaSlice<f32>,
streams: usize,
t: usize,
hidden: usize,
) -> Res<()> {
let f = e.func("hc_mix_epilogue_f32");
let cfg = LaunchConfig::for_num_elems((t * hidden) as u32);
let (si, ti, hi) = (streams as i32, t as i32, hidden as i32);
let inv = 1.0f32 / streams as f32;
let stream = e.gpu.stream();
let mut b = stream.launch_builder(&f);
b.arg(gates)
.arg(normed)
.arg(mixed)
.arg(&si)
.arg(&ti)
.arg(&hi)
.arg(&inv);
unsafe {
b.launch(cfg)?;
}
Ok(())
}
fn launch_hc_inject_gates(
e: &Engine,
normed: &CudaSlice<f32>,
w: &CudaSlice<f32>,
out: &mut CudaSlice<f32>,
streams: usize,
t: usize,
hidden: usize,
) -> Res<()> {
let f = e.func("hc_inject_gates_f32");
let cfg = LaunchConfig {
grid_dim: (streams as u32, t as u32, 1),
block_dim: (256, 1, 1),
shared_mem_bytes: 0,
};
let (si, ti, hi) = (streams as i32, t as i32, hidden as i32);
let inv = 1.0f32 / streams as f32;
let stream = e.gpu.stream();
let mut b = stream.launch_builder(&f);
b.arg(normed)
.arg(w)
.arg(out)
.arg(&si)
.arg(&ti)
.arg(&hi)
.arg(&inv);
unsafe {
b.launch(cfg)?;
}
Ok(())
}
enum InjectOut {
Rows(Vec<CudaSlice<f32>>),
Slab(CudaSlice<f32>),
}
fn put_inject(ws: &mut StepPool, inject: InjectOut) {
match inject {
InjectOut::Rows(rows) => {
for (s, row) in rows.into_iter().enumerate() {
ws.put_f32(INJECT_SLOTS[s], row);
}
}
InjectOut::Slab(slab) => ws.put_f32("hc.inj_all", slab),
}
}
fn take_inject(e: &Engine, ws: &mut StepPool, streams: usize, t: usize) -> Res<InjectOut> {
if micro_inj_on() && hc_fused_gate_on() {
Ok(InjectOut::Slab(ws.take_f32(
e,
"hc.inj_all",
streams * t,
0,
)?))
} else {
let mut rows = Vec::with_capacity(streams);
for s in 0..streams {
rows.push(ws.take_f32(e, INJECT_SLOTS[s], t, 0)?);
}
Ok(InjectOut::Rows(rows))
}
}
#[allow(clippy::too_many_arguments)]
fn launch_hc_norm_planes(
e: &Engine,
ptrs: &CudaSlice<u64>,
w_stack: &CudaSlice<f32>,
dst: &mut CudaSlice<f32>,
hidden: usize,
t: usize,
streams: usize,
eps: f32,
) -> Res<()> {
let f = e.func("hc_norm_planes_f32");
let cfg = LaunchConfig {
grid_dim: (t as u32, streams as u32, 1),
block_dim: (256, 1, 1),
shared_mem_bytes: 0,
};
let (hi, ti) = (hidden as i32, t as i32);
let stream = e.gpu.stream();
let mut b = stream.launch_builder(&f);
b.arg(ptrs)
.arg(w_stack)
.arg(dst)
.arg(&hi)
.arg(&ti)
.arg(&eps);
unsafe {
b.launch(cfg)?;
}
Ok(())
}
#[allow(clippy::too_many_arguments)]
fn launch_hc_inject_two_stage(
e: &Engine,
normed: &CudaSlice<f32>,
w_f32: &CudaSlice<f32>,
w_b16: Option<&CudaSlice<u8>>,
partials: &mut CudaSlice<f32>,
out: &mut CudaSlice<f32>,
streams: usize,
t: usize,
hidden: usize,
chunks: usize,
) -> Res<()> {
let cfg = LaunchConfig {
grid_dim: (streams as u32, t as u32, chunks as u32),
block_dim: (256, 1, 1),
shared_mem_bytes: 0,
};
let (si, ti, hi, ci) = (streams as i32, t as i32, hidden as i32, chunks as i32);
let stream = e.gpu.stream();
if let Some(w) = w_b16 {
let f = e.func("hc_inject_partials_bf16w_f32");
let mut b = stream.launch_builder(&f);
b.arg(normed)
.arg(w)
.arg(&mut *partials)
.arg(&si)
.arg(&ti)
.arg(&hi)
.arg(&ci);
unsafe {
b.launch(cfg)?;
}
} else {
let f = e.func("hc_inject_partials_f32");
let mut b = stream.launch_builder(&f);
b.arg(normed)
.arg(w_f32)
.arg(&mut *partials)
.arg(&si)
.arg(&ti)
.arg(&hi)
.arg(&ci);
unsafe {
b.launch(cfg)?;
}
}
let rows = (streams * t) as i32;
let inv = 1.0f32 / streams as f32;
let f = e.func("hc_inject_reduce_f32");
let cfg = LaunchConfig::for_num_elems((streams * t) as u32);
let mut b = stream.launch_builder(&f);
b.arg(&*partials).arg(out).arg(&rows).arg(&ci).arg(&inv);
unsafe {
b.launch(cfg)?;
}
Ok(())
}
#[allow(clippy::too_many_arguments)]
fn launch_hc_diet_stage1(
e: &Engine,
ptrs: &CudaSlice<u64>,
nw_stack: &CudaSlice<f32>,
wdown_b16: &CudaSlice<u8>,
winj_b16: Option<&CudaSlice<u8>>,
parts: &mut CudaSlice<f32>,
inj_parts: &mut CudaSlice<f32>,
inv_out: &mut CudaSlice<f32>,
hidden: usize,
rank: usize,
streams: usize,
t: usize,
eps: f32,
) -> Res<()> {
if hidden % 8 != 0 {
return Err("hc_diet_stage1_f32: hidden % 8 != 0".into());
}
let n_inj = if winj_b16.is_some() { streams } else { 0 };
const ROWS_PB: usize = 4;
let total_rows = rank + n_inj;
if parts.len() < t * streams * rank
|| (n_inj > 0 && inj_parts.len() < t * n_inj * streams)
|| inv_out.len() < t * streams
{
return Err("hc_diet_stage1_f32: output buffers too short".into());
}
let f = e.func("hc_diet_stage1_f32");
let cfg = LaunchConfig {
grid_dim: (
total_rows.div_ceil(ROWS_PB) as u32,
t as u32,
streams as u32,
),
block_dim: (256, 1, 1),
shared_mem_bytes: (hidden * 4) as u32,
};
let (hi, ri, si, nji, rpb) = (
hidden as i32,
rank as i32,
streams as i32,
n_inj as i32,
ROWS_PB as i32,
);
let winj = winj_b16.unwrap_or(wdown_b16); let stream = e.gpu.stream();
let mut b = stream.launch_builder(&f);
b.arg(ptrs)
.arg(nw_stack)
.arg(wdown_b16)
.arg(winj)
.arg(&mut *parts)
.arg(&mut *inj_parts)
.arg(&mut *inv_out)
.arg(&hi)
.arg(&ri)
.arg(&si)
.arg(&nji)
.arg(&rpb)
.arg(&eps);
unsafe {
b.launch(cfg)?;
}
Ok(())
}
#[allow(clippy::too_many_arguments)]
fn launch_hc_diet_stage2(
e: &Engine,
parts: &CudaSlice<f32>,
inj_parts: &CudaSlice<f32>,
low_act: &mut CudaSlice<f32>,
inj_all: &mut CudaSlice<f32>,
rank: usize,
streams: usize,
t: usize,
with_inject: bool,
) -> Res<()> {
let n_inj = if with_inject { streams } else { 0 };
if low_act.len() < t * rank || (n_inj > 0 && inj_all.len() < n_inj * t) {
return Err("hc_diet_stage2_f32: output buffers too short".into());
}
let f = e.func("hc_diet_stage2_f32");
let cfg = LaunchConfig {
grid_dim: (((rank + n_inj) as u32).div_ceil(256), t as u32, 1),
block_dim: (256, 1, 1),
shared_mem_bytes: 0,
};
let (ri, si, nji, ti) = (rank as i32, streams as i32, n_inj as i32, t as i32);
let inv = 1.0f32 / streams as f32;
let stream = e.gpu.stream();
let mut b = stream.launch_builder(&f);
b.arg(parts)
.arg(inj_parts)
.arg(&mut *low_act)
.arg(&mut *inj_all)
.arg(&ri)
.arg(&si)
.arg(&nji)
.arg(&ti)
.arg(&inv);
unsafe {
b.launch(cfg)?;
}
Ok(())
}
#[allow(clippy::too_many_arguments)]
fn launch_hc_diet_stage3(
e: &Engine,
ptrs: &CudaSlice<u64>,
nw_stack: &CudaSlice<f32>,
inv_in: &CudaSlice<f32>,
wup_b16: &CudaSlice<u8>,
low_act: &CudaSlice<f32>,
mixed: &mut CudaSlice<f32>,
hidden: usize,
rank: usize,
streams: usize,
t: usize,
) -> Res<()> {
const DIMS_PB: usize = 8;
if mixed.len() < t * hidden {
return Err("hc_diet_stage3_f32: output buffer too short".into());
}
let f = e.func("hc_diet_stage3_f32");
let cfg = LaunchConfig {
grid_dim: (hidden.div_ceil(DIMS_PB) as u32, t as u32, 1),
block_dim: (256, 1, 1),
shared_mem_bytes: ((rank + DIMS_PB * streams) * 4) as u32,
};
let (hi, ri, si, dpb) = (hidden as i32, rank as i32, streams as i32, DIMS_PB as i32);
let inv_streams = 1.0f32 / streams as f32;
let stream = e.gpu.stream();
let mut b = stream.launch_builder(&f);
b.arg(ptrs)
.arg(nw_stack)
.arg(inv_in)
.arg(wup_b16)
.arg(low_act)
.arg(&mut *mixed)
.arg(&hi)
.arg(&ri)
.arg(&si)
.arg(&dpb)
.arg(&inv_streams);
unsafe {
b.launch(cfg)?;
}
Ok(())
}
fn launch_hc_diet_stage0_mt(
e: &Engine,
ptrs: &CudaSlice<u64>,
inv_out: &mut CudaSlice<f32>,
hidden: usize,
streams: usize,
t: usize,
eps: f32,
) -> Res<()> {
if inv_out.len() < t * streams {
return Err("hc_diet_stage0_mt_f32: inv buffer too short".into());
}
let f = e.func("hc_diet_stage0_mt_f32");
let cfg = LaunchConfig {
grid_dim: (t as u32, streams as u32, 1),
block_dim: (256, 1, 1),
shared_mem_bytes: 0,
};
let (hi, si, ti) = (hidden as i32, streams as i32, t as i32);
let stream = e.gpu.stream();
let mut b = stream.launch_builder(&f);
b.arg(ptrs)
.arg(&mut *inv_out)
.arg(&hi)
.arg(&si)
.arg(&ti)
.arg(&eps);
unsafe {
b.launch(cfg)?;
}
Ok(())
}
#[allow(clippy::too_many_arguments)]
fn launch_hc_diet_stage1_mt(
e: &Engine,
ptrs: &CudaSlice<u64>,
nw_stack: &CudaSlice<f32>,
inv_in: &CudaSlice<f32>,
wdown_b16: &CudaSlice<u8>,
winj_b16: Option<&CudaSlice<u8>>,
parts: &mut CudaSlice<f32>,
inj_parts: &mut CudaSlice<f32>,
hidden: usize,
rank: usize,
streams: usize,
t: usize,
) -> Res<()> {
if hidden % 8 != 0 || !(2..=12).contains(&t) {
return Err("hc_diet_stage1_mt_f32: geometry".into());
}
let n_inj = if winj_b16.is_some() { streams } else { 0 };
const ROWS_PB: usize = 4;
let total_rows = rank + n_inj;
if parts.len() < t * streams * rank || (n_inj > 0 && inj_parts.len() < t * n_inj * streams) {
return Err("hc_diet_stage1_mt_f32: output buffers too short".into());
}
let f = e.func("hc_diet_stage1_mt_f32");
let cfg = LaunchConfig {
grid_dim: (total_rows.div_ceil(ROWS_PB) as u32, 1, streams as u32),
block_dim: (256, 1, 1),
shared_mem_bytes: 0,
};
let (hi, ri, si, nji, rpb, ti) = (
hidden as i32,
rank as i32,
streams as i32,
n_inj as i32,
ROWS_PB as i32,
t as i32,
);
let winj = winj_b16.unwrap_or(wdown_b16);
let stream = e.gpu.stream();
let mut b = stream.launch_builder(&f);
b.arg(ptrs)
.arg(nw_stack)
.arg(inv_in)
.arg(wdown_b16)
.arg(winj)
.arg(&mut *parts)
.arg(&mut *inj_parts)
.arg(&hi)
.arg(&ri)
.arg(&si)
.arg(&nji)
.arg(&rpb)
.arg(&ti);
unsafe {
b.launch(cfg)?;
}
Ok(())
}
#[allow(clippy::too_many_arguments)]
fn launch_hc_diet_stage3_mt(
e: &Engine,
ptrs: &CudaSlice<u64>,
nw_stack: &CudaSlice<f32>,
inv_in: &CudaSlice<f32>,
wup_b16: &CudaSlice<u8>,
low_act: &CudaSlice<f32>,
mixed: &mut CudaSlice<f32>,
hidden: usize,
rank: usize,
streams: usize,
t: usize,
) -> Res<()> {
const DIMS_PB: usize = 8;
if !(2..=12).contains(&t) || mixed.len() < t * hidden {
return Err("hc_diet_stage3_mt_f32: geometry".into());
}
let smem = ((t * rank + DIMS_PB * streams * t) * 4) as u32;
if smem > 96 * 1024 {
return Err("hc_diet_stage3_mt_f32: smem over budget".into());
}
let f = e.func("hc_diet_stage3_mt_f32");
let cfg = LaunchConfig {
grid_dim: (hidden.div_ceil(DIMS_PB) as u32, 1, 1),
block_dim: (256, 1, 1),
shared_mem_bytes: smem,
};
let (hi, ri, si, dpb, ti) = (
hidden as i32,
rank as i32,
streams as i32,
DIMS_PB as i32,
t as i32,
);
let inv_streams = 1.0f32 / streams as f32;
let stream = e.gpu.stream();
let mut b = stream.launch_builder(&f);
b.arg(ptrs)
.arg(nw_stack)
.arg(inv_in)
.arg(wup_b16)
.arg(low_act)
.arg(&mut *mixed)
.arg(&hi)
.arg(&ri)
.arg(&si)
.arg(&dpb)
.arg(&ti)
.arg(&inv_streams);
unsafe {
b.launch(cfg)?;
}
Ok(())
}
fn launch_hc_write_planes(
e: &Engine,
ptrs: &CudaSlice<u64>,
block_out: &CudaSlice<f32>,
inj: &CudaSlice<f32>,
hidden: usize,
t: usize,
streams: usize,
) -> Res<()> {
let f = e.func("hc_write_planes_f32");
let n = (t * hidden) as u32;
let cfg = LaunchConfig {
grid_dim: (n.div_ceil(256), streams as u32, 1),
block_dim: (256, 1, 1),
shared_mem_bytes: 0,
};
let (hi, ti) = (hidden as i32, t as i32);
let stream = e.gpu.stream();
let mut b = stream.launch_builder(&f);
b.arg(ptrs).arg(block_out).arg(inj).arg(&hi).arg(&ti);
unsafe {
b.launch(cfg)?;
}
Ok(())
}
fn launch_hc_inject_gates_b16(
e: &Engine,
normed: &CudaSlice<f32>,
w: &CudaSlice<u8>,
out: &mut CudaSlice<f32>,
streams: usize,
t: usize,
hidden: usize,
) -> Res<()> {
let f = e.func("hc_inject_gates_bf16w_f32");
let cfg = LaunchConfig {
grid_dim: (streams as u32, t as u32, 1),
block_dim: (256, 1, 1),
shared_mem_bytes: 0,
};
let (si, ti, hi) = (streams as i32, t as i32, hidden as i32);
let inv = 1.0f32 / streams as f32;
let stream = e.gpu.stream();
let mut b = stream.launch_builder(&f);
b.arg(normed)
.arg(w)
.arg(out)
.arg(&si)
.arg(&ti)
.arg(&hi)
.arg(&inv);
unsafe {
b.launch(cfg)?;
}
Ok(())
}
fn bf16_twin(e: &Engine, data: &[f32], in_f: usize) -> Res<Option<CudaSlice<u8>>> {
if in_f % 8 != 0 {
return Ok(None);
}
let mut bytes = Vec::with_capacity(data.len() * 2);
for &v in data {
let bits = v.to_bits();
if bits & 0xFFFF != 0 {
return Ok(None);
}
bytes.extend_from_slice(&((bits >> 16) as u16).to_le_bytes());
}
Ok(Some(e.htod_bytes(&bytes)?))
}
#[allow(clippy::too_many_arguments)]
fn launch_qmatvec_bf16w(
e: &Engine,
w: &CudaSlice<u8>,
x: &CudaSlice<f32>,
y: &mut CudaSlice<f32>,
in_f: usize,
out_f: usize,
t: usize,
batch: usize,
w_bstride: usize,
x_bstride: usize,
x_tstride: usize,
y_bstride: usize,
) -> Res<()> {
if in_f % 8 != 0 || x_bstride % 8 != 0 || x_tstride % 8 != 0 {
return Err("qmatvec_bf16w_f32: stride breaks the uint4/float4 vector width".into());
}
if y.len() < (batch - 1) * y_bstride + t * out_f {
return Err("qmatvec_bf16w_f32: output buffer too short".into());
}
let f = e.func("qmatvec_bf16w_f32");
let cfg = LaunchConfig {
grid_dim: (out_f as u32, t as u32, batch as u32),
block_dim: (128, 1, 1),
shared_mem_bytes: 0,
};
let (inf, outf, ti) = (in_f as i32, out_f as i32, t as i32);
let (wb, xb, xt, yb) = (
w_bstride as i64,
x_bstride as i64,
x_tstride as i64,
y_bstride as i64,
);
let stream = e.gpu.stream();
let mut b = stream.launch_builder(&f);
b.arg(w)
.arg(x)
.arg(y)
.arg(&inf)
.arg(&outf)
.arg(&ti)
.arg(&wb)
.arg(&xb)
.arg(&xt)
.arg(&yb);
unsafe {
b.launch(cfg)?;
}
Ok(())
}
fn bf16_stack_twin(e: &Engine, parts: &[&[f32]], in_f: usize) -> Res<Option<CudaSlice<u8>>> {
let mut cat: Vec<f32> = Vec::with_capacity(parts.iter().map(|p| p.len()).sum());
for p in parts {
cat.extend_from_slice(p);
}
bf16_twin(e, &cat, in_f)
}
fn need_stack_twin(e: &Engine, parts: &[&[f32]], in_f: usize, what: &str) -> Res<CudaSlice<u8>> {
bf16_stack_twin(e, parts, in_f)?.ok_or_else(|| {
format!("qwen4exp_gpu tp2: {what} has no exact bf16 stack twin (in_f {in_f})").into()
})
}
#[allow(clippy::too_many_arguments)]
fn launch_qmatvec_bf16w_off(
e: &Engine,
w_stack: &CudaSlice<u8>,
row_off: usize,
x: &CudaSlice<f32>,
y: &mut CudaSlice<f32>,
in_f: usize,
out_f: usize,
t: usize,
) -> Res<()> {
if in_f % 8 != 0 {
return Err("qmatvec_bf16w_f32: stride breaks the uint4/float4 vector width".into());
}
if y.len() < t * out_f {
return Err("qmatvec_bf16w_f32: output buffer too short".into());
}
let byte_off = row_off * in_f * 2;
if w_stack.len() < byte_off + out_f * in_f * 2 {
return Err("qmatvec_bf16w_f32: stacked twin shorter than the row window".into());
}
let wv = w_stack.slice(byte_off..w_stack.len());
let f = e.func("qmatvec_bf16w_f32");
let cfg = LaunchConfig {
grid_dim: (out_f as u32, t as u32, 1),
block_dim: (128, 1, 1),
shared_mem_bytes: 0,
};
let (inf, outf, ti) = (in_f as i32, out_f as i32, t as i32);
let (wb, xb, xt, yb) = (0i64, 0i64, in_f as i64, 0i64);
let stream = e.gpu.stream();
let mut b = stream.launch_builder(&f);
b.arg(&wv)
.arg(x)
.arg(y)
.arg(&inf)
.arg(&outf)
.arg(&ti)
.arg(&wb)
.arg(&xb)
.arg(&xt)
.arg(&yb);
unsafe {
b.launch(cfg)?;
}
Ok(())
}
#[allow(clippy::too_many_arguments)]
fn launch_qmatvec_bf16w_off_into(
e: &Engine,
w_stack: &CudaSlice<u8>,
w_row_off: usize,
x: &CudaSlice<f32>,
x_off: usize,
y: &mut CudaSlice<f32>,
y_off: usize,
in_f: usize,
out_f: usize,
) -> Res<()> {
if in_f % 8 != 0 {
return Err("qmatvec_bf16w_f32: stride breaks the uint4/float4 vector width".into());
}
let byte_off = w_row_off * in_f * 2;
if w_stack.len() < byte_off + out_f * in_f * 2 {
return Err("qmatvec_bf16w_f32: bank shorter than the expert row window".into());
}
if x.len() < x_off + in_f || y.len() < y_off + out_f {
return Err("qmatvec_bf16w_f32: operand views out of range".into());
}
let wv = w_stack.slice(byte_off..w_stack.len());
let xv = x.slice(x_off..x_off + in_f);
let mut yv = y.slice_mut(y_off..y_off + out_f);
let f = e.func("qmatvec_bf16w_f32");
let cfg = LaunchConfig {
grid_dim: (out_f as u32, 1, 1),
block_dim: (128, 1, 1),
shared_mem_bytes: 0,
};
let (inf, outf, ti) = (in_f as i32, out_f as i32, 1i32);
let (wb, xb, xt, yb) = (0i64, 0i64, in_f as i64, 0i64);
let stream = e.gpu.stream();
let mut b = stream.launch_builder(&f);
b.arg(&wv)
.arg(&xv)
.arg(&mut yv)
.arg(&inf)
.arg(&outf)
.arg(&ti)
.arg(&wb)
.arg(&xb)
.arg(&xt)
.arg(&yb);
unsafe {
b.launch(cfg)?;
}
Ok(())
}
#[allow(clippy::too_many_arguments)]
fn launch_qmatvec_bf16w_sel(
e: &Engine,
bank: &CudaSlice<u8>,
sel: &CudaSlice<i32>,
sel_off: usize,
x: &CudaSlice<f32>,
x_off: usize,
x_sstride: usize,
y: &mut CudaSlice<f32>,
n_sel: usize,
in_f: usize,
out_f: usize,
) -> Res<()> {
if in_f % 8 != 0 {
return Err("qmatvec_bf16w_sel_f32: stride breaks the uint4/float4 vector width".into());
}
if sel.len() < sel_off + n_sel
|| x.len() < x_off + (n_sel - 1) * x_sstride + in_f
|| y.len() < n_sel * out_f
|| n_sel == 0
{
return Err("qmatvec_bf16w_sel_f32: operand views out of range".into());
}
let sv = sel.slice(sel_off..sel_off + n_sel);
let xv = x.slice(x_off..x.len());
let f = e.func("qmatvec_bf16w_sel_f32");
let cfg = LaunchConfig {
grid_dim: (out_f as u32, 1, n_sel as u32),
block_dim: (128, 1, 1),
shared_mem_bytes: 0,
};
let (inf, outf, ns) = (in_f as i32, out_f as i32, n_sel as i32);
let xs = x_sstride as i64;
let stream = e.gpu.stream();
let mut b = stream.launch_builder(&f);
b.arg(bank)
.arg(&sv)
.arg(&xv)
.arg(&mut *y)
.arg(&inf)
.arg(&outf)
.arg(&ns)
.arg(&xs);
unsafe {
b.launch(cfg)?;
}
Ok(())
}
#[allow(clippy::too_many_arguments)]
fn launch_qmatvec_bf16w_mt(
e: &Engine,
w_stack: &CudaSlice<u8>,
w_row_off: usize,
x: &CudaSlice<f32>,
y: &mut CudaSlice<f32>,
in_f: usize,
out_f: usize,
t: usize,
) -> Res<()> {
if in_f % 8 != 0 {
return Err("qmatvec_bf16w_mt_f32: in_f % 8 != 0".into());
}
if !(2..=12).contains(&t) {
return Err("qmatvec_bf16w_mt_f32: t out of range (2..=12)".into());
}
let byte_off = w_row_off * in_f * 2;
if w_stack.len() < byte_off + out_f * in_f * 2 || y.len() < t * out_f || x.len() < t * in_f {
return Err("qmatvec_bf16w_mt_f32: operands out of range".into());
}
let wv = w_stack.slice(byte_off..w_stack.len());
let f = e.func("qmatvec_bf16w_mt_f32");
let cfg = LaunchConfig {
grid_dim: (out_f as u32, 1, 1),
block_dim: (128, 1, 1),
shared_mem_bytes: 0,
};
let (inf, outf, ti) = (in_f as i32, out_f as i32, t as i32);
let (wb, xb, xt, yb) = (0i64, 0i64, in_f as i64, 0i64);
let stream = e.gpu.stream();
let mut b = stream.launch_builder(&f);
b.arg(&wv)
.arg(x)
.arg(y)
.arg(&inf)
.arg(&outf)
.arg(&ti)
.arg(&wb)
.arg(&xb)
.arg(&xt)
.arg(&yb);
unsafe {
b.launch(cfg)?;
}
Ok(())
}
#[allow(clippy::too_many_arguments)]
fn linear_trunk_stacked_into(
e: &Engine,
w_f32: &CudaSlice<f32>,
stack_b16: &Option<CudaSlice<u8>>,
row_off: usize,
x: &CudaSlice<f32>,
y: &mut CudaSlice<f32>,
t: usize,
in_f: usize,
out_f: usize,
) -> Res<()> {
if trunk_bf16_on() {
if let Some(w) = stack_b16 {
if (2..=12).contains(&t) && verify_mt_on() {
return launch_qmatvec_bf16w_mt(e, w, row_off, x, y, in_f, out_f, t);
}
return launch_qmatvec_bf16w_off(e, w, row_off, x, y, in_f, out_f, t);
}
}
if w_f32.len() < in_f * out_f {
return Err(
"qwen4exp_gpu: trunk f32 original dropped (trunk_f32_diet) — the bf16 \
twin path is required (keep trunk seams ON)"
.into(),
);
}
e.linear_device_into(x, w_f32, y, t, in_f, out_f)
}
fn launch_qmatvec_bf16w_multi4(
e: &Engine,
w_stack: &CudaSlice<u8>,
x: &CudaSlice<f32>,
parts: &[(&CudaSlice<f32>, usize)],
in_f: usize,
) -> Res<()> {
if in_f % 8 != 0 {
return Err("qmatvec_bf16w_multi4_f32: in_f % 8 != 0".into());
}
if parts.is_empty() || parts.len() > 4 {
return Err("qmatvec_bf16w_multi4_f32: 1..=4 parts".into());
}
let total: usize = parts.iter().map(|&(_, r)| r).sum();
if w_stack.len() < total * in_f * 2 {
return Err("qmatvec_bf16w_multi4_f32: stacked twin shorter than the row plan".into());
}
let stream = e.gpu.stream();
let mut ptrs = [0u64; 4];
let mut rows = [0i32; 4];
for (i, &(buf, r)) in parts.iter().enumerate() {
if buf.len() < r {
return Err("qmatvec_bf16w_multi4_f32: destination shorter than its rows".into());
}
ptrs[i] = buf.device_ptr(&stream).0;
rows[i] = r as i32;
}
let f = e.func("qmatvec_bf16w_multi4_f32");
let cfg = LaunchConfig {
grid_dim: (total as u32, 1, 1),
block_dim: (128, 1, 1),
shared_mem_bytes: 0,
};
let inf = in_f as i32;
let mut b = stream.launch_builder(&f);
b.arg(w_stack)
.arg(x)
.arg(&ptrs[0])
.arg(&rows[0])
.arg(&ptrs[1])
.arg(&rows[1])
.arg(&ptrs[2])
.arg(&rows[2])
.arg(&ptrs[3])
.arg(&rows[3])
.arg(&inf);
unsafe {
b.launch(cfg)?;
}
Ok(())
}
#[allow(clippy::too_many_arguments)]
fn linear_trunk_into(
e: &Engine,
w_f32: &CudaSlice<f32>,
w_b16: &Option<CudaSlice<u8>>,
x: &CudaSlice<f32>,
y: &mut CudaSlice<f32>,
t: usize,
in_f: usize,
out_f: usize,
) -> Res<()> {
if trunk_bf16_on() {
if let Some(w) = w_b16 {
if (2..=12).contains(&t) && verify_mt_on() {
return launch_qmatvec_bf16w_mt(e, w, 0, x, y, in_f, out_f, t);
}
return launch_qmatvec_bf16w(e, w, x, y, in_f, out_f, t, 1, 0, 0, in_f, 0);
}
}
if w_f32.len() < in_f * out_f {
return Err(
"qwen4exp_gpu: trunk f32 original dropped (trunk_f32_diet) — the bf16 \
twin path is required (keep trunk seams ON)"
.into(),
);
}
e.linear_device_into(x, w_f32, y, t, in_f, out_f)
}
#[allow(clippy::too_many_arguments)]
fn launch_nvfp4_sel_matvec(
e: &Engine,
codes: &CudaSlice<u8>,
scales: &CudaSlice<u8>,
macros_dev: &CudaSlice<f32>,
sel: &CudaSlice<i32>,
x: &CudaSlice<f32>,
y: &mut CudaSlice<f32>,
n_sel: usize,
in_f: usize,
out_f: usize,
x_stride: usize,
) -> Res<()> {
if in_f % 16 != 0 {
return Err("qmatvec_nvfp4_modelopt_sel_f32: in_f % 16 != 0".into());
}
if y.len() < n_sel * out_f {
return Err("qmatvec_nvfp4_modelopt_sel_f32: output shorter than n_sel*out_f".into());
}
let grp = sel_group_resolve(sel_group_dn(), in_f, out_f);
let v3 = grp.is_none() && sel_v3_on() && in_f % 32 == 0 && out_f % 4 == 0;
let v2 = grp.is_none() && !v3 && sel_v2_on() && in_f % 32 == 0 && out_f % 2 == 0;
let f = e.func(if grp.is_some() {
"qmatvec_nvfp4_modelopt_sel_g_f32"
} else if v3 {
"qmatvec_nvfp4_modelopt_sel_f32_v3"
} else if v2 {
"qmatvec_nvfp4_modelopt_sel_f32_v2"
} else {
"qmatvec_nvfp4_modelopt_sel_f32"
});
let grid_x = match grp {
Some((g, rows)) => out_f / ((32 / g) * rows),
None if v3 => out_f / 4,
None if v2 => out_f / 2,
None => out_f,
};
let cfg = LaunchConfig {
grid_dim: (grid_x as u32, n_sel as u32, 1),
block_dim: (32, 1, 1),
shared_mem_bytes: 0,
};
let (inf, outf) = (in_f as i32, out_f as i32);
let xs = x_stride as i64;
let (gi, ri) = grp.map_or((0i32, 0i32), |(g, rows)| (g as i32, rows as i32));
let stream = e.gpu.stream();
let mut b = stream.launch_builder(&f);
b.arg(codes)
.arg(scales)
.arg(macros_dev)
.arg(sel)
.arg(x)
.arg(y)
.arg(&inf)
.arg(&outf)
.arg(&xs);
if grp.is_some() {
b.arg(&gi).arg(&ri);
}
unsafe {
b.launch(cfg)?;
}
Ok(())
}
#[allow(clippy::too_many_arguments)]
fn launch_nvfp4_sel_gu_silu(
e: &Engine,
gate: (&CudaSlice<u8>, &CudaSlice<u8>, &CudaSlice<f32>),
up: (&CudaSlice<u8>, &CudaSlice<u8>, &CudaSlice<f32>),
sel: Option<&CudaSlice<i32>>,
pack_raw: u64,
n_sel: usize,
x: &CudaSlice<f32>,
act: &mut CudaSlice<f32>,
in_f: usize,
ff: usize,
tok: Option<(&CudaSlice<i32>, usize)>,
) -> Res<()> {
if in_f % 32 != 0 || ff % 4 != 0 {
return Err("qmatvec_nvfp4_modelopt_sel_gu_silu_f32: geometry".into());
}
if act.len() < n_sel * ff {
return Err("qmatvec_nvfp4_modelopt_sel_gu_silu_f32: act buffer too short".into());
}
if sel.is_none() == (pack_raw == 0) {
return Err("qmatvec_nvfp4_modelopt_sel_gu_silu_f32: exactly one of sel/pack".into());
}
let grp = sel_group_resolve(sel_group_gu(), in_f, ff);
let f = e.func(if grp.is_some() {
"qmatvec_nvfp4_modelopt_sel_gu_silu_g_f32"
} else {
"qmatvec_nvfp4_modelopt_sel_gu_silu_f32"
});
let grid_x = match grp {
Some((g, rows)) => ff / ((32 / g) * rows),
None => ff / 4,
};
let cfg = LaunchConfig {
grid_dim: (grid_x as u32, n_sel as u32, 1),
block_dim: (32, 1, 1),
shared_mem_bytes: 0,
};
let (inf, ffi, ms) = (in_f as i32, ff as i32, n_sel as i32);
let (gi, ri) = grp.map_or((0i32, 0i32), |(g, rows)| (g as i32, rows as i32));
let stream = e.gpu.stream();
let mut b = stream.launch_builder(&f);
b.arg(gate.0)
.arg(gate.1)
.arg(gate.2)
.arg(up.0)
.arg(up.1)
.arg(up.2);
match sel {
Some(s) => {
b.arg(s);
}
None => {
b.arg(gate.2);
}
}
let stream2 = e.gpu.stream();
let (tok_raw, x_tstride) = match tok {
Some((tm, stride)) => (tm.device_ptr(&stream2).0, stride as i64),
None => (0u64, 0i64),
};
b.arg(&pack_raw)
.arg(&ms)
.arg(x)
.arg(&mut *act)
.arg(&inf)
.arg(&ffi)
.arg(&tok_raw)
.arg(&x_tstride);
if grp.is_some() {
b.arg(&gi).arg(&ri);
}
unsafe {
b.launch(cfg)?;
}
Ok(())
}
#[derive(Debug, Clone, Copy)]
pub struct MoeUnionRow {
pub t: usize,
pub slots: usize,
pub union_size: usize,
pub gu_us: f64,
pub down_us: f64,
pub gu_spread_rel: f64,
pub down_spread_rel: f64,
}
pub fn moe_union_cost_probe(
e: &Engine,
experts: usize,
hidden: usize,
ff: usize,
selected: usize,
t: usize,
reps: usize,
) -> Res<Vec<MoeUnionRow>> {
if hidden % 32 != 0 || ff % 4 != 0 {
return Err("moe_union_cost_probe: needs the gufuse geometry (hidden%32, ff%4)".into());
}
if selected == 0 || t == 0 || reps == 0 {
return Err("moe_union_cost_probe: selected/t/reps must be non-zero".into());
}
let code_byte = |i: usize| -> u8 { (i.wrapping_mul(2_654_435_761) >> 13) as u8 };
let scale_byte = |i: usize| -> u8 { 0x38 | ((i.wrapping_mul(40_503) >> 7) & 0x07) as u8 };
let mk = |n: usize, f: &dyn Fn(usize) -> u8| -> Res<CudaSlice<u8>> {
let host: Vec<u8> = (0..n).map(f).collect();
let d = e.htod_bytes(&host)?;
drop(host);
Ok(d)
};
let gu_codes_n = experts * ff * (hidden / 2);
let gu_scales_n = experts * ff * (hidden / 16);
let dn_codes_n = experts * hidden * (ff / 2);
let dn_scales_n = experts * hidden * (ff / 16);
let gc = mk(gu_codes_n, &code_byte)?;
let gs = mk(gu_scales_n, &scale_byte)?;
let uc = mk(gu_codes_n, &|i| code_byte(i ^ 0x5A5A_5A5A))?;
let us = mk(gu_scales_n, &|i| scale_byte(i ^ 0x3C3C_3C3C))?;
let dc = mk(dn_codes_n, &|i| code_byte(i ^ 0x0F0F_0F0F))?;
let ds = mk(dn_scales_n, &|i| scale_byte(i ^ 0x1111_1111))?;
let gm = e.htod(&vec![1.0f32; experts])?;
let um = e.htod(&vec![1.0f32; experts])?;
let dm = e.htod(&vec![1.0f32; experts])?;
let mixed_h: Vec<f32> = (0..t * hidden)
.map(|i| ((i.wrapping_mul(40_503) % 1000) as f32) / 4000.0 - 0.125)
.collect();
let mixed = e.htod(&mixed_h)?;
let pool: Vec<i32> = {
let stride = (experts / (t * selected).max(1)).max(1);
(0..t * selected)
.map(|i| ((i * stride) % experts) as i32)
.collect()
};
let mut rows: Vec<MoeUnionRow> = Vec::new();
let mut cells: Vec<(usize, usize)> = vec![(1, selected)];
for new in 0..=selected {
cells.push((t, selected + (t - 1) * new));
}
struct Cell {
t: usize,
slots: usize,
union_size: usize,
sel: CudaSlice<i32>,
tokm: CudaSlice<i32>,
act: CudaSlice<f32>,
partial: CudaSlice<f32>,
gu: Vec<f64>,
dn: Vec<f64>,
}
let mut built: Vec<Cell> = Vec::with_capacity(cells.len());
for (cells_t, want_union) in cells {
let slots = cells_t * selected;
let new = if cells_t > 1 {
(want_union - selected) / (cells_t - 1)
} else {
0
};
let shared = selected - new;
let mut sel_h: Vec<i32> = Vec::with_capacity(slots);
let mut tok_h: Vec<i32> = Vec::with_capacity(slots);
let mut fresh = selected;
for col in 0..cells_t {
if col == 0 {
sel_h.extend_from_slice(&pool[0..selected]);
} else {
sel_h.extend_from_slice(&pool[0..shared]);
for _ in 0..new {
sel_h.push(pool[fresh % pool.len()]);
fresh += 1;
}
}
for _ in 0..selected {
tok_h.push(col as i32);
}
}
let union_size = {
let mut u: Vec<i32> = sel_h.clone();
u.sort_unstable();
u.dedup();
u.len()
};
built.push(Cell {
t: cells_t,
slots,
union_size,
sel: e.htod_i32(&sel_h)?,
tokm: e.htod_i32(&tok_h)?,
act: e.zeros(slots * ff)?,
partial: e.zeros(slots * hidden)?,
gu: Vec::with_capacity(reps),
dn: Vec::with_capacity(reps),
});
}
for rep in 0..(reps + 1) {
for c in built.iter_mut() {
let tok_arg = if c.t > 1 {
Some((&c.tokm, hidden))
} else {
None
};
e.stream().synchronize()?;
let t0 = std::time::Instant::now();
launch_nvfp4_sel_gu_silu(
e,
(&gc, &gs, &gm),
(&uc, &us, &um),
Some(&c.sel),
0,
c.slots,
&mixed,
&mut c.act,
hidden,
ff,
tok_arg,
)?;
e.stream().synchronize()?;
let t1 = std::time::Instant::now();
launch_nvfp4_sel_matvec(
e,
&dc,
&ds,
&dm,
&c.sel,
&c.act,
&mut c.partial,
c.slots,
ff,
hidden,
ff,
)?;
e.stream().synchronize()?;
let t2 = std::time::Instant::now();
if rep > 0 {
c.gu.push(t1.duration_since(t0).as_secs_f64() * 1e6);
c.dn.push(t2.duration_since(t1).as_secs_f64() * 1e6);
}
}
}
let stat = |v: &[f64]| -> (f64, f64) {
let mut s = v.to_vec();
s.sort_by(|a, b| a.partial_cmp(b).unwrap());
let med = s[s.len() / 2];
let spread = if med > 0.0 {
(s[s.len() - 1] - s[0]) / med
} else {
0.0
};
(med, spread)
};
for c in &built {
let (gu_us, gu_spread) = stat(&c.gu);
let (down_us, down_spread) = stat(&c.dn);
rows.push(MoeUnionRow {
t: c.t,
slots: c.slots,
union_size: c.union_size,
gu_us,
down_us,
gu_spread_rel: gu_spread,
down_spread_rel: down_spread,
});
}
Ok(rows)
}
#[derive(Debug, Clone)]
pub struct SelShapeRow {
pub t: usize,
pub slots: usize,
pub arm: String,
pub gu_shape: String,
pub dn_shape: String,
pub gu_grid_x: usize,
pub dn_grid_x: usize,
pub gu_us: f64,
pub down_us: f64,
pub gu_spread_rel: f64,
pub down_spread_rel: f64,
}
#[allow(clippy::too_many_arguments)]
pub fn sel_shape_cost_probe(
e: &Engine,
experts: usize,
hidden: usize,
ff: usize,
selected: usize,
t: usize,
reps: usize,
arms: &[String],
) -> Res<Vec<SelShapeRow>> {
if hidden % 32 != 0 || ff % 4 != 0 {
return Err("sel_shape_cost_probe: needs the gufuse geometry (hidden%32, ff%4)".into());
}
if selected == 0 || t == 0 || reps == 0 || arms.is_empty() {
return Err("sel_shape_cost_probe: selected/t/reps/arms must be non-empty".into());
}
let saved = sel_group_spec();
let out = sel_shape_cost_probe_inner(e, experts, hidden, ff, selected, t, reps, arms);
set_sel_group(&saved);
out
}
#[allow(clippy::too_many_arguments)]
fn sel_shape_cost_probe_inner(
e: &Engine,
experts: usize,
hidden: usize,
ff: usize,
selected: usize,
t: usize,
reps: usize,
arms: &[String],
) -> Res<Vec<SelShapeRow>> {
let code_byte = |i: usize| -> u8 { (i.wrapping_mul(2_654_435_761) >> 13) as u8 };
let scale_byte = |i: usize| -> u8 { 0x38 | ((i.wrapping_mul(40_503) >> 7) & 0x07) as u8 };
let mk = |n: usize, f: &dyn Fn(usize) -> u8| -> Res<CudaSlice<u8>> {
let host: Vec<u8> = (0..n).map(f).collect();
let d = e.htod_bytes(&host)?;
drop(host);
Ok(d)
};
let gc = mk(experts * ff * (hidden / 2), &code_byte)?;
let gs = mk(experts * ff * (hidden / 16), &scale_byte)?;
let uc = mk(experts * ff * (hidden / 2), &|i| code_byte(i ^ 0x5A5A_5A5A))?;
let us = mk(experts * ff * (hidden / 16), &|i| {
scale_byte(i ^ 0x3C3C_3C3C)
})?;
let dc = mk(experts * hidden * (ff / 2), &|i| code_byte(i ^ 0x0F0F_0F0F))?;
let ds = mk(experts * hidden * (ff / 16), &|i| {
scale_byte(i ^ 0x1111_1111)
})?;
let gm = e.htod(&vec![1.0f32; experts])?;
let um = e.htod(&vec![1.0f32; experts])?;
let dm = e.htod(&vec![1.0f32; experts])?;
let mixed_h: Vec<f32> = (0..t * hidden)
.map(|i| ((i.wrapping_mul(40_503) % 1000) as f32) / 4000.0 - 0.125)
.collect();
let mixed = e.htod(&mixed_h)?;
let slots = t * selected;
let stride = (experts / slots.max(1)).max(1);
let sel_h: Vec<i32> = (0..slots)
.map(|i| ((i * stride) % experts) as i32)
.collect();
let tok_h: Vec<i32> = (0..slots).map(|i| (i / selected) as i32).collect();
let sel = e.htod_i32(&sel_h)?;
let tokm = e.htod_i32(&tok_h)?;
let mut act = e.zeros(slots * ff)?;
let mut partial = e.zeros(slots * hidden)?;
struct Arm {
spec: String,
gu_shape: String,
dn_shape: String,
gu_grid_x: usize,
dn_grid_x: usize,
gu: Vec<f64>,
dn: Vec<f64>,
}
let describe = |code: u32, in_f: usize, out_f: usize| -> (String, usize) {
match sel_group_resolve(code, in_f, out_f) {
Some((g, rows)) => {
let rpw = (32 / g) * rows;
(format!("g{g}r{rows}/rpw{rpw}"), out_f / rpw)
}
None => ("shipped".to_string(), out_f / 4),
}
};
let mut built: Vec<Arm> = Vec::with_capacity(arms.len());
for spec in arms {
if !set_sel_group(spec) {
return Err(format!("sel_shape_cost_probe: bad arm spec {spec:?}").into());
}
let (gu_shape, gu_grid_x) = describe(sel_group_gu(), hidden, ff);
let (dn_shape, dn_grid_x) = describe(sel_group_dn(), ff, hidden);
built.push(Arm {
spec: spec.clone(),
gu_shape,
dn_shape,
gu_grid_x,
dn_grid_x,
gu: Vec::with_capacity(reps),
dn: Vec::with_capacity(reps),
});
}
let tok_arg = if t > 1 { Some((&tokm, hidden)) } else { None };
for rep in 0..(reps + 1) {
for a in built.iter_mut() {
set_sel_group(&a.spec);
e.stream().synchronize()?;
let t0 = std::time::Instant::now();
launch_nvfp4_sel_gu_silu(
e,
(&gc, &gs, &gm),
(&uc, &us, &um),
Some(&sel),
0,
slots,
&mixed,
&mut act,
hidden,
ff,
tok_arg,
)?;
e.stream().synchronize()?;
let t1 = std::time::Instant::now();
launch_nvfp4_sel_matvec(
e,
&dc,
&ds,
&dm,
&sel,
&act,
&mut partial,
slots,
ff,
hidden,
ff,
)?;
e.stream().synchronize()?;
let t2 = std::time::Instant::now();
if rep > 0 {
a.gu.push(t1.duration_since(t0).as_secs_f64() * 1e6);
a.dn.push(t2.duration_since(t1).as_secs_f64() * 1e6);
}
}
}
let stat = |v: &[f64]| -> (f64, f64) {
let mut s = v.to_vec();
s.sort_by(|a, b| a.partial_cmp(b).unwrap());
let med = s[s.len() / 2];
let spread = if med > 0.0 {
(s[s.len() - 1] - s[0]) / med
} else {
0.0
};
(med, spread)
};
Ok(built
.iter()
.map(|a| {
let (gu_us, gu_spread_rel) = stat(&a.gu);
let (down_us, down_spread_rel) = stat(&a.dn);
SelShapeRow {
t,
slots,
arm: a.spec.clone(),
gu_shape: a.gu_shape.clone(),
dn_shape: a.dn_shape.clone(),
gu_grid_x: a.gu_grid_x,
dn_grid_x: a.dn_grid_x,
gu_us,
down_us,
gu_spread_rel,
down_spread_rel,
}
})
.collect())
}
#[allow(clippy::too_many_arguments)]
fn launch_axpy_rows_seq_at(
e: &Engine,
x: &CudaSlice<f32>,
x_row0: usize,
w: &CudaSlice<f32>,
w_off: usize,
y: &mut CudaSlice<f32>,
y_row: usize,
width: usize,
n_rows: usize,
) -> Res<()> {
if x.len() < (x_row0 + n_rows) * width
|| w.len() < w_off + n_rows
|| y.len() < (y_row + 1) * width
{
return Err("axpy_rows_seq_f32: window out of range".into());
}
let xv = x.slice(x_row0 * width..(x_row0 + n_rows) * width);
let wv = w.slice(w_off..w_off + n_rows);
let mut yv = y.slice_mut(y_row * width..(y_row + 1) * width);
let f = e.func("axpy_rows_seq_f32");
let cfg = LaunchConfig::for_num_elems(width as u32);
let (wi, nr) = (width as i32, n_rows as i32);
let stream = e.gpu.stream();
let mut b = stream.launch_builder(&f);
b.arg(&xv).arg(&wv).arg(&mut yv).arg(&wi).arg(&nr);
unsafe {
b.launch(cfg)?;
}
Ok(())
}
pub fn gate_nvfp4_sel_matvec(e: &Engine) -> Res<String> {
let mut lcg = 0x2545_f491_u64;
let mut next_u32 = move || -> u32 {
lcg = lcg
.wrapping_mul(6364136223846793005)
.wrapping_add(1442695040888963407);
(lcg >> 33) as u32
};
let macros = [
1.0f32,
0.5,
5.9945243e-5, 2.0,
0.25,
3.7e-3,
1.0,
8.0,
];
let sel_host: Vec<i32> = vec![3, 5, 3, 0]; let n_sel = sel_host.len();
let mut worst = (0.0f32, 0.0f32); for (mode, out_f, in_f) in [
("gate_up_v1", 16usize, 48usize),
("down_v1", 32, 16),
("gate_up_v1_oddrows", 7, 64),
("gate_up_v2", 16, 64),
("down_v2", 32, 32),
("gate_up_v3", 16, 64),
("down_v3", 32, 32),
("gate_up_v3_v2rows", 6, 64), ] {
set_sel_v3(mode.contains("v3"));
let n_expert = macros.len();
let mut codes = vec![0u8; n_expert * out_f * in_f / 2];
for byte in &mut codes {
*byte = next_u32() as u8;
}
let mut scales = vec![0u8; n_expert * out_f * in_f / 16];
for byte in &mut scales {
*byte = (next_u32() as u8) & 0xBF; }
scales[0] = 0x7F; scales[3] = 0xFF; let x_stride = if mode.starts_with("down") { in_f } else { 0 };
let x_rows = if x_stride == 0 { 1 } else { n_sel };
let x_host: Vec<f32> = (0..x_rows * in_f)
.map(|_| (next_u32() % 2000) as f32 / 1000.0 - 1.0)
.collect();
let codes_dev = e.htod_bytes(&codes)?;
let scales_dev = e.htod_bytes(&scales)?;
let macros_dev = e.htod(¯os)?;
let sel_dev = e.htod_i32(&sel_host)?;
let x_dev = e.htod(&x_host)?;
let mut y_dev = e.uninit(n_sel * out_f)?;
launch_nvfp4_sel_matvec(
e,
&codes_dev,
&scales_dev,
¯os_dev,
&sel_dev,
&x_dev,
&mut y_dev,
n_sel,
in_f,
out_f,
x_stride,
)?;
let y = e.dtoh(&y_dev)?;
let wbytes = out_f * in_f / 2;
let sbytes = out_f * in_f / 16;
for (slot, &expert) in sel_host.iter().enumerate() {
let expert = expert as usize;
let w = memra_gguf::dsv4::dequant_nvfp4_expert(
&codes[expert * wbytes..(expert + 1) * wbytes],
&scales[expert * sbytes..(expert + 1) * sbytes],
macros[expert],
out_f,
in_f,
);
let xrow = &x_host[slot * x_stride..slot * x_stride + in_f];
for o in 0..out_f {
let mut want = 0.0f32;
for i in 0..in_f {
want += w[o * in_f + i] * xrow[i];
}
let got = y[slot * out_f + o];
let abs = (want - got).abs();
let rel = abs / want.abs().max(1.0);
if abs > worst.0 {
worst.0 = abs;
}
if rel > worst.1 {
worst.1 = rel;
}
if rel > 1e-5 {
return Err(format!(
"nvfp4-sel-matvec oracle: {mode} slot {slot} row {o}: want {want} \
got {got} (rel {rel:.3e})"
)
.into());
}
}
}
}
set_sel_v3(SEL_V3_DEFAULT);
{
set_sel_v3(true);
let (ff, in_f) = (16usize, 64usize);
let n_expert = macros.len();
let mut mk = |seed: u8| -> (Vec<u8>, Vec<u8>) {
let mut codes = vec![0u8; n_expert * ff * in_f / 2];
for byte in &mut codes {
*byte = (next_u32() as u8) ^ seed;
}
let mut scales = vec![0u8; n_expert * ff * in_f / 16];
for byte in &mut scales {
*byte = (next_u32() as u8) & 0xBF;
}
scales[1] = 0x7F; (codes, scales)
};
let (g_codes, g_scales) = mk(0x00);
let (u_codes, u_scales) = mk(0x5A);
let gmac: Vec<f32> = macros.to_vec();
let umac: Vec<f32> = macros.iter().map(|m| m * 0.5).collect();
let x_host: Vec<f32> = (0..in_f)
.map(|_| (next_u32() % 2000) as f32 / 1000.0 - 1.0)
.collect();
let gc = e.htod_bytes(&g_codes)?;
let gs = e.htod_bytes(&g_scales)?;
let gm = e.htod(&gmac)?;
let uc = e.htod_bytes(&u_codes)?;
let us = e.htod_bytes(&u_scales)?;
let um = e.htod(&umac)?;
let sel_dev = e.htod_i32(&sel_host)?;
let x_dev = e.htod(&x_host)?;
let mut yg = e.uninit(n_sel * ff)?;
let mut yu = e.uninit(n_sel * ff)?;
launch_nvfp4_sel_matvec(
e, &gc, &gs, &gm, &sel_dev, &x_dev, &mut yg, n_sel, in_f, ff, 0,
)?;
launch_nvfp4_sel_matvec(
e, &uc, &us, &um, &sel_dev, &x_dev, &mut yu, n_sel, in_f, ff, 0,
)?;
let mut act_chain = e.zeros(n_sel * ff)?;
e.silu_mul(&yg, &yu, &mut act_chain, n_sel * ff)?;
let mut act_fused = e.zeros(n_sel * ff)?;
launch_nvfp4_sel_gu_silu(
e,
(&gc, &gs, &gm),
(&uc, &us, &um),
Some(&sel_dev),
0,
n_sel,
&x_dev,
&mut act_fused,
in_f,
ff,
None,
)?;
let a = e.dtoh(&act_chain)?;
let b = e.dtoh(&act_fused)?;
for (i, (&x1, &x2)) in a.iter().zip(&b).enumerate() {
if x1.to_bits() != x2.to_bits() {
return Err(format!(
"nvfp4-sel-matvec oracle: gufuse idx {i} not bit-identical \
(chain {x1} fused {x2})"
)
.into());
}
}
let pack_bytes = tp2_pack_bytes(&sel_host[..2], &[0.5, 0.25], n_sel);
let pack = e.htod_bytes(&pack_bytes)?;
let pack_raw = {
let stream = e.gpu.stream();
pack.device_ptr(&stream).0
};
let sentinel = vec![-777.0f32; n_sel * ff];
let mut act_pack = e.htod(&sentinel)?;
launch_nvfp4_sel_gu_silu(
e,
(&gc, &gs, &gm),
(&uc, &us, &um),
None,
pack_raw,
n_sel,
&x_dev,
&mut act_pack,
in_f,
ff,
None,
)?;
let c = e.dtoh(&act_pack)?;
for slot in 0..n_sel {
for o in 0..ff {
let got = c[slot * ff + o];
if slot < 2 {
if got.to_bits() != a[slot * ff + o].to_bits() {
return Err(format!(
"nvfp4-sel-matvec oracle: gufuse pack slot {slot} o {o} \
not bit-identical"
)
.into());
}
} else if got != -777.0 {
return Err(format!(
"nvfp4-sel-matvec oracle: gufuse pack dead slot {slot} written"
)
.into());
}
}
}
{
let t2 = 2usize;
let x2_host: Vec<f32> = (0..t2 * in_f)
.map(|_| (next_u32() % 2000) as f32 / 1000.0 - 1.0)
.collect();
let x2 = e.htod(&x2_host)?;
let tok_host: Vec<i32> = (0..n_sel).map(|s| (s % t2) as i32).collect();
let tokm = e.htod_i32(&tok_host)?;
let mut act_map = e.zeros(n_sel * ff)?;
launch_nvfp4_sel_gu_silu(
e,
(&gc, &gs, &gm),
(&uc, &us, &um),
Some(&sel_dev),
0,
n_sel,
&x2,
&mut act_map,
in_f,
ff,
Some((&tokm, in_f)),
)?;
let got = e.dtoh(&act_map)?;
for tok in 0..t2 {
let slots: Vec<usize> = (0..n_sel).filter(|s| s % t2 == tok).collect();
let sel_tok: Vec<i32> = slots.iter().map(|&s| sel_host[s]).collect();
let sel_tok_dev = e.htod_i32(&sel_tok)?;
let xrow = e.htod(&x2_host[tok * in_f..(tok + 1) * in_f])?;
let mut act_tok = e.zeros(sel_tok.len() * ff)?;
launch_nvfp4_sel_gu_silu(
e,
(&gc, &gs, &gm),
(&uc, &us, &um),
Some(&sel_tok_dev),
0,
sel_tok.len(),
&xrow,
&mut act_tok,
in_f,
ff,
None,
)?;
let want = e.dtoh(&act_tok)?;
for (local, &slot) in slots.iter().enumerate() {
for o in 0..ff {
let a = got[slot * ff + o];
let b = want[local * ff + o];
if a.to_bits() != b.to_bits() {
return Err(format!(
"nvfp4-sel-matvec oracle: gufuse tok_map slot {slot} o {o} \
not bit-identical (map {a} per-token {b})"
)
.into());
}
}
}
}
}
set_sel_v3(SEL_V3_DEFAULT);
}
Ok(format!(
"nvfp4-sel-matvec kernel oracle: worst abs {:.3e} rel {:.3e} over gate_up+down \
v1/v2/v3 modes, NaN scales + non-pow2 macros + duplicate slots; gufuse \
BIT-IDENTICAL to the v3+silu chain incl. the count-gated pack twin + the \
tok_map verify merge",
worst.0, worst.1
))
}
pub fn gate_nvfp4_sel_group(e: &Engine) -> Res<String> {
let saved = sel_group_spec();
let out = gate_nvfp4_sel_group_inner(e);
set_sel_group(&saved);
set_sel_v3(SEL_V3_DEFAULT);
out
}
fn gate_nvfp4_sel_group_inner(e: &Engine) -> Res<String> {
let mut lcg = 0x2545_f491_u64; let mut next_u32 = move || -> u32 {
lcg = lcg
.wrapping_mul(6364136223846793005)
.wrapping_add(1442695040888963407);
(lcg >> 33) as u32
};
let macros = [
1.0f32,
0.5,
5.9945243e-5, 2.0,
0.25,
3.7e-3,
1.0,
8.0,
];
let n_expert = macros.len();
let sel_host: Vec<i32> = vec![3, 5, 3, 0]; let n_sel = sel_host.len();
let mut worst = (0.0f32, 0.0f32);
let mut shapes_checked = 0usize;
let mut bits_checked = 0usize;
let mut calib: Vec<String> = Vec::new();
for (geom, out_f, in_f, per_slot_x) in [
("down_real", 2560usize, 640usize, true),
("gateup_real", 640, 2560, false),
("down_tiny", 32, 32, true),
("gateup_tiny", 16, 64, false),
] {
let mut codes = vec![0u8; n_expert * out_f * in_f / 2];
for byte in &mut codes {
*byte = next_u32() as u8;
}
let mut scales = vec![0u8; n_expert * out_f * in_f / 16];
for byte in &mut scales {
*byte = (next_u32() as u8) & 0xBF; }
scales[0] = 0x7F; scales[3] = 0xFF; let x_stride = if per_slot_x { in_f } else { 0 };
let x_rows = if per_slot_x { n_sel } else { 1 };
let x_host: Vec<f32> = (0..x_rows * in_f)
.map(|_| (next_u32() % 2000) as f32 / 1000.0 - 1.0)
.collect();
let codes_dev = e.htod_bytes(&codes)?;
let scales_dev = e.htod_bytes(&scales)?;
let macros_dev = e.htod(¯os)?;
let sel_dev = e.htod_i32(&sel_host)?;
let x_dev = e.htod(&x_host)?;
let run = |spec: &str| -> Res<Vec<f32>> {
set_sel_group(spec);
let mut y = e.uninit(n_sel * out_f)?;
launch_nvfp4_sel_matvec(
e,
&codes_dev,
&scales_dev,
¯os_dev,
&sel_dev,
&x_dev,
&mut y,
n_sel,
in_f,
out_f,
x_stride,
)?;
e.dtoh(&y)
};
set_sel_v3(true);
let shipped = run("off")?;
let wbytes = out_f * in_f / 2;
let sbytes = out_f * in_f / 16;
let mut want = vec![0.0f32; n_sel * out_f];
for (slot, &expert) in sel_host.iter().enumerate() {
let expert = expert as usize;
let w = memra_gguf::dsv4::dequant_nvfp4_expert(
&codes[expert * wbytes..(expert + 1) * wbytes],
&scales[expert * sbytes..(expert + 1) * sbytes],
macros[expert],
out_f,
in_f,
);
let xrow = &x_host[slot * x_stride..slot * x_stride + in_f];
for o in 0..out_f {
let mut acc = 0.0f32;
for i in 0..in_f {
acc += w[o * in_f + i] * xrow[i];
}
want[slot * out_f + o] = acc;
}
}
let ship_vs_host = want
.iter()
.zip(&shipped)
.map(|(&w, &s)| (w - s).abs() / w.abs().max(1.0))
.fold(0.0f32, f32::max);
let class_tol = (4.0 * ship_vs_host).max(1e-5);
calib.push(format!(
"{geom} ship_vs_host={ship_vs_host:.3e} tol={class_tol:.3e}"
));
for spec in [
"dn:32:4", "dn:auto", "dn:16:4", "dn:16:2", "dn:8:4", "dn:8:2", "dn:8:1", "dn:4:4",
"dn:4:2", "dn:4:1", "dn:2:4", "dn:2:2", "dn:2:1", "dn:1:1",
] {
let Some((g, rows)) = sel_group_resolve(
match spec {
"dn:auto" => SEL_GROUP_AUTO,
_ => {
let (gs, rs) = spec.trim_start_matches("dn:").split_once(':').unwrap();
(gs.parse::<u32>().unwrap() << 8) | rs.parse::<u32>().unwrap()
}
},
in_f,
out_f,
) else {
continue; };
let got = run(spec)?;
shapes_checked += 1;
if (g, rows) == (32, 4) {
for (i, (&a, &b)) in shipped.iter().zip(&got).enumerate() {
if a.to_bits() != b.to_bits() {
return Err(format!(
"sel-group oracle: {geom} g=32 rows=4 idx {i} NOT bit-identical to \
the shipped v3 kernel (v3 {a} group {b}) — the sub-warp form must \
degenerate to v3 exactly"
)
.into());
}
}
bits_checked += shipped.len();
}
for (i, (&w, &got)) in want.iter().zip(&got).enumerate() {
let abs = (w - got).abs();
let rel = abs / w.abs().max(1.0);
worst.0 = worst.0.max(abs);
worst.1 = worst.1.max(rel);
if rel > class_tol {
return Err(format!(
"sel-group oracle: {geom} {spec} (g={g} rows={rows}) idx {i} vs HOST \
chain: want {w} got {got} (rel {rel:.3e} > tol {class_tol:.3e}, \
shipped v3 itself is {ship_vs_host:.3e})"
)
.into());
}
}
for (i, (&s, &got)) in shipped.iter().zip(&got).enumerate() {
let rel = (s - got).abs() / s.abs().max(1.0);
if rel > class_tol {
return Err(format!(
"sel-group oracle: {geom} {spec} (g={g} rows={rows}) idx {i} vs SHIPPED \
v3: v3 {s} group {got} (rel {rel:.3e} > tol {class_tol:.3e})"
)
.into());
}
}
}
set_sel_group("off");
}
for (geom, ff, in_f) in [("gu_real", 640usize, 2560usize), ("gu_tiny", 16, 64)] {
let mut mk = |seed: u8| -> (Vec<u8>, Vec<u8>) {
let mut codes = vec![0u8; n_expert * ff * in_f / 2];
for byte in &mut codes {
*byte = (next_u32() as u8) ^ seed;
}
let mut scales = vec![0u8; n_expert * ff * in_f / 16];
for byte in &mut scales {
*byte = (next_u32() as u8) & 0xBF;
}
scales[1] = 0x7F; (codes, scales)
};
let (g_codes, g_scales) = mk(0x00);
let (u_codes, u_scales) = mk(0x5A);
let gmac: Vec<f32> = macros.to_vec();
let umac: Vec<f32> = macros.iter().map(|m| m * 0.5).collect();
let x_host: Vec<f32> = (0..in_f)
.map(|_| (next_u32() % 2000) as f32 / 1000.0 - 1.0)
.collect();
let gc = e.htod_bytes(&g_codes)?;
let gs = e.htod_bytes(&g_scales)?;
let gm = e.htod(&gmac)?;
let uc = e.htod_bytes(&u_codes)?;
let us = e.htod_bytes(&u_scales)?;
let um = e.htod(&umac)?;
let sel_dev = e.htod_i32(&sel_host)?;
let x_dev = e.htod(&x_host)?;
set_sel_group("off");
set_sel_v3(true);
let shipped_fused = {
let mut act = e.zeros(n_sel * ff)?;
launch_nvfp4_sel_gu_silu(
e,
(&gc, &gs, &gm),
(&uc, &us, &um),
Some(&sel_dev),
0,
n_sel,
&x_dev,
&mut act,
in_f,
ff,
None,
)?;
e.dtoh(&act)?
};
for spec in ["32:4", "auto", "16:4", "16:2", "8:4", "8:1", "4:4"] {
let Some((g, rows)) = sel_group_resolve(
match spec {
"auto" => SEL_GROUP_AUTO,
_ => {
let (gs, rs) = spec.split_once(':').unwrap();
(gs.parse::<u32>().unwrap() << 8) | rs.parse::<u32>().unwrap()
}
},
in_f,
ff,
) else {
continue;
};
set_sel_group(&format!("dn:{spec}+gu:off"));
let mut yg = e.uninit(n_sel * ff)?;
let mut yu = e.uninit(n_sel * ff)?;
launch_nvfp4_sel_matvec(
e, &gc, &gs, &gm, &sel_dev, &x_dev, &mut yg, n_sel, in_f, ff, 0,
)?;
launch_nvfp4_sel_matvec(
e, &uc, &us, &um, &sel_dev, &x_dev, &mut yu, n_sel, in_f, ff, 0,
)?;
let mut act_chain = e.zeros(n_sel * ff)?;
e.silu_mul(&yg, &yu, &mut act_chain, n_sel * ff)?;
let chain = e.dtoh(&act_chain)?;
set_sel_group(&format!("dn:off+gu:{spec}"));
let mut act_fused = e.zeros(n_sel * ff)?;
launch_nvfp4_sel_gu_silu(
e,
(&gc, &gs, &gm),
(&uc, &us, &um),
Some(&sel_dev),
0,
n_sel,
&x_dev,
&mut act_fused,
in_f,
ff,
None,
)?;
let fused = e.dtoh(&act_fused)?;
for (i, (&a, &b)) in chain.iter().zip(&fused).enumerate() {
if a.to_bits() != b.to_bits() {
return Err(format!(
"sel-group oracle: {geom} gu {spec} (g={g} rows={rows}) idx {i} fused \
NOT bit-identical to the same-shape chain (chain {a} fused {b})"
)
.into());
}
}
bits_checked += chain.len();
shapes_checked += 1;
if (g, rows) == (32, 4) {
for (i, (&a, &b)) in shipped_fused.iter().zip(&fused).enumerate() {
if a.to_bits() != b.to_bits() {
return Err(format!(
"sel-group oracle: {geom} gu g=32 rows=4 idx {i} NOT bit-identical \
to the shipped gufuse kernel (gufuse {a} group {b})"
)
.into());
}
}
bits_checked += shipped_fused.len();
}
}
set_sel_group("dn:off+gu:auto");
if sel_group_resolve(SEL_GROUP_AUTO, in_f, ff).is_some() {
let auto_plain = {
let mut act = e.zeros(n_sel * ff)?;
launch_nvfp4_sel_gu_silu(
e,
(&gc, &gs, &gm),
(&uc, &us, &um),
Some(&sel_dev),
0,
n_sel,
&x_dev,
&mut act,
in_f,
ff,
None,
)?;
e.dtoh(&act)?
};
let pack_bytes = tp2_pack_bytes(&sel_host[..2], &[0.5, 0.25], n_sel);
let pack = e.htod_bytes(&pack_bytes)?;
let pack_raw = {
let stream = e.gpu.stream();
pack.device_ptr(&stream).0
};
let sentinel = vec![-777.0f32; n_sel * ff];
let mut act_pack = e.htod(&sentinel)?;
launch_nvfp4_sel_gu_silu(
e,
(&gc, &gs, &gm),
(&uc, &us, &um),
None,
pack_raw,
n_sel,
&x_dev,
&mut act_pack,
in_f,
ff,
None,
)?;
let packed = e.dtoh(&act_pack)?;
for slot in 0..n_sel {
for o in 0..ff {
let got = packed[slot * ff + o];
if slot < 2 {
if got.to_bits() != auto_plain[slot * ff + o].to_bits() {
return Err(format!(
"sel-group oracle: {geom} gu auto pack slot {slot} o {o} not \
bit-identical to the sel-array arm"
)
.into());
}
} else if got != -777.0 {
return Err(format!(
"sel-group oracle: {geom} gu auto pack dead slot {slot} written"
)
.into());
}
}
}
let t2 = 2usize;
let x2_host: Vec<f32> = (0..t2 * in_f)
.map(|_| (next_u32() % 2000) as f32 / 1000.0 - 1.0)
.collect();
let x2 = e.htod(&x2_host)?;
let tok_host: Vec<i32> = (0..n_sel).map(|s| (s % t2) as i32).collect();
let tokm = e.htod_i32(&tok_host)?;
let mut act_map = e.zeros(n_sel * ff)?;
launch_nvfp4_sel_gu_silu(
e,
(&gc, &gs, &gm),
(&uc, &us, &um),
Some(&sel_dev),
0,
n_sel,
&x2,
&mut act_map,
in_f,
ff,
Some((&tokm, in_f)),
)?;
let mapped = e.dtoh(&act_map)?;
for tok in 0..t2 {
let slots: Vec<usize> = (0..n_sel).filter(|s| s % t2 == tok).collect();
let sel_tok: Vec<i32> = slots.iter().map(|&s| sel_host[s]).collect();
let sel_tok_dev = e.htod_i32(&sel_tok)?;
let xrow = e.htod(&x2_host[tok * in_f..(tok + 1) * in_f])?;
let mut act_tok = e.zeros(sel_tok.len() * ff)?;
launch_nvfp4_sel_gu_silu(
e,
(&gc, &gs, &gm),
(&uc, &us, &um),
Some(&sel_tok_dev),
0,
sel_tok.len(),
&xrow,
&mut act_tok,
in_f,
ff,
None,
)?;
let want = e.dtoh(&act_tok)?;
for (local, &slot) in slots.iter().enumerate() {
for o in 0..ff {
let a = mapped[slot * ff + o];
let b = want[local * ff + o];
if a.to_bits() != b.to_bits() {
return Err(format!(
"sel-group oracle: {geom} gu auto tok_map slot {slot} o {o} not \
bit-identical (map {a} per-token {b})"
)
.into());
}
}
}
}
bits_checked += auto_plain.len() + mapped.len();
}
set_sel_group("off");
}
Ok(format!(
"nvfp4-sel-GROUP kernel oracle: {shapes_checked} (geometry, shape) cells over REAL \
MoE geometry (down 2560x640 pairs=20, gate_up 640x2560 pairs=80) + tiny, worst abs \
{:.3e} rel {:.3e} vs the host decoder chain; (g=32,rows=4) BIT-IDENTICAL to the \
shipped v3 and gufuse kernels and every shape's fused arm BIT-IDENTICAL to its \
same-shape chain ({bits_checked} f32 byte-compared), incl. the count-gated pack \
twin + the tok_map verify merge; NaN scales + non-pow2 macros + duplicate slots; \
per-geometry class calibration [{}]",
worst.0,
worst.1,
calib.join("; ")
))
}
pub fn gate_hc_diet_kernels(e: &Engine) -> Res<String> {
let (streams, hidden, rank, t) = (4usize, 2560usize, 320usize, 1usize);
let wide = streams * hidden;
let mut lcg = 0x8badf00d_u64;
let mut next_f32 = move || -> f32 {
lcg = lcg
.wrapping_mul(6364136223846793005)
.wrapping_add(1442695040888963407);
(((lcg >> 33) as u32) % 2000) as f32 / 1000.0 - 1.0
};
let mut rand_vec = |n: usize| -> Vec<f32> { (0..n).map(|_| next_f32()).collect() };
let to_b16_vals = |v: Vec<f32>| -> Vec<f32> {
v.into_iter()
.map(|x| f32::from_bits(x.to_bits() & 0xFFFF_0000))
.collect()
};
let planes_host: Vec<Vec<f32>> = (0..streams).map(|_| rand_vec(t * hidden)).collect();
let planes: Vec<CudaSlice<f32>> = planes_host
.iter()
.map(|v| e.htod(v))
.collect::<Result<_, _>>()?;
let ptr_vals: Vec<u64> = {
let stream = e.gpu.stream();
planes.iter().map(|p| p.device_ptr(&stream).0).collect()
};
let ptrs = e.htod_u64(&ptr_vals)?;
let norm_stack_host = rand_vec(wide);
let norm_stack = e.htod(&norm_stack_host)?;
let down_host = to_b16_vals(rand_vec(streams * rank * hidden));
let up_host = to_b16_vals(rand_vec(streams * hidden * rank));
let inj_host = to_b16_vals(rand_vec(streams * wide));
let down_b16 = bf16_twin(e, &down_host, hidden)?.ok_or("hc-diet oracle: down twin")?;
let up_b16 = bf16_twin(e, &up_host, rank)?.ok_or("hc-diet oracle: up twin")?;
let inj_b16 = bf16_twin(e, &inj_host, hidden)?.ok_or("hc-diet oracle: inject twin")?;
let inj_f32 = e.htod(&inj_host)?;
let eps = 1e-6f32;
let mut normed = e.zeros(streams * t * hidden)?;
launch_hc_norm_planes(e, &ptrs, &norm_stack, &mut normed, hidden, t, streams, eps)?;
let mut parts_c = e.zeros(streams * t * rank)?;
launch_qmatvec_bf16w(
e,
&down_b16,
&normed,
&mut parts_c,
hidden,
rank,
t,
streams,
rank * hidden,
t * hidden,
hidden,
t * rank,
)?;
let mut low_c = e.zeros(t * rank)?;
launch_hc_lowrank_reduce(e, &parts_c, &mut low_c, streams, t, rank)?;
let mut gates_c = e.zeros(streams * t * hidden)?;
launch_qmatvec_bf16w(
e,
&up_b16,
&low_c,
&mut gates_c,
rank,
hidden,
t,
streams,
hidden * rank,
0,
rank,
t * hidden,
)?;
let mut mixed_c = e.zeros(t * hidden)?;
launch_hc_mix_epilogue(e, &gates_c, &normed, &mut mixed_c, streams, t, hidden)?;
let mut partials_c = e.zeros(streams * t * 16)?;
let mut all_c = e.zeros(streams * t)?;
launch_hc_inject_two_stage(
e,
&normed,
&inj_f32,
Some(&inj_b16),
&mut partials_c,
&mut all_c,
streams,
t,
hidden,
16,
)?;
let mut parts_d = e.zeros(streams * rank)?;
let mut injp_d = e.zeros(streams * streams)?;
let mut inv_d = e.zeros(streams)?;
launch_hc_diet_stage1(
e,
&ptrs,
&norm_stack,
&down_b16,
Some(&inj_b16),
&mut parts_d,
&mut injp_d,
&mut inv_d,
hidden,
rank,
streams,
1,
eps,
)?;
let mut low_d = e.zeros(rank)?;
let mut all_d = e.zeros(streams)?;
launch_hc_diet_stage2(
e, &parts_d, &injp_d, &mut low_d, &mut all_d, rank, streams, 1, true,
)?;
let mut mixed_d = e.zeros(hidden)?;
launch_hc_diet_stage3(
e,
&ptrs,
&norm_stack,
&inv_d,
&up_b16,
&low_d,
&mut mixed_d,
hidden,
rank,
streams,
1,
)?;
let mut worst = 0.0f32;
let check = |name: &str, a: &[f32], b: &[f32], worst: &mut f32| -> Res<()> {
for (i, (&x, &y)) in a.iter().zip(b).enumerate() {
let rel = (x - y).abs() / y.abs().max(1.0);
if rel > *worst {
*worst = rel;
}
if rel > 1e-4 {
return Err(format!(
"hc-diet oracle: {name} idx {i}: diet {x} classic {y} (rel {rel:.3e})"
)
.into());
}
}
Ok(())
};
check("low_act", &e.dtoh(&low_d)?, &e.dtoh(&low_c)?, &mut worst)?;
check("inject", &e.dtoh(&all_d)?, &e.dtoh(&all_c)?, &mut worst)?;
check("mixed", &e.dtoh(&mixed_d)?, &e.dtoh(&mixed_c)?, &mut worst)?;
{
let t3 = 3usize;
let planes3_host: Vec<Vec<f32>> = (0..streams).map(|_| rand_vec(t3 * hidden)).collect();
let planes3: Vec<CudaSlice<f32>> = planes3_host
.iter()
.map(|v| e.htod(v))
.collect::<Result<_, _>>()?;
let ptr_vals3: Vec<u64> = {
let stream = e.gpu.stream();
planes3.iter().map(|p| p.device_ptr(&stream).0).collect()
};
let ptrs3 = e.htod_u64(&ptr_vals3)?;
let mut parts3 = e.zeros(t3 * streams * rank)?;
let mut injp3 = e.zeros(t3 * streams * streams)?;
let mut inv3 = e.zeros(t3 * streams)?;
launch_hc_diet_stage1(
e,
&ptrs3,
&norm_stack,
&down_b16,
Some(&inj_b16),
&mut parts3,
&mut injp3,
&mut inv3,
hidden,
rank,
streams,
t3,
eps,
)?;
let mut low3 = e.zeros(t3 * rank)?;
let mut all3 = e.zeros(streams * t3)?;
launch_hc_diet_stage2(
e, &parts3, &injp3, &mut low3, &mut all3, rank, streams, t3, true,
)?;
let mut mixed3 = e.zeros(t3 * hidden)?;
launch_hc_diet_stage3(
e,
&ptrs3,
&norm_stack,
&inv3,
&up_b16,
&low3,
&mut mixed3,
hidden,
rank,
streams,
t3,
)?;
let low3_h = e.dtoh(&low3)?;
let all3_h = e.dtoh(&all3)?;
let mixed3_h = e.dtoh(&mixed3)?;
{
let mut inv_mt = e.zeros(t3 * streams)?;
launch_hc_diet_stage0_mt(e, &ptrs3, &mut inv_mt, hidden, streams, t3, eps)?;
let mut parts_mt = e.zeros(t3 * streams * rank)?;
let mut injp_mt = e.zeros(t3 * streams * streams)?;
launch_hc_diet_stage1_mt(
e,
&ptrs3,
&norm_stack,
&inv_mt,
&down_b16,
Some(&inj_b16),
&mut parts_mt,
&mut injp_mt,
hidden,
rank,
streams,
t3,
)?;
let mut low_mt = e.zeros(t3 * rank)?;
let mut all_mt = e.zeros(streams * t3)?;
launch_hc_diet_stage2(
e,
&parts_mt,
&injp_mt,
&mut low_mt,
&mut all_mt,
rank,
streams,
t3,
true,
)?;
let mut mixed_mt = e.zeros(t3 * hidden)?;
launch_hc_diet_stage3_mt(
e,
&ptrs3,
&norm_stack,
&inv_mt,
&up_b16,
&low_mt,
&mut mixed_mt,
hidden,
rank,
streams,
t3,
)?;
let bit_check_mt = |name: &str, a: &[f32], b: &[f32]| -> Res<()> {
for (i, (&x, &y)) in a.iter().zip(b).enumerate() {
if x.to_bits() != y.to_bits() {
return Err(format!(
"hc-diet mt oracle: {name} idx {i}: mt {x} vs grid {y} NOT \
bit-identical"
)
.into());
}
}
Ok(())
};
bit_check_mt("inv", &e.dtoh(&inv_mt)?, &e.dtoh(&inv3)?)?;
bit_check_mt("low_act", &e.dtoh(&low_mt)?, &low3_h)?;
bit_check_mt("inject", &e.dtoh(&all_mt)?, &all3_h)?;
bit_check_mt("mixed", &e.dtoh(&mixed_mt)?, &mixed3_h)?;
}
let bit_check = |name: &str, a: &[f32], b: &[f32]| -> Res<()> {
for (i, (&x, &y)) in a.iter().zip(b).enumerate() {
if x.to_bits() != y.to_bits() {
return Err(format!(
"hc-diet t-ext oracle: {name} idx {i}: t3 {x} vs t1 {y} NOT bit-identical"
)
.into());
}
}
Ok(())
};
for tok in 0..t3 {
let ptr_tok: Vec<u64> = ptr_vals3
.iter()
.map(|&base| base + (tok * hidden * 4) as u64)
.collect();
let ptrs_tok = e.htod_u64(&ptr_tok)?;
let mut parts1 = e.zeros(streams * rank)?;
let mut injp1 = e.zeros(streams * streams)?;
let mut inv1 = e.zeros(streams)?;
launch_hc_diet_stage1(
e,
&ptrs_tok,
&norm_stack,
&down_b16,
Some(&inj_b16),
&mut parts1,
&mut injp1,
&mut inv1,
hidden,
rank,
streams,
1,
eps,
)?;
let mut low1 = e.zeros(rank)?;
let mut all1 = e.zeros(streams)?;
launch_hc_diet_stage2(
e, &parts1, &injp1, &mut low1, &mut all1, rank, streams, 1, true,
)?;
let mut mixed1 = e.zeros(hidden)?;
launch_hc_diet_stage3(
e,
&ptrs_tok,
&norm_stack,
&inv1,
&up_b16,
&low1,
&mut mixed1,
hidden,
rank,
streams,
1,
)?;
bit_check(
"low_act",
&low3_h[tok * rank..(tok + 1) * rank],
&e.dtoh(&low1)?,
)?;
let all1_h = e.dtoh(&all1)?;
let col: Vec<f32> = (0..streams).map(|s| all3_h[s * t3 + tok]).collect();
bit_check("inject", &col, &all1_h)?;
bit_check(
"mixed",
&mixed3_h[tok * hidden..(tok + 1) * hidden],
&e.dtoh(&mixed1)?,
)?;
}
}
Ok(format!(
"hc-diet real-geometry oracle: streams 4 hidden 2560 rank 320, worst rel \
{worst:.3e} vs the classic fused chain at t 1; t 3 token-dim AND the mt \
weight-shared stages BIT-IDENTICAL to per-token t 1 launches"
))
}
pub fn gate_hc_micro_kernels(e: &Engine) -> Res<String> {
let (streams, hidden, t) = (4usize, 2560usize, 10usize);
let wide = streams * hidden;
let mut lcg = 0x1357_9bdf_u64;
let mut next_f32 = move || -> f32 {
lcg = lcg
.wrapping_mul(6364136223846793005)
.wrapping_add(1442695040888963407);
(((lcg >> 33) as u32) % 2000) as f32 / 1000.0 - 1.0
};
let mut rand_vec = |n: usize| -> Vec<f32> { (0..n).map(|_| next_f32()).collect() };
let planes_host: Vec<Vec<f32>> = (0..streams).map(|_| rand_vec(t * hidden)).collect();
let planes: Vec<CudaSlice<f32>> = planes_host
.iter()
.map(|v| e.htod(v))
.collect::<Result<_, _>>()?;
let ptr_vals: Vec<u64> = {
let stream = e.gpu.stream();
planes.iter().map(|p| p.device_ptr(&stream).0).collect()
};
let ptrs = e.htod_u64(&ptr_vals)?;
let mut worst = 0.0f32;
let check = |name: &str, a: &[f32], b: &[f32], worst: &mut f32| -> Res<()> {
for (i, (&x, &y)) in a.iter().zip(b).enumerate() {
let rel = (x - y).abs() / y.abs().max(1.0);
if rel > *worst {
*worst = rel;
}
if rel > 1e-4 {
return Err(format!(
"hc-micro oracle: {name} idx {i}: micro {x} classic {y} (rel {rel:.3e})"
)
.into());
}
}
Ok(())
};
let norm_stack_host = rand_vec(wide);
let norm_stack = e.htod(&norm_stack_host)?;
let eps = 1e-6f32;
let mut normed_a = e.zeros(streams * t * hidden)?;
launch_hc_norm_planes(
e,
&ptrs,
&norm_stack,
&mut normed_a,
hidden,
t,
streams,
eps,
)?;
let mut normed_b = e.zeros(streams * t * hidden)?;
for s in 0..streams {
let w = e.htod(&norm_stack_host[s * hidden..(s + 1) * hidden])?;
let mut dst = normed_b.slice_mut(s * t * hidden..(s + 1) * t * hidden);
launch_rms_norm_into_view(e, &planes[s], &w, &mut dst, hidden, t, eps)?;
}
check("norm", &e.dtoh(&normed_a)?, &e.dtoh(&normed_b)?, &mut worst)?;
let inj_w_host = rand_vec(streams * wide);
let inj_w = e.htod(&inj_w_host)?;
let mut all_a = e.zeros(streams * t)?;
let mut partials = e.zeros(streams * t * 16)?;
launch_hc_inject_two_stage(
e,
&normed_b,
&inj_w,
None,
&mut partials,
&mut all_a,
streams,
t,
hidden,
16,
)?;
let mut all_b = e.zeros(streams * t)?;
launch_hc_inject_gates(e, &normed_b, &inj_w, &mut all_b, streams, t, hidden)?;
check("inject", &e.dtoh(&all_a)?, &e.dtoh(&all_b)?, &mut worst)?;
let block_out = e.htod(&rand_vec(t * hidden))?;
launch_hc_write_planes(e, &ptrs, &block_out, &all_b, hidden, t, streams)?;
let mut expect: Vec<Vec<f32>> = Vec::with_capacity(streams);
let all_host = e.dtoh(&all_b)?;
let bo_host = e.dtoh(&block_out)?;
for (s, base) in planes_host.iter().enumerate() {
let mut rows = base.clone();
for tok in 0..t {
let g = all_host[s * t + tok];
for d in 0..hidden {
rows[tok * hidden + d] += bo_host[tok * hidden + d] * g;
}
}
expect.push(rows);
}
for (s, plane) in planes.iter().enumerate() {
check(
&format!("write plane {s}"),
&e.dtoh(plane)?,
&expect[s],
&mut worst,
)?;
}
Ok(format!(
"hc-micro real-geometry oracle: streams 4 hidden 2560 t 10, worst rel {worst:.3e} \
over norm/inject/write vs the classic composition"
))
}
pub fn gate_sdpa_blocklist(e: &Engine) -> Res<String> {
let mut lcg = 0x51ee_7bad_u64;
let mut next_f32 = move || -> f32 {
lcg = lcg
.wrapping_mul(6364136223846793005)
.wrapping_add(1442695040888963407);
(((lcg >> 33) as u32) % 2000) as f32 / 1000.0 - 1.0
};
let (hd, nh, nkv, t) = (256usize, 24usize, 2usize, 3usize);
let block_size = 4usize;
let scale = 1.0 / (hd as f32).sqrt();
let mut bit_rows = 0usize;
let mut worst_rel = 0.0f32;
for (t_kv, vs_masked) in [(4096usize, true), (16384usize, false)] {
let q_host: Vec<f32> = (0..t * nh * hd).map(|_| next_f32()).collect();
let k_host: Vec<f32> = (0..t_kv * nkv * hd).map(|_| next_f32()).collect();
let v_host: Vec<f32> = (0..t_kv * nkv * hd).map(|_| next_f32()).collect();
let sels: Vec<RowSel> = (0..t)
.map(|qt| {
let visible = t_kv - t + qt + 1;
let complete = visible / block_size;
if qt == 0 && vs_masked {
return RowSel {
full: true,
blocks: Vec::new(),
visible,
};
}
let stride = if qt == 1 { 3 } else { 7 };
let blocks: Vec<u32> = (0..complete as u32)
.rev()
.step_by(stride)
.take(512)
.collect::<Vec<_>>()
.into_iter()
.rev()
.collect();
RowSel {
full: false,
blocks,
visible,
}
})
.collect();
let (pos_flat, meta, max_count) = rowsel_positions(&sels, block_size);
let q = e.htod(&q_host)?;
let k = e.htod(&k_host)?;
let v = e.htod(&v_host)?;
let pos = e.htod_i32(&pos_flat)?;
let meta_dev = e.htod_i32(&meta)?;
let mut o_list = e.zeros(t * nh * hd)?;
launch_sdpa_blocklist(
e,
&q,
&k.slice(0..t_kv * nkv * hd),
&v.slice(0..t_kv * nkv * hd),
&mut o_list,
&pos,
&meta_dev,
hd,
nh,
nkv,
t,
max_count,
scale,
)?;
let ours = e.dtoh(&o_list)?;
if vs_masked {
let mask = rowsel_to_mask(&sels, block_size, t_kv);
let mask_dev = e.htod_bytes(&mask)?;
let mut o_mask = e.zeros(t * nh * hd)?;
launch_sdpa_mask(
e,
&q,
&k.slice(0..t_kv * nkv * hd),
&v.slice(0..t_kv * nkv * hd),
&mut o_mask,
&mask_dev,
hd,
nh,
nkv,
t,
t_kv,
scale,
)?;
let masked = e.dtoh(&o_mask)?;
for (i, (a, b)) in masked.iter().zip(ours.iter()).enumerate() {
if a.to_bits() != b.to_bits() {
return Err(format!(
"sdpa_blocklist vs masked: bit mismatch at {i}: {a} vs {b} (t_kv {t_kv})"
)
.into());
}
}
bit_rows = t * nh * hd;
} else {
for qt in 0..t {
let off = meta[2 * qt] as usize;
let count = meta[2 * qt + 1] as usize;
for head in 0..nh {
let kvh = head / (nh / nkv);
let qrow = &q_host[(qt * nh + head) * hd..(qt * nh + head + 1) * hd];
let mut scores: Vec<f32> = (0..count)
.map(|i| {
let p = pos_flat[off + i] as usize;
let krow = &k_host[(p * nkv + kvh) * hd..(p * nkv + kvh + 1) * hd];
let mut acc = 0.0f32;
for d in 0..hd {
acc += qrow[d] * krow[d];
}
acc * scale
})
.collect();
let mx = scores.iter().copied().fold(-1e30f32, f32::max);
let mut sum = 0.0f32;
for s in scores.iter_mut() {
*s = (*s - mx).exp();
sum += *s;
}
let inv = 1.0 / sum;
for s in scores.iter_mut() {
*s *= inv;
}
for d in 0..hd {
let mut acc = 0.0f32;
for (i, s) in scores.iter().enumerate() {
let p = pos_flat[off + i] as usize;
acc += s * v_host[(p * nkv + kvh) * hd + d];
}
let got = ours[(qt * nh + head) * hd + d];
let rel = (got - acc).abs() / acc.abs().max(1e-3);
worst_rel = worst_rel.max(rel);
if rel > 1e-4 {
return Err(format!(
"sdpa_blocklist vs host twin: rel {rel} at row {qt} head {head} \
dim {d} (t_kv {t_kv})"
)
.into());
}
}
}
}
}
}
Ok(format!(
"sdpa-blocklist oracle: BIT-IDENTICAL to the masked kernel over {bit_rows} values \
(t_kv 4096, full+stride selections); past the mask bound (t_kv 16384) worst rel \
{worst_rel:.3e} vs the host twin"
))
}
pub fn gate_kvq_kernels(e: &Engine) -> Res<String> {
let mut lcg = 0x6b_7671_5eed_u64; let mut next_f32 = move || -> f32 {
lcg = lcg
.wrapping_mul(6364136223846793005)
.wrapping_add(1442695040888963407);
(((lcg >> 33) as u32) % 2000) as f32 / 1000.0 - 1.0
};
let mut report = Vec::new();
for &dim in &[512usize, 40usize] {
let rows = 9usize;
let mut host_rows_f: Vec<f32> = (0..rows * dim).map(|_| next_f32()).collect();
for v in host_rows_f[0..dim].iter_mut() {
*v = 0.0;
}
for v in host_rows_f[dim..2 * dim].iter_mut() {
*v = 0.75;
}
for (i, v) in host_rows_f[2 * dim..3 * dim].iter_mut().enumerate() {
*v = if i == 0 {
1.0
} else {
(i as f32) * 0.5 / 127.0
};
}
for v in host_rows_f[3 * dim..4 * dim].iter_mut() {
*v *= 1e-40;
}
let dev_rows = e.htod(&host_rows_f)?;
let mut kq = e.alloc_u8(rows * q8_row_bytes(dim))?;
let mut vq = e.alloc_u8(rows * q5_row_bytes(dim))?;
launch_q4e_kv_append(e, &dev_rows, &dev_rows, &mut kq, &mut vq, 0, rows, dim)?;
let kq_host = e.dtoh_u8(&kq)?;
let vq_host = e.dtoh_u8(&vq)?;
let mut k_twin = Vec::new();
let mut v_twin = Vec::new();
for r in 0..rows {
host_quant_q8_row(&host_rows_f[r * dim..(r + 1) * dim], dim, &mut k_twin);
host_quant_q5_row(&host_rows_f[r * dim..(r + 1) * dim], dim, &mut v_twin);
}
if kq_host != k_twin {
let i = kq_host.iter().zip(&k_twin).position(|(a, b)| a != b);
return Err(format!("kvq q8 quantize twin: byte mismatch at {i:?} (dim {dim})").into());
}
if vq_host != v_twin {
let i = vq_host.iter().zip(&v_twin).position(|(a, b)| a != b);
return Err(format!("kvq q5 quantize twin: byte mismatch at {i:?} (dim {dim})").into());
}
let mut kf = e.zeros(rows * dim)?;
let mut vf = e.zeros(rows * dim)?;
launch_q4e_kv_dequant_rows(e, &kq, &vq, &mut kf, &mut vf, 0, rows, dim)?;
let kf_host = e.dtoh(&kf)?;
let vf_host = e.dtoh(&vf)?;
let mut kf_twin = Vec::new();
let mut vf_twin = Vec::new();
host_deq_q8_rows(&kq_host, 0, rows, dim, &mut kf_twin);
host_deq_q5_rows(&vq_host, 0, rows, dim, &mut vf_twin);
for (i, (a, b)) in kf_host.iter().zip(&kf_twin).enumerate() {
if a.to_bits() != b.to_bits() {
return Err(format!("kvq q8 dequant twin: bit mismatch at {i} (dim {dim})").into());
}
}
for (i, (a, b)) in vf_host.iter().zip(&vf_twin).enumerate() {
if a.to_bits() != b.to_bits() {
return Err(format!("kvq q5 dequant twin: bit mismatch at {i} (dim {dim})").into());
}
}
report.push(format!("quant+dequant twins dim {dim}: BYTE/BIT-IDENTICAL"));
}
{
let (hd, nh, nkv, t) = (256usize, 24usize, 2usize, 3usize);
let kv_dim = nkv * hd;
let block_size = 4usize;
let scale = 1.0 / (hd as f32).sqrt();
let t_kv = 4096usize;
let q_host: Vec<f32> = (0..t * nh * hd).map(|_| next_f32()).collect();
let k_host: Vec<f32> = (0..t_kv * kv_dim).map(|_| next_f32()).collect();
let v_host: Vec<f32> = (0..t_kv * kv_dim).map(|_| next_f32()).collect();
let k_rows = e.htod(&k_host)?;
let v_rows = e.htod(&v_host)?;
let mut kq = e.alloc_u8(t_kv * q8_row_bytes(kv_dim))?;
let mut vq = e.alloc_u8(t_kv * q5_row_bytes(kv_dim))?;
launch_q4e_kv_append(e, &k_rows, &v_rows, &mut kq, &mut vq, 0, t_kv, kv_dim)?;
let sels: Vec<RowSel> = (0..t)
.map(|qt| {
let visible = (t_kv - t + qt + 1).min(2052);
if qt == 0 {
return RowSel {
full: true,
blocks: Vec::new(),
visible,
};
}
let complete = (t_kv - t + qt + 1) / block_size;
let stride = if qt == 1 { 3 } else { 7 };
let blocks: Vec<u32> = (0..complete as u32)
.rev()
.step_by(stride)
.take(512)
.collect::<Vec<_>>()
.into_iter()
.rev()
.collect();
RowSel {
full: false,
blocks,
visible: t_kv - t + qt + 1,
}
})
.collect();
let (pos_flat, meta, max_count) = rowsel_positions(&sels, block_size);
let q = e.htod(&q_host)?;
let pos = e.htod_i32(&pos_flat)?;
let meta_dev = e.htod_i32(&meta)?;
let mut o_fused = e.zeros(t * nh * hd)?;
launch_q4e_sdpa_blocklist_q8q5(
e,
&q,
&kq,
&vq,
&mut o_fused,
&pos,
&meta_dev,
hd,
nh,
nkv,
t,
max_count,
scale,
)?;
let mut k_deq = e.zeros(t_kv * kv_dim)?;
let mut v_deq = e.zeros(t_kv * kv_dim)?;
launch_q4e_kv_dequant_rows(e, &kq, &vq, &mut k_deq, &mut v_deq, 0, t_kv, kv_dim)?;
let mut o_comp = e.zeros(t * nh * hd)?;
launch_sdpa_blocklist(
e,
&q,
&k_deq.slice(0..t_kv * kv_dim),
&v_deq.slice(0..t_kv * kv_dim),
&mut o_comp,
&pos,
&meta_dev,
hd,
nh,
nkv,
t,
max_count,
scale,
)?;
let fused = e.dtoh(&o_fused)?;
let comp = e.dtoh(&o_comp)?;
for (i, (a, b)) in fused.iter().zip(&comp).enumerate() {
if a.to_bits() != b.to_bits() {
return Err(format!(
"kvq fused attention vs dequant composition: bit mismatch at {i}: {a} vs {b}"
)
.into());
}
}
{
let was = kv_hoist_on();
set_kv_hoist(true);
let mut o_hoist = e.zeros(t * nh * hd)?;
let launched = launch_q4e_sdpa_blocklist_q8q5(
e,
&q,
&kq,
&vq,
&mut o_hoist,
&pos,
&meta_dev,
hd,
nh,
nkv,
t,
max_count,
scale,
);
set_kv_hoist(was);
launched?;
let hoist = e.dtoh(&o_hoist)?;
let mut worst: Option<(usize, f32, f32)> = None;
for (i, (a, b)) in hoist.iter().zip(&fused).enumerate() {
if a.to_bits() != b.to_bits() && worst.is_none() {
worst = Some((i, *a, *b));
}
}
if let Some((i, a, b)) = worst {
return Err(format!(
"kvhoist vs un-hoisted q8q5 blocklist: bit mismatch at {i}: {a} vs {b} \
(hd={hd} nh={nh} nkv={nkv} t={t} t_kv={t_kv} max_count={max_count})"
)
.into());
}
report.push(format!(
"kvhoist vs un-hoisted q8q5 blocklist: BIT-IDENTICAL over {} values \
(real geometry hd={hd} nh={nh} nkv={nkv}, {} blocks/head slice, max_count={max_count})",
t * nh * hd,
hd / 32
));
}
report.push(format!(
"fused q8q5 blocklist vs dequant+f32 composition: BIT-IDENTICAL over {} values",
t * nh * hd
));
}
{
let idx_dim = 128usize;
let qk_width = 5 * idx_dim; let rows = 7usize;
let src_host: Vec<f32> = (0..rows * qk_width).map(|_| next_f32()).collect();
let src = e.htod(&src_host)?;
let q_off = 4 * idx_dim;
let mut dst_q8 = e.alloc_u8((rows + 2) * q8_row_bytes(idx_dim))?;
launch_q4e_idx_append_q8(e, &src, &mut dst_q8, rows, idx_dim, qk_width, q_off, 2)?;
let got = e.dtoh_u8(&dst_q8)?;
let mut twin = vec![0u8; 2 * q8_row_bytes(idx_dim)];
for r in 0..rows {
host_quant_q8_row(
&src_host[r * qk_width + q_off..(r + 1) * qk_width],
idx_dim,
&mut twin,
);
}
if got[2 * q8_row_bytes(idx_dim)..] != twin[2 * q8_row_bytes(idx_dim)..] {
return Err("idxq q8 append twin: byte mismatch".into());
}
let mut dst_bf = unsafe { e.gpu.stream().alloc::<u16>((rows + 2) * idx_dim)? };
e.gpu.stream().memset_zeros(&mut dst_bf)?;
launch_q4e_idx_append_bf16(e, &src, &mut dst_bf, rows, idx_dim, qk_width, q_off, 2)?;
let got_bf: Vec<u16> = {
let v = e
.gpu
.stream()
.clone_dtoh(&dst_bf.slice(0..(rows + 2) * idx_dim))?;
e.gpu.stream().synchronize()?;
v
};
for r in 0..rows {
for c in 0..idx_dim {
let want = f32_to_bf16_rne(src_host[r * qk_width + q_off + c]);
if got_bf[(2 + r) * idx_dim + c] != want {
return Err(format!("idxq bf16 append twin: mismatch row {r} col {c}").into());
}
}
}
report.push("idx q8/bf16 appenders vs host twins: BYTE-IDENTICAL".to_string());
}
Ok(format!("kvq kernel oracles: {}", report.join("; ")))
}
pub fn gate_qsa_index_score(e: &Engine) -> Res<String> {
let mut lcg = 0xfeed_1234_u64;
let mut next_f32 = move || -> f32 {
lcg = lcg
.wrapping_mul(6364136223846793005)
.wrapping_add(1442695040888963407);
(((lcg >> 33) as u32) % 4000) as f32 / 2000.0 - 1.0
};
let (heads, head_dim) = (4usize, 128usize);
let scale = (head_dim as f32).sqrt();
let budget = 512usize;
let mut worst_rows = 0usize;
for (rows, n_blocks) in [(1usize, 4096usize), (7, 1031)] {
let q_host: Vec<f32> = (0..rows * heads * head_dim).map(|_| next_f32()).collect();
let pooled_host: Vec<f32> = (0..n_blocks * head_dim).map(|_| next_f32()).collect();
let q = e.htod(&q_host)?;
let pooled = e.htod(&pooled_host)?;
let mut scores_dev = e.uninit(rows * n_blocks)?;
launch_qsa_index_score(
e,
&q,
&pooled,
&mut scores_dev,
heads,
head_dim,
n_blocks,
rows,
scale,
)?;
let got = e.dtoh(&scores_dev)?;
for row in 0..rows {
let qr = &q_host[row * heads * head_dim..(row + 1) * heads * head_dim];
let host = score_blocks(qr, &pooled_host, heads, head_dim, n_blocks, scale, 1);
for (b, want) in host.iter().enumerate() {
let g = got[row * n_blocks + b];
if g.to_bits() != want.to_bits() {
return Err(format!(
"qsa_index_score: bit mismatch row {row} block {b}: {g} vs host {want}"
)
.into());
}
}
let a = top_blocks_ascending(&host, budget, 1);
let b = top_blocks_ascending(&got[row * n_blocks..(row + 1) * n_blocks], budget, 1);
if a != b {
return Err(format!("qsa_index_score: top-k set differs at row {row}").into());
}
worst_rows += 1;
}
}
let mut pool_t_rows = 0usize;
for (rows, n_blocks, cap_rows) in [(1usize, 1031usize, 4096usize), (5, 2048, 2048)] {
let q_host: Vec<f32> = (0..rows * heads * head_dim).map(|_| next_f32()).collect();
let pooled_host: Vec<f32> = (0..n_blocks * head_dim).map(|_| next_f32()).collect();
let q = e.htod(&q_host)?;
let mut mirror = e.zeros(cap_rows * head_dim * POOL_PLANES)?;
{
let mut view = mirror.slice_mut(0..n_blocks * head_dim);
e.gpu.stream().memcpy_htod(&pooled_host, &mut view)?;
}
launch_qsa_pooled_transpose(e, &mut mirror, 0, n_blocks, head_dim, cap_rows)?;
let was = pool_t_on();
set_pool_t(true);
let mut scores_dev = e.uninit(rows * n_blocks)?;
let launched = launch_qsa_index_score(
e,
&q,
&mirror,
&mut scores_dev,
heads,
head_dim,
n_blocks,
rows,
scale,
);
set_pool_t(was);
launched?;
let got = e.dtoh(&scores_dev)?;
for row in 0..rows {
let qr = &q_host[row * heads * head_dim..(row + 1) * heads * head_dim];
let host = score_blocks(qr, &pooled_host, heads, head_dim, n_blocks, scale, 1);
for (b, want) in host.iter().enumerate() {
let g = got[row * n_blocks + b];
if g.to_bits() != want.to_bits() {
return Err(format!(
"poolT qsa_index_score_f32_t: bit mismatch row {row} block {b}: \
{g} vs host {want} (n_blocks={n_blocks} cap_rows={cap_rows})"
)
.into());
}
}
if top_blocks_ascending(&host, budget, 1)
!= top_blocks_ascending(&got[row * n_blocks..(row + 1) * n_blocks], budget, 1)
{
return Err(format!(
"poolT qsa_index_score_f32_t: top-{budget} set differs at row {row} \
(n_blocks={n_blocks} cap_rows={cap_rows})"
)
.into());
}
pool_t_rows += 1;
}
}
Ok(format!(
"qsa-index-score oracle: device scores BIT-IDENTICAL to the host twin over \
{worst_rows} rows (4096 + 1031 blocks, real 4x128 geometry) and top-512 sets equal; \
poolT dim-major plane (transpose + transposed kernel) BIT-IDENTICAL to the SAME host \
twin over {pool_t_rows} rows, incl. the pitch-trap case cap_rows=4096 != n_blocks=1031"
))
}
pub fn gate_ple_ngram_cache() -> Res<String> {
let max_ngram = 3usize;
let heads_per_ngram = 8usize;
let total_heads = (max_ngram - 1) * heads_per_ngram;
let multipliers: Vec<i64> = vec![
0x2545_F491_4F6C_DD1D,
0x9E37_79B9_7F4A_7C15u64 as i64,
0x1234_5678_9ABC_DEF1,
];
let sizes: Vec<i64> = (0..total_heads)
.map(|i| 2_500_012_160 - (i as i64) * 7)
.collect();
let offsets: Vec<i64> = (0..total_heads)
.map(|i| (i as i64) * 2_500_012_160)
.collect();
let eos = 248_046u32;
let full = |ids: &[u32]| -> Vec<i64> {
host_ngram_ids(
ids,
&multipliers,
&sizes,
&offsets,
max_ngram,
heads_per_ngram,
eos,
)
};
let mut lcg = 0x0be1_10ca_u64;
let mut next_tok = move || -> u32 {
lcg = lcg
.wrapping_mul(6364136223846793005)
.wrapping_add(1442695040888963407);
((lcg >> 33) as u32) % 250_000
};
let mut checks = 0usize;
let run = |label: &str, steps: Vec<Vec<u32>>| -> Res<usize> {
let (mut ci, mut ch, mut ce) = (Vec::new(), Vec::new(), -1i64);
let mut n = 0usize;
for seq in &steps {
host_ngram_ids_cached(
&mut ci,
&mut ch,
&mut ce,
seq,
&multipliers,
&sizes,
&offsets,
max_ngram,
heads_per_ngram,
eos,
);
let want = full(seq);
if ci.len() != want.len() {
return Err(format!(
"plecache oracle {label}: cache has {} ids, twin {} at len {}",
ci.len(),
want.len(),
seq.len()
)
.into());
}
if let Some(i) = ci.iter().zip(&want).position(|(a, b)| a != b) {
return Err(format!(
"plecache oracle {label}: id {i} differs at len {} (token {}, head {}): \
cache {} vs twin {}",
seq.len(),
i / total_heads,
i % total_heads,
ci[i],
want[i]
)
.into());
}
n += seq.len();
}
Ok(n)
};
{
let base: Vec<u32> = (0..200).map(|_| next_tok()).collect();
let steps: Vec<Vec<u32>> = (1..=base.len()).map(|n| base[..n].to_vec()).collect();
checks += run("decode-growth", steps)?;
}
{
let base: Vec<u32> = (0..600).map(|_| next_tok()).collect();
let mut steps = Vec::new();
let mut n = 0usize;
for step in [7usize, 1, 64, 3, 128, 2, 200, 195] {
n = (n + step).min(base.len());
steps.push(base[..n].to_vec());
}
checks += run("prefill-chunks", steps)?;
}
{
let mut base: Vec<u32> = (0..300).map(|_| next_tok()).collect();
for p in [0usize, 1, 2, 37, 38, 100, 101, 102, 299] {
base[p] = eos;
}
let steps: Vec<Vec<u32>> = (1..=base.len()).map(|n| base[..n].to_vec()).collect();
checks += run("eos-segments", steps)?;
}
{
let base: Vec<u32> = vec![eos; 40];
let steps: Vec<Vec<u32>> = (1..=base.len()).map(|n| base[..n].to_vec()).collect();
checks += run("all-eos", steps)?;
}
{
let a: Vec<u32> = (0..300).map(|_| next_tok()).collect();
let mut b = a.clone();
b[150] = a[150].wrapping_add(1) % 250_000;
let mut c = b.clone();
c[7] = b[7].wrapping_add(3) % 250_000;
let mut d = c.clone();
d.truncate(9);
d.extend((0..100).map(|_| next_tok()));
checks += run(
"rewind-divergent",
vec![
a.clone(),
a[..151].to_vec(),
b.clone(),
b[..8].to_vec(),
c.clone(),
d.clone(),
a.clone(),
],
)?;
}
{
let a: Vec<u32> = (0..250).map(|_| next_tok()).collect();
let mut s: Vec<u32> = (0..11).map(|_| next_tok()).collect();
s[0] = eos;
checks += run("state-reuse", vec![a.clone(), s.clone(), a.clone(), s])?;
}
Ok(format!(
"plecache oracle: incremental n-gram ids EXACT vs the full host_ngram_ids twin over \
{checks} cumulative-sequence comparisons across 6 case families (decode one-at-a-time \
growth, ragged prefill chunks, eos segment resets incl. adjacent + leading + trailing \
eos, all-eos, repeated rewinds to DIVERGING prefixes, and shorter-unrelated-sequence \
state reuse)"
))
}
pub fn gate_seam_table() -> Res<String> {
let all: &[&str] = seam_names();
let boolean: Vec<&str> = all
.iter()
.copied()
.filter(|n| seam_state(n).is_some())
.collect();
if boolean.len() < 20 || all.len() < boolean.len() + 2 {
return Err(format!(
"seam-table oracle: refusing to report on {} boolean names out of {} total — the \
seam list collapsed, so every assertion below would be vacuous",
boolean.len(),
all.len()
)
.into());
}
let names: &[&str] = &boolean;
let snapshot = || -> Res<Vec<bool>> {
names
.iter()
.map(|n| {
seam_state(n).ok_or_else(|| {
Box::<dyn std::error::Error>::from(format!(
"seam-table oracle: seam_state({n:?}) is None — the name is in set_seam \
but not in seam_state, so save/restore around a measurement would \
silently not restore it"
))
})
})
.collect()
};
let restore = |v: &[bool]| {
for (n, &b) in names.iter().zip(v) {
set_seam(n, b, None);
}
};
let entry = snapshot()?;
let mut checks = 0usize;
for (i, name) in names.iter().enumerate() {
for &want in &[true, false, true] {
let before = snapshot()?;
if !set_seam(name, want, None) {
restore(&entry);
return Err(format!("seam-table oracle: set_seam({name:?}) refused").into());
}
let after = snapshot()?;
if after[i] != want {
restore(&entry);
return Err(format!(
"seam-table oracle: set_seam({name:?}, {want}) then seam_state read {} — the \
two tables disagree on this name",
after[i]
)
.into());
}
for (j, other) in names.iter().enumerate() {
if j != i && after[j] != before[j] {
restore(&entry);
return Err(format!(
"seam-table oracle: arming {name:?} also changed {other:?} ({} -> {}) — \
two names share one switch",
before[j], after[j]
)
.into());
}
}
checks += 1;
}
}
for name in all {
if !seam_exists(name) {
restore(&entry);
return Err(format!(
"seam-table oracle: seam_names() lists {name:?} but seam_exists refuses it"
)
.into());
}
if !set_seam(name, seam_state(name).unwrap_or(false), None) {
restore(&entry);
return Err(format!(
"seam-table oracle: seam_names() lists {name:?} but set_seam refuses it"
)
.into());
}
}
if seam_exists("definitely-not-a-seam") || set_seam("definitely-not-a-seam", true, None) {
restore(&entry);
return Err("seam-table oracle: an unknown seam name was accepted".into());
}
let before = snapshot()?;
for name in all {
let _ = seam_exists(name);
}
if snapshot()? != before {
restore(&entry);
return Err("seam-table oracle: seam_exists mutated a seam (it must be name-only)".into());
}
restore(&entry);
if snapshot()? != entry {
return Err("seam-table oracle: the gate did not restore the entry state".into());
}
Ok(format!(
"seam-table oracle: {} boolean seam names of {} total, {checks} set/read cycles, each \
verified to change its OWN state and NO other (the cross-check that catches an arm \
wired to a neighbour's switch), every listed name accepted by both entry points, \
unknown names refused by both, seam_exists proven side-effect-free, entry state restored",
names.len(),
all.len()
))
}
pub fn gate_qsa_index_topk(e: &Engine) -> Res<String> {
let budget = 512usize;
let mut lcg = 0x1d5e_10ca_u64;
let mut rows_checked = 0usize;
let mut deepest = 0usize;
let mut cases: Vec<(String, Vec<usize>, usize, Vec<f32>)> = Vec::new();
let mut next_f32 = move || -> f32 {
lcg = lcg
.wrapping_mul(6364136223846793005)
.wrapping_add(1442695040888963407);
let r = ((lcg >> 33) as u32) % 1000;
if r < 250 { 0.0 } else { (r as f32) / 250.0 }
};
for (label, counts) in [
("real-262k-depth", vec![65_536usize]),
("real-131k-depth", vec![32_768usize, 32_768]),
("shallow", vec![513usize, 1_031, 4_096]),
("ragged-batch", vec![2_049usize, 8_191, 65_536, 4_097]),
] {
let stride = *counts.iter().max().unwrap();
let slab: Vec<f32> = (0..counts.len() * stride).map(|_| next_f32()).collect();
cases.push((label.to_string(), counts, stride, slab));
}
cases.push((
"all-zero".into(),
vec![65_536usize],
65_536,
vec![0.0f32; 65_536],
));
{
let n = 4_096usize;
let mut v = vec![0.0f32; n];
for (i, slot) in v.iter_mut().enumerate() {
*slot = if i < 12 {
100.0 - i as f32
} else if i % 7 == 0 {
2.5 } else {
(i % 3) as f32 * 0.25
};
}
cases.push(("dup-straddle".into(), vec![n], n, v));
}
{
let n = 2_048usize;
let mut v = vec![0.0f32; n];
for (i, slot) in v.iter_mut().enumerate() {
*slot = match i % 8 {
0 => 0.0,
1 => -0.0,
2 => f32::from_bits(1), 3 => -f32::from_bits(1), 4 => -(i as f32) * 0.5,
5 => f32::NAN,
6 => -f32::NAN,
_ => (i % 5) as f32,
};
}
cases.push(("total-cmp-domain".into(), vec![n], n, v));
}
for (label, counts, stride, slab) in &cases {
let scores = e.htod(slab)?;
let picked = launch_qsa_index_topk(e, &scores, counts, *stride, budget)?;
if picked.len() != counts.len() {
return Err(format!("idxsel oracle {label}: {} rows back", picked.len()).into());
}
for (r, &complete) in counts.iter().enumerate() {
let row = &slab[r * *stride..r * *stride + complete];
let twin = top_blocks_ascending(row, budget, 1);
if twin != picked[r] {
let first = twin
.iter()
.zip(picked[r].iter())
.position(|(a, b)| a != b)
.unwrap_or(twin.len().min(picked[r].len()));
return Err(format!(
"idxsel oracle {label}: selection differs at row {r} (blocks {complete}), \
first differing slot {first}: host {:?} vs device {:?}",
twin.get(first),
picked[r].get(first)
)
.into());
}
rows_checked += 1;
deepest = deepest.max(complete);
}
}
Ok(format!(
"qsa-index-topk oracle: device selection ids + ASCENDING order EXACT vs \
top_blocks_ascending over {rows_checked} rows / {} cases at budget {budget}, \
deepest {deepest} blocks (= the 262,144-token window), incl. the all-zero, \
boundary-straddling-duplicate and total_cmp-domain (signed zero / subnormal / \
negative / NaN) tie classes",
cases.len()
))
}
pub fn gate_route_kernel(e: &Engine) -> Res<String> {
let mut lcg = 0x00de_7710_u64;
let mut next_f32 = move || -> f32 {
lcg = lcg
.wrapping_mul(6364136223846793005)
.wrapping_add(1442695040888963407);
(((lcg >> 33) as u32) % 8000) as f32 / 200.0 - 20.0 };
const ULP_BOUND: u32 = 2;
let mut worst_ulp: u32 = 0;
let mut rows_checked = 0usize;
let run = |e: &Engine,
label: &str,
logits_host: &[f32],
experts: usize,
selected: usize,
rows: usize,
worst_ulp: &mut u32|
-> Res<()> {
let logits = e.htod(logits_host)?;
let mut sel = e.alloc_uninit::<i32>(rows * selected)?;
let mut w = e.uninit(rows * selected)?;
let mut tok = e.alloc_uninit::<i32>(rows * selected)?;
launch_route_topk(
e,
&logits,
&mut sel,
&mut w,
Some((&mut tok, 3)),
experts,
selected,
rows,
)?;
let sel_h = e.gpu.stream().clone_dtoh(&sel.slice(0..rows * selected))?;
let w_h = e.dtoh(&w)?;
let tok_h = e.gpu.stream().clone_dtoh(&tok.slice(0..rows * selected))?;
let k = selected.min(experts);
for row in 0..rows {
let twin =
host_route_softmax_topk(&logits_host[row * experts..(row + 1) * experts], selected);
if twin.len() != k {
return Err(format!("route oracle {label}: host twin width {}", twin.len()).into());
}
for (j, &(idx, wt)) in twin.iter().enumerate() {
let ds = sel_h[row * selected + j];
let dw = w_h[row * selected + j];
if ds != idx as i32 {
return Err(format!(
"route oracle {label}: selection mismatch row {row} slot {j}: \
device {ds} vs host {idx}"
)
.into());
}
let ulp = (dw.to_bits() as i64 - wt.to_bits() as i64).unsigned_abs();
let ulp = u32::try_from(ulp).unwrap_or(u32::MAX);
if ulp > ULP_BOUND {
return Err(format!(
"route oracle {label}: weight ULP {ulp} > {ULP_BOUND} at row {row} \
slot {j}: device {dw:e} vs host {wt:e}"
)
.into());
}
*worst_ulp = (*worst_ulp).max(ulp);
if tok_h[row * selected + j] != (3 + row) as i32 {
return Err(format!(
"route oracle {label}: tok map wrong at row {row} slot {j}"
)
.into());
}
}
}
Ok(())
};
let (experts, selected) = (512usize, 10usize);
for rows in [1usize, 6, 16] {
let logits: Vec<f32> = (0..rows * experts).map(|_| next_f32()).collect();
run(e, "real", &logits, experts, selected, rows, &mut worst_ulp)?;
rows_checked += rows;
}
let mut tie_rows: Vec<(String, Vec<f32>)> = Vec::new();
{
let mut v: Vec<f32> = (0..experts).map(|i| -30.0 - (i as f32) * 0.01).collect();
for (rank, slot) in [40usize, 7, 300, 11].iter().enumerate() {
v[*slot] = 10.0 - rank as f32;
}
for slot in [500usize, 3, 77, 210, 8, 401, 129, 64, 255, 380, 17, 450] {
v[slot] = 2.5;
}
tie_rows.push(("dup-straddle".into(), v));
tie_rows.push(("all-equal".into(), vec![0.125f32; experts]));
let mut v = vec![-200.0f32; experts];
v[100] = 5.0;
for (i, slot) in [479usize, 2, 33].iter().enumerate() {
v[*slot] = -80.0 - i as f32; }
tie_rows.push(("underflow".into(), v));
}
for (label, v) in &tie_rows {
run(e, label, v, experts, selected, 1, &mut worst_ulp)?;
rows_checked += 1;
}
for (ex, se) in [(64usize, 4usize), (16, 16), (128, 32)] {
let logits: Vec<f32> = (0..3 * ex).map(|_| next_f32()).collect();
run(e, "geom", &logits, ex, se, 3, &mut worst_ulp)?;
rows_checked += 3;
}
Ok(format!(
"route oracle: device selection ids+order EXACT vs host twin over {rows_checked} rows \
(real 512/10 + tie straddle/all-equal/underflow + geometry edges), worst weight \
ULP {worst_ulp} (bound {ULP_BOUND}), tok map exact"
))
}
pub fn gate_gdn_step_kernels(e: &Engine) -> Res<String> {
let mut lcg = 0x0bad_cafe_u64;
let mut next_f32 = move || -> f32 {
lcg = lcg
.wrapping_mul(6364136223846793005)
.wrapping_add(1442695040888963407);
(((lcg >> 33) as u32) % 2000) as f32 / 1000.0 - 1.0
};
let mut worst = 0.0f32;
for (nk, nv, hk, hv) in [(16usize, 48usize, 128usize, 128usize), (2, 4, 32, 8)] {
let conv_dim = 2 * nk * hk + nv * hv;
let qkv_host: Vec<f32> = (0..conv_dim).map(|_| next_f32()).collect();
let g_log_host: Vec<f32> = (0..nv).map(|_| next_f32().abs() * -2.0).collect();
let beta_host: Vec<f32> = (0..nv).map(|_| next_f32()).collect();
let state_host: Vec<f32> = (0..nv * hv * hk).map(|_| next_f32()).collect();
let qkv = e.htod(&qkv_host)?;
let g_log = e.htod(&g_log_host)?;
let beta = e.htod(&beta_host)?;
let scale = 1.0 / (hk as f32).sqrt();
let eps = 1e-6f32;
let mut state_a = e.htod(&state_host)?;
let mut o_a = e.zeros(nv * hv)?;
launch_gdn_scan(
e,
&qkv,
&g_log,
&beta,
&mut state_a,
&mut o_a,
nk,
nv,
hk,
hv,
1,
scale,
eps,
)?;
let mut state_b = e.htod(&state_host)?;
let mut o_b = e.zeros(nv * hv)?;
launch_gdn_scan_step(
e,
&qkv,
&g_log,
&beta,
&mut state_b,
&mut o_b,
nk,
nv,
hk,
hv,
scale,
eps,
)?;
for (name, reference, candidate) in [
("o", e.dtoh(&o_a)?, e.dtoh(&o_b)?),
("state", e.dtoh(&state_a)?, e.dtoh(&state_b)?),
] {
for (i, (&r, &c)) in reference.iter().zip(&candidate).enumerate() {
let rel = (r - c).abs() / r.abs().max(1.0);
if rel > worst {
worst = rel;
}
if rel > 1e-4 {
return Err(format!(
"gdn-step oracle: nk{nk}/nv{nv}/hk{hk}/hv{hv} {name} idx {i}: \
naive {r} step {c} (rel {rel:.3e})"
)
.into());
}
}
}
}
let (rows, cols) = (48usize, 128usize);
let x = e.htod(&(0..rows * cols).map(|_| next_f32()).collect::<Vec<_>>())?;
let w = e.htod(&(0..cols).map(|_| next_f32()).collect::<Vec<_>>())?;
let z = e.htod(&(0..rows * cols).map(|_| next_f32()).collect::<Vec<_>>())?;
let eps = 1e-6f32;
let mut normed = e.zeros(rows * cols)?;
e.rms_norm(&x, &w, &mut normed, cols, rows, eps)?;
let mut sg = e.zeros(rows * cols)?;
e.sigmoid(&z, &mut sg, rows * cols)?;
let mut chain = e.zeros(rows * cols)?;
e.mul(&normed, &sg, &mut chain, rows * cols)?;
let mut fused = e.zeros(rows * cols)?;
launch_rms_sigmul(e, &x, &w, &z, &mut fused, cols, rows, eps)?;
let (chain_h, fused_h) = (e.dtoh(&chain)?, e.dtoh(&fused)?);
for (i, (&a, &b)) in chain_h.iter().zip(&fused_h).enumerate() {
if a.to_bits() != b.to_bits() {
return Err(format!(
"rms_sigmul oracle: idx {i} not bit-identical: chain {a:?} fused {b:?}"
)
.into());
}
}
Ok(format!(
"gdn-step kernel oracle: scan step twin worst rel {worst:.3e} over artifact + \
hk32 geometries; rms_sigmul bit-identical to the norm/sigmoid/mul chain ({rows}x{cols})"
))
}
pub fn gate_qmatvec_bf16(e: &Engine) -> Res<String> {
let mut lcg = 0x9e37_79b9_u64;
let mut next_u32 = move || -> u32 {
lcg = lcg
.wrapping_mul(6364136223846793005)
.wrapping_add(1442695040888963407);
(lcg >> 33) as u32
};
let mut worst = (0.0f32, 0.0f32);
for (mode, batch, t, out_f, in_f, x_bstride) in [
("per_batch_x", 3usize, 2usize, 5usize, 48usize, 2 * 48usize),
("shared_x", 4, 3, 7, 16, 0usize),
] {
let w_elems = batch * out_f * in_f;
let mut w_bytes = Vec::with_capacity(w_elems * 2);
let mut w_host = Vec::with_capacity(w_elems);
for _ in 0..w_elems {
let h = ((next_u32() % 0x4000) as u16) | (((next_u32() & 1) as u16) << 15);
w_bytes.extend_from_slice(&h.to_le_bytes());
w_host.push(f32::from_bits(u32::from(h) << 16));
}
let x_rows = if x_bstride == 0 { t } else { batch * t };
let x_host: Vec<f32> = (0..x_rows * in_f)
.map(|_| (next_u32() % 2000) as f32 / 1000.0 - 1.0)
.collect();
let w_dev = e.htod_bytes(&w_bytes)?;
let x_dev = e.htod(&x_host)?;
let mut y_dev = e.uninit(batch * t * out_f)?;
launch_qmatvec_bf16w(
e,
&w_dev,
&x_dev,
&mut y_dev,
in_f,
out_f,
t,
batch,
out_f * in_f,
x_bstride,
in_f,
t * out_f,
)?;
let y = e.dtoh(&y_dev)?;
for b in 0..batch {
for tok in 0..t {
let xrow = &x_host[b * x_bstride + tok * in_f..][..in_f];
for o in 0..out_f {
let wrow = &w_host[(b * out_f + o) * in_f..][..in_f];
let mut want = 0.0f32;
for i in 0..in_f {
want += wrow[i] * xrow[i];
}
let got = y[(b * t + tok) * out_f + o];
let abs = (want - got).abs();
let rel = abs / want.abs().max(1.0);
worst.0 = worst.0.max(abs);
worst.1 = worst.1.max(rel);
if rel > 1e-5 {
return Err(format!(
"bf16-matvec oracle: {mode} b {b} tok {tok} row {o}: want {want} \
got {got} (rel {rel:.3e})"
)
.into());
}
}
}
}
}
{
let (out_f, in_f, t) = (33usize, 64usize, 5usize);
let w_elems = out_f * in_f;
let mut w_bytes = Vec::with_capacity(w_elems * 2);
for _ in 0..w_elems {
let h = ((next_u32() % 0x4000) as u16) | (((next_u32() & 1) as u16) << 15);
w_bytes.extend_from_slice(&h.to_le_bytes());
}
let x_host: Vec<f32> = (0..t * in_f)
.map(|_| (next_u32() % 2000) as f32 / 1000.0 - 1.0)
.collect();
let w_dev = e.htod_bytes(&w_bytes)?;
let x_dev = e.htod(&x_host)?;
let mut y_grid = e.uninit(t * out_f)?;
launch_qmatvec_bf16w(
e,
&w_dev,
&x_dev,
&mut y_grid,
in_f,
out_f,
t,
1,
0,
0,
in_f,
0,
)?;
let mut y_mt = e.uninit(t * out_f)?;
launch_qmatvec_bf16w_mt(e, &w_dev, 0, &x_dev, &mut y_mt, in_f, out_f, t)?;
let (a, b) = (e.dtoh(&y_grid)?, e.dtoh(&y_mt)?);
for (i, (&x1, &x2)) in a.iter().zip(&b).enumerate() {
if x1.to_bits() != x2.to_bits() {
return Err(format!(
"bf16-matvec mt oracle: idx {i}: grid {x1} vs mt {x2} NOT bit-identical"
)
.into());
}
}
}
{
let (experts, out_f, in_f, n_sel) = (16usize, 24usize, 32usize, 6usize);
let w_elems = experts * out_f * in_f;
let mut w_bytes = Vec::with_capacity(w_elems * 2);
for _ in 0..w_elems {
let h = ((next_u32() % 0x4000) as u16) | (((next_u32() & 1) as u16) << 15);
w_bytes.extend_from_slice(&h.to_le_bytes());
}
let sel_host: Vec<i32> = vec![7, 0, 15, 7, 3, 9]; let bank = e.htod_bytes(&w_bytes)?;
let sel = e.htod_i32(&sel_host)?;
for (label, x_rows, x_sstride) in [("shared-x", 1usize, 0usize), ("slot-x", n_sel, in_f)] {
let x_host: Vec<f32> = (0..x_rows * in_f)
.map(|_| (next_u32() % 2000) as f32 / 1000.0 - 1.0)
.collect();
let x_dev = e.htod(&x_host)?;
let mut y_sel = e.uninit(n_sel * out_f)?;
launch_qmatvec_bf16w_sel(
e, &bank, &sel, 0, &x_dev, 0, x_sstride, &mut y_sel, n_sel, in_f, out_f,
)?;
let mut y_ref = e.uninit(n_sel * out_f)?;
for (slot, &eid) in sel_host.iter().enumerate() {
launch_qmatvec_bf16w_off_into(
e,
&bank,
eid as usize * out_f,
&x_dev,
slot * x_sstride,
&mut y_ref,
slot * out_f,
in_f,
out_f,
)?;
}
let (a, b) = (e.dtoh(&y_sel)?, e.dtoh(&y_ref)?);
for (i, (&x1, &x2)) in a.iter().zip(&b).enumerate() {
if x1.to_bits() != x2.to_bits() {
return Err(format!(
"bf16-matvec sel oracle ({label}): idx {i}: sel {x1} vs off_into {x2} \
NOT bit-identical"
)
.into());
}
}
}
}
Ok(format!(
"bf16-matvec kernel oracle: worst abs {:.3e} rel {:.3e} over per-batch + shared-x \
modes, batch>1, t>1, signed/denormal bf16; mt weight-shared twin BIT-IDENTICAL \
at t 5; sel grouped twin BIT-IDENTICAL to the off_into chain (shared-x + slot-x, \
duplicate slots)",
worst.0, worst.1
))
}
fn dequant_nvfp4_expert_f32(
e: &Engine,
codes: &CudaSlice<u8>,
scales: &CudaSlice<u8>,
macro_scale: f32,
expert: usize,
rows: usize,
cols: usize,
) -> Res<CudaSlice<f32>> {
let wbytes = rows * cols / 2;
let sbytes = rows * cols / 16;
let bf = e.alloc_u8(rows * cols * 2)?;
let stream = e.gpu.stream();
let wp = (codes.device_ptr(&stream).0 as usize + expert * wbytes) as *const c_void;
let scp = (scales.device_ptr(&stream).0 as usize + expert * sbytes) as *const c_void;
let dst = bf.device_ptr(&stream).0 as usize as *mut c_void;
let rc = unsafe {
crate::dsv4_ffi::memra_dsv4_nvfp4_deq_bf16(
wp,
scp,
1.0, rows as i32,
cols as i32,
dst,
stream.cu_stream() as *mut c_void,
)
};
if rc != 0 {
return Err(format!("memra_dsv4_nvfp4_deq_bf16 rc={rc}").into());
}
let mut out = e.bf16_to_f32(&bf.slice(0..rows * cols * 2), rows * cols)?;
if macro_scale != 1.0 {
e.scale_inplace(&mut out, macro_scale, rows * cols)?;
}
Ok(out)
}
fn expect(weights: &ReferenceWeights, id: &TensorId) -> Res<ReferenceTensor> {
weights
.get(id)
.cloned()
.ok_or_else(|| format!("qwen4exp_gpu: missing weight {id:?}").into())
}
fn family_id(key: String) -> TensorId {
TensorId::Family {
family: "qwen4_exp",
key,
}
}
fn layer_id(index: u32, tensor: LayerTensor) -> TensorId {
TensorId::Layer { index, tensor }
}
fn upload(e: &Engine, tensor: &ReferenceTensor) -> Res<CudaSlice<f32>> {
e.htod(&tensor.data)
}
fn split_columns(data: &[f32], rows: usize, streams: usize, hidden: usize) -> Vec<Vec<f32>> {
let wide = streams * hidden;
(0..streams)
.map(|s| {
let mut out = Vec::with_capacity(rows * hidden);
for row in 0..rows {
out.extend_from_slice(
&data[row * wide + s * hidden..row * wide + (s + 1) * hidden],
);
}
out
})
.collect()
}
fn split_rows(data: &[f32], streams: usize, hidden: usize, cols: usize) -> Vec<Vec<f32>> {
(0..streams)
.map(|s| data[s * hidden * cols..(s + 1) * hidden * cols].to_vec())
.collect()
}
fn load_gate(
e: &Engine,
weights: &ReferenceWeights,
prefix: &str,
sublayer: &str,
streams: usize,
hidden: usize,
rank: usize,
with_inject: bool,
) -> Res<GateW> {
let wide = streams * hidden;
let norm = expect(
weights,
&family_id(format!("{prefix}{sublayer}hc_norm.weight")),
)?;
let down = expect(
weights,
&family_id(format!("{prefix}{sublayer}input_mix_weight_down.weight")),
)?;
let up = expect(
weights,
&family_id(format!("{prefix}{sublayer}input_mix_weight_up.weight")),
)?;
if norm.data.len() != wide || down.data.len() != rank * wide || up.data.len() != wide * rank {
return Err(format!("qwen4exp_gpu: gate {prefix}{sublayer} shape mismatch").into());
}
let norm_slices = split_rows(&norm.data, streams, hidden, 1);
let down_slices = split_columns(&down.data, rank, streams, hidden);
let up_slices = split_rows(&up.data, streams, hidden, rank);
let stack = |slices: &[Vec<f32>]| -> Vec<f32> {
let mut out = Vec::with_capacity(slices.len() * slices[0].len());
for s in slices {
out.extend_from_slice(s);
}
out
};
let down_b16 = bf16_twin(e, &stack(&down_slices), hidden)?;
let up_b16 = bf16_twin(e, &stack(&up_slices), rank)?;
let (inject, inject_b16) = if with_inject {
let inject = expect(
weights,
&family_id(format!("{prefix}{sublayer}block_inject_weight.weight")),
)?;
if inject.data.len() != streams * wide {
return Err(format!("qwen4exp_gpu: inject {prefix}{sublayer} shape mismatch").into());
}
(
Some(e.htod(&inject.data)?),
bf16_twin(e, &inject.data, hidden)?,
)
} else {
(None, None)
};
Ok(GateW {
norm_stack: e.htod(&stack(&norm_slices))?,
norm: norm_slices
.into_iter()
.map(|v| e.htod(&v))
.collect::<Result<_, _>>()?,
down: down_slices
.into_iter()
.map(|v| e.htod(&v))
.collect::<Result<_, _>>()?,
up: up_slices
.into_iter()
.map(|v| e.htod(&v))
.collect::<Result<_, _>>()?,
inject,
down_b16,
up_b16,
inject_b16,
})
}
#[derive(Default)]
pub struct ExternalParts {
expert_banks: std::collections::BTreeMap<u32, ExpertBank>,
ngram_tables: std::collections::BTreeMap<u32, NgramTable>,
}
#[allow(clippy::too_many_arguments)]
fn build_layer_w(
e: &Engine,
weights: &ReferenceWeights,
layer: &memra_gguf::model_plan::LayerPlan,
prefix: &str,
streams: usize,
hidden: usize,
rank: usize,
bank_override: Option<ExpertBank>,
table_override: Option<NgramTable>,
) -> Res<LayerW> {
let ResidualTopology::GatedResidual { .. } = layer.residual else {
return Err(format!("qwen4exp_gpu: layer {} is not gated-residual", layer.index).into());
};
let attn_gate = load_gate(
e,
weights,
prefix,
"attn_hyper_connection.",
streams,
hidden,
rank,
true,
)?;
let mlp_gate = load_gate(
e,
weights,
prefix,
"mlp_hyper_connection.",
streams,
hidden,
rank,
true,
)?;
let mixer = match &layer.attention {
AttentionPlan::Full(attn) => {
let overlay = layer.sparse_overlay.ok_or_else(|| {
format!(
"qwen4exp_gpu: QSA layer {} has no indexer overlay",
layer.index
)
})?;
let yarn = build_yarn(e, &attn.rope, Some(&overlay), layer.index)?;
if attn.key_head_dim != attn.value_head_dim {
return Err(format!(
"qwen4exp_gpu: layer {} key_head_dim {} != value_head_dim {}",
layer.index, attn.key_head_dim, attn.value_head_dim
)
.into());
}
let load_opt_norm = |tensor: LayerTensor| -> Res<Option<CudaSlice<f32>>> {
match weights.get(&layer_id(layer.index, tensor)) {
Some(t) => Ok(Some(e.htod(&t.data)?)),
None if attn.qk_norm == TensorPresence::Required => {
Err(format!("qwen4exp_gpu: layer {} missing qk norm", layer.index).into())
}
None => Ok(None),
}
};
let wq_t = expect(weights, &layer_id(layer.index, LayerTensor::Query))?;
let wk_t = expect(weights, &layer_id(layer.index, LayerTensor::Key))?;
let wv_t = expect(weights, &layer_id(layer.index, LayerTensor::Value))?;
let wo_t = expect(
weights,
&layer_id(layer.index, LayerTensor::AttentionOutput),
)?;
let o_in = (attn.query_heads * attn.key_head_dim) as usize;
MixerW::Qsa(QsaW {
attn: attn.clone(),
overlay,
yarn,
proj_b16: bf16_stack_twin(e, &[&wq_t.data, &wk_t.data, &wv_t.data], hidden)?,
wo_b16: bf16_twin(e, &wo_t.data, o_in)?,
wq: upload(e, &wq_t)?,
wk: upload(e, &wk_t)?,
wv: upload(e, &wv_t)?,
wo: upload(e, &wo_t)?,
q_norm: load_opt_norm(LayerTensor::QueryNorm)?,
k_norm: load_opt_norm(LayerTensor::KeyNorm)?,
idx_proj: upload(
e,
&expect(
weights,
&family_id(format!("{prefix}self_attn.indexer.index_qk_proj.weight")),
)?,
)?,
idx_q_norm: expect(
weights,
&family_id(format!("{prefix}self_attn.indexer.q_layernorm.weight")),
)?
.data,
idx_k_norm: expect(
weights,
&family_id(format!("{prefix}self_attn.indexer.k_layernorm.weight")),
)?
.data,
})
}
AttentionPlan::GatedDeltaNet(gdn) => {
let qkv_t = expect(weights, &layer_id(layer.index, LayerTensor::GdnQkv))?;
let z_t = expect(weights, &layer_id(layer.index, LayerTensor::GdnGate))?;
let beta_t = expect(weights, &layer_id(layer.index, LayerTensor::GdnBeta))?;
let alpha_t = expect(weights, &layer_id(layer.index, LayerTensor::GdnAlpha))?;
let out_t = expect(weights, &layer_id(layer.index, LayerTensor::GdnOutput))?;
let o_in = (gdn.value_heads * gdn.value_head_dim) as usize;
MixerW::Gdn(GdnW {
plan: *gdn,
proj_b16: bf16_stack_twin(
e,
&[&qkv_t.data, &z_t.data, &beta_t.data, &alpha_t.data],
hidden,
)?,
out_b16: bf16_twin(e, &out_t.data, o_in)?,
qkv: upload(e, &qkv_t)?,
z: upload(e, &z_t)?,
beta: upload(e, &beta_t)?,
alpha: upload(e, &alpha_t)?,
conv_w: upload(
e,
&expect(weights, &layer_id(layer.index, LayerTensor::GdnConv1d))?,
)?,
a: upload(
e,
&expect(weights, &layer_id(layer.index, LayerTensor::GdnA))?,
)?,
dt: upload(
e,
&expect(weights, &layer_id(layer.index, LayerTensor::GdnDtBias))?,
)?,
norm: upload(
e,
&expect(weights, &layer_id(layer.index, LayerTensor::GdnNorm))?,
)?,
out: upload(e, &out_t)?,
})
}
other => {
return Err(format!(
"qwen4exp_gpu: unsupported mixer {other:?} at layer {}",
layer.index
)
.into());
}
};
let MlpPlan::Moe(moe_plan) = &layer.mlp else {
return Err(format!("qwen4exp_gpu: layer {} is not MoE", layer.index).into());
};
if !matches!(moe_plan.router, RouterPlan::Softmax) {
return Err("qwen4exp_gpu: only the softmax router arm is implemented".into());
}
let shared = moe_plan
.shared
.as_ref()
.ok_or("qwen4exp_gpu: missing shared expert plan")?;
let bank = match bank_override {
Some(bank) => bank,
None => {
let gate = expect(
weights,
&layer_id(layer.index, LayerTensor::MoeExpertGateBank),
)?;
let up = expect(
weights,
&layer_id(layer.index, LayerTensor::MoeExpertUpBank),
)?;
let down = expect(
weights,
&layer_id(layer.index, LayerTensor::MoeExpertDownBank),
)?;
let experts = moe_plan.expert_count as usize;
let ff = moe_plan.expert_intermediate_size as usize;
if gate.data.len() != experts * ff * hidden
|| up.data.len() != experts * ff * hidden
|| down.data.len() != experts * hidden * ff
{
return Err(format!(
"qwen4exp_gpu: layer {} expert bank shape mismatch",
layer.index
)
.into());
}
ExpertBank {
gate: BankHalf::F32(e.htod(&gate.data)?),
up: BankHalf::F32(e.htod(&up.data)?),
down: BankHalf::F32(e.htod(&down.data)?),
}
}
};
let sh_gate_t = expect(weights, &layer_id(layer.index, LayerTensor::SharedMlpGate))?;
let sh_up_t = expect(weights, &layer_id(layer.index, LayerTensor::SharedMlpUp))?;
let sh_down_t = expect(weights, &layer_id(layer.index, LayerTensor::SharedMlpDown))?;
let sff = shared.intermediate_size as usize;
let router_t = expect(weights, &layer_id(layer.index, LayerTensor::MoeRouter))?;
let moe = MoeW {
plan: moe_plan.clone(),
router_b16: bf16_twin(e, &router_t.data, hidden)?,
router: upload(e, &router_t)?,
bank,
shared_gu_b16: bf16_stack_twin(e, &[&sh_gate_t.data, &sh_up_t.data], hidden)?,
shared_down_b16: bf16_twin(e, &sh_down_t.data, sff)?,
shared_gate: upload(e, &sh_gate_t)?,
shared_up: upload(e, &sh_up_t)?,
shared_down: upload(e, &sh_down_t)?,
shared_input_gate: if shared.gated {
Some(upload(
e,
&expect(
weights,
&layer_id(layer.index, LayerTensor::SharedMlpInputGate),
)?,
)?)
} else {
None
},
};
let ple = match layer.ple.as_ref() {
None => None,
Some(ple_plan) => {
let embed_dim = ple_plan.embed_dim as usize;
let head_dim = ple_plan.head_embed_dim as usize;
let wide = streams * hidden;
let key_proj = expect(weights, &family_id(format!("{prefix}ple.key_proj.weight")))?;
let conv_w = expect(weights, &family_id(format!("{prefix}ple.conv1d.weight")))?;
if key_proj.data.len() != wide * embed_dim {
return Err("qwen4exp_gpu: ple key_proj shape mismatch".into());
}
let norm_slices = |name: &str| -> Res<Vec<CudaSlice<f32>>> {
let t = expect(weights, &family_id(format!("{prefix}ple.{name}.weight")))?;
split_rows(&t.data, streams, hidden, 1)
.into_iter()
.map(|v| e.htod(&v))
.collect::<Result<_, _>>()
};
let ints = |name: &str| -> Res<Vec<i64>> {
let t = expect(
weights,
&family_id(format!("{prefix}ple.ple_embedding.{name}")),
)?;
t.ints
.clone()
.ok_or_else(|| "qwen4exp_gpu: n-gram buffer must be I64".into())
};
let table = match table_override {
Some(table) => table,
None => {
let t = expect(
weights,
&family_id(format!("{prefix}ple.ple_embedding.ngram_embedding")),
)?;
if t.shape.len() != 2 || t.shape[1] != head_dim {
return Err("qwen4exp_gpu: n-gram table shape mismatch".into());
}
NgramTable::F32(t.data)
}
};
Some(PleW {
plan: *ple_plan,
key_proj: split_rows(&key_proj.data, streams, hidden, embed_dim)
.into_iter()
.map(|v| e.htod(&v))
.collect::<Result<_, _>>()?,
value_proj: upload(
e,
&expect(
weights,
&family_id(format!("{prefix}ple.value_proj.weight")),
)?,
)?,
norm_key: norm_slices("norm_key")?,
norm_query: norm_slices("norm_query")?,
norm_conv: norm_slices("norm_conv")?,
conv_w: split_rows(&conv_w.data, streams, hidden, ple_plan.conv_kernel as usize)
.into_iter()
.map(|v| e.htod(&v))
.collect::<Result<_, _>>()?,
multipliers: ints("layer_multipliers")?,
sizes: ints("ngram_heads_vocab_sizes")?,
offsets: ints("ngram_heads_offsets")?,
table,
})
}
};
Ok(LayerW {
index: layer.index,
eps_attn: layer.pre_attention_norm.epsilon,
eps_mlp: layer.pre_mlp_norm.epsilon,
attn_gate,
mlp_gate,
mixer,
moe,
ple,
})
}
#[allow(clippy::too_many_arguments)]
fn build_mtp_w(
e: &Engine,
weights: &ReferenceWeights,
block: &memra_gguf::model_plan::MtpBlockPlan,
streams: usize,
hidden: usize,
rank: usize,
bank_override: Option<ExpertBank>,
) -> Res<MtpW> {
use memra_gguf::tensor_contract::MtpTensor;
if block.input.fusion != memra_gguf::model_plan::MtpFusionPlan::SeparateProjections {
return Err("qwen4exp_gpu: MTP block is not the separate-projections family".into());
}
let wide = streams * hidden;
let depth = block.depth;
let mtp_id = |tensor: MtpTensor| TensorId::Mtp { depth, tensor };
let pre_e = expect(weights, &mtp_id(MtpTensor::EmbeddingNorm))?;
let pre_h = expect(weights, &mtp_id(MtpTensor::HiddenNorm))?;
let fc_e = expect(weights, &mtp_id(MtpTensor::EmbeddingProjection))?;
let fc_h = expect(weights, &mtp_id(MtpTensor::HiddenProjection))?;
if pre_e.data.len() != hidden
|| pre_h.data.len() != wide
|| fc_e.data.len() != hidden * hidden
|| fc_h.data.len() != hidden * hidden
{
return Err("qwen4exp_gpu: MTP fusion tensor shape mismatch".into());
}
let prefix = format!("mtp.layers.{depth}.");
let layer = build_layer_w(
e,
weights,
&block.layer,
&prefix,
streams,
hidden,
rank,
bank_override,
None,
)?;
let mixer = load_gate(
e,
weights,
"mtp.hyper_connection_mixer.",
"",
streams,
hidden,
rank,
false,
)?;
Ok(MtpW {
eps_embed: block.input.embedding_norm.epsilon,
eps_hidden: block.input.hidden_norm.epsilon,
fc_embed_b16: bf16_twin(e, &fc_e.data, hidden)?,
fc_hidden_b16: bf16_twin(e, &fc_h.data, hidden)?,
pre_norm_embed: upload(e, &pre_e)?,
pre_norm_hidden: upload(e, &pre_h)?,
fc_embed: upload(e, &fc_e)?,
fc_hidden: upload(e, &fc_h)?,
layer,
mixer,
})
}
impl Qwen4ExpGpu {
pub fn from_reference_weights(
e: &Engine,
plan: &ModelPlan,
weights: &ReferenceWeights,
) -> Res<Self> {
Self::from_reference_weights_with(e, None, plan, weights, ExternalParts::default())
}
fn from_reference_weights_with(
e: &Engine,
draft_e: Option<&Engine>,
plan: &ModelPlan,
weights: &ReferenceWeights,
mut parts: ExternalParts,
) -> Res<Self> {
let hidden = plan.hidden_size as usize;
let vocab = plan.vocab_size as usize;
let Some(mixer_plan) = plan.exit_mixer else {
return Err("qwen4exp_gpu requires the gated-residual exit mixer".into());
};
let streams = mixer_plan.streams as usize;
if streams > PLANE_SLOTS.len() {
return Err("qwen4exp_gpu: hc_count exceeds the step-workspace slot table".into());
}
let rank = mixer_plan.bottleneck_rank as usize;
if !plan.logits.is_empty() {
return Err("qwen4exp_gpu: logits transforms are not part of this family".into());
}
let embed = expect(weights, &TensorId::TokenEmbedding)?;
if embed.data.len() != vocab * hidden {
return Err("qwen4exp_gpu: embedding shape mismatch".into());
}
let (output, output_b16) = match weights.get(&TensorId::OutputProjection) {
Some(tensor) => (e.htod(&tensor.data)?, bf16_twin(e, &tensor.data, hidden)?),
None => (e.htod(&embed.data)?, bf16_twin(e, &embed.data, hidden)?),
};
let mut layers = Vec::with_capacity(plan.layers.len());
for layer in &plan.layers {
let prefix = format!("trunk.layers.{}.", layer.index);
layers.push(build_layer_w(
e,
weights,
layer,
&prefix,
streams,
hidden,
rank,
parts.expert_banks.remove(&layer.index),
parts.ngram_tables.remove(&layer.index),
)?);
}
let mtp = match plan.mtp_blocks.first() {
Some(block)
if weights
.get(&TensorId::Mtp {
depth: block.depth,
tensor: memra_gguf::tensor_contract::MtpTensor::EmbeddingProjection,
})
.is_some() =>
{
Some(build_mtp_w(
draft_e.unwrap_or(e),
weights,
block,
streams,
hidden,
rank,
parts.expert_banks.remove(&block.layer.index),
)?)
}
_ => None,
};
let mtp_dev1 = match (draft_e, mtp.as_ref()) {
(Some(de), Some(_)) => {
let head_data: &[f32] = match weights.get(&TensorId::OutputProjection) {
Some(tensor) => &tensor.data,
None => &embed.data,
};
Some(MtpDev1 {
dev: de.ctx().ordinal(),
output: de.htod(head_data)?,
output_b16: bf16_twin(de, head_data, hidden)?,
})
}
(Some(_), None) => {
return Err(
"qwen4exp_gpu: a draft engine was given but no mtp.* rows were \
materialized (LoadOptions::load_mtp)"
.into(),
);
}
_ => None,
};
let exit_mixer = load_gate(
e,
weights,
"trunk.hyper_connection_mixer.",
"",
streams,
hidden,
rank,
false,
)?;
Ok(Self {
plan: plan.clone(),
hidden,
streams,
vocab,
embed_host: embed.data,
output,
output_b16,
layers,
exit_mixer,
exit_eps: plan.output_norm.epsilon,
mtp,
mtp_dev1,
draft_trim: None,
draft_trim_parked: None,
chain_embed: None,
})
}
pub fn build_draft_trim(&mut self, e: &Engine, ids: &[u32]) -> Res<()> {
self.check_draft_engine(e)?;
let n = ids.len();
if n == 0 || n > self.vocab {
return Err(format!("qwen4exp_gpu: draft trim wants 1..={} ids", self.vocab).into());
}
let mut seen = vec![false; self.vocab];
for &id in ids {
let id = id as usize;
if id >= self.vocab {
return Err(format!("qwen4exp_gpu: draft trim id {id} out of vocab").into());
}
if std::mem::replace(&mut seen[id], true) {
return Err(format!("qwen4exp_gpu: draft trim id {id} repeats").into());
}
}
let hidden = self.hidden;
let (src_f32, src_b16) = match self.mtp_dev1.as_ref() {
Some(d) => (&d.output, d.output_b16.as_ref()),
None => (&self.output, self.output_b16.as_ref()),
};
let (head_b16, head) = match src_b16 {
Some(full) => {
let mut trim = e.alloc_u8_uninit(n * hidden * 2)?;
for (row, &id) in ids.iter().enumerate() {
e.copy_u8_range_into(
&mut trim,
row * hidden * 2,
full,
id as usize * hidden * 2,
hidden * 2,
)?;
}
(Some(trim), None)
}
None => {
let mut head = e.uninit(n * hidden)?;
for (row, &id) in ids.iter().enumerate() {
e.copy_range_into(
&mut head,
row * hidden,
src_f32,
id as usize * hidden,
hidden,
)?;
}
(None, Some(head))
}
};
self.draft_trim = Some(DraftTrim {
n,
d2t: ids.to_vec(),
head,
head_b16,
});
self.draft_trim_parked = None;
Ok(())
}
pub fn set_draft_trim(&mut self, on: bool) {
if on {
if let Some(t) = self.draft_trim_parked.take() {
self.draft_trim = Some(t);
}
} else if let Some(t) = self.draft_trim.take() {
self.draft_trim_parked = Some(t);
}
}
pub fn clear_draft_trim(&mut self) {
self.draft_trim = None;
self.draft_trim_parked = None;
}
pub fn arm_spec_devchain(&mut self, de: &Engine) -> Res<()> {
self.check_draft_engine(de)?;
let hidden = self.hidden;
let (rows, for_trim) = match self.draft_trim.as_ref() {
Some(tr) => (tr.n, true),
None => (self.vocab, false),
};
let src_row = |r: usize| -> &[f32] {
let id = match self.draft_trim.as_ref() {
Some(tr) => tr.d2t[r] as usize,
None => r,
};
&self.embed_host[id * hidden..(id + 1) * hidden]
};
let clean = (0..rows).all(|r| src_row(r).iter().all(|x| x.to_bits() & 0xFFFF == 0));
let (bytes, qt, row_bytes) = if clean {
let mut b = vec![0u8; rows * hidden * 2];
for r in 0..rows {
for (j, &x) in src_row(r).iter().enumerate() {
let h = (x.to_bits() >> 16) as u16;
b[(r * hidden + j) * 2..(r * hidden + j) * 2 + 2]
.copy_from_slice(&h.to_le_bytes());
}
}
(b, crate::QT_BF16, hidden * 2)
} else {
let mut b = vec![0u8; rows * hidden * 4];
for r in 0..rows {
for (j, &x) in src_row(r).iter().enumerate() {
b[(r * hidden + j) * 4..(r * hidden + j) * 4 + 4]
.copy_from_slice(&x.to_le_bytes());
}
}
(b, crate::QT_F32, hidden * 4)
};
let table = de.upload_u8(&bytes)?;
println!(
"[qwen4exp-spec] deferred-chain embed table armed: {} rows x {hidden} ({}, {:.1} MiB, dev {}{})",
rows,
if clean {
"bf16 bit-clean"
} else {
"f32 fallback"
},
(rows * row_bytes) as f64 / (1024.0 * 1024.0),
de.ctx().ordinal(),
if for_trim { ", trim-rank order" } else { "" },
);
self.chain_embed = Some(ChainEmbed {
table,
qt,
row_bytes,
rows,
for_trim,
dev: de.ctx().ordinal(),
});
Ok(())
}
pub fn clear_spec_devchain(&mut self) {
self.chain_embed = None;
}
pub fn draft_logits_width(&self) -> usize {
match self.draft_trim.as_ref() {
Some(t) => t.n,
None => self.vocab,
}
}
fn draft_token(&self, row: u32) -> Res<u32> {
match self.draft_trim.as_ref() {
Some(t) => t
.d2t
.get(row as usize)
.copied()
.ok_or_else(|| format!("qwen4exp_gpu: draft row {row} outside the trim").into()),
None => Ok(row),
}
}
pub fn trunk_f32_diet(&mut self, e: &Engine) -> Res<usize> {
if !trunk_bf16_on() || !hc_fused_gate_on() {
return Err(
"qwen4exp_gpu: trunk_f32_diet requires the trunk-bf16 + fused-gate seams ON \
(the bf16 paths must be the ones serving)"
.into(),
);
}
let mut freed = 0usize;
fn stub(e: &Engine, s: &mut CudaSlice<f32>, freed: &mut usize) -> Res<()> {
if s.len() > 1 {
*freed += s.len() * 4;
*s = e.zeros(1)?;
}
Ok(())
}
fn diet_gate(e: &Engine, g: &mut GateW, freed: &mut usize) -> Res<()> {
if g.down_b16.is_none()
|| g.up_b16.is_none()
|| (g.inject.is_some() && g.inject_b16.is_none())
{
return Ok(()); }
for s in g.down.iter_mut() {
stub(e, s, freed)?;
}
for s in g.up.iter_mut() {
stub(e, s, freed)?;
}
if let Some(inj) = g.inject.as_mut() {
stub(e, inj, freed)?;
}
Ok(())
}
for layer in self.layers.iter_mut() {
diet_gate(e, &mut layer.attn_gate, &mut freed)?;
diet_gate(e, &mut layer.mlp_gate, &mut freed)?;
match &mut layer.mixer {
MixerW::Qsa(q) => {
if q.proj_b16.is_some() {
stub(e, &mut q.wq, &mut freed)?;
stub(e, &mut q.wk, &mut freed)?;
stub(e, &mut q.wv, &mut freed)?;
}
if q.wo_b16.is_some() {
stub(e, &mut q.wo, &mut freed)?;
}
}
MixerW::Gdn(g) => {
if g.proj_b16.is_some() {
stub(e, &mut g.qkv, &mut freed)?;
stub(e, &mut g.z, &mut freed)?;
stub(e, &mut g.beta, &mut freed)?;
stub(e, &mut g.alpha, &mut freed)?;
}
if g.out_b16.is_some() {
stub(e, &mut g.out, &mut freed)?;
}
}
}
let moe = &mut layer.moe;
if moe.router_b16.is_some() {
stub(e, &mut moe.router, &mut freed)?;
}
if moe.shared_gu_b16.is_some() {
stub(e, &mut moe.shared_gate, &mut freed)?;
stub(e, &mut moe.shared_up, &mut freed)?;
}
if moe.shared_down_b16.is_some() {
stub(e, &mut moe.shared_down, &mut freed)?;
}
}
diet_gate(e, &mut self.exit_mixer, &mut freed)?;
if self.output_b16.is_some() {
stub(e, &mut self.output, &mut freed)?;
}
Ok(freed)
}
pub fn alloc_state(&self, e: &Engine, capacity: usize) -> Res<Qwen4ExpState> {
self.alloc_state_reserve(e, capacity, capacity, None)
}
pub fn alloc_state_reserve(
&self,
e: &Engine,
capacity: usize,
reserve: usize,
kv_engine: Option<&Engine>,
) -> Res<Qwen4ExpState> {
let kv_e = kv_engine.unwrap_or(e);
if kv_e.ctx().ordinal() != e.ctx().ordinal() {
let limit = peer_kv_max_cap();
if capacity > limit {
return Err(format!(
"qwen4exp_gpu: peer-resident QSA KV refused — capacity {capacity} rows on \
device {} while the attention runs on device {} (ceiling {limit} rows, \
MEMRA_Q4E_PEER_KV_MAX_CAP). The block-list form is the only read path for a \
quantized cache and it is a scatter reader (q4e_sdpa_blocklist_q8q5 phase 1 \
is thread-per-position: 32 lanes on 32 rows, 32 sectors per load \
instruction). Peer memory is not cached in the reading card's L2, so at this \
depth ONE 2,048-token prefill chunk asks ~523 GB across the link and the run \
never finishes — it does not deadlock, it just never arrives. Keep the QSA KV \
on the compute card: it is 10,368 B/row across the 12 QSA layers (2.7 GiB at \
262,144), while the allocation that forces a second card is the MTP draft \
state (~17.6 GiB), which --mtp-dev1 / load_from_dir_dev1 already places \
there.",
kv_e.ctx().ordinal(),
e.ctx().ordinal(),
)
.into());
}
}
let mut layers = Vec::with_capacity(self.layers.len());
for layer in &self.layers {
let mixer = match &layer.mixer {
MixerW::Qsa(qsa) => {
let kv_width = qsa.attn.kv_heads as usize * qsa.attn.key_head_dim as usize;
let v_width = qsa.attn.kv_heads as usize * qsa.attn.value_head_dim as usize;
let kv = if kv_quant_on() {
QsaKvStore::Q8Q5 {
k: kv_e.alloc_u8(capacity * q8_row_bytes(kv_width))?,
v: kv_e.alloc_u8(capacity * q5_row_bytes(v_width))?,
}
} else {
QsaKvStore::F32 {
k: kv_e.zeros(capacity * kv_width)?,
v: kv_e.zeros(capacity * v_width)?,
}
};
MixerState::Qsa {
kv,
raw_keys: IdxRawCache::new(idxq_mode()),
pooled_keys: Vec::new(),
pooled_dev: None,
pooled_dev_rows: 0,
raw_dev: None,
raw_dev_rows: 0,
idx_audit: (idxq_mode() != IdxQMode::F32 && idxq_audit_on()).then(|| {
Box::new(IdxAudit {
raw_f32: IdxRawCache::F32(Vec::new()),
pooled_f32: Vec::new(),
})
}),
}
}
MixerW::Gdn(gdn) => {
let p = &gdn.plan;
let conv_dim = 2 * (p.key_heads * p.key_head_dim) as usize
+ (p.value_heads * p.value_head_dim) as usize;
let pad = p.conv_kernel as usize - 1;
MixerState::Gdn {
conv: e.zeros(pad * conv_dim)?,
state: e
.zeros((p.value_heads * p.value_head_dim * p.key_head_dim) as usize)?,
}
}
};
let ple = match layer.ple.as_ref() {
None => None,
Some(ple) => {
let pad = (ple.plan.conv_kernel as usize - 1) * ple.plan.max_ngram as usize;
let mut conv_hist = Vec::with_capacity(self.streams);
for _ in 0..self.streams {
conv_hist.push(e.zeros(pad * self.hidden)?);
}
Some(PleState {
conv_hist,
ngram_ids: Vec::new(),
ngram_history: Vec::new(),
ngram_last_eos: -1,
})
}
};
layers.push(LayerState { mixer, ple });
}
Ok(Qwen4ExpState {
pos: 0,
capacity,
reserve,
tokens: Vec::new(),
layers,
ws: StepPool::default(),
graphs: StepGraphs::default(),
tp2: None,
verify: None,
})
}
pub fn prefill(&self, e: &Engine, ids: &[u32], state: &mut Qwen4ExpState) -> Res<Vec<f32>> {
self.forward(e, ids, state, None)
}
pub fn prefill_extend(
&self,
e: &Engine,
ids: &[u32],
state: &mut Qwen4ExpState,
chunk: usize,
) -> Res<Vec<f32>> {
if ids.is_empty() || chunk == 0 {
return Err("qwen4exp_gpu: prefill_extend needs ids and a chunk size".into());
}
let mut last = Vec::new();
for piece in ids.chunks(chunk) {
let is_last =
piece.as_ptr() as usize + piece.len() * 4 == ids.as_ptr() as usize + ids.len() * 4;
let head = if is_last {
HeadMode::LastRow
} else {
HeadMode::Skip
};
last = self.forward_with(e, piece, state, None, head)?;
}
Ok(last)
}
pub fn decode_step(&self, e: &Engine, token: u32, state: &mut Qwen4ExpState) -> Res<Vec<f32>> {
self.forward(e, &[token], state, None)
}
pub fn prefill_captured(
&self,
e: &Engine,
ids: &[u32],
state: &mut Qwen4ExpState,
) -> Res<(Vec<f32>, PrefillCapture)> {
let mut capture = PrefillCapture {
layer_wide: Vec::with_capacity(self.layers.len()),
exit_mixed: Vec::new(),
};
let logits = self.forward(e, ids, state, Some(&mut capture))?;
Ok((logits, capture))
}
fn planes_to_wide(&self, e: &Engine, planes: &[CudaSlice<f32>], t: usize) -> Res<Vec<f32>> {
let hidden = self.hidden;
let wide = self.streams * hidden;
let mut out = vec![0.0f32; t * wide];
for (s, plane) in planes.iter().enumerate() {
let host = e.dtoh_view(&plane.slice(0..t * hidden))?;
for row in 0..t {
out[row * wide + s * hidden..row * wide + (s + 1) * hidden]
.copy_from_slice(&host[row * hidden..(row + 1) * hidden]);
}
}
Ok(out)
}
fn forward(
&self,
e: &Engine,
ids: &[u32],
state: &mut Qwen4ExpState,
capture: Option<&mut PrefillCapture>,
) -> Res<Vec<f32>> {
self.forward_with(e, ids, state, capture, HeadMode::All)
}
fn forward_with(
&self,
e: &Engine,
ids: &[u32],
state: &mut Qwen4ExpState,
mut capture: Option<&mut PrefillCapture>,
head: HeadMode,
) -> Res<Vec<f32>> {
let t = ids.len();
let hidden = self.hidden;
if t == 0 {
return Err("qwen4exp_gpu: empty input".into());
}
if head != HeadMode::All {
if capture.is_some() {
return Err("qwen4exp_gpu: prefill capture wants every logits row".into());
}
if let Some(v) = state.verify.as_ref()
&& (t == 1 || t <= v.k_cap)
{
return Err(
"qwen4exp_gpu: head-skipping forward on a verify-exact chunk shape".into(),
);
}
}
if state.pos + t > state.capacity {
return Err("qwen4exp_gpu: state capacity exceeded".into());
}
if state.tp2.is_some() {
return Err(
"qwen4exp_gpu: state already decoded in TP2 mode; single-card forward \
requires a fresh state (the half-state migration is one-way)"
.into(),
);
}
let base_pos = state.pos;
state.tokens.extend_from_slice(ids);
if t > 1 {
state.graphs = StepGraphs::default();
}
let graphs_mode = t == 1
&& decode_graphs_on()
&& step_ws_on()
&& hc_fused_gate_on()
&& !prof::on()
&& capture.is_none()
&& state.verify.is_none();
let tokens = &state.tokens;
let ws = &mut state.ws;
let verify = state.verify.as_mut();
let (exact, vfused, stash_gdn, stash_ple, stash_wide, argmax_sink, last_row_only) =
match verify {
Some(v) => {
let vchunk = base_pos > 0 && t > 1 && t <= v.k_cap;
let vfused = vchunk && verify_fused_on();
let exact = vchunk && !vfused;
let amx_t1 = t == 1 && v.want_argmax_t1;
if exact {
v.chunk = Some((base_pos, t));
v.argmax.clear();
} else if (vfused || amx_t1) && v.want_argmax {
v.argmax.clear();
}
if vfused {
v.fused_chunk = Some((base_pos, t));
}
(
exact,
vfused,
Some(&mut v.gdn),
Some(&mut v.ple),
Some((&mut v.wide, v.ring_rows)),
if (exact || vfused || amx_t1) && v.want_argmax {
Some((&mut v.argmax, &mut v.toks))
} else {
None
},
v.last_row_only && t > 1 && !exact && !vfused,
)
}
None => (false, false, None, None, None, None, false),
};
let mut stash_gdn = stash_gdn;
let mut stash_ple = stash_ple;
let cap = state.reserve.max(t);
let mut planes = prof_section(e, "entry.embed", || {
let mut embedded = vec![0.0f32; t * hidden];
for (row, &token) in ids.iter().enumerate() {
let token = token as usize;
if token >= self.vocab {
return Err(format!("qwen4exp_gpu: token {token} out of range").into());
}
embedded[row * hidden..(row + 1) * hidden]
.copy_from_slice(&self.embed_host[token * hidden..(token + 1) * hidden]);
}
let embedded_dev = ws.take_f32_h2d(e, "entry.embed", &embedded, cap * hidden)?;
let mut planes: Vec<CudaSlice<f32>> = Vec::with_capacity(self.streams);
for s in 0..self.streams {
let mut plane = ws.take_f32(e, PLANE_SLOTS[s], t * hidden, cap * hidden)?;
e.copy_into(&mut plane, 0, &embedded_dev, t * hidden)?;
planes.push(plane);
}
ws.put_f32("entry.embed", embedded_dev);
Ok(planes)
})?;
let ptr_vals: Vec<u64> = {
let stream = e.gpu.stream();
planes.iter().map(|p| p.device_ptr(&stream).0).collect()
};
let ptrs = ws.take_u64_h2d(e, "hc.ptrs", &ptr_vals, 0)?;
if graphs_mode {
if state.graphs.warm {
return self.forward_graphs_tail(e, state, planes, ptrs, base_pos);
}
state.graphs.warm = true;
}
for (li, (layer, lstate)) in self.layers.iter().zip(state.layers.iter_mut()).enumerate() {
if let (Some(ple), Some(ple_state)) = (layer.ple.as_ref(), lstate.ple.as_mut()) {
let ps = if exact {
stash_ple
.as_mut()
.and_then(|v| v.get_mut(li))
.and_then(|s| s.as_mut())
} else {
None
};
self.ple_block(
e,
layer,
ple,
&ple.table,
ple_state,
&mut planes,
tokens,
t,
exact,
ps,
)?;
}
let (mixed, inject) = prof_section(e, "hyper.read", || {
self.gate_read(
e,
ws,
&ptrs,
&layer.attn_gate,
&planes,
t,
layer.eps_attn,
exact,
)
})?;
let block_out = match &layer.mixer {
MixerW::Qsa(qsa) => self.qsa_forward(
e,
ws,
layer,
qsa,
&mixed,
&mut lstate.mixer,
base_pos,
t,
0,
exact,
)?,
MixerW::Gdn(gdn) => {
let gs = if exact {
stash_gdn
.as_mut()
.and_then(|v| v.get_mut(li))
.and_then(|s| s.as_mut())
} else {
None
};
self.gdn_forward(e, ws, layer, gdn, &mixed, &mut lstate.mixer, t, gs)?
}
};
ws.put_f32("hc.mixed", mixed);
prof_section(e, "hyper.write", || {
self.gate_write(e, &mut planes, &ptrs, &block_out, &inject, t)
})?;
ws.put_f32("mixer.out", block_out);
put_inject(ws, inject);
let (mixed, inject) = prof_section(e, "hyper.read", || {
self.gate_read(
e,
ws,
&ptrs,
&layer.mlp_gate,
&planes,
t,
layer.eps_mlp,
exact,
)
})?;
let grouped = exact || vfused || head != HeadMode::All || prefill_grouped_all_on();
let mlp = self.moe_forward(e, ws, &layer.moe, &mixed, t, grouped, layer.index)?;
ws.put_f32("hc.mixed", mixed);
prof_section(e, "hyper.write", || {
self.gate_write(e, &mut planes, &ptrs, &mlp, &inject, t)
})?;
ws.put_f32("moe.out", mlp);
put_inject(ws, inject);
if let Some(capture) = capture.as_deref_mut() {
capture.layer_wide.push(self.planes_to_wide(e, &planes, t)?);
}
}
if let Some((wide_buf, ring_rows)) = stash_wide {
let wide = self.streams * hidden;
for (s, plane) in planes.iter().enumerate() {
for tok in 0..t {
e.copy_range_into(
wide_buf,
((base_pos + tok) % ring_rows) * wide + s * hidden,
plane,
tok * hidden,
hidden,
)?;
}
}
}
if head == HeadMode::Skip {
ws.put_u64("hc.ptrs", ptrs);
state.pos += t;
for (s, plane) in planes.into_iter().enumerate() {
ws.put_f32(PLANE_SLOTS[s], plane);
}
return Ok(Vec::new());
}
let x = prof_section(e, "exit.mixer", || {
Ok(self
.gate_read_inner(
e,
ws,
&ptrs,
&self.exit_mixer,
&planes,
t,
self.exit_eps,
false,
exact,
)?
.0)
})?;
ws.put_u64("hc.ptrs", ptrs);
if let Some(capture) = capture.as_deref_mut() {
capture.exit_mixed = e.dtoh(&x)?;
}
let head_rows = if head == HeadMode::LastRow { 1 } else { t };
let logits = prof_section(e, "lm_head", || {
let mut logits =
ws.take_f32(e, "logits", head_rows * self.vocab, head_rows * self.vocab)?;
let x_head = if head == HeadMode::LastRow {
let mut last = ws.take_f32(e, "exit.last", hidden, hidden)?;
e.copy_range_into(&mut last, 0, &x, (t - 1) * hidden, hidden)?;
last
} else {
x
};
linear_trunk_into(
e,
&self.output,
&self.output_b16,
&x_head,
&mut logits,
head_rows,
hidden,
self.vocab,
)?;
ws.put_f32(
if head == HeadMode::LastRow {
"exit.last"
} else {
"hc.mixed"
},
x_head,
);
Ok(logits)
})?;
state.pos += t;
if head == HeadMode::LastRow {
let out = prof_section(e, "logits.dtoh", || {
Ok(e.dtoh_view(&logits.slice(0..self.vocab))?)
})?;
ws.put_f32("logits", logits);
for (s, plane) in planes.into_iter().enumerate() {
ws.put_f32(PLANE_SLOTS[s], plane);
}
return Ok(out);
}
let out = if let Some((argmax_rows, toks)) = argmax_sink {
prof_section(e, "logits.argmax", || {
for row in 0..t {
e.argmax_token_device_col(&logits, row, self.vocab, toks, row)?;
}
let host = e.gpu.stream().clone_dtoh(&toks.slice(0..t))?;
argmax_rows.extend_from_slice(&host);
Ok(Vec::new())
})?
} else if last_row_only {
prof_section(e, "logits.dtoh", || {
Ok(e.dtoh_view(&logits.slice((t - 1) * self.vocab..t * self.vocab))?)
})?
} else {
prof_section(e, "logits.dtoh", || {
Ok(e.dtoh_view(&logits.slice(0..t * self.vocab))?)
})?
};
ws.put_f32("logits", logits);
for (s, plane) in planes.into_iter().enumerate() {
ws.put_f32(PLANE_SLOTS[s], plane);
}
Ok(out)
}
#[allow(clippy::too_many_arguments)]
fn gate_read(
&self,
e: &Engine,
ws: &mut StepPool,
ptrs: &CudaSlice<u64>,
gate: &GateW,
planes: &[CudaSlice<f32>],
t: usize,
eps: f32,
exact: bool,
) -> Res<(CudaSlice<f32>, InjectOut)> {
self.gate_read_inner(e, ws, ptrs, gate, planes, t, eps, true, exact)
}
#[allow(clippy::too_many_arguments)]
fn gate_read_inner(
&self,
e: &Engine,
ws: &mut StepPool,
ptrs: &CudaSlice<u64>,
gate: &GateW,
planes: &[CudaSlice<f32>],
t: usize,
eps: f32,
with_inject: bool,
exact: bool,
) -> Res<(CudaSlice<f32>, InjectOut)> {
if !hc_fused_gate_on() {
return self.gate_read_legacy(e, ws, gate, planes, t, eps, with_inject);
}
let hidden = self.hidden;
let streams = self.streams;
let rank = gate_rank(gate, hidden, streams)?;
let micro_norm = micro_norm_on();
let micro_inj = micro_inj_on();
if hc_diet_on()
&& (t == 1 || exact)
&& trunk_bf16_on()
&& micro_inj
&& hidden % 8 == 0
&& rank % 8 == 0
&& gate.down_b16.is_some()
&& gate.up_b16.is_some()
&& (!with_inject || gate.inject_b16.is_some())
{
let mut parts = ws.take_f32(e, "hc.parts", t * streams * rank, 0)?;
let mut injp = ws.take_f32(e, "hc.injp", t * streams * streams, 0)?;
let mut inv = ws.take_f32(e, "hc.inv", t * streams, 0)?;
let winj = if with_inject {
gate.inject_b16.as_ref()
} else {
None
};
let mt = t > 1 && verify_mt_on() && (2..=12).contains(&t);
if mt {
launch_hc_diet_stage0_mt(e, ptrs, &mut inv, hidden, streams, t, eps)?;
launch_hc_diet_stage1_mt(
e,
ptrs,
&gate.norm_stack,
&inv,
gate.down_b16.as_ref().expect("guarded above"),
winj,
&mut parts,
&mut injp,
hidden,
rank,
streams,
t,
)?;
} else {
launch_hc_diet_stage1(
e,
ptrs,
&gate.norm_stack,
gate.down_b16.as_ref().expect("guarded above"),
winj,
&mut parts,
&mut injp,
&mut inv,
hidden,
rank,
streams,
t,
eps,
)?;
}
let mut low_act = ws.take_f32(e, "hc.low_act", t * rank, 0)?;
let mut all = ws.take_f32(e, "hc.inj_all", streams * t, 0)?;
launch_hc_diet_stage2(
e,
&parts,
&injp,
&mut low_act,
&mut all,
rank,
streams,
t,
with_inject,
)?;
let mut mixed = ws.take_f32(e, "hc.mixed", t * hidden, 0)?;
if mt && (t * rank + 8 * streams * t) * 4 <= 96 * 1024 {
launch_hc_diet_stage3_mt(
e,
ptrs,
&gate.norm_stack,
&inv,
gate.up_b16.as_ref().expect("guarded above"),
&low_act,
&mut mixed,
hidden,
rank,
streams,
t,
)?;
} else {
launch_hc_diet_stage3(
e,
ptrs,
&gate.norm_stack,
&inv,
gate.up_b16.as_ref().expect("guarded above"),
&low_act,
&mut mixed,
hidden,
rank,
streams,
t,
)?;
}
ws.put_f32("hc.parts", parts);
ws.put_f32("hc.injp", injp);
ws.put_f32("hc.inv", inv);
ws.put_f32("hc.low_act", low_act);
let inject_out = if with_inject {
InjectOut::Slab(all)
} else {
ws.put_f32("hc.inj_all", all);
InjectOut::Rows(Vec::new())
};
return Ok((mixed, inject_out));
}
let mut normed = ws.take_f32(e, "hc.normed", streams * t * hidden, 0)?;
if micro_norm {
launch_hc_norm_planes(
e,
ptrs,
&gate.norm_stack,
&mut normed,
hidden,
t,
streams,
eps,
)?;
} else {
for s in 0..streams {
let mut dst = normed.slice_mut(s * t * hidden..(s + 1) * t * hidden);
launch_rms_norm_into_view(e, &planes[s], &gate.norm[s], &mut dst, hidden, t, eps)?;
}
}
let trunk_b16 = trunk_bf16_on();
let mut parts = ws.take_f32(e, "hc.parts", streams * t * rank, 0)?;
if let (true, Some(w)) = (trunk_b16, gate.down_b16.as_ref()) {
launch_qmatvec_bf16w(
e,
w,
&normed,
&mut parts,
hidden,
rank,
t,
streams,
rank * hidden,
t * hidden,
hidden,
t * rank,
)?;
} else {
if gate.down[0].len() < rank * hidden {
return Err(
"qwen4exp_gpu: gate down f32 dropped (trunk_f32_diet) — keep the \
trunk-bf16 seam ON"
.into(),
);
}
for s in 0..streams {
let x = normed.slice(s * t * hidden..(s + 1) * t * hidden);
let w = gate.down[s].slice(0..rank * hidden);
let mut out = parts.slice_mut(s * t * rank..(s + 1) * t * rank);
e.linear_device_into(&x, &w, &mut out, t, hidden, rank)?;
}
}
let mut low_act = ws.take_f32(e, "hc.low_act", t * rank, 0)?;
launch_hc_lowrank_reduce(e, &parts, &mut low_act, streams, t, rank)?;
ws.put_f32("hc.parts", parts);
let mut gates = ws.take_f32(e, "hc.gates", streams * t * hidden, 0)?;
if let (true, Some(w)) = (trunk_b16, gate.up_b16.as_ref()) {
launch_qmatvec_bf16w(
e,
w,
&low_act,
&mut gates,
rank,
hidden,
t,
streams,
hidden * rank,
0,
rank,
t * hidden,
)?;
} else {
if gate.up[0].len() < hidden * rank {
return Err(
"qwen4exp_gpu: gate up f32 dropped (trunk_f32_diet) — keep the \
trunk-bf16 seam ON"
.into(),
);
}
for s in 0..streams {
let x = low_act.slice(0..t * rank);
let w = gate.up[s].slice(0..hidden * rank);
let mut out = gates.slice_mut(s * t * hidden..(s + 1) * t * hidden);
e.linear_device_into(&x, &w, &mut out, t, rank, hidden)?;
}
}
let mut mixed = ws.take_f32(e, "hc.mixed", t * hidden, 0)?;
launch_hc_mix_epilogue(e, &gates, &normed, &mut mixed, streams, t, hidden)?;
ws.put_f32("hc.gates", gates);
ws.put_f32("hc.low_act", low_act);
let mut inject_out = InjectOut::Rows(Vec::new());
if with_inject {
let inject = gate
.inject
.as_ref()
.ok_or("qwen4exp_gpu: read gate missing inject weights")?;
let inject_dropped = inject.len() < streams * streams * hidden;
let inject_guard = || -> Res<()> {
if inject_dropped {
return Err("qwen4exp_gpu: inject f32 dropped (trunk_f32_diet) — keep \
the trunk-bf16 seam ON"
.into());
}
Ok(())
};
let mut all = ws.take_f32(e, "hc.inj_all", streams * t, 0)?;
if micro_inj {
const CHUNKS: usize = 16;
let mut partials = ws.take_f32(e, "hc.inj_part", streams * t * CHUNKS, 0)?;
let w_b16 = if trunk_b16 {
gate.inject_b16.as_ref()
} else {
None
};
if w_b16.is_none() {
inject_guard()?;
}
launch_hc_inject_two_stage(
e,
&normed,
inject,
w_b16,
&mut partials,
&mut all,
streams,
t,
hidden,
CHUNKS,
)?;
ws.put_f32("hc.inj_part", partials);
inject_out = InjectOut::Slab(all);
} else {
if let (true, Some(w)) = (trunk_b16, gate.inject_b16.as_ref()) {
launch_hc_inject_gates_b16(e, &normed, w, &mut all, streams, t, hidden)?;
} else {
inject_guard()?;
launch_hc_inject_gates(e, &normed, inject, &mut all, streams, t, hidden)?;
}
let mut rows = Vec::with_capacity(streams);
for s in 0..streams {
let mut row = ws.take_f32(e, INJECT_SLOTS[s], t, 0)?;
e.copy_range_into(&mut row, 0, &all, s * t, t)?;
rows.push(row);
}
ws.put_f32("hc.inj_all", all);
inject_out = InjectOut::Rows(rows);
}
}
ws.put_f32("hc.normed", normed);
Ok((mixed, inject_out))
}
#[allow(clippy::too_many_arguments)]
fn gate_read_legacy(
&self,
e: &Engine,
_ws: &mut StepPool,
gate: &GateW,
planes: &[CudaSlice<f32>],
t: usize,
eps: f32,
with_inject: bool,
) -> Res<(CudaSlice<f32>, InjectOut)> {
let hidden = self.hidden;
let streams = self.streams;
let rank = gate_rank(gate, hidden, streams)?;
if gate.down[0].len() < rank * hidden {
return Err(
"qwen4exp_gpu: gate f32 originals dropped (trunk_f32_diet) — the \
legacy gate path needs them (keep hc seams ON)"
.into(),
);
}
let inv_streams = 1.0 / streams as f32;
let mut normed = Vec::with_capacity(streams);
for s in 0..streams {
let mut dst = e.uninit(t * hidden)?;
e.rms_norm(&planes[s], &gate.norm[s], &mut dst, hidden, t, eps)?;
normed.push(dst);
}
let mut low = e.linear(&normed[0], &gate.down[0], t, hidden, rank)?;
for s in 1..streams {
let part = e.linear(&normed[s], &gate.down[s], t, hidden, rank)?;
let mut view = low.slice_mut(0..t * rank);
e.axpy_into(&part, 1.0, &mut view, t * rank)?;
}
e.scale_inplace(&mut low, inv_streams, t * rank)?;
let ones = e.htod(&vec![1.0f32; t * rank.max(1)])?;
let mut low_act = e.uninit(t * rank)?;
e.silu_mul(&low, &ones, &mut low_act, t * rank)?;
let mut mixed = e.zeros(t * hidden)?;
let mut gate_buf = e.uninit(t * hidden)?;
let mut prod = e.uninit(t * hidden)?;
for s in 0..streams {
let g = e.linear(&low_act, &gate.up[s], t, rank, hidden)?;
e.sigmoid(&g, &mut gate_buf, t * hidden)?;
e.mul(&gate_buf, &normed[s], &mut prod, t * hidden)?;
let mut view = mixed.slice_mut(0..t * hidden);
e.axpy_into(&prod, 1.0, &mut view, t * hidden)?;
}
e.scale_inplace(&mut mixed, inv_streams, t * hidden)?;
let mut inject_out = Vec::new();
if with_inject {
let inject = gate
.inject
.as_ref()
.ok_or("qwen4exp_gpu: read gate missing inject weights")?;
let wide = streams * hidden;
for s in 0..streams {
let mut acc = {
let w = inject.slice(s * wide..s * wide + hidden);
let x = normed[0].slice(0..t * hidden);
let mut out = e.uninit(t)?;
e.linear_device_into(&x, &w, &mut out, t, hidden, 1)?;
out
};
for s2 in 1..streams {
let w = inject.slice(s * wide + s2 * hidden..s * wide + (s2 + 1) * hidden);
let x = normed[s2].slice(0..t * hidden);
let mut part = e.uninit(t)?;
e.linear_device_into(&x, &w, &mut part, t, hidden, 1)?;
let mut view = acc.slice_mut(0..t);
e.axpy_into(&part, 1.0, &mut view, t)?;
}
e.scale_inplace(&mut acc, inv_streams, t)?;
let mut sg = e.uninit(t)?;
e.sigmoid(&acc, &mut sg, t)?;
e.scale_inplace(&mut sg, 2.0, t)?;
inject_out.push(sg);
}
}
Ok((mixed, InjectOut::Rows(inject_out)))
}
fn gate_write(
&self,
e: &Engine,
planes: &mut [CudaSlice<f32>],
ptrs: &CudaSlice<u64>,
block_out: &CudaSlice<f32>,
inject: &InjectOut,
t: usize,
) -> Res<()> {
match inject {
InjectOut::Rows(rows) => {
for (plane, inj) in planes.iter_mut().zip(rows) {
e.add_scaled_rows(block_out, inj, plane, self.hidden, t)?;
}
Ok(())
}
InjectOut::Slab(slab) => {
launch_hc_write_planes(e, ptrs, block_out, slab, self.hidden, t, self.streams)
}
}
}
#[allow(clippy::too_many_arguments)]
fn layer_interior(
&self,
e: &Engine,
ws: &mut StepPool,
ptrs: &CudaSlice<u64>,
layer: &LayerW,
lstate: &mut LayerState,
planes: &mut [CudaSlice<f32>],
tokens: &[u32],
base_pos: usize,
) -> Res<()> {
if let (Some(ple), Some(ple_state)) = (layer.ple.as_ref(), lstate.ple.as_mut()) {
self.ple_block(
e, layer, ple, &ple.table, ple_state, planes, tokens, 1, false, None,
)?;
}
let (mixed, inject) = self.gate_read(
e,
ws,
ptrs,
&layer.attn_gate,
planes,
1,
layer.eps_attn,
false,
)?;
let block_out = match &layer.mixer {
MixerW::Qsa(qsa) => self.qsa_forward(
e,
ws,
layer,
qsa,
&mixed,
&mut lstate.mixer,
base_pos,
1,
0,
false,
)?,
MixerW::Gdn(gdn) => {
self.gdn_forward(e, ws, layer, gdn, &mixed, &mut lstate.mixer, 1, None)?
}
};
ws.put_f32("hc.mixed", mixed);
self.gate_write(e, planes, ptrs, &block_out, &inject, 1)?;
ws.put_f32("mixer.out", block_out);
put_inject(ws, inject);
let (mixed, inject) = self.gate_read(
e,
ws,
ptrs,
&layer.mlp_gate,
planes,
1,
layer.eps_mlp,
false,
)?;
ws.put_f32("hc.mixed", mixed);
put_inject(ws, inject);
Ok(())
}
fn moe_route_slots(&self, e: &Engine, ws: &mut StepPool, moe: &MoeW, layer: u32) -> Res<()> {
let hidden = self.hidden;
let experts = moe.plan.expert_count as usize;
let selected = moe.plan.experts_per_token as usize;
let mixed = ws.take_f32(e, "hc.mixed", hidden, 0)?;
let mut router_out = ws.take_f32(e, "moe.router", experts, 0)?;
let none: Option<CudaSlice<u8>> = None;
let rb = if router_bf16_on() {
&moe.router_b16
} else {
&none
};
linear_trunk_into(
e,
&moe.router,
rb,
&mixed,
&mut router_out,
1,
hidden,
experts,
)?;
if router_dev_on() && route_dev_geometry(experts, selected) {
let mut sel = ws.take_i32_slot(e, "moe.sel", selected, 0)?;
let mut w = ws.take_f32(e, "moe.w", selected, 0)?;
route_topk_device(
e,
&router_out,
&mut sel,
&mut w,
None,
experts,
selected,
1,
layer,
)?;
if route_sync_diag() {
e.gpu.stream().synchronize()?;
}
ws.put_i32("moe.sel", sel);
ws.put_f32("moe.w", w);
ws.put_f32("moe.router", router_out);
ws.put_f32("hc.mixed", mixed);
return Ok(());
}
let logits = e.dtoh_view(&router_out.slice(0..experts))?;
ws.put_f32("moe.router", router_out);
ws.put_f32("hc.mixed", mixed);
let route = host_route_softmax_topk(&logits, selected);
let sel_host: Vec<i32> = route.iter().map(|&(x, _)| x as i32).collect();
let w_host: Vec<f32> = route.iter().map(|&(_, w)| w).collect();
ws.write_i32(e, "moe.sel", &sel_host)?;
ws.write_f32(e, "moe.w", &w_host)?;
Ok(())
}
fn moe_grouped_tail_slots(
&self,
e: &Engine,
ws: &mut StepPool,
ptrs: &CudaSlice<u64>,
moe: &MoeW,
planes: &mut [CudaSlice<f32>],
) -> Res<()> {
let hidden = self.hidden;
let ff = moe.plan.expert_intermediate_size as usize;
let n_sel = moe.plan.experts_per_token as usize;
let (
BankHalf::Nvfp4 {
codes: gc,
scales: gs,
macros_dev: gm,
..
},
BankHalf::Nvfp4 {
codes: uc,
scales: us,
macros_dev: um,
..
},
BankHalf::Nvfp4 {
codes: dc,
scales: ds,
macros_dev: dm,
..
},
) = (&moe.bank.gate, &moe.bank.up, &moe.bank.down)
else {
return Err("qwen4exp_gpu: grouped tail on a non-NVFP4 bank".into());
};
let mixed = ws.take_f32(e, "hc.mixed", hidden, 0)?;
let sel = ws
.i32s
.remove("moe.sel")
.ok_or("step workspace: moe.sel is not parked")?;
let w_dev = ws.take_f32(e, "moe.w", n_sel, 0)?;
let mut act = ws.take_f32(e, "moe.act", n_sel * ff, 0)?;
if sel_gufuse_on() && hidden % 32 == 0 && ff % 4 == 0 {
launch_nvfp4_sel_gu_silu(
e,
(gc, gs, gm),
(uc, us, um),
Some(&sel),
0,
n_sel,
&mixed,
&mut act,
hidden,
ff,
None,
)?;
} else {
let mut yg = ws.take_f32(e, "moe.yg", n_sel * ff, 0)?;
let mut yu = ws.take_f32(e, "moe.yu", n_sel * ff, 0)?;
launch_nvfp4_sel_matvec(e, gc, gs, gm, &sel, &mixed, &mut yg, n_sel, hidden, ff, 0)?;
launch_nvfp4_sel_matvec(e, uc, us, um, &sel, &mixed, &mut yu, n_sel, hidden, ff, 0)?;
e.silu_mul(&yg, &yu, &mut act, n_sel * ff)?;
ws.put_f32("moe.yg", yg);
ws.put_f32("moe.yu", yu);
}
let mut partial = ws.take_f32(e, "moe.partial", n_sel * hidden, 0)?;
launch_nvfp4_sel_matvec(
e,
dc,
ds,
dm,
&sel,
&act,
&mut partial,
n_sel,
ff,
hidden,
ff,
)?;
let mut out = ws.take_f32(e, "moe.out", hidden, 0)?;
e.axpy_rows_seq_into(&partial, &w_dev, &mut out, hidden, n_sel)?;
ws.put_i32("moe.sel", sel);
ws.put_f32("moe.w", w_dev);
ws.put_f32("moe.act", act);
ws.put_f32("moe.partial", partial);
let out = self.moe_shared_tail(e, ws, moe, &mixed, out, 1)?;
ws.put_f32("hc.mixed", mixed);
let inject = take_inject(e, ws, self.streams, 1)?;
self.gate_write(e, planes, ptrs, &out, &inject, 1)?;
ws.put_f32("moe.out", out);
put_inject(ws, inject);
Ok(())
}
fn forward_graphs_tail(
&self,
e: &Engine,
state: &mut Qwen4ExpState,
mut planes: Vec<CudaSlice<f32>>,
ptrs: CudaSlice<u64>,
base_pos: usize,
) -> Res<Vec<f32>> {
let mut graphs = std::mem::take(&mut state.graphs);
if graphs.a.len() != self.layers.len() {
graphs.a = (0..self.layers.len()).map(|_| None).collect();
graphs.b = (0..self.layers.len()).map(|_| None).collect();
}
let ws = &mut state.ws;
let tokens = &state.tokens;
for (li, (layer, lstate)) in self.layers.iter().zip(state.layers.iter_mut()).enumerate() {
let a_ok = matches!(layer.mixer, MixerW::Gdn(_)) && layer.ple.is_none();
if a_ok {
if graphs.a[li].is_none() {
graphs.a[li] = Some(e.capture_graph_retained_nowarm(|eng| {
self.layer_interior(
eng,
ws,
&ptrs,
layer,
lstate,
&mut planes,
tokens,
base_pos,
)
})?);
}
graphs.a[li].as_ref().unwrap().0.launch()?;
} else {
self.layer_interior(e, ws, &ptrs, layer, lstate, &mut planes, tokens, base_pos)?;
}
let b_ok = moe_sel_path_on()
&& matches!(
(
&layer.moe.bank.gate,
&layer.moe.bank.up,
&layer.moe.bank.down
),
(
BankHalf::Nvfp4 { .. },
BankHalf::Nvfp4 { .. },
BankHalf::Nvfp4 { .. }
)
);
if b_ok {
self.moe_route_slots(e, ws, &layer.moe, layer.index)?;
if graphs.b[li].is_none() {
graphs.b[li] = Some(e.capture_graph_retained_nowarm(|eng| {
self.moe_grouped_tail_slots(eng, ws, &ptrs, &layer.moe, &mut planes)
})?);
}
graphs.b[li].as_ref().unwrap().0.launch()?;
} else {
let mixed = ws.take_f32(e, "hc.mixed", self.hidden, 0)?;
let mlp = self.moe_forward(e, ws, &layer.moe, &mixed, 1, false, layer.index)?;
ws.put_f32("hc.mixed", mixed);
let inject = take_inject(e, ws, self.streams, 1)?;
self.gate_write(e, &mut planes, &ptrs, &mlp, &inject, 1)?;
ws.put_f32("moe.out", mlp);
put_inject(ws, inject);
}
}
if graphs.exit.is_none() {
graphs.exit = Some(e.capture_graph_retained_nowarm(|eng| {
let x = self
.gate_read_inner(
eng,
ws,
&ptrs,
&self.exit_mixer,
&planes,
1,
self.exit_eps,
false,
false,
)?
.0;
let mut logits = ws.take_f32(eng, "logits", self.vocab, 0)?;
linear_trunk_into(
eng,
&self.output,
&self.output_b16,
&x,
&mut logits,
1,
self.hidden,
self.vocab,
)?;
ws.put_f32("hc.mixed", x);
ws.put_f32("logits", logits);
Ok(())
})?);
}
graphs.exit.as_ref().unwrap().0.launch()?;
let out = {
let logits = ws.peek_f32("logits")?;
e.dtoh_view(&logits.slice(0..self.vocab))?
};
for (s, plane) in planes.into_iter().enumerate() {
ws.put_f32(PLANE_SLOTS[s], plane);
}
ws.put_u64("hc.ptrs", ptrs);
state.pos += 1;
state.graphs = graphs;
Ok(out)
}
#[allow(clippy::too_many_arguments)]
fn qsa_update_select(
&self,
e: &Engine,
ws: &mut StepPool,
qsa: &QsaW,
eps: f32,
mixed: &CudaSlice<f32>,
raw_keys: &mut IdxRawCache,
pooled_keys: &mut Vec<f32>,
pooled_dev: &mut Option<CudaSlice<f32>>,
pooled_dev_rows: &mut usize,
raw_dev: &mut Option<IdxRawDev>,
raw_dev_rows: &mut usize,
mut idx_audit: Option<&mut Box<IdxAudit>>,
base_pos: usize,
t: usize,
pos_off: usize,
exact: bool,
) -> Res<Vec<RowSel>> {
let hidden = self.hidden;
let base = qsa.attn.rope.base;
let t_kv = base_pos + t;
let overlay = &qsa.overlay;
let idx_dim = overlay.head_dim as usize;
let qk_width = (overlay.query_heads as usize + overlay.kv_heads as usize) * idx_dim;
if overlay.kv_heads != 1 {
return Err("qwen4exp_gpu: indexer with more than one key head".into());
}
let idx_proj = prof_section(e, "qsa.idx_proj", || {
let mut idx_proj = ws.take_f32(e, "qsa.idxp", t * qk_width, 0)?;
if exact && t > 1 {
let wv = qsa.idx_proj.slice(0..qsa.idx_proj.len());
for tok in 0..t {
let xv = mixed.slice(tok * hidden..(tok + 1) * hidden);
let mut yv = idx_proj.slice_mut(tok * qk_width..(tok + 1) * qk_width);
e.linear_device_into(&xv, &wv, &mut yv, 1, hidden, qk_width)?;
}
} else {
e.linear_device_into(mixed, &qsa.idx_proj, &mut idx_proj, t, hidden, qk_width)?;
}
Ok(idx_proj)
})?;
let dev_cache = idx_cache_on();
let block_size = overlay.block_size as usize;
let all_full = (base_pos + t) / block_size <= overlay.budget_blocks as usize;
let host_rows = raw_keys.rows(idx_dim);
if *raw_dev_rows > host_rows && !(dev_cache && all_full) {
idx_materialize_host(e, raw_keys, raw_dev, *raw_dev_rows, idx_dim)?;
}
if dev_cache {
let host_rows = raw_keys.rows(idx_dim);
let base_rows = (*raw_dev_rows).max(host_rows);
let cap_rows = (base_rows + t).next_power_of_two().max(64);
let q_off = overlay.query_heads as usize * idx_dim;
match &mut *raw_keys {
IdxRawCache::F32(h) => {
let want = (base_rows + t) * idx_dim;
let grow = match raw_dev.as_ref() {
Some(IdxRawDev::F32(m)) => m.len() < want,
Some(_) => return Err("idxcache: device format lag on f32".into()),
None => true,
};
if grow {
let mut fresh = e.uninit(cap_rows * idx_dim)?;
if let (Some(IdxRawDev::F32(old)), rows) = (raw_dev.as_ref(), *raw_dev_rows)
{
if rows > 0 {
e.copy_range_into(&mut fresh, 0, old, 0, rows * idx_dim)?;
}
}
*raw_dev = Some(IdxRawDev::F32(fresh));
}
let Some(IdxRawDev::F32(m)) = raw_dev.as_mut() else {
unreachable!("allocated above");
};
if host_rows > *raw_dev_rows {
let mut view = m.slice_mut(*raw_dev_rows * idx_dim..host_rows * idx_dim);
e.gpu
.stream()
.memcpy_htod(&h[*raw_dev_rows * idx_dim..], &mut view)?;
*raw_dev_rows = host_rows;
}
launch_copy_rows_col(
e,
&idx_proj,
m,
t,
idx_dim,
qk_width,
q_off,
*raw_dev_rows,
)?;
}
IdxRawCache::Q8(h) => {
let rb = q8_row_bytes(idx_dim);
let want = (base_rows + t) * rb;
let grow = match raw_dev.as_ref() {
Some(IdxRawDev::Q8(m)) => m.len() < want,
Some(_) => return Err("idxcache: device format lag on q8".into()),
None => true,
};
if grow {
let mut fresh = e.alloc_u8_uninit(cap_rows * rb)?;
if let (Some(IdxRawDev::Q8(old)), rows) = (raw_dev.as_ref(), *raw_dev_rows)
{
if rows > 0 {
let mut dst = fresh.slice_mut(0..rows * rb);
e.gpu
.stream()
.memcpy_dtod(&old.slice(0..rows * rb), &mut dst)?;
}
}
*raw_dev = Some(IdxRawDev::Q8(fresh));
}
let Some(IdxRawDev::Q8(m)) = raw_dev.as_mut() else {
unreachable!("allocated above");
};
if host_rows > *raw_dev_rows {
let mut view = m.slice_mut(*raw_dev_rows * rb..host_rows * rb);
e.gpu
.stream()
.memcpy_htod(&h[*raw_dev_rows * rb..host_rows * rb], &mut view)?;
*raw_dev_rows = host_rows;
}
launch_q4e_idx_append_q8(
e,
&idx_proj,
m,
t,
idx_dim,
qk_width,
q_off,
*raw_dev_rows,
)?;
}
IdxRawCache::Bf16(h) => {
let want = (base_rows + t) * idx_dim;
let grow = match raw_dev.as_ref() {
Some(IdxRawDev::Bf16(m)) => m.len() < want,
Some(_) => return Err("idxcache: device format lag on bf16".into()),
None => true,
};
if grow {
let mut fresh = unsafe { e.gpu.stream().alloc::<u16>(cap_rows * idx_dim)? };
if let (Some(IdxRawDev::Bf16(old)), rows) =
(raw_dev.as_ref(), *raw_dev_rows)
{
if rows > 0 {
let mut dst = fresh.slice_mut(0..rows * idx_dim);
e.gpu
.stream()
.memcpy_dtod(&old.slice(0..rows * idx_dim), &mut dst)?;
}
}
*raw_dev = Some(IdxRawDev::Bf16(fresh));
}
let Some(IdxRawDev::Bf16(m)) = raw_dev.as_mut() else {
unreachable!("allocated above");
};
if host_rows > *raw_dev_rows {
let mut view = m.slice_mut(*raw_dev_rows * idx_dim..host_rows * idx_dim);
e.gpu.stream().memcpy_htod(
&h[*raw_dev_rows * idx_dim..host_rows * idx_dim],
&mut view,
)?;
*raw_dev_rows = host_rows;
}
launch_q4e_idx_append_bf16(
e,
&idx_proj,
m,
t,
idx_dim,
qk_width,
q_off,
*raw_dev_rows,
)?;
}
}
*raw_dev_rows += t;
}
if let Some(audit) = idx_audit.as_deref_mut() {
let q_off = overlay.query_heads as usize * idx_dim;
let rows_f = e.dtoh_view(&idx_proj.slice(0..t * qk_width))?;
let IdxRawCache::F32(twin) = &mut audit.raw_f32 else {
return Err("idxq audit: twin cache is not f32".into());
};
for row in 0..t {
twin.extend_from_slice(&rows_f[row * qk_width + q_off..(row + 1) * qk_width]);
}
}
let sels: Vec<RowSel> = if dev_cache && all_full {
ws.put_f32("qsa.idxp", idx_proj);
(0..t)
.map(|qt| RowSel {
full: true,
blocks: Vec::new(),
visible: base_pos + qt + 1,
})
.collect()
} else {
let idx_rows = e.dtoh_view(&idx_proj.slice(0..t * qk_width))?;
ws.put_f32("qsa.idxp", idx_proj);
let q_off = overlay.query_heads as usize * idx_dim;
for row in 0..t {
raw_keys.append_rows_f32(
&idx_rows[row * qk_width + q_off..(row + 1) * qk_width],
1,
idx_dim,
);
}
let dev_scorer = idx_dev_on();
let sels = prof_section(e, "qsa.idx_host", || {
indexer_select_rows(
overlay,
base,
qsa.yarn.as_ref().map(|y| (y.ff_host.as_slice(), y.mscale)),
eps,
&qsa.idx_q_norm,
&qsa.idx_k_norm,
&idx_rows,
raw_keys,
pooled_keys,
if dev_scorer {
Some((e, pooled_dev, pooled_dev_rows))
} else {
None
},
base_pos,
t,
t_kv,
pos_off,
)
})?;
if let Some(audit) = idx_audit.as_deref_mut() {
if t <= 8 && sels.iter().any(|s| !s.full) {
let twin_sels = indexer_select_rows(
overlay,
base,
qsa.yarn.as_ref().map(|y| (y.ff_host.as_slice(), y.mscale)),
eps,
&qsa.idx_q_norm,
&qsa.idx_k_norm,
&idx_rows,
&audit.raw_f32,
&mut audit.pooled_f32,
None,
base_pos,
t,
t_kv,
pos_off,
)?;
use std::sync::atomic::Ordering::Relaxed;
for (a, b) in sels.iter().zip(&twin_sels) {
if a.full && b.full {
continue;
}
IDXQ_AUDIT_ROWS.fetch_add(1, Relaxed);
if a.full != b.full || a.blocks != b.blocks {
IDXQ_AUDIT_FLIPPED.fetch_add(1, Relaxed);
let mut diff = 0u64;
let (sa, sb) = (&a.blocks, &b.blocks);
let seta: std::collections::BTreeSet<_> = sa.iter().collect();
let setb: std::collections::BTreeSet<_> = sb.iter().collect();
diff += seta.symmetric_difference(&setb).count() as u64;
IDXQ_AUDIT_BLOCKS.fetch_add(diff, Relaxed);
}
}
}
}
sels
};
Ok(sels)
}
fn qsa_forward(
&self,
e: &Engine,
ws: &mut StepPool,
layer: &LayerW,
qsa: &QsaW,
mixed: &CudaSlice<f32>,
mstate: &mut MixerState,
base_pos: usize,
t: usize,
pos_off: usize,
exact: bool,
) -> Res<CudaSlice<f32>> {
let MixerState::Qsa {
kv,
raw_keys,
pooled_keys,
pooled_dev,
pooled_dev_rows,
raw_dev,
raw_dev_rows,
idx_audit,
} = mstate
else {
return Err(format!(
"qwen4exp_gpu: QSA layer {} bound to non-QSA state",
layer.index
)
.into());
};
let hidden = self.hidden;
let nh = qsa.attn.query_heads as usize;
let nkv = qsa.attn.kv_heads as usize;
let hd = qsa.attn.key_head_dim as usize;
let eps = layer.eps_attn;
let cap = kv.capacity_rows(nkv * hd);
let n_rot = qsa.attn.rope.dimensions as usize;
let base = qsa.attn.rope.base;
let (q, gate) = prof_section(e, "qsa.proj", || {
let mut q_fused = ws.take_f32(e, "qsa.qf", t * 2 * nh * hd, 0)?;
let mut k_new = ws.take_f32(e, "qsa.k", t * nkv * hd, 0)?;
let mut v_new = ws.take_f32(e, "qsa.v", t * nkv * hd, 0)?;
if let (true, Some(stack)) = (
t == 1 && proj_stack_on() && trunk_bf16_on(),
qsa.proj_b16.as_ref(),
) {
launch_qmatvec_bf16w_multi4(
e,
stack,
mixed,
&[
(&q_fused, 2 * nh * hd),
(&k_new, nkv * hd),
(&v_new, nkv * hd),
],
hidden,
)?;
} else {
linear_trunk_stacked_into(
e,
&qsa.wq,
&qsa.proj_b16,
0,
mixed,
&mut q_fused,
t,
hidden,
2 * nh * hd,
)?;
linear_trunk_stacked_into(
e,
&qsa.wk,
&qsa.proj_b16,
2 * nh * hd,
mixed,
&mut k_new,
t,
hidden,
nkv * hd,
)?;
linear_trunk_stacked_into(
e,
&qsa.wv,
&qsa.proj_b16,
2 * nh * hd + nkv * hd,
mixed,
&mut v_new,
t,
hidden,
nkv * hd,
)?;
}
let mut q = ws.take_f32(e, "qsa.q", t * nh * hd, 0)?;
let mut gate = ws.take_f32(e, "qsa.gate", t * nh * hd, 0)?;
e.q_gate_split(&q_fused, &mut q, &mut gate, hd, nh, t)?;
ws.put_f32("qsa.qf", q_fused);
let mut q = if let Some(norm) = qsa.q_norm.as_ref() {
let mut dst = ws.take_f32(e, "qsa.qn", t * nh * hd, 0)?;
e.rms_norm(&q, norm, &mut dst, hd, t * nh, eps)?;
ws.put_f32("qsa.q", q);
dst
} else {
q
};
let mut k_new = if let Some(norm) = qsa.k_norm.as_ref() {
let mut dst = ws.take_f32(e, "qsa.kn", t * nkv * hd, 0)?;
e.rms_norm(&k_new, norm, &mut dst, hd, t * nkv, eps)?;
ws.put_f32("qsa.k", k_new);
dst
} else {
k_new
};
let positions: Vec<i32> = (0..t).map(|i| (base_pos + i + pos_off) as i32).collect();
let pos_dev = ws.take_i32(e, "qsa.pos", &positions, 0)?;
if let Some(yarn) = qsa.yarn.as_ref() {
e.rope_neox_ffm(
&mut q,
&pos_dev,
hd,
n_rot,
nh,
t,
base,
1.0,
&yarn.ff,
yarn.mscale,
)?;
e.rope_neox_ffm(
&mut k_new,
&pos_dev,
hd,
n_rot,
nkv,
t,
base,
1.0,
&yarn.ff,
yarn.mscale,
)?;
} else {
e.rope_neox(&mut q, &pos_dev, hd, n_rot, nh, t, base, 1.0)?;
e.rope_neox(&mut k_new, &pos_dev, hd, n_rot, nkv, t, base, 1.0)?;
}
ws.put_i32("qsa.pos", pos_dev);
match kv {
QsaKvStore::F32 { k, v } => {
e.copy_range_into(k, base_pos * nkv * hd, &k_new, 0, t * nkv * hd)?;
e.copy_range_into(v, base_pos * nkv * hd, &v_new, 0, t * nkv * hd)?;
}
QsaKvStore::Q8Q5 { k, v } => {
launch_q4e_kv_append(e, &k_new, &v_new, k, v, base_pos, t, nkv * hd)?;
}
}
ws.put_f32(
if qsa.k_norm.is_some() {
"qsa.kn"
} else {
"qsa.k"
},
k_new,
);
ws.put_f32("qsa.v", v_new);
Ok((q, gate))
})?;
let t_kv = base_pos + t;
let sels = self.qsa_update_select(
e,
ws,
qsa,
eps,
mixed,
raw_keys,
pooled_keys,
pooled_dev,
pooled_dev_rows,
raw_dev,
raw_dev_rows,
idx_audit.as_mut(),
base_pos,
t,
pos_off,
exact,
)?;
let overlay = &qsa.overlay;
let scale = match qsa.attn.scale {
memra_gguf::model_plan::AttentionScale::InverseSqrtKeyDim => 1.0 / (hd as f32).sqrt(),
memra_gguf::model_plan::AttentionScale::Fixed(scale) => scale,
};
let long_att = if kv.is_quant() {
if longatt_mode() == LongAttMode::Off {
return Err(
"qwen4exp_gpu: kvq requires the block-list attention form (longatt=off)".into(),
);
}
true
} else {
match longatt_mode() {
LongAttMode::Force => true,
LongAttMode::Auto => t_kv > SDPA_MASK_TKV_BOUND || sels.iter().any(|s| !s.full),
LongAttMode::Off => false,
}
};
let block_size = overlay.block_size as usize;
let attended = if long_att {
let (pos_flat, meta, max_count) = rowsel_positions(&sels, block_size);
let pos_dev = prof_section(e, "qsa.mask_h2d", || {
ws.take_i32(e, "qsa.selpos", &pos_flat, 0)
})?;
let meta_dev = ws.take_i32(e, "qsa.selmeta", &meta, 0)?;
let attended = prof_section(e, "qsa.sdpa", || {
let mut attended = ws.take_f32(e, "qsa.att", t * nh * hd, 0)?;
match kv {
QsaKvStore::F32 { k, v } => {
let k_view = k.slice(0..t_kv * nkv * hd);
let v_view = v.slice(0..t_kv * nkv * hd);
launch_sdpa_blocklist(
e,
&q,
&k_view,
&v_view,
&mut attended,
&pos_dev,
&meta_dev,
hd,
nh,
nkv,
t,
max_count,
scale,
)?;
}
QsaKvStore::Q8Q5 { k, v } => {
launch_q4e_sdpa_blocklist_q8q5(
e,
&q,
k,
v,
&mut attended,
&pos_dev,
&meta_dev,
hd,
nh,
nkv,
t,
max_count,
scale,
)?;
}
}
Ok(attended)
})?;
ws.put_i32("qsa.selpos", pos_dev);
ws.put_i32("qsa.selmeta", meta_dev);
attended
} else {
let QsaKvStore::F32 { k, v } = &*kv else {
return Err("qwen4exp_gpu: masked SDPA reached with a quantized cache".into());
};
let mask = rowsel_to_mask(&sels, block_size, t_kv);
let mask_dev = prof_section(e, "qsa.mask_h2d", || {
ws.take_u8_h2d(e, "qsa.mask", &mask, t * cap.min(SDPA_MASK_TKV_BOUND))
})?;
let attended = prof_section(e, "qsa.sdpa", || {
let mut attended = ws.take_f32(e, "qsa.att", t * nh * hd, 0)?;
let k_view = k.slice(0..t_kv * nkv * hd);
let v_view = v.slice(0..t_kv * nkv * hd);
launch_sdpa_mask(
e,
&q,
&k_view,
&v_view,
&mut attended,
&mask_dev,
hd,
nh,
nkv,
t,
t_kv,
scale,
)?;
Ok(attended)
})?;
ws.put_u8("qsa.mask", mask_dev);
attended
};
ws.put_f32(
if qsa.q_norm.is_some() {
"qsa.qn"
} else {
"qsa.q"
},
q,
);
let out = prof_section(e, "qsa.gate_wo", || {
let mut sg = ws.take_f32(e, "qsa.sg", t * nh * hd, 0)?;
e.sigmoid(&gate, &mut sg, t * nh * hd)?;
let mut gated = ws.take_f32(e, "qsa.gated", t * nh * hd, 0)?;
e.mul(&attended, &sg, &mut gated, t * nh * hd)?;
let mut out = ws.take_f32(e, "mixer.out", t * hidden, 0)?;
linear_trunk_into(
e,
&qsa.wo,
&qsa.wo_b16,
&gated,
&mut out,
t,
nh * hd,
hidden,
)?;
ws.put_f32("qsa.sg", sg);
ws.put_f32("qsa.gated", gated);
Ok(out)
})?;
ws.put_f32("qsa.att", attended);
ws.put_f32("qsa.gate", gate);
Ok(out)
}
#[allow(clippy::too_many_arguments)]
fn gdn_forward(
&self,
e: &Engine,
ws: &mut StepPool,
layer: &LayerW,
gdn: &GdnW,
mixed: &CudaSlice<f32>,
mstate: &mut MixerState,
t: usize,
mut stash: Option<&mut GdnStash>,
) -> Res<CudaSlice<f32>> {
let MixerState::Gdn { conv, state } = mstate else {
return Err(format!(
"qwen4exp_gpu: GDN layer {} bound to non-GDN state",
layer.index
)
.into());
};
let hidden = self.hidden;
let p = &gdn.plan;
let (nk, nv) = (p.key_heads as usize, p.value_heads as usize);
let (hk, hv) = (p.key_head_dim as usize, p.value_head_dim as usize);
let kernel = p.conv_kernel as usize;
let pad = kernel - 1;
let conv_dim = 2 * nk * hk + nv * hv;
let eps = layer.eps_attn;
let (qkv, z, beta_raw, g_log) = prof_section(e, "gdn.proj", || {
let mut qkv = ws.take_f32(e, "gdn.qkv", t * conv_dim, 0)?;
let mut z = ws.take_f32(e, "gdn.z", t * nv * hv, 0)?;
let mut beta_raw = ws.take_f32(e, "gdn.beta", t * nv, 0)?;
let mut alpha = ws.take_f32(e, "gdn.alpha", t * nv, 0)?;
if let (true, Some(stack)) = (
t == 1 && proj_stack_on() && trunk_bf16_on(),
gdn.proj_b16.as_ref(),
) {
launch_qmatvec_bf16w_multi4(
e,
stack,
mixed,
&[
(&qkv, conv_dim),
(&z, nv * hv),
(&beta_raw, nv),
(&alpha, nv),
],
hidden,
)?;
} else {
linear_trunk_stacked_into(
e,
&gdn.qkv,
&gdn.proj_b16,
0,
mixed,
&mut qkv,
t,
hidden,
conv_dim,
)?;
linear_trunk_stacked_into(
e,
&gdn.z,
&gdn.proj_b16,
conv_dim,
mixed,
&mut z,
t,
hidden,
nv * hv,
)?;
linear_trunk_stacked_into(
e,
&gdn.beta,
&gdn.proj_b16,
conv_dim + nv * hv,
mixed,
&mut beta_raw,
t,
hidden,
nv,
)?;
linear_trunk_stacked_into(
e,
&gdn.alpha,
&gdn.proj_b16,
conv_dim + nv * hv + nv,
mixed,
&mut alpha,
t,
hidden,
nv,
)?;
}
let mut g_log = ws.take_f32(e, "gdn.glog", t * nv, 0)?;
e.gdn_glog_v(&alpha.slice(0..t * nv), &gdn.dt, &gdn.a, &mut g_log, nv, t)?;
ws.put_f32("gdn.alpha", alpha);
Ok((qkv, z, beta_raw, g_log))
})?;
let o = prof_section(e, "gdn.conv_scan", || {
if let Some(st) = stash.as_deref_mut() {
e.copy_range_into(&mut st.conv_pre, 0, conv, 0, pad * conv_dim)?;
e.copy_range_into(&mut st.qkv_rows, 0, &qkv, 0, t * conv_dim)?;
}
let mut conv_out = ws.take_f32(e, "gdn.conv_out", t * conv_dim, 0)?;
let mut o = ws.take_f32(e, "gdn.o", t * nv * hv, 0)?;
let mut tmp = if t >= pad {
None
} else {
Some(ws.take_f32(e, "gdn.tmp", (pad - t) * conv_dim, 0)?)
};
let scale = 1.0 / (hk as f32).sqrt();
let step_ok = gdn_step_on() && hk % 32 == 0 && hk <= 1024;
let chain = |eng: &Engine,
conv: &mut CudaSlice<f32>,
state: &mut CudaSlice<f32>,
states_snap: Option<&mut CudaSlice<f32>>,
conv_out: &mut CudaSlice<f32>,
o: &mut CudaSlice<f32>,
tmp: Option<&mut CudaSlice<f32>>|
-> Res<()> {
launch_dwconv(
eng,
&qkv,
conv,
&gdn.conv_w,
conv_out,
t,
pad,
conv_dim,
kernel,
1,
1,
)?;
match states_snap {
Some(states) => {
let state_len = nv * hv * hk;
for tok in 0..t {
if step_ok {
launch_gdn_scan_step_at(
eng, conv_out, &g_log, &beta_raw, state, o, tok, nk, nv, hk,
hv, scale, eps,
)?;
} else {
launch_gdn_scan_at(
eng, conv_out, &g_log, &beta_raw, state, o, tok, nk, nv, hk,
hv, scale, eps,
)?;
}
eng.copy_range_into(states, tok * state_len, state, 0, state_len)?;
}
}
None if t == 1 && step_ok => {
launch_gdn_scan_step(
eng, conv_out, &g_log, &beta_raw, state, o, nk, nv, hk, hv, scale, eps,
)?;
}
None => {
launch_gdn_scan(
eng, conv_out, &g_log, &beta_raw, state, o, nk, nv, hk, hv, t, scale,
eps,
)?;
}
}
if t >= pad {
eng.copy_range_into(conv, 0, &qkv, (t - pad) * conv_dim, pad * conv_dim)?;
} else {
let keep = pad - t;
let tmp = tmp.ok_or("qwen4exp_gpu: gdn conv roll needs the tmp slot")?;
eng.copy_range_into(tmp, 0, conv, t * conv_dim, keep * conv_dim)?;
eng.copy_range_into(conv, 0, tmp, 0, keep * conv_dim)?;
eng.copy_range_into(conv, keep * conv_dim, &qkv, 0, t * conv_dim)?;
}
Ok(())
};
let graphable = stash.is_some() && verify_graphs_on() && step_ws_on() && !prof::on();
match stash.as_deref_mut() {
Some(st) if graphable => {
let entry = match st.scan_graph.take() {
Some((gt, g)) if gt == t => Some(g),
_ => None,
};
let warm = st.scan_warm == Some(t);
st.scan_warm = Some(t);
let entry = match (warm, entry) {
(false, _) => {
chain(
e,
conv,
state,
Some(&mut st.states),
&mut conv_out,
&mut o,
tmp.as_mut(),
)?;
None
}
(true, Some(g)) => {
g.0.launch()?;
Some(g)
}
(true, None) => {
let states = &mut st.states;
let mut tmp_ref = tmp.as_mut();
let g = e.capture_graph_retained_nowarm(|eng| {
chain(
eng,
conv,
state,
Some(states),
&mut conv_out,
&mut o,
tmp_ref.as_deref_mut(),
)
})?;
g.0.launch()?;
Some(g)
}
};
if let Some(g) = entry {
st.scan_graph = Some((t, g));
}
}
Some(st) => chain(
e,
conv,
state,
Some(&mut st.states),
&mut conv_out,
&mut o,
tmp.as_mut(),
)?,
None => chain(e, conv, state, None, &mut conv_out, &mut o, tmp.as_mut())?,
}
ws.put_f32("gdn.conv_out", conv_out);
if let Some(tmp) = tmp {
ws.put_f32("gdn.tmp", tmp);
}
Ok(o)
})?;
ws.put_f32("gdn.qkv", qkv);
ws.put_f32("gdn.beta", beta_raw);
ws.put_f32("gdn.glog", g_log);
let out = prof_section(e, "gdn.norm_gate_out", || {
let mut gated = ws.take_f32(e, "gdn.gated", t * nv * hv, 0)?;
match p.gate_activation {
GdnGateActivation::Sigmoid if gdn_fuse_on() => {
launch_rms_sigmul(e, &o, &gdn.norm, &z, &mut gated, hv, t * nv, eps)?;
}
GdnGateActivation::Sigmoid => {
let mut normed = ws.take_f32(e, "gdn.normed", t * nv * hv, 0)?;
e.rms_norm(&o, &gdn.norm, &mut normed, hv, t * nv, eps)?;
let mut sg = ws.take_f32(e, "gdn.sg", t * nv * hv, 0)?;
e.sigmoid(&z, &mut sg, t * nv * hv)?;
e.mul(&normed, &sg, &mut gated, t * nv * hv)?;
ws.put_f32("gdn.sg", sg);
ws.put_f32("gdn.normed", normed);
}
GdnGateActivation::Silu => {
let mut normed = ws.take_f32(e, "gdn.normed", t * nv * hv, 0)?;
e.rms_norm(&o, &gdn.norm, &mut normed, hv, t * nv, eps)?;
e.silu_mul(&z, &normed, &mut gated, t * nv * hv)?;
ws.put_f32("gdn.normed", normed);
}
}
let mut out = ws.take_f32(e, "mixer.out", t * hidden, 0)?;
linear_trunk_into(
e,
&gdn.out,
&gdn.out_b16,
&gated,
&mut out,
t,
nv * hv,
hidden,
)?;
ws.put_f32("gdn.gated", gated);
Ok(out)
})?;
ws.put_f32("gdn.z", z);
ws.put_f32("gdn.o", o);
Ok(out)
}
fn moe_forward(
&self,
e: &Engine,
ws: &mut StepPool,
moe: &MoeW,
mixed: &CudaSlice<f32>,
t: usize,
rows_grouped: bool,
layer: u32,
) -> Res<CudaSlice<f32>> {
let hidden = self.hidden;
let experts = moe.plan.expert_count as usize;
let selected = moe.plan.experts_per_token as usize;
let ff = moe.plan.expert_intermediate_size as usize;
let nvfp4_bank = matches!(
(&moe.bank.gate, &moe.bank.up, &moe.bank.down),
(
BankHalf::Nvfp4 { .. },
BankHalf::Nvfp4 { .. },
BankHalf::Nvfp4 { .. }
)
);
let devbf16_bank = matches!(
(&moe.bank.gate, &moe.bank.up, &moe.bank.down),
(
BankHalf::DeviceBf16(_),
BankHalf::DeviceBf16(_),
BankHalf::DeviceBf16(_)
)
);
let use_dev_router = router_dev_on()
&& moe_sel_path_on()
&& route_dev_geometry(experts, selected)
&& ((nvfp4_bank
&& hidden % 32 == 0
&& ff % 4 == 0
&& (t == 1
|| (rows_grouped
&& verify_mt_on()
&& sel_gufuse_on()
&& t * selected <= 8192)))
|| (devbf16_bank && hidden % 8 == 0 && ff % 8 == 0 && (t == 1 || rows_grouped)));
type DevRoute = (CudaSlice<i32>, CudaSlice<f32>, Option<CudaSlice<i32>>);
let (routes, mut dev_route): (Vec<Vec<(usize, f32)>>, Option<DevRoute>) =
prof_section(e, "moe.router", || {
let mut router_out = ws.take_f32(e, "moe.router", t * experts, 0)?;
let none: Option<CudaSlice<u8>> = None;
let rb = if router_bf16_on() {
&moe.router_b16
} else {
&none
};
linear_trunk_into(
e,
&moe.router,
rb,
mixed,
&mut router_out,
t,
hidden,
experts,
)?;
if use_dev_router {
let mut sel = ws.take_i32_slot(e, "moe.sel", t * selected, 0)?;
let mut w = ws.take_f32(e, "moe.w", t * selected, 0)?;
let mut tokm = if t > 1 {
Some(ws.take_i32_slot(e, "moe.tok", t * selected, 0)?)
} else {
None
};
route_topk_device(
e,
&router_out,
&mut sel,
&mut w,
tokm.as_mut().map(|m| (m, 0)),
experts,
selected,
t,
layer,
)?;
ws.put_f32("moe.router", router_out);
return Ok((Vec::new(), Some((sel, w, tokm))));
}
let logits = e.dtoh_view(&router_out.slice(0..t * experts))?;
ws.put_f32("moe.router", router_out);
let mut routes: Vec<Vec<(usize, f32)>> = Vec::with_capacity(t);
for token in 0..t {
routes.push(host_route_softmax_topk(
&logits[token * experts..(token + 1) * experts],
selected,
));
}
Ok((routes, None))
})?;
if (t == 1 || rows_grouped) && moe_sel_path_on() {
if let (
BankHalf::Nvfp4 {
codes: gc,
scales: gs,
macros_dev: gm,
..
},
BankHalf::Nvfp4 {
codes: uc,
scales: us,
macros_dev: um,
..
},
BankHalf::Nvfp4 {
codes: dc,
scales: ds,
macros_dev: dm,
..
},
) = (&moe.bank.gate, &moe.bank.up, &moe.bank.down)
{
if t > 1 && verify_mt_on() && sel_gufuse_on() && hidden % 32 == 0 && ff % 4 == 0 {
if let Some((sel, w_dev, tokm)) = dev_route.take() {
let tokm =
tokm.ok_or("moe_forward: device route at t > 1 without a tok map")?;
let out = prof_section(e, "moe.sel_grouped", || {
let mut out = ws.take_f32(e, "moe.out", t * hidden, 0)?;
let s_total = t * selected;
let mut act = ws.take_f32(e, "moe.act", s_total * ff, 0)?;
launch_nvfp4_sel_gu_silu(
e,
(gc, gs, gm),
(uc, us, um),
Some(&sel),
0,
s_total,
mixed,
&mut act,
hidden,
ff,
Some((&tokm, hidden)),
)?;
let mut partial = ws.take_f32(e, "moe.partial", s_total * hidden, 0)?;
launch_nvfp4_sel_matvec(
e,
dc,
ds,
dm,
&sel,
&act,
&mut partial,
s_total,
ff,
hidden,
ff,
)?;
for tok in 0..t {
launch_axpy_rows_seq_at(
e,
&partial,
tok * selected,
&w_dev,
tok * selected,
&mut out,
tok,
hidden,
selected,
)?;
}
ws.put_i32("moe.sel", sel);
ws.put_i32("moe.tok", tokm);
ws.put_f32("moe.w", w_dev);
ws.put_f32("moe.act", act);
ws.put_f32("moe.partial", partial);
Ok(out)
})?;
return self.moe_shared_tail(e, ws, moe, mixed, out, t);
}
let out = prof_section(e, "moe.sel_grouped", || {
let mut out = ws.take_f32(e, "moe.out", t * hidden, 0)?;
const SLOT_CAP: usize = 8192;
let tok_step = (SLOT_CAP / selected.max(1)).max(1);
let mut tok0 = 0usize;
while tok0 < t {
let tok_n = tok_step.min(t - tok0);
let batch = &routes[tok0..tok0 + tok_n];
let mut sel_all: Vec<i32> = Vec::with_capacity(tok_n * selected);
let mut w_all: Vec<f32> = Vec::with_capacity(tok_n * selected);
let mut tok_all: Vec<i32> = Vec::with_capacity(tok_n * selected);
let mut ranges: Vec<(usize, usize)> = Vec::with_capacity(tok_n);
for (i, route) in batch.iter().enumerate() {
ranges.push((sel_all.len(), route.len()));
for &(eid, wgt) in route {
sel_all.push(eid as i32);
w_all.push(wgt);
tok_all.push((tok0 + i) as i32);
}
}
let s_total = sel_all.len();
let sel = ws.take_i32(e, "moe.sel", &sel_all, 0)?;
let w_dev = ws.take_f32_h2d(e, "moe.w", &w_all, 0)?;
let tokm = ws.take_i32(e, "moe.tok", &tok_all, 0)?;
let mut act = ws.take_f32(e, "moe.act", s_total * ff, 0)?;
launch_nvfp4_sel_gu_silu(
e,
(gc, gs, gm),
(uc, us, um),
Some(&sel),
0,
s_total,
mixed,
&mut act,
hidden,
ff,
Some((&tokm, hidden)),
)?;
let mut partial = ws.take_f32(e, "moe.partial", s_total * hidden, 0)?;
launch_nvfp4_sel_matvec(
e,
dc,
ds,
dm,
&sel,
&act,
&mut partial,
s_total,
ff,
hidden,
ff,
)?;
for (i, &(start, len)) in ranges.iter().enumerate() {
launch_axpy_rows_seq_at(
e,
&partial,
start,
&w_dev,
start,
&mut out,
tok0 + i,
hidden,
len,
)?;
}
ws.put_i32("moe.sel", sel);
ws.put_i32("moe.tok", tokm);
ws.put_f32("moe.w", w_dev);
ws.put_f32("moe.act", act);
ws.put_f32("moe.partial", partial);
tok0 += tok_n;
}
Ok(out)
})?;
return self.moe_shared_tail(e, ws, moe, mixed, out, t);
}
if let Some((sel, w_dev, _)) = dev_route.take() {
let out = prof_section(e, "moe.sel_grouped", || {
let mut out = ws.take_f32(e, "moe.out", hidden, 0)?;
let mut act = ws.take_f32(e, "moe.act", selected * ff, 0)?;
if sel_gufuse_on() && hidden % 32 == 0 && ff % 4 == 0 {
launch_nvfp4_sel_gu_silu(
e,
(gc, gs, gm),
(uc, us, um),
Some(&sel),
0,
selected,
mixed,
&mut act,
hidden,
ff,
None,
)?;
} else {
let mut yg = ws.take_f32(e, "moe.yg", selected * ff, 0)?;
let mut yu = ws.take_f32(e, "moe.yu", selected * ff, 0)?;
launch_nvfp4_sel_matvec(
e, gc, gs, gm, &sel, mixed, &mut yg, selected, hidden, ff, 0,
)?;
launch_nvfp4_sel_matvec(
e, uc, us, um, &sel, mixed, &mut yu, selected, hidden, ff, 0,
)?;
e.silu_mul(&yg, &yu, &mut act, selected * ff)?;
ws.put_f32("moe.yg", yg);
ws.put_f32("moe.yu", yu);
}
let mut partial = ws.take_f32(e, "moe.partial", selected * hidden, 0)?;
launch_nvfp4_sel_matvec(
e,
dc,
ds,
dm,
&sel,
&act,
&mut partial,
selected,
ff,
hidden,
ff,
)?;
e.axpy_rows_seq_into(&partial, &w_dev, &mut out, hidden, selected)?;
ws.put_i32("moe.sel", sel);
ws.put_f32("moe.w", w_dev);
ws.put_f32("moe.act", act);
ws.put_f32("moe.partial", partial);
Ok(out)
})?;
return self.moe_shared_tail(e, ws, moe, mixed, out, t);
}
let out = prof_section(e, "moe.sel_grouped", || {
let mut out = ws.take_f32(e, "moe.out", t * hidden, 0)?;
for (tok, route) in routes.iter().enumerate() {
let n_sel = route.len();
let sel_host: Vec<i32> = route.iter().map(|&(x, _)| x as i32).collect();
let w_host: Vec<f32> = route.iter().map(|&(_, w)| w).collect();
let sel = ws.take_i32(e, "moe.sel", &sel_host, 0)?;
let w_dev = ws.take_f32_h2d(e, "moe.w", &w_host, 0)?;
let x_tok = if t == 1 {
None
} else {
let mut x = ws.take_f32(e, "moe.x", hidden, 0)?;
e.copy_range_into(&mut x, 0, mixed, tok * hidden, hidden)?;
Some(x)
};
let x_ref = x_tok.as_ref().unwrap_or(mixed);
let mut act = ws.take_f32(e, "moe.act", n_sel * ff, 0)?;
if sel_gufuse_on() && hidden % 32 == 0 && ff % 4 == 0 {
launch_nvfp4_sel_gu_silu(
e,
(gc, gs, gm),
(uc, us, um),
Some(&sel),
0,
n_sel,
x_ref,
&mut act,
hidden,
ff,
None,
)?;
} else {
let mut yg = ws.take_f32(e, "moe.yg", n_sel * ff, 0)?;
let mut yu = ws.take_f32(e, "moe.yu", n_sel * ff, 0)?;
launch_nvfp4_sel_matvec(
e, gc, gs, gm, &sel, x_ref, &mut yg, n_sel, hidden, ff, 0,
)?;
launch_nvfp4_sel_matvec(
e, uc, us, um, &sel, x_ref, &mut yu, n_sel, hidden, ff, 0,
)?;
e.silu_mul(&yg, &yu, &mut act, n_sel * ff)?;
ws.put_f32("moe.yg", yg);
ws.put_f32("moe.yu", yu);
}
let mut partial = ws.take_f32(e, "moe.partial", n_sel * hidden, 0)?;
launch_nvfp4_sel_matvec(
e,
dc,
ds,
dm,
&sel,
&act,
&mut partial,
n_sel,
ff,
hidden,
ff,
)?;
if t == 1 {
e.axpy_rows_seq_into(&partial, &w_dev, &mut out, hidden, n_sel)?;
} else {
let mut row = ws.take_f32(e, "moe.row", hidden, 0)?;
e.axpy_rows_seq_into(&partial, &w_dev, &mut row, hidden, n_sel)?;
e.copy_range_into(&mut out, tok * hidden, &row, 0, hidden)?;
ws.put_f32("moe.row", row);
}
ws.put_i32("moe.sel", sel);
ws.put_f32("moe.w", w_dev);
ws.put_f32("moe.act", act);
ws.put_f32("moe.partial", partial);
if let Some(x) = x_tok {
ws.put_f32("moe.x", x);
}
}
Ok(out)
})?;
return self.moe_shared_tail(e, ws, moe, mixed, out, t);
}
if let (BankHalf::DeviceBf16(gb), BankHalf::DeviceBf16(ub), BankHalf::DeviceBf16(db)) =
(&moe.bank.gate, &moe.bank.up, &moe.bank.down)
{
if let Some((sel, w_dev, _)) = dev_route.take() {
let out = prof_section(e, "moe.sel_bf16", || {
let mut out = ws.take_f32(e, "moe.out", t * hidden, 0)?;
for tok in 0..t {
let mut yg = ws.take_f32(e, "moe.yg", selected * ff, 0)?;
let mut yu = ws.take_f32(e, "moe.yu", selected * ff, 0)?;
launch_qmatvec_bf16w_sel(
e,
gb,
&sel,
tok * selected,
mixed,
tok * hidden,
0,
&mut yg,
selected,
hidden,
ff,
)?;
launch_qmatvec_bf16w_sel(
e,
ub,
&sel,
tok * selected,
mixed,
tok * hidden,
0,
&mut yu,
selected,
hidden,
ff,
)?;
let mut act = ws.take_f32(e, "moe.act", selected * ff, 0)?;
e.silu_mul(&yg, &yu, &mut act, selected * ff)?;
let mut partial =
ws.take_f32(e, "moe.partial", selected * hidden, 0)?;
launch_qmatvec_bf16w_sel(
e,
db,
&sel,
tok * selected,
&act,
0,
ff,
&mut partial,
selected,
ff,
hidden,
)?;
launch_axpy_rows_seq_at(
e,
&partial,
0,
&w_dev,
tok * selected,
&mut out,
tok,
hidden,
selected,
)?;
ws.put_f32("moe.yg", yg);
ws.put_f32("moe.yu", yu);
ws.put_f32("moe.act", act);
ws.put_f32("moe.partial", partial);
}
ws.put_i32("moe.sel", sel);
ws.put_f32("moe.w", w_dev);
Ok(out)
})?;
return self.moe_shared_tail(e, ws, moe, mixed, out, t);
}
let out = prof_section(e, "moe.sel_bf16", || {
let mut out = ws.take_f32(e, "moe.out", t * hidden, 0)?;
for (tok, route) in routes.iter().enumerate() {
let n_sel = route.len();
let w_host: Vec<f32> = route.iter().map(|&(_, w)| w).collect();
let w_dev = ws.take_f32_h2d(e, "moe.w", &w_host, 0)?;
let mut yg = ws.take_f32(e, "moe.yg", n_sel * ff, 0)?;
let mut yu = ws.take_f32(e, "moe.yu", n_sel * ff, 0)?;
for (slot, &(eid, _)) in route.iter().enumerate() {
launch_qmatvec_bf16w_off_into(
e,
gb,
eid * ff,
mixed,
tok * hidden,
&mut yg,
slot * ff,
hidden,
ff,
)?;
launch_qmatvec_bf16w_off_into(
e,
ub,
eid * ff,
mixed,
tok * hidden,
&mut yu,
slot * ff,
hidden,
ff,
)?;
}
let mut act = ws.take_f32(e, "moe.act", n_sel * ff, 0)?;
e.silu_mul(&yg, &yu, &mut act, n_sel * ff)?;
let mut partial = ws.take_f32(e, "moe.partial", n_sel * hidden, 0)?;
for (slot, &(eid, _)) in route.iter().enumerate() {
launch_qmatvec_bf16w_off_into(
e,
db,
eid * hidden,
&act,
slot * ff,
&mut partial,
slot * hidden,
ff,
hidden,
)?;
}
let mut row = ws.take_f32(e, "moe.row", hidden, 0)?;
e.axpy_rows_seq_into(&partial, &w_dev, &mut row, hidden, n_sel)?;
e.copy_range_into(&mut out, tok * hidden, &row, 0, hidden)?;
ws.put_f32("moe.row", row);
ws.put_f32("moe.w", w_dev);
ws.put_f32("moe.yg", yg);
ws.put_f32("moe.yu", yu);
ws.put_f32("moe.act", act);
ws.put_f32("moe.partial", partial);
}
Ok(out)
})?;
return self.moe_shared_tail(e, ws, moe, mixed, out, t);
}
}
if dev_route.is_some() {
return Err(
"moe_forward: device route left unconsumed (engage guard drifted from the \
dispatch arms)"
.into(),
);
}
let mut by_expert: Vec<Vec<(i32, i32, f32)>> = vec![Vec::new(); experts];
for (token, token_routes) in routes.iter().enumerate() {
for (slot, &(expert, weight)) in token_routes.iter().enumerate() {
by_expert[expert].push((token as i32, slot as i32, weight));
}
}
let mut slots = e.zeros(t * selected * hidden)?;
let mut wbuf = e.zeros(t * selected)?;
for (expert, entries) in by_expert.iter().enumerate() {
if entries.is_empty() {
continue;
}
let m_e = entries.len();
let (tok_dev, slot_dev, w_dev, xg) = prof_section(e, "moe.idx_gather", || {
let tok_idx: Vec<i32> = entries.iter().map(|&(tok, _, _)| tok).collect();
let slot_idx: Vec<i32> = entries.iter().map(|&(_, slot, _)| slot).collect();
let weights: Vec<f32> = entries.iter().map(|&(_, _, w)| w).collect();
let tok_dev = e.htod_i32(&tok_idx)?;
let slot_dev = e.htod_i32(&slot_idx)?;
let w_dev = e.htod(&weights)?;
let mut xg = e.uninit(m_e * hidden)?;
e.gather_rows(mixed, &tok_dev, &mut xg, hidden, m_e)?;
Ok((tok_dev, slot_dev, w_dev, xg))
})?;
let resolve = |half: &BankHalf,
out_f: usize,
in_f: usize|
-> Res<(Option<CudaSlice<f32>>, usize)> {
match half {
BankHalf::F32(_) => Ok((None, expert * out_f * in_f)),
BankHalf::Nvfp4 {
codes,
scales,
macros,
..
} => Ok((
Some(dequant_nvfp4_expert_f32(
e,
codes,
scales,
macros[expert],
expert,
out_f,
in_f,
)?),
0,
)),
BankHalf::HostBf16(bytes) => {
let row_bytes = out_f * in_f * 2;
let dev =
e.htod_bytes(&bytes[expert * row_bytes..(expert + 1) * row_bytes])?;
Ok((
Some(e.bf16_to_f32(&dev.slice(0..row_bytes), out_f * in_f)?),
0,
))
}
BankHalf::DeviceBf16(bytes) => {
let row_bytes = out_f * in_f * 2;
let view = bytes.slice(expert * row_bytes..(expert + 1) * row_bytes);
Ok((Some(e.bf16_to_f32(&view, out_f * in_f)?), 0))
}
}
};
let ((gate_owned, gate_base), (up_owned, up_base), (down_owned, down_base)) =
prof_section(e, "moe.dequant", || {
Ok((
resolve(&moe.bank.gate, ff, hidden)?,
resolve(&moe.bank.up, ff, hidden)?,
resolve(&moe.bank.down, hidden, ff)?,
))
})?;
let gate_view = match (&moe.bank.gate, &gate_owned) {
(_, Some(owned)) => owned.slice(0..ff * hidden),
(BankHalf::F32(bank), None) => bank.slice(gate_base..gate_base + ff * hidden),
(
BankHalf::Nvfp4 { .. } | BankHalf::HostBf16(_) | BankHalf::DeviceBf16(_),
None,
) => {
unreachable!("quantized/host/device-bf16 halves always resolve owned")
}
};
let up_view = match (&moe.bank.up, &up_owned) {
(_, Some(owned)) => owned.slice(0..ff * hidden),
(BankHalf::F32(bank), None) => bank.slice(up_base..up_base + ff * hidden),
(
BankHalf::Nvfp4 { .. } | BankHalf::HostBf16(_) | BankHalf::DeviceBf16(_),
None,
) => {
unreachable!("quantized/host/device-bf16 halves always resolve owned")
}
};
let down_view = match (&moe.bank.down, &down_owned) {
(_, Some(owned)) => owned.slice(0..hidden * ff),
(BankHalf::F32(bank), None) => bank.slice(down_base..down_base + hidden * ff),
(
BankHalf::Nvfp4 { .. } | BankHalf::HostBf16(_) | BankHalf::DeviceBf16(_),
None,
) => {
unreachable!("quantized/host/device-bf16 halves always resolve owned")
}
};
prof_section(e, "moe.expert_gemms", || {
let down_out =
run_routed_expert(e, &xg, &gate_view, &up_view, &down_view, m_e, hidden, ff)?;
e.scatter_slot(
&down_out, &tok_dev, &slot_dev, &w_dev, &mut slots, &mut wbuf, hidden,
selected, m_e,
)
})?;
}
let out = prof_section(e, "moe.reduce", || {
let mut out = e.zeros(t * hidden)?;
e.reduce_slots(&slots, &wbuf, &mut out, hidden, selected, t)?;
Ok(out)
})?;
self.moe_shared_tail(e, ws, moe, mixed, out, t)
}
fn moe_shared_tail(
&self,
e: &Engine,
ws: &mut StepPool,
moe: &MoeW,
mixed: &CudaSlice<f32>,
mut out: CudaSlice<f32>,
t: usize,
) -> Res<CudaSlice<f32>> {
let hidden = self.hidden;
let sff = moe
.plan
.shared
.as_ref()
.map(|s| s.intermediate_size as usize)
.unwrap_or(0);
if sff > 0 {
prof_section(e, "moe.shared", || {
let none: Option<CudaSlice<u8>> = None;
let (gu, db) = if micro_shexp_on() {
(&moe.shared_gu_b16, &moe.shared_down_b16)
} else {
(&none, &none)
};
let mut gate = ws.take_f32(e, "moe.sh_gate", t * sff, 0)?;
let mut up = ws.take_f32(e, "moe.sh_up", t * sff, 0)?;
if let (true, Some(stack)) =
(t == 1 && proj_stack_on() && trunk_bf16_on(), gu.as_ref())
{
launch_qmatvec_bf16w_multi4(
e,
stack,
mixed,
&[(&gate, sff), (&up, sff)],
hidden,
)?;
} else {
linear_trunk_stacked_into(
e,
&moe.shared_gate,
gu,
0,
mixed,
&mut gate,
t,
hidden,
sff,
)?;
linear_trunk_stacked_into(
e,
&moe.shared_up,
gu,
sff,
mixed,
&mut up,
t,
hidden,
sff,
)?;
}
let mut act = ws.take_f32(e, "moe.sh_act", t * sff, 0)?;
e.silu_mul(&gate, &up, &mut act, t * sff)?;
let mut shared = ws.take_f32(e, "moe.sh_down", t * hidden, 0)?;
linear_trunk_into(e, &moe.shared_down, db, &act, &mut shared, t, sff, hidden)?;
if let Some(input_gate) = moe.shared_input_gate.as_ref() {
let mut g = ws.take_f32(e, "moe.g", t, 0)?;
e.sigmoid_dot_rows_into(mixed, input_gate, &mut g, hidden, t)?;
e.add_scaled_rows(&shared, &g, &mut out, hidden, t)?;
ws.put_f32("moe.g", g);
} else {
let mut view = out.slice_mut(0..t * hidden);
e.axpy_into(&shared, 1.0, &mut view, t * hidden)?;
}
ws.put_f32("moe.sh_gate", gate);
ws.put_f32("moe.sh_up", up);
ws.put_f32("moe.sh_act", act);
ws.put_f32("moe.sh_down", shared);
Ok(())
})?;
}
Ok(out)
}
#[allow(clippy::too_many_arguments)]
fn ple_block(
&self,
e: &Engine,
layer: &LayerW,
ple: &PleW,
table: &NgramTable,
ple_state: &mut PleState,
planes: &mut [CudaSlice<f32>],
tokens: &[u32],
t: usize,
exact: bool,
mut stash: Option<&mut PleStash>,
) -> Res<()> {
let hidden = self.hidden;
let streams = self.streams;
let plan = &ple.plan;
let heads = plan.ngram_heads as usize;
let head_dim = plan.head_embed_dim as usize;
let embed_dim = plan.embed_dim as usize;
let kernel = plan.conv_kernel as usize;
let max_ngram = plan.max_ngram as usize;
let dilation = max_ngram;
let pad = (kernel - 1) * dilation;
let eps = layer.eps_attn;
let gathered = prof_section(e, "ple.host_ngram_gather", || {
let total_heads = heads;
let mut ids: Vec<i64> = Vec::new();
if ple_cache_on() {
host_ngram_ids_cached(
&mut ple_state.ngram_ids,
&mut ple_state.ngram_history,
&mut ple_state.ngram_last_eos,
tokens,
&ple.multipliers,
&ple.sizes,
&ple.offsets,
max_ngram,
heads / (max_ngram - 1),
plan.eos_token_id,
);
if ple_cache_audit_on() {
let twin = host_ngram_ids(
tokens,
&ple.multipliers,
&ple.sizes,
&ple.offsets,
max_ngram,
heads / (max_ngram - 1),
plan.eos_token_id,
);
let from = (tokens.len() - t) * total_heads;
let mism = twin[from..]
.iter()
.zip(&ple_state.ngram_ids[from..])
.filter(|(a, b)| a != b)
.count() as u64;
PLE_CACHE_AUDIT_ROWS.fetch_add(t as u64, std::sync::atomic::Ordering::Relaxed);
PLE_CACHE_AUDIT_MISMATCH.fetch_add(mism, std::sync::atomic::Ordering::Relaxed);
PLE_CACHE_AUDIT_MAX_FILL
.fetch_max(tokens.len() as u64, std::sync::atomic::Ordering::Relaxed);
if mism > 0 {
return Err(format!(
"plecache audit: {mism} cached n-gram ids differ from the full twin \
at history {} (t={t})",
tokens.len()
)
.into());
}
}
} else {
ids = host_ngram_ids(
tokens,
&ple.multipliers,
&ple.sizes,
&ple.offsets,
max_ngram,
heads / (max_ngram - 1),
plan.eos_token_id,
);
}
let all_ids: &[i64] = if ple_cache_on() {
&ple_state.ngram_ids
} else {
&ids
};
let chunk_ids = &all_ids[(tokens.len() - t) * total_heads..];
let table_rows = table.rows(head_dim);
let mut gathered = vec![0.0f32; t * embed_dim];
for token in 0..t {
for head in 0..heads {
let id = chunk_ids[token * total_heads + head];
if id < 0 || id as usize >= table_rows {
return Err("qwen4exp_gpu: n-gram id outside the embedding table".into());
}
table.gather_into(
id as usize,
head_dim,
&mut gathered[token * embed_dim + head * head_dim
..token * embed_dim + (head + 1) * head_dim],
);
}
}
Ok(gathered)
})?;
let emb = prof_section(e, "ple.h2d", || e.htod(&gathered))?;
let lin_rows = |x: &CudaSlice<f32>,
w: &CudaSlice<f32>,
in_f: usize,
out_f: usize|
-> Res<CudaSlice<f32>> {
let mut out = e.uninit(t * out_f)?;
if exact && t > 1 {
let wv = w.slice(0..w.len());
for tok in 0..t {
let xv = x.slice(tok * in_f..(tok + 1) * in_f);
let mut yv = out.slice_mut(tok * out_f..(tok + 1) * out_f);
e.linear_device_into(&xv, &wv, &mut yv, 1, in_f, out_f)?;
}
} else {
e.linear_device_into(x, w, &mut out, t, in_f, out_f)?;
}
Ok(out)
};
let (value, mut dots_host) = prof_section(e, "ple.key_gate", || {
let value = lin_rows(&emb, &ple.value_proj, embed_dim, hidden)?;
let ones = e.htod(&vec![1.0f32; hidden])?;
let mut dots_host = vec![0.0f32; streams * t];
for s in 0..streams {
let key = lin_rows(&emb, &ple.key_proj[s], embed_dim, hidden)?;
let mut key_normed = e.uninit(t * hidden)?;
e.rms_norm(&key, &ple.norm_key[s], &mut key_normed, hidden, t, eps)?;
let mut query = e.uninit(t * hidden)?;
e.rms_norm(&planes[s], &ple.norm_query[s], &mut query, hidden, t, eps)?;
let mut prod = e.uninit(t * hidden)?;
e.mul(&key_normed, &query, &mut prod, t * hidden)?;
let dots = lin_rows(&prod, &ones, hidden, 1)?;
dots_host[s * t..(s + 1) * t].copy_from_slice(&e.dtoh(&dots)?);
}
Ok((value, dots_host))
})?;
for dot in dots_host.iter_mut() {
let gate = *dot / (hidden as f32).sqrt();
let magnitude = gate.abs().max(1e-6).sqrt();
let signed = if gate > 0.0 {
magnitude
} else if gate < 0.0 {
-magnitude
} else {
0.0
};
*dot = host_sigmoid(signed);
}
prof_section(e, "ple.conv_write", || {
for s in 0..streams {
let g = e.htod(&dots_host[s * t..(s + 1) * t])?;
let mut gated = e.zeros(t * hidden)?;
e.add_scaled_rows(&value, &g, &mut gated, hidden, t)?;
let mut normed = e.uninit(t * hidden)?;
e.rms_norm(&gated, &ple.norm_conv[s], &mut normed, hidden, t, eps)?;
if let Some(st) = stash.as_deref_mut() {
e.copy_range_into(
&mut st.hist_pre[s],
0,
&ple_state.conv_hist[s],
0,
pad * hidden,
)?;
e.copy_range_into(&mut st.normed_rows[s], 0, &normed, 0, t * hidden)?;
}
launch_dwconv(
e,
&normed,
&ple_state.conv_hist[s],
&ple.conv_w[s],
&mut gated,
t,
pad,
hidden,
kernel,
dilation,
2,
)?;
let hist = &mut ple_state.conv_hist[s];
if t >= pad {
e.copy_range_into(hist, 0, &normed, (t - pad) * hidden, pad * hidden)?;
} else {
let keep = pad - t;
let mut tmp = e.uninit(keep * hidden)?;
e.copy_range_into(&mut tmp, 0, hist, t * hidden, keep * hidden)?;
e.copy_range_into(hist, 0, &tmp, 0, keep * hidden)?;
e.copy_range_into(hist, keep * hidden, &normed, 0, t * hidden)?;
}
let mut view = planes[s].slice_mut(0..t * hidden);
e.axpy_into(&gated, 1.0, &mut view, t * hidden)?;
}
Ok(())
})
}
}
pub struct MtpDraftState {
mixer: MixerState,
rows: usize,
pub committed: usize,
capacity: usize,
ws: StepPool,
}
impl MtpDraftState {
pub fn rows(&self) -> usize {
self.rows
}
}
#[derive(Clone, Copy)]
enum DraftTokSrc<'a> {
Host(&'a [u32]),
HostDev(&'a [u32]),
DevSlot(&'a CudaSlice<u32>, usize),
}
impl Qwen4ExpGpu {
pub fn has_mtp(&self) -> bool {
self.mtp.is_some()
}
pub fn mtp_on_dev1(&self) -> bool {
self.mtp_dev1.is_some()
}
fn check_draft_engine(&self, e: &Engine) -> Res<()> {
if let Some(d) = self.mtp_dev1.as_ref() {
if e.ctx().ordinal() != d.dev {
return Err(format!(
"qwen4exp_gpu: the draft lives on device {} (card-1 placement); \
this call presented device {}",
d.dev,
e.ctx().ordinal()
)
.into());
}
}
Ok(())
}
pub fn mtp_state(&self, e: &Engine, capacity: usize) -> Res<MtpDraftState> {
self.check_draft_engine(e)?;
let mtp = self
.mtp
.as_ref()
.ok_or("qwen4exp_gpu: no MTP block loaded (LoadOptions::load_mtp)")?;
let MixerW::Qsa(qsa) = &mtp.layer.mixer else {
return Err("qwen4exp_gpu: MTP mixer is not QSA".into());
};
let kv_width = qsa.attn.kv_heads as usize * qsa.attn.key_head_dim as usize;
let v_width = qsa.attn.kv_heads as usize * qsa.attn.value_head_dim as usize;
let kv = if kv_quant_on() {
QsaKvStore::Q8Q5 {
k: e.alloc_u8(capacity * q8_row_bytes(kv_width))?,
v: e.alloc_u8(capacity * q5_row_bytes(v_width))?,
}
} else {
QsaKvStore::F32 {
k: e.zeros(capacity * kv_width)?,
v: e.zeros(capacity * v_width)?,
}
};
Ok(MtpDraftState {
mixer: MixerState::Qsa {
kv,
raw_keys: IdxRawCache::new(idxq_mode()),
pooled_keys: Vec::new(),
pooled_dev: None,
pooled_dev_rows: 0,
raw_dev: None,
raw_dev_rows: 0,
idx_audit: None,
},
rows: 0,
committed: 0,
capacity,
ws: StepPool::default(),
})
}
pub fn mtp_rewind(&self, dstate: &mut MtpDraftState, rows: usize) -> Res<()> {
if rows > dstate.rows {
return Err("qwen4exp_gpu: mtp_rewind past the cache".into());
}
let mtp = self.mtp.as_ref().ok_or("qwen4exp_gpu: no MTP block")?;
let MixerW::Qsa(qsa) = &mtp.layer.mixer else {
return Err("qwen4exp_gpu: MTP mixer is not QSA".into());
};
let MixerState::Qsa {
raw_keys,
pooled_keys,
pooled_dev_rows,
raw_dev_rows,
..
} = &mut dstate.mixer
else {
return Err("qwen4exp_gpu: MTP state is not QSA".into());
};
let idx_dim = qsa.overlay.head_dim as usize;
raw_keys.truncate_rows(rows, idx_dim);
let block = qsa.overlay.block_size as usize;
pooled_keys.truncate((rows / block) * idx_dim);
*pooled_dev_rows = (*pooled_dev_rows).min(pooled_keys.len() / idx_dim);
*raw_dev_rows = (*raw_dev_rows).min(rows);
dstate.rows = rows;
dstate.committed = dstate.committed.min(rows);
Ok(())
}
#[allow(clippy::too_many_arguments)]
pub fn mtp_draft_forward(
&self,
e: &Engine,
tokens: &[u32],
hidden_wide: &CudaSlice<f32>,
wide_off: usize,
dstate: &mut MtpDraftState,
pos_off: usize,
logits_all: bool,
) -> Res<(CudaSlice<f32>, CudaSlice<f32>)> {
self.mtp_draft_forward_impl(
e,
DraftTokSrc::Host(tokens),
hidden_wide,
wide_off,
dstate,
pos_off,
logits_all,
)
}
fn mtp_draft_forward_devslot(
&self,
e: &Engine,
toks: &CudaSlice<u32>,
slot: usize,
hidden_wide: &CudaSlice<f32>,
wide_off: usize,
dstate: &mut MtpDraftState,
) -> Res<(CudaSlice<f32>, CudaSlice<f32>)> {
self.mtp_draft_forward_impl(
e,
DraftTokSrc::DevSlot(toks, slot),
hidden_wide,
wide_off,
dstate,
1,
false,
)
}
fn mtp_draft_forward_spec(
&self,
e: &Engine,
tokens: &[u32],
dev_embed: bool,
hidden_wide: &CudaSlice<f32>,
wide_off: usize,
dstate: &mut MtpDraftState,
) -> Res<(CudaSlice<f32>, CudaSlice<f32>)> {
let src = if dev_embed {
DraftTokSrc::HostDev(tokens)
} else {
DraftTokSrc::Host(tokens)
};
self.mtp_draft_forward_impl(e, src, hidden_wide, wide_off, dstate, 1, false)
}
fn draft_consume_ring(
&self,
de: &Engine,
tokens: &[u32],
dev_embed: bool,
seed: &CudaSlice<f32>,
ring: usize,
first_row: usize,
dstate: &mut MtpDraftState,
) -> Res<(CudaSlice<f32>, CudaSlice<f32>, usize)> {
let mut out: Option<(CudaSlice<f32>, CudaSlice<f32>, usize)> = None;
let mut done = 0usize;
while done < tokens.len() {
let slot = (first_row + done) % ring;
let len = (tokens.len() - done).min(ring - slot);
let (l, c) = self.mtp_draft_forward_spec(
de,
&tokens[done..done + len],
dev_embed,
seed,
slot,
dstate,
)?;
if let Some((pl, pc, _)) = out.take() {
self.mtp_recycle(dstate, pl, pc);
}
out = Some((l, c, len));
done += len;
}
out.ok_or("qwen4exp_gpu: empty draft consume".into())
}
#[allow(clippy::too_many_arguments)]
fn mtp_draft_forward_impl(
&self,
e: &Engine,
tok_src: DraftTokSrc<'_>,
hidden_wide: &CudaSlice<f32>,
wide_off: usize,
dstate: &mut MtpDraftState,
pos_off: usize,
logits_all: bool,
) -> Res<(CudaSlice<f32>, CudaSlice<f32>)> {
self.check_draft_engine(e)?;
let mtp = self
.mtp
.as_ref()
.ok_or("qwen4exp_gpu: no MTP block loaded (LoadOptions::load_mtp)")?;
let t = match tok_src {
DraftTokSrc::Host(tokens) | DraftTokSrc::HostDev(tokens) => tokens.len(),
DraftTokSrc::DevSlot(..) => 1,
};
let hidden = self.hidden;
let streams = self.streams;
let wide = streams * hidden;
if t == 0 {
return Err("qwen4exp_gpu: empty draft input".into());
}
if dstate.rows + t > dstate.capacity {
return Err("qwen4exp_gpu: draft state capacity exceeded".into());
}
if hidden_wide.len() < (wide_off + t) * wide {
return Err("qwen4exp_gpu: draft hidden seed rows out of range".into());
}
let base = dstate.rows;
let ws = &mut dstate.ws;
let cap = dstate.capacity;
let mut planes = prof_section(e, "mtp.fuse", || {
let emb = match tok_src {
DraftTokSrc::Host(tokens) => {
let mut embedded = vec![0.0f32; t * hidden];
for (row, &token) in tokens.iter().enumerate() {
let token = token as usize;
if token >= self.vocab {
return Err(
format!("qwen4exp_gpu: draft token {token} out of range").into()
);
}
embedded[row * hidden..(row + 1) * hidden].copy_from_slice(
&self.embed_host[token * hidden..(token + 1) * hidden],
);
}
ws.take_f32_h2d(e, "mtp.emb", &embedded, cap * hidden)?
}
DraftTokSrc::HostDev(tokens) => {
let ce = self
.chain_embed
.as_ref()
.filter(|ce| !ce.for_trim && ce.rows == self.vocab)
.ok_or("qwen4exp_gpu: HostDev embed needs the full-vocab chain table")?;
for &token in tokens {
if token as usize >= self.vocab {
return Err(
format!("qwen4exp_gpu: draft token {token} out of range").into()
);
}
}
let tok_d = e.gpu.stream().clone_htod(tokens)?;
let mut emb = ws.take_f32(e, "mtp.emb", t * hidden, cap * hidden)?;
let tv = tok_d.slice(0..t);
embed_gather_rows_into(
e,
&ce.table,
&tv,
&mut emb,
t,
hidden,
ce.qt,
ce.row_bytes,
)?;
emb
}
DraftTokSrc::DevSlot(toks, slot) => {
let ce = self
.chain_embed
.as_ref()
.ok_or("qwen4exp_gpu: deferred draft step without arm_spec_devchain")?;
let mut emb = ws.take_f32(e, "mtp.emb", hidden, cap * hidden)?;
let tv = toks.slice(slot..slot + 1);
embed_gather_rows_into(
e,
&ce.table,
&tv,
&mut emb,
1,
hidden,
ce.qt,
ce.row_bytes,
)?;
emb
}
};
let mut enorm = ws.take_f32(e, "mtp.enorm", t * hidden, 0)?;
e.rms_norm(
&emb,
&mtp.pre_norm_embed,
&mut enorm,
hidden,
t,
mtp.eps_embed,
)?;
let mut evec = ws.take_f32(e, "mtp.evec", t * hidden, 0)?;
linear_trunk_into(
e,
&mtp.fc_embed,
&mtp.fc_embed_b16,
&enorm,
&mut evec,
t,
hidden,
hidden,
)?;
let mut hin = ws.take_f32(e, "mtp.hin", t * wide, 0)?;
e.copy_range_into(&mut hin, 0, hidden_wide, wide_off * wide, t * wide)?;
let mut hnorm = ws.take_f32(e, "mtp.hnorm", t * wide, 0)?;
e.rms_norm(
&hin,
&mtp.pre_norm_hidden,
&mut hnorm,
wide,
t,
mtp.eps_hidden,
)?;
let mut fused = ws.take_f32(e, "mtp.fused", t * wide, 0)?;
linear_trunk_into(
e,
&mtp.fc_hidden,
&mtp.fc_hidden_b16,
&hnorm,
&mut fused,
t * streams,
hidden,
hidden,
)?;
let mut planes: Vec<CudaSlice<f32>> = Vec::with_capacity(streams);
for s in 0..streams {
let mut plane = ws.take_f32(e, PLANE_SLOTS[s], t * hidden, cap * hidden)?;
for tok in 0..t {
e.copy_range_into(
&mut plane,
tok * hidden,
&fused,
(tok * streams + s) * hidden,
hidden,
)?;
}
let mut view = plane.slice_mut(0..t * hidden);
e.axpy_into(&evec, 1.0, &mut view, t * hidden)?;
planes.push(plane);
}
ws.put_f32("mtp.emb", emb);
ws.put_f32("mtp.enorm", enorm);
ws.put_f32("mtp.evec", evec);
ws.put_f32("mtp.hin", hin);
ws.put_f32("mtp.hnorm", hnorm);
ws.put_f32("mtp.fused", fused);
Ok(planes)
})?;
let ptr_vals: Vec<u64> = {
let stream = e.gpu.stream();
planes.iter().map(|p| p.device_ptr(&stream).0).collect()
};
let ptrs = ws.take_u64_h2d(e, "hc.ptrs", &ptr_vals, 0)?;
let layer = &mtp.layer;
let (mixed, inject) = prof_section(e, "mtp.hyper.read", || {
self.gate_read(
e,
ws,
&ptrs,
&layer.attn_gate,
&planes,
t,
layer.eps_attn,
false,
)
})?;
let MixerW::Qsa(qsa) = &layer.mixer else {
return Err("qwen4exp_gpu: MTP mixer is not QSA".into());
};
let block_out = prof_section(e, "mtp.qsa", || {
self.qsa_forward(
e,
ws,
layer,
qsa,
&mixed,
&mut dstate.mixer,
base,
t,
pos_off,
false,
)
})?;
ws.put_f32("hc.mixed", mixed);
prof_section(e, "mtp.hyper.write", || {
self.gate_write(e, &mut planes, &ptrs, &block_out, &inject, t)
})?;
ws.put_f32("mixer.out", block_out);
put_inject(ws, inject);
let (mixed, inject) = prof_section(e, "mtp.hyper.read", || {
self.gate_read(
e,
ws,
&ptrs,
&layer.mlp_gate,
&planes,
t,
layer.eps_mlp,
false,
)
})?;
let mlp = prof_section(e, "mtp.moe", || {
self.moe_forward(e, ws, &layer.moe, &mixed, t, t <= 32, layer.index)
})?;
ws.put_f32("hc.mixed", mixed);
prof_section(e, "mtp.hyper.write", || {
self.gate_write(e, &mut planes, &ptrs, &mlp, &inject, t)
})?;
ws.put_f32("moe.out", mlp);
put_inject(ws, inject);
let mut carrier = ws.take_f32(e, "mtp.carrier", t * wide, 0)?;
for (s, plane) in planes.iter().enumerate() {
for tok in 0..t {
e.copy_range_into(
&mut carrier,
tok * wide + s * hidden,
plane,
tok * hidden,
hidden,
)?;
}
}
let x = prof_section(e, "mtp.exit", || {
Ok(self
.gate_read_inner(
e,
ws,
&ptrs,
&mtp.mixer,
&planes,
t,
self.exit_eps,
false,
false,
)?
.0)
})?;
ws.put_u64("hc.ptrs", ptrs);
let trim = self.draft_trim.as_ref();
let out_f = trim.map_or(self.vocab, |t| t.n);
let (head_w, head_b16) = match self.mtp_dev1.as_ref() {
Some(d) => (&d.output, &d.output_b16),
None => (&self.output, &self.output_b16),
};
let head_into =
|e: &Engine, x: &CudaSlice<f32>, y: &mut CudaSlice<f32>, rows: usize| -> Res<()> {
match trim {
Some(trim) => linear_trim_into(e, trim, x, y, rows, hidden),
None => linear_trunk_into(e, head_w, head_b16, x, y, rows, hidden, self.vocab),
}
};
let logits = prof_section(e, "mtp.lm_head", || {
if logits_all {
let mut logits = ws.take_f32(e, "mtp.logits", t * out_f, 0)?;
head_into(e, &x, &mut logits, t)?;
return Ok(logits);
}
let mut logits = ws.take_f32(e, "mtp.logits", out_f, 0)?;
let mut x_last = ws.take_f32(e, "mtp.xlast", hidden, 0)?;
e.copy_range_into(&mut x_last, 0, &x, (t - 1) * hidden, hidden)?;
head_into(e, &x_last, &mut logits, 1)?;
ws.put_f32("mtp.xlast", x_last);
Ok(logits)
})?;
ws.put_f32("hc.mixed", x);
for (s, plane) in planes.into_iter().enumerate() {
ws.put_f32(PLANE_SLOTS[s], plane);
}
dstate.rows += t;
Ok((logits, carrier))
}
pub fn mtp_recycle(
&self,
dstate: &mut MtpDraftState,
logits: CudaSlice<f32>,
carrier: CudaSlice<f32>,
) {
dstate.ws.put_f32("mtp.logits", logits);
dstate.ws.put_f32("mtp.carrier", carrier);
}
}
#[derive(Clone, Copy)]
pub struct SpecSamplerCfg {
pub temperature: f32,
pub top_p: f32,
pub top_k: usize,
pub seed: u64,
}
struct SpecRng(u64);
impl SpecRng {
fn next_f32(&mut self) -> f32 {
let mut x = self.0;
x ^= x >> 12;
x ^= x << 25;
x ^= x >> 27;
self.0 = x;
let bits = (x.wrapping_mul(0x2545_F491_4F6C_DD1D) >> 40) as u32;
bits as f32 / (1u64 << 24) as f32
}
}
fn sample_row(cfg: &SpecSamplerCfg, rng: &mut SpecRng, row: &[f32]) -> u32 {
let k = cfg.top_k.max(1).min(row.len());
let mut idx: Vec<u32> = (0..row.len() as u32).collect();
idx.select_nth_unstable_by(k - 1, |&a, &b| row[b as usize].total_cmp(&row[a as usize]));
let mut top: Vec<(u32, f32)> = idx[..k].iter().map(|&i| (i, row[i as usize])).collect();
top.sort_by(|a, b| b.1.total_cmp(&a.1));
let temp = cfg.temperature.max(1e-6);
let mx = top[0].1;
let mut probs: Vec<f32> = top.iter().map(|&(_, v)| ((v - mx) / temp).exp()).collect();
let sum: f32 = probs.iter().sum();
for p in &mut probs {
*p /= sum;
}
let mut cut = probs.len();
let mut acc = 0.0f32;
for (i, &p) in probs.iter().enumerate() {
acc += p;
if acc >= cfg.top_p {
cut = i + 1;
break;
}
}
let renorm: f32 = probs[..cut].iter().sum();
let draw = rng.next_f32() * renorm;
let mut acc = 0.0f32;
for (i, &p) in probs[..cut].iter().enumerate() {
acc += p;
if draw < acc {
return top[i].0;
}
}
top[cut - 1].0
}
fn host_argmax(row: &[f32]) -> usize {
let mut best = 0usize;
for (i, &v) in row.iter().enumerate() {
if v > row[best] {
best = i;
}
}
best
}
fn ring_pieces(ring: usize, off: usize, t: usize) -> Vec<(usize, usize)> {
debug_assert!(t <= ring, "wide-ring consumer wider than the ring");
let slot = off % ring;
if slot + t <= ring {
vec![(slot, t)]
} else {
vec![(slot, ring - slot), (0, t - (ring - slot))]
}
}
fn cross_wide_rows(
e: &Engine,
de: &Engine,
src: &CudaSlice<f32>,
dst: &mut CudaSlice<f32>,
off: usize,
t: usize,
wide: usize,
) -> Res<f64> {
let t0 = std::time::Instant::now();
let stream = de.gpu.stream();
let bytes = t * wide * 4;
let byte_off = (off * wide * 4) as u64;
let (sp, _g0) = src.device_ptr(&stream);
let (dp, _g1) = dst.device_ptr_mut(&stream);
unsafe {
cudarc::driver::result::memcpy_peer_async(
de.ctx().cu_ctx(),
dp + byte_off,
e.ctx().cu_ctx(),
sp + byte_off,
bytes,
stream.cu_stream(),
)?;
}
stream.synchronize()?;
Ok(t0.elapsed().as_secs_f64() * 1e3)
}
fn embed_gather_rows_into(
e: &Engine,
table: &CudaSlice<u8>,
tok_v: &CudaView<u32>,
x_out: &mut CudaSlice<f32>,
t: usize,
n_embd: usize,
qtype: i32,
row_bytes: usize,
) -> Res<()> {
let f = e.func("embed_gather_u32_t");
let cfg = LaunchConfig {
grid_dim: (((n_embd as u32).div_ceil(256)).max(1), t as u32, 1),
block_dim: (256, 1, 1),
shared_mem_bytes: 0,
};
let (ne, qt, rb, ti) = (n_embd as i32, qtype, row_bytes as i64, t as i32);
let stream = e.gpu.stream();
let mut b = stream.launch_builder(&f);
b.arg(table)
.arg(tok_v)
.arg(x_out)
.arg(&ne)
.arg(&qt)
.arg(&rb)
.arg(&ti);
unsafe {
b.launch(cfg)?;
}
Ok(())
}
#[derive(Debug, Default, Clone)]
pub struct SpecReport {
pub tokens: Vec<u32>,
pub rounds: usize,
pub drafted: u64,
pub accepted: u64,
pub accept_hist: Vec<u64>,
pub draft_ms: f64,
pub verify_ms: f64,
pub prefill_ms: f64,
pub total_ms: f64,
pub chain_ms: f64,
pub replay_ms: f64,
pub draft_prefill_ms: f64,
pub cross_ms: f64,
pub cross_bytes: u64,
pub k_decays: Vec<(usize, usize)>,
pub spec_off_at: Option<usize>,
pub plain_steps: usize,
pub round_wall: Vec<(usize, f64)>,
pub zero_draft_rounds: usize,
pub guard_stops: usize,
}
#[derive(Clone, Copy, Debug)]
pub struct DynKCfg {
pub window: usize,
pub thr: f64,
pub k_floor: usize,
}
#[derive(Clone, Copy, Debug, Default)]
pub struct SpecOpts {
pub dynk: Option<DynKCfg>,
pub adapt_k_lo: Option<usize>,
pub pmin: f32,
pub defer: bool,
pub defer_guard_sync: bool,
pub prefill_chunk: Option<usize>,
pub wide_ring: Option<usize>,
}
pub fn spec_guard_trunc(probs: &[f32], pmin: f32) -> usize {
probs.iter().position(|&p| p < pmin).unwrap_or(probs.len())
}
#[derive(Debug, Default, Clone)]
pub struct SpecTraceRound {
pub round: usize,
pub gen_pos: usize,
pub base: usize,
pub k: usize,
pub a: usize,
pub drafts: Vec<u32>,
pub targets: Vec<u32>,
pub draft_top1: f32,
pub draft_top2: f32,
pub draft_tgt_logit: f32,
pub draft_tgt_rank: usize,
pub target_top1: f32,
pub target_top2: f32,
pub target_draft_logit: f32,
pub target_entropy: f64,
pub carrier_rel_l2: Vec<f32>,
pub carrier_cos: Vec<f32>,
}
impl SpecReport {
pub fn accept_rate(&self) -> f64 {
if self.drafted == 0 {
0.0
} else {
self.accepted as f64 / self.drafted as f64
}
}
pub fn mean_accept_len(&self) -> f64 {
if self.rounds == 0 {
0.0
} else {
self.tokens.len() as f64 / self.rounds as f64
}
}
}
impl Qwen4ExpGpu {
pub fn spec_arm(&self, e: &Engine, state: &mut Qwen4ExpState, k_cap: usize) -> Res<()> {
self.spec_arm_ring(e, state, k_cap, state.capacity)
}
pub fn spec_arm_ring(
&self,
e: &Engine,
state: &mut Qwen4ExpState,
k_cap: usize,
ring_rows: usize,
) -> Res<()> {
let ring_rows = ring_rows.min(state.capacity).max(k_cap + 2);
if let Some(v) = state.verify.as_ref() {
if v.k_cap == k_cap && v.ring_rows == ring_rows {
return Ok(());
}
}
let wide = self.streams * self.hidden;
let mut gdn = Vec::with_capacity(self.layers.len());
let mut ple = Vec::with_capacity(self.layers.len());
for layer in &self.layers {
gdn.push(match &layer.mixer {
MixerW::Gdn(g) => {
let p = &g.plan;
let (nk, nv) = (p.key_heads as usize, p.value_heads as usize);
let (hk, hv) = (p.key_head_dim as usize, p.value_head_dim as usize);
let conv_dim = 2 * nk * hk + nv * hv;
let pad = p.conv_kernel as usize - 1;
Some(GdnStash {
states: e.zeros(k_cap * nv * hv * hk)?,
conv_pre: e.zeros(pad * conv_dim)?,
qkv_rows: e.zeros(k_cap * conv_dim)?,
scan_graph: None,
scan_warm: None,
})
}
MixerW::Qsa(_) => None,
});
ple.push(match layer.ple.as_ref() {
Some(pw) => {
let pad = (pw.plan.conv_kernel as usize - 1) * pw.plan.max_ngram as usize;
let mut hist_pre = Vec::with_capacity(self.streams);
let mut normed_rows = Vec::with_capacity(self.streams);
for _ in 0..self.streams {
hist_pre.push(e.zeros(pad * self.hidden)?);
normed_rows.push(e.zeros(k_cap * self.hidden)?);
}
Some(PleStash {
hist_pre,
normed_rows,
})
}
None => None,
});
}
state.verify = Some(VerifyStash {
k_cap,
chunk: None,
fused_chunk: None,
gdn,
ple,
wide: e.zeros(ring_rows * wide)?,
ring_rows,
wide_dev1: None,
argmax: Vec::new(),
toks: unsafe { e.gpu.stream().alloc::<u32>(k_cap)? },
want_argmax: false,
want_argmax_t1: false,
last_row_only: false,
});
Ok(())
}
pub fn spec_disarm(&self, state: &mut Qwen4ExpState) {
state.verify = None;
}
pub fn set_verify_want_argmax(&self, state: &mut Qwen4ExpState, on: bool) -> Res<()> {
state
.verify
.as_mut()
.ok_or("qwen4exp_gpu: verify not armed")?
.want_argmax = on;
Ok(())
}
pub fn verify_argmax_rows<'s>(&self, state: &'s Qwen4ExpState) -> Res<&'s [u32]> {
Ok(&state
.verify
.as_ref()
.ok_or("qwen4exp_gpu: verify not armed")?
.argmax)
}
pub fn verify_rewind(&self, e: &Engine, state: &mut Qwen4ExpState, keep: usize) -> Res<()> {
let Some(v) = state.verify.as_mut() else {
return Err("qwen4exp_gpu: verify not armed".into());
};
let Some((base, t)) = v.chunk.take() else {
if let Some((fb, ft)) = v.fused_chunk.take() {
return Err(format!(
"qwen4exp_gpu: verify chunk (base {fb}, t {ft}) ran the FUSED program \
(`vfuse` cost instrument) — no per-column GDN/PLE stash exists, so it \
cannot be rewound. vfuse is a timing probe on a throwaway state; drop \
the seam to run a spec loop."
)
.into());
}
return Err("qwen4exp_gpu: no live verify chunk to rewind".into());
};
if keep == 0 || keep > t {
return Err("qwen4exp_gpu: rewind keep out of range".into());
}
if keep == t {
return Ok(());
}
state.pos = base + keep;
state.tokens.truncate(base + keep);
for (li, (layer, lstate)) in self.layers.iter().zip(state.layers.iter_mut()).enumerate() {
match (&layer.mixer, &mut lstate.mixer) {
(
MixerW::Qsa(qsa),
MixerState::Qsa {
raw_keys,
pooled_keys,
pooled_dev_rows,
raw_dev_rows,
idx_audit,
..
},
) => {
let idx_dim = qsa.overlay.head_dim as usize;
raw_keys.truncate_rows(base + keep, idx_dim);
let block = qsa.overlay.block_size as usize;
pooled_keys.truncate(((base + keep) / block) * idx_dim);
*pooled_dev_rows = (*pooled_dev_rows).min(pooled_keys.len() / idx_dim);
*raw_dev_rows = (*raw_dev_rows).min(base + keep);
if let Some(audit) = idx_audit.as_deref_mut() {
audit.raw_f32.truncate_rows(base + keep, idx_dim);
audit.pooled_f32.truncate(((base + keep) / block) * idx_dim);
}
}
(MixerW::Gdn(g), MixerState::Gdn { conv, state: rec }) => {
let st = v.gdn[li]
.as_mut()
.ok_or("qwen4exp_gpu: GDN layer without a verify stash")?;
let p = &g.plan;
let (nk, nv) = (p.key_heads as usize, p.value_heads as usize);
let (hk, hv) = (p.key_head_dim as usize, p.value_head_dim as usize);
let conv_dim = 2 * nk * hk + nv * hv;
let pad = p.conv_kernel as usize - 1;
let state_len = nv * hv * hk;
e.copy_range_into(rec, 0, &st.states, (keep - 1) * state_len, state_len)?;
if keep >= pad {
e.copy_range_into(
conv,
0,
&st.qkv_rows,
(keep - pad) * conv_dim,
pad * conv_dim,
)?;
} else {
let keep_hist = pad - keep;
e.copy_range_into(
conv,
0,
&st.conv_pre,
keep * conv_dim,
keep_hist * conv_dim,
)?;
e.copy_range_into(
conv,
keep_hist * conv_dim,
&st.qkv_rows,
0,
keep * conv_dim,
)?;
}
}
_ => return Err("qwen4exp_gpu: mixer/state mismatch in rewind".into()),
}
if let (Some(pw), Some(ps)) = (layer.ple.as_ref(), lstate.ple.as_mut()) {
let st = v.ple[li]
.as_mut()
.ok_or("qwen4exp_gpu: PLE layer without a verify stash")?;
let pad = (pw.plan.conv_kernel as usize - 1) * pw.plan.max_ngram as usize;
let hidden = self.hidden;
for s in 0..self.streams {
let hist = &mut ps.conv_hist[s];
if keep >= pad {
e.copy_range_into(
hist,
0,
&st.normed_rows[s],
(keep - pad) * hidden,
pad * hidden,
)?;
} else {
let keep_hist = pad - keep;
e.copy_range_into(
hist,
0,
&st.hist_pre[s],
keep * hidden,
keep_hist * hidden,
)?;
e.copy_range_into(
hist,
keep_hist * hidden,
&st.normed_rows[s],
0,
keep * hidden,
)?;
}
}
}
}
Ok(())
}
fn draft_row_argmax(
&self,
e: &Engine,
logits: &CudaSlice<f32>,
row: usize,
conf: bool,
) -> Res<(u32, f32)> {
let width = self.draft_logits_width();
let mut tok = unsafe { e.gpu.stream().alloc::<u32>(1)? };
e.argmax_token_device_col(logits, row, width, &mut tok, 0)?;
let p = if conf {
if row != 0 {
return Err("qwen4exp_gpu: draft confidence reads row 0 (the chain shape)".into());
}
let pd = e.prob_of_token_device(logits, &tok, width)?;
e.gpu.stream().clone_dtoh(&pd)?[0]
} else {
1.0
};
Ok((self.draft_token(e.gpu.stream().clone_dtoh(&tok)?[0])?, p))
}
#[allow(clippy::too_many_arguments)]
pub fn spec_generate(
&self,
e: &Engine,
prompt: &[u32],
max_new: usize,
k: usize,
state: &mut Qwen4ExpState,
dstate: &mut MtpDraftState,
sampler: Option<SpecSamplerCfg>,
) -> Res<SpecReport> {
self.spec_generate_ext(
e,
e,
prompt,
max_new,
k,
state,
dstate,
sampler,
SpecOpts::default(),
None,
)
}
#[allow(clippy::too_many_arguments)]
pub fn spec_generate_ext(
&self,
e: &Engine,
de: &Engine,
prompt: &[u32],
max_new: usize,
k: usize,
state: &mut Qwen4ExpState,
dstate: &mut MtpDraftState,
sampler: Option<SpecSamplerCfg>,
opts: SpecOpts,
mut trace: Option<&mut Vec<SpecTraceRound>>,
) -> Res<SpecReport> {
use std::time::Instant;
if k == 0 {
return Err("qwen4exp_gpu: spec needs k >= 1".into());
}
let n = prompt.len();
if n < 2 {
return Err("qwen4exp_gpu: spec needs a >= 2 token prompt".into());
}
if state.pos != 0 || dstate.rows != 0 {
return Err("qwen4exp_gpu: spec_generate wants FRESH trunk + draft states".into());
}
if state.capacity < n + max_new + k + 2 || dstate.capacity < n + max_new + k + 2 {
return Err("qwen4exp_gpu: state capacity too small for prompt + max_new + k".into());
}
self.check_draft_engine(de)?;
let dev1 = self.mtp_dev1.is_some();
if !dev1 && de.ctx().ordinal() != e.ctx().ordinal() {
return Err(
"qwen4exp_gpu: draft engine on another card, but the draft was not \
built there (load_from_dir_dev1)"
.into(),
);
}
let vocab = self.vocab;
let wide_w = self.streams * self.hidden;
let greedy = sampler.is_none();
let tracing = trace.is_some();
let guard = opts.pmin > 0.0;
let deferred = opts.defer;
if deferred && tracing {
return Err(
"qwen4exp_gpu: spec defer + trace are mutually exclusive (the trace \
instrument reads per-step host rows); run the trace on the host-chain arm"
.into(),
);
}
if deferred {
let ce = self.chain_embed.as_ref().ok_or(
"qwen4exp_gpu: SpecOpts::defer needs arm_spec_devchain on the draft engine",
)?;
if ce.dev != de.ctx().ordinal() {
return Err(format!(
"qwen4exp_gpu: the chain-embed table lives on device {} but the \
draft engine is device {} — re-arm arm_spec_devchain",
ce.dev,
de.ctx().ordinal()
)
.into());
}
if ce.for_trim != self.draft_trim.is_some() || ce.rows != self.draft_logits_width() {
return Err(
"qwen4exp_gpu: the chain-embed table was armed for a different trim \
state — re-arm arm_spec_devchain after trim changes"
.into(),
);
}
}
let (mut chain_toks_d, mut chain_probs_d) = if deferred {
(
Some(unsafe { de.gpu.stream().alloc::<u32>(k)? }),
Some(de.zeros(k)?),
)
} else {
(None, None)
};
let mut rng = sampler
.as_ref()
.map(|cfg| SpecRng(cfg.seed | 1))
.unwrap_or(SpecRng(1));
let t_total = Instant::now();
let mut report = SpecReport {
accept_hist: vec![0; k + 1],
..Default::default()
};
match opts.wide_ring {
Some(ring) => {
let chunk = opts
.prefill_chunk
.ok_or("qwen4exp_gpu: SpecOpts::wide_ring needs prefill_chunk")?;
if ring < 2 * chunk || ring < 2 * (k + 2) {
return Err("qwen4exp_gpu: wide_ring must cover 2 prefill chunks".into());
}
self.spec_arm_ring(e, state, k + 1, ring)?;
}
None => self.spec_arm(e, state, k + 1)?,
}
self.set_verify_want_argmax(state, false)?;
if let Some(v) = state.verify.as_mut() {
v.want_argmax_t1 = deferred && greedy && !tracing;
v.last_row_only = deferred;
}
let ring = state.verify.as_ref().expect("armed above").ring_rows;
if dev1 {
let v = state.verify.as_mut().expect("armed above");
if v.wide_dev1.as_ref().is_none_or(|m| m.len() < ring * wide_w) {
v.wide_dev1 = Some(de.zeros(ring * wide_w)?);
}
}
let dev_embed = deferred
&& self
.chain_embed
.as_ref()
.is_some_and(|ce| !ce.for_trim && ce.rows == self.vocab);
let t_prefill = Instant::now();
let mut draft_prefill_ms = 0f64;
let x0: u32 = match opts.prefill_chunk {
Some(chunk) if n > chunk => {
let mut b = 0usize;
let mut last = Vec::new();
let mut chunks = 0usize;
while b < n {
let mut t = chunk.min(n - b);
if n - (b + t) > 0 && n - (b + t) <= k + 1 {
t = n - b;
}
let is_last = b + t == n;
let head = if is_last {
HeadMode::LastRow
} else {
HeadMode::Skip
};
let piece = self.forward_with(e, &prompt[b..b + t], state, None, head)?;
let t_draft = Instant::now();
if dev1 {
let v = state.verify.as_mut().expect("armed above");
let VerifyStash {
wide, wide_dev1, ..
} = v;
let mirror = wide_dev1.as_mut().expect("allocated above");
for (slot, len) in ring_pieces(ring, b, t) {
report.cross_ms +=
cross_wide_rows(e, de, wide, mirror, slot, len, wide_w)?;
}
report.cross_bytes += (t * wide_w * 4) as u64;
}
let p0 = b.max(1);
if b + t > p0 {
let v = state.verify.as_ref().expect("armed above");
let seed: &CudaSlice<f32> = v.wide_dev1.as_ref().unwrap_or(&v.wide);
let (ld, cd, _) = self.draft_consume_ring(
de,
&prompt[p0..b + t],
dev_embed,
seed,
ring,
p0 - 1,
dstate,
)?;
self.mtp_recycle(dstate, ld, cd);
}
draft_prefill_ms += t_draft.elapsed().as_secs_f64() * 1e3;
b += t;
chunks += 1;
if chunks % 8 == 0 || is_last {
println!(
"# spec-prefill-progress\tfill={b}/{n}\tchunks={chunks}\t\
elapsed_s={:.1}\tdraft_s={:.1}",
t_prefill.elapsed().as_secs_f64(),
draft_prefill_ms / 1e3,
);
}
if is_last {
last = piece;
}
}
dstate.committed = n - 1;
let shed_bytes = state.ws.shed();
state.graphs = StepGraphs::default();
println!(
"# spec-prefill-shed\tworkspace_mib={:.1}\tchunks={chunks}",
shed_bytes as f64 / (1024.0 * 1024.0),
);
debug_assert_eq!(last.len(), vocab);
match sampler.as_ref() {
None => host_argmax(&last) as u32,
Some(cfg) => sample_row(cfg, &mut rng, &last),
}
}
_ => {
let prefill = self.forward(e, prompt, state, None)?;
let last = &prefill[prefill.len() - vocab..];
let x0 = match sampler.as_ref() {
None => host_argmax(last) as u32,
Some(cfg) => sample_row(cfg, &mut rng, last),
};
let t_draft0 = Instant::now();
if dev1 {
let v = state.verify.as_mut().expect("armed above");
let VerifyStash {
wide, wide_dev1, ..
} = v;
let mirror = wide_dev1.as_mut().expect("allocated above");
for (slot, len) in ring_pieces(ring, 0, n) {
report.cross_ms += cross_wide_rows(e, de, wide, mirror, slot, len, wide_w)?;
}
report.cross_bytes += (n * wide_w * 4) as u64;
}
{
let v = state.verify.as_ref().expect("armed above");
let seed: &CudaSlice<f32> = v.wide_dev1.as_ref().unwrap_or(&v.wide);
if n >= 2 {
let (ld, cd, _) = self.draft_consume_ring(
de,
&prompt[1..],
dev_embed,
seed,
ring,
0,
dstate,
)?;
self.mtp_recycle(dstate, ld, cd);
}
dstate.committed = n - 1;
}
draft_prefill_ms += t_draft0.elapsed().as_secs_f64() * 1e3;
x0
}
};
report.prefill_ms = t_prefill.elapsed().as_secs_f64() * 1e3 - draft_prefill_ms;
self.set_verify_want_argmax(state, greedy && !tracing)?;
report.tokens.push(x0);
let t_boot = Instant::now();
let (mut tip_logits, mut tip_carrier) = {
let v = state.verify.as_ref().expect("armed above");
let seed: &CudaSlice<f32> = v.wide_dev1.as_ref().unwrap_or(&v.wide);
self.mtp_draft_forward_spec(de, &[x0], dev_embed, seed, (n - 1) % ring, dstate)?
};
let mut tip_rows = 1usize;
dstate.committed = dstate.rows;
report.draft_prefill_ms = draft_prefill_ms + t_boot.elapsed().as_secs_f64() * 1e3;
report.draft_ms += report.draft_prefill_ms;
prof::split_prefill();
let mut m = n; let mut tip = x0;
let mut k_cur = k;
let mut k_next = k;
let mut window: Vec<usize> = Vec::new();
let mut round_idx = 0usize;
while report.tokens.len() < max_new {
if k_cur == 0 {
let row = self.forward(e, &[tip], state, None)?;
let next: u32 = match sampler.as_ref() {
None if deferred => self.verify_argmax_rows(state)?[0],
None => host_argmax(&row) as u32,
Some(cfg) => sample_row(cfg, &mut rng, &row),
};
report.tokens.push(next);
report.plain_steps += 1;
report
.round_wall
.push((report.tokens.len(), t_total.elapsed().as_secs_f64() * 1e3));
m += 1;
tip = next;
continue;
}
let k_round = k_next.min(k_cur).max(1);
let t_draft = Instant::now();
let mut drafts: Vec<u32> = Vec::with_capacity(k_round);
let mut chain_rows_h: Vec<Vec<f32>> = Vec::new(); let mut seeds_h: Vec<Vec<f32>> = Vec::new(); if let (Some(toks), Some(probs)) = (chain_toks_d.as_mut(), chain_probs_d.as_mut()) {
let width = self.draft_logits_width();
de.argmax_token_device_col(&tip_logits, 0, width, toks, 0)?;
if guard {
de.prob_of_token_device_col(&tip_logits, toks, 0, probs, 0, width)?;
}
let mut prev_logits = tip_logits;
let mut prev_carrier = tip_carrier;
let mut prev_rows = tip_rows;
let mut drafted = 1usize;
let mut stopped = false;
if guard && opts.defer_guard_sync {
let p = de.gpu.stream().clone_dtoh(&probs.slice(0..1))?[0];
if p < opts.pmin {
drafted = 0;
stopped = true;
report.guard_stops += 1;
}
}
while !stopped && drafted < k_round {
let (l2, c2) = self.mtp_draft_forward_devslot(
de,
toks,
drafted - 1,
&prev_carrier,
prev_rows - 1,
dstate,
)?;
self.mtp_recycle(dstate, prev_logits, prev_carrier);
prev_logits = l2;
prev_carrier = c2;
prev_rows = 1;
de.argmax_token_device_col(&prev_logits, 0, width, toks, drafted)?;
if guard {
de.prob_of_token_device_col(
&prev_logits,
toks,
drafted,
probs,
drafted,
width,
)?;
if opts.defer_guard_sync {
let p = de
.gpu
.stream()
.clone_dtoh(&probs.slice(drafted..drafted + 1))?[0];
if p < opts.pmin {
report.guard_stops += 1;
break;
}
}
}
drafted += 1;
}
self.mtp_recycle(dstate, prev_logits, prev_carrier);
if drafted > 0 {
let raw = de.gpu.stream().clone_dtoh(&toks.slice(0..drafted))?;
let trunc = if guard && !opts.defer_guard_sync {
let pw = de.gpu.stream().clone_dtoh(&probs.slice(0..drafted))?;
let trunc = spec_guard_trunc(&pw, opts.pmin);
if trunc < drafted {
report.guard_stops += 1;
}
trunc
} else {
drafted
};
for &r in raw.iter().take(trunc) {
drafts.push(self.draft_token(r)?);
}
}
} else {
let (d1, c1) = self.draft_row_argmax(de, &tip_logits, 0, guard)?;
if !(guard && c1 < opts.pmin) {
drafts.push(d1);
if tracing {
chain_rows_h
.push(de.dtoh_view(&tip_logits.slice(0..self.draft_logits_width()))?);
}
} else {
report.guard_stops += 1;
}
let mut prev_logits = tip_logits;
let mut prev_carrier = tip_carrier;
let mut prev_rows = tip_rows;
while !drafts.is_empty() && drafts.len() < k_round {
if tracing {
seeds_h.push(de.dtoh_view(
&prev_carrier.slice((prev_rows - 1) * wide_w..prev_rows * wide_w),
)?);
}
let lastd = *drafts.last().expect("non-empty");
let (l2, c2) = self.mtp_draft_forward(
de,
&[lastd],
&prev_carrier,
prev_rows - 1,
dstate,
1,
false,
)?;
self.mtp_recycle(dstate, prev_logits, prev_carrier);
prev_logits = l2;
prev_carrier = c2;
prev_rows = 1;
let (dn, cn) = self.draft_row_argmax(de, &prev_logits, 0, guard)?;
if guard && cn < opts.pmin {
report.guard_stops += 1;
break;
}
drafts.push(dn);
if tracing {
chain_rows_h
.push(de.dtoh_view(&prev_logits.slice(0..self.draft_logits_width()))?);
}
}
self.mtp_recycle(dstate, prev_logits, prev_carrier);
}
let chain_ms = t_draft.elapsed().as_secs_f64() * 1e3;
report.chain_ms += chain_ms;
report.draft_ms += chain_ms;
let t_ver = Instant::now();
let mut chunk = Vec::with_capacity(drafts.len() + 1);
chunk.push(tip);
chunk.extend_from_slice(&drafts);
let tlen = chunk.len();
let host_logits = self.forward(e, &chunk, state, None)?;
let targets: Vec<u32> = if greedy && !tracing && (tlen > 1 || deferred) {
self.verify_argmax_rows(state)?.to_vec()
} else if greedy {
(0..tlen)
.map(|row| host_argmax(&host_logits[row * vocab..(row + 1) * vocab]) as u32)
.collect()
} else {
let cfg = sampler.as_ref().expect("sampled mode");
(0..tlen)
.map(|row| {
sample_row(cfg, &mut rng, &host_logits[row * vocab..(row + 1) * vocab])
})
.collect()
};
report.verify_ms += t_ver.elapsed().as_secs_f64() * 1e3;
if targets.len() != tlen {
return Err("qwen4exp_gpu: verify produced the wrong row count".into());
}
let mut a = 0usize;
while a < drafts.len() && drafts[a] == targets[a] {
a += 1;
}
report.rounds += 1;
report.drafted += drafts.len() as u64;
report.accepted += a as u64;
report.accept_hist[a] += 1;
if drafts.is_empty() {
report.zero_draft_rounds += 1;
}
report.tokens.extend_from_slice(&targets[0..=a]);
if let Some(tr) = trace.as_deref_mut() {
let mut rec = SpecTraceRound {
round: round_idx,
gen_pos: report.tokens.len() - (a + 1),
base: m,
k: drafts.len(),
a,
drafts: drafts.clone(),
targets: targets.clone(),
draft_top1: f32::NAN,
draft_top2: f32::NAN,
draft_tgt_logit: f32::NAN,
draft_tgt_rank: 0,
target_top1: f32::NAN,
target_top2: f32::NAN,
target_draft_logit: f32::NAN,
target_entropy: 0.0,
carrier_rel_l2: Vec::new(),
carrier_cos: Vec::new(),
};
if a < drafts.len() {
let drow = &chain_rows_h[a];
let trow = &host_logits[a * vocab..(a + 1) * vocab];
let tgt = targets[a] as usize;
let dtok = drafts[a] as usize;
let (mut d1v, mut d2v) = (f32::NEG_INFINITY, f32::NEG_INFINITY);
let mut rank = 0usize;
let dt = drow[tgt];
for &v in drow.iter() {
if v > d1v {
d2v = d1v;
d1v = v;
} else if v > d2v {
d2v = v;
}
if v > dt {
rank += 1;
}
}
let (mut t1v, mut t2v) = (f32::NEG_INFINITY, f32::NEG_INFINITY);
for &v in trow.iter() {
if v > t1v {
t2v = t1v;
t1v = v;
} else if v > t2v {
t2v = v;
}
}
let mx = t1v as f64;
let mut z = 0.0f64;
let mut sxl = 0.0f64;
for &v in trow.iter() {
let ev = ((v as f64) - mx).exp();
z += ev;
sxl += ev * ((v as f64) - mx);
}
rec.draft_top1 = d1v;
rec.draft_top2 = d2v;
rec.draft_tgt_logit = dt;
rec.draft_tgt_rank = rank;
rec.target_top1 = t1v;
rec.target_top2 = t2v;
rec.target_draft_logit = trow[dtok];
rec.target_entropy = z.ln() - sxl / z;
}
let v = state.verify.as_ref().expect("armed above");
for (j, seed) in seeds_h.iter().enumerate() {
let slot = (m + j) % ring;
let truth = e.dtoh_view(&v.wide.slice(slot * wide_w..(slot + 1) * wide_w))?;
let mut dd = 0.0f64;
let mut tt = 0.0f64;
let mut st = 0.0f64;
let mut ss = 0.0f64;
for (&s, &t) in seed.iter().zip(truth.iter()) {
let (s, t) = (s as f64, t as f64);
dd += (s - t) * (s - t);
tt += t * t;
st += s * t;
ss += s * s;
}
rec.carrier_rel_l2
.push((dd.sqrt() / tt.sqrt().max(1e-30)) as f32);
rec.carrier_cos
.push((st / (ss.sqrt() * tt.sqrt()).max(1e-30)) as f32);
}
tr.push(rec);
}
if tlen > 1 {
self.verify_rewind(e, state, a + 1)?;
}
self.mtp_rewind(dstate, m)?;
let t_draft2 = Instant::now();
let x_next = targets[a];
let mut replay: Vec<u32> = drafts[0..a].to_vec();
replay.push(x_next);
if dev1 {
let v = state.verify.as_mut().expect("armed above");
let VerifyStash {
wide, wide_dev1, ..
} = v;
let mirror = wide_dev1.as_mut().expect("allocated above");
for (slot, len) in ring_pieces(ring, m, replay.len()) {
report.cross_ms += cross_wide_rows(e, de, wide, mirror, slot, len, wide_w)?;
}
report.cross_bytes += (replay.len() * wide_w * 4) as u64;
}
let (l, c, last_len) = {
let v = state.verify.as_ref().expect("armed above");
let seed: &CudaSlice<f32> = v.wide_dev1.as_ref().unwrap_or(&v.wide);
self.draft_consume_ring(de, &replay, dev_embed, seed, ring, m, dstate)?
};
tip_logits = l;
tip_carrier = c;
tip_rows = last_len;
dstate.committed = dstate.rows;
let replay_ms = t_draft2.elapsed().as_secs_f64() * 1e3;
report.replay_ms += replay_ms;
report.draft_ms += replay_ms;
m += a + 1;
tip = x_next;
report
.round_wall
.push((report.tokens.len(), t_total.elapsed().as_secs_f64() * 1e3));
if let Some(lo) = opts.adapt_k_lo {
k_next = (a + 1).clamp(lo.max(1), k);
}
if let Some(cfg) = opts.dynk {
window.push(a);
if window.len() >= cfg.window.max(1) {
let mean = window.iter().sum::<usize>() as f64 / window.len() as f64;
if mean < cfg.thr {
let new_k = k_cur.saturating_sub(1).max(cfg.k_floor);
if new_k < k_cur {
k_cur = new_k;
report.k_decays.push((round_idx, k_cur));
if k_cur == 0 {
report.spec_off_at = Some(report.tokens.len());
}
}
}
window.clear();
}
}
round_idx += 1;
}
report.tokens.truncate(max_new);
report.total_ms = t_total.elapsed().as_secs_f64() * 1e3;
Ok(report)
}
}
enum BankTensorSrc {
F32(Vec<f32>),
Nvfp4 {
codes: Vec<u8>,
scales: Vec<u8>,
macros: Vec<f32>,
act_scale: Option<f32>,
},
Bf16(Vec<u8>),
}
struct BankSrc {
gate: BankTensorSrc, up: BankTensorSrc, down: BankTensorSrc, n_expert: usize,
ff: usize,
hidden: usize,
}
struct FusedTensorPlan {
name: String,
shape: [usize; 3],
}
struct PerExpertPlan {
names: Vec<String>, out_f: usize,
in_f: usize,
quant: memra_gguf::tensor_contract::QuantConstraint,
}
enum BankPlanSrc {
Fused {
gate_up: FusedTensorPlan,
down: FusedTensorPlan,
keep_bf16: bool,
},
PerExpert {
gate: PerExpertPlan,
up: PerExpertPlan,
down: PerExpertPlan,
},
}
struct BankPlan {
n_expert: usize,
ff: usize,
hidden: usize,
src: BankPlanSrc,
}
impl BankPlan {
fn read(&self, model: &memra_gguf::safetensors::StModel) -> Res<BankSrc> {
let (gate, up, down) = match &self.src {
BankPlanSrc::Fused {
gate_up,
down,
keep_bf16,
} => {
let fused = read_bank_tensor(
model,
&gate_up.name,
gate_up.shape[0],
gate_up.shape[1],
gate_up.shape[2],
*keep_bf16,
)?;
let (gate, up) = split_fused_gate_up(fused, self.n_expert, self.ff, self.hidden)?;
let down = read_bank_tensor(
model,
&down.name,
down.shape[0],
down.shape[1],
down.shape[2],
*keep_bf16,
)?;
(gate, up, down)
}
BankPlanSrc::PerExpert { gate, up, down } => (
read_per_expert_bank(model, gate)?,
read_per_expert_bank(model, up)?,
read_per_expert_bank(model, down)?,
),
};
Ok(BankSrc {
gate,
up,
down,
n_expert: self.n_expert,
ff: self.ff,
hidden: self.hidden,
})
}
}
fn check_bank_header(model: &memra_gguf::safetensors::StModel, name: &str) -> Res<()> {
let info = model
.info(name)
.ok_or_else(|| format!("qwen4exp_gpu: checkpoint is missing {name}"))?;
match info.dtype.as_str() {
"BF16" | "F32" | "U8" => Ok(()),
other => Err(format!("qwen4exp_gpu: {name} bank dtype {other} unsupported").into()),
}
}
fn read_per_expert_bank(
model: &memra_gguf::safetensors::StModel,
plan: &PerExpertPlan,
) -> Res<BankTensorSrc> {
let mut experts = Vec::with_capacity(plan.names.len());
for name in &plan.names {
experts.push(read_per_expert(
model, name, plan.out_f, plan.in_f, plan.quant,
)?);
}
assemble_per_expert_bank(experts)
}
pub struct LoadedCheckpoint {
pub plan: ModelPlan,
pub weights: ReferenceWeights,
model: memra_gguf::safetensors::StModel,
bank_plans: std::collections::BTreeMap<u32, BankPlan>,
tables: std::collections::BTreeMap<u32, Vec<u8>>, }
fn norm_fold_add_one(name: &str) -> bool {
name.contains("norm") && name.ends_with(".weight") && !name.ends_with("linear_attn.norm.weight")
}
fn indexer_layernorm(name: &str) -> bool {
name.contains(".indexer.")
&& (name.ends_with("q_layernorm.weight") || name.ends_with("k_layernorm.weight"))
}
#[derive(Default, Clone, Copy)]
pub struct LoadOptions {
pub host_bf16_banks: bool,
pub indexer_norm_raw: bool,
pub load_mtp: bool,
}
fn bridge_transform(
transform: memra_gguf::tensor_contract::TensorTransform,
) -> Res<memra_gguf::hf_mapping::TransformKind> {
use memra_gguf::hf_mapping::TransformKind as K;
use memra_gguf::tensor_contract::TensorTransform as T;
Ok(match transform {
T::Identity => K::Identity,
T::NormAddOne => K::NormPlusOne,
T::QkvVReorderRows => K::QkvVReorderRows,
T::ZReorderRows => K::ZReorderRows,
T::AbReorderRows => K::AbReorderRows,
T::NegExpReorderHeads => K::NegExpReorderHeads,
T::ReorderHeads => K::ReorderHeads,
T::Conv1dSqueezeReorder => K::Conv1dSqueezeReorder,
T::OutReorderColumns => K::OutReorderCols,
other => return Err(format!("qwen4exp_gpu: unsupported transform {other:?}").into()),
})
}
fn dequant_float(
name: &str,
info: &memra_gguf::safetensors::StInfo,
bytes: &[u8],
) -> Res<Vec<f32>> {
let elements: usize = info.shape.iter().map(|&d| d as usize).product();
match info.dtype.as_str() {
"BF16" | "F32" => Ok(memra_gguf::dequant::dequantize(
info.ggml_type()
.map_err(|error| format!("qwen4exp_gpu: {name}: {error}"))?,
bytes,
elements,
)),
other => Err(format!("qwen4exp_gpu: {name} has unsupported float dtype {other}").into()),
}
}
fn read_i64(name: &str, info: &memra_gguf::safetensors::StInfo, bytes: &[u8]) -> Res<Vec<i64>> {
if info.dtype != "I64" {
return Err(format!("qwen4exp_gpu: {name} must be I64, got {}", info.dtype).into());
}
Ok(bytes
.chunks_exact(8)
.map(|chunk| i64::from_le_bytes(chunk.try_into().unwrap()))
.collect())
}
fn validate_macro(stem: &str, value: f32) -> Res<()> {
if !(value.is_finite() && value > 0.0) {
return Err(format!(
"qwen4exp_gpu: {stem}.weight_scale_2 carries a non-finite/non-positive \
macro {value}"
)
.into());
}
Ok(())
}
fn read_bank_tensor(
model: &memra_gguf::safetensors::StModel,
name: &str,
n_expert: usize,
out_f: usize,
in_f: usize,
host_bf16: bool,
) -> Res<BankTensorSrc> {
let (info, bytes) = model
.raw(name)
.ok_or_else(|| format!("qwen4exp_gpu: checkpoint is missing {name}"))?;
match info.dtype.as_str() {
"BF16" | "F32" => {
if info.shape != [n_expert as u64, out_f as u64, in_f as u64] {
return Err(format!("qwen4exp_gpu: {name} bank shape mismatch").into());
}
if host_bf16 && info.dtype == "BF16" {
if bytes.len() != n_expert * out_f * in_f * 2 {
return Err(format!("qwen4exp_gpu: {name} bank byte-length mismatch").into());
}
return Ok(BankTensorSrc::Bf16(bytes.to_vec()));
}
Ok(BankTensorSrc::F32(dequant_float(name, info, bytes)?))
}
"U8" => {
if in_f % 16 != 0
|| info.shape != [n_expert as u64, out_f as u64, (in_f / 2) as u64]
|| bytes.len() != n_expert * out_f * in_f / 2
{
return Err(format!("qwen4exp_gpu: {name} NVFP4 code shape mismatch").into());
}
let stem = name.strip_suffix(".weight").unwrap_or(name);
let scale_name = format!("{stem}.weight_scale");
let (scale_info, scale_bytes) = model
.raw(&scale_name)
.ok_or_else(|| format!("qwen4exp_gpu: missing {scale_name}"))?;
if scale_info.dtype != "F8_E4M3"
|| scale_info.shape != [n_expert as u64, out_f as u64, (in_f / 16) as u64]
|| scale_bytes.len() != n_expert * out_f * in_f / 16
{
return Err(format!("qwen4exp_gpu: {scale_name} shape mismatch").into());
}
let macros = match model.raw(&format!("{stem}.weight_scale_2")) {
Some((macro_info, macro_bytes))
if macro_info.dtype == "F32" && macro_bytes.len() == n_expert * 4 =>
{
macro_bytes
.chunks_exact(4)
.map(|chunk| f32::from_le_bytes(chunk.try_into().unwrap()))
.collect()
}
None => vec![1.0; n_expert],
_ => return Err(format!("qwen4exp_gpu: {stem}.weight_scale_2 malformed").into()),
};
for &m in ¯os {
validate_macro(stem, m)?;
}
let act_scale = match model.raw(&format!("{stem}.input_scale")) {
Some((is_info, is_bytes))
if is_info.dtype == "F32" && is_bytes.len() == n_expert * 4 =>
{
let mut mx = 0.0f32;
for chunk in is_bytes.chunks_exact(4) {
let v = f32::from_le_bytes(chunk.try_into().unwrap());
if !(v.is_finite() && v > 0.0) {
return Err(format!(
"qwen4exp_gpu: {stem}.input_scale carries a non-finite/\
non-positive value {v}"
)
.into());
}
mx = mx.max(v);
}
Some(mx)
}
Some(_) => {
return Err(format!("qwen4exp_gpu: {stem}.input_scale malformed").into());
}
None => None,
};
Ok(BankTensorSrc::Nvfp4 {
codes: bytes.to_vec(),
scales: scale_bytes.to_vec(),
macros,
act_scale,
})
}
other => Err(format!("qwen4exp_gpu: {name} bank dtype {other} unsupported").into()),
}
}
fn split_fused_gate_up(
fused: BankTensorSrc,
n_expert: usize,
ff: usize,
hidden: usize,
) -> Res<(BankTensorSrc, BankTensorSrc)> {
match fused {
BankTensorSrc::F32(data) => {
if data.len() != n_expert * 2 * ff * hidden {
return Err("qwen4exp_gpu: fused gate_up bank size mismatch".into());
}
let mut gate = Vec::with_capacity(n_expert * ff * hidden);
let mut up = Vec::with_capacity(n_expert * ff * hidden);
for expert in 0..n_expert {
let base = expert * 2 * ff * hidden;
gate.extend_from_slice(&data[base..base + ff * hidden]);
up.extend_from_slice(&data[base + ff * hidden..base + 2 * ff * hidden]);
}
Ok((BankTensorSrc::F32(gate), BankTensorSrc::F32(up)))
}
BankTensorSrc::Bf16(bytes) => {
let row = hidden * 2; if bytes.len() != n_expert * 2 * ff * row {
return Err("qwen4exp_gpu: fused bf16 gate_up bank size mismatch".into());
}
let mut gate = Vec::with_capacity(n_expert * ff * row);
let mut up = Vec::with_capacity(n_expert * ff * row);
for expert in 0..n_expert {
let base = expert * 2 * ff * row;
gate.extend_from_slice(&bytes[base..base + ff * row]);
up.extend_from_slice(&bytes[base + ff * row..base + 2 * ff * row]);
}
Ok((BankTensorSrc::Bf16(gate), BankTensorSrc::Bf16(up)))
}
BankTensorSrc::Nvfp4 {
codes,
scales,
macros,
act_scale,
} => {
let code_row = hidden / 2;
let scale_row = hidden / 16;
let mut gate_codes = Vec::with_capacity(n_expert * ff * code_row);
let mut up_codes = Vec::with_capacity(n_expert * ff * code_row);
let mut gate_scales = Vec::with_capacity(n_expert * ff * scale_row);
let mut up_scales = Vec::with_capacity(n_expert * ff * scale_row);
for expert in 0..n_expert {
let cbase = expert * 2 * ff * code_row;
gate_codes.extend_from_slice(&codes[cbase..cbase + ff * code_row]);
up_codes
.extend_from_slice(&codes[cbase + ff * code_row..cbase + 2 * ff * code_row]);
let sbase = expert * 2 * ff * scale_row;
gate_scales.extend_from_slice(&scales[sbase..sbase + ff * scale_row]);
up_scales
.extend_from_slice(&scales[sbase + ff * scale_row..sbase + 2 * ff * scale_row]);
}
Ok((
BankTensorSrc::Nvfp4 {
codes: gate_codes,
scales: gate_scales,
macros: macros.clone(),
act_scale,
},
BankTensorSrc::Nvfp4 {
codes: up_codes,
scales: up_scales,
macros,
act_scale,
},
))
}
}
}
enum PerExpertSrc {
F32(Vec<f32>),
Nvfp4 {
codes: Vec<u8>,
scales: Vec<u8>,
macro_scale: f32,
input_scale: Option<f32>,
},
}
fn read_per_expert(
model: &memra_gguf::safetensors::StModel,
name: &str,
out_f: usize,
in_f: usize,
quant: memra_gguf::tensor_contract::QuantConstraint,
) -> Res<PerExpertSrc> {
use memra_gguf::tensor_contract::QuantConstraint;
let (info, bytes) = model
.raw(name)
.ok_or_else(|| format!("qwen4exp_gpu: checkpoint is missing {name}"))?;
match quant {
QuantConstraint::ExactFloat(_) => {
if info.shape != [out_f as u64, in_f as u64] {
return Err(format!("qwen4exp_gpu: {name} shape mismatch").into());
}
Ok(PerExpertSrc::F32(dequant_float(name, info, bytes)?))
}
QuantConstraint::Nvfp4 => {
if info.dtype != "U8"
|| in_f % 16 != 0
|| info.shape != [out_f as u64, (in_f / 2) as u64]
|| bytes.len() != out_f * in_f / 2
{
return Err(format!("qwen4exp_gpu: {name} NVFP4 code shape mismatch").into());
}
let stem = name.strip_suffix(".weight").unwrap_or(name);
let (scale_info, scale_bytes) = model
.raw(&format!("{stem}.weight_scale"))
.ok_or_else(|| format!("qwen4exp_gpu: missing {stem}.weight_scale"))?;
if scale_info.dtype != "F8_E4M3"
|| scale_info.shape != [out_f as u64, (in_f / 16) as u64]
|| scale_bytes.len() != out_f * in_f / 16
{
return Err(format!("qwen4exp_gpu: {stem}.weight_scale shape mismatch").into());
}
let macro_scale = match model.raw(&format!("{stem}.weight_scale_2")) {
Some((macro_info, macro_bytes))
if macro_info.dtype == "F32" && macro_bytes.len() == 4 =>
{
f32::from_le_bytes(macro_bytes.try_into().unwrap())
}
None => 1.0,
_ => return Err(format!("qwen4exp_gpu: {stem}.weight_scale_2 malformed").into()),
};
validate_macro(stem, macro_scale)?;
let input_scale = match model.raw(&format!("{stem}.input_scale")) {
Some((input_info, input_bytes)) => {
if input_info.dtype != "F32" || input_bytes.len() != 4 {
return Err(format!("qwen4exp_gpu: {stem}.input_scale malformed").into());
}
let v = f32::from_le_bytes(input_bytes.try_into().unwrap());
if !(v.is_finite() && v > 0.0) {
return Err(format!(
"qwen4exp_gpu: {stem}.input_scale carries a non-finite/non-positive \
value {v}"
)
.into());
}
Some(v)
}
None => None,
};
Ok(PerExpertSrc::Nvfp4 {
codes: bytes.to_vec(),
scales: scale_bytes.to_vec(),
macro_scale,
input_scale,
})
}
other => Err(format!("qwen4exp_gpu: per-expert quant {other:?} unsupported").into()),
}
}
fn assemble_per_expert_bank(experts: Vec<PerExpertSrc>) -> Res<BankTensorSrc> {
let mut f32_data: Vec<f32> = Vec::new();
let mut codes: Vec<u8> = Vec::new();
let mut scales: Vec<u8> = Vec::new();
let mut macros: Vec<f32> = Vec::new();
let mut act_scale: Option<f32> = None;
let mut act_scale_complete = true;
let mut kinds = (false, false);
for expert in experts {
match expert {
PerExpertSrc::F32(data) => {
kinds.0 = true;
f32_data.extend_from_slice(&data);
}
PerExpertSrc::Nvfp4 {
codes: c,
scales: s,
macro_scale,
input_scale,
} => {
kinds.1 = true;
codes.extend_from_slice(&c);
scales.extend_from_slice(&s);
macros.push(macro_scale);
match input_scale {
Some(v) => act_scale = Some(act_scale.map_or(v, |a: f32| a.max(v))),
None => act_scale_complete = false,
}
}
}
}
match kinds {
(true, false) => Ok(BankTensorSrc::F32(f32_data)),
(false, true) => Ok(BankTensorSrc::Nvfp4 {
codes,
scales,
macros,
act_scale: if act_scale_complete { act_scale } else { None },
}),
_ => Err("qwen4exp_gpu: mixed per-expert kinds within one projection".into()),
}
}
fn family_layer_index(key: &str) -> Option<u32> {
key.strip_prefix("trunk.layers.")?
.split('.')
.next()?
.parse()
.ok()
}
fn plan_layer_at(plan: &ModelPlan, index: u32) -> Option<&memra_gguf::model_plan::LayerPlan> {
let n_trunk = plan.layers.len() as u32;
if index < n_trunk {
plan.layers.get(index as usize)
} else {
plan.mtp_blocks
.iter()
.find(|block| block.layer.index == index)
.map(|block| &block.layer)
}
}
pub fn read_checkpoint(dir: &std::path::Path) -> Res<LoadedCheckpoint> {
read_checkpoint_with(dir, LoadOptions::default())
}
pub fn read_checkpoint_with(dir: &std::path::Path, opts: LoadOptions) -> Res<LoadedCheckpoint> {
use memra_gguf::model_packs::qwen4_exp::{ExpertDialect, tensor_contract_for};
use memra_gguf::tensor_contract::{TensorMatch, TensorOwner};
let config = std::fs::read_to_string(dir.join("config.json"))?;
let cfg =
memra_gguf::config::ModelConfig::from_hf(&memra_gguf::config::HfConfig::parse(&config));
let pack = memra_gguf::model_packs::for_config(&cfg)
.ok_or("qwen4exp_gpu: no model pack matches this config")?;
if pack.family != "qwen4_exp" {
return Err(format!("qwen4exp_gpu: config resolves to pack {}", pack.family).into());
}
let plan = pack.compile_plan(&cfg)?;
let model = memra_gguf::safetensors::StModel::open(dir)?;
let dialect = if model
.raw("model.language_model.layers.0.mlp.experts.0.gate_proj.weight")
.is_some()
{
ExpertDialect::PerExpertModelopt
} else {
ExpertDialect::FusedBanks
};
let contract = tensor_contract_for(&cfg, &plan, dialect)?;
let mut weights = ReferenceWeights::new();
let mut gate_up_banks: std::collections::BTreeMap<u32, FusedTensorPlan> = Default::default();
let mut per_expert: std::collections::BTreeMap<
(u32, u8),
std::collections::BTreeMap<
u32,
(
String,
usize,
usize,
memra_gguf::tensor_contract::QuantConstraint,
),
>,
> = Default::default();
let mut down_banks: std::collections::BTreeMap<u32, FusedTensorPlan> = Default::default();
let mut tables: std::collections::BTreeMap<u32, Vec<u8>> = Default::default();
let n_trunk = plan.layers.len() as u32;
for requirement in &contract.requirements {
match requirement.owner {
TensorOwner::Mtp(_) if !opts.load_mtp => continue,
TensorOwner::Vision(_) => continue,
TensorOwner::Global | TensorOwner::Layer(_) | TensorOwner::Mtp(_) => {}
}
if requirement.match_mode == TensorMatch::All {
let TensorId::Family { key, .. } = &requirement.id else {
return Err("qwen4exp_gpu: unexpected All-mode requirement".into());
};
let layer =
family_layer_index(key).ok_or("qwen4exp_gpu: n-gram bank outside a trunk layer")?;
let mut bytes = Vec::new();
for name in &requirement.names {
let (info, shard) = model
.raw(name)
.ok_or_else(|| format!("qwen4exp_gpu: checkpoint is missing {name}"))?;
if info.dtype != "BF16" || info.shape != requirement.shape {
return Err(format!("qwen4exp_gpu: {name} shard shape/dtype mismatch").into());
}
bytes.extend_from_slice(shard);
}
tables.insert(layer, bytes);
continue;
}
let name = &requirement.names[0];
if let TensorId::Family { key, .. } = &requirement.id {
if key.ends_with(".ple_embedding.ngram_embedding") {
let layer = family_layer_index(key)
.ok_or("qwen4exp_gpu: n-gram table outside a trunk layer")?;
let (info, bytes) = model
.raw(name)
.ok_or_else(|| format!("qwen4exp_gpu: checkpoint is missing {name}"))?;
if info.dtype != "BF16" || info.shape != requirement.shape {
return Err(format!("qwen4exp_gpu: {name} table shape/dtype mismatch").into());
}
tables.insert(layer, bytes.to_vec());
continue;
}
}
if let TensorId::Expert {
layer,
expert,
tensor,
} = requirement.id
{
let (out_f, in_f) = (requirement.shape[0] as usize, requirement.shape[1] as usize);
check_bank_header(&model, name)?;
let proj = match tensor {
memra_gguf::tensor_contract::ExpertTensor::Gate => 0u8,
memra_gguf::tensor_contract::ExpertTensor::Up => 1,
memra_gguf::tensor_contract::ExpertTensor::Down => 2,
};
if per_expert
.entry((layer, proj))
.or_default()
.insert(expert, (name.clone(), out_f, in_f, requirement.quant))
.is_some()
{
return Err(format!(
"qwen4exp_gpu: duplicate per-expert row layer {layer} expert {expert}"
)
.into());
}
continue;
}
if let TensorId::Layer { index, tensor } = requirement.id {
if matches!(
tensor,
LayerTensor::MoeExpertGateUpBank | LayerTensor::MoeExpertDownBank
) {
let shape = [
requirement.shape[0] as usize,
requirement.shape[1] as usize,
requirement.shape[2] as usize,
];
check_bank_header(&model, name)?;
let address = FusedTensorPlan {
name: name.clone(),
shape,
};
if tensor == LayerTensor::MoeExpertGateUpBank {
gate_up_banks.insert(index, address);
} else {
down_banks.insert(index, address);
}
continue;
}
}
let (info, bytes) = model
.raw(name)
.ok_or_else(|| format!("qwen4exp_gpu: checkpoint is missing {name}"))?;
if info.shape != requirement.shape {
return Err(format!(
"qwen4exp_gpu: {name} shape {:?} != contract {:?}",
info.shape, requirement.shape
)
.into());
}
if info.dtype == "I64" {
let ints = read_i64(name, info, bytes)?;
let shape: Vec<usize> = info.shape.iter().map(|&d| d as usize).collect();
weights.insert(
requirement.id.clone(),
ReferenceTensor::new_i64(shape, ints)?,
);
continue;
}
let mut data = dequant_float(name, info, bytes)?;
if norm_fold_add_one(name) && !(opts.indexer_norm_raw && indexer_layernorm(name)) {
for value in &mut data {
*value += 1.0;
}
}
let kind = bridge_transform(requirement.transform)?;
let (ne_out, out_bytes) = kind.apply(&mut data, info.ne(), &cfg);
let data: Vec<f32> = out_bytes
.chunks_exact(4)
.map(|chunk| f32::from_le_bytes(chunk.try_into().unwrap()))
.collect();
let mut shape: Vec<usize> = ne_out.iter().rev().map(|&d| d as usize).collect();
if name.ends_with("ple.conv1d.weight") && shape.len() == 3 && shape[1] == 1 {
shape = vec![shape[0], shape[2]];
}
if name.ends_with("mlp.shared_expert_gate.weight") && shape.len() == 2 && shape[0] == 1 {
shape = vec![shape[1]];
}
weights.insert(requirement.id.clone(), ReferenceTensor::new(shape, data)?);
}
let mut bank_plans = std::collections::BTreeMap::new();
for (index, gate_up) in gate_up_banks {
let down = down_banks
.remove(&index)
.ok_or_else(|| format!("qwen4exp_gpu: layer {index} has gate_up but no down bank"))?;
let layer_plan = plan_layer_at(&plan, index)
.ok_or_else(|| format!("qwen4exp_gpu: bank at unknown layer index {index}"))?;
let MlpPlan::Moe(moe) = &layer_plan.mlp else {
return Err(format!("qwen4exp_gpu: bank on non-MoE layer {index}").into());
};
bank_plans.insert(
index,
BankPlan {
n_expert: moe.expert_count as usize,
ff: moe.expert_intermediate_size as usize,
hidden: plan.hidden_size as usize,
src: BankPlanSrc::Fused {
gate_up,
down,
keep_bf16: opts.host_bf16_banks || index >= n_trunk,
},
},
);
}
if !down_banks.is_empty() {
return Err("qwen4exp_gpu: down bank without a gate_up twin".into());
}
let mut per_layer: std::collections::BTreeMap<u32, [Option<PerExpertPlan>; 3]> =
Default::default();
for ((layer, proj), experts) in per_expert {
let layer_plan = plan_layer_at(&plan, layer)
.ok_or_else(|| format!("qwen4exp_gpu: per-expert rows at unknown layer {layer}"))?;
let MlpPlan::Moe(moe) = &layer_plan.mlp else {
return Err(format!("qwen4exp_gpu: per-expert rows on non-MoE layer {layer}").into());
};
let count = moe.expert_count as usize;
if experts.len() != count || experts.keys().last().copied() != Some(count as u32 - 1) {
return Err(format!(
"qwen4exp_gpu: layer {layer} proj {proj} has {} experts, plan says {count}",
experts.len()
)
.into());
}
let mut names = Vec::with_capacity(count);
let mut geometry: Option<(usize, usize, memra_gguf::tensor_contract::QuantConstraint)> =
None;
for (name, out_f, in_f, quant) in experts.into_values() {
match geometry {
None => geometry = Some((out_f, in_f, quant)),
Some((o, i, q)) if (o, i) == (out_f, in_f) && q == quant => {}
Some((o, i, _)) => {
return Err(format!(
"qwen4exp_gpu: layer {layer} proj {proj} mixes expert geometry \
({out_f}, {in_f}) vs ({o}, {i}) or quant classes"
)
.into());
}
}
names.push(name);
}
let (out_f, in_f, quant) = geometry
.ok_or_else(|| format!("qwen4exp_gpu: layer {layer} proj {proj} has no expert rows"))?;
per_layer.entry(layer).or_default()[proj as usize] = Some(PerExpertPlan {
names,
out_f,
in_f,
quant,
});
}
for (layer, mut projections) in per_layer {
let MlpPlan::Moe(moe) = &plan_layer_at(&plan, layer).expect("checked above").mlp else {
unreachable!("checked above");
};
let take = |slot: &mut Option<PerExpertPlan>, what: &str| -> Res<PerExpertPlan> {
slot.take()
.ok_or_else(|| format!("qwen4exp_gpu: layer {layer} missing {what} experts").into())
};
bank_plans.insert(
layer,
BankPlan {
n_expert: moe.expert_count as usize,
ff: moe.expert_intermediate_size as usize,
hidden: plan.hidden_size as usize,
src: BankPlanSrc::PerExpert {
gate: take(&mut projections[0], "gate")?,
up: take(&mut projections[1], "up")?,
down: take(&mut projections[2], "down")?,
},
},
);
}
Ok(LoadedCheckpoint {
plan,
weights,
model,
bank_plans,
tables,
})
}
pub struct BankFingerprint {
pub layer: u32,
pub projection: &'static str,
pub kind: &'static str,
pub bytes: usize,
pub digest: String,
}
impl LoadedCheckpoint {
fn read_bank(&self, index: u32) -> Res<BankSrc> {
self.bank_plans
.get(&index)
.ok_or_else(|| format!("qwen4exp_gpu: no bank source for layer {index}"))?
.read(&self.model)
}
pub fn bank_fingerprints(&self) -> Res<Vec<BankFingerprint>> {
use sha2::{Digest, Sha256};
let mut out = Vec::new();
for (&layer, plan) in &self.bank_plans {
let bank = plan.read(&self.model)?;
for (projection, src) in [("gate", &bank.gate), ("up", &bank.up), ("down", &bank.down)]
{
let mut hasher = Sha256::new();
let (kind, bytes) = match src {
BankTensorSrc::F32(data) => {
for value in data {
hasher.update(value.to_le_bytes());
}
("f32", data.len() * 4)
}
BankTensorSrc::Bf16(raw) => {
hasher.update(raw);
("bf16", raw.len())
}
BankTensorSrc::Nvfp4 {
codes,
scales,
macros,
..
} => {
hasher.update(codes);
hasher.update(scales);
for m in macros {
hasher.update(m.to_le_bytes());
}
("nvfp4", codes.len() + scales.len() + macros.len() * 4)
}
};
out.push(BankFingerprint {
layer,
projection,
kind,
bytes,
digest: hasher
.finalize()
.iter()
.map(|b| format!("{b:02x}"))
.collect(),
});
}
}
Ok(out)
}
pub fn into_reference_weights(mut self) -> Res<ReferenceWeights> {
let bank_plans = std::mem::take(&mut self.bank_plans);
for (index, plan) in bank_plans {
let bank = plan.read(&self.model)?;
let gate = bank_to_f32(&bank.gate, bank.n_expert, bank.ff, bank.hidden)?;
let up = bank_to_f32(&bank.up, bank.n_expert, bank.ff, bank.hidden)?;
let down = bank_to_f32(&bank.down, bank.n_expert, bank.hidden, bank.ff)?;
self.weights.insert(
layer_id(index, LayerTensor::MoeExpertGateBank),
ReferenceTensor::new(vec![bank.n_expert, bank.ff, bank.hidden], gate)?,
);
self.weights.insert(
layer_id(index, LayerTensor::MoeExpertUpBank),
ReferenceTensor::new(vec![bank.n_expert, bank.ff, bank.hidden], up)?,
);
self.weights.insert(
layer_id(index, LayerTensor::MoeExpertDownBank),
ReferenceTensor::new(vec![bank.n_expert, bank.hidden, bank.ff], down)?,
);
}
for (index, bytes) in self.tables {
let ple = self.plan.layers[index as usize]
.ple
.as_ref()
.ok_or("qwen4exp_gpu: table on a non-PLE layer")?;
let head_dim = ple.head_embed_dim as usize;
let table = NgramTable::Bf16(bytes);
let rows = table.rows(head_dim);
let mut data = vec![0.0f32; rows * head_dim];
for row in 0..rows {
table.gather_into(
row,
head_dim,
&mut data[row * head_dim..(row + 1) * head_dim],
);
}
self.weights.insert(
family_id(format!(
"trunk.layers.{index}.ple.ple_embedding.ngram_embedding"
)),
ReferenceTensor::new(vec![rows, head_dim], data)?,
);
}
Ok(self.weights)
}
}
impl LoadedCheckpoint {
pub fn mtp_reference_weights(&self) -> Res<ReferenceWeights> {
let mut weights = self.weights.clone();
let n_trunk = self.plan.layers.len() as u32;
for (index, plan) in &self.bank_plans {
if *index < n_trunk {
continue;
}
let bank = plan.read(&self.model)?;
let gate = bank_to_f32(&bank.gate, bank.n_expert, bank.ff, bank.hidden)?;
let up = bank_to_f32(&bank.up, bank.n_expert, bank.ff, bank.hidden)?;
let down = bank_to_f32(&bank.down, bank.n_expert, bank.hidden, bank.ff)?;
weights.insert(
layer_id(*index, LayerTensor::MoeExpertGateBank),
ReferenceTensor::new(vec![bank.n_expert, bank.ff, bank.hidden], gate)?,
);
weights.insert(
layer_id(*index, LayerTensor::MoeExpertUpBank),
ReferenceTensor::new(vec![bank.n_expert, bank.ff, bank.hidden], up)?,
);
weights.insert(
layer_id(*index, LayerTensor::MoeExpertDownBank),
ReferenceTensor::new(vec![bank.n_expert, bank.hidden, bank.ff], down)?,
);
}
Ok(weights)
}
}
fn bank_to_f32(bank: &BankTensorSrc, n_expert: usize, out_f: usize, in_f: usize) -> Res<Vec<f32>> {
match bank {
BankTensorSrc::F32(data) => Ok(data.clone()),
BankTensorSrc::Bf16(bytes) => Ok(bytes
.chunks_exact(2)
.map(|b| f32::from_bits(u32::from(u16::from_le_bytes([b[0], b[1]])) << 16))
.collect()),
BankTensorSrc::Nvfp4 {
codes,
scales,
macros,
..
} => {
let mut out = Vec::with_capacity(n_expert * out_f * in_f);
let wbytes = out_f * in_f / 2;
let sbytes = out_f * in_f / 16;
for expert in 0..n_expert {
out.extend(memra_gguf::dsv4::dequant_nvfp4_expert(
&codes[expert * wbytes..(expert + 1) * wbytes],
&scales[expert * sbytes..(expert + 1) * sbytes],
macros[expert],
out_f,
in_f,
));
}
Ok(out)
}
}
}
impl Qwen4ExpGpu {
pub fn load_from_dir(e: &Engine, dir: &std::path::Path) -> Res<Self> {
Self::from_loaded_checkpoint(e, read_checkpoint(dir)?)
}
pub fn load_from_dir_with(e: &Engine, dir: &std::path::Path, opts: LoadOptions) -> Res<Self> {
Self::from_loaded_checkpoint(e, read_checkpoint_with(dir, opts)?)
}
pub fn load_from_dir_dev1(
e: &Engine,
draft_e: &Engine,
dir: &std::path::Path,
opts: LoadOptions,
) -> Res<Self> {
Self::from_loaded_checkpoint_dual(e, Some(draft_e), read_checkpoint_with(dir, opts)?)
}
pub fn from_loaded_checkpoint(e: &Engine, checkpoint: LoadedCheckpoint) -> Res<Self> {
Self::from_loaded_checkpoint_dual(e, None, checkpoint)
}
pub fn from_loaded_checkpoint_dual(
e: &Engine,
draft_e: Option<&Engine>,
checkpoint: LoadedCheckpoint,
) -> Res<Self> {
let LoadedCheckpoint {
plan,
weights,
model,
bank_plans,
tables,
} = checkpoint;
let mut parts = ExternalParts::default();
let n_trunk = plan.layers.len() as u32;
let upload_half = |e: &Engine, src: BankTensorSrc, device_bf16: bool| -> Res<BankHalf> {
Ok(match src {
BankTensorSrc::F32(data) => BankHalf::F32(e.htod(&data)?),
BankTensorSrc::Nvfp4 {
codes,
scales,
macros,
..
} => BankHalf::Nvfp4 {
codes: e.htod_bytes(&codes)?,
scales: e.htod_bytes(&scales)?,
macros_dev: e.htod(¯os)?,
macros,
},
BankTensorSrc::Bf16(bytes) if device_bf16 => {
BankHalf::DeviceBf16(e.htod_bytes(&bytes)?)
}
BankTensorSrc::Bf16(bytes) => BankHalf::HostBf16(bytes),
})
};
for (index, bank_plan) in bank_plans {
let bank = bank_plan.read(&model)?;
let device_bf16 = index >= n_trunk;
let bank_e = if device_bf16 { draft_e.unwrap_or(e) } else { e };
parts.expert_banks.insert(
index,
ExpertBank {
gate: upload_half(bank_e, bank.gate, device_bf16)?,
up: upload_half(bank_e, bank.up, device_bf16)?,
down: upload_half(bank_e, bank.down, device_bf16)?,
},
);
}
drop(model);
for (index, bytes) in tables {
parts.ngram_tables.insert(index, NgramTable::Bf16(bytes));
}
Self::from_reference_weights_with(e, draft_e, &plan, &weights, parts)
}
}
struct GdnHalfW {
nk_h: usize,
nv_h: usize,
hk: usize,
hv: usize,
kernel: usize,
gate_activation: GdnGateActivation,
proj_b16: CudaSlice<u8>,
out_b16: CudaSlice<u8>, conv_w: CudaSlice<f32>, a: CudaSlice<f32>, dt: CudaSlice<f32>, norm: CudaSlice<f32>, }
struct QsaHalfW {
nh_h: usize,
nkv_h: usize,
hd: usize,
n_rot: usize,
rope_base: f32,
scale: f32,
proj_b16: CudaSlice<u8>,
wo_b16: CudaSlice<u8>, q_norm: Option<CudaSlice<f32>>,
k_norm: Option<CudaSlice<f32>>,
yarn: Option<YarnRopeW>,
}
enum MixerHalfW {
Gdn(GdnHalfW),
Qsa(QsaHalfW),
}
struct Nvfp4Half {
codes: CudaSlice<u8>,
scales: CudaSlice<u8>,
macros_dev: CudaSlice<f32>,
}
struct MoeHalfW {
gate1: Nvfp4Half,
up1: Nvfp4Half,
down1: Nvfp4Half,
shared_down0: CudaSlice<u8>, shared_down1: CudaSlice<u8>, shared_input_gate1: Option<CudaSlice<f32>>, shared_gu0_b16: CudaSlice<u8>,
shared_gu1_b16: CudaSlice<u8>,
}
struct Tp2LayerW {
attn_gate1: GateW,
mlp_gate1: GateW,
mixer0: MixerHalfW,
mixer1: MixerHalfW,
moe: MoeHalfW,
ple1: Option<PleW>,
place: LayerPlacement,
}
pub struct Tp2Shard {
layers: Vec<Tp2LayerW>,
exit_gate1: GateW,
lm_head1: CudaSlice<u8>, vsplit: usize,
stage0: [CudaSlice<f32>; 2], stage1: [CudaSlice<f32>; 2], stage0_raw: [u64; 2],
stage1_raw: [u64; 2],
ev0: [cudarc::driver::CudaEvent; 2], ev1: [cudarc::driver::CudaEvent; 2], }
enum MixerHalfState {
Gdn {
conv: CudaSlice<f32>, state: CudaSlice<f32>, },
Qsa {
kv: QsaKvStore,
},
}
struct Tp2LayerState {
m0: MixerHalfState,
m1: MixerHalfState,
ple1: Option<PleState>,
}
struct Tp2State {
ws1: StepPool,
layers: Vec<Tp2LayerState>,
graphs: Tp2Graphs,
pf_stage0: Option<[CudaSlice<f32>; 2]>, pf_stage1: Option<[CudaSlice<f32>; 2]>,
pf_stage0_raw: [u64; 2],
pf_stage1_raw: [u64; 2],
pf_rows: usize,
}
#[derive(Default)]
struct Tp2Graphs {
warm: bool,
a: [Vec<Option<GraphEntry>>; 2],
b: [Vec<Option<GraphEntry>>; 2],
c: [Vec<Option<GraphEntry>>; 2],
d: [Vec<Option<GraphEntry>>; 2],
exit: [Option<GraphEntry>; 2],
}
fn tp2_gdn_head_map(d: usize, nk: usize, nv: usize) -> Vec<usize> {
let nk_h = nk / 2;
let nv_h = nv / 2;
(0..nv_h)
.map(|j| (j / nk_h) * nk + (j % nk_h) + d * nk_h)
.collect()
}
fn gather_rows_host(src: &[f32], in_f: usize, rows: &[usize]) -> Vec<f32> {
let mut out = Vec::with_capacity(rows.len() * in_f);
for &r in rows {
out.extend_from_slice(&src[r * in_f..(r + 1) * in_f]);
}
out
}
fn gather_cols_host(
src: &[f32],
nrows: usize,
ncols: usize,
blocks: &[(usize, usize)],
) -> Vec<f32> {
let width: usize = blocks.iter().map(|&(_, l)| l).sum();
let mut out = Vec::with_capacity(nrows * width);
for r in 0..nrows {
for &(start, len) in blocks {
out.extend_from_slice(&src[r * ncols + start..r * ncols + start + len]);
}
}
out
}
fn need_twin(e: &Engine, data: &[f32], in_f: usize, what: &str) -> Res<CudaSlice<u8>> {
bf16_twin(e, data, in_f)?.ok_or_else(|| {
format!("qwen4exp_gpu tp2: {what} has no exact bf16 twin (in_f {in_f})").into()
})
}
fn launch_push(e: &Engine, src: &CudaSlice<f32>, dst_raw: u64, n: usize) -> Res<()> {
let f = e.func("q4e_push_f32");
let cfg = LaunchConfig::for_num_elems(n as u32);
let nl = n as i64;
let stream = e.gpu.stream();
let mut b = stream.launch_builder(&f);
b.arg(src).arg(&dst_raw).arg(&nl);
unsafe {
b.launch(cfg)?;
}
Ok(())
}
pub fn tp2_enable_p2p(e0: &Engine, e1: &Engine) -> Res<()> {
use cudarc::driver::sys;
for (src, dst) in [(e0, e1), (e1, e0)] {
let mut can = 0i32;
unsafe {
sys::cuDeviceCanAccessPeer(&mut can, src.ctx().cu_device(), dst.ctx().cu_device())
.result()?;
}
if can == 0 {
return Err(format!(
"qwen4exp_gpu tp2: dev{} cannot access dev{} over P2P",
src.ctx().ordinal(),
dst.ctx().ordinal()
)
.into());
}
src.ctx().bind_to_thread()?;
let rc = unsafe { sys::cuCtxEnablePeerAccess(dst.ctx().cu_ctx(), 0) };
use cudarc::driver::sys::cudaError_enum as E;
if rc != E::CUDA_SUCCESS && rc != E::CUDA_ERROR_PEER_ACCESS_ALREADY_ENABLED {
return Err(format!("qwen4exp_gpu tp2: cuCtxEnablePeerAccess failed: {rc:?}").into());
}
}
for (owner, accessor) in [(e0, e1), (e1, e0)] {
let device = cudarc::driver::result::device::get(owner.ctx().ordinal() as i32)?;
let mut pool: sys::CUmemoryPool = std::ptr::null_mut();
unsafe {
sys::cuDeviceGetDefaultMemPool(&mut pool, device).result()?;
}
let desc = sys::CUmemAccessDesc {
location: sys::CUmemLocation {
type_: sys::CUmemLocationType::CU_MEM_LOCATION_TYPE_DEVICE,
id: accessor.ctx().ordinal() as i32,
},
flags: sys::CUmemAccess_flags::CU_MEM_ACCESS_FLAGS_PROT_READWRITE,
};
let rc = unsafe { sys::cuMemPoolSetAccess(pool, &desc, 1) };
if rc != sys::cudaError_enum::CUDA_SUCCESS {
return Err(format!("qwen4exp_gpu tp2: cuMemPoolSetAccess failed: {rc:?}").into());
}
}
Ok(())
}
fn build_ple_replica(
e: &Engine,
weights: &ReferenceWeights,
prefix: &str,
ple_plan: &PleEmbeddingPlan,
streams: usize,
hidden: usize,
) -> Res<PleW> {
let embed_dim = ple_plan.embed_dim as usize;
let key_proj = expect(weights, &family_id(format!("{prefix}ple.key_proj.weight")))?;
let conv_w = expect(weights, &family_id(format!("{prefix}ple.conv1d.weight")))?;
let norm_slices = |name: &str| -> Res<Vec<CudaSlice<f32>>> {
let t = expect(weights, &family_id(format!("{prefix}ple.{name}.weight")))?;
split_rows(&t.data, streams, hidden, 1)
.into_iter()
.map(|v| e.htod(&v))
.collect::<Result<_, _>>()
};
let ints = |name: &str| -> Res<Vec<i64>> {
let t = expect(
weights,
&family_id(format!("{prefix}ple.ple_embedding.{name}")),
)?;
t.ints
.clone()
.ok_or_else(|| "qwen4exp_gpu: n-gram buffer must be I64".into())
};
Ok(PleW {
plan: *ple_plan,
key_proj: split_rows(&key_proj.data, streams, hidden, embed_dim)
.into_iter()
.map(|v| e.htod(&v))
.collect::<Result<_, _>>()?,
value_proj: upload(
e,
&expect(
weights,
&family_id(format!("{prefix}ple.value_proj.weight")),
)?,
)?,
norm_key: norm_slices("norm_key")?,
norm_query: norm_slices("norm_query")?,
norm_conv: norm_slices("norm_conv")?,
conv_w: split_rows(&conv_w.data, streams, hidden, ple_plan.conv_kernel as usize)
.into_iter()
.map(|v| e.htod(&v))
.collect::<Result<_, _>>()?,
multipliers: ints("layer_multipliers")?,
sizes: ints("ngram_heads_vocab_sizes")?,
offsets: ints("ngram_heads_offsets")?,
table: NgramTable::F32(Vec::new()), })
}
#[allow(clippy::too_many_arguments)]
fn build_gdn_half(
e: &Engine,
weights: &ReferenceWeights,
index: u32,
gdn: &GatedDeltaNetPlan,
hidden: usize,
d: usize,
) -> Res<GdnHalfW> {
let (nk, nv) = (gdn.key_heads as usize, gdn.value_heads as usize);
let (hk, hv) = (gdn.key_head_dim as usize, gdn.value_head_dim as usize);
if nk % 2 != 0 || nv % nk != 0 {
return Err(format!(
"qwen4exp_gpu tp2: GDN layer {index} nk {nk} / nv {nv} does not split by key-head halves"
)
.into());
}
let (nk_h, nv_h) = (nk / 2, nv / 2);
let head_map = tp2_gdn_head_map(d, nk, nv);
let qkv = expect(weights, &layer_id(index, LayerTensor::GdnQkv))?;
let z = expect(weights, &layer_id(index, LayerTensor::GdnGate))?;
let beta = expect(weights, &layer_id(index, LayerTensor::GdnBeta))?;
let alpha = expect(weights, &layer_id(index, LayerTensor::GdnAlpha))?;
let out = expect(weights, &layer_id(index, LayerTensor::GdnOutput))?;
let conv_w = expect(weights, &layer_id(index, LayerTensor::GdnConv1d))?;
let a = expect(weights, &layer_id(index, LayerTensor::GdnA))?;
let dt = expect(weights, &layer_id(index, LayerTensor::GdnDtBias))?;
let norm = expect(weights, &layer_id(index, LayerTensor::GdnNorm))?;
let kernel = gdn.conv_kernel as usize;
let mut qkv_rows: Vec<usize> = Vec::with_capacity(2 * nk_h * hk + nv_h * hv);
qkv_rows.extend(d * nk_h * hk..(d + 1) * nk_h * hk);
qkv_rows.extend(nk * hk + d * nk_h * hk..nk * hk + (d + 1) * nk_h * hk);
for &hm in &head_map {
qkv_rows.extend(2 * nk * hk + hm * hv..2 * nk * hk + (hm + 1) * hv);
}
let mut z_rows: Vec<usize> = Vec::with_capacity(nv_h * hv);
for &hm in &head_map {
z_rows.extend(hm * hv..(hm + 1) * hv);
}
let out_blocks: Vec<(usize, usize)> = head_map.iter().map(|&hm| (hm * hv, hv)).collect();
let qkv_c = gather_rows_host(&qkv.data, hidden, &qkv_rows);
let z_c = gather_rows_host(&z.data, hidden, &z_rows);
let beta_c = gather_rows_host(&beta.data, hidden, &head_map);
let alpha_c = gather_rows_host(&alpha.data, hidden, &head_map);
let out_c = gather_cols_host(&out.data, hidden, nv * hv, &out_blocks);
let conv_c = gather_rows_host(&conv_w.data, kernel, &qkv_rows);
let a_c: Vec<f32> = head_map.iter().map(|&hm| a.data[hm]).collect();
let dt_c: Vec<f32> = head_map.iter().map(|&hm| dt.data[hm]).collect();
Ok(GdnHalfW {
nk_h,
nv_h,
hk,
hv,
kernel,
gate_activation: gdn.gate_activation,
proj_b16: need_stack_twin(
e,
&[&qkv_c, &z_c, &beta_c, &alpha_c],
hidden,
"tp2 gdn proj half",
)?,
out_b16: need_twin(e, &out_c, nv_h * hv, "tp2 gdn out half")?,
conv_w: e.htod(&conv_c)?,
a: e.htod(&a_c)?,
dt: e.htod(&dt_c)?,
norm: e.htod(&norm.data)?,
})
}
#[allow(clippy::too_many_arguments)]
fn build_qsa_half(
e: &Engine,
weights: &ReferenceWeights,
index: u32,
attn: &FullAttentionPlan,
hidden: usize,
d: usize,
) -> Res<QsaHalfW> {
let nh = attn.query_heads as usize;
let nkv = attn.kv_heads as usize;
let hd = attn.key_head_dim as usize;
if nh % 2 != 0 || nkv % 2 != 0 || nh % nkv != 0 {
return Err(format!(
"qwen4exp_gpu tp2: QSA layer {index} heads {nh}/{nkv} do not split in halves"
)
.into());
}
let (nh_h, nkv_h) = (nh / 2, nkv / 2);
let wq = expect(weights, &layer_id(index, LayerTensor::Query))?;
let wk = expect(weights, &layer_id(index, LayerTensor::Key))?;
let wv = expect(weights, &layer_id(index, LayerTensor::Value))?;
let wo = expect(weights, &layer_id(index, LayerTensor::AttentionOutput))?;
let q_rows: Vec<usize> = (d * nh_h * 2 * hd..(d + 1) * nh_h * 2 * hd).collect();
let kv_rows: Vec<usize> = (d * nkv_h * hd..(d + 1) * nkv_h * hd).collect();
let wq_c = gather_rows_host(&wq.data, hidden, &q_rows);
let wk_c = gather_rows_host(&wk.data, hidden, &kv_rows);
let wv_c = gather_rows_host(&wv.data, hidden, &kv_rows);
let wo_c = gather_cols_host(&wo.data, hidden, nh * hd, &[(d * nh_h * hd, nh_h * hd)]);
let opt_norm = |tensor: LayerTensor| -> Res<Option<CudaSlice<f32>>> {
match weights.get(&layer_id(index, tensor)) {
Some(t) => Ok(Some(e.htod(&t.data)?)),
None => Ok(None),
}
};
let scale = match attn.scale {
memra_gguf::model_plan::AttentionScale::InverseSqrtKeyDim => 1.0 / (hd as f32).sqrt(),
memra_gguf::model_plan::AttentionScale::Fixed(scale) => scale,
};
Ok(QsaHalfW {
nh_h,
nkv_h,
hd,
n_rot: attn.rope.dimensions as usize,
rope_base: attn.rope.base,
scale,
proj_b16: need_stack_twin(e, &[&wq_c, &wk_c, &wv_c], hidden, "tp2 qsa proj half")?,
wo_b16: need_twin(e, &wo_c, nh_h * hd, "tp2 qsa o half")?,
q_norm: opt_norm(LayerTensor::QueryNorm)?,
k_norm: opt_norm(LayerTensor::KeyNorm)?,
yarn: build_yarn(e, &attn.rope, None, index)?,
})
}
pub fn build_tp2_shard(e0: &Engine, e1: &Engine, ckpt: &LoadedCheckpoint) -> Res<Tp2Shard> {
let plan = &ckpt.plan;
let weights = &ckpt.weights;
let hidden = plan.hidden_size as usize;
let vocab = plan.vocab_size as usize;
if vocab % 2 != 0 {
return Err("qwen4exp_gpu tp2: odd vocab".into());
}
let mixer_plan = plan
.exit_mixer
.ok_or("qwen4exp_gpu tp2: missing exit mixer")?;
let streams = mixer_plan.streams as usize;
let rank = mixer_plan.bottleneck_rank as usize;
let plan_experts = plan
.layers
.iter()
.find_map(|l| match &l.mlp {
MlpPlan::Moe(m) => Some(m.expert_count as usize),
_ => None,
})
.ok_or("qwen4exp_gpu tp2: no MoE layer in the plan")?;
let placement = match Tp2Placement::from_env(plan_experts)? {
Some(p) => p,
None => Tp2Placement::even(plan_experts),
};
println!(
"# tp2-placement\tstrategy={}\tentry_rank={}\texperts={plan_experts}\tsource={}",
placement.strategy(),
placement.entry_rank(),
placement.source()
);
let mut layers = Vec::with_capacity(plan.layers.len());
for layer in &plan.layers {
let prefix = format!("trunk.layers.{}.", layer.index);
let _g1 = e1.gpu.enter_main()?;
let attn_gate1 = load_gate(
e1,
weights,
&prefix,
"attn_hyper_connection.",
streams,
hidden,
rank,
true,
)?;
let mlp_gate1 = load_gate(
e1,
weights,
&prefix,
"mlp_hyper_connection.",
streams,
hidden,
rank,
true,
)?;
let ple1 = match layer.ple.as_ref() {
None => None,
Some(ple_plan) => Some(build_ple_replica(
e1, weights, &prefix, ple_plan, streams, hidden,
)?),
};
drop(_g1);
let (mixer0, mixer1) = match &layer.attention {
AttentionPlan::GatedDeltaNet(gdn) => {
let _g0 = e0.gpu.enter_main()?;
let m0 = MixerHalfW::Gdn(build_gdn_half(e0, weights, layer.index, gdn, hidden, 0)?);
drop(_g0);
let _g1 = e1.gpu.enter_main()?;
let m1 = MixerHalfW::Gdn(build_gdn_half(e1, weights, layer.index, gdn, hidden, 1)?);
(m0, m1)
}
AttentionPlan::Full(attn) => {
let _g0 = e0.gpu.enter_main()?;
let m0 =
MixerHalfW::Qsa(build_qsa_half(e0, weights, layer.index, attn, hidden, 0)?);
drop(_g0);
let _g1 = e1.gpu.enter_main()?;
let m1 =
MixerHalfW::Qsa(build_qsa_half(e1, weights, layer.index, attn, hidden, 1)?);
(m0, m1)
}
other => {
return Err(format!("qwen4exp_gpu tp2: unsupported mixer {other:?}").into());
}
};
let MlpPlan::Moe(moe_plan) = &layer.mlp else {
return Err("qwen4exp_gpu tp2: non-MoE layer".into());
};
let experts = moe_plan.expert_count as usize;
let ff = moe_plan.expert_intermediate_size as usize;
if experts % 2 != 0 {
return Err("qwen4exp_gpu tp2: odd expert count".into());
}
let bank = ckpt.read_bank(layer.index)?;
let bank = &bank;
let place = placement.layer(layer.index, experts)?;
let upper =
|src: &BankTensorSrc, out_f: usize, in_f: usize, what: &str| -> Res<Nvfp4Half> {
let BankTensorSrc::Nvfp4 {
codes,
scales,
macros,
..
} = src
else {
return Err(format!("qwen4exp_gpu tp2: {what} bank is not NVFP4").into());
};
let wbytes = out_f * in_f / 2;
let sbytes = out_f * in_f / 16;
let need_codes = place.card1.len() * wbytes;
let need_scales = place.card1.len() * sbytes;
if codes.len() < experts * wbytes || scales.len() < experts * sbytes {
return Err(format!(
"qwen4exp_gpu tp2: {what} bank is {} code / {} scale bytes, too \
small for {experts} experts x ({wbytes}, {sbytes})",
codes.len(),
scales.len()
)
.into());
}
let mut gcodes = Vec::with_capacity(need_codes);
let mut gscales = Vec::with_capacity(need_scales);
let mut gmacros = Vec::with_capacity(place.card1.len());
for &eid in &place.card1 {
let e = eid as usize;
gcodes.extend_from_slice(&codes[e * wbytes..(e + 1) * wbytes]);
gscales.extend_from_slice(&scales[e * sbytes..(e + 1) * sbytes]);
gmacros.push(macros[e]);
}
Ok(Nvfp4Half {
codes: e1.htod_bytes(&gcodes)?,
scales: e1.htod_bytes(&gscales)?,
macros_dev: e1.htod(&gmacros)?,
})
};
let shared = moe_plan
.shared
.as_ref()
.ok_or("qwen4exp_gpu tp2: missing shared expert")?;
let sff = shared.intermediate_size as usize;
if sff % 2 != 0 {
return Err("qwen4exp_gpu tp2: odd shared ff".into());
}
let sffh = sff / 2;
let sh_gate = expect(weights, &layer_id(layer.index, LayerTensor::SharedMlpGate))?;
let sh_up = expect(weights, &layer_id(layer.index, LayerTensor::SharedMlpUp))?;
let sh_down = expect(weights, &layer_id(layer.index, LayerTensor::SharedMlpDown))?;
let sh_ig = if shared.gated {
Some(expect(
weights,
&layer_id(layer.index, LayerTensor::SharedMlpInputGate),
)?)
} else {
None
};
let moe = {
let _g1 = e1.gpu.enter_main()?;
let gate1 = upper(&bank.gate, ff, hidden, "gate")?;
let up1 = upper(&bank.up, ff, hidden, "up")?;
let down1 = upper(&bank.down, hidden, ff, "down")?;
let shared_gu1_b16 = need_stack_twin(
e1,
&[&sh_gate.data[sffh * hidden..], &sh_up.data[sffh * hidden..]],
hidden,
"tp2 shared gate/up (card1)",
)?;
let down1_c = gather_cols_host(&sh_down.data, hidden, sff, &[(sffh, sffh)]);
let shared_down1 = need_twin(e1, &down1_c, sffh, "tp2 shared down (card1)")?;
let shared_input_gate1 = match sh_ig.as_ref() {
Some(t) => Some(e1.htod(&t.data)?),
None => None,
};
drop(_g1);
let _g0 = e0.gpu.enter_main()?;
let down0_c = gather_cols_host(&sh_down.data, hidden, sff, &[(0, sffh)]);
let shared_down0 = need_twin(e0, &down0_c, sffh, "tp2 shared down (card0)")?;
let shared_gu0_b16 = need_stack_twin(
e0,
&[&sh_gate.data[..sffh * hidden], &sh_up.data[..sffh * hidden]],
hidden,
"tp2 shared gate/up (card0)",
)?;
MoeHalfW {
gate1,
up1,
down1,
shared_down0,
shared_down1,
shared_input_gate1,
shared_gu0_b16,
shared_gu1_b16,
}
};
layers.push(Tp2LayerW {
attn_gate1,
mlp_gate1,
mixer0,
mixer1,
moe,
ple1,
place,
});
}
let _g1 = e1.gpu.enter_main()?;
let exit_gate1 = load_gate(
e1,
weights,
"trunk.hyper_connection_mixer.",
"",
streams,
hidden,
rank,
false,
)?;
let vsplit = vocab / 2;
let head = match weights.get(&TensorId::OutputProjection) {
Some(t) => &t.data,
None => &expect(weights, &TensorId::TokenEmbedding)?.data.clone(),
};
let lm_head1 = need_twin(
e1,
&head[vsplit * hidden..],
hidden,
"tp2 lm_head upper half",
)?;
let stage1 = [e1.zeros(hidden)?, e1.zeros(hidden)?];
let ev1 = [e1.ctx().new_event(None)?, e1.ctx().new_event(None)?];
let stage1_raw = {
let s = e1.gpu.stream();
[stage1[0].device_ptr(&s).0, stage1[1].device_ptr(&s).0]
};
drop(_g1);
let _g0 = e0.gpu.enter_main()?;
let stage0 = [e0.zeros(hidden)?, e0.zeros(hidden)?];
let ev0 = [e0.ctx().new_event(None)?, e0.ctx().new_event(None)?];
let stage0_raw = {
let s = e0.gpu.stream();
[stage0[0].device_ptr(&s).0, stage0[1].device_ptr(&s).0]
};
Ok(Tp2Shard {
layers,
exit_gate1,
lm_head1,
vsplit,
stage0,
stage1,
stage0_raw,
stage1_raw,
ev0,
ev1,
})
}
impl Qwen4ExpGpu {
pub fn load_from_dir_tp2(
e0: &Engine,
e1: &Engine,
dir: &std::path::Path,
opts: LoadOptions,
) -> Res<(Self, Tp2Shard)> {
tp2_enable_p2p(e0, e1)?;
let checkpoint = read_checkpoint_with(dir, opts)?;
let shard = build_tp2_shard(e0, e1, &checkpoint)?;
let model = Self::from_loaded_checkpoint(e0, checkpoint)?;
Ok((model, shard))
}
fn tp2_migrate(
&self,
e0: &Engine,
e1: &Engine,
_shard: &Tp2Shard,
state: &mut Qwen4ExpState,
) -> Res<()> {
let cap = state.capacity;
let pos = state.pos;
let mut tlayers = Vec::with_capacity(self.layers.len());
for (layer, lstate) in self.layers.iter().zip(state.layers.iter_mut()) {
let (m0, m1) = match (&layer.mixer, &mut lstate.mixer) {
(
MixerW::Gdn(gdn),
MixerState::Gdn {
conv,
state: gstate,
},
) => {
let p = &gdn.plan;
let (nk, nv) = (p.key_heads as usize, p.value_heads as usize);
let (hk, hv) = (p.key_head_dim as usize, p.value_head_dim as usize);
let (nk_h, nv_h) = (nk / 2, nv / 2);
let pad = p.conv_kernel as usize - 1;
let conv_dim = 2 * nk * hk + nv * hv;
let conv_dim_h = 2 * nk_h * hk + nv_h * hv;
let state_host = {
let _g = e0.gpu.enter_main()?;
e0.dtoh(gstate)?
};
let conv_host = {
let _g = e0.gpu.enter_main()?;
e0.dtoh(conv)?
};
let mut halves = Vec::with_capacity(2);
for d in 0..2 {
let head_map = tp2_gdn_head_map(d, nk, nv);
let state_c = gather_rows_host(&state_host, hv * hk, &head_map);
let mut blocks: Vec<(usize, usize)> = vec![
(d * nk_h * hk, nk_h * hk),
(nk * hk + d * nk_h * hk, nk_h * hk),
];
blocks.extend(head_map.iter().map(|&hm| (2 * nk * hk + hm * hv, hv)));
let conv_c = gather_cols_host(&conv_host, pad, conv_dim, &blocks);
let e = if d == 0 { e0 } else { e1 };
let _g = e.gpu.enter_main()?;
let state_dev = e.htod(&state_c)?;
let conv_dev = e.htod(&conv_c)?;
debug_assert_eq!(conv_c.len(), pad * conv_dim_h);
halves.push(MixerHalfState::Gdn {
conv: conv_dev,
state: state_dev,
});
}
let m1 = halves.pop().expect("two halves");
let m0 = halves.pop().expect("two halves");
(m0, m1)
}
(MixerW::Qsa(qsa), MixerState::Qsa { kv, .. }) => {
let nkv = qsa.attn.kv_heads as usize;
let hd = qsa.attn.key_head_dim as usize;
let nkv_h = nkv / 2;
let mut halves = Vec::with_capacity(2);
match &*kv {
QsaKvStore::F32 { k, v } => {
let (k_host, v_host) = {
let _g = e0.gpu.enter_main()?;
(
e0.dtoh_view(&k.slice(0..pos * nkv * hd))?,
e0.dtoh_view(&v.slice(0..pos * nkv * hd))?,
)
};
for d in 0..2 {
let block = [(d * nkv_h * hd, nkv_h * hd)];
let k_c = gather_cols_host(&k_host, pos, nkv * hd, &block);
let v_c = gather_cols_host(&v_host, pos, nkv * hd, &block);
let e = if d == 0 { e0 } else { e1 };
let _g = e.gpu.enter_main()?;
let mut k_dev = e.zeros(cap * nkv_h * hd)?;
let mut v_dev = e.zeros(cap * nkv_h * hd)?;
if pos > 0 {
let mut kv_view = k_dev.slice_mut(0..pos * nkv_h * hd);
e.gpu.stream().memcpy_htod(&k_c, &mut kv_view)?;
let mut vv_view = v_dev.slice_mut(0..pos * nkv_h * hd);
e.gpu.stream().memcpy_htod(&v_c, &mut vv_view)?;
}
halves.push(MixerHalfState::Qsa {
kv: QsaKvStore::F32 { k: k_dev, v: v_dev },
});
}
}
QsaKvStore::Q8Q5 { k, v } => {
if hd % 32 != 0 {
return Err("qwen4exp_gpu tp2: quantized halves need \
hd % 32 == 0 (byte-aligned head blocks)"
.into());
}
let (krb, vrb) = (q8_row_bytes(nkv * hd), q5_row_bytes(nkv * hd));
let (krb_h, vrb_h) =
(q8_row_bytes(nkv_h * hd), q5_row_bytes(nkv_h * hd));
let (k_host, v_host) = {
let _g = e0.gpu.enter_main()?;
(
e0.dtoh_u8_view(&k.slice(0..pos * krb))?,
e0.dtoh_u8_view(&v.slice(0..pos * vrb))?,
)
};
for d in 0..2 {
let mut k_c = Vec::with_capacity(pos * krb_h);
let mut v_c = Vec::with_capacity(pos * vrb_h);
for r in 0..pos {
let ko = r * krb + d * krb_h;
k_c.extend_from_slice(&k_host[ko..ko + krb_h]);
let vo = r * vrb + d * vrb_h;
v_c.extend_from_slice(&v_host[vo..vo + vrb_h]);
}
let e = if d == 0 { e0 } else { e1 };
let _g = e.gpu.enter_main()?;
let mut k_dev = e.alloc_u8(cap * krb_h)?;
let mut v_dev = e.alloc_u8(cap * vrb_h)?;
if pos > 0 {
let mut kv_view = k_dev.slice_mut(0..pos * krb_h);
e.gpu.stream().memcpy_htod(&k_c, &mut kv_view)?;
let mut vv_view = v_dev.slice_mut(0..pos * vrb_h);
e.gpu.stream().memcpy_htod(&v_c, &mut vv_view)?;
}
halves.push(MixerHalfState::Qsa {
kv: QsaKvStore::Q8Q5 { k: k_dev, v: v_dev },
});
}
}
}
{
let _g = e0.gpu.enter_main()?;
*kv = match &*kv {
QsaKvStore::F32 { .. } => QsaKvStore::F32 {
k: e0.zeros(1)?,
v: e0.zeros(1)?,
},
QsaKvStore::Q8Q5 { .. } => QsaKvStore::Q8Q5 {
k: e0.alloc_u8(34)?,
v: e0.alloc_u8(24)?,
},
};
}
let m1 = halves.pop().expect("two halves");
let m0 = halves.pop().expect("two halves");
(m0, m1)
}
_ => return Err("qwen4exp_gpu tp2: layer/state mixer mismatch".into()),
};
let ple1 = match lstate.ple.as_ref() {
None => None,
Some(ps) => {
let mut conv_hist = Vec::with_capacity(ps.conv_hist.len());
for h in &ps.conv_hist {
let host = {
let _g = e0.gpu.enter_main()?;
e0.dtoh(h)?
};
let _g = e1.gpu.enter_main()?;
conv_hist.push(e1.htod(&host)?);
}
Some(PleState {
conv_hist,
ngram_ids: Vec::new(),
ngram_history: Vec::new(),
ngram_last_eos: -1,
})
}
};
tlayers.push(Tp2LayerState { m0, m1, ple1 });
}
state.tp2 = Some(Tp2State {
ws1: StepPool::default(),
layers: tlayers,
graphs: Tp2Graphs::default(),
pf_stage0: None,
pf_stage1: None,
pf_stage0_raw: [0; 2],
pf_stage1_raw: [0; 2],
pf_rows: 0,
});
Ok(())
}
#[allow(clippy::too_many_arguments)]
fn gdn_forward_half(
&self,
e: &Engine,
ws: &mut StepPool,
eps: f32,
h: &GdnHalfW,
mixed: &CudaSlice<f32>,
hstate: &mut MixerHalfState,
t: usize,
) -> Res<CudaSlice<f32>> {
let MixerHalfState::Gdn { conv, state } = hstate else {
return Err("qwen4exp_gpu tp2: GDN half bound to non-GDN state".into());
};
let hidden = self.hidden;
let (nk, nv, hk, hv) = (h.nk_h, h.nv_h, h.hk, h.hv);
let kernel = h.kernel;
let pad = kernel - 1;
let conv_dim = 2 * nk * hk + nv * hv;
let mut qkv = ws.take_f32(e, "gdn.qkv", t * conv_dim, 0)?;
let mut z = ws.take_f32(e, "gdn.z", t * nv * hv, 0)?;
let mut beta_raw = ws.take_f32(e, "gdn.beta", t * nv, 0)?;
let mut alpha = ws.take_f32(e, "gdn.alpha", t * nv, 0)?;
if t == 1 && proj_stack_on() {
launch_qmatvec_bf16w_multi4(
e,
&h.proj_b16,
mixed,
&[
(&qkv, conv_dim),
(&z, nv * hv),
(&beta_raw, nv),
(&alpha, nv),
],
hidden,
)?;
} else {
launch_qmatvec_bf16w_off(e, &h.proj_b16, 0, mixed, &mut qkv, hidden, conv_dim, t)?;
launch_qmatvec_bf16w_off(e, &h.proj_b16, conv_dim, mixed, &mut z, hidden, nv * hv, t)?;
launch_qmatvec_bf16w_off(
e,
&h.proj_b16,
conv_dim + nv * hv,
mixed,
&mut beta_raw,
hidden,
nv,
t,
)?;
launch_qmatvec_bf16w_off(
e,
&h.proj_b16,
conv_dim + nv * hv + nv,
mixed,
&mut alpha,
hidden,
nv,
t,
)?;
}
let mut g_log = ws.take_f32(e, "gdn.glog", t * nv, 0)?;
e.gdn_glog_v(&alpha.slice(0..t * nv), &h.dt, &h.a, &mut g_log, nv, t)?;
ws.put_f32("gdn.alpha", alpha);
let mut conv_out = ws.take_f32(e, "gdn.conv_out", t * conv_dim, 0)?;
launch_dwconv(
e,
&qkv,
conv,
&h.conv_w,
&mut conv_out,
t,
pad,
conv_dim,
kernel,
1,
1,
)?;
let mut o = ws.take_f32(e, "gdn.o", t * nv * hv, 0)?;
let scale = 1.0 / (hk as f32).sqrt();
if t == 1 && gdn_step_on() && hk % 32 == 0 && hk <= 1024 {
launch_gdn_scan_step(
e, &conv_out, &g_log, &beta_raw, state, &mut o, nk, nv, hk, hv, scale, eps,
)?;
} else {
launch_gdn_scan(
e, &conv_out, &g_log, &beta_raw, state, &mut o, nk, nv, hk, hv, t, scale, eps,
)?;
}
ws.put_f32("gdn.conv_out", conv_out);
if t >= pad {
e.copy_range_into(conv, 0, &qkv, (t - pad) * conv_dim, pad * conv_dim)?;
} else {
let keep = pad - t;
let mut tmp = ws.take_f32(e, "gdn.tmp", keep * conv_dim, 0)?;
e.copy_range_into(&mut tmp, 0, conv, t * conv_dim, keep * conv_dim)?;
e.copy_range_into(conv, 0, &tmp, 0, keep * conv_dim)?;
e.copy_range_into(conv, keep * conv_dim, &qkv, 0, t * conv_dim)?;
ws.put_f32("gdn.tmp", tmp);
}
ws.put_f32("gdn.qkv", qkv);
ws.put_f32("gdn.beta", beta_raw);
ws.put_f32("gdn.glog", g_log);
let mut gated = ws.take_f32(e, "gdn.gated", t * nv * hv, 0)?;
match h.gate_activation {
GdnGateActivation::Sigmoid if gdn_fuse_on() => {
launch_rms_sigmul(e, &o, &h.norm, &z, &mut gated, hv, t * nv, eps)?;
}
GdnGateActivation::Sigmoid => {
let mut normed = ws.take_f32(e, "gdn.normed", t * nv * hv, 0)?;
e.rms_norm(&o, &h.norm, &mut normed, hv, t * nv, eps)?;
let mut sg = ws.take_f32(e, "gdn.sg", t * nv * hv, 0)?;
e.sigmoid(&z, &mut sg, t * nv * hv)?;
e.mul(&normed, &sg, &mut gated, t * nv * hv)?;
ws.put_f32("gdn.sg", sg);
ws.put_f32("gdn.normed", normed);
}
GdnGateActivation::Silu => {
let mut normed = ws.take_f32(e, "gdn.normed", t * nv * hv, 0)?;
e.rms_norm(&o, &h.norm, &mut normed, hv, t * nv, eps)?;
e.silu_mul(&z, &normed, &mut gated, t * nv * hv)?;
ws.put_f32("gdn.normed", normed);
}
}
let mut partial = ws.take_f32(e, "mixer.out", t * hidden, 0)?;
launch_qmatvec_bf16w(
e,
&h.out_b16,
&gated,
&mut partial,
nv * hv,
hidden,
t,
1,
0,
0,
nv * hv,
0,
)?;
ws.put_f32("gdn.gated", gated);
ws.put_f32("gdn.z", z);
ws.put_f32("gdn.o", o);
Ok(partial)
}
#[allow(clippy::too_many_arguments)]
fn qsa_half_proj(
&self,
e: &Engine,
ws: &mut StepPool,
eps: f32,
h: &QsaHalfW,
mixed: &CudaSlice<f32>,
hstate: &mut MixerHalfState,
base_pos: usize,
t: usize,
) -> Res<(CudaSlice<f32>, CudaSlice<f32>)> {
let MixerHalfState::Qsa { kv } = hstate else {
return Err("qwen4exp_gpu tp2: QSA half bound to non-QSA state".into());
};
let hidden = self.hidden;
let (nh, nkv, hd) = (h.nh_h, h.nkv_h, h.hd);
let mut q_fused = ws.take_f32(e, "qsa.qf", t * 2 * nh * hd, 0)?;
let mut k_new = ws.take_f32(e, "qsa.k", t * nkv * hd, 0)?;
let mut v_new = ws.take_f32(e, "qsa.v", t * nkv * hd, 0)?;
if t == 1 && proj_stack_on() {
launch_qmatvec_bf16w_multi4(
e,
&h.proj_b16,
mixed,
&[
(&q_fused, 2 * nh * hd),
(&k_new, nkv * hd),
(&v_new, nkv * hd),
],
hidden,
)?;
} else {
launch_qmatvec_bf16w_off(
e,
&h.proj_b16,
0,
mixed,
&mut q_fused,
hidden,
2 * nh * hd,
t,
)?;
launch_qmatvec_bf16w_off(
e,
&h.proj_b16,
2 * nh * hd,
mixed,
&mut k_new,
hidden,
nkv * hd,
t,
)?;
launch_qmatvec_bf16w_off(
e,
&h.proj_b16,
2 * nh * hd + nkv * hd,
mixed,
&mut v_new,
hidden,
nkv * hd,
t,
)?;
}
let mut q = ws.take_f32(e, "qsa.q", t * nh * hd, 0)?;
let mut gate = ws.take_f32(e, "qsa.gate", t * nh * hd, 0)?;
e.q_gate_split(&q_fused, &mut q, &mut gate, hd, nh, t)?;
ws.put_f32("qsa.qf", q_fused);
let mut q = if let Some(norm) = h.q_norm.as_ref() {
let mut dst = ws.take_f32(e, "qsa.qn", t * nh * hd, 0)?;
e.rms_norm(&q, norm, &mut dst, hd, t * nh, eps)?;
ws.put_f32("qsa.q", q);
dst
} else {
q
};
let mut k_new = if let Some(norm) = h.k_norm.as_ref() {
let mut dst = ws.take_f32(e, "qsa.kn", t * nkv * hd, 0)?;
e.rms_norm(&k_new, norm, &mut dst, hd, t * nkv, eps)?;
ws.put_f32("qsa.k", k_new);
dst
} else {
k_new
};
let positions: Vec<i32> = (0..t).map(|i| (base_pos + i) as i32).collect();
let pos_dev = ws.take_i32(e, "qsa.pos", &positions, 0)?;
if let Some(yarn) = h.yarn.as_ref() {
e.rope_neox_ffm(
&mut q,
&pos_dev,
hd,
h.n_rot,
nh,
t,
h.rope_base,
1.0,
&yarn.ff,
yarn.mscale,
)?;
e.rope_neox_ffm(
&mut k_new,
&pos_dev,
hd,
h.n_rot,
nkv,
t,
h.rope_base,
1.0,
&yarn.ff,
yarn.mscale,
)?;
} else {
e.rope_neox(&mut q, &pos_dev, hd, h.n_rot, nh, t, h.rope_base, 1.0)?;
e.rope_neox(&mut k_new, &pos_dev, hd, h.n_rot, nkv, t, h.rope_base, 1.0)?;
}
ws.put_i32("qsa.pos", pos_dev);
match kv {
QsaKvStore::F32 { k, v } => {
e.copy_range_into(k, base_pos * nkv * hd, &k_new, 0, t * nkv * hd)?;
e.copy_range_into(v, base_pos * nkv * hd, &v_new, 0, t * nkv * hd)?;
}
QsaKvStore::Q8Q5 { k, v } => {
launch_q4e_kv_append(e, &k_new, &v_new, k, v, base_pos, t, nkv * hd)?;
}
}
ws.put_f32(
if h.k_norm.is_some() {
"qsa.kn"
} else {
"qsa.k"
},
k_new,
);
ws.put_f32("qsa.v", v_new);
Ok((q, gate))
}
#[allow(clippy::too_many_arguments)]
fn qsa_half_attend(
&self,
e: &Engine,
ws: &mut StepPool,
h: &QsaHalfW,
hstate: &MixerHalfState,
q: CudaSlice<f32>,
gate: CudaSlice<f32>,
pos_dev: &CudaSlice<i32>,
meta_dev: &CudaSlice<i32>,
max_count: usize,
t: usize,
t_kv: usize,
) -> Res<CudaSlice<f32>> {
let MixerHalfState::Qsa { kv } = hstate else {
return Err("qwen4exp_gpu tp2: QSA half bound to non-QSA state".into());
};
let hidden = self.hidden;
let (nh, nkv, hd) = (h.nh_h, h.nkv_h, h.hd);
let mut attended = ws.take_f32(e, "qsa.att", t * nh * hd, 0)?;
match kv {
QsaKvStore::F32 { k, v } => {
let k_view = k.slice(0..t_kv * nkv * hd);
let v_view = v.slice(0..t_kv * nkv * hd);
launch_sdpa_blocklist(
e,
&q,
&k_view,
&v_view,
&mut attended,
pos_dev,
meta_dev,
hd,
nh,
nkv,
t,
max_count,
h.scale,
)?;
}
QsaKvStore::Q8Q5 { k, v } => {
launch_q4e_sdpa_blocklist_q8q5(
e,
&q,
k,
v,
&mut attended,
pos_dev,
meta_dev,
hd,
nh,
nkv,
t,
max_count,
h.scale,
)?;
}
}
ws.put_f32(
if h.q_norm.is_some() {
"qsa.qn"
} else {
"qsa.q"
},
q,
);
let mut sg = ws.take_f32(e, "qsa.sg", t * nh * hd, 0)?;
e.sigmoid(&gate, &mut sg, t * nh * hd)?;
let mut gated = ws.take_f32(e, "qsa.gated", t * nh * hd, 0)?;
e.mul(&attended, &sg, &mut gated, t * nh * hd)?;
let mut partial = ws.take_f32(e, "mixer.out", t * hidden, 0)?;
launch_qmatvec_bf16w(
e,
&h.wo_b16,
&gated,
&mut partial,
nh * hd,
hidden,
t,
1,
0,
0,
nh * hd,
0,
)?;
ws.put_f32("qsa.sg", sg);
ws.put_f32("qsa.gated", gated);
ws.put_f32("qsa.att", attended);
ws.put_f32("qsa.gate", gate);
Ok(partial)
}
#[allow(dead_code)]
#[allow(clippy::too_many_arguments)]
fn qsa_indexer_mask(
&self,
e: &Engine,
ws: &mut StepPool,
qsa: &QsaW,
eps: f32,
mixed: &CudaSlice<f32>,
raw_keys: &mut IdxRawCache,
pooled_keys: &mut Vec<f32>,
base_pos: usize,
) -> Res<Vec<u8>> {
let overlay = &qsa.overlay;
let idx_dim = overlay.head_dim as usize;
let qk_width = (overlay.query_heads as usize + overlay.kv_heads as usize) * idx_dim;
let hidden = self.hidden;
let mut idx_proj = ws.take_f32(e, "qsa.idxp", qk_width, 0)?;
e.linear_device_into(mixed, &qsa.idx_proj, &mut idx_proj, 1, hidden, qk_width)?;
let rows = e.dtoh_view(&idx_proj.slice(0..qk_width))?;
ws.put_f32("qsa.idxp", idx_proj);
raw_keys.append_rows_f32(
&rows[overlay.query_heads as usize * idx_dim..qk_width],
1,
idx_dim,
);
indexer_mask_rows(
overlay,
qsa.attn.rope.base,
qsa.yarn.as_ref().map(|y| (y.ff_host.as_slice(), y.mscale)),
eps,
&qsa.idx_q_norm,
&qsa.idx_k_norm,
&rows,
raw_keys,
pooled_keys,
base_pos,
1,
base_pos + 1,
0,
)
}
#[allow(clippy::too_many_arguments)]
fn tp2_shared_half(
&self,
e: &Engine,
ws: &mut StepPool,
gu_b16: &CudaSlice<u8>,
down_b16: &CudaSlice<u8>,
input_gate: Option<&CudaSlice<f32>>,
mixed: &CudaSlice<f32>,
sffh: usize,
t: usize,
) -> Res<(CudaSlice<f32>, Option<CudaSlice<f32>>)> {
let hidden = self.hidden;
let mut sh_gate = ws.take_f32(e, "moe.sh_gate", t * sffh, 0)?;
let mut sh_up = ws.take_f32(e, "moe.sh_up", t * sffh, 0)?;
if t == 1 && proj_stack_on() {
launch_qmatvec_bf16w_multi4(
e,
gu_b16,
mixed,
&[(&sh_gate, sffh), (&sh_up, sffh)],
hidden,
)?;
} else {
launch_qmatvec_bf16w_off(e, gu_b16, 0, mixed, &mut sh_gate, hidden, sffh, t)?;
launch_qmatvec_bf16w_off(e, gu_b16, sffh, mixed, &mut sh_up, hidden, sffh, t)?;
}
let mut act = ws.take_f32(e, "moe.sh_act", t * sffh, 0)?;
e.silu_mul(&sh_gate, &sh_up, &mut act, t * sffh)?;
let mut shared = ws.take_f32(e, "moe.sh_down", t * hidden, 0)?;
launch_qmatvec_bf16w(
e,
down_b16,
&act,
&mut shared,
sffh,
hidden,
t,
1,
0,
0,
sffh,
0,
)?;
let g = match input_gate {
Some(w) => {
let mut g = ws.take_f32(e, "moe.g", t, 0)?;
e.sigmoid_dot_rows_into(mixed, w, &mut g, hidden, t)?;
Some(g)
}
None => None,
};
ws.put_f32("moe.sh_gate", sh_gate);
ws.put_f32("moe.sh_up", sh_up);
ws.put_f32("moe.sh_act", act);
Ok((shared, g))
}
}
impl Qwen4ExpGpu {
pub fn decode_step_tp2(
&self,
e0: &Engine,
e1: &Engine,
shard: &Tp2Shard,
token: u32,
state: &mut Qwen4ExpState,
) -> Res<Vec<f32>> {
if !trunk_bf16_on() || !hc_fused_gate_on() {
return Err(
"qwen4exp_gpu tp2: requires set_trunk_bf16(true) and set_hc_fused_gate(true) \
(replicated compute must run deterministic kernels)"
.into(),
);
}
if state.pos + 1 > state.capacity {
return Err("qwen4exp_gpu: state capacity exceeded".into());
}
if state.tp2.is_none() {
self.tp2_migrate(e0, e1, shard, state)?;
state.graphs = StepGraphs::default();
}
let hidden = self.hidden;
let vocab = self.vocab;
let vsplit = shard.vsplit;
let base_pos = state.pos;
let reserve = state.reserve;
state.tokens.push(token);
let Qwen4ExpState {
ref tokens,
ws: ref mut ws0,
ref mut tp2,
layers: ref mut lstates,
..
} = *state;
let Tp2State {
ws1,
layers: tlayers,
graphs: tgraphs,
..
} = tp2.as_mut().expect("migrated above");
let cap = reserve.max(1);
let token_us = token as usize;
if token_us >= vocab {
return Err(format!("qwen4exp_gpu: token {token_us} out of range").into());
}
let embedded = &self.embed_host[token_us * hidden..(token_us + 1) * hidden];
let mut planes1: Vec<CudaSlice<f32>> = Vec::with_capacity(self.streams);
let ptrs1 = {
let _g = e1.gpu.enter_main()?;
let embedded_dev = ws1.take_f32_h2d(e1, "entry.embed", embedded, cap * hidden)?;
for s in 0..self.streams {
let mut plane = ws1.take_f32(e1, PLANE_SLOTS[s], hidden, cap * hidden)?;
e1.copy_into(&mut plane, 0, &embedded_dev, hidden)?;
planes1.push(plane);
}
ws1.put_f32("entry.embed", embedded_dev);
let ptr_vals: Vec<u64> = {
let stream = e1.gpu.stream();
planes1.iter().map(|p| p.device_ptr(&stream).0).collect()
};
ws1.take_u64_h2d(e1, "hc.ptrs", &ptr_vals, 0)?
};
let mut planes0: Vec<CudaSlice<f32>> = Vec::with_capacity(self.streams);
let ptrs0 = {
let _g = e0.gpu.enter_main()?;
let embedded_dev = ws0.take_f32_h2d(e0, "entry.embed", embedded, cap * hidden)?;
for s in 0..self.streams {
let mut plane = ws0.take_f32(e0, PLANE_SLOTS[s], hidden, cap * hidden)?;
e0.copy_into(&mut plane, 0, &embedded_dev, hidden)?;
planes0.push(plane);
}
ws0.put_f32("entry.embed", embedded_dev);
let ptr_vals: Vec<u64> = {
let stream = e0.gpu.stream();
planes0.iter().map(|p| p.device_ptr(&stream).0).collect()
};
ws0.take_u64_h2d(e0, "hc.ptrs", &ptr_vals, 0)?
};
let use_graphs = decode_graphs_on() && step_ws_on();
let graphs_live = use_graphs && tgraphs.warm;
if use_graphs && !tgraphs.warm {
tgraphs.warm = true;
}
if graphs_live && tgraphs.a[0].len() != self.layers.len() {
for d in 0..2 {
tgraphs.a[d] = (0..self.layers.len()).map(|_| None).collect();
tgraphs.b[d] = (0..self.layers.len()).map(|_| None).collect();
tgraphs.c[d] = (0..self.layers.len()).map(|_| None).collect();
tgraphs.d[d] = (0..self.layers.len()).map(|_| None).collect();
}
}
for (li, layer) in self.layers.iter().enumerate() {
let lstate = &mut lstates[li];
let tw = &shard.layers[li];
let ts = &mut tlayers[li];
let eps_a = layer.eps_attn;
let eps_m = layer.eps_mlp;
let moe = &layer.moe;
let ff = moe.plan.expert_intermediate_size as usize;
let experts = moe.plan.expert_count as usize;
let selected = moe.plan.experts_per_token as usize;
let sff = moe
.plan
.shared
.as_ref()
.map(|s| s.intermediate_size as usize)
.unwrap_or(0);
let sffh = sff / 2;
match (&layer.mixer, &tw.mixer0, &tw.mixer1) {
(MixerW::Gdn(_), MixerHalfW::Gdn(h0), MixerHalfW::Gdn(h1)) => {
{
let _g = e1.gpu.enter_main()?;
if let (Some(ple1), Some(ps1)) = (tw.ple1.as_ref(), ts.ple1.as_mut()) {
let table = &layer.ple.as_ref().expect("ple plan").table;
self.ple_block(
e1,
layer,
ple1,
table,
ps1,
&mut planes1,
tokens,
1,
false,
None,
)?;
}
if graphs_live && layer.ple.is_none() {
if tgraphs.a[1][li].is_none() {
tgraphs.a[1][li] =
Some(e1.capture_graph_retained_nowarm(|eng| {
self.tp2_gdn_seg_a(
eng,
ws1,
&ptrs1,
&tw.attn_gate1,
h1,
&mut ts.m1,
&planes1,
eps_a,
shard.stage0_raw[0],
)
})?);
}
tgraphs.a[1][li].as_ref().unwrap().0.launch()?;
} else {
self.tp2_gdn_seg_a(
e1,
ws1,
&ptrs1,
&tw.attn_gate1,
h1,
&mut ts.m1,
&planes1,
eps_a,
shard.stage0_raw[0],
)?;
}
shard.ev1[0].record(&e1.gpu.stream())?;
}
{
let _g = e0.gpu.enter_main()?;
if let (Some(ple), Some(ps)) = (layer.ple.as_ref(), lstate.ple.as_mut()) {
self.ple_block(
e0,
layer,
ple,
&ple.table,
ps,
&mut planes0,
tokens,
1,
false,
None,
)?;
}
if graphs_live && layer.ple.is_none() {
if tgraphs.a[0][li].is_none() {
tgraphs.a[0][li] =
Some(e0.capture_graph_retained_nowarm(|eng| {
self.tp2_gdn_seg_a(
eng,
ws0,
&ptrs0,
&layer.attn_gate,
h0,
&mut ts.m0,
&planes0,
eps_a,
shard.stage1_raw[0],
)
})?);
}
tgraphs.a[0][li].as_ref().unwrap().0.launch()?;
} else {
self.tp2_gdn_seg_a(
e0,
ws0,
&ptrs0,
&layer.attn_gate,
h0,
&mut ts.m0,
&planes0,
eps_a,
shard.stage1_raw[0],
)?;
}
shard.ev0[0].record(&e0.gpu.stream())?;
}
}
(MixerW::Qsa(qsa), MixerHalfW::Qsa(h0), MixerHalfW::Qsa(h1)) => {
let (q1, g1, inj1) = {
let _g = e1.gpu.enter_main()?;
let (mixed1, inj1) = self.gate_read(
e1,
ws1,
&ptrs1,
&tw.attn_gate1,
&planes1,
1,
eps_a,
false,
)?;
let (q1, g1) = self
.qsa_half_proj(e1, ws1, eps_a, h1, &mixed1, &mut ts.m1, base_pos, 1)?;
ws1.put_f32("hc.mixed", mixed1);
(q1, g1, inj1)
};
let (sels, q0, g0, inj0) = {
let _g = e0.gpu.enter_main()?;
let (mixed0, inj0) = self.gate_read(
e0,
ws0,
&ptrs0,
&layer.attn_gate,
&planes0,
1,
eps_a,
false,
)?;
let (q0, g0) = self
.qsa_half_proj(e0, ws0, eps_a, h0, &mixed0, &mut ts.m0, base_pos, 1)?;
let MixerState::Qsa {
raw_keys,
pooled_keys,
pooled_dev,
pooled_dev_rows,
raw_dev,
raw_dev_rows,
idx_audit,
..
} = &mut lstate.mixer
else {
return Err("qwen4exp_gpu tp2: QSA layer without raw-key cache".into());
};
let sels = self.qsa_update_select(
e0,
ws0,
qsa,
eps_a,
&mixed0,
raw_keys,
pooled_keys,
pooled_dev,
pooled_dev_rows,
raw_dev,
raw_dev_rows,
idx_audit.as_mut(),
base_pos,
1,
0,
false,
)?;
ws0.put_f32("hc.mixed", mixed0);
(sels, q0, g0, inj0)
};
let t_kv = base_pos + 1;
let block_size = qsa.overlay.block_size as usize;
let (pos_flat, meta, max_count) = rowsel_positions(&sels, block_size);
{
let _g = e1.gpu.enter_main()?;
let pos_dev = ws1.take_i32(e1, "qsa.selpos", &pos_flat, 0)?;
let meta_dev = ws1.take_i32(e1, "qsa.selmeta", &meta, 0)?;
let p1 = self.qsa_half_attend(
e1, ws1, h1, &ts.m1, q1, g1, &pos_dev, &meta_dev, max_count, 1, t_kv,
)?;
ws1.put_i32("qsa.selpos", pos_dev);
ws1.put_i32("qsa.selmeta", meta_dev);
launch_push(e1, &p1, shard.stage0_raw[0], hidden)?;
ws1.put_f32("mixer.out", p1);
put_inject(ws1, inj1);
shard.ev1[0].record(&e1.gpu.stream())?;
}
{
let _g = e0.gpu.enter_main()?;
let pos_dev = ws0.take_i32(e0, "qsa.selpos", &pos_flat, 0)?;
let meta_dev = ws0.take_i32(e0, "qsa.selmeta", &meta, 0)?;
let p0 = self.qsa_half_attend(
e0, ws0, h0, &ts.m0, q0, g0, &pos_dev, &meta_dev, max_count, 1, t_kv,
)?;
ws0.put_i32("qsa.selpos", pos_dev);
ws0.put_i32("qsa.selmeta", meta_dev);
launch_push(e0, &p0, shard.stage1_raw[0], hidden)?;
ws0.put_f32("mixer.out", p0);
put_inject(ws0, inj0);
shard.ev0[0].record(&e0.gpu.stream())?;
}
}
_ => return Err("qwen4exp_gpu tp2: mixer/shard shape mismatch".into()),
}
{
let _g = e0.gpu.enter_main()?;
e0.gpu.stream().wait(&shard.ev1[0])?;
}
{
let _g = e1.gpu.enter_main()?;
e1.gpu.stream().wait(&shard.ev0[0])?;
}
{
let _g = e1.gpu.enter_main()?;
if graphs_live {
if tgraphs.b[1][li].is_none() {
tgraphs.b[1][li] = Some(e1.capture_graph_retained_nowarm(|eng| {
self.tp2_seg_b(
eng,
ws1,
&ptrs1,
&tw.mlp_gate1,
&mut planes1,
&shard.stage1[0],
false,
eps_m,
Some((
&tw.moe.shared_gu1_b16,
&tw.moe.shared_down1,
tw.moe.shared_input_gate1.as_ref(),
sffh,
)),
)
})?);
}
tgraphs.b[1][li].as_ref().unwrap().0.launch()?;
} else {
self.tp2_seg_b(
e1,
ws1,
&ptrs1,
&tw.mlp_gate1,
&mut planes1,
&shard.stage1[0],
false,
eps_m,
Some((
&tw.moe.shared_gu1_b16,
&tw.moe.shared_down1,
tw.moe.shared_input_gate1.as_ref(),
sffh,
)),
)?;
}
}
{
let _g = e0.gpu.enter_main()?;
if graphs_live {
if tgraphs.b[0][li].is_none() {
tgraphs.b[0][li] = Some(e0.capture_graph_retained_nowarm(|eng| {
self.tp2_seg_b(
eng,
ws0,
&ptrs0,
&layer.mlp_gate,
&mut planes0,
&shard.stage0[0],
true,
eps_m,
None,
)
})?);
}
tgraphs.b[0][li].as_ref().unwrap().0.launch()?;
} else {
self.tp2_seg_b(
e0,
ws0,
&ptrs0,
&layer.mlp_gate,
&mut planes0,
&shard.stage0[0],
true,
eps_m,
None,
)?;
}
}
let route = {
let _g = e0.gpu.enter_main()?;
let mixed0 = ws0.take_f32(e0, "hc.mixed", hidden, 0)?;
let mut router_out = ws0.take_f32(e0, "moe.router", experts, 0)?;
let none: Option<CudaSlice<u8>> = None;
let rb = if router_bf16_on() {
&moe.router_b16
} else {
&none
};
linear_trunk_into(
e0,
&moe.router,
rb,
&mixed0,
&mut router_out,
1,
hidden,
experts,
)?;
let logits = e0.dtoh_view(&router_out.slice(0..experts))?;
ws0.put_f32("moe.router", router_out);
ws0.put_f32("hc.mixed", mixed0);
host_route_softmax_topk(&logits, selected)
};
let place = &tw.place;
let mut sel0: Vec<i32> = Vec::with_capacity(selected);
let mut w0: Vec<f32> = Vec::with_capacity(selected);
let mut sel1: Vec<i32> = Vec::with_capacity(selected);
let mut w1: Vec<f32> = Vec::with_capacity(selected);
for &(expert, weight) in &route {
if place.rank(expert) == 0 {
sel0.push(place.local(expert) as i32);
w0.push(weight);
} else {
sel1.push(place.local(expert) as i32);
w1.push(weight);
}
}
{
let r0: Vec<Vec<(usize, f32)>> = vec![
sel0.iter()
.zip(&w0)
.map(|(&s, &w)| (s as usize, w))
.collect(),
];
let r1: Vec<Vec<(usize, f32)>> = vec![
sel1.iter()
.zip(&w1)
.map(|(&s, &w)| (s as usize, w))
.collect(),
];
tp2_count_split(&r0, &r1);
trace_moe_routes(layer.index, 1, std::slice::from_ref(&route));
}
match tp2_gate_red()? {
Tp2GateRed::None => {}
Tp2GateRed::SkipPeerMoe => {
sel1.clear();
w1.clear();
}
Tp2GateRed::PeerLocalIds => {
sel0.extend(sel1.drain(..));
w0.extend(w1.drain(..));
}
Tp2GateRed::ReverseePeerWeights => w1.reverse(),
}
let max_sel = selected;
{
let _g = e1.gpu.enter_main()?;
ws1.upsert_u8(e1, "moe.pack", &tp2_pack_bytes(&sel1, &w1, max_sel), 0)?;
if graphs_live {
if tgraphs.c[1][li].is_none() {
tgraphs.c[1][li] = Some(e1.capture_graph_retained_nowarm(|eng| {
self.tp2_seg_c(
eng,
ws1,
(
&tw.moe.gate1.codes,
&tw.moe.gate1.scales,
&tw.moe.gate1.macros_dev,
),
(
&tw.moe.up1.codes,
&tw.moe.up1.scales,
&tw.moe.up1.macros_dev,
),
(
&tw.moe.down1.codes,
&tw.moe.down1.scales,
&tw.moe.down1.macros_dev,
),
ff,
max_sel,
None,
tw.moe.shared_input_gate1.is_some(),
shard.stage0_raw[1],
)
})?);
}
tgraphs.c[1][li].as_ref().unwrap().0.launch()?;
} else {
self.tp2_seg_c(
e1,
ws1,
(
&tw.moe.gate1.codes,
&tw.moe.gate1.scales,
&tw.moe.gate1.macros_dev,
),
(
&tw.moe.up1.codes,
&tw.moe.up1.scales,
&tw.moe.up1.macros_dev,
),
(
&tw.moe.down1.codes,
&tw.moe.down1.scales,
&tw.moe.down1.macros_dev,
),
ff,
max_sel,
None,
tw.moe.shared_input_gate1.is_some(),
shard.stage0_raw[1],
)?;
}
shard.ev1[1].record(&e1.gpu.stream())?;
}
{
let _g = e0.gpu.enter_main()?;
let (
BankHalf::Nvfp4 {
codes: gc,
scales: gs,
macros_dev: gm,
..
},
BankHalf::Nvfp4 {
codes: uc,
scales: us,
macros_dev: um,
..
},
BankHalf::Nvfp4 {
codes: dc,
scales: ds,
macros_dev: dm,
..
},
) = (&moe.bank.gate, &moe.bank.up, &moe.bank.down)
else {
return Err("qwen4exp_gpu tp2: card0 bank is not NVFP4".into());
};
ws0.upsert_u8(e0, "moe.pack", &tp2_pack_bytes(&sel0, &w0, max_sel), 0)?;
if graphs_live {
if tgraphs.c[0][li].is_none() {
tgraphs.c[0][li] = Some(e0.capture_graph_retained_nowarm(|eng| {
self.tp2_seg_c(
eng,
ws0,
(gc, gs, gm),
(uc, us, um),
(dc, ds, dm),
ff,
max_sel,
Some((
&tw.moe.shared_gu0_b16,
&tw.moe.shared_down0,
moe.shared_input_gate.as_ref(),
sffh,
)),
false,
shard.stage1_raw[1],
)
})?);
}
tgraphs.c[0][li].as_ref().unwrap().0.launch()?;
} else {
self.tp2_seg_c(
e0,
ws0,
(gc, gs, gm),
(uc, us, um),
(dc, ds, dm),
ff,
max_sel,
Some((
&tw.moe.shared_gu0_b16,
&tw.moe.shared_down0,
moe.shared_input_gate.as_ref(),
sffh,
)),
false,
shard.stage1_raw[1],
)?;
}
shard.ev0[1].record(&e0.gpu.stream())?;
}
{
let _g = e1.gpu.enter_main()?;
e1.gpu.stream().wait(&shard.ev0[1])?;
if graphs_live {
if tgraphs.d[1][li].is_none() {
tgraphs.d[1][li] = Some(e1.capture_graph_retained_nowarm(|eng| {
self.tp2_seg_d(eng, ws1, &ptrs1, &mut planes1, &shard.stage1[1], false)
})?);
}
tgraphs.d[1][li].as_ref().unwrap().0.launch()?;
} else {
self.tp2_seg_d(e1, ws1, &ptrs1, &mut planes1, &shard.stage1[1], false)?;
}
}
{
let _g = e0.gpu.enter_main()?;
e0.gpu.stream().wait(&shard.ev1[1])?;
if graphs_live {
if tgraphs.d[0][li].is_none() {
tgraphs.d[0][li] = Some(e0.capture_graph_retained_nowarm(|eng| {
self.tp2_seg_d(eng, ws0, &ptrs0, &mut planes0, &shard.stage0[1], true)
})?);
}
tgraphs.d[0][li].as_ref().unwrap().0.launch()?;
} else {
self.tp2_seg_d(e0, ws0, &ptrs0, &mut planes0, &shard.stage0[1], true)?;
}
}
}
{
let _g = e1.gpu.enter_main()?;
if graphs_live {
if tgraphs.exit[1].is_none() {
tgraphs.exit[1] = Some(e1.capture_graph_retained_nowarm(|eng| {
self.tp2_seg_exit(
eng,
ws1,
&ptrs1,
&shard.exit_gate1,
&planes1,
&shard.lm_head1,
vocab - vsplit,
1, )
})?);
}
tgraphs.exit[1].as_ref().unwrap().0.launch()?;
} else {
self.tp2_seg_exit(
e1,
ws1,
&ptrs1,
&shard.exit_gate1,
&planes1,
&shard.lm_head1,
vocab - vsplit,
1, )?;
}
}
{
let _g = e0.gpu.enter_main()?;
let head0 = self
.output_b16
.as_ref()
.ok_or("qwen4exp_gpu tp2: lm_head has no bf16 twin")?;
if graphs_live {
if tgraphs.exit[0].is_none() {
tgraphs.exit[0] = Some(e0.capture_graph_retained_nowarm(|eng| {
self.tp2_seg_exit(
eng,
ws0,
&ptrs0,
&self.exit_mixer,
&planes0,
head0,
vsplit,
1, )
})?);
}
tgraphs.exit[0].as_ref().unwrap().0.launch()?;
} else {
self.tp2_seg_exit(
e0,
ws0,
&ptrs0,
&self.exit_mixer,
&planes0,
head0,
vsplit,
1, )?;
}
}
let mut out = vec![0.0f32; vocab];
{
let _g = e0.gpu.enter_main()?;
let logits0 = ws0.peek_f32("logits")?;
let host0 = e0.dtoh_view(&logits0.slice(0..vsplit))?;
out[..vsplit].copy_from_slice(&host0);
}
{
let _g = e1.gpu.enter_main()?;
let logits1 = ws1.peek_f32("logits")?;
let host1 = e1.dtoh_view(&logits1.slice(0..vocab - vsplit))?;
out[vsplit..].copy_from_slice(&host1);
}
for (s, plane) in planes0.into_iter().enumerate() {
ws0.put_f32(PLANE_SLOTS[s], plane);
}
for (s, plane) in planes1.into_iter().enumerate() {
ws1.put_f32(PLANE_SLOTS[s], plane);
}
ws0.put_u64("hc.ptrs", ptrs0);
ws1.put_u64("hc.ptrs", ptrs1);
state.pos += 1;
Ok(out)
}
}
impl Qwen4ExpGpu {
pub fn alloc_state_tp2(
&self,
e0: &Engine,
e1: &Engine,
shard: &Tp2Shard,
capacity: usize,
reserve: usize,
) -> Res<Qwen4ExpState> {
let mut state = {
let mut st = self.alloc_state_reserve(e0, 1, 1, None)?;
st.capacity = capacity;
st.reserve = reserve;
st
};
let mut tlayers = Vec::with_capacity(self.layers.len());
for (layer, tw) in self.layers.iter().zip(shard.layers.iter()) {
let mk_half = |e: &Engine, hw: &MixerHalfW| -> Res<MixerHalfState> {
let _g = e.gpu.enter_main()?;
match hw {
MixerHalfW::Gdn(h) => {
let conv_dim = 2 * h.nk_h * h.hk + h.nv_h * h.hv;
let pad = h.kernel - 1;
Ok(MixerHalfState::Gdn {
conv: e.zeros(pad * conv_dim)?,
state: e.zeros(h.nv_h * h.hv * h.hk)?,
})
}
MixerHalfW::Qsa(h) => {
let kv_dim = h.nkv_h * h.hd;
let kv = if kv_quant_on() {
QsaKvStore::Q8Q5 {
k: e.alloc_u8(capacity * q8_row_bytes(kv_dim))?,
v: e.alloc_u8(capacity * q5_row_bytes(kv_dim))?,
}
} else {
QsaKvStore::F32 {
k: e.zeros(capacity * kv_dim)?,
v: e.zeros(capacity * kv_dim)?,
}
};
Ok(MixerHalfState::Qsa { kv })
}
}
};
let m0 = mk_half(e0, &tw.mixer0)?;
let m1 = mk_half(e1, &tw.mixer1)?;
let ple1 = match layer.ple.as_ref() {
None => None,
Some(ple) => {
let pad = (ple.plan.conv_kernel as usize - 1) * ple.plan.max_ngram as usize;
let _g = e1.gpu.enter_main()?;
let mut conv_hist = Vec::with_capacity(self.streams);
for _ in 0..self.streams {
conv_hist.push(e1.zeros(pad * self.hidden)?);
}
Some(PleState {
conv_hist,
ngram_ids: Vec::new(),
ngram_history: Vec::new(),
ngram_last_eos: -1,
})
}
};
tlayers.push(Tp2LayerState { m0, m1, ple1 });
}
state.tp2 = Some(Tp2State {
ws1: StepPool::default(),
layers: tlayers,
graphs: Tp2Graphs::default(),
pf_stage0: None,
pf_stage1: None,
pf_stage0_raw: [0; 2],
pf_stage1_raw: [0; 2],
pf_rows: 0,
});
Ok(state)
}
pub fn prefill_extend_tp2(
&self,
e0: &Engine,
e1: &Engine,
shard: &Tp2Shard,
ids: &[u32],
state: &mut Qwen4ExpState,
chunk: usize,
) -> Res<Vec<f32>> {
if ids.is_empty() || chunk == 0 {
return Err("qwen4exp_gpu: prefill_extend_tp2 needs ids and a chunk size".into());
}
let mut last = Vec::new();
for piece in ids.chunks(chunk) {
let is_last =
piece.as_ptr() as usize + piece.len() * 4 == ids.as_ptr() as usize + ids.len() * 4;
let head = if is_last {
HeadMode::LastRow
} else {
HeadMode::Skip
};
last = self.forward_tp2(e0, e1, shard, piece, state, head)?;
}
Ok(last)
}
#[allow(clippy::too_many_arguments)]
pub fn forward_tp2(
&self,
e0: &Engine,
e1: &Engine,
shard: &Tp2Shard,
ids: &[u32],
state: &mut Qwen4ExpState,
head: HeadMode,
) -> Res<Vec<f32>> {
if !trunk_bf16_on() || !hc_fused_gate_on() {
return Err(
"qwen4exp_gpu tp2: requires set_trunk_bf16(true) and set_hc_fused_gate(true)"
.into(),
);
}
let t = ids.len();
if t == 0 {
return Err("qwen4exp_gpu tp2: empty chunk".into());
}
if state.pos + t > state.capacity {
return Err("qwen4exp_gpu: state capacity exceeded".into());
}
if state.tp2.is_none() {
self.tp2_migrate(e0, e1, shard, state)?;
state.graphs = StepGraphs::default();
}
let hidden = self.hidden;
let vocab = self.vocab;
let vsplit = shard.vsplit;
let base_pos = state.pos;
let reserve = state.reserve;
state.tokens.extend_from_slice(ids);
let Qwen4ExpState {
ref tokens,
ws: ref mut ws0,
ref mut tp2,
layers: ref mut lstates,
..
} = *state;
let tp2s = tp2.as_mut().expect("alloc'd or migrated above");
if tp2s.pf_rows < t {
{
let _g = e1.gpu.enter_main()?;
let s1 = [e1.zeros(t * hidden)?, e1.zeros(t * hidden)?];
let s = e1.gpu.stream();
tp2s.pf_stage1_raw = [s1[0].device_ptr(&s).0, s1[1].device_ptr(&s).0];
tp2s.pf_stage1 = Some(s1);
}
{
let _g = e0.gpu.enter_main()?;
let s0 = [e0.zeros(t * hidden)?, e0.zeros(t * hidden)?];
let s = e0.gpu.stream();
tp2s.pf_stage0_raw = [s0[0].device_ptr(&s).0, s0[1].device_ptr(&s).0];
tp2s.pf_stage0 = Some(s0);
}
tp2s.pf_rows = t;
}
let Tp2State {
ws1,
layers: tlayers,
pf_stage0,
pf_stage1,
pf_stage0_raw,
pf_stage1_raw,
..
} = tp2s;
let pf_stage0 = pf_stage0.as_ref().expect("sized above");
let pf_stage1 = pf_stage1.as_ref().expect("sized above");
let resv = reserve.max(t);
let mut embedded = vec![0.0f32; t * hidden];
for (row, &token) in ids.iter().enumerate() {
let token = token as usize;
if token >= vocab {
return Err(format!("qwen4exp_gpu: token {token} out of range").into());
}
embedded[row * hidden..(row + 1) * hidden]
.copy_from_slice(&self.embed_host[token * hidden..(token + 1) * hidden]);
}
let mut planes1: Vec<CudaSlice<f32>> = Vec::with_capacity(self.streams);
let ptrs1 = {
let _g = e1.gpu.enter_main()?;
let embedded_dev = ws1.take_f32_h2d(e1, "entry.embed", &embedded, resv * hidden)?;
for s in 0..self.streams {
let mut plane = ws1.take_f32(e1, PLANE_SLOTS[s], t * hidden, resv * hidden)?;
e1.copy_into(&mut plane, 0, &embedded_dev, t * hidden)?;
planes1.push(plane);
}
ws1.put_f32("entry.embed", embedded_dev);
let ptr_vals: Vec<u64> = {
let stream = e1.gpu.stream();
planes1.iter().map(|p| p.device_ptr(&stream).0).collect()
};
ws1.take_u64_h2d(e1, "hc.ptrs", &ptr_vals, 0)?
};
let mut planes0: Vec<CudaSlice<f32>> = Vec::with_capacity(self.streams);
let ptrs0 = {
let _g = e0.gpu.enter_main()?;
let embedded_dev = ws0.take_f32_h2d(e0, "entry.embed", &embedded, resv * hidden)?;
for s in 0..self.streams {
let mut plane = ws0.take_f32(e0, PLANE_SLOTS[s], t * hidden, resv * hidden)?;
e0.copy_into(&mut plane, 0, &embedded_dev, t * hidden)?;
planes0.push(plane);
}
ws0.put_f32("entry.embed", embedded_dev);
let ptr_vals: Vec<u64> = {
let stream = e0.gpu.stream();
planes0.iter().map(|p| p.device_ptr(&stream).0).collect()
};
ws0.take_u64_h2d(e0, "hc.ptrs", &ptr_vals, 0)?
};
for (li, layer) in self.layers.iter().enumerate() {
let lstate = &mut lstates[li];
let tw = &shard.layers[li];
let ts = &mut tlayers[li];
let eps_a = layer.eps_attn;
let eps_m = layer.eps_mlp;
let moe = &layer.moe;
let ff = moe.plan.expert_intermediate_size as usize;
let experts = moe.plan.expert_count as usize;
let selected = moe.plan.experts_per_token as usize;
let sff = moe
.plan
.shared
.as_ref()
.map(|s| s.intermediate_size as usize)
.unwrap_or(0);
let sffh = sff / 2;
if let (Some(ple), Some(ps)) = (layer.ple.as_ref(), lstate.ple.as_mut()) {
let _g = e0.gpu.enter_main()?;
self.ple_block(
e0,
layer,
ple,
&ple.table,
ps,
&mut planes0,
tokens,
t,
false,
None,
)?;
}
if let (Some(ple1), Some(ps1)) = (tw.ple1.as_ref(), ts.ple1.as_mut()) {
let table = &layer.ple.as_ref().expect("ple plan").table;
let _g = e1.gpu.enter_main()?;
self.ple_block(
e1,
layer,
ple1,
table,
ps1,
&mut planes1,
tokens,
t,
false,
None,
)?;
}
match (&layer.mixer, &tw.mixer0, &tw.mixer1) {
(MixerW::Gdn(_), MixerHalfW::Gdn(h0), MixerHalfW::Gdn(h1)) => {
{
let _g = e1.gpu.enter_main()?;
let (mixed1, inj1) = self.gate_read(
e1,
ws1,
&ptrs1,
&tw.attn_gate1,
&planes1,
t,
eps_a,
false,
)?;
let p1 =
self.gdn_forward_half(e1, ws1, eps_a, h1, &mixed1, &mut ts.m1, t)?;
ws1.put_f32("hc.mixed", mixed1);
launch_push(e1, &p1, pf_stage0_raw[0], t * hidden)?;
ws1.put_f32("mixer.out", p1);
put_inject(ws1, inj1);
shard.ev1[0].record(&e1.gpu.stream())?;
}
{
let _g = e0.gpu.enter_main()?;
let (mixed0, inj0) = self.gate_read(
e0,
ws0,
&ptrs0,
&layer.attn_gate,
&planes0,
t,
eps_a,
false,
)?;
let p0 =
self.gdn_forward_half(e0, ws0, eps_a, h0, &mixed0, &mut ts.m0, t)?;
ws0.put_f32("hc.mixed", mixed0);
launch_push(e0, &p0, pf_stage1_raw[0], t * hidden)?;
ws0.put_f32("mixer.out", p0);
put_inject(ws0, inj0);
shard.ev0[0].record(&e0.gpu.stream())?;
}
}
(MixerW::Qsa(qsa), MixerHalfW::Qsa(h0), MixerHalfW::Qsa(h1)) => {
let (q1, g1, inj1) = {
let _g = e1.gpu.enter_main()?;
let (mixed1, inj1) = self.gate_read(
e1,
ws1,
&ptrs1,
&tw.attn_gate1,
&planes1,
t,
eps_a,
false,
)?;
let (q1, g1) = self
.qsa_half_proj(e1, ws1, eps_a, h1, &mixed1, &mut ts.m1, base_pos, t)?;
ws1.put_f32("hc.mixed", mixed1);
(q1, g1, inj1)
};
let (sels, q0, g0, inj0) = {
let _g = e0.gpu.enter_main()?;
let (mixed0, inj0) = self.gate_read(
e0,
ws0,
&ptrs0,
&layer.attn_gate,
&planes0,
t,
eps_a,
false,
)?;
let (q0, g0) = self
.qsa_half_proj(e0, ws0, eps_a, h0, &mixed0, &mut ts.m0, base_pos, t)?;
let MixerState::Qsa {
raw_keys,
pooled_keys,
pooled_dev,
pooled_dev_rows,
raw_dev,
raw_dev_rows,
idx_audit,
..
} = &mut lstate.mixer
else {
return Err("qwen4exp_gpu tp2: QSA layer without raw-key cache".into());
};
let sels = self.qsa_update_select(
e0,
ws0,
qsa,
eps_a,
&mixed0,
raw_keys,
pooled_keys,
pooled_dev,
pooled_dev_rows,
raw_dev,
raw_dev_rows,
idx_audit.as_mut(),
base_pos,
t,
0,
false,
)?;
ws0.put_f32("hc.mixed", mixed0);
(sels, q0, g0, inj0)
};
let t_kv = base_pos + t;
let block_size = qsa.overlay.block_size as usize;
let (pos_flat, meta, max_count) = rowsel_positions(&sels, block_size);
{
let _g = e1.gpu.enter_main()?;
let pos_dev = ws1.take_i32(e1, "qsa.selpos", &pos_flat, 0)?;
let meta_dev = ws1.take_i32(e1, "qsa.selmeta", &meta, 0)?;
let p1 = self.qsa_half_attend(
e1, ws1, h1, &ts.m1, q1, g1, &pos_dev, &meta_dev, max_count, t, t_kv,
)?;
ws1.put_i32("qsa.selpos", pos_dev);
ws1.put_i32("qsa.selmeta", meta_dev);
launch_push(e1, &p1, pf_stage0_raw[0], t * hidden)?;
ws1.put_f32("mixer.out", p1);
put_inject(ws1, inj1);
shard.ev1[0].record(&e1.gpu.stream())?;
}
{
let _g = e0.gpu.enter_main()?;
let pos_dev = ws0.take_i32(e0, "qsa.selpos", &pos_flat, 0)?;
let meta_dev = ws0.take_i32(e0, "qsa.selmeta", &meta, 0)?;
let p0 = self.qsa_half_attend(
e0, ws0, h0, &ts.m0, q0, g0, &pos_dev, &meta_dev, max_count, t, t_kv,
)?;
ws0.put_i32("qsa.selpos", pos_dev);
ws0.put_i32("qsa.selmeta", meta_dev);
launch_push(e0, &p0, pf_stage1_raw[0], t * hidden)?;
ws0.put_f32("mixer.out", p0);
put_inject(ws0, inj0);
shard.ev0[0].record(&e0.gpu.stream())?;
}
}
_ => return Err("qwen4exp_gpu tp2: mixer/shard shape mismatch".into()),
}
{
let _g = e0.gpu.enter_main()?;
e0.gpu.stream().wait(&shard.ev1[0])?;
}
{
let _g = e1.gpu.enter_main()?;
e1.gpu.stream().wait(&shard.ev0[0])?;
}
let join_write = |e: &Engine,
ws: &mut StepPool,
ptrs: &CudaSlice<u64>,
planes: &mut [CudaSlice<f32>],
stage: &CudaSlice<f32>,
rank0: bool|
-> Res<()> {
let p = ws.take_f32(e, "mixer.out", t * hidden, 0)?;
let mut out = ws.take_f32(e, "join.out", t * hidden, 0)?;
if rank0 {
e.add(&p, stage, &mut out, t * hidden)?;
} else {
e.add(stage, &p, &mut out, t * hidden)?;
}
let inj = take_inject(e, ws, self.streams, t)?;
self.gate_write(e, planes, ptrs, &out, &inj, t)?;
ws.put_f32("mixer.out", p);
ws.put_f32("join.out", out);
put_inject(ws, inj);
Ok(())
};
{
let _g = e1.gpu.enter_main()?;
join_write(e1, ws1, &ptrs1, &mut planes1, &pf_stage1[0], false)?;
}
{
let _g = e0.gpu.enter_main()?;
join_write(e0, ws0, &ptrs0, &mut planes0, &pf_stage0[0], true)?;
}
let mixed1 = {
let _g = e1.gpu.enter_main()?;
let (mixed1, injm1) =
self.gate_read(e1, ws1, &ptrs1, &tw.mlp_gate1, &planes1, t, eps_m, false)?;
put_inject(ws1, injm1);
mixed1
};
let mixed0 = {
let _g = e0.gpu.enter_main()?;
let (mixed0, injm0) =
self.gate_read(e0, ws0, &ptrs0, &layer.mlp_gate, &planes0, t, eps_m, false)?;
put_inject(ws0, injm0);
mixed0
};
let routes: Vec<Vec<(usize, f32)>> = {
let _g = e0.gpu.enter_main()?;
let mut router_out = ws0.take_f32(e0, "moe.router", t * experts, 0)?;
let none: Option<CudaSlice<u8>> = None;
let rb = if router_bf16_on() {
&moe.router_b16
} else {
&none
};
linear_trunk_into(
e0,
&moe.router,
rb,
&mixed0,
&mut router_out,
t,
hidden,
experts,
)?;
let logits = e0.dtoh_view(&router_out.slice(0..t * experts))?;
ws0.put_f32("moe.router", router_out);
let mut routes = Vec::with_capacity(t);
for token in 0..t {
routes.push(host_route_softmax_topk(
&logits[token * experts..(token + 1) * experts],
selected,
));
}
routes
};
let place = &tw.place;
let split_half = |home: bool| -> Vec<Vec<(usize, f32)>> {
routes
.iter()
.map(|r| {
r.iter()
.filter(|&&(eid, _)| (place.rank(eid) == 0) == home)
.map(|&(eid, w)| (place.local(eid), w))
.collect()
})
.collect()
};
let mut routes0 = split_half(true);
let mut routes1 = split_half(false);
tp2_count_split(&routes0, &routes1);
trace_moe_routes(layer.index, t, &routes);
match tp2_gate_red()? {
Tp2GateRed::None => {}
Tp2GateRed::SkipPeerMoe => routes1.iter_mut().for_each(|r| r.clear()),
Tp2GateRed::PeerLocalIds => {
for (r0, r1) in routes0.iter_mut().zip(routes1.iter_mut()) {
r0.append(r1);
}
}
Tp2GateRed::ReverseePeerWeights => {
for r in routes1.iter_mut() {
let n = r.len();
for i in 0..n / 2 {
let (a, b) = (r[i].1, r[n - 1 - i].1);
r[i].1 = b;
r[n - 1 - i].1 = a;
}
}
}
}
let (routes0, routes1) = (routes0, routes1);
{
let _g = e1.gpu.enter_main()?;
let mut out1 = self.tp2_moe_rows(
e1,
ws1,
(
&tw.moe.gate1.codes,
&tw.moe.gate1.scales,
&tw.moe.gate1.macros_dev,
),
(
&tw.moe.up1.codes,
&tw.moe.up1.scales,
&tw.moe.up1.macros_dev,
),
(
&tw.moe.down1.codes,
&tw.moe.down1.scales,
&tw.moe.down1.macros_dev,
),
&routes1,
&mixed1,
t,
ff,
)?;
let (sh, g) = self.tp2_shared_half(
e1,
ws1,
&tw.moe.shared_gu1_b16,
&tw.moe.shared_down1,
tw.moe.shared_input_gate1.as_ref(),
&mixed1,
sffh,
t,
)?;
match g.as_ref() {
Some(g) => e1.add_scaled_rows(&sh, g, &mut out1, hidden, t)?,
None => {
let mut summed = ws1.take_f32(e1, "moe.sum", t * hidden, 0)?;
e1.add(&out1, &sh, &mut summed, t * hidden)?;
ws1.put_f32("moe.out", out1);
out1 = summed;
}
}
ws1.put_f32("moe.sh_down", sh);
if let Some(g) = g {
ws1.put_f32("moe.g", g);
}
launch_push(e1, &out1, pf_stage0_raw[1], t * hidden)?;
ws1.put_f32("moe.out", out1);
ws1.put_f32("hc.mixed", mixed1);
shard.ev1[1].record(&e1.gpu.stream())?;
}
{
let _g = e0.gpu.enter_main()?;
let (
BankHalf::Nvfp4 {
codes: gc,
scales: gs,
macros_dev: gm,
..
},
BankHalf::Nvfp4 {
codes: uc,
scales: us,
macros_dev: um,
..
},
BankHalf::Nvfp4 {
codes: dc,
scales: ds,
macros_dev: dm,
..
},
) = (&moe.bank.gate, &moe.bank.up, &moe.bank.down)
else {
return Err("qwen4exp_gpu tp2: card0 bank is not NVFP4".into());
};
let mut out0 = self.tp2_moe_rows(
e0,
ws0,
(gc, gs, gm),
(uc, us, um),
(dc, ds, dm),
&routes0,
&mixed0,
t,
ff,
)?;
let (sh, g) = self.tp2_shared_half(
e0,
ws0,
&tw.moe.shared_gu0_b16,
&tw.moe.shared_down0,
moe.shared_input_gate.as_ref(),
&mixed0,
sffh,
t,
)?;
match g.as_ref() {
Some(g) => e0.add_scaled_rows(&sh, g, &mut out0, hidden, t)?,
None => {
let mut summed = ws0.take_f32(e0, "moe.sum", t * hidden, 0)?;
e0.add(&out0, &sh, &mut summed, t * hidden)?;
ws0.put_f32("moe.out", out0);
out0 = summed;
}
}
ws0.put_f32("moe.sh_down", sh);
if let Some(g) = g {
ws0.put_f32("moe.g", g);
}
launch_push(e0, &out0, pf_stage1_raw[1], t * hidden)?;
ws0.put_f32("moe.out", out0);
ws0.put_f32("hc.mixed", mixed0);
shard.ev0[1].record(&e0.gpu.stream())?;
}
{
let _g = e1.gpu.enter_main()?;
e1.gpu.stream().wait(&shard.ev0[1])?;
let p = ws1.take_f32(e1, "moe.out", t * hidden, 0)?;
let mut out = ws1.take_f32(e1, "join.out", t * hidden, 0)?;
e1.add(&pf_stage1[1], &p, &mut out, t * hidden)?;
let inj = take_inject(e1, ws1, self.streams, t)?;
self.gate_write(e1, &mut planes1, &ptrs1, &out, &inj, t)?;
ws1.put_f32("moe.out", p);
ws1.put_f32("join.out", out);
put_inject(ws1, inj);
}
{
let _g = e0.gpu.enter_main()?;
e0.gpu.stream().wait(&shard.ev1[1])?;
let p = ws0.take_f32(e0, "moe.out", t * hidden, 0)?;
let mut out = ws0.take_f32(e0, "join.out", t * hidden, 0)?;
e0.add(&p, &pf_stage0[1], &mut out, t * hidden)?;
let inj = take_inject(e0, ws0, self.streams, t)?;
self.gate_write(e0, &mut planes0, &ptrs0, &out, &inj, t)?;
ws0.put_f32("moe.out", p);
ws0.put_f32("join.out", out);
put_inject(ws0, inj);
}
}
let head_rows = match head {
HeadMode::All => t,
_ => 1,
};
let mut out = vec![
0.0f32;
if head == HeadMode::Skip {
0
} else {
head_rows * vocab
}
];
if head != HeadMode::Skip {
{
let _g = e1.gpu.enter_main()?;
let mut exit_planes: Vec<CudaSlice<f32>> = Vec::with_capacity(self.streams);
if head != HeadMode::All {
for (s, plane) in planes1.iter().enumerate() {
let mut row = ws1.take_f32(e1, EXIT_PLANE_SLOTS[s], hidden, hidden)?;
e1.copy_range_into(&mut row, 0, plane, (t - 1) * hidden, hidden)?;
exit_planes.push(row);
}
}
let use_planes: &[CudaSlice<f32>] = if head == HeadMode::All {
&planes1
} else {
&exit_planes
};
let ptr_vals: Vec<u64> = {
let stream = e1.gpu.stream();
use_planes.iter().map(|p| p.device_ptr(&stream).0).collect()
};
let eptrs = ws1.take_u64_h2d(e1, "exit.ptrs", &ptr_vals, 0)?;
self.tp2_seg_exit(
e1,
ws1,
&eptrs,
&shard.exit_gate1,
use_planes,
&shard.lm_head1,
vocab - vsplit,
head_rows,
)?;
ws1.put_u64("exit.ptrs", eptrs);
for (s, p) in exit_planes.into_iter().enumerate() {
ws1.put_f32(EXIT_PLANE_SLOTS[s], p);
}
let logits1 = ws1.peek_f32("logits")?;
let half1 = vocab - vsplit;
let host1 = e1.dtoh_view(&logits1.slice(0..head_rows * half1))?;
for r in 0..head_rows {
out[r * vocab + vsplit..(r + 1) * vocab]
.copy_from_slice(&host1[r * half1..(r + 1) * half1]);
}
}
{
let _g = e0.gpu.enter_main()?;
let head0 = self
.output_b16
.as_ref()
.ok_or("qwen4exp_gpu tp2: lm_head has no bf16 twin")?;
let mut exit_planes: Vec<CudaSlice<f32>> = Vec::with_capacity(self.streams);
if head != HeadMode::All {
for (s, plane) in planes0.iter().enumerate() {
let mut row = ws0.take_f32(e0, EXIT_PLANE_SLOTS[s], hidden, hidden)?;
e0.copy_range_into(&mut row, 0, plane, (t - 1) * hidden, hidden)?;
exit_planes.push(row);
}
}
let use_planes: &[CudaSlice<f32>] = if head == HeadMode::All {
&planes0
} else {
&exit_planes
};
let ptr_vals: Vec<u64> = {
let stream = e0.gpu.stream();
use_planes.iter().map(|p| p.device_ptr(&stream).0).collect()
};
let eptrs = ws0.take_u64_h2d(e0, "exit.ptrs", &ptr_vals, 0)?;
self.tp2_seg_exit(
e0,
ws0,
&eptrs,
&self.exit_mixer,
use_planes,
head0,
vsplit,
head_rows,
)?;
ws0.put_u64("exit.ptrs", eptrs);
for (s, p) in exit_planes.into_iter().enumerate() {
ws0.put_f32(EXIT_PLANE_SLOTS[s], p);
}
let logits0 = ws0.peek_f32("logits")?;
let host0 = e0.dtoh_view(&logits0.slice(0..head_rows * vsplit))?;
for r in 0..head_rows {
out[r * vocab..r * vocab + vsplit]
.copy_from_slice(&host0[r * vsplit..(r + 1) * vsplit]);
}
}
} else {
{
let _g = e0.gpu.enter_main()?;
e0.gpu.stream().synchronize()?;
}
{
let _g = e1.gpu.enter_main()?;
e1.gpu.stream().synchronize()?;
}
}
for (s, plane) in planes0.into_iter().enumerate() {
ws0.put_f32(PLANE_SLOTS[s], plane);
}
for (s, plane) in planes1.into_iter().enumerate() {
ws1.put_f32(PLANE_SLOTS[s], plane);
}
ws0.put_u64("hc.ptrs", ptrs0);
ws1.put_u64("hc.ptrs", ptrs1);
state.pos += t;
Ok(out)
}
#[allow(clippy::too_many_arguments)]
fn tp2_moe_rows(
&self,
e: &Engine,
ws: &mut StepPool,
gate: (&CudaSlice<u8>, &CudaSlice<u8>, &CudaSlice<f32>),
up: (&CudaSlice<u8>, &CudaSlice<u8>, &CudaSlice<f32>),
down: (&CudaSlice<u8>, &CudaSlice<u8>, &CudaSlice<f32>),
routes: &[Vec<(usize, f32)>],
mixed: &CudaSlice<f32>,
t: usize,
ff: usize,
) -> Res<CudaSlice<f32>> {
let hidden = self.hidden;
if !(sel_gufuse_on() && hidden % 32 == 0 && ff % 4 == 0) {
return Err(
"qwen4exp_gpu tp2: prefill MoE needs the gufuse geometry (hidden%32, ff%4)".into(),
);
}
let mut out = ws.take_f32(e, "moe.out", t * hidden, 0)?;
{
let mut view = out.slice_mut(0..t * hidden);
e.memset_zeros_view(&mut view)?;
}
const SLOT_CAP: usize = 8192;
let mut tok0 = 0usize;
while tok0 < t {
let mut tok_n = 0usize;
let mut slots = 0usize;
while tok0 + tok_n < t {
let n = routes[tok0 + tok_n].len();
if tok_n > 0 && slots + n > SLOT_CAP {
break;
}
slots += n;
tok_n += 1;
}
let batch = &routes[tok0..tok0 + tok_n];
let mut sel_all: Vec<i32> = Vec::with_capacity(slots);
let mut w_all: Vec<f32> = Vec::with_capacity(slots);
let mut tok_all: Vec<i32> = Vec::with_capacity(slots);
let mut ranges: Vec<(usize, usize)> = Vec::with_capacity(tok_n);
for (i, route) in batch.iter().enumerate() {
ranges.push((sel_all.len(), route.len()));
for &(eid, wgt) in route {
sel_all.push(eid as i32);
w_all.push(wgt);
tok_all.push((tok0 + i) as i32);
}
}
let s_total = sel_all.len();
if s_total > 0 {
let sel = ws.take_i32(e, "moe.sel", &sel_all, 0)?;
let w_dev = ws.take_f32_h2d(e, "moe.w", &w_all, 0)?;
let tokm = ws.take_i32(e, "moe.tok", &tok_all, 0)?;
let mut act = ws.take_f32(e, "moe.act", s_total * ff, 0)?;
launch_nvfp4_sel_gu_silu(
e,
gate,
up,
Some(&sel),
0,
s_total,
mixed,
&mut act,
hidden,
ff,
Some((&tokm, hidden)),
)?;
let mut partial = ws.take_f32(e, "moe.partial", s_total * hidden, 0)?;
launch_nvfp4_sel_matvec(
e,
down.0,
down.1,
down.2,
&sel,
&act,
&mut partial,
s_total,
ff,
hidden,
ff,
)?;
for (i, &(start, len)) in ranges.iter().enumerate() {
if len > 0 {
launch_axpy_rows_seq_at(
e,
&partial,
start,
&w_dev,
start,
&mut out,
tok0 + i,
hidden,
len,
)?;
}
}
ws.put_i32("moe.sel", sel);
ws.put_i32("moe.tok", tokm);
ws.put_f32("moe.w", w_dev);
ws.put_f32("moe.act", act);
ws.put_f32("moe.partial", partial);
}
tok0 += tok_n;
}
Ok(out)
}
#[allow(clippy::too_many_arguments)]
fn tp2_gdn_seg_a(
&self,
e: &Engine,
ws: &mut StepPool,
ptrs: &CudaSlice<u64>,
attn_gate: &GateW,
h: &GdnHalfW,
hstate: &mut MixerHalfState,
planes: &[CudaSlice<f32>],
eps: f32,
push_raw: u64,
) -> Res<()> {
let (mixed, inj) = self.gate_read(e, ws, ptrs, attn_gate, planes, 1, eps, false)?;
let p = self.gdn_forward_half(e, ws, eps, h, &mixed, hstate, 1)?;
ws.put_f32("hc.mixed", mixed);
launch_push(e, &p, push_raw, self.hidden)?;
ws.put_f32("mixer.out", p);
put_inject(ws, inj);
Ok(())
}
#[allow(clippy::too_many_arguments)]
fn tp2_seg_b(
&self,
e: &Engine,
ws: &mut StepPool,
ptrs: &CudaSlice<u64>,
mlp_gate: &GateW,
planes: &mut [CudaSlice<f32>],
stage_recv: &CudaSlice<f32>,
rank0: bool,
eps_m: f32,
shared: Option<(
&CudaSlice<u8>,
&CudaSlice<u8>,
Option<&CudaSlice<f32>>,
usize,
)>,
) -> Res<()> {
let hidden = self.hidden;
let p = ws.take_f32(e, "mixer.out", hidden, 0)?;
let mut out = ws.take_f32(e, "join.out", hidden, 0)?;
if rank0 {
e.add(&p, stage_recv, &mut out, hidden)?;
} else {
e.add(stage_recv, &p, &mut out, hidden)?;
}
let inj = take_inject(e, ws, self.streams, 1)?;
self.gate_write(e, planes, ptrs, &out, &inj, 1)?;
ws.put_f32("mixer.out", p);
ws.put_f32("join.out", out);
put_inject(ws, inj);
let (mixed, injm) = self.gate_read(e, ws, ptrs, mlp_gate, planes, 1, eps_m, false)?;
if let Some((gu_b16, d_b16, ig, sffh)) = shared {
let (sh, gg) = self.tp2_shared_half(e, ws, gu_b16, d_b16, ig, &mixed, sffh, 1)?;
ws.put_f32("moe.sh_down", sh);
if let Some(gg) = gg {
ws.put_f32("moe.g", gg);
}
}
ws.put_f32("hc.mixed", mixed);
put_inject(ws, injm);
Ok(())
}
#[allow(clippy::too_many_arguments)]
fn tp2_seg_exit(
&self,
e: &Engine,
ws: &mut StepPool,
ptrs: &CudaSlice<u64>,
gate: &GateW,
planes: &[CudaSlice<f32>],
head_b16: &CudaSlice<u8>,
out_f: usize,
rows: usize,
) -> Res<()> {
let x = self
.gate_read_inner(e, ws, ptrs, gate, planes, rows, self.exit_eps, false, false)?
.0;
let mut logits = ws.take_f32(e, "logits", rows * out_f, rows * out_f)?;
launch_qmatvec_bf16w(
e,
head_b16,
&x,
&mut logits,
self.hidden,
out_f,
rows,
1,
0,
0,
self.hidden,
0,
)?;
ws.put_f32("hc.mixed", x);
ws.put_f32("logits", logits);
Ok(())
}
}
#[allow(clippy::too_many_arguments)]
fn launch_nvfp4_sel_matvec_pack(
e: &Engine,
codes: &CudaSlice<u8>,
scales: &CudaSlice<u8>,
macros_dev: &CudaSlice<f32>,
pack_raw: u64,
max_sel: usize,
x: &CudaSlice<f32>,
y: &mut CudaSlice<f32>,
in_f: usize,
out_f: usize,
x_stride: usize,
) -> Res<()> {
if in_f % 32 != 0 || out_f % 4 != 0 {
return Err(
"qmatvec_nvfp4_modelopt_sel_f32_v3c: geometry needs in_f%32==0 && out_f%4==0".into(),
);
}
let f = e.func("qmatvec_nvfp4_modelopt_sel_f32_v3c");
let cfg = LaunchConfig {
grid_dim: ((out_f / 4) as u32, max_sel as u32, 1),
block_dim: (32, 1, 1),
shared_mem_bytes: 0,
};
let (inf, outf, ms) = (in_f as i32, out_f as i32, max_sel as i32);
let xs = x_stride as i64;
let stream = e.gpu.stream();
let mut b = stream.launch_builder(&f);
b.arg(codes)
.arg(scales)
.arg(macros_dev)
.arg(&pack_raw)
.arg(&ms)
.arg(x)
.arg(y)
.arg(&inf)
.arg(&outf)
.arg(&xs);
unsafe {
b.launch(cfg)?;
}
Ok(())
}
fn launch_axpy_rows_seq_pack(
e: &Engine,
x: &CudaSlice<f32>,
pack_raw: u64,
max_sel: usize,
y: &mut CudaSlice<f32>,
width: usize,
) -> Res<()> {
let f = e.func("axpy_rows_seq_pack_f32");
let cfg = LaunchConfig::for_num_elems(width as u32);
let (ms, wi) = (max_sel as i32, width as i32);
let stream = e.gpu.stream();
let mut b = stream.launch_builder(&f);
b.arg(x).arg(&pack_raw).arg(&ms).arg(y).arg(&wi);
unsafe {
b.launch(cfg)?;
}
Ok(())
}
fn tp2_pack_bytes(sel: &[i32], w: &[f32], max_sel: usize) -> Vec<u8> {
let mut out = Vec::with_capacity((2 * max_sel + 1) * 4);
for i in 0..max_sel {
out.extend_from_slice(&sel.get(i).copied().unwrap_or(0).to_le_bytes());
}
for i in 0..max_sel {
out.extend_from_slice(&w.get(i).copied().unwrap_or(0.0).to_le_bytes());
}
out.extend_from_slice(&(sel.len() as i32).to_le_bytes());
out
}
impl Qwen4ExpGpu {
#[allow(clippy::too_many_arguments)]
fn tp2_seg_c(
&self,
e: &Engine,
ws: &mut StepPool,
gate: (&CudaSlice<u8>, &CudaSlice<u8>, &CudaSlice<f32>),
up: (&CudaSlice<u8>, &CudaSlice<u8>, &CudaSlice<f32>),
down: (&CudaSlice<u8>, &CudaSlice<u8>, &CudaSlice<f32>),
ff: usize,
max_sel: usize,
shared_compute: Option<(
&CudaSlice<u8>,
&CudaSlice<u8>,
Option<&CudaSlice<f32>>,
usize,
)>,
shared_gated: bool,
push_raw: u64,
) -> Res<()> {
let hidden = self.hidden;
let pack_raw = {
let pack = ws.peek_u8("moe.pack")?;
let stream = e.gpu.stream();
pack.device_ptr(&stream).0
};
let mixed = ws.take_f32(e, "hc.mixed", hidden, 0)?;
let mut act = ws.take_f32(e, "moe.act", max_sel * ff, 0)?;
if sel_gufuse_on() && hidden % 32 == 0 && ff % 4 == 0 {
launch_nvfp4_sel_gu_silu(
e, gate, up, None, pack_raw, max_sel, &mixed, &mut act, hidden, ff, None,
)?;
} else {
let mut yg = ws.take_f32(e, "moe.yg", max_sel * ff, 0)?;
let mut yu = ws.take_f32(e, "moe.yu", max_sel * ff, 0)?;
launch_nvfp4_sel_matvec_pack(
e, gate.0, gate.1, gate.2, pack_raw, max_sel, &mixed, &mut yg, hidden, ff, 0,
)?;
launch_nvfp4_sel_matvec_pack(
e, up.0, up.1, up.2, pack_raw, max_sel, &mixed, &mut yu, hidden, ff, 0,
)?;
e.silu_mul(&yg, &yu, &mut act, max_sel * ff)?;
ws.put_f32("moe.yg", yg);
ws.put_f32("moe.yu", yu);
}
let mut partial = ws.take_f32(e, "moe.partial", max_sel * hidden, 0)?;
launch_nvfp4_sel_matvec_pack(
e,
down.0,
down.1,
down.2,
pack_raw,
max_sel,
&act,
&mut partial,
ff,
hidden,
ff,
)?;
let mut r = ws.take_f32(e, "moe.out", hidden, 0)?;
launch_axpy_rows_seq_pack(e, &partial, pack_raw, max_sel, &mut r, hidden)?;
ws.put_f32("moe.act", act);
ws.put_f32("moe.partial", partial);
let (sh, g) = match shared_compute {
Some((gu_b16, d_b16, ig, sffh)) => {
self.tp2_shared_half(e, ws, gu_b16, d_b16, ig, &mixed, sffh, 1)?
}
None => {
let sh = ws.take_f32(e, "moe.sh_down", hidden, 0)?;
let g = if shared_gated {
Some(ws.take_f32(e, "moe.g", 1, 0)?)
} else {
None
};
(sh, g)
}
};
match g.as_ref() {
Some(g) => e.add_scaled_rows(&sh, g, &mut r, hidden, 1)?,
None => {
let mut view = r.slice_mut(0..hidden);
e.axpy_into(&sh, 1.0, &mut view, hidden)?;
}
}
ws.put_f32("moe.sh_down", sh);
if let Some(g) = g {
ws.put_f32("moe.g", g);
}
launch_push(e, &r, push_raw, hidden)?;
ws.put_f32("moe.out", r);
ws.put_f32("hc.mixed", mixed);
Ok(())
}
#[allow(clippy::too_many_arguments)]
fn tp2_seg_d(
&self,
e: &Engine,
ws: &mut StepPool,
ptrs: &CudaSlice<u64>,
planes: &mut [CudaSlice<f32>],
stage_recv: &CudaSlice<f32>,
rank0: bool,
) -> Res<()> {
let hidden = self.hidden;
let mp = ws.take_f32(e, "moe.out", hidden, 0)?;
let mut mo = ws.take_f32(e, "join.out", hidden, 0)?;
if rank0 {
e.add(&mp, stage_recv, &mut mo, hidden)?;
} else {
e.add(stage_recv, &mp, &mut mo, hidden)?;
}
let injm = take_inject(e, ws, self.streams, 1)?;
self.gate_write(e, planes, ptrs, &mo, &injm, 1)?;
ws.put_f32("moe.out", mp);
ws.put_f32("join.out", mo);
put_inject(ws, injm);
Ok(())
}
}
#[cfg(test)]
mod sel_group_tests {
use super::*;
static SEAM: std::sync::Mutex<()> = std::sync::Mutex::new(());
#[test]
fn auto_resolves_the_measured_serving_shapes() {
assert_eq!(sel_group_resolve(SEL_GROUP_AUTO, 640, 2560), Some((4, 4)));
assert_eq!(sel_group_resolve(SEL_GROUP_AUTO, 2560, 640), Some((16, 4)));
}
#[test]
fn off_and_odd_in_f_take_the_shipped_kernel() {
assert_eq!(sel_group_resolve(SEL_GROUP_OFF, 640, 2560), None);
assert_eq!(sel_group_resolve(SEL_GROUP_AUTO, 48, 2560), None);
}
#[test]
fn auto_backs_off_rows_then_refuses_rather_than_tiling_raggedly() {
assert_eq!(sel_group_resolve(SEL_GROUP_AUTO, 64, 32), Some((2, 2)));
assert_eq!(sel_group_resolve(SEL_GROUP_AUTO, 64, 24), None);
for &(in_f, out_f) in &[(640usize, 2560usize), (2560, 640), (32, 32), (64, 16)] {
let (g, rows) = sel_group_resolve(SEL_GROUP_AUTO, in_f, out_f)
.unwrap_or_else(|| panic!("auto refused {in_f}x{out_f}"));
assert_eq!(
out_f % ((32 / g) * rows),
0,
"{in_f}x{out_f} -> g{g} rows{rows}"
);
}
}
#[test]
fn explicit_pins_are_verbatim_and_still_tile_checked() {
let _g = SEAM.lock().unwrap();
assert!(set_sel_group("dn:8:1+gu:16:4"));
assert_eq!(sel_group_resolve(sel_group_dn(), 640, 2560), Some((8, 1)));
assert_eq!(sel_group_resolve(sel_group_gu(), 2560, 640), Some((16, 4)));
assert!(set_sel_group("dn:1:4"));
assert_eq!(sel_group_resolve(sel_group_dn(), 2560, 640), Some((1, 4)));
assert_eq!(sel_group_resolve(sel_group_dn(), 2560, 96), None);
set_sel_group("off");
}
#[test]
fn malformed_specs_apply_nothing_and_refuse() {
let _g = SEAM.lock().unwrap();
assert!(set_sel_group("dn:4:4+gu:16:4"));
let before = sel_group_spec();
for bad in [
"dn:3:4", "dn:4:3", "dn:64:4", "dn:4", "xx:4:4", "dn:4:4+xx:1", "+",
] {
assert!(!set_sel_group(bad), "{bad:?} was accepted");
assert_eq!(
sel_group_spec(),
before,
"{bad:?} mutated state while refusing"
);
}
set_sel_group("off");
}
#[test]
fn seam_round_trips_through_the_boolean_harness() {
let _g = SEAM.lock().unwrap();
set_sel_group("off");
assert_eq!(seam_state("selgroup"), Some(false));
assert!(set_seam("selgroup", true, None));
assert_eq!(seam_state("selgroup"), Some(true));
assert_eq!(sel_group_spec(), "dn:auto+gu:auto");
assert!(set_seam("selgroup", false, None));
assert_eq!(seam_state("selgroup"), Some(false));
assert_eq!(sel_group_spec(), "dn:off+gu:off");
assert!(set_sel_group("dn:4:4+gu:off"));
assert_eq!(seam_state("selgroup"), Some(true));
set_sel_group("off");
assert!(seam_names().contains(&"selgroup"));
assert!(seam_exists("selgroup"));
}
}
#[cfg(test)]
mod tp2_placement_tests {
use super::{LayerPlacement, Tp2Placement};
fn write_map(name: &str, body: &str) -> std::path::PathBuf {
let path = std::env::temp_dir().join(format!(
"memra-q4e-ep-map-{name}-{}.json",
std::process::id()
));
std::fs::write(&path, body).expect("write map fixture");
path
}
fn load(name: &str, body: &str, expert_count: usize) -> Result<Tp2Placement, String> {
let path = write_map(name, body);
let out = Tp2Placement::load(&path, expert_count).map_err(|e| e.to_string());
let _ = std::fs::remove_file(&path);
out
}
fn doc(body: &str) -> String {
format!(
"{{\"format\": \"memra-ep-map-v1\", \"strategy\": \"coactivation\", \
\"ranks\": 2, \"entry_rank\": 0, \"expert_count\": 4, {body}}}"
)
}
fn assert_refuses(name: &str, body: &str, expert_count: usize, clause: &str) {
match load(name, body, expert_count) {
Ok(_) => panic!("{name}: expected a refusal naming {clause:?}, but the map loaded"),
Err(msg) => {
assert!(
msg.contains(clause),
"{name}: refusal must name the broken clause {clause:?}, got: {msg}"
);
assert!(
msg.contains("MEMRA_Q4E_EP_MAP"),
"{name}: refusal must name the flag/file, got: {msg}"
);
}
}
}
#[test]
fn even_split_is_the_contiguous_suffix_control_arm() {
let p = Tp2Placement::even(512);
assert_eq!(p.strategy(), "even");
assert_eq!(p.entry_rank(), 0);
assert!(p.source().contains("MEMRA_Q4E_EP_MAP unset"));
let l = p.layer(0, 512).expect("even split resolves every layer");
assert!(l.is_even(), "the built-in even split must report as even");
assert_eq!(l.card1.len(), 256);
assert_eq!(l.card1, (256u32..512).collect::<Vec<_>>());
for e in 0..256 {
assert_eq!(l.rank(e), 0, "expert {e} belongs to card 0");
assert_eq!(l.local(e), e, "card-0 local slot IS the global id");
}
for (slot, e) in (256..512).enumerate() {
assert_eq!(l.rank(e), 1, "expert {e} belongs to card 1");
assert_eq!(l.local(e), slot, "card-1 local slot is the gather position");
}
let l47 = p.layer(47, 512).expect("layer 47");
assert_eq!(l47.card1, l.card1);
}
#[test]
fn measured_map_resolves_ascending_gather_and_local_slots() {
let body = "\"layers\": [{\"layer\": 0, \"assignment\": [1, 0, 0, 1]}]";
let p = load("measured", &doc(body), 4).expect("balanced map loads");
assert_eq!(p.strategy(), "coactivation");
let l = p.layer(0, 4).expect("layer 0");
assert!(
!l.is_even(),
"a non-suffix placement is not the control arm"
);
assert_eq!(l.card1, vec![0u32, 3]);
assert_eq!((l.rank(0), l.rank(1), l.rank(2), l.rank(3)), (1, 0, 0, 1));
assert_eq!(l.local(0), 0);
assert_eq!(l.local(3), 1);
assert_eq!(l.local(1), 1);
assert_eq!(l.local(2), 2);
}
#[test]
fn is_even_rejects_a_balanced_permutation() {
let body = "\"layers\": [{\"layer\": 0, \"assignment\": [0, 1, 1, 0]}]";
let p = load("perm", &doc(body), 4).expect("balanced map loads");
let l = p.layer(0, 4).expect("layer 0");
assert_eq!(l.card1, vec![1u32, 2]);
assert!(!l.is_even());
}
#[test]
fn an_explicit_even_map_matches_the_builtin_even_split() {
let body = "\"layers\": [{\"layer\": 0, \"assignment\": [0, 0, 1, 1]}]";
let p = load("explicit-even", &doc(body), 4).expect("even map loads");
let l = p.layer(0, 4).expect("layer 0");
let builtin = Tp2Placement::even(4).layer(0, 4).expect("builtin");
assert!(l.is_even());
assert_eq!(l.card1, builtin.card1);
for e in 0..4 {
assert_eq!(l.rank(e), builtin.rank(e), "rank of expert {e}");
assert_eq!(l.local(e), builtin.local(e), "local slot of expert {e}");
}
}
#[test]
fn refuses_a_foreign_format() {
let body = "\"layers\": [{\"layer\": 0, \"assignment\": [0, 0, 1, 1]}]";
let text = format!(
"{{\"format\": \"memra-ep-map-v2\", \"ranks\": 2, \"expert_count\": 4, {body}}}"
);
assert_refuses("format", &text, 4, "memra-ep-map-v1");
}
#[test]
fn refuses_a_rank_count_that_is_not_two() {
let text = "{\"format\": \"memra-ep-map-v1\", \"ranks\": 4, \"expert_count\": 4, \
\"layers\": [{\"layer\": 0, \"assignment\": [0, 0, 1, 1]}]}";
assert_refuses("ranks", text, 4, "exactly two cards");
}
#[test]
fn refuses_an_expert_count_that_is_not_the_plans() {
let body = "\"layers\": [{\"layer\": 0, \"assignment\": [0, 0, 1, 1]}]";
assert_refuses("experts", &doc(body), 8, "expert_count=4");
}
#[test]
fn refuses_an_entry_rank_outside_the_two_cards() {
let text = "{\"format\": \"memra-ep-map-v1\", \"ranks\": 2, \"entry_rank\": 2, \
\"expert_count\": 4, \
\"layers\": [{\"layer\": 0, \"assignment\": [0, 0, 1, 1]}]}";
assert_refuses("entry", text, 4, "entry_rank=2");
}
#[test]
fn refuses_a_document_with_no_layers_array() {
assert_refuses("nolayers", &doc("\"strategy2\": 0"), 4, "no `layers` array");
}
#[test]
fn refuses_an_empty_layers_array() {
assert_refuses(
"emptylayers",
&doc("\"layers\": []"),
4,
"`layers` is empty",
);
}
#[test]
fn refuses_an_assignment_of_the_wrong_length() {
let body = "\"layers\": [{\"layer\": 0, \"assignment\": [0, 1]}]";
assert_refuses("shortassign", &doc(body), 4, "expected 4");
}
#[test]
fn refuses_a_rank_id_outside_the_two_cards() {
let body = "\"layers\": [{\"layer\": 0, \"assignment\": [0, 0, 1, 7]}]";
assert_refuses("badrank", &doc(body), 4, "expert 3");
}
#[test]
fn refuses_an_unbalanced_layer_and_names_the_rebalance_knob() {
let body = "\"layers\": [{\"layer\": 0, \"assignment\": [0, 1, 1, 1]}]";
assert_refuses("unbalanced", &doc(body), 4, "card 1 owns 3 experts");
let body = "\"layers\": [{\"layer\": 0, \"assignment\": [0, 1, 1, 1]}]";
assert_refuses("unbalanced2", &doc(body), 4, "--balance-tolerance");
}
#[test]
fn refuses_a_layer_the_map_does_not_cover() {
let body = "\"layers\": [{\"layer\": 0, \"assignment\": [0, 0, 1, 1]}]";
let p = load("partial", &doc(body), 4).expect("map loads");
assert!(p.layer(0, 4).is_ok(), "the covered layer resolves");
let msg = p
.layer(1, 4)
.expect_err("an uncovered MoE layer must refuse")
.to_string();
assert!(msg.contains("does not cover MoE layer 1"), "got: {msg}");
assert!(
msg.contains("partly-applied map is not a placement"),
"the refusal must say WHY it is fail-closed, got: {msg}"
);
}
#[test]
fn refuses_a_layer_whose_expert_count_disagrees_with_the_map() {
let p = Tp2Placement::even(512);
let msg = p
.layer(0, 256)
.expect_err("a layer geometry the map is not for must refuse")
.to_string();
assert!(msg.contains("map is for 512"), "got: {msg}");
}
#[test]
fn refuses_an_unreadable_map_path() {
let missing = std::env::temp_dir().join(format!(
"memra-q4e-ep-map-absent-{}.json",
std::process::id()
));
let _ = std::fs::remove_file(&missing);
let msg = Tp2Placement::load(&missing, 4)
.expect_err("an unreadable map must refuse at the load preflight")
.to_string();
assert!(msg.contains("MEMRA_Q4E_EP_MAP"), "got: {msg}");
}
#[test]
fn refuses_an_odd_routed_bank_on_both_paths() {
let msg = Tp2Placement::even(5)
.layer(0, 5)
.expect_err("an odd bank has no two-card placement")
.to_string();
assert!(msg.contains("ODD"), "got: {msg}");
assert!(msg.contains("EQUAL-size"), "got: {msg}");
let balanced_odd = "{\"format\": \"memra-ep-map-v1\", \"ranks\": 2, \
\"expert_count\": 5, \
\"layers\": [{\"layer\": 0, \"assignment\": [0, 0, 0, 1, 1]}]}";
assert_refuses("odd-balanced", balanced_odd, 5, "ODD");
let unbalanced_odd = "{\"format\": \"memra-ep-map-v1\", \"ranks\": 2, \
\"expert_count\": 5, \
\"layers\": [{\"layer\": 0, \"assignment\": [0, 0, 1, 1, 1]}]}";
assert_refuses("odd-unbalanced", unbalanced_odd, 5, "ODD");
}
#[test]
fn local_slots_are_a_bijection_at_the_serving_geometry() {
let experts = 512usize;
let half = experts / 2;
let assignment: Vec<String> = (0..experts).map(|e| (e % 2).to_string()).collect();
let body = format!(
"\"layers\": [{{\"layer\": 0, \"assignment\": [{}]}}]",
assignment.join(", ")
);
let text = format!(
"{{\"format\": \"memra-ep-map-v1\", \"ranks\": 2, \"expert_count\": {experts}, \
{body}}}"
);
let p = load("bijection", &text, experts).expect("balanced alternating map loads");
let l: LayerPlacement = p.layer(0, experts).expect("layer 0");
assert_eq!(l.card1.len(), half, "card 1 must own exactly half the bank");
assert!(
l.card1.windows(2).all(|w| w[0] < w[1]),
"the gather order must be strictly ascending"
);
let mut seen = vec![false; half];
for e in 0..experts {
match l.rank(e) {
0 => assert_eq!(l.local(e), e, "card-0 slot is the global id"),
1 => {
let slot = l.local(e);
assert!(slot < half, "card-1 slot {slot} outside its half-bank");
assert!(!seen[slot], "card-1 slot {slot} claimed twice");
seen[slot] = true;
}
r => panic!("expert {e} has rank {r}"),
}
}
assert!(seen.into_iter().all(|s| s), "card-1 slots must be dense");
assert!(!l.is_even());
}
}