use std::sync::atomic::{AtomicBool, AtomicUsize, Ordering};
use std::sync::{Arc, Mutex, OnceLock};
use cudarc::driver::{CudaContext, CudaEvent, CudaSlice, CudaStream};
use crate::Engine;
pub fn pp_cuts(n_layers: usize) -> Option<Vec<usize>> {
let n_st: usize = match std::env::var("MEMRA_PP_STAGES") {
Ok(v) if v.is_empty() || v == "0" || v == "1" => return None,
Ok(v) => match v.parse::<usize>() {
Ok(n) => n,
Err(_) => {
warn_bad_once(&format!("MEMRA_PP_STAGES={v} unparseable; door stays OFF"));
return None;
}
},
Err(_) => return None,
};
if n_st < 2 || n_st > n_layers {
warn_bad_once(&format!(
"MEMRA_PP_STAGES={n_st} outside [2, n_layers={n_layers}]; door stays OFF"
));
return None;
}
let mut fence = Vec::with_capacity(n_st + 1);
fence.push(0usize);
if let Ok(s) = std::env::var("MEMRA_PP_SPLITS") {
let parts: Result<Vec<usize>, _> =
s.split(',').map(|p| p.trim().parse::<usize>()).collect();
match parts {
Ok(cuts) if cuts.len() == n_st - 1 => fence.extend(cuts),
_ => {
warn_bad_once(&format!(
"MEMRA_PP_SPLITS={s} invalid (want {} comma-separated cuts); door stays OFF",
n_st - 1
));
return None;
}
}
} else if let Ok(v) = std::env::var("MEMRA_PP_SPLIT") {
if n_st != 2 {
warn_bad_once(&format!(
"MEMRA_PP_SPLIT={v} set with MEMRA_PP_STAGES={n_st}; use MEMRA_PP_SPLITS \
for N>2 — door stays OFF"
));
return None;
}
match v.parse::<usize>() {
Ok(c) => fence.push(c),
Err(_) => {
warn_bad_once(&format!("MEMRA_PP_SPLIT={v} unparseable; door stays OFF"));
return None;
}
}
} else {
for s in 1..n_st {
fence.push(s * n_layers / n_st);
}
}
fence.push(n_layers);
for w in fence.windows(2) {
if w[0] >= w[1] {
warn_bad_once(&format!(
"pp stage fence {fence:?} not strictly increasing over [0, {n_layers}]; \
door stays OFF"
));
return None;
}
}
Some(fence)
}
pub fn pp2_split(n_layers: usize) -> Option<usize> {
pp_cuts(n_layers).filter(|f| f.len() == 3).map(|f| f[1])
}
pub fn stage_of(fence: &[usize], il: usize) -> usize {
debug_assert!(fence.len() >= 2);
match fence[1..fence.len() - 1].binary_search(&il) {
Ok(k) => k + 1,
Err(k) => k,
}
}
pub fn pp2_streams_off() -> bool {
matches!(std::env::var("MEMRA_PP_STREAMS").as_deref(), Ok("0"))
}
pub fn pp_multi_stream_same_device() -> bool {
let stages_open = std::env::var("MEMRA_PP_STAGES")
.map(|v| v.parse::<usize>().map(|n| n >= 2).unwrap_or(false))
.unwrap_or(false);
let devices = std::env::var("MEMRA_PP_DEVICES").ok().filter(|v| !v.is_empty());
if (!stages_open && devices.is_none()) || pp2_streams_off() {
return false;
}
match devices {
None => true, Some(s) => {
let mut v: Vec<&str> = s.split(',').map(|p| p.trim()).collect();
let n = v.len();
v.sort_unstable();
v.dedup();
v.len() < n }
}
}
pub fn pp2_overlap() -> bool {
matches!(std::env::var("MEMRA_PP_OVERLAP").as_deref(), Ok("1"))
}
pub fn pp_shard_off() -> bool {
matches!(std::env::var("MEMRA_PP_SHARD").as_deref(), Ok("0"))
}
fn pp2_devices_env() -> Option<String> {
std::env::var("MEMRA_PP_DEVICES").ok().filter(|v| !v.is_empty())
}
static WARNED_BAD: AtomicBool = AtomicBool::new(false);
fn warn_bad_once(msg: &str) {
if !WARNED_BAD.swap(true, Ordering::Relaxed) {
eprintln!("[pp] {msg}");
}
}
static WARNED_UNWIRED: AtomicBool = AtomicBool::new(false);
pub fn warn_unwired_once(path: &str) {
let open = std::env::var("MEMRA_PP_STAGES")
.map(|v| !v.is_empty() && v != "0" && v != "1")
.unwrap_or(false);
if open && !WARNED_UNWIRED.swap(true, Ordering::Relaxed) {
eprintln!(
"[pp] MEMRA_PP_STAGES set but `{path}` has no pp arm at this N; running unsplit"
);
}
}
pub struct StageRt {
pub dev: usize,
pub ctx: Arc<CudaContext>,
pub stream: Arc<CudaStream>,
engine: Option<Engine>,
}
struct BoundarySlot {
buf: Mutex<Option<CudaSlice<f32>>>,
ev_tx: CudaEvent,
ev_rx: CudaEvent,
}
struct BoundaryRt {
slots: [BoundarySlot; 2],
step: AtomicUsize,
cross: bool,
}
pub struct PpNRt {
stages: Vec<StageRt>,
boundaries: Vec<BoundaryRt>,
cross_any: bool,
readback: Arc<CudaStream>,
}
pub type Pp2Rt = PpNRt;
static RTN: OnceLock<Result<PpNRt, String>> = OnceLock::new();
impl PpNRt {
pub fn get(e: &Engine) -> Result<&'static PpNRt, Box<dyn std::error::Error>> {
RTN.get_or_init(|| Self::build(e).map_err(|err| err.to_string()))
.as_ref()
.map_err(|s| -> Box<dyn std::error::Error> { s.clone().into() })
}
fn build(e: &Engine) -> Result<PpNRt, Box<dyn std::error::Error>> {
let primary_dev = e.ctx().ordinal();
let devices: Vec<usize> = match pp2_devices_env() {
Some(s) => {
let parts: Result<Vec<usize>, _> =
s.split(',').map(|p| p.trim().parse::<usize>()).collect();
match parts {
Ok(v) if v.len() >= 2 => v,
_ => {
return Err(format!(
"MEMRA_PP_DEVICES={s} unparseable (want <d0>,..,<dN-1> e.g. 0,1,2,3)"
)
.into())
}
}
}
None => {
let n_st = std::env::var("MEMRA_PP_STAGES")
.ok()
.and_then(|v| v.parse::<usize>().ok())
.filter(|&n| n >= 2)
.unwrap_or(2);
vec![primary_dev; n_st]
}
};
if let Ok(v) = std::env::var("MEMRA_PP_STAGES") {
if let Ok(n) = v.parse::<usize>() {
if n >= 2 && n != devices.len() {
return Err(format!(
"MEMRA_PP_DEVICES lists {} devices but MEMRA_PP_STAGES={n} — \
refusing an ambiguous placement",
devices.len()
)
.into());
}
}
}
let n_st = devices.len();
let cross_any = devices.iter().any(|&d| d != devices[0]);
let mut used: Vec<usize> = devices.clone();
used.push(primary_dev);
used.sort_unstable();
used.dedup();
if used.len() > 1 {
let n = cudarc::driver::result::device::get_count()? as usize;
for &d in &used {
if d >= n {
return Err(format!(
"MEMRA_PP_DEVICES={devices:?} but only {n} CUDA device(s) present"
)
.into());
}
}
for &a in &used {
for &b in &used {
if a == b {
continue;
}
let da = cudarc::driver::result::device::get(a as i32)?;
let db = cudarc::driver::result::device::get(b as i32)?;
let mut can: i32 = 0;
unsafe {
cudarc::driver::sys::cuDeviceCanAccessPeer(&mut can, da, db).result()?;
}
if can == 0 {
return Err(format!(
"device {a} cannot peer-access device {b} (cuDeviceCanAccessPeer=0); \
ppN cross-device needs P2P — refusing a silently-staged path"
)
.into());
}
}
}
}
let mk_stage = |dev: usize, s: usize| -> Result<StageRt, Box<dyn std::error::Error>> {
if dev == primary_dev && s == 0 {
let ctx = e.ctx().clone();
let stream = ctx.new_stream()?;
Ok(StageRt { dev, ctx, stream, engine: None })
} else {
let eng = Engine::new(dev)?;
let ctx = eng.ctx().clone();
let stream = ctx.new_stream()?;
Ok(StageRt { dev, ctx, stream, engine: Some(eng) })
}
};
let mut stages = Vec::with_capacity(n_st);
for (s, &d) in devices.iter().enumerate() {
stages.push(mk_stage(d, s)?);
}
if used.len() > 1 {
let ctx_of = |d: usize| -> &Arc<CudaContext> {
if d == primary_dev {
e.ctx()
} else {
&stages.iter().find(|s| s.dev == d).unwrap().ctx
}
};
for &a in &used {
for &b in &used {
if a == b {
continue;
}
ctx_of(a).bind_to_thread()?;
let rc = unsafe {
cudarc::driver::sys::cuCtxEnablePeerAccess(ctx_of(b).cu_ctx(), 0)
};
use cudarc::driver::sys::cudaError_enum as E;
if rc != E::CUDA_SUCCESS && rc != E::CUDA_ERROR_PEER_ACCESS_ALREADY_ENABLED {
return Err(format!(
"cuCtxEnablePeerAccess(dev{a} -> dev{b}) failed: {rc:?}"
)
.into());
}
}
}
for &owner in &used {
for &accessor in &used {
if owner == accessor {
continue;
}
let dev = cudarc::driver::result::device::get(owner as i32)?;
let mut pool: cudarc::driver::sys::CUmemoryPool = std::ptr::null_mut();
unsafe {
cudarc::driver::sys::cuDeviceGetDefaultMemPool(&mut pool, dev).result()?;
}
let desc = cudarc::driver::sys::CUmemAccessDesc {
location: cudarc::driver::sys::CUmemLocation {
type_: cudarc::driver::sys::CUmemLocationType::CU_MEM_LOCATION_TYPE_DEVICE,
id: accessor as i32,
},
flags: cudarc::driver::sys::CUmemAccess_flags::CU_MEM_ACCESS_FLAGS_PROT_READWRITE,
};
let rc = unsafe { cudarc::driver::sys::cuMemPoolSetAccess(pool, &desc, 1) };
if rc != cudarc::driver::sys::cudaError_enum::CUDA_SUCCESS {
return Err(format!(
"cuMemPoolSetAccess(dev{owner} pool -> dev{accessor}) failed: {rc:?}"
)
.into());
}
}
}
for (owner, accessor) in [(stages[0].dev, stages[1].dev), (stages[1].dev, stages[0].dev)] {
let dev = cudarc::driver::result::device::get(owner as i32)?;
let mut pool: cudarc::driver::sys::CUmemoryPool = std::ptr::null_mut();
unsafe {
cudarc::driver::sys::cuDeviceGetDefaultMemPool(&mut pool, dev).result()?;
}
let desc = cudarc::driver::sys::CUmemAccessDesc {
location: cudarc::driver::sys::CUmemLocation {
type_: cudarc::driver::sys::CUmemLocationType::CU_MEM_LOCATION_TYPE_DEVICE,
id: accessor as i32,
},
flags: cudarc::driver::sys::CUmemAccess_flags::CU_MEM_ACCESS_FLAGS_PROT_READWRITE,
};
let rc = unsafe { cudarc::driver::sys::cuMemPoolSetAccess(pool, &desc, 1) };
if rc != cudarc::driver::sys::cudaError_enum::CUDA_SUCCESS {
return Err(format!(
"cuMemPoolSetAccess(dev{owner} pool -> dev{accessor}) failed: {rc:?}"
)
.into());
}
}
e.ctx().bind_to_thread()?;
eprintln!(
"[pp] cross-device transport: {} (cudaMemcpyPeerAsync per cross boundary; \
peer + default-pool access granted all pairs over {used:?}; weight home: {})",
devices
.iter()
.enumerate()
.map(|(s, d)| format!("stage{s}=dev{d}"))
.collect::<Vec<_>>()
.join(" "),
if pp_shard_off() {
format!("dev{primary_dev} (MEMRA_PP_SHARD=0 bring-up placement)")
} else {
"per-stage (sharded loader)".to_string()
}
);
}
let mk_slot = |tx: &StageRt, rx: &StageRt| -> Result<BoundarySlot, Box<dyn std::error::Error>> {
Ok(BoundarySlot {
buf: Mutex::new(None),
ev_tx: tx.ctx.new_event(None)?,
ev_rx: rx.ctx.new_event(None)?,
})
};
let mut boundaries = Vec::with_capacity(n_st - 1);
for b in 0..n_st - 1 {
let (tx, rx) = (&stages[b], &stages[b + 1]);
boundaries.push(BoundaryRt {
slots: [mk_slot(tx, rx)?, mk_slot(tx, rx)?],
step: AtomicUsize::new(0),
cross: tx.dev != rx.dev,
});
}
let readback = stages[n_st - 1].ctx.new_stream()?;
Ok(PpNRt { stages, boundaries, cross_any, readback })
}
pub fn n_stages(&self) -> usize {
self.stages.len()
}
pub fn cross_device(&self) -> bool {
self.cross_any
}
pub fn engine<'a>(&'a self, s: usize, primary: &'a Engine) -> &'a Engine {
self.stages[s].engine.as_ref().unwrap_or(primary)
}
pub fn enter(&self, s: usize) -> memra_runtime::StreamOverride {
memra_runtime::push_stream_override(self.stages[s].stream.clone())
}
pub fn tx(&self, b: usize, x: &CudaSlice<f32>, n: usize)
-> Result<usize, Box<dyn std::error::Error>> {
assert_eq!(x.len(), n, "pp tx: residual length mismatch");
let bd = &self.boundaries[b];
let slot_idx = if pp2_overlap() {
bd.step.fetch_add(1, Ordering::Relaxed) % 2
} else {
0
};
let sl = &bd.slots[slot_idx];
let s_tx = &self.stages[b].stream;
s_tx.wait(&sl.ev_rx)?;
let mut guard = sl.buf.lock().unwrap();
if guard.as_ref().map(|bf| bf.len() != n).unwrap_or(true) {
let s_rx = &self.stages[b + 1].stream;
*guard = Some(s_rx.alloc_zeros::<f32>(n)?);
s_rx.synchronize()?;
}
let buf = guard.as_mut().unwrap();
if !bd.cross {
s_tx.memcpy_dtod(x, buf)?;
} else {
use cudarc::driver::{DevicePtr, DevicePtrMut};
let (sp, _g0) = x.device_ptr(s_tx);
let (dp, _g1) = buf.device_ptr_mut(s_tx);
self.stages[b].ctx.bind_to_thread()?;
unsafe {
cudarc::driver::result::memcpy_peer_async(
self.stages[b + 1].ctx.cu_ctx(),
dp,
self.stages[b].ctx.cu_ctx(),
sp,
n * std::mem::size_of::<f32>(),
s_tx.cu_stream(),
)?;
}
}
sl.ev_tx.record(s_tx)?;
Ok(slot_idx)
}
pub fn rx(&self, b: usize, slot_idx: usize, n: usize)
-> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
let sl = &self.boundaries[b].slots[slot_idx];
let s_rx = &self.stages[b + 1].stream;
s_rx.wait(&sl.ev_tx)?;
let guard = sl.buf.lock().unwrap();
let buf = guard.as_ref().expect("pp rx before tx");
let mut work = unsafe { s_rx.alloc::<f32>(n)? };
s_rx.memcpy_dtod(buf, &mut work)?;
sl.ev_rx.record(s_rx)?;
Ok(work)
}
pub fn record_done(&self) -> Result<CudaEvent, Box<dyn std::error::Error>> {
let last = &self.stages[self.stages.len() - 1];
let ev = last.ctx.new_event(None)?;
ev.record(&last.stream)?;
Ok(ev)
}
pub fn readback_stream(&self) -> &Arc<CudaStream> {
&self.readback
}
}
pub struct PendingLogits {
logits: CudaSlice<f32>,
ev: CudaEvent,
rb: Arc<CudaStream>,
}
impl PendingLogits {
pub fn new(logits: CudaSlice<f32>, ev: CudaEvent, rb: Arc<CudaStream>) -> Self {
PendingLogits { logits, ev, rb }
}
pub fn wait(self) -> Result<Vec<f32>, Box<dyn std::error::Error>> {
self.rb.wait(&self.ev)?;
let host = self.rb.clone_dtoh(&self.logits)?;
self.rb.synchronize()?;
Ok(host)
}
}
pub fn new_cache(e: &Engine, cfg: &memra_gguf::config::ModelConfig, max_ctx: usize)
-> Result<crate::cache::Cache, Box<dyn std::error::Error>> {
let n_trunk = (cfg.n_layer - cfg.nextn_predict_layers) as usize;
if let Some(fence) = pp_cuts(n_trunk) {
if pp2_devices_env().is_some() && !pp2_streams_off() {
let rt = PpNRt::get(e)?;
let n_st = fence.len() - 1;
assert_eq!(
rt.n_stages(), n_st,
"PpNRt stage count {} != fence stages {n_st}", rt.n_stages()
);
let devs: Vec<&dyn memra_kv::KvDev> =
(0..n_st).map(|s| rt.engine(s, e) as &dyn memra_kv::KvDev).collect();
let cache = crate::cache::Cache::new_ppn(&devs, &fence, cfg, max_ctx)?;
sync_stages_after_load(e, n_trunk)?;
return Ok(cache);
}
if !pp2_streams_off() {
let cache = crate::cache::Cache::new(e, cfg, max_ctx)?;
sync_stages_after_load(e, n_trunk)?;
return Ok(cache);
}
}
crate::cache::Cache::new(e, cfg, max_ctx)
}
pub fn sync_stages_after_load(e: &Engine, n_trunk: usize)
-> Result<(), Box<dyn std::error::Error>> {
if pp2_streams_off() || pp_cuts(n_trunk).is_none() {
return Ok(());
}
let rt = PpNRt::get(e)?;
for s in 0..rt.n_stages() {
rt.stages[s].ctx.bind_to_thread()?;
unsafe {
cudarc::driver::sys::cuCtxSynchronize().result()?;
}
}
e.ctx().bind_to_thread()?;
unsafe {
cudarc::driver::sys::cuCtxSynchronize().result()?;
}
Ok(())
}
pub fn layer_engine<'a>(e: &'a Engine, n_trunk: usize, il: usize)
-> Result<&'a Engine, Box<dyn std::error::Error>> {
if pp_shard_off() || pp2_devices_env().is_none() || pp2_streams_off() {
return Ok(e);
}
let Some(fence) = pp_cuts(n_trunk) else { return Ok(e) };
let rt = PpNRt::get(e)?;
let s = stage_of(&fence, il.min(n_trunk - 1));
Ok(rt.engine(s, e))
}