#![allow(clippy::needless_range_loop)]
use std::ops::Range;
use std::sync::Arc;
use std::sync::atomic::{AtomicBool, AtomicU64, Ordering};
use cudarc::driver::CudaSlice;
use crate::Engine;
use crate::kda::{ConvArm, KdaAttnLayer};
use crate::model::GpuTensor;
use memra_kv::{Cache, LatentKvLayer, RecurLayer};
pub const GLM5_TP_ALLOWED_RANKS: [usize; 2] = [2, 4];
pub const GLM5_TP_TRANSPORT_TAG: &str = "glm5-tp-transport";
pub type Glm5TpLayerSpec = crate::tp::StepEpLayerSpec;
pub fn glm5_tp_env_raw() -> Option<String> {
std::env::var("MEMRA_GLM5_TP").ok()
}
pub fn glm5_tp_armed() -> bool {
matches!(glm5_tp_env_raw().as_deref(), Some(v) if !v.is_empty() && v != "0")
}
pub fn parse_glm5_tp_layer_specs(
value: Option<&str>,
trunk_layers: usize,
) -> Result<Vec<Glm5TpLayerSpec>, String> {
crate::tp::parse_layer_specs_for_trunk("MEMRA_GLM5_TP", value, Some(trunk_layers))
}
pub fn gate_same_device() -> bool {
std::env::var("MEMRA_GLM5_TP_GATE_SAME_DEV").as_deref() == Ok("1")
}
#[derive(Clone, Copy, PartialEq, Eq, Debug)]
pub enum GateRed {
SwapWo,
SwapEpGateUp,
SkipPeerCombine,
CorruptEpMap,
}
pub fn gate_red() -> Result<Option<GateRed>, String> {
match std::env::var("MEMRA_GLM5_TP_GATE_RED").ok().as_deref() {
None | Some("") => Ok(None),
Some("swap-wo") => Ok(Some(GateRed::SwapWo)),
Some("swap-ep-gateup") => Ok(Some(GateRed::SwapEpGateUp)),
Some("skip-peer-combine") => Ok(Some(GateRed::SkipPeerCombine)),
Some("corrupt-ep-map") => Ok(Some(GateRed::CorruptEpMap)),
Some(other) => Err(format!(
"MEMRA_GLM5_TP_GATE_RED={other:?} is not a known red arm \
(swap-wo | swap-ep-gateup | skip-peer-combine | corrupt-ep-map)"
)),
}
}
pub struct Glm5TpRt {
pub peers: Vec<Engine>,
pub root_dev: usize,
pub peer_devs: Vec<usize>,
pub same_device_gate: bool,
pub transport: crate::tp_transport::TpTransport,
link: Option<crate::tp_transport::PeerPullLink>,
}
impl Glm5TpRt {
pub fn new(devices: &[usize]) -> Result<Self, Box<dyn std::error::Error>> {
let root_dev = devices[0];
let peer_devs: Vec<usize> = devices[1..].to_vec();
for &d in &peer_devs {
if d == root_dev || peer_devs.iter().filter(|&&x| x == d).count() > 1 {
return Err(format!(
"MEMRA_GLM5_TP rank devices must be distinct in serving; got {devices:?} \
(the same-device form exists only for the rig gate binary)"
)
.into());
}
}
let mut peers = Vec::with_capacity(peer_devs.len());
for &d in &peer_devs {
peers.push(Engine::new(d)?);
}
Ok(Self {
peers,
root_dev,
peer_devs,
same_device_gate: false,
transport: crate::tp_transport::TpTransport::HostCanonical,
link: None,
})
}
pub fn new_gate_same_device(
root_dev: usize,
ranks: usize,
) -> Result<Self, Box<dyn std::error::Error>> {
let mut peers = Vec::with_capacity(ranks - 1);
for _ in 1..ranks {
peers.push(Engine::new(root_dev)?);
}
Ok(Self {
peers,
root_dev,
peer_devs: vec![root_dev; ranks - 1],
same_device_gate: true,
transport: crate::tp_transport::TpTransport::HostCanonical,
link: None,
})
}
pub fn ranks(&self) -> usize {
self.peers.len() + 1
}
pub fn arm_transport(&mut self, root: &Engine) -> Result<(), Box<dyn std::error::Error>> {
let (transport, armed_flag) = crate::tp_transport::transport_env()?;
let engines: Vec<&Engine> = std::iter::once(root).chain(self.peers.iter()).collect();
let link = crate::tp_transport::arm_transport(
transport,
armed_flag,
GLM5_TP_TRANSPORT_TAG,
&engines,
self.same_device_gate,
)?;
self.transport = transport;
self.link = link;
Ok(())
}
pub fn hop<'a>(&'a self, root: &'a Engine) -> crate::tp_transport::Hop<'a> {
crate::tp_transport::Hop {
engines: std::iter::once(root).chain(self.peers.iter()).collect(),
transport: self.transport,
link: self.link.as_ref(),
}
}
}
pub struct Glm5TpModelView {
pub trunk_layers: usize,
pub layer_class: Vec<Glm5LayerClass>,
pub layer_is_moe: Vec<bool>,
pub kda_heads: usize,
pub kda_head_dim: usize,
pub mla_heads: usize,
pub n_routed_experts: usize,
pub top_k: usize,
}
#[derive(Clone, Copy, PartialEq, Eq, Debug)]
pub enum Glm5LayerClass {
Kda,
Mla,
}
pub struct Glm5TpLoadPlan {
pub rt: Arc<Glm5TpRt>,
pub layers: std::collections::BTreeSet<usize>,
pub ep_map: Option<crate::ep_map::EpMap>,
}
fn load_glm5_ep_map(
view: &Glm5TpModelView,
layers: &std::collections::BTreeSet<usize>,
ranks: usize,
) -> Result<Option<crate::ep_map::EpMap>, Box<dyn std::error::Error>> {
let Some((flag, path)) = crate::ep_map::ep_map_env()? else {
return Ok(None);
};
if path.is_empty() {
return Err(format!(
"{flag} is set but empty (fail-closed: unset the flag for \
the even split; an empty value never silently means default)"
)
.into());
}
let text = std::fs::read_to_string(&path)
.map_err(|e| format!("{flag}={path}: cannot read the map file ({e}) — refused by name"))?;
let map = crate::ep_map::EpMap::parse(&text).map_err(|e| format!("{flag}={path}: {e}"))?;
if map.ranks != ranks {
return Err(format!(
"{flag}={path}: map declares ranks={}, this load is TP-{ranks} \
(re-mint the map for the armed rank count)",
map.ranks
)
.into());
}
if map.n_experts != view.n_routed_experts {
return Err(format!(
"{flag}={path}: map declares expert_count={}, the model routes {}",
map.n_experts, view.n_routed_experts
)
.into());
}
if map.entry_rank != 0 {
return Err(format!(
"{flag}={path}: entry_rank={} but the glm5 TP first-hop card is \
rank 0 (root: router + combine + shared expert) — re-mint with \
--entry-rank 0 (refused rather than silently remapping ranks)",
map.entry_rank
)
.into());
}
let ep_layers: Vec<usize> = layers
.iter()
.copied()
.filter(|&il| view.layer_is_moe[il])
.collect();
map.validate_layer_cover(&ep_layers)
.map_err(|e| format!("{flag}={path}: {e}"))?;
let digest = {
use sha2::{Digest, Sha256};
let mut h = Sha256::new();
h.update(text.as_bytes());
let out = h.finalize();
out.iter().map(|b| format!("{b:02x}")).collect::<String>()
};
eprintln!(
"[glm5-tp-preflight] ep-map armed path={path} sha256={digest} layers={} \
experts={} ranks={} entry_rank={} performance_claim=false",
map.layers.len(),
map.n_experts,
map.ranks,
map.entry_rank,
);
Ok(Some(map))
}
pub const GLM5_TP_REFUSED_DOOR_FLAGS: [(&str, &str); 4] = [
(
"MEMRA_HC_FUSED_PRE",
"the fused mHC pre-chain is gated on the unsharded walk only",
),
(
"MEMRA_HC_DECODE_WS",
"the workspace decode walk carries no TP mixer branches",
),
(
"MEMRA_KDA_FUSED_PROJ",
"the fused six-projection door (either operand arm) is gated on full-width \
projections, never head shards",
),
(
"MEMRA_MLA_DECODE_SPLIT",
"the absorb/decompress split is gated on the full-head geometry",
),
];
pub fn refuse_glm5_tp_door_composition(armed: impl Fn(&str) -> bool) -> Result<(), String> {
crate::tp::refuse_door_composition("MEMRA_GLM5_TP", &GLM5_TP_REFUSED_DOOR_FLAGS, armed)
}
pub fn prepare_glm5_tp_load(
e: &Engine,
view: &Glm5TpModelView,
) -> Result<Option<Glm5TpLoadPlan>, Box<dyn std::error::Error>> {
let raw = glm5_tp_env_raw();
let specs = parse_glm5_tp_layer_specs(raw.as_deref(), view.trunk_layers)?;
if specs.is_empty() {
return Ok(None);
}
if crate::pp::pp_cuts(view.trunk_layers).is_some() {
return Err(
"MEMRA_GLM5_TP + MEMRA_PP_STAGES>1: the TP x PP composition is unwired and \
refuses until its own gate exists (stage 5 of the tp2 lane names it)"
.into(),
);
}
if !crate::tp::step_tp_layer_specs()?.is_empty()
|| !crate::tp::step_ep_layer_specs()?.is_empty()
{
return Err(
"MEMRA_GLM5_TP + MEMRA_STEP_TP/MEMRA_STEP_EP: the step and glm5 parallel \
contracts never co-arm"
.into(),
);
}
refuse_glm5_tp_door_composition(|flag| std::env::var(flag).as_deref() == Ok("1"))?;
let devices = specs[0].devices.clone();
let ranks = devices.len();
if !GLM5_TP_ALLOWED_RANKS.contains(&ranks) {
return Err(format!(
"MEMRA_GLM5_TP names {ranks} devices per layer; the qualified rank envelope is \
{GLM5_TP_ALLOWED_RANKS:?} (TP-3 is a head-padding question, not built — see the \
module doc)"
)
.into());
}
if view.layer_class.len() != view.trunk_layers || view.layer_is_moe.len() != view.trunk_layers {
return Err(format!(
"glm5-tp preflight: layer class map ({}/{}) does not cover the {}-layer trunk",
view.layer_class.len(),
view.layer_is_moe.len(),
view.trunk_layers
)
.into());
}
if !view.kda_heads.is_multiple_of(ranks) || view.kda_heads == 0 {
return Err(format!(
"glm5-tp: {} KDA heads do not shard across {ranks} ranks",
view.kda_heads
)
.into());
}
if view.kda_head_dim != crate::kda::KDA_HEAD_DIM {
return Err(format!(
"glm5-tp: KDA head_dim {} is not the {} the scan kernel is instantiated for",
view.kda_head_dim,
crate::kda::KDA_HEAD_DIM
)
.into());
}
if !view.mla_heads.is_multiple_of(ranks) || view.mla_heads == 0 {
return Err(format!(
"glm5-tp: {} MLA heads do not shard across {ranks} ranks",
view.mla_heads
)
.into());
}
if !view.n_routed_experts.is_multiple_of(ranks) || view.n_routed_experts == 0 {
return Err(format!(
"glm5-tp: {} routed experts do not partition across {ranks} ranks",
view.n_routed_experts
)
.into());
}
if view.top_k > view.n_routed_experts {
return Err("glm5-tp: top_k exceeds the routed expert count".into());
}
for s in &specs {
if s.devices != devices {
return Err(format!(
"MEMRA_GLM5_TP carries ONE runtime group: layer {} names devices {:?}, \
the first spec names {:?}",
s.layer, s.devices, devices
)
.into());
}
if s.layer >= view.trunk_layers {
return Err(format!(
"MEMRA_GLM5_TP layer {} outside the {}-layer trunk",
s.layer, view.trunk_layers
)
.into());
}
}
let root_dev = e.ctx().ordinal();
if devices[0] != root_dev {
return Err(format!(
"MEMRA_GLM5_TP rank list {:?} must start with the owning device {root_dev} \
(the owner-first rank law)",
devices
)
.into());
}
let red = gate_red()?;
let same_dev = gate_same_device();
if let Some(red) = red {
eprintln!("[glm5-tp-preflight] GATE RED ARM armed: {red:?} — outputs MUST diverge");
}
let mut rt = if same_dev {
eprintln!(
"[glm5-tp-preflight] GATE same-device emulation: {} peer ranks are additional \
contexts on device {root_dev} (spec devices {:?} are logical rank ids)",
ranks - 1,
&devices[1..],
);
Glm5TpRt::new_gate_same_device(root_dev, ranks)?
} else {
Glm5TpRt::new(&devices)?
};
rt.arm_transport(e)?;
let rt = Arc::new(rt);
let layers: std::collections::BTreeSet<usize> = specs.iter().map(|s| s.layer).collect();
let ep_map = load_glm5_ep_map(view, &layers, ranks)?;
let (mut kda_n, mut mla_n, mut moe_n) = (0usize, 0usize, 0usize);
for &il in &layers {
match view.layer_class[il] {
Glm5LayerClass::Kda => kda_n += 1,
Glm5LayerClass::Mla => mla_n += 1,
}
if view.layer_is_moe[il] {
moe_n += 1;
}
}
eprintln!(
"[glm5-tp-preflight] armed ranks={ranks} devices={devices:?} layers={} \
kda_shard={kda_n} mla_shard={mla_n} moe_ep={moe_n} kda_heads_per_rank={} \
mla_heads_per_rank={} experts_per_rank={} transport={} \
weights_loaded=false performance_claim=false",
layers.len(),
view.kda_heads / ranks,
view.mla_heads / ranks,
view.n_routed_experts / ranks,
rt.transport.name(),
);
Ok(Some(Glm5TpLoadPlan { rt, layers, ep_map }))
}
fn outer_rows(ne: &[u64]) -> (usize, usize) {
let outer = *ne.last().expect("tensor has at least one axis") as usize;
let inner: usize = ne[..ne.len() - 1].iter().map(|&d| d as usize).product();
(outer, inner.max(1))
}
fn shard_rows(
src_engine: &Engine,
dst: &Engine,
t: &GpuTensor,
rows: Range<usize>,
) -> Result<GpuTensor, Box<dyn std::error::Error>> {
match t {
GpuTensor::Float { data, ne } => {
let (outer, inner) = outer_rows(ne);
if rows.end > outer {
return Err(format!("shard rows {rows:?} exceed outer axis {outer}").into());
}
let host = src_engine.dtoh(data)?;
let piece = &host[rows.start * inner..rows.end * inner];
let mut ne2 = ne.clone();
*ne2.last_mut().unwrap() = (rows.end - rows.start) as u64;
Ok(GpuTensor::Float {
data: dst.htod(piece)?,
ne: ne2,
})
}
GpuTensor::FloatBf16 { data, ne } => {
let (outer, inner) = outer_rows(ne);
if rows.end > outer {
return Err(format!("shard rows {rows:?} exceed outer axis {outer}").into());
}
let host = src_engine.dtoh_u8(data)?;
let piece = &host[rows.start * inner * 2..rows.end * inner * 2];
let mut ne2 = ne.clone();
*ne2.last_mut().unwrap() = (rows.end - rows.start) as u64;
Ok(GpuTensor::FloatBf16 {
data: dst.htod_bytes(piece)?,
ne: ne2,
})
}
GpuTensor::Quant {
bytes,
qtype,
row_bytes,
ne,
scale,
rp,
fp8,
rp4,
blk,
f16,
#[cfg(memra_cutlass)]
cutlass,
} => {
if *rp {
return Err(
"glm5-tp shard: rp split-plane mirror layout is unwired — load \
the TP-armed tensor with MEMRA_RP=0 (raw layout is bit-identical \
by the mirror's own contract)"
.into(),
);
}
if fp8.is_some() || rp4.is_some() || blk.is_some() || f16.is_some() {
return Err(
"glm5-tp shard: a decode/prefill mirror (fp8/rp4/blk/f16) is present on a \
TP-armed tensor — mirrors are unwired for shards in v1; disable the \
mirror door for this load"
.into(),
);
}
#[cfg(memra_cutlass)]
if cutlass.is_some() {
return Err("glm5-tp shard: cutlass prefill operand unwired for shards".into());
}
let (outer, inner) = outer_rows(ne);
if ne.len() != 2 {
return Err("glm5-tp shard: quantized shards are 2D-only in v1".into());
}
let _ = inner;
if rows.end > outer {
return Err(format!("shard rows {rows:?} exceed outer axis {outer}").into());
}
let host = src_engine.dtoh_u8(bytes)?;
let piece = &host[rows.start * row_bytes..rows.end * row_bytes];
let mut ne2 = ne.clone();
*ne2.last_mut().unwrap() = (rows.end - rows.start) as u64;
Ok(GpuTensor::Quant {
bytes: dst.htod_bytes(piece)?,
qtype: *qtype,
row_bytes: *row_bytes,
ne: ne2,
scale: *scale,
rp: false,
fp8: None,
rp4: None,
blk: None,
f16: None,
#[cfg(memra_cutlass)]
cutlass: None,
})
}
}
}
fn replicate(
src_engine: &Engine,
dst: &Engine,
t: &GpuTensor,
) -> Result<GpuTensor, Box<dyn std::error::Error>> {
let (outer, _) = outer_rows(t.ne());
shard_rows(src_engine, dst, t, 0..outer)
}
pub(crate) fn rank_engine<'a>(e: &'a Engine, rt: &'a Glm5TpRt, r: usize) -> &'a Engine {
if r == 0 { e } else { &rt.peers[r - 1] }
}
pub struct Glm5TpKda {
pub rt: Arc<Glm5TpRt>,
pub peers: Vec<KdaAttnLayer>,
pub full_qkv: usize,
pub n_embd: usize,
}
impl Glm5TpKda {
pub fn ranks(&self) -> usize {
self.peers.len() + 1
}
}
static KDA_MARKED: AtomicBool = AtomicBool::new(false);
pub(crate) fn shard_kda_layer(
e: &Engine,
rt: &Arc<Glm5TpRt>,
la: KdaAttnLayer,
) -> Result<KdaAttnLayer, Box<dyn std::error::Error>> {
if la.tp.is_some() {
return Err("shard_kda_layer: layer is already sharded".into());
}
let ranks = rt.ranks();
let heads = la.heads();
let head_dim = la.head_dim();
let qkv = la.qkv();
let kernel = la.conv_kernel();
if !heads.is_multiple_of(ranks) {
return Err(format!("KDA heads {heads} do not shard across {ranks} ranks").into());
}
let hl = heads / ranks; let ql = qkv / ranks; let n_embd = la.wo.out_features();
if !n_embd.is_multiple_of(ranks) {
return Err(format!("KDA wo out {n_embd} does not split across ranks").into());
}
let hh = n_embd / ranks;
let mut shard_plan = la.plan;
shard_plan.num_heads = hl as u32;
let conv_host = e.dtoh(&la.conv)?;
let conv_rank =
|dst: &Engine, r: usize| -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
let mut piece = Vec::with_capacity(3 * ql * kernel);
for p in 0..3 {
let a = (p * qkv + r * ql) * kernel;
piece.extend_from_slice(&conv_host[a..a + ql * kernel]);
}
dst.htod(&piece)
};
let wo_rank = |r: usize| -> usize {
match gate_red() {
Ok(Some(GateRed::SwapWo)) => (r + 1) % ranks,
_ => r,
}
};
let rank_shard = |dst: &Engine, r: usize| -> Result<KdaAttnLayer, Box<dyn std::error::Error>> {
let wr = wo_rank(r);
Ok(KdaAttnLayer {
plan: shard_plan,
wq: shard_rows(e, dst, &la.wq, r * ql..(r + 1) * ql)?,
wk: shard_rows(e, dst, &la.wk, r * ql..(r + 1) * ql)?,
wv: shard_rows(e, dst, &la.wv, r * ql..(r + 1) * ql)?,
f_a: replicate(e, dst, &la.f_a)?,
f_b: shard_rows(e, dst, &la.f_b, r * ql..(r + 1) * ql)?,
g_a: replicate(e, dst, &la.g_a)?,
g_b: shard_rows(e, dst, &la.g_b, r * ql..(r + 1) * ql)?,
b_proj: shard_rows(e, dst, &la.b_proj, r * hl..(r + 1) * hl)?,
wo: shard_rows(e, dst, &la.wo, wr * hh..(wr + 1) * hh)?,
conv: conv_rank(dst, r)?,
a_log: shard_rows(e, dst, &la.a_log, r * hl..(r + 1) * hl)?,
dt_bias: shard_rows(e, dst, &la.dt_bias, r * ql..(r + 1) * ql)?,
o_norm: replicate(e, dst, &la.o_norm)?,
tp: None,
})
};
let mut root = rank_shard(e, 0)?;
let mut peers = Vec::with_capacity(ranks - 1);
for r in 1..ranks {
peers.push(rank_shard(&rt.peers[r - 1], r)?);
}
if !KDA_MARKED.swap(true, Ordering::Relaxed) {
eprintln!(
"[glm5-tp-kda] head shard armed: ranks={ranks} heads_per_rank={hl} \
head_dim={head_dim} wo=column-over-gather transport={} performance_claim=false",
rt.transport.name(),
);
}
root.tp = Some(Box::new(Glm5TpKda {
rt: Arc::clone(rt),
peers,
full_qkv: qkv,
n_embd,
}));
Ok(root)
}
fn ensure_kda_tp_state<'c>(
e: &Engine,
rt: &Glm5TpRt,
la_root: &KdaAttnLayer,
cache: &'c mut Cache,
il: usize,
) -> Result<&'c mut Vec<RecurLayer>, Box<dyn std::error::Error>> {
if cache.glm5_tp_recur.len() <= il {
return Err(format!("glm5-tp: cache carries no TP recur slot for layer {il}").into());
}
if cache.glm5_tp_recur[il].is_none() {
let conv_pad = la_root.conv_width() * (la_root.conv_kernel() - 1);
let state = la_root.state_width();
let mk = |dev: &Engine| -> Result<RecurLayer, Box<dyn std::error::Error>> {
Ok(RecurLayer {
conv_state: dev.zeros(conv_pad)?,
ssm_state: dev.zeros(state)?,
ssm_state_alt: dev.zeros(state)?,
})
};
let mut planes = Vec::with_capacity(rt.ranks());
planes.push(mk(e)?);
for p in &rt.peers {
planes.push(mk(p)?);
}
cache.glm5_tp_recur[il] = Some(planes);
}
Ok(cache.glm5_tp_recur[il].as_mut().unwrap())
}
#[allow(clippy::too_many_arguments)] fn kda_tp_core(
e: &Engine,
la_root: &KdaAttnLayer,
x: &CudaSlice<f32>,
t: usize,
eps: f32,
cache: &mut Cache,
il: usize,
arm: ConvArm,
verify_stash: Option<&mut Glm5TpKdaVerifyStash>,
mut scan_clock: Option<&mut u64>,
) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
let tp = la_root
.tp
.as_ref()
.ok_or("kda_tp_core called on an unsharded layer")?;
let rt = &tp.rt;
let ranks = rt.ranks();
let ql = la_root.qkv(); let full = tp.full_qkv;
let n_embd = tp.n_embd;
let hh = n_embd / ranks;
let rows_exact = verify_stash.is_some();
let mut captured: Vec<Option<(CudaSlice<f32>, crate::kda::KdaRowsStash)>> =
(0..ranks).map(|_| None).collect();
let hop = rt.hop(e);
let x_peers = crate::tp_transport::fanout_f32(&hop, x, x.len())?;
let states = ensure_kda_tp_state(e, rt, la_root, cache, il)?;
let mut gated: Vec<Option<CudaSlice<f32>>> = (0..ranks).map(|_| None).collect();
for r in (1..ranks).chain(std::iter::once(0)) {
let dev = if r == 0 { e } else { &rt.peers[r - 1] };
let la = if r == 0 { la_root } else { &tp.peers[r - 1] };
let xin = if r == 0 { x } else { &x_peers[r - 1] };
let snap = if rows_exact {
Some(dev.clone_dtod(&states[r].ssm_state)?)
} else {
None
};
let mut rank_stash: Option<crate::kda::KdaRowsStash> = None;
let mut rank_scan_ns = 0u64;
let out = {
let RecurLayer {
conv_state,
ssm_state,
ssm_state_alt,
} = &mut states[r];
let out = crate::kda::kda_core_gated(
dev,
la,
xin,
t,
eps,
conv_state,
ssm_state,
ssm_state_alt,
arm,
if rows_exact {
crate::kda::KdaStash::Rows(&mut rank_stash)
} else {
crate::kda::KdaStash::None
},
scan_clock.as_deref_mut().map(|_| &mut rank_scan_ns),
)?;
std::mem::swap(ssm_state, ssm_state_alt);
out
};
if let Some(clock) = scan_clock.as_deref_mut() {
*clock += rank_scan_ns;
}
if rows_exact {
let snap = snap.expect("verify arm cloned the snapshot above");
let rank_stash = rank_stash
.ok_or("kda_core_gated returned without filling the requested rows stash")?;
captured[r] = Some((snap, rank_stash));
}
gated[r] = Some(out);
}
if let Some(stash_vec) = verify_stash {
stash_vec.clear();
for c in captured {
stash_vec.push(c.expect("every rank captured on the verify arm"));
}
}
debug_assert_eq!(full, ranks * ql);
let gated_refs: Vec<&CudaSlice<f32>> = gated
.iter()
.map(|g| g.as_ref().expect("filled above"))
.collect();
let fulls = crate::tp_transport::gather_parts(&hop, &gated_refs, t, ql)?;
let mut ys = Vec::with_capacity(ranks);
if rows_exact {
ys.push(e.matmul_rows_exact(&la_root.wo, &fulls[0], t)?);
for r in 1..ranks {
ys.push(rt.peers[r - 1].matmul_rows_exact(&tp.peers[r - 1].wo, &fulls[r], t)?);
}
} else {
ys.push(e.matmul(&la_root.wo, &fulls[0], t)?);
for r in 1..ranks {
ys.push(rt.peers[r - 1].matmul(&tp.peers[r - 1].wo, &fulls[r], t)?);
}
}
let y_refs: Vec<&CudaSlice<f32>> = ys.iter().collect();
crate::tp_transport::concat_parts_on_root(&hop, &y_refs, t, hh)
}
#[allow(clippy::too_many_arguments)] pub(crate) fn kda_tp_cached(
e: &Engine,
la_root: &KdaAttnLayer,
x: &CudaSlice<f32>,
t: usize,
eps: f32,
cache: &mut Cache,
il: usize,
arm: ConvArm,
) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
kda_tp_core(e, la_root, x, t, eps, cache, il, arm, None, None)
}
pub type Glm5TpKdaVerifyStash = Vec<(CudaSlice<f32>, crate::kda::KdaRowsStash)>;
#[allow(clippy::too_many_arguments)] pub(crate) fn kda_tp_verify_rows(
e: &Engine,
la_root: &KdaAttnLayer,
x: &CudaSlice<f32>,
t: usize,
eps: f32,
cache: &mut Cache,
il: usize,
scan_clock: Option<&mut u64>,
) -> Result<(CudaSlice<f32>, Glm5TpKdaVerifyStash), Box<dyn std::error::Error>> {
let mut stash: Glm5TpKdaVerifyStash = Vec::new();
let out = kda_tp_core(
e,
la_root,
x,
t,
eps,
cache,
il,
ConvArm::Prefill,
Some(&mut stash),
scan_clock,
)?;
Ok((out, stash))
}
pub(crate) fn kda_tp_verify_rollback(
e: &Engine,
la_root: &KdaAttnLayer,
stash: &Glm5TpKdaVerifyStash,
keep: usize,
cache: &mut Cache,
il: usize,
) -> Result<(), Box<dyn std::error::Error>> {
let tp = la_root
.tp
.as_ref()
.ok_or("kda_tp_verify_rollback called on an unsharded layer")?;
let rt = &tp.rt;
let ranks = rt.ranks();
if stash.len() != ranks {
return Err(format!(
"glm5-tp verify rollback: stash carries {} ranks, the runtime has {ranks}",
stash.len()
)
.into());
}
let states = cache.glm5_tp_recur[il]
.as_mut()
.ok_or_else(|| format!("glm5-tp verify rollback: layer {il} has no per-rank state"))?;
for r in 0..ranks {
let dev = if r == 0 { e } else { &rt.peers[r - 1] };
let la = if r == 0 { la_root } else { &tp.peers[r - 1] };
let (snap, rows) = &stash[r];
crate::kda::kda_verify_rollback_rows_on(dev, la, snap, rows, keep, &mut states[r], il)?;
}
Ok(())
}
pub struct Glm5TpMla {
pub rt: Arc<Glm5TpRt>,
pub peers: Vec<crate::hybrid::MlaAttnLayer>,
pub full_heads: usize,
pub n_embd: usize,
}
impl Glm5TpMla {
pub fn ranks(&self) -> usize {
self.peers.len() + 1
}
}
static MLA_MARKED: AtomicBool = AtomicBool::new(false);
pub(crate) fn shard_mla_layer(
e: &Engine,
rt: &Arc<Glm5TpRt>,
la: crate::hybrid::MlaAttnLayer,
) -> Result<crate::hybrid::MlaAttnLayer, Box<dyn std::error::Error>> {
use crate::hybrid::{MlaAttnLayer, MlaIndexer};
if la.tp.is_some() {
return Err("shard_mla_layer: layer is already sharded".into());
}
let ranks = rt.ranks();
let g = la.geom;
let nh = g.n_head;
if !nh.is_multiple_of(ranks) {
return Err(format!("MLA heads {nh} do not shard across {ranks} ranks").into());
}
let hl = nh / ranks;
let head_q = g.d_nope + g.d_rope; let n_embd = la.wo.out_features();
if !n_embd.is_multiple_of(ranks) {
return Err(format!("MLA wo out {n_embd} does not split across ranks").into());
}
let hh = n_embd / ranks;
let mut shard_geom = g;
shard_geom.n_head = hl;
let replicate_indexer =
|dst: &Engine, ix: &MlaIndexer| -> Result<MlaIndexer, Box<dyn std::error::Error>> {
Ok(MlaIndexer {
wq_b: replicate(e, dst, &ix.wq_b)?,
wk: replicate(e, dst, &ix.wk)?,
k_norm_w: replicate(e, dst, &ix.k_norm_w)?,
k_norm_b: replicate(e, dst, &ix.k_norm_b)?,
weights_proj: replicate(e, dst, &ix.weights_proj)?,
kpool_gate: replicate(e, dst, &ix.kpool_gate)?,
kpool_ape: replicate(e, dst, &ix.kpool_ape)?,
geom: ix.geom,
})
};
let wo_rank = |r: usize| -> usize {
match gate_red() {
Ok(Some(GateRed::SwapWo)) => (r + 1) % ranks,
_ => r,
}
};
let rank_shard = |dst: &Engine, r: usize| -> Result<MlaAttnLayer, Box<dyn std::error::Error>> {
let wr = wo_rank(r);
Ok(MlaAttnLayer {
wq_a: replicate(e, dst, &la.wq_a)?,
q_a_norm: replicate(e, dst, &la.q_a_norm)?,
wq_b: shard_rows(e, dst, &la.wq_b, r * hl * head_q..(r + 1) * hl * head_q)?,
wkv_a: replicate(e, dst, &la.wkv_a)?,
kv_a_norm: replicate(e, dst, &la.kv_a_norm)?,
wk_b: shard_rows(e, dst, &la.wk_b, r * hl..(r + 1) * hl)?,
wv_b: shard_rows(e, dst, &la.wv_b, r * hl..(r + 1) * hl)?,
wo: shard_rows(e, dst, &la.wo, wr * hh..(wr + 1) * hh)?,
geom: shard_geom,
index: match &la.index {
Some(ix) => Some(replicate_indexer(dst, ix)?),
None => None,
},
tp: None,
tp_shard: true,
})
};
let mut root = rank_shard(e, 0)?;
let mut peers = Vec::with_capacity(ranks - 1);
for r in 1..ranks {
peers.push(rank_shard(&rt.peers[r - 1], r)?);
}
if !MLA_MARKED.swap(true, Ordering::Relaxed) {
eprintln!(
"[glm5-tp-mla] head shard armed: ranks={ranks} heads_per_rank={hl} kv_rank={} \
latent=replicated indexer=replicated wo=column-over-gather transport={} \
performance_claim=false",
g.kv_rank,
rt.transport.name(),
);
}
root.tp = Some(Box::new(Glm5TpMla {
rt: Arc::clone(rt),
peers,
full_heads: nh,
n_embd,
}));
Ok(root)
}
pub(crate) fn ensure_mla_peer_latent(
rt: &Glm5TpRt,
canonical: &LatentKvLayer,
cache_slot: &mut Option<Vec<LatentKvLayer>>,
) -> Result<(), Box<dyn std::error::Error>> {
if cache_slot.is_some() {
return Ok(());
}
let mut planes = Vec::with_capacity(rt.peers.len());
for dev in &rt.peers {
let rows = dev.zeros(canonical.rows.len())?;
let len_d = dev.htod_i32(&[0])?;
let index_rows = match &canonical.index_rows {
Some(p) => Some(dev.zeros(p.len())?),
None => None,
};
planes.push(LatentKvLayer {
rows,
width: canonical.width,
index_width: canonical.index_width,
len: 0,
len_d,
index_rows,
index_ring_rows: canonical.index_ring_rows,
index_pool_keys: None, index_pools_ready: 0,
index_pool: canonical.index_pool,
});
}
*cache_slot = Some(planes);
Ok(())
}
pub struct EpRankSlab {
pub gate: CudaSlice<u8>,
pub up: CudaSlice<u8>,
pub down: CudaSlice<u8>,
pub n_experts: usize,
}
pub struct Glm5EpExps {
pub rt: Arc<Glm5TpRt>,
pub slabs: Vec<EpRankSlab>,
pub owner_of: Vec<u8>,
pub local_of: Vec<u32>,
pub ptr_rows: Vec<CudaSlice<u64>>,
}
impl Glm5EpExps {
pub fn owner(&self, expert: usize) -> usize {
self.owner_of[expert] as usize
}
pub fn ranks(&self) -> usize {
self.slabs.len()
}
}
static EP_MARKED: AtomicBool = AtomicBool::new(false);
pub static GLM5_EP_PEER_SLOT_DISPATCHES: AtomicU64 = AtomicU64::new(0);
pub fn glm5_ep_peer_slot_dispatches() -> u64 {
GLM5_EP_PEER_SLOT_DISPATCHES.load(Ordering::Relaxed)
}
pub static GLM5_EP_DIET_DISPATCHES: AtomicU64 = AtomicU64::new(0);
pub fn glm5_ep_diet_dispatches() -> u64 {
GLM5_EP_DIET_DISPATCHES.load(Ordering::Relaxed)
}
pub static GLM5_EP_DIET_BULK_RETURNS: AtomicU64 = AtomicU64::new(0);
pub fn glm5_ep_diet_bulk_returns() -> u64 {
GLM5_EP_DIET_BULK_RETURNS.load(Ordering::Relaxed)
}
pub static GLM5_EP_DIET_PEER_ROUNDTRIPS_AVOIDED: AtomicU64 = AtomicU64::new(0);
pub fn glm5_ep_diet_peer_roundtrips_avoided() -> u64 {
GLM5_EP_DIET_PEER_ROUNDTRIPS_AVOIDED.load(Ordering::Relaxed)
}
pub static GLM5_EP_DIET_FANOUT_UPLOADS_AVOIDED: AtomicU64 = AtomicU64::new(0);
pub fn glm5_ep_diet_fanout_uploads_avoided() -> u64 {
GLM5_EP_DIET_FANOUT_UPLOADS_AVOIDED.load(Ordering::Relaxed)
}
pub static GLM5_EP_GROUPED_PRIME_DISPATCHES: AtomicU64 = AtomicU64::new(0);
pub fn glm5_ep_grouped_prime_dispatches() -> u64 {
GLM5_EP_GROUPED_PRIME_DISPATCHES.load(Ordering::Relaxed)
}
pub(crate) fn arm_moe_ep(
e: &Engine,
rt: &Arc<Glm5TpRt>,
m: &mut crate::hybrid::MoeWeights,
placement: Option<&[u8]>,
) -> Result<(), Box<dyn std::error::Error>> {
if m.glm5_ep.is_some() {
return Err("arm_moe_ep: layer is already EP-armed".into());
}
let ranks = rt.ranks();
let n_expert = m.gate_exps.n_expert;
if !n_expert.is_multiple_of(ranks) {
return Err(format!(
"glm5-tp EP: {n_expert} experts do not partition across {ranks} ranks"
)
.into());
}
if m.gate_exps.layouts.is_some() || m.up_exps.layouts.is_some() || m.down_exps.layouts.is_some()
{
return Err("glm5-tp EP: per-expert mixed layouts are unwired for EP shards".into());
}
let owner_of: Vec<u8> = match placement {
Some(owners) => {
if owners.len() != n_expert {
return Err(format!(
"glm5-tp EP: placement row carries {} owners for a {n_expert}-expert bank",
owners.len()
)
.into());
}
if owners.iter().any(|&r| (r as usize) >= ranks) {
return Err(
format!("glm5-tp EP: placement row names a rank outside TP-{ranks}").into(),
);
}
owners.to_vec()
}
None => crate::ep_map::EpMap::even_owners(n_expert, ranks),
};
let mut local_of = vec![0u32; n_expert];
let mut owned: Vec<Vec<usize>> = vec![Vec::new(); ranks];
for ex in 0..n_expert {
let r = owner_of[ex] as usize;
local_of[ex] = owned[r].len() as u32;
owned[r].push(ex);
}
if owned.iter().any(|o| o.is_empty()) {
return Err("glm5-tp EP: placement leaves a rank with zero experts (refused)".into());
}
let slab =
|dev: &Engine, experts: &[usize]| -> Result<EpRankSlab, Box<dyn std::error::Error>> {
let cut = |h: &crate::model::HostExps,
pad: usize|
-> Result<CudaSlice<u8>, Box<dyn std::error::Error>> {
let stride = h.expert_stride;
let bytes = h.bytes.as_bytes();
let contiguous = experts.windows(2).all(|w| w[1] == w[0] + 1);
if contiguous {
let a = experts[0] * stride;
let b = (experts[experts.len() - 1] + 1) * stride;
return dev.htod_bytes_padded(&bytes[a..b], pad);
}
let mut staged = Vec::with_capacity(experts.len() * stride);
for &ex in experts {
staged.extend_from_slice(&bytes[ex * stride..(ex + 1) * stride]);
}
dev.htod_bytes_padded(&staged, pad)
};
Ok(EpRankSlab {
gate: cut(&m.gate_exps, 8)?,
up: cut(&m.up_exps, 8)?,
down: cut(&m.down_exps, 144)?,
n_experts: experts.len(),
})
};
let mut slabs = Vec::with_capacity(ranks);
for r in 0..ranks {
slabs.push(slab(rank_engine(e, rt, r), &owned[r])?);
}
if matches!(gate_red(), Ok(Some(GateRed::SwapEpGateUp))) {
let root = &mut slabs[0];
std::mem::swap(&mut root.gate, &mut root.up);
}
if matches!(gate_red(), Ok(Some(GateRed::CorruptEpMap))) {
let n0 = owned[0].len() as u32;
for &ex in &owned[0] {
local_of[ex] = n0 - 1 - local_of[ex];
}
}
let ptr_table = |dev: &Engine,
slab: &EpRankSlab,
rank: u8|
-> Result<CudaSlice<u64>, Box<dyn std::error::Error>> {
use cudarc::driver::DevicePtr;
let (pg, pu, pd) = {
let s = dev.stream();
let (pg, _g0) = slab.gate.device_ptr(&s);
let (pu, _g1) = slab.up.device_ptr(&s);
let (pd, _g2) = slab.down.device_ptr(&s);
(pg, pu, pd)
};
let mut host = vec![0u64; 3 * n_expert];
for ex in 0..n_expert {
if owner_of[ex] != rank {
continue; }
let local = local_of[ex] as usize;
host[ex] = pg + (local * m.gate_exps.expert_stride) as u64;
host[n_expert + ex] = pu + (local * m.up_exps.expert_stride) as u64;
host[2 * n_expert + ex] = pd + (local * m.down_exps.expert_stride) as u64;
}
dev.htod_u64(&host)
};
let mut ptr_rows = Vec::with_capacity(ranks);
for r in 0..ranks {
ptr_rows.push(ptr_table(rank_engine(e, rt, r), &slabs[r], r as u8)?);
}
if !EP_MARKED.swap(true, Ordering::Relaxed) {
eprintln!(
"[glm5-tp-ep] expert-parallel armed: experts_per_rank={:?} ownership={} \
router=root combine=slot-ordered-fmaf transport={} \
performance_claim=false",
owned.iter().map(Vec::len).collect::<Vec<_>>(),
if placement.is_some() {
"measured-map"
} else {
"even-split"
},
rt.transport.name(),
);
}
m.dev_exps = None;
m.glm5_ep = Some(Glm5EpExps {
rt: Arc::clone(rt),
slabs,
owner_of,
local_of,
ptr_rows,
});
Ok(())
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn parse_is_literal_and_fail_closed() {
assert!(parse_glm5_tp_layer_specs(None, 45).unwrap().is_empty());
assert!(parse_glm5_tp_layer_specs(Some(""), 45).unwrap().is_empty());
assert!(parse_glm5_tp_layer_specs(Some("0"), 45).unwrap().is_empty());
let all = parse_glm5_tp_layer_specs(Some("all@0,1"), 45).unwrap();
assert_eq!(all.len(), 45);
assert_eq!(all[0].devices, vec![0, 1]);
let all4 = parse_glm5_tp_layer_specs(Some("all@0,1"), 4).unwrap();
assert_eq!(all4.len(), 4);
let quad = parse_glm5_tp_layer_specs(Some("all@0,1,2,3"), 45).unwrap();
assert_eq!(quad.len(), 45);
assert_eq!(quad[0].devices, vec![0, 1, 2, 3]);
let r = parse_glm5_tp_layer_specs(Some("0-2@0,1;4@0,1"), 45).unwrap();
assert_eq!(
r.iter().map(|s| s.layer).collect::<Vec<_>>(),
vec![0, 1, 2, 4]
);
assert!(parse_glm5_tp_layer_specs(Some("0@0,0"), 45).is_err());
assert!(parse_glm5_tp_layer_specs(Some("0@0,1;0@0,1"), 45).is_err());
assert!(parse_glm5_tp_layer_specs(Some("banana"), 45).is_err());
}
fn fixture_view() -> Glm5TpModelView {
Glm5TpModelView {
trunk_layers: 4,
layer_class: vec![
Glm5LayerClass::Kda,
Glm5LayerClass::Mla,
Glm5LayerClass::Kda,
Glm5LayerClass::Mla,
],
layer_is_moe: vec![false, true, true, true],
kda_heads: 4,
kda_head_dim: 128,
mla_heads: 4,
n_routed_experts: 4,
top_k: 2,
}
}
#[test]
fn preflight_geometry_laws_are_dimension_derived() {
let v = fixture_view();
for ranks in GLM5_TP_ALLOWED_RANKS {
assert_eq!(v.kda_heads % ranks, 0);
assert_eq!(v.mla_heads % ranks, 0);
assert_eq!(v.n_routed_experts % ranks, 0);
}
let odd = Glm5TpModelView {
kda_heads: 3,
..fixture_view()
};
assert_ne!(odd.kda_heads % 2, 0);
let bad_dim = Glm5TpModelView {
kda_head_dim: 64,
..fixture_view()
};
assert_ne!(bad_dim.kda_head_dim, crate::kda::KDA_HEAD_DIM);
let odd_experts = Glm5TpModelView {
n_routed_experts: 5,
..fixture_view()
};
assert_ne!(odd_experts.n_routed_experts % 2, 0);
assert!(!GLM5_TP_ALLOWED_RANKS.contains(&3));
}
#[test]
fn armed_check_counts_parse_errors_as_armed() {
for (v, armed) in [
("", false),
("0", false),
("all@0,1", true),
("all@0,1,2,3", true),
("junk", true),
] {
let is_armed = !v.is_empty() && v != "0";
assert_eq!(is_armed, armed);
}
}
#[test]
fn every_refused_door_composition_bites_by_name() {
for (flag, _) in GLM5_TP_REFUSED_DOOR_FLAGS {
let err = refuse_glm5_tp_door_composition(|f| f == flag)
.expect_err("an armed door must refuse");
assert!(err.contains("MEMRA_GLM5_TP"), "{err}");
assert!(err.contains(flag), "{err}");
assert!(err.contains("unproven composition"), "{err}");
}
refuse_glm5_tp_door_composition(|_| false).expect("cold doors must pass");
refuse_glm5_tp_door_composition(|f| f == "MEMRA_GLM5_VERIFY_BATCH")
.expect("verify-batch is refused via the spec co-refusal, not here");
}
}