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);
fn hc_fused_pre_on() -> bool {
std::env::var("MEMRA_HC_FUSED_PRE").as_deref() == Ok("1")
}
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);
e.linear_t1_into(&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();
if hc_fused_pre_on() && streams <= 8 {
unsafe {
ck(
"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),
),
)?;
}
if HC_FUSED_PRE_DISPATCHES.fetch_add(1, std::sync::atomic::Ordering::Relaxed) == 0 {
eprintln!(
"[hc-fused-pre] engaged streams={streams} hidden={hidden} t={t} (one launch \
replaces rowsq_scale + sinkhorn + collapse per site; MEMRA_HC_FUSED_PRE=1)"
);
}
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);
e.linear_t1_into(&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,
)
}
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());
}
}