use crate::Engine;
use crate::glm5_decode_graph::{CapCtx, capture_error, step};
use crate::hybrid::{HybridModel, Mixer};
use crate::hyper::{HyperDecodeWs, HyperTopology};
use cudarc::driver::{CudaGraph, CudaSlice};
use memra_kv::Cache;
type Res<T> = Result<T, Box<dyn std::error::Error>>;
pub fn on() -> bool {
std::env::var("MEMRA_GLM5_TP_SYM_GRAPH").as_deref() == Ok("1")
}
pub static GLM5_TP_SYM_GRAPH_TOKENS: std::sync::atomic::AtomicU64 =
std::sync::atomic::AtomicU64::new(0);
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
enum Seg {
Kda(usize),
MlaPre(usize),
MlaFfn(usize),
}
enum Piece {
Graph {
segs: Vec<Seg>,
graphs: [[CudaGraph; 2]; 2],
},
MlaMid(usize),
Eager(usize),
}
#[derive(Debug, PartialEq, Eq)]
enum Item {
Seg(Seg),
MlaMid(usize),
Eager(usize),
}
pub fn mla_pieces_on() -> bool {
std::env::var("MEMRA_GLM5_TP_SYM_GRAPH_MLA").as_deref() != Ok("0")
}
unsafe impl Send for SymGraphState {}
pub(crate) struct SymGraphState {
lo: usize,
hi: usize,
x_io: CudaSlice<f32>,
x_peer_io: CudaSlice<f32>,
pos_peer_io: CudaSlice<i32>,
ws_root: HyperDecodeWs,
ws_peer: HyperDecodeWs,
mla_a_root: Option<CudaSlice<f32>>,
mla_a_peer: Option<CudaSlice<f32>>,
pieces: Vec<Piece>,
warm: bool,
phase: usize,
failed: bool,
}
fn plan_items(m: &HybridModel, lo: usize, hi: usize) -> Vec<Item> {
let mla_split = mla_pieces_on() && Engine::mla_seg_ws_on();
let mut out = Vec::with_capacity(hi.saturating_sub(lo) * 3);
for il in lo..hi {
let moe = matches!(&m.layers[il].ffn, crate::hybrid::Ffn::Moe(_));
let dense_local = m.layers[il]
.tp_glue
.first()
.is_some_and(|g| g.dense.is_some());
match &m.layers[il].mixer {
Mixer::Kda(la) if la.tp.is_some() && (moe || dense_local) => {
out.push(Item::Seg(Seg::Kda(il)))
}
Mixer::Mla(mla) if mla.tp.is_some() && moe && mla_split => {
out.push(Item::Seg(Seg::MlaPre(il)));
out.push(Item::MlaMid(il));
out.push(Item::Seg(Seg::MlaFfn(il)));
}
_ => out.push(Item::Eager(il)),
}
}
out
}
fn group_runs(items: &[Item]) -> Vec<Vec<Seg>> {
let mut runs = Vec::new();
let mut cur: Vec<Seg> = Vec::new();
for it in items {
match it {
Item::Seg(sg) => cur.push(*sg),
_ => {
if !cur.is_empty() {
runs.push(std::mem::take(&mut cur));
}
}
}
}
if !cur.is_empty() {
runs.push(cur);
}
runs
}
fn seg_layer(sg: Seg) -> usize {
match sg {
Seg::Kda(il) | Seg::MlaPre(il) | Seg::MlaFfn(il) => il,
}
}
fn capture_pair<F>(e: &Engine, peer: &Engine, ctx: &CapCtx, mut body: F) -> Res<[CudaGraph; 2]>
where
F: FnMut() -> Res<()>,
{
use cudarc::driver::sys::{CUgraphInstantiate_flags, CUstreamCaptureMode};
for eng in [e, peer] {
let _m = eng.gpu.enter_main()?;
step(
eng,
ctx,
"synchronize(before begin_capture)",
eng.stream().synchronize().map_err(Into::into),
)?;
}
let tracking: Vec<bool> = [e, peer]
.iter()
.map(|eng| eng.ctx().is_event_tracking())
.collect();
for (i, eng) in [e, peer].iter().enumerate() {
if tracking[i] {
unsafe { eng.ctx().disable_event_tracking() };
}
}
let out = (|| -> Res<[CudaGraph; 2]> {
for eng in [e, peer] {
let _m = eng.gpu.enter_main()?;
step(
eng,
ctx,
"cuStreamBeginCapture(RELAXED)",
eng.stream()
.begin_capture(CUstreamCaptureMode::CU_STREAM_CAPTURE_MODE_RELAXED)
.map_err(Into::into),
)?;
}
let body_res = body();
let mut graphs = Vec::with_capacity(2);
let mut end_err: Option<Box<dyn std::error::Error>> = None;
for eng in [e, peer] {
let _m = eng.gpu.enter_main()?;
match eng.stream().end_capture(
CUgraphInstantiate_flags::CUDA_GRAPH_INSTANTIATE_FLAG_AUTO_FREE_ON_LAUNCH,
) {
Ok(Some(g)) => graphs.push(g),
Ok(None) => {
end_err.get_or_insert_with(|| "capture produced no graph".into());
}
Err(err) => {
end_err.get_or_insert_with(|| Box::new(err) as Box<dyn std::error::Error>);
}
}
}
if let Err(err) = body_res {
capture_error(e, ctx, "capture body", &err);
return Err(err);
}
if let Some(err) = end_err {
capture_error(e, ctx, "cuStreamEndCapture+cuGraphInstantiate", &err);
return Err(err);
}
let mut it = graphs.into_iter();
let g0 = it.next().ok_or("no root graph")?;
let g1 = it.next().ok_or("no peer graph")?;
{
let _m = e.gpu.enter_main()?;
step(
e,
ctx,
"cuGraphUpload(root)",
g0.upload().map_err(Into::into),
)?;
}
{
let _m = peer.gpu.enter_main()?;
step(
peer,
ctx,
"cuGraphUpload(peer)",
g1.upload().map_err(Into::into),
)?;
}
Ok([g0, g1])
})();
for (i, eng) in [e, peer].iter().enumerate() {
if tracking[i] {
unsafe { eng.ctx().enable_event_tracking() };
}
}
out
}
fn take_state(
m: &HybridModel,
e: &Engine,
peer: &Engine,
topology: &HyperTopology,
cache: &mut Cache,
lo: usize,
hi: usize,
) -> Res<Box<SymGraphState>> {
let n_embd = m.cfg.n_embd as usize;
let width = topology.streams * n_embd;
if let Some(b) = cache.glm5_tp_sym_graph.take() {
match b.downcast::<SymGraphState>() {
Ok(st) if st.lo == lo && st.hi == hi => return Ok(st),
Ok(_) | Err(_) => {}
}
}
let mla_a = (lo..hi).find_map(|il| match &m.layers[il].mixer {
Mixer::Mla(mla) if mla.tp.is_some() => Some(mla.wo.in_features()),
_ => None,
});
Ok(Box::new(SymGraphState {
lo,
hi,
x_io: e.zeros(width)?,
x_peer_io: peer.zeros(width)?,
pos_peer_io: {
let _m = peer.gpu.enter_main()?;
peer.htod_i32(&[0])?
},
ws_root: HyperDecodeWs::new(e, topology, n_embd)?,
ws_peer: {
let _m = peer.gpu.enter_main()?;
HyperDecodeWs::new(peer, topology, n_embd)?
},
mla_a_root: mla_a.map(|n| e.zeros(n)).transpose()?,
mla_a_peer: mla_a
.map(|n| {
let _m = peer.gpu.enter_main()?;
peer.zeros(n)
})
.transpose()?,
pieces: Vec::new(),
warm: false,
phase: 0,
failed: false,
}))
}
#[allow(clippy::too_many_arguments)] pub(crate) fn walk_graphed(
m: &HybridModel,
e: &Engine,
peer: &Engine,
rt: &std::sync::Arc<crate::glm5_tp::Glm5TpRt>,
topology: &HyperTopology,
x: &mut CudaSlice<f32>,
x_peer: &mut CudaSlice<f32>,
_ws: &mut HyperDecodeWs,
_ws_peer: &mut HyperDecodeWs,
pos_d: &CudaSlice<i32>,
pos_peer: &CudaSlice<i32>,
lo: usize,
hi: usize,
cache: &mut Cache,
) -> Res<bool> {
let n_embd = m.cfg.n_embd as usize;
let width = topology.streams * n_embd;
if x.len() != width || x_peer.len() != width || pos_d.len() != 1 {
static ONCE: std::sync::Once = std::sync::Once::new();
ONCE.call_once(|| {
eprintln!(
"[glm5-tp-sym-graph] declined: t=1 walks only (x.len()={} pos.len()={})",
x.len(),
pos_d.len()
)
});
return Ok(false);
}
if let Some(reason) = crate::glm5_decode_graph::glm5_graph_process_refusal(e) {
static ONCE: std::sync::Once = std::sync::Once::new();
ONCE.call_once(|| {
eprintln!(
"[glm5-tp-sym-graph] declined: {reason}; this session's symmetric walk stays EAGER"
)
});
return Ok(false);
}
let items = plan_items(m, lo, hi);
if !items.iter().any(|it| matches!(it, Item::Seg(_))) {
return Ok(false);
}
rt.ar_prepare(e, n_embd)?;
let dev = e.ctx().ordinal();
let mut st = take_state(m, e, peer, topology, cache, lo, hi)?;
let r = walk_token(
m, e, peer, rt, topology, x, x_peer, pos_d, pos_peer, lo, hi, cache, dev, &items, &mut st,
);
cache.glm5_tp_sym_graph = Some(st);
r
}
#[allow(clippy::too_many_arguments)] fn walk_token(
m: &HybridModel,
e: &Engine,
peer: &Engine,
rt: &crate::glm5_tp::Glm5TpRt,
topology: &HyperTopology,
x: &mut CudaSlice<f32>,
x_peer: &mut CudaSlice<f32>,
pos_d: &CudaSlice<i32>,
pos_peer: &CudaSlice<i32>,
lo: usize,
hi: usize,
cache: &mut Cache,
dev: usize,
items: &[Item],
st: &mut SymGraphState,
) -> Res<bool> {
let width = topology.streams * m.cfg.n_embd as usize;
if st.failed {
return Ok(false);
}
{
let src = x.slice(0..width);
let mut dst = st.x_io.slice_mut(0..width);
e.stream().memcpy_dtod(&src, &mut dst)?;
}
{
let _m = peer.gpu.enter_main()?;
let src = x_peer.slice(0..width);
let mut dst = st.x_peer_io.slice_mut(0..width);
peer.stream().memcpy_dtod(&src, &mut dst)?;
let ps = pos_peer.slice(0..1);
let mut pd = st.pos_peer_io.slice_mut(0..1);
peer.stream().memcpy_dtod(&ps, &mut pd)?;
}
if !st.warm {
let SymGraphState {
x_io,
x_peer_io,
pos_peer_io,
ws_root,
ws_peer,
..
} = &mut *st;
for il in lo..hi {
m.sym_layer_step(
e,
peer,
rt,
topology,
il,
x_io,
x_peer_io,
ws_root,
ws_peer,
pos_d,
pos_peer_io,
cache,
)?;
}
st.warm = true;
let src = st.x_io.slice(0..width);
let mut dst = x.slice_mut(0..width);
e.stream().memcpy_dtod(&src, &mut dst)?;
return Ok(true);
}
if st.pieces.is_empty() {
let runs = group_runs(items);
let n = runs.len();
let mut pieces: Vec<Piece> = Vec::with_capacity(items.len());
let mut ri = 0usize;
let mut i = 0usize;
while i < items.len() {
match &items[i] {
Item::MlaMid(il) => {
pieces.push(Piece::MlaMid(*il));
i += 1;
}
Item::Eager(il) => {
pieces.push(Piece::Eager(*il));
i += 1;
}
Item::Seg(_) => {
let segs = runs[ri].clone();
let a = seg_layer(segs[0]);
let b = seg_layer(*segs.last().expect("a run has segments")) + 1;
let mut phases: Vec<[CudaGraph; 2]> = Vec::with_capacity(2);
for phase in 0..2 {
let ctx = CapCtx {
dev,
lo,
hi,
run: ri,
runs: n,
a,
b,
phase,
recapture: false,
};
let SymGraphState {
x_io,
x_peer_io,
pos_peer_io,
ws_root,
ws_peer,
mla_a_root,
mla_a_peer,
..
} = &mut *st;
let cap = capture_pair(e, peer, &ctx, || {
for sg in &segs {
match *sg {
Seg::Kda(il) => m.sym_layer_step(
e,
peer,
rt,
topology,
il,
x_io,
x_peer_io,
ws_root,
ws_peer,
pos_d,
pos_peer_io,
cache,
)?,
Seg::MlaPre(il) => m.sym_mla_pre_piece(
e,
peer,
topology,
il,
x_io,
x_peer_io,
ws_root,
ws_peer,
pos_d,
pos_peer_io,
)?,
Seg::MlaFfn(il) => {
let (ar, ap) = match (&*mla_a_root, &*mla_a_peer) {
(Some(ar), Some(ap)) => (ar, ap),
_ => {
return Err(format!(
"layer {il}: sym MLA FFN piece without handoff buffers"
)
.into());
}
};
m.sym_mla_ffn_piece(
e, peer, rt, topology, il, ar, ap, x_io, x_peer_io,
ws_root, ws_peer,
)?
}
}
}
Ok(())
});
match cap {
Ok(g) => phases.push(g),
Err(err) => {
eprintln!(
"[glm5-tp-sym-graph] capture refused for run [{a}, {b}) phase {phase}: {err}; \
this session's symmetric walk stays EAGER (byte-identical)"
);
st.failed = true;
st.pieces.clear();
return Ok(false);
}
}
}
let mut it = phases.into_iter();
let p0 = it.next().ok_or("phase 0 missing")?;
let p1 = it.next().ok_or("phase 1 missing")?;
pieces.push(Piece::Graph {
graphs: [p0, p1],
segs,
});
i += runs[ri].len();
ri += 1;
}
}
}
let n_graph = pieces
.iter()
.filter(|p| matches!(p, Piece::Graph { .. }))
.count();
let n_mid = pieces
.iter()
.filter(|p| matches!(p, Piece::MlaMid(_)))
.count();
let n_eager = pieces
.iter()
.filter(|p| matches!(p, Piece::Eager(_)))
.count();
let n_segs: usize = pieces
.iter()
.map(|p| match p {
Piece::Graph { segs, .. } => segs.len(),
_ => 0,
})
.sum();
st.pieces = pieces;
eprintln!(
"[glm5-tp-sym-graph] engaged: {n_graph} graph piece(s) over {n_segs} segment(s) recorded \
per rank in both ping-pong phases over layers [{lo}, {hi}); {n_mid} MLA middle(s) eager \
between pieces, {n_eager} whole layer(s) eager"
);
}
let phase = st.phase;
let SymGraphState {
x_io,
x_peer_io,
pos_peer_io,
ws_root,
ws_peer,
mla_a_root,
mla_a_peer,
pieces,
..
} = &mut *st;
for piece in pieces.iter() {
match piece {
Piece::Graph { graphs, .. } => {
{
let _m = e.gpu.enter_main()?;
graphs[phase][0].launch()?;
}
{
let _m = peer.gpu.enter_main()?;
graphs[phase][1].launch()?;
}
}
Piece::MlaMid(il) => {
let (ar, ap) = match (mla_a_root.as_mut(), mla_a_peer.as_mut()) {
(Some(ar), Some(ap)) => (ar, ap),
_ => {
return Err(
format!("layer {il}: sym MLA middle without handoff buffers").into(),
);
}
};
m.sym_mla_mid_eager(
e,
peer,
*il,
ws_root,
ws_peer,
pos_d,
pos_peer_io,
cache,
ar,
ap,
)?;
}
Piece::Eager(il) => {
m.sym_layer_step(
e,
peer,
rt,
topology,
*il,
x_io,
x_peer_io,
ws_root,
ws_peer,
pos_d,
pos_peer_io,
cache,
)?;
}
}
}
st.phase ^= 1;
GLM5_TP_SYM_GRAPH_TOKENS.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
{
let src = st.x_io.slice(0..width);
let mut dst = x.slice_mut(0..width);
e.stream().memcpy_dtod(&src, &mut dst)?;
}
Ok(true)
}