#![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("") | Some("0") => 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>,
ar: std::sync::Mutex<Option<crate::tp_ar::ArLink>>,
}
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,
ar: std::sync::Mutex::new(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,
ar: std::sync::Mutex::new(None),
})
}
pub fn ar_1stage_available(&self) -> bool {
tp_ar_1stage_on() && self.ranks() == 2 && !self.same_device_gate
}
pub fn ar_prepare(&self, root: &Engine, n: usize) -> Result<(), Box<dyn std::error::Error>> {
if !self.ar_1stage_available() {
return Ok(());
}
let engines: Vec<&Engine> = std::iter::once(root).chain(self.peers.iter()).collect();
let mut guard = self.ar.lock().map_err(|_| "tp ar link poisoned")?;
match guard.as_mut() {
Some(link) => link.ensure_stage(&engines, n),
None => Ok(()),
}
}
pub fn ar_1stage(
&self,
root: &Engine,
x: &mut [&mut CudaSlice<f32>],
n: usize,
) -> Result<bool, Box<dyn std::error::Error>> {
if !self.ar_1stage_available() {
return Ok(false);
}
let engines: Vec<&Engine> = std::iter::once(root).chain(self.peers.iter()).collect();
let mut guard = self
.ar
.lock()
.map_err(|_| "glm5-tp one-shot all-reduce: the link mutex is poisoned")?;
if guard.is_none() {
*guard = Some(crate::tp_ar::ArLink::new(&engines)?);
}
guard
.as_mut()
.expect("built above")
.all_reduce_1stage(&engines, x, n)?;
Ok(true)
}
pub fn ar_1stage_into(
&self,
root: &Engine,
inputs: &[&CudaSlice<f32>],
outs: &mut [&mut CudaSlice<f32>],
n: usize,
) -> Result<bool, Box<dyn std::error::Error>> {
if !self.ar_1stage_available() {
return Ok(false);
}
let engines: Vec<&Engine> = std::iter::once(root).chain(self.peers.iter()).collect();
let mut guard = self
.ar
.lock()
.map_err(|_| "glm5-tp one-shot all-reduce: the link mutex is poisoned")?;
if guard.is_none() {
*guard = Some(crate::tp_ar::ArLink::new(&engines)?);
}
guard
.as_mut()
.expect("built above")
.all_reduce_1stage_into(&engines, inputs, outs, n)?;
Ok(true)
}
#[allow(clippy::too_many_arguments)] pub fn ar_1stage_hcpost(
&self,
root: &Engine,
inputs: &[&CudaSlice<f32>],
residuals: &[&CudaSlice<f32>],
posts: &[&CudaSlice<f32>],
combs: &[&CudaSlice<f32>],
outs: &mut [&mut CudaSlice<f32>],
hc: usize,
d: usize,
) -> Result<bool, Box<dyn std::error::Error>> {
if !self.ar_1stage_available() {
return Ok(false);
}
let engines: Vec<&Engine> = std::iter::once(root).chain(self.peers.iter()).collect();
let mut guard = self
.ar
.lock()
.map_err(|_| "glm5-tp one-shot all-reduce: the link mutex is poisoned")?;
if guard.is_none() {
*guard = Some(crate::tp_ar::ArLink::new(&engines)?);
}
guard
.as_mut()
.expect("built above")
.all_reduce_1stage_hcpost(&engines, inputs, residuals, posts, combs, outs, hc, d)?;
Ok(true)
}
pub fn ranks(&self) -> usize {
self.peers.len() + 1
}
pub fn devices(&self) -> Vec<usize> {
let mut devs = Vec::with_capacity(1 + self.peer_devs.len());
devs.push(self.root_dev);
for &d in &self.peer_devs {
if !devs.contains(&d) {
devs.push(d);
}
}
devs
}
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); 3] = [
(
"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_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_cols(
src_engine: &Engine,
dst: &Engine,
t: &GpuTensor,
cols: Range<usize>,
) -> Result<GpuTensor, Box<dyn std::error::Error>> {
let take = cols.end - cols.start;
match t {
GpuTensor::Float { data, ne } => {
let (outer, inner) = outer_rows(ne);
if cols.end > inner {
return Err(format!("shard cols {cols:?} exceed inner axis {inner}").into());
}
let host = src_engine.dtoh(data)?;
let mut out = Vec::with_capacity(outer * take);
for r in 0..outer {
out.extend_from_slice(&host[r * inner + cols.start..r * inner + cols.end]);
}
let mut ne2 = ne.clone();
ne2[0] = take as u64;
Ok(GpuTensor::Float {
data: dst.htod(&out)?,
ne: ne2,
})
}
GpuTensor::FloatBf16 { data, ne } => {
let (outer, inner) = outer_rows(ne);
if cols.end > inner {
return Err(format!("shard cols {cols:?} exceed inner axis {inner}").into());
}
let host = src_engine.dtoh_u8(data)?;
let mut out = Vec::with_capacity(outer * take * 2);
for r in 0..outer {
let base = r * inner * 2;
out.extend_from_slice(&host[base + cols.start * 2..base + cols.end * 2]);
}
let mut ne2 = ne.clone();
ne2[0] = take as u64;
Ok(GpuTensor::FloatBf16 {
data: dst.htod_bytes(&out)?,
ne: ne2,
})
}
GpuTensor::Quant {
bytes,
qtype,
row_bytes,
ne,
scale,
rp,
fp8,
rp4,
blk,
f16,
#[cfg(memra_cutlass)]
cutlass: _,
} => {
if *rp || fp8.is_some() || rp4.is_some() || blk.is_some() || f16.is_some() {
return Err("glm5-tp shard cols: a mirror layout is unwired for K slices".into());
}
if ne.len() != 2 {
return Err("glm5-tp shard cols: quantized K slices are 2D-only".into());
}
let (outer, inner) = outer_rows(ne);
if cols.end > inner {
return Err(format!("shard cols {cols:?} exceed inner axis {inner}").into());
}
if !(cols.start * row_bytes).is_multiple_of(inner)
|| !(take * row_bytes).is_multiple_of(inner)
{
return Err(format!(
"glm5-tp shard cols: {cols:?} of {inner} does not land on a block boundary \
(row_bytes {row_bytes})"
)
.into());
}
let off = cols.start * row_bytes / inner;
let keep = take * row_bytes / inner;
let host = src_engine.dtoh_u8(bytes)?;
let mut out = Vec::with_capacity(outer * keep);
for r in 0..outer {
let base = r * row_bytes + off;
out.extend_from_slice(&host[base..base + keep]);
}
let mut ne2 = ne.clone();
ne2[0] = take as u64;
Ok(GpuTensor::Quant {
bytes: dst.htod_bytes(&out)?,
qtype: *qtype,
row_bytes: keep,
ne: ne2,
scale: *scale,
rp: false,
fp8: None,
rp4: None,
blk: None,
f16: None,
#[cfg(memra_cutlass)]
cutlass: None,
})
}
}
}
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: if glm5_tp_symmetric_on() {
shard_cols(e, dst, &la.wo, wr * ql..(wr + 1) * ql)?
} else {
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),
None,
None,
)?;
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();
if glm5_tp_symmetric_on() {
let mut partials = Vec::with_capacity(ranks);
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] };
partials.push(if rows_exact {
dev.matmul_rows_exact(&la.wo, gated_refs[r], t)?
} else {
dev.matmul(&la.wo, gated_refs[r], t)?
});
}
let mut y = partials.remove(0);
let reduced = if partials.len() == 1 && rt.ar_1stage_available() {
let mut peer_y = partials.remove(0);
let done = rt.ar_1stage(e, &mut [&mut y, &mut peer_y], t * n_embd)?;
if !done {
partials.insert(0, peer_y);
}
done
} else {
false
};
if !reduced {
for (r, peer_y) in partials.iter().enumerate() {
let landed =
crate::tp_transport::return_row_to_root(&hop, r + 1, peer_y, t * n_embd)?;
let mut dst = y.slice_mut(0..t * n_embd);
e.axpy_into(&landed, 1.0, &mut dst, t * n_embd)?;
}
}
return Ok(y);
}
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_partials_sym(
e: &Engine,
la_root: &KdaAttnLayer,
x_by_rank: &[&CudaSlice<f32>],
t: usize,
eps: f32,
cache: &mut Cache,
il: usize,
arm: ConvArm,
) -> Result<Vec<CudaSlice<f32>>, Box<dyn std::error::Error>> {
let tp = la_root
.tp
.as_ref()
.ok_or("kda_tp_cached_sym called on an unsharded layer")?;
let rt = &tp.rt;
let ranks = rt.ranks();
if ranks != 2 || x_by_rank.len() != 2 {
return Err("kda_tp_cached_sym: this arm is two ranks".into());
}
if !glm5_tp_symmetric_on() || !rt.ar_1stage_available() {
return Err(
"kda_tp_cached_sym: needs MEMRA_GLM5_TP_SYMMETRIC and a real two-device group".into(),
);
}
let states = ensure_kda_tp_state(e, rt, la_root, cache, il)?;
let mut partials: Vec<CudaSlice<f32>> = Vec::with_capacity(ranks);
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 RecurLayer {
conv_state,
ssm_state,
ssm_state_alt,
} = &mut states[r];
let gated = crate::kda::kda_core_gated(
dev,
la,
x_by_rank[r],
t,
eps,
conv_state,
ssm_state,
ssm_state_alt,
arm,
crate::kda::KdaStash::None,
None,
None,
None,
)?;
std::mem::swap(ssm_state, ssm_state_alt);
partials.push(dev.matmul(&la.wo, &gated, t)?);
}
Ok(partials)
}
#[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)?,
wk_b16: None,
wv_b16: None,
wo: if glm5_tp_expert_split_on() {
shard_cols(e, dst, &la.wo, wr * hl * g.d_v..(wr + 1) * hl * g.d_v)?
} else {
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) {
let wo_shape = if glm5_tp_expert_split_on() {
"row-parallel-partial-sums"
} else {
"column-over-gather"
};
eprintln!(
"[glm5-tp-mla] head shard armed: ranks={ranks} heads_per_rank={hl} kv_rank={} \
latent=replicated indexer=replicated wo={wo_shape} 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() || m.glm5_tp_split.is_some() {
return Err("arm_moe_ep: layer is already expert-armed".into());
}
if glm5_tp_expert_split_on() {
let n_embd = m.gate_inp.in_features();
m.dev_exps = None;
m.glm5_tp_split = Some(shard_moe_layer_split(e, rt, m, n_embd)?);
return Ok(());
}
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(())
}
pub struct Glm5TpGlue {
pub attn_norm: GpuTensor,
pub post_attn_norm: GpuTensor,
pub hyper: crate::hyper::HyperLayer,
pub router: Option<Glm5TpPeerRouter>,
pub dense: Option<Glm5TpPeerDense>,
pub shexp_root: Option<Glm5TpShexpHalf>,
pub shexp_peer: Option<Glm5TpShexpHalf>,
}
pub struct Glm5TpShexpHalf {
pub gate: GpuTensor,
pub up: GpuTensor,
pub down: GpuTensor,
}
pub struct Glm5TpPeerDense {
pub ffn_gate: GpuTensor,
pub ffn_up: GpuTensor,
pub ffn_down: GpuTensor,
pub ffn_down_pqs: Option<GpuTensor>,
}
pub struct Glm5TpPeerRouter {
pub gate_inp: GpuTensor,
pub exp_probs_b_dev: CudaSlice<f32>,
pub active_experts_dev: CudaSlice<u8>,
}
fn replicate_tensor(
src_engine: &Engine,
dst: &Engine,
t: &GpuTensor,
) -> Result<GpuTensor, Box<dyn std::error::Error>> {
let outer = match t {
GpuTensor::Float { ne, .. }
| GpuTensor::FloatBf16 { ne, .. }
| GpuTensor::Quant { ne, .. } => outer_rows(ne).0,
};
shard_rows(src_engine, dst, t, 0..outer)
}
fn replicate_f32(
src_engine: &Engine,
dst: &Engine,
x: &CudaSlice<f32>,
) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
let host = src_engine.dtoh(x)?;
dst.htod(&host)
}
fn replicate_site(
src_engine: &Engine,
dst: &Engine,
s: &crate::hyper::HyperSite,
) -> Result<crate::hyper::HyperSite, Box<dyn std::error::Error>> {
Ok(crate::hyper::HyperSite {
fn_w: replicate_f32(src_engine, dst, &s.fn_w)?,
base: replicate_f32(src_engine, dst, &s.base)?,
scale: replicate_f32(src_engine, dst, &s.scale)?,
})
}
pub(crate) fn replicate_layer_glue(
e: &Engine,
rt: &Arc<Glm5TpRt>,
attn_norm: &GpuTensor,
post_attn_norm: &GpuTensor,
hyper: &crate::hyper::HyperLayer,
ffn: &crate::hybrid::Ffn,
) -> Result<Vec<Glm5TpGlue>, Box<dyn std::error::Error>> {
let moe = match ffn {
crate::hybrid::Ffn::Moe(m) => Some(m),
crate::hybrid::Ffn::Dense { .. } => None,
};
let mut out = Vec::with_capacity(rt.peers.len());
for peer in &rt.peers {
let dense = None;
let router = match moe {
Some(m) => Some(Glm5TpPeerRouter {
gate_inp: replicate_tensor(e, peer, &m.gate_inp)?,
exp_probs_b_dev: replicate_f32(e, peer, &m.exp_probs_b_dev)?,
active_experts_dev: {
let host = e.dtoh_u8(&m.active_experts_dev)?;
peer.htod_bytes(&host)?
},
}),
None => None,
};
out.push(Glm5TpGlue {
attn_norm: replicate_tensor(e, peer, attn_norm)?,
post_attn_norm: replicate_tensor(e, peer, post_attn_norm)?,
hyper: crate::hyper::HyperLayer {
attn: replicate_site(e, peer, &hyper.attn)?,
mlp: replicate_site(e, peer, &hyper.mlp)?,
},
router,
dense,
shexp_root: None,
shexp_peer: None,
});
}
Ok(out)
}
pub fn glm5_tp_symmetric_on() -> bool {
std::env::var("MEMRA_GLM5_TP_SYMMETRIC").as_deref() == Ok("1")
}
pub struct Glm5TpSplitExps {
pub rt: Arc<Glm5TpRt>,
pub slabs: Vec<EpRankSlab>,
pub half_ff: usize,
pub gate_stride: usize,
pub up_stride: usize,
pub down_stride: usize,
pub down_row_bytes: usize,
pub ptr_rows: Vec<CudaSlice<u64>>,
}
impl Glm5TpSplitExps {
pub fn ranks(&self) -> usize {
self.slabs.len()
}
}
pub fn tp_ar_1stage_on() -> bool {
std::env::var("MEMRA_TP_AR_1STAGE").as_deref() == Ok("1")
}
pub fn glm5_tp_expert_split_on() -> bool {
std::env::var("MEMRA_GLM5_TP_EXPERT_SPLIT").as_deref() == Ok("1")
}
static SPLIT_MARKED: AtomicBool = AtomicBool::new(false);
pub(crate) fn shard_moe_layer_split(
e: &Engine,
rt: &Arc<Glm5TpRt>,
m: &crate::hybrid::MoeWeights,
n_embd: usize,
) -> Result<Glm5TpSplitExps, Box<dyn std::error::Error>> {
use crate::tp_expert_split::{split_cols, split_rows};
let ranks = rt.ranks();
let mut slabs = Vec::with_capacity(ranks);
let (mut gate_stride, mut up_stride, mut down_stride, mut down_row_bytes, mut half_ff) =
(0usize, 0usize, 0usize, 0usize, 0usize);
for r in 0..ranks {
let dev = rank_engine(e, rt, r);
let g = split_rows(&m.gate_exps, ranks, r)?;
let u = split_rows(&m.up_exps, ranks, r)?;
let d = split_cols(&m.down_exps, ranks, r)?;
if g.out_f != u.out_f || g.out_f != d.in_f {
return Err(format!(
"glm5-tp expert split: rank {r} halves disagree (gate out {} up out {} down in {})",
g.out_f, u.out_f, d.in_f
)
.into());
}
if d.out_f != n_embd {
return Err(format!(
"glm5-tp expert split: down out {} is not the hidden width {n_embd}",
d.out_f
)
.into());
}
half_ff = g.out_f;
gate_stride = g.expert_stride;
up_stride = u.expert_stride;
down_stride = d.expert_stride;
down_row_bytes = d.row_bytes;
slabs.push(EpRankSlab {
gate: dev.htod_bytes_padded(&g.bytes, 8)?,
up: dev.htod_bytes_padded(&u.bytes, 8)?,
down: dev.htod_bytes_padded(&d.bytes, 144)?,
n_experts: m.gate_exps.n_expert,
});
}
if !SPLIT_MARKED.swap(true, Ordering::Relaxed) {
eprintln!(
"[glm5-tp-split] expert TENSOR-parallel armed: every rank holds half of ALL {} \
experts, half_ff={half_ff} (whole-expert ownership pays the busier rank: \
E[max]=5.094 of 8, a 1.571x expert half where the split is a deterministic 2x) \
transport={} performance_claim=false",
m.gate_exps.n_expert,
rt.transport.name(),
);
}
let n_expert_all = m.gate_exps.n_expert;
let mut ptr_rows = Vec::with_capacity(ranks);
for (r, slab) in slabs.iter().enumerate() {
use cudarc::driver::DevicePtr;
let dev = rank_engine(e, rt, r);
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_all];
for ex in 0..n_expert_all {
host[ex] = pg + (ex * gate_stride) as u64;
host[n_expert_all + ex] = pu + (ex * up_stride) as u64;
host[2 * n_expert_all + ex] = pd + (ex * down_stride) as u64;
}
ptr_rows.push(dev.htod_u64(&host)?);
}
Ok(Glm5TpSplitExps {
rt: rt.clone(),
slabs,
half_ff,
gate_stride,
up_stride,
down_stride,
down_row_bytes,
ptr_rows,
})
}
#[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");
}
}