use crate::Engine;
use crate::dsv4_ffi as k;
use crate::dsv4_ffi::ck;
use crate::model::GpuTensor;
use cudarc::driver::{CudaSlice, CudaStream, DevicePtr, DevicePtrMut};
use memra_gguf::model_plan::{HcCollapse, ModelPlan, ResidualTopology};
use memra_gguf::source::TensorSource;
use std::os::raw::c_void;
type Res<T> = Result<T, Box<dyn std::error::Error>>;
fn sp(stream: &CudaStream) -> *mut c_void {
stream.cu_stream() as *mut c_void
}
macro_rules! dpf {
($slice:expr, $stream:expr) => {{ $slice.device_ptr($stream).0 as *const f32 }};
}
macro_rules! dpm {
($slice:expr, $stream:expr) => {{ $slice.device_ptr_mut($stream).0 as *mut f32 }};
}
#[derive(Debug, Clone, Copy, PartialEq)]
pub struct HyperTopology {
pub streams: usize,
pub epsilon: f32,
pub sinkhorn_iterations: u32,
pub collapse: HcCollapse,
}
impl HyperTopology {
pub fn rows(&self) -> usize {
(2 + self.streams) * self.streams
}
pub fn from_plan(plan: &ModelPlan) -> Result<Option<Self>, String> {
let mut found: Option<Self> = None;
for layer in &plan.layers {
let ResidualTopology::HyperConnections {
streams,
epsilon,
sinkhorn_iterations,
collapse,
} = layer.residual
else {
if found.is_some() {
return Err(format!(
"layer {} declares a serial/gemma residual while an earlier trunk layer \
declares HyperConnections; the topology must be uniform across the trunk",
layer.index
));
}
continue;
};
let this = Self {
streams: streams as usize,
epsilon,
sinkhorn_iterations,
collapse,
};
if streams == 0 || epsilon <= 0.0 || sinkhorn_iterations == 0 {
return Err(format!(
"layer {}: HyperConnections need streams > 0, epsilon > 0 and \
sinkhorn_iterations > 0, got streams={streams} epsilon={epsilon} \
iterations={sinkhorn_iterations}",
layer.index
));
}
match found {
None if layer.index != plan.layers[0].index => {
return Err(format!(
"layer {} declares HyperConnections but earlier trunk layers do not; the \
topology must be uniform across the trunk",
layer.index
));
}
None => found = Some(this),
Some(first) if first != this => {
return Err(format!(
"layer {} declares {this:?} but the trunk opened with {first:?}; the \
topology must be uniform across the trunk",
layer.index
));
}
Some(_) => {}
}
}
Ok(found)
}
}
pub struct HyperSite {
pub fn_w: CudaSlice<f32>,
pub base: CudaSlice<f32>,
pub scale: CudaSlice<f32>,
}
pub struct HyperLayer {
pub attn: HyperSite,
pub mlp: HyperSite,
}
pub struct HyperHead {
pub fn_w: CudaSlice<f32>,
pub base: CudaSlice<f32>,
pub scale: CudaSlice<f32>,
}
fn float_data<'a>(name: &str, t: &'a GpuTensor, want: usize) -> Result<&'a CudaSlice<f32>, String> {
let data = match t {
GpuTensor::Float { data, .. } => data,
GpuTensor::Quant { .. } => {
return Err(format!(
"{name}: hyper-connection parameters must be f32-resident, got a quantized \
tensor; re-mint this tensor unquantized (the whole hc program is an f32 island)"
));
}
GpuTensor::FloatBf16 { .. } => {
return Err(format!(
"{name}: hyper-connection parameters must be f32-resident, got a bf16-resident \
matmul weight"
));
}
};
if data.len() != want {
return Err(format!(
"{name}: {} elements, the plan's HyperConnections require {want}",
data.len()
));
}
Ok(data)
}
fn load_site(
e: &Engine,
src: &dyn TensorSource,
il: u32,
topology: &HyperTopology,
hidden: usize,
site: &str,
) -> Res<HyperSite> {
let rows = topology.rows();
let width = topology.streams * hidden;
let mut out: Vec<CudaSlice<f32>> = Vec::with_capacity(3);
for (suffix, want) in [
("fn", rows * width),
("base", rows),
("scale", 3),
] {
let name = format!("blk.{il}.{site}_{suffix}");
if !src.has(&name) {
return Err(format!(
"{name} is absent, but the compiled ModelPlan declares \
ResidualTopology::HyperConnections{{ streams: {} }} for layer {il}. Refusing to \
load: a serial residual would compute a different model, silently.",
topology.streams
)
.into());
}
let loaded = GpuTensor::load_from_source(e, src, &name)?;
out.push(e.clone_dtod(float_data(&name, &loaded, want)?)?);
}
let mut out = out.into_iter();
Ok(HyperSite {
fn_w: out.next().expect("function"),
base: out.next().expect("base"),
scale: out.next().expect("scale"),
})
}
impl HyperLayer {
pub fn load(
e: &Engine,
src: &dyn TensorSource,
il: u32,
topology: &HyperTopology,
hidden: usize,
) -> Res<Self> {
Ok(Self {
attn: load_site(e, src, il, topology, hidden, "hc_attn")?,
mlp: load_site(e, src, il, topology, hidden, "hc_ffn")?,
})
}
}
impl HyperHead {
pub fn load(
e: &Engine,
src: &dyn TensorSource,
topology: &HyperTopology,
hidden: usize,
) -> Res<Option<Self>> {
if topology.collapse != HcCollapse::GatedHead {
return Ok(None);
}
let streams = topology.streams;
let mut out: Vec<CudaSlice<f32>> = Vec::with_capacity(3);
for (name, want) in [
("hc_head_fn", streams * streams * hidden),
("hc_head_base", streams),
("hc_head_scale", 1),
] {
if !src.has(name) {
return Err(format!(
"{name} is absent, but the compiled ModelPlan declares \
HcCollapse::GatedHead. Refusing to load: collapsing with an unweighted mean \
instead would compute a different model, silently."
)
.into());
}
let loaded = GpuTensor::load_from_source(e, src, name)?;
out.push(e.clone_dtod(float_data(name, &loaded, want)?)?);
}
let mut out = out.into_iter();
Ok(Some(Self {
fn_w: out.next().expect("function"),
base: out.next().expect("base"),
scale: out.next().expect("scale"),
}))
}
}
pub struct HcMix {
pub post: CudaSlice<f32>,
pub comb: CudaSlice<f32>,
}
pub static HC_FUSED_PRE_DISPATCHES: std::sync::atomic::AtomicU64 =
std::sync::atomic::AtomicU64::new(0);
pub static HC_FUSED_PRE_V2_DISPATCHES: std::sync::atomic::AtomicU64 =
std::sync::atomic::AtomicU64::new(0);
#[derive(Clone, Copy, PartialEq, Eq, Debug)]
pub enum HcFusedPreArm {
Off,
V1,
V2,
}
fn hc_fused_pre_arm() -> HcFusedPreArm {
hc_fused_pre_arm_from(
std::env::var("MEMRA_HC_FUSED_PRE").ok().as_deref(),
env!("MEMRA_BUILT_CUDA_ARCH"),
)
}
pub fn hc_fused_pre_arm_from(v: Option<&str>, built_arch: &str) -> HcFusedPreArm {
match v.map(str::trim) {
Some("1") => HcFusedPreArm::V1,
Some("2") => HcFusedPreArm::V2,
Some("0") => HcFusedPreArm::Off,
_ if built_arch == "100a" => HcFusedPreArm::V2,
_ => HcFusedPreArm::Off,
}
}
pub fn expand(
e: &Engine,
topology: &HyperTopology,
embedded: &CudaSlice<f32>,
t: usize,
hidden: usize,
) -> Res<CudaSlice<f32>> {
let streams = topology.streams;
let mut out = e.uninit(t * streams * hidden)?;
let stream = e.stream();
unsafe {
ck(
"hc_expand",
k::memra_dsv4_hc_expand(
dpf!(embedded, &stream),
dpm!(out, &stream),
t as i32,
streams as i32,
hidden as i32,
sp(&stream),
),
)?;
}
Ok(out)
}
pub fn pre(
e: &Engine,
topology: &HyperTopology,
site: &HyperSite,
x: &CudaSlice<f32>,
t: usize,
hidden: usize,
) -> Res<(CudaSlice<f32>, HcMix)> {
let width = topology.streams * hidden;
let mixes = e.linear(x, &site.fn_w, t, width, topology.rows())?;
pre_finish(e, topology, site, x, mixes, t, hidden)
}
pub fn pre_exact(
e: &Engine,
topology: &HyperTopology,
site: &HyperSite,
x: &CudaSlice<f32>,
t: usize,
hidden: usize,
) -> Res<(CudaSlice<f32>, HcMix)> {
let rows = topology.rows();
let width = topology.streams * hidden;
let mut mixes = e.uninit(t * rows)?;
for r in 0..t {
let xr = x.slice(r * width..(r + 1) * width);
let wv = site.fn_w.slice(0..site.fn_w.len());
let mut yr = mixes.slice_mut(r * rows..(r + 1) * rows);
hc_mixes_into(e, &xr, &wv, &mut yr, width, rows)
.map_err(|err| format!("hc pre_exact row {r}: {err}"))?;
}
pre_finish(e, topology, site, x, mixes, t, hidden)
}
fn pre_finish(
e: &Engine,
topology: &HyperTopology,
site: &HyperSite,
x: &CudaSlice<f32>,
mut mixes: CudaSlice<f32>,
t: usize,
hidden: usize,
) -> Res<(CudaSlice<f32>, HcMix)> {
let streams = topology.streams;
let mut pre_gates = e.uninit(t * streams)?;
let mut post = e.uninit(t * streams)?;
let mut comb = e.uninit(t * streams * streams)?;
let mut y = e.uninit(t * hidden)?;
pre_finish_into(
e,
topology,
site,
x,
&mut mixes,
&mut pre_gates,
&mut post,
&mut comb,
&mut y,
t,
hidden,
)?;
Ok((y, HcMix { post, comb }))
}
#[allow(clippy::too_many_arguments)] fn pre_finish_into(
e: &Engine,
topology: &HyperTopology,
site: &HyperSite,
x: &CudaSlice<f32>,
mixes: &mut CudaSlice<f32>,
pre_gates: &mut CudaSlice<f32>,
post: &mut CudaSlice<f32>,
comb: &mut CudaSlice<f32>,
y: &mut CudaSlice<f32>,
t: usize,
hidden: usize,
) -> Res<()> {
let streams = topology.streams;
let rows = topology.rows();
let width = streams * hidden;
let eps = topology.epsilon;
let stream = e.stream();
let fused_arm = hc_fused_pre_arm();
if fused_arm != HcFusedPreArm::Off && streams <= 8 {
let (label, rc) = unsafe {
match fused_arm {
HcFusedPreArm::V1 => (
"hc_pre_fused",
k::memra_dsv4_hc_pre_fused(
dpf!(x, &stream),
dpf!(mixes, &stream),
dpf!(site.scale, &stream),
dpf!(site.base, &stream),
dpm!(pre_gates, &stream),
dpm!(post, &stream),
dpm!(comb, &stream),
dpm!(y, &stream),
t as i32,
streams as i32,
hidden as i32,
topology.sinkhorn_iterations as i32,
eps,
std::ptr::null_mut(),
sp(&stream),
),
),
HcFusedPreArm::V2 if crate::hc_pre_block() != 128 || crate::hc_pre_sink_reg() => {
let v4 = if crate::hc_pre_v4_on() {
let rc = k::memra_dsv4_hc_pre_v4(
dpf!(x, &stream),
dpf!(mixes, &stream),
dpf!(site.scale, &stream),
dpf!(site.base, &stream),
dpm!(pre_gates, &stream),
dpm!(post, &stream),
dpm!(comb, &stream),
dpm!(y, &stream),
t as i32,
streams as i32,
hidden as i32,
topology.sinkhorn_iterations as i32,
eps,
std::ptr::null_mut(),
crate::hc_pre_block() as i32,
sp(&stream),
);
if rc == 40025 {
None
} else {
if rc == 0 {
HC_PRE_V4_DISPATCHES
.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
}
Some(("hc_pre_v4", rc))
}
} else {
None
};
match v4 {
Some(done) => done,
None => (
"hc_pre_fused_v3",
k::memra_dsv4_hc_pre_fused_v3(
dpf!(x, &stream),
dpf!(mixes, &stream),
dpf!(site.scale, &stream),
dpf!(site.base, &stream),
dpm!(pre_gates, &stream),
dpm!(post, &stream),
dpm!(comb, &stream),
dpm!(y, &stream),
t as i32,
streams as i32,
hidden as i32,
topology.sinkhorn_iterations as i32,
eps,
std::ptr::null_mut(),
crate::hc_pre_block() as i32,
crate::hc_pre_sink_reg() as i32,
sp(&stream),
),
),
}
}
HcFusedPreArm::V2 => (
"hc_pre_fused_v2",
k::memra_dsv4_hc_pre_fused_v2(
dpf!(x, &stream),
dpf!(mixes, &stream),
dpf!(site.scale, &stream),
dpf!(site.base, &stream),
dpm!(pre_gates, &stream),
dpm!(post, &stream),
dpm!(comb, &stream),
dpm!(y, &stream),
t as i32,
streams as i32,
hidden as i32,
topology.sinkhorn_iterations as i32,
eps,
std::ptr::null_mut(),
sp(&stream),
),
),
HcFusedPreArm::Off => unreachable!("guarded by the enclosing if"),
}
};
ck(label, rc)?;
let block = crate::hc_pre_block();
let (counter, tag) = match fused_arm {
HcFusedPreArm::V1 => (&HC_FUSED_PRE_DISPATCHES, "1"),
HcFusedPreArm::V2 => (&HC_FUSED_PRE_V2_DISPATCHES, "2"),
HcFusedPreArm::Off => unreachable!("guarded by the enclosing if"),
};
if counter.fetch_add(1, std::sync::atomic::Ordering::Relaxed) == 0 {
let kern = if fused_arm == HcFusedPreArm::V2
&& HC_PRE_V4_DISPATCHES.load(std::sync::atomic::Ordering::Relaxed) > 0
{
"hc_pre_v4"
} else if fused_arm == HcFusedPreArm::V2 && (block != 128 || crate::hc_pre_sink_reg()) {
"hc_pre_fused_v3"
} else if fused_arm == HcFusedPreArm::V2 {
"hc_pre_fused_v2"
} else {
"hc_pre_fused"
};
eprintln!(
"[hc-fused-pre] engaged streams={streams} hidden={hidden} t={t} arm={tag} \
kernel={kern} block={block} sinkhorn={} (one launch replaces rowsq_scale + \
sinkhorn + collapse per site; MEMRA_HC_FUSED_PRE={tag}, MEMRA_HC_PRE_BLOCK={block})",
if crate::hc_pre_sink_reg() {
"registers"
} else {
"shared"
}
);
}
return Ok(());
}
unsafe {
ck(
"hc rowsq_scale",
k::memra_dsv4_rowsq_scale(
dpf!(x, &stream),
dpm!(mixes, &stream),
t as i32,
width as i32,
rows as i32,
eps,
sp(&stream),
),
)?;
ck(
"hc_sinkhorn",
k::memra_dsv4_hc_sinkhorn_m(
dpf!(mixes, &stream),
dpf!(site.scale, &stream),
dpf!(site.base, &stream),
dpm!(pre_gates, &stream),
dpm!(post, &stream),
dpm!(comb, &stream),
t as i32,
streams as i32,
topology.sinkhorn_iterations as i32,
eps,
sp(&stream),
),
)?;
ck(
"hc_collapse",
k::memra_dsv4_hc_collapse(
dpf!(x, &stream),
dpf!(pre_gates, &stream),
dpm!(y, &stream),
t as i32,
streams as i32,
hidden as i32,
sp(&stream),
),
)?;
}
Ok(())
}
pub struct HyperDecodeWs {
pub mixes: CudaSlice<f32>,
pub pre: CudaSlice<f32>,
pub post: CudaSlice<f32>,
pub comb: CudaSlice<f32>,
pub y: CudaSlice<f32>,
pub h: CudaSlice<f32>,
pub z: CudaSlice<f32>,
pub xb: CudaSlice<f32>,
streams: usize,
hidden: usize,
}
impl HyperDecodeWs {
pub fn new(e: &Engine, topology: &HyperTopology, hidden: usize) -> Res<Self> {
let streams = topology.streams;
Ok(Self {
mixes: e.uninit(topology.rows())?,
pre: e.uninit(streams)?,
post: e.uninit(streams)?,
comb: e.uninit(streams * streams)?,
y: e.uninit(hidden)?,
h: e.uninit(hidden)?,
z: e.uninit(hidden)?,
xb: e.uninit(streams * hidden)?,
streams,
hidden,
})
}
pub fn matches(&self, topology: &HyperTopology, hidden: usize) -> bool {
self.streams == topology.streams && self.hidden == hidden
}
}
pub fn pre_t1_ws(
e: &Engine,
topology: &HyperTopology,
site: &HyperSite,
x: &CudaSlice<f32>,
ws: &mut HyperDecodeWs,
hidden: usize,
) -> Res<()> {
let rows = topology.rows();
let width = topology.streams * hidden;
{
let xr = x.slice(0..width);
let wv = site.fn_w.slice(0..site.fn_w.len());
let mut yr = ws.mixes.slice_mut(0..rows);
hc_mixes_into(e, &xr, &wv, &mut yr, width, rows)
.map_err(|err| format!("hc pre_t1_ws mixes: {err}"))?;
}
let ws = &mut *ws;
pre_finish_into(
e,
topology,
site,
x,
&mut ws.mixes,
&mut ws.pre,
&mut ws.post,
&mut ws.comb,
&mut ws.y,
1,
hidden,
)
}
#[derive(Clone, Copy, PartialEq, Eq, Debug)]
pub enum NormDst {
H,
Z,
}
pub static HC_MIXES_KERNEL_DISPATCHES: std::sync::atomic::AtomicU64 =
std::sync::atomic::AtomicU64::new(0);
fn hc_mixes_into(
e: &Engine,
x: &cudarc::driver::CudaView<'_, f32>,
w: &cudarc::driver::CudaView<'_, f32>,
y: &mut cudarc::driver::CudaViewMut<'_, f32>,
in_f: usize,
out_f: usize,
) -> Result<(), Box<dyn std::error::Error>> {
if Engine::hc_mixes_kernel_on() && e.hc_mixes_gemv_into(x, w, y, in_f, out_f)? {
if HC_MIXES_KERNEL_DISPATCHES.fetch_add(1, std::sync::atomic::Ordering::Relaxed) == 0 {
eprintln!(
"[hc-mixes-kernel] engaged in_f={in_f} out_f={out_f} (native hc_mixes_gemv_f32 \
in place of cuBLASLt dot+reduce; MEMRA_HC_MIXES_KERNEL=1, numeric class)"
);
}
return Ok(());
}
e.linear_t1_into(x, w, y, in_f, out_f)
}
pub static HC_PRE_ZQ8_DISPATCHES: std::sync::atomic::AtomicU64 =
std::sync::atomic::AtomicU64::new(0);
pub static HC_PRE_V4_DISPATCHES: std::sync::atomic::AtomicU64 =
std::sync::atomic::AtomicU64::new(0);
pub static HC_PRE_V4Z_DISPATCHES: std::sync::atomic::AtomicU64 =
std::sync::atomic::AtomicU64::new(0);
fn hc_pre_zq8_selfcheck() -> bool {
static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
*ON.get_or_init(|| std::env::var("MEMRA_HC_PRE_ZQ8").as_deref() == Ok("2"))
}
static HC_PRE_ZQ8_CHECK_SITES: std::sync::atomic::AtomicU64 = std::sync::atomic::AtomicU64::new(0);
static HC_PRE_ZQ8_CHECK_BAD: std::sync::atomic::AtomicU64 = std::sync::atomic::AtomicU64::new(0);
#[allow(clippy::too_many_arguments)]
fn pre_t1_ws_zq8_selfcheck(
e: &Engine,
topology: &HyperTopology,
site: &HyperSite,
x: &CudaSlice<f32>,
ws: &mut HyperDecodeWs,
hidden: usize,
norm_w: &CudaSlice<f32>,
dst: NormDst,
eps_norm: f32,
) -> Res<()> {
let streams = topology.streams;
let rows = topology.rows();
let width = streams * hidden;
let block = crate::hc_pre_block();
let rms_bd = crate::rms_block() as usize;
let stream = e.stream();
{
let xr = x.slice(0..width);
let wv = site.fn_w.slice(0..site.fn_w.len());
let mut yr = ws.mixes.slice_mut(0..rows);
e.linear_t1_into(&xr, &wv, &mut yr, width, rows)
.map_err(|err| format!("hc pre_t1_ws_zq8 selfcheck mixes: {err}"))?;
}
let mut s_pre = e.uninit(streams)?;
let mut s_post = e.uninit(streams)?;
let mut s_comb = e.uninit(streams * streams)?;
let mut s_y = e.uninit(hidden)?;
let mut s_z = e.uninit(hidden)?;
let mut s_q = e.uninit_i8(hidden)?;
let mut s_d = e.uninit(hidden / 32)?;
unsafe {
ck(
"hc_pre_zq8 (selfcheck fused)",
if crate::hc_pre_v4z_on() {
k::memra_dsv4_hc_pre_v4z(
dpf!(x, &stream),
dpf!(&ws.mixes, &stream),
dpf!(site.scale, &stream),
dpf!(site.base, &stream),
dpm!(&mut s_pre, &stream),
dpm!(&mut s_post, &stream),
dpm!(&mut s_comb, &stream),
dpm!(&mut s_y, &stream),
1,
streams as i32,
hidden as i32,
topology.sinkhorn_iterations as i32,
topology.epsilon,
std::ptr::null_mut(),
dpf!(norm_w, &stream),
dpm!(&mut s_z, &stream),
s_q.device_ptr_mut(&stream).0 as *mut i8,
s_d.device_ptr_mut(&stream).0 as *mut f32,
eps_norm,
rms_bd as i32,
sp(&stream),
)
} else {
k::memra_dsv4_hc_pre_zq8(
dpf!(x, &stream),
dpf!(&ws.mixes, &stream),
dpf!(site.scale, &stream),
dpf!(site.base, &stream),
dpm!(&mut s_pre, &stream),
dpm!(&mut s_post, &stream),
dpm!(&mut s_comb, &stream),
dpm!(&mut s_y, &stream),
1,
streams as i32,
hidden as i32,
topology.sinkhorn_iterations as i32,
topology.epsilon,
std::ptr::null_mut(),
block as i32,
crate::hc_pre_sink_reg() as i32,
dpf!(norm_w, &stream),
dpm!(&mut s_z, &stream),
s_q.device_ptr_mut(&stream).0 as *mut i8,
s_d.device_ptr_mut(&stream).0 as *mut f32,
rms_bd as i32,
eps_norm,
sp(&stream),
)
},
)?;
}
{
let ws2 = &mut *ws;
pre_finish_into(
e,
topology,
site,
x,
&mut ws2.mixes,
&mut ws2.pre,
&mut ws2.post,
&mut ws2.comb,
&mut ws2.y,
1,
hidden,
)?;
}
let (r_q, r_d) = {
let zdst: &mut CudaSlice<f32> = match dst {
NormDst::H => &mut ws.h,
NormDst::Z => &mut ws.z,
};
e.rms_norm_zq8_f32(&ws.y, norm_w, zdst, hidden, 1, eps_norm)?
};
stream.synchronize()?;
let ord = HC_PRE_ZQ8_CHECK_SITES.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
let f = |a: &[f32], b: &[f32], name: &str| -> Vec<String> {
a.iter()
.zip(b)
.enumerate()
.filter(|(_, (p, q))| p.to_bits() != q.to_bits())
.take(8)
.map(|(i, (p, q))| {
format!(
"{name}[{i}] fused={:#010x} two={:#010x}",
p.to_bits(),
q.to_bits()
)
})
.collect()
};
let zref = match dst {
NormDst::H => e.dtoh(&ws.h)?,
NormDst::Z => e.dtoh(&ws.z)?,
};
let mut bad: Vec<String> = Vec::new();
bad.extend(f(&e.dtoh(&s_pre)?, &e.dtoh(&ws.pre)?[..streams], "pre"));
bad.extend(f(&e.dtoh(&s_post)?, &e.dtoh(&ws.post)?[..streams], "post"));
bad.extend(f(
&e.dtoh(&s_comb)?,
&e.dtoh(&ws.comb)?[..streams * streams],
"comb",
));
bad.extend(f(&e.dtoh(&s_y)?, &e.dtoh(&ws.y)?[..hidden], "y"));
bad.extend(f(&e.dtoh(&s_z)?, &zref[..hidden], "z"));
bad.extend(f(&e.dtoh(&s_d)?, &e.dtoh(&r_d)?[..hidden / 32], "d"));
let (sq, rq) = (e.dtoh_i8(&s_q)?, e.dtoh_i8(&r_q)?);
let qbad: Vec<String> = sq
.iter()
.zip(rq.iter())
.enumerate()
.filter(|(_, (p, q))| p != q)
.take(8)
.map(|(i, (p, q))| format!("q[{i}] fused={p} two={q}"))
.collect();
bad.extend(qbad);
if !bad.is_empty() {
HC_PRE_ZQ8_CHECK_BAD.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
eprintln!(
"[hc-pre-zq8-check] site #{ord} dst={dst:?} MISMATCH: {}",
bad.join("; ")
);
}
if ord.is_multiple_of(256) {
eprintln!(
"[hc-pre-zq8-check] {} sites compared, {} with a mismatch",
ord + 1,
HC_PRE_ZQ8_CHECK_BAD.load(std::sync::atomic::Ordering::Relaxed)
);
}
Ok(())
}
#[allow(clippy::too_many_arguments)]
pub fn pre_t1_ws_zq8(
e: &Engine,
topology: &HyperTopology,
site: &HyperSite,
x: &CudaSlice<f32>,
ws: &mut HyperDecodeWs,
hidden: usize,
norm_w: &CudaSlice<f32>,
dst: NormDst,
eps_norm: f32,
) -> Res<Option<(CudaSlice<i8>, CudaSlice<f32>)>> {
let streams = topology.streams;
let rows = topology.rows();
let width = streams * hidden;
let block = crate::hc_pre_block();
let rms_bd = crate::rms_block() as usize;
if hc_fused_pre_arm() != HcFusedPreArm::V2
|| streams > 8
|| rows > 32
|| rms_bd > block
|| !rms_bd.is_multiple_of(32)
|| !hidden.is_multiple_of(32)
{
return Ok(None);
}
if hc_pre_zq8_selfcheck() {
return pre_t1_ws_zq8_selfcheck(e, topology, site, x, ws, hidden, norm_w, dst, eps_norm)
.map(|()| None);
}
let stream = e.stream();
{
let xr = x.slice(0..width);
let wv = site.fn_w.slice(0..site.fn_w.len());
let mut yr = ws.mixes.slice_mut(0..rows);
hc_mixes_into(e, &xr, &wv, &mut yr, width, rows)
.map_err(|err| format!("hc pre_t1_ws_zq8 mixes: {err}"))?;
}
let mut q = e.uninit_i8(hidden)?;
let mut d = e.uninit(hidden / 32)?;
let HyperDecodeWs {
mixes,
pre,
post,
comb,
y,
h,
z,
..
} = ws;
let zdst: &mut CudaSlice<f32> = match dst {
NormDst::H => h,
NormDst::Z => z,
};
unsafe {
ck("hc_pre_zq8", {
let v4z = if crate::hc_pre_v4z_on() {
let rc = k::memra_dsv4_hc_pre_v4z(
dpf!(x, &stream),
dpf!(mixes, &stream),
dpf!(site.scale, &stream),
dpf!(site.base, &stream),
dpm!(pre, &stream),
dpm!(post, &stream),
dpm!(comb, &stream),
dpm!(y, &stream),
1,
streams as i32,
hidden as i32,
topology.sinkhorn_iterations as i32,
topology.epsilon,
std::ptr::null_mut(),
dpf!(norm_w, &stream),
dpm!(zdst, &stream),
q.device_ptr_mut(&stream).0 as *mut i8,
d.device_ptr_mut(&stream).0 as *mut f32,
eps_norm,
rms_bd as i32,
sp(&stream),
);
if rc == 40025 { None } else { Some(rc) }
} else {
None
};
match v4z {
Some(rc) => {
HC_PRE_V4Z_DISPATCHES.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
rc
}
None => k::memra_dsv4_hc_pre_zq8(
dpf!(x, &stream),
dpf!(mixes, &stream),
dpf!(site.scale, &stream),
dpf!(site.base, &stream),
dpm!(pre, &stream),
dpm!(post, &stream),
dpm!(comb, &stream),
dpm!(y, &stream),
1,
streams as i32,
hidden as i32,
topology.sinkhorn_iterations as i32,
topology.epsilon,
std::ptr::null_mut(),
block as i32,
crate::hc_pre_sink_reg() as i32,
dpf!(norm_w, &stream),
dpm!(zdst, &stream),
q.device_ptr_mut(&stream).0 as *mut i8,
d.device_ptr_mut(&stream).0 as *mut f32,
rms_bd as i32,
eps_norm,
sp(&stream),
),
}
})?;
}
if HC_PRE_ZQ8_DISPATCHES.fetch_add(1, std::sync::atomic::Ordering::Relaxed) == 0 {
eprintln!(
"[hc-pre-zq8] engaged streams={streams} hidden={hidden} block={block} rms_bd={rms_bd} \
(one launch replaces hc_pre_fused_v3 + rms_norm_zq8_f32_v2 per site; MEMRA_HC_PRE_ZQ8=1)"
);
}
Ok(Some((q, d)))
}
pub fn post_t1_ws(
e: &Engine,
topology: &HyperTopology,
f: &CudaSlice<f32>,
residual: &CudaSlice<f32>,
ws: &mut HyperDecodeWs,
hidden: usize,
) -> Res<()> {
let stream = e.stream();
let ws = &mut *ws;
unsafe {
ck(
"hc_post",
k::memra_dsv4_hc_post(
dpf!(f, &stream),
dpf!(residual, &stream),
dpf!(ws.post, &stream),
dpf!(ws.comb, &stream),
dpm!(ws.xb, &stream),
1,
topology.streams as i32,
hidden as i32,
sp(&stream),
),
)?;
}
Ok(())
}
pub fn post(
e: &Engine,
topology: &HyperTopology,
f: &CudaSlice<f32>,
residual: &CudaSlice<f32>,
mix: &HcMix,
t: usize,
hidden: usize,
) -> Res<CudaSlice<f32>> {
let streams = topology.streams;
let mut out = e.uninit(t * streams * hidden)?;
let stream = e.stream();
unsafe {
ck(
"hc_post",
k::memra_dsv4_hc_post(
dpf!(f, &stream),
dpf!(residual, &stream),
dpf!(mix.post, &stream),
dpf!(mix.comb, &stream),
dpm!(out, &stream),
t as i32,
streams as i32,
hidden as i32,
sp(&stream),
),
)?;
}
Ok(out)
}
pub fn contract_mean(
e: &Engine,
topology: &HyperTopology,
x: &CudaSlice<f32>,
t: usize,
hidden: usize,
) -> Res<CudaSlice<f32>> {
let streams = topology.streams;
let stream = e.stream();
let mut out = e.uninit(t * hidden)?;
unsafe {
ck(
"hc_mean",
k::memra_dsv4_hc_mean(
dpf!(x, &stream),
dpm!(out, &stream),
t as i32,
streams as i32,
hidden as i32,
sp(&stream),
),
)?;
}
Ok(out)
}
pub fn collapse(
e: &Engine,
topology: &HyperTopology,
head: Option<&HyperHead>,
x: &CudaSlice<f32>,
t: usize,
hidden: usize,
) -> Res<CudaSlice<f32>> {
let streams = topology.streams;
let stream = e.stream();
let mut out = e.uninit(t * hidden)?;
match topology.collapse {
HcCollapse::Mean => unsafe {
ck(
"hc_mean",
k::memra_dsv4_hc_mean(
dpf!(x, &stream),
dpm!(out, &stream),
t as i32,
streams as i32,
hidden as i32,
sp(&stream),
),
)?;
},
HcCollapse::GatedHead => {
let head = head.ok_or_else(|| {
"HcCollapse::GatedHead reached the trunk exit with no head trio loaded".to_string()
})?;
let width = streams * hidden;
let mut mixes = e.linear(x, &head.fn_w, t, width, streams)?;
let mut gates = e.uninit(t * streams)?;
unsafe {
ck(
"hc_head rowsq_scale",
k::memra_dsv4_rowsq_scale(
dpf!(x, &stream),
dpm!(mixes, &stream),
t as i32,
width as i32,
streams as i32,
topology.epsilon,
sp(&stream),
),
)?;
ck(
"hc_head_pre",
k::memra_dsv4_hc_head_pre_m(
dpf!(mixes, &stream),
dpf!(head.scale, &stream),
dpf!(head.base, &stream),
dpm!(gates, &stream),
t as i32,
streams as i32,
topology.epsilon,
sp(&stream),
),
)?;
ck(
"hc_head collapse",
k::memra_dsv4_hc_collapse(
dpf!(x, &stream),
dpf!(gates, &stream),
dpm!(out, &stream),
t as i32,
streams as i32,
hidden as i32,
sp(&stream),
),
)?;
}
}
}
Ok(out)
}
#[cfg(test)]
mod tests {
use super::*;
use memra_gguf::model_plan::{
ActivationPlan, AttentionPlan, DenseMlpPlan, DraftSourcePlan, KimiDeltaNetPlan, LayerPlan,
MlpPlan, NormKind, NormPlan, StatePlan, WeightTransform,
};
fn norm() -> NormPlan {
NormPlan {
kind: NormKind::Rms,
epsilon: 1e-5,
weight_transform: WeightTransform::Identity,
}
}
fn layer(index: u32, residual: ResidualTopology) -> LayerPlan {
LayerPlan {
index,
pre_attention_norm: norm(),
attention: AttentionPlan::KimiDeltaNet(KimiDeltaNetPlan {
num_heads: 1,
head_dim: 128,
conv_kernel: 4,
gate_lower_bound: -5.0,
}),
pre_mlp_norm: norm(),
mlp: MlpPlan::Dense(DenseMlpPlan {
intermediate_size: 16,
activation: ActivationPlan::Silu,
}),
residual,
state: StatePlan::Recurrent {
conv_width: 384,
conv_kernel: 4,
state_width: 16384,
},
ple: None,
sparse_overlay: None,
}
}
fn plan(residuals: [ResidualTopology; 2]) -> ModelPlan {
ModelPlan {
arch: memra_gguf::config::Arch::Glm5Next,
hidden_size: 8,
vocab_size: 16,
context_length: 32,
embedding_scale: 1.0,
vision: None,
multimodal: None,
layers: vec![layer(0, residuals[0]), layer(1, residuals[1])],
output_norm: norm(),
logits: Vec::new(),
mtp_blocks: Vec::new(),
drafter: None,
exit_mixer: None,
draft_source: DraftSourcePlan::Embedded,
sampling_defaults: None,
partition_boundaries: Vec::new(),
}
}
fn hc(streams: u32) -> ResidualTopology {
ResidualTopology::HyperConnections {
streams,
epsilon: 1e-6,
sinkhorn_iterations: 20,
collapse: HcCollapse::Mean,
}
}
#[test]
fn serial_trunk_has_no_topology() {
let plan = plan([ResidualTopology::Serial, ResidualTopology::Serial]);
assert!(HyperTopology::from_plan(&plan).unwrap().is_none());
}
#[test]
fn uniform_trunk_yields_the_plans_constants() {
let plan = plan([hc(4), hc(4)]);
let topology = HyperTopology::from_plan(&plan).unwrap().unwrap();
assert_eq!(topology.streams, 4);
assert_eq!(topology.sinkhorn_iterations, 20);
assert_eq!(topology.collapse, HcCollapse::Mean);
assert_eq!(topology.rows(), 24);
}
#[test]
fn a_mixed_trunk_is_refused_in_both_orders() {
for residuals in [
[hc(4), ResidualTopology::Serial],
[ResidualTopology::Serial, hc(4)],
[hc(4), hc(2)],
] {
assert!(
HyperTopology::from_plan(&plan(residuals)).is_err(),
"a non-uniform trunk must be refused, not silently keyed off layer 0"
);
}
}
#[test]
fn zero_iterations_are_refused() {
let bad = ResidualTopology::HyperConnections {
streams: 4,
epsilon: 1e-6,
sinkhorn_iterations: 0,
collapse: HcCollapse::Mean,
};
assert!(HyperTopology::from_plan(&plan([bad, bad])).is_err());
}
}