use crate::Engine;
use crate::mmq_ffi::{DeviceExpertCsr, ExpertCsr, Fp8GroupedWorkspace};
use crate::parallel::{PRODUCT_MAX_CARDS, STEP37_TRUNK_LAYERS};
use cudarc::driver::{CudaEvent, CudaSlice, DeviceSlice};
use std::ops::Range;
const FP8_BLOCK: usize = 128;
const NATIVE_P2P_PROBE_WORDS: usize = 4096;
const STEP_GROUPED_FP8_EXPERTS: usize = 288;
const STEP_GROUPED_FP8_TOP_K: usize = 8;
const STEP_GROUPED_FP8_WIDTH: usize = 1280;
fn validate_step_expert_activation_limit(limit: Option<f32>) -> Result<(), String> {
if let Some(limit) = limit {
if !limit.is_finite() || limit <= 0.0 {
return Err(format!(
"Step routed-expert activation limit must be positive and finite, got {limit}"
));
}
}
Ok(())
}
pub(crate) fn routes_prestage_on() -> bool {
static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
*ON.get_or_init(|| std::env::var("MEMRA_ROUTES_PRESTAGE").as_deref() == Ok("1"))
}
pub(crate) fn oproj_tail_on() -> bool {
static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
*ON.get_or_init(|| std::env::var("MEMRA_OPROJ_TAIL").as_deref() == Ok("1"))
}
thread_local! {
static OPROJ_TAIL_PENDING: std::cell::Cell<Option<(u64, u64)>> =
const { std::cell::Cell::new(None) };
}
thread_local! {
static OPROJ_TAIL_ELIGIBLE: std::cell::Cell<bool> = const { std::cell::Cell::new(false) };
}
pub(crate) struct OprojTailScope(());
pub(crate) fn oproj_tail_scope() -> OprojTailScope {
OPROJ_TAIL_ELIGIBLE.with(|c| c.set(true));
OprojTailScope(())
}
impl Drop for OprojTailScope {
fn drop(&mut self) {
OPROJ_TAIL_ELIGIBLE.with(|c| c.set(false));
OPROJ_TAIL_PENDING.with(|c| c.set(None));
}
}
thread_local! {
static VERIFY_TCOL: std::cell::Cell<Option<usize>> = const { std::cell::Cell::new(None) };
}
pub(crate) fn set_verify_tcol(c: Option<usize>) {
VERIFY_TCOL.with(|x| x.set(c));
}
pub(crate) fn take_verify_tcol() -> Option<usize> {
VERIFY_TCOL.with(|x| x.take())
}
pub(crate) fn tcol_oproj_on() -> bool {
static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
*ON.get_or_init(|| std::env::var("MEMRA_TCOL_OPROJ").as_deref() == Ok("1"))
}
thread_local! {
static TCOL_OPROJ_DEFER: std::cell::Cell<Option<usize>> = const { std::cell::Cell::new(None) };
static TCOL_OPROJ_STASHED: std::cell::Cell<bool> = const { std::cell::Cell::new(false) };
}
pub(crate) fn set_tcol_oproj_defer(c: Option<usize>) {
TCOL_OPROJ_DEFER.with(|x| x.set(c));
}
pub(crate) fn take_tcol_oproj_defer() -> Option<usize> {
TCOL_OPROJ_DEFER.with(|x| x.take())
}
pub(crate) fn set_tcol_oproj_stashed() {
TCOL_OPROJ_STASHED.with(|x| x.set(true));
}
pub(crate) fn take_tcol_oproj_stashed() -> bool {
TCOL_OPROJ_STASHED.with(|x| x.replace(false))
}
pub(crate) fn oproj_tail_eligible() -> bool {
OPROJ_TAIL_ELIGIBLE.with(|c| c.get())
}
pub(crate) fn take_oproj_tail() -> Option<(u64, u64)> {
OPROJ_TAIL_PENDING.with(|c| c.take())
}
pub(crate) fn set_oproj_tail(v: (u64, u64)) {
OPROJ_TAIL_PENDING.with(|c| c.set(Some(v)));
}
pub(crate) fn rank0_merge_on() -> bool {
static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
*ON.get_or_init(|| std::env::var("MEMRA_RANK0_MERGE").as_deref() == Ok("1"))
}
pub(crate) fn len_mirror_lazy_on() -> bool {
static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
*ON.get_or_init(|| std::env::var("MEMRA_LEN_MIRROR_LAZY").as_deref() == Ok("1"))
}
pub(crate) fn fence_memops_on() -> bool {
static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
*ON.get_or_init(|| std::env::var("MEMRA_FENCE_MEMOPS").as_deref() == Ok("1"))
}
pub(crate) fn moe_direct_on() -> bool {
static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
*ON.get_or_init(|| std::env::var("MEMRA_MOE_DIRECT").as_deref() == Ok("1"))
}
pub(crate) fn oproj_direct_on() -> bool {
static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
*ON.get_or_init(|| std::env::var("MEMRA_OPROJ_DIRECT").as_deref() == Ok("1"))
}
pub(crate) fn raw_copy_bytes(
dst: u64,
src: u64,
bytes: usize,
engine: &Engine,
) -> Result<(), Box<dyn std::error::Error>> {
use cudarc::driver::sys;
let r = unsafe {
sys::cuMemcpyAsync(
dst as sys::CUdeviceptr,
src as sys::CUdeviceptr,
bytes,
engine.stream().cu_stream() as sys::CUstream,
)
};
if r == sys::CUresult::CUDA_SUCCESS {
Ok(())
} else {
Err(format!("raw_copy_bytes: {r:?}").into())
}
}
pub fn step_expert_activation_host(gate: f32, up: f32, limit: Option<f32>) -> f32 {
let silu = gate / (1.0 + (-gate).exp());
match limit {
Some(limit) => silu.min(limit) * up.clamp(-limit, limit),
None => silu * up,
}
}
#[derive(Debug, Clone, PartialEq, Eq)]
struct ExpertOwnerRoutes {
rank: usize,
selected: Vec<usize>,
token_rows: Vec<usize>,
global_pairs: Vec<usize>,
}
fn partition_expert_owner_routes(
expert_count: usize,
ranks: usize,
tokens: usize,
experts_per_token: usize,
selected: &[usize],
) -> Result<Vec<ExpertOwnerRoutes>, String> {
if expert_count == 0
|| ranks == 0
|| tokens == 0
|| experts_per_token == 0
|| expert_count % ranks != 0
{
return Err(format!(
"invalid expert-owner route geometry experts={expert_count} ranks={ranks} \
tokens={tokens} experts_per_token={experts_per_token}"
));
}
let pairs = tokens
.checked_mul(experts_per_token)
.ok_or("expert-owner route count overflow")?;
if selected.len() != pairs {
return Err(format!(
"expert-owner routes {} != {tokens}x{experts_per_token} ({pairs})",
selected.len()
));
}
let per_rank = expert_count / ranks;
let mut owners = (0..ranks)
.map(|rank| ExpertOwnerRoutes {
rank,
selected: Vec::new(),
token_rows: Vec::new(),
global_pairs: Vec::new(),
})
.collect::<Vec<_>>();
for (pair, &expert) in selected.iter().enumerate() {
if expert >= expert_count {
return Err(format!(
"expert-owner route {pair} selects expert {expert} outside 0..{expert_count}"
));
}
let rank = expert / per_rank;
owners[rank].selected.push(expert - rank * per_rank);
owners[rank].token_rows.push(pair / experts_per_token);
owners[rank].global_pairs.push(pair);
}
Ok(owners)
}
fn validate_step_grouped_owner_routes(
expert_count: usize,
tokens: usize,
selected: &[usize],
) -> Result<usize, String> {
if expert_count != STEP_GROUPED_FP8_EXPERTS || tokens == 0 {
return Err(format!(
"official Step owner-grouped FP8 requires {} experts and nonzero tokens, got \
experts={expert_count} tokens={tokens}",
STEP_GROUPED_FP8_EXPERTS
));
}
let pairs = tokens
.checked_mul(STEP_GROUPED_FP8_TOP_K)
.ok_or("official Step owner-grouped FP8 route count overflow")?;
if selected.len() != pairs {
return Err(format!(
"official Step owner-grouped FP8 routes {} != {tokens}x{} ({pairs})",
selected.len(),
STEP_GROUPED_FP8_TOP_K,
));
}
for (token, routes) in selected.chunks_exact(STEP_GROUPED_FP8_TOP_K).enumerate() {
let mut unique = routes.to_vec();
unique.sort_unstable();
unique.dedup();
if unique.len() != STEP_GROUPED_FP8_TOP_K {
return Err(format!(
"official Step owner-grouped FP8 token {token} routes are not top-8 unique: \
{routes:?}"
));
}
}
Ok(pairs)
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
struct WeightedRouteCombineShape {
pairs: usize,
max_pairs: usize,
}
fn validate_weighted_route_combine(
width: usize,
experts_per_token: usize,
max_tokens: usize,
tokens: usize,
owner_global_pairs: &[&[usize]],
route_weights: &[f32],
) -> Result<WeightedRouteCombineShape, String> {
if width == 0
|| experts_per_token == 0
|| max_tokens == 0
|| tokens == 0
|| tokens > max_tokens
|| width > i32::MAX as usize
|| experts_per_token > i32::MAX as usize
|| tokens > i32::MAX as usize
{
return Err(format!(
"invalid weighted route combine geometry width={width} experts_per_token=\
{experts_per_token} tokens={tokens}/{max_tokens}"
));
}
let pairs = tokens
.checked_mul(experts_per_token)
.ok_or("weighted route combine pair count overflow")?;
let max_pairs = max_tokens
.checked_mul(experts_per_token)
.ok_or("weighted route combine capacity overflow")?;
if route_weights.len() != pairs || !route_weights.iter().all(|weight| weight.is_finite()) {
return Err(format!(
"weighted route combine weights {} != pairs {pairs} or contain a non-finite value",
route_weights.len()
));
}
let mut seen = vec![false; pairs];
let mut observed = 0usize;
for pairs_for_owner in owner_global_pairs {
observed = observed
.checked_add(pairs_for_owner.len())
.ok_or("weighted route combine observed pair count overflow")?;
for &pair in *pairs_for_owner {
if pair >= pairs || std::mem::replace(&mut seen[pair], true) {
return Err(format!(
"weighted route combine pair {pair} is outside 0..{pairs} or duplicated"
));
}
}
}
if observed != pairs || seen.iter().any(|present| !present) {
return Err(format!(
"weighted route combine owner schedules cover {observed} of {pairs} canonical pairs"
));
}
Ok(WeightedRouteCombineShape { pairs, max_pairs })
}
fn cache_rank_rows(
rows: &[u8],
tokens: usize,
local_token_bytes: usize,
ranks: usize,
rank: usize,
) -> Result<Vec<u8>, String> {
if ranks == 0 || rank >= ranks {
return Err(format!(
"TP cache rank {rank} is outside a {ranks}-rank layout"
));
}
let global_token_bytes = local_token_bytes
.checked_mul(ranks)
.ok_or("TP cache global token-byte overflow")?;
let expected = tokens
.checked_mul(global_token_bytes)
.ok_or("TP cache row-byte overflow")?;
if rows.len() != expected {
return Err(format!(
"TP cache rows contain {} bytes, expected {tokens}x{global_token_bytes}={expected}",
rows.len()
));
}
let mut shard = Vec::with_capacity(tokens * local_token_bytes);
for token in 0..tokens {
let start = token * global_token_bytes + rank * local_token_bytes;
shard.extend_from_slice(&rows[start..start + local_token_bytes]);
}
Ok(shard)
}
fn parse_step_tp_native_p2p(value: Option<&str>) -> Result<bool, String> {
match value {
None | Some("") | Some("0") => Ok(false),
Some("1") => Ok(true),
Some(value) => Err(format!(
"MEMRA_STEP_TP_NATIVE_P2P={value:?} is invalid; expected 0 or 1"
)),
}
}
pub fn step_tp_native_p2p_enabled() -> Result<bool, String> {
parse_step_tp_native_p2p(std::env::var("MEMRA_STEP_TP_NATIVE_P2P").ok().as_deref())
}
fn parse_step_tp_bulk_p2p(value: Option<&str>) -> Result<bool, String> {
match value {
None | Some("") | Some("0") => Ok(false),
Some("1") => Ok(true),
Some(value) => Err(format!(
"MEMRA_STEP_TP_BULK_P2P={value:?} is invalid; expected 0 or 1"
)),
}
}
pub fn step_tp_bulk_p2p_enabled() -> Result<bool, String> {
parse_step_tp_bulk_p2p(std::env::var("MEMRA_STEP_TP_BULK_P2P").ok().as_deref())
}
fn parse_step_ep_device_arithmetic(value: Option<&str>) -> Result<bool, String> {
match value {
None | Some("") | Some("0") => Ok(false),
Some("1") => Ok(true),
Some(value) => Err(format!(
"MEMRA_STEP_EP_DEVICE_ARITHMETIC={value:?} is invalid; expected 0 or 1"
)),
}
}
fn parse_step_nvfp4_dev_routes(value: Option<&str>) -> Result<bool, String> {
match value {
None | Some("") | Some("0") => Ok(false),
Some("1") => Ok(true),
Some(value) => Err(format!(
"MEMRA_STEP_NVFP4_DEV_ROUTES={value:?} is invalid; expected 0 or 1"
)),
}
}
pub fn step_nvfp4_dev_routes_enabled() -> Result<bool, String> {
parse_step_nvfp4_dev_routes(std::env::var("MEMRA_STEP_NVFP4_DEV_ROUTES").ok().as_deref())
}
pub fn step_ep_device_arithmetic_enabled() -> Result<bool, String> {
parse_step_ep_device_arithmetic(
std::env::var("MEMRA_STEP_EP_DEVICE_ARITHMETIC")
.ok()
.as_deref(),
)
}
fn parse_step_tp_f32_mirror(value: Option<&str>) -> Result<bool, String> {
match value {
None | Some("") | Some("0") => Ok(false),
Some("1") => Ok(true),
Some(value) => Err(format!(
"MEMRA_STEP_TP_F32_MIRROR={value:?} is invalid; expected 0 or 1"
)),
}
}
pub fn step_tp_f32_mirror_enabled() -> Result<bool, String> {
parse_step_tp_f32_mirror(std::env::var("MEMRA_STEP_TP_F32_MIRROR").ok().as_deref())
}
fn parse_step_tp_decode_v2(value: Option<&str>) -> Result<bool, String> {
match value {
None | Some("") | Some("0") => Ok(false),
Some("1") => Ok(true),
Some(value) => Err(format!(
"MEMRA_STEP_TP_DECODE_V2={value:?} is invalid; expected 0 or 1"
)),
}
}
pub fn step_tp_decode_v2_enabled() -> Result<bool, String> {
parse_step_tp_decode_v2(std::env::var("MEMRA_STEP_TP_DECODE_V2").ok().as_deref())
}
fn parse_step_tp_qkv_fused(value: Option<&str>) -> Result<bool, String> {
match value {
None | Some("") | Some("0") => Ok(false),
Some("1") => Ok(true),
Some(value) => Err(format!(
"MEMRA_STEP_TP_QKV_FUSED={value:?} is invalid; expected 0 or 1"
)),
}
}
fn parse_step_tp_dev_router(value: Option<&str>) -> Result<bool, String> {
match value {
None | Some("") | Some("0") => Ok(false),
Some("1") => Ok(true),
Some(value) => Err(format!(
"MEMRA_STEP_TP_DEV_ROUTER={value:?} is invalid; expected 0 or 1"
)),
}
}
pub fn step_tp_dev_router_enabled() -> Result<bool, String> {
parse_step_tp_dev_router(std::env::var("MEMRA_STEP_TP_DEV_ROUTER").ok().as_deref())
}
fn parse_step_tp_graph(value: Option<&str>) -> Result<bool, String> {
match value {
None | Some("") | Some("0") => Ok(false),
Some("1") => Ok(true),
Some(value) => Err(format!(
"MEMRA_STEP_TP_GRAPH={value:?} is invalid; expected 0 or 1"
)),
}
}
fn parse_step_tp_dcw(value: Option<&str>) -> Result<bool, String> {
match value {
None | Some("") | Some("0") => Ok(false),
Some("1") => Ok(true),
Some(value) => Err(format!(
"MEMRA_STEP_TP_DCW={value:?} is invalid; expected 0 or 1"
)),
}
}
pub fn step_tp_dcw_enabled() -> Result<bool, String> {
parse_step_tp_dcw(std::env::var("MEMRA_STEP_TP_DCW").ok().as_deref())
}
pub fn step_tp_graph_enabled() -> Result<bool, String> {
parse_step_tp_graph(std::env::var("MEMRA_STEP_TP_GRAPH").ok().as_deref())
}
pub fn step_tp_qkv_fused_enabled() -> Result<bool, String> {
parse_step_tp_qkv_fused(std::env::var("MEMRA_STEP_TP_QKV_FUSED").ok().as_deref())
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct StepEpLayerSpec {
pub layer: usize,
pub devices: Vec<usize>,
}
pub type StepTpLayerSpec = StepEpLayerSpec;
fn parse_step_layer_specs(
flag: &str,
value: Option<&str>,
allow_full_model: bool,
) -> Result<Vec<StepEpLayerSpec>, String> {
let Some(value) = value else {
return Ok(Vec::new());
};
if value.is_empty() || value == "0" {
return Ok(Vec::new());
}
let mut specs = Vec::new();
for item in value.split(';') {
let (layers, devices) = item.split_once('@').ok_or_else(|| {
let layers = if allow_full_model {
"LAYER[-LAYER] or all"
} else {
"LAYER[-LAYER]"
};
format!("{flag} must be {layers}@DEVICE,DEVICE[;...]")
})?;
let (first, last) = if layers == "all" {
if !allow_full_model {
return Err(format!(
"{flag} does not support the full-model shorthand; assign routed layers \
explicitly"
));
}
(0, STEP37_TRUNK_LAYERS - 1)
} else {
match layers.split_once('-') {
Some((first, last)) => {
let first = first
.parse::<usize>()
.map_err(|_| format!("{flag} layer {first:?} is not an integer"))?;
let last = last
.parse::<usize>()
.map_err(|_| format!("{flag} layer {last:?} is not an integer"))?;
if first > last {
return Err(format!("{flag} layer range {first}-{last} is reversed"));
}
if last - first + 1 > 128 {
return Err(format!(
"{flag} layer range {first}-{last} exceeds the 128-layer parser cap"
));
}
(first, last)
}
None => {
let layer = layers
.parse::<usize>()
.map_err(|_| format!("{flag} layer {layers:?} is not an integer"))?;
(layer, layer)
}
}
};
let devices = devices
.split(',')
.map(|device| {
device
.parse::<usize>()
.map_err(|_| format!("{flag} device {device:?} is not an integer"))
})
.collect::<Result<Vec<_>, _>>()?;
if !(2..=8).contains(&devices.len()) {
return Err(format!(
"{flag} requires 2..=8 devices, got {}",
devices.len()
));
}
let mut unique = devices.clone();
unique.sort_unstable();
unique.dedup();
if unique.len() != devices.len() {
return Err(format!("{flag} devices must be distinct, got {devices:?}"));
}
for layer in first..=last {
if specs
.iter()
.any(|existing: &StepEpLayerSpec| existing.layer == layer)
{
return Err(format!("{flag} assigns layer {layer} more than once"));
}
specs.push(StepEpLayerSpec {
layer,
devices: devices.clone(),
});
}
}
Ok(specs)
}
pub fn parse_step_ep_layer_specs(value: Option<&str>) -> Result<Vec<StepEpLayerSpec>, String> {
parse_step_layer_specs("MEMRA_STEP_EP", value, false)
}
pub fn step_ep_layer_specs() -> Result<Vec<StepEpLayerSpec>, String> {
parse_step_ep_layer_specs(std::env::var("MEMRA_STEP_EP").ok().as_deref())
}
pub fn parse_step_tp_layer_specs(value: Option<&str>) -> Result<Vec<StepTpLayerSpec>, String> {
parse_step_layer_specs("MEMRA_STEP_TP", value, true)
}
pub fn step_tp_layer_specs() -> Result<Vec<StepTpLayerSpec>, String> {
parse_step_tp_layer_specs(std::env::var("MEMRA_STEP_TP").ok().as_deref())
}
#[derive(Clone, Copy)]
pub struct E4m3BlockMatrix<'a> {
pub codes: &'a [u8],
pub scales: &'a [f32],
pub out_features: usize,
pub in_features: usize,
}
impl E4m3BlockMatrix<'_> {
fn validate(&self) -> Result<(), String> {
let code_count = self
.out_features
.checked_mul(self.in_features)
.ok_or_else(|| "E4M3 matrix size overflow".to_string())?;
if self.codes.len() != code_count {
return Err(format!(
"E4M3 code count {} != {}x{} ({code_count})",
self.codes.len(),
self.out_features,
self.in_features,
));
}
let scale_count =
self.out_features.div_ceil(FP8_BLOCK) * self.in_features.div_ceil(FP8_BLOCK);
if self.scales.len() != scale_count {
return Err(format!(
"E4M3 scale count {} != {scale_count} for {}x{}",
self.scales.len(),
self.out_features,
self.in_features,
));
}
if !self
.scales
.iter()
.all(|scale| scale.is_finite() && *scale > 0.0)
{
return Err("E4M3 scale grid contains a non-finite or non-positive value".to_string());
}
Ok(())
}
}
#[derive(Clone, Copy)]
pub struct E4m3ExpertBank<'a> {
pub codes: &'a [u8],
pub scales: &'a [f32],
pub expert_count: usize,
pub out_features: usize,
pub in_features: usize,
}
impl E4m3ExpertBank<'_> {
fn validate(&self) -> Result<(), String> {
if self.expert_count == 0 {
return Err("E4M3 expert bank is empty".to_string());
}
let code_stride = self
.out_features
.checked_mul(self.in_features)
.ok_or_else(|| "E4M3 expert code stride overflow".to_string())?;
let code_count = self
.expert_count
.checked_mul(code_stride)
.ok_or_else(|| "E4M3 expert code count overflow".to_string())?;
if self.codes.len() != code_count {
return Err(format!(
"E4M3 expert code count {} != {}x{} ({code_count})",
self.codes.len(),
self.expert_count,
code_stride,
));
}
let scale_stride =
self.out_features.div_ceil(FP8_BLOCK) * self.in_features.div_ceil(FP8_BLOCK);
let scale_count = self
.expert_count
.checked_mul(scale_stride)
.ok_or_else(|| "E4M3 expert scale count overflow".to_string())?;
if self.scales.len() != scale_count {
return Err(format!(
"E4M3 expert scale count {} != {}x{} ({scale_count})",
self.scales.len(),
self.expert_count,
scale_stride,
));
}
if !self
.scales
.iter()
.all(|scale| scale.is_finite() && *scale > 0.0)
{
return Err(
"E4M3 expert scale grid contains a non-finite or non-positive value".to_string(),
);
}
Ok(())
}
pub fn expert(&self, expert: usize) -> Result<E4m3BlockMatrix<'_>, String> {
if expert >= self.expert_count {
return Err(format!("expert {expert} outside 0..{}", self.expert_count));
}
let code_stride = self.out_features * self.in_features;
let scale_stride =
self.out_features.div_ceil(FP8_BLOCK) * self.in_features.div_ceil(FP8_BLOCK);
Ok(E4m3BlockMatrix {
codes: &self.codes[expert * code_stride..(expert + 1) * code_stride],
scales: &self.scales[expert * scale_stride..(expert + 1) * scale_stride],
out_features: self.out_features,
in_features: self.in_features,
})
}
}
pub struct ColumnParallelResult {
pub gathered: Vec<f32>,
pub rank_outputs: Vec<Vec<f32>>,
}
pub struct RowParallelResult {
pub reduced: Vec<f32>,
pub rank_partials: Vec<Vec<f32>>,
}
#[derive(Clone, Copy)]
pub struct Bf16Matrix<'a> {
pub bytes: &'a [u8],
pub out_features: usize,
pub in_features: usize,
}
impl Bf16Matrix<'_> {
pub fn validate(&self) -> Result<(), String> {
if self.out_features == 0 || self.in_features == 0 {
return Err("BF16 matrix dimensions must be nonzero".into());
}
let expected = self
.out_features
.checked_mul(self.in_features)
.and_then(|values| values.checked_mul(2))
.ok_or("BF16 matrix byte count overflow")?;
if self.bytes.len() != expected {
return Err(format!(
"BF16 matrix bytes {} != {}x{}x2 ({expected})",
self.bytes.len(),
self.out_features,
self.in_features,
));
}
Ok(())
}
}
struct ResidentE4m3Rank {
codes: CudaSlice<u8>,
scales: CudaSlice<f32>,
out_features: usize,
in_features: usize,
}
enum ResidentBf16Weight {
Bf16(CudaSlice<u8>),
F32(CudaSlice<f32>),
}
impl ResidentBf16Weight {
fn ordinal(&self) -> usize {
match self {
Self::Bf16(bytes) => bytes.ordinal(),
Self::F32(values) => values.ordinal(),
}
}
}
struct ResidentBf16Rank {
weight: ResidentBf16Weight,
out_features: usize,
in_features: usize,
}
pub struct ResidentColumnParallel {
ranks: Vec<ResidentE4m3Rank>,
out_features: usize,
in_features: usize,
}
pub struct ResidentRowParallel {
ranks: Vec<ResidentE4m3Rank>,
out_features: usize,
in_features: usize,
}
pub struct ResidentBf16ColumnParallel {
ranks: Vec<ResidentBf16Rank>,
out_features: usize,
in_features: usize,
canonical_chunk_rows: Option<usize>,
}
pub struct ResidentBf16RowParallel {
ranks: Vec<ResidentBf16Rank>,
out_features: usize,
in_features: usize,
}
pub struct ResidentStepBf16RowParallel {
ranks: Vec<Vec<ResidentBf16Rank>>,
out_features: usize,
in_features: usize,
canonical_chunk_cols: usize,
}
pub struct ResidentSigmoidTopKRouter {
weight: CudaSlice<f32>,
correction_bias: CudaSlice<f32>,
active: CudaSlice<u8>,
root_device: usize,
input_width: usize,
expert_count: usize,
experts_per_token: usize,
active_count: usize,
scaling_factor: f32,
route_norm: bool,
}
pub struct SigmoidTopKHostOutput {
pub logits: Vec<f32>,
pub selected: Vec<u32>,
pub weights: Vec<f32>,
}
pub struct ResidentReplicatedBf16SwiGlu {
gate: Vec<ResidentBf16Rank>,
up: Vec<ResidentBf16Rank>,
down: Vec<ResidentBf16Rank>,
input_width: usize,
intermediate_width: usize,
}
pub struct ResidentReplicatedDeviceRows {
ranks: Vec<CudaSlice<f32>>,
tokens: usize,
width: usize,
}
impl ResidentReplicatedDeviceRows {
pub fn tokens(&self) -> usize {
self.tokens
}
pub fn width(&self) -> usize {
self.width
}
pub fn ranks(&self) -> usize {
self.ranks.len()
}
}
pub fn moe_residual_host(
residual: &[f32],
routed: &[f32],
shared: &[f32],
) -> Result<Vec<f32>, String> {
if residual.len() != routed.len() || residual.len() != shared.len() {
return Err(format!(
"MoE residual lengths residual={} routed={} shared={}",
residual.len(),
routed.len(),
shared.len()
));
}
let ffn = routed
.iter()
.zip(shared)
.map(|(&routed, &shared)| routed + shared)
.collect::<Vec<_>>();
Ok(residual
.iter()
.zip(ffn)
.map(|(&residual, ffn)| residual + ffn)
.collect())
}
pub use memra_kv::{
KvRingAppend, ResidentTpKvCache, ResidentTpKvCacheRank, TpKvAppendPlan, TpKvTransaction,
};
pub struct ResidentTpExpert {
gate: ResidentColumnParallel,
up: ResidentColumnParallel,
down: ResidentRowParallel,
input_width: usize,
expert_width: usize,
}
struct ResidentE4m3ExpertBankRank {
codes: CudaSlice<u8>,
scales: CudaSlice<f32>,
expert_range: Range<usize>,
out_features: usize,
in_features: usize,
code_stride: usize,
scale_stride: usize,
k_blocks: Option<usize>,
}
struct PackedE4m3ExpertBankRank {
codes: Vec<u8>,
scales: Vec<f32>,
expert_range: Range<usize>,
out_features: usize,
in_features: usize,
code_stride: usize,
scale_stride: usize,
k_blocks: Option<usize>,
}
struct ResidentEpRank {
gate: ResidentE4m3ExpertBankRank,
up: ResidentE4m3ExpertBankRank,
down: ResidentE4m3ExpertBankRank,
}
pub struct ResidentExpertParallel {
ranks: Vec<ResidentEpRank>,
expert_count: usize,
input_width: usize,
expert_width: usize,
}
pub struct StepGroupedFp8ProjectionOutput {
pub gate: Vec<f32>,
pub up: Vec<f32>,
pub down: Vec<f32>,
}
pub struct PreparedStepGroupedFp8Gate {
device: usize,
gate: ResidentE4m3ExpertBankRank,
up: ResidentE4m3ExpertBankRank,
down: ResidentE4m3ExpertBankRank,
input: CudaSlice<f32>,
route_csr: DeviceExpertCsr,
down_csr: DeviceExpertCsr,
gate_workspace: Fp8GroupedWorkspace,
up_workspace: Fp8GroupedWorkspace,
down_workspace: Fp8GroupedWorkspace,
activation: CudaSlice<f32>,
activation_limit: Option<f32>,
tokens: usize,
pairs: usize,
}
impl PreparedStepGroupedFp8Gate {
pub fn tokens(&self) -> usize {
self.tokens
}
pub fn pairs(&self) -> usize {
self.pairs
}
}
struct PreparedStepGroupedExpertOwner {
rank: usize,
global_pairs: Vec<usize>,
route_csr: DeviceExpertCsr,
down_csr: DeviceExpertCsr,
gate_workspace: Fp8GroupedWorkspace,
up_workspace: Fp8GroupedWorkspace,
down_workspace: Fp8GroupedWorkspace,
activation: CudaSlice<f32>,
}
struct StepGroupedExpertOwnerSchedule {
global_pairs: Vec<usize>,
route_csr: ExpertCsr,
down_csr: ExpertCsr,
}
pub struct PreparedStepGroupedExpertParallelGate {
rank_inputs: Vec<CudaSlice<f32>>,
owners: Vec<PreparedStepGroupedExpertOwner>,
activation_limit: Option<f32>,
tokens: usize,
pairs: usize,
max_tokens: usize,
max_pairs: usize,
input_width: usize,
expert_width: usize,
generation: u64,
executed_generation: Option<u64>,
ready: bool,
}
impl PreparedStepGroupedExpertParallelGate {
pub fn tokens(&self) -> usize {
self.tokens
}
pub fn pairs(&self) -> usize {
self.pairs
}
pub fn max_tokens(&self) -> usize {
self.max_tokens
}
pub fn input_width(&self) -> usize {
self.input_width
}
pub fn expert_width(&self) -> usize {
self.expert_width
}
pub fn set_activation_limit(&mut self, limit: Option<f32>) -> Result<(), String> {
validate_step_expert_activation_limit(limit)?;
self.activation_limit = limit;
self.executed_generation = None;
Ok(())
}
pub fn active_owners(&self) -> usize {
self.owners
.iter()
.filter(|owner| !owner.global_pairs.is_empty())
.count()
}
pub fn owner_pair_counts(&self) -> Vec<usize> {
self.owners
.iter()
.map(|owner| owner.global_pairs.len())
.collect()
}
pub fn generation(&self) -> u64 {
self.generation
}
}
struct PreparedPeerWeightedRouteOwner {
token_rows: CudaSlice<i32>,
slots: CudaSlice<i32>,
weights: CudaSlice<f32>,
active_pairs: usize,
}
pub struct PreparedPeerWeightedRouteCombine {
root_device: usize,
owners: Vec<PreparedPeerWeightedRouteOwner>,
peer_staging: CudaSlice<f32>,
slots: CudaSlice<f32>,
weights: CudaSlice<f32>,
output: CudaSlice<f32>,
peer_devices: Vec<usize>,
peer_outputs: Vec<CudaSlice<f32>>,
width: usize,
experts_per_token: usize,
max_tokens: usize,
max_pairs: usize,
tokens: usize,
pairs: usize,
projection_generation: u64,
output_generation: Option<u64>,
broadcast_generation: Option<u64>,
ready: bool,
}
impl PreparedPeerWeightedRouteCombine {
pub fn tokens(&self) -> usize {
self.tokens
}
pub fn pairs(&self) -> usize {
self.pairs
}
pub fn owner_pair_counts(&self) -> Vec<usize> {
self.owners.iter().map(|owner| owner.active_pairs).collect()
}
pub fn distributed_ranks(&self) -> usize {
1 + self.peer_outputs.len()
}
}
struct ResidentTpExpertBank {
gate: Vec<ResidentE4m3ExpertBankRank>,
up: Vec<ResidentE4m3ExpertBankRank>,
down: Vec<ResidentE4m3ExpertBankRank>,
expert_count: usize,
input_width: usize,
expert_width: usize,
}
pub struct ResidentTensorParallel {
bank: ResidentTpExpertBank,
}
pub struct TpE4m3HostBounce {
devices: Vec<usize>,
ranks: Vec<Engine>,
native_p2p: bool,
ep_device_arithmetic: bool,
bulk_p2p: bool,
decode_v2: std::sync::Mutex<Vec<StepTpDecodeV2Ws>>,
}
pub enum StepTpGateShards<'a> {
F32(&'a [crate::CudaSlice<f32>]),
Bf16(&'a [crate::CudaSlice<u8>]),
}
pub struct StepTpDecodeV2Ws {
pub(crate) tcol_q: Vec<CudaSlice<f32>>,
pub(crate) tcol_k: Vec<CudaSlice<f32>>,
pub(crate) tcol_v: Vec<CudaSlice<f32>>,
pub(crate) tcol_g: Vec<CudaSlice<f32>>,
pub(crate) tcol_in: Vec<CudaSlice<f32>>,
pub(crate) tcol_cap: usize,
tcol_gated: Vec<CudaSlice<f32>>,
tcol_opart: Vec<CudaSlice<f32>>,
tcol_opeer: Option<CudaSlice<f32>>,
tcol_omix: Option<CudaSlice<f32>>,
tcol_ocap: usize,
pub(crate) q_raw: Vec<CudaSlice<f32>>,
pub(crate) k_raw: Vec<CudaSlice<f32>>,
pub(crate) v_raw: Vec<CudaSlice<f32>>,
pub(crate) q: Vec<CudaSlice<f32>>,
pub(crate) k: Vec<CudaSlice<f32>>,
pub(crate) pos: Vec<CudaSlice<i32>>,
pub(crate) fuse_ctr: Vec<CudaSlice<u32>>,
pub(crate) gate: Vec<CudaSlice<f32>>,
pub(crate) attn_out: Vec<CudaSlice<f32>>,
pub(crate) gated: Vec<CudaSlice<f32>>,
o_partials: Vec<Vec<CudaSlice<f32>>>,
ev_rank: Vec<CudaEvent>,
peer_partial: CudaSlice<f32>,
reduce_a: CudaSlice<f32>,
reduce_b: CudaSlice<f32>,
zeros: CudaSlice<f32>,
pub(crate) k_shadow: CudaSlice<f32>,
pub(crate) v_shadow: CudaSlice<f32>,
ev_refresh: CudaEvent,
ev_oproj: CudaEvent,
gate_e: CudaSlice<f32>,
pub(crate) h_stage: Option<CudaSlice<f32>>,
pub(crate) pos_stage: Option<CudaSlice<i32>>,
attn_in: Vec<CudaSlice<f32>>,
raw_h_stage: u64,
raw_pos_stage: u64,
raw_attn_in: Vec<u64>,
raw_pos: Vec<u64>,
raw_o_partial1: u64,
raw_peer_partial: u64,
raw_k1: u64,
raw_v1: u64,
raw_k_shadow: u64,
raw_v_shadow: u64,
raw_mixed_stage_e: u64,
raw_reduce_a: u64,
raw_shadow_stage_e: (u64, u64),
ev_entry: CudaEvent,
e_device: usize,
local_q_dim: usize,
local_kv_dim: usize,
heads: usize,
pub(crate) o_out: usize,
o_block_cols: usize,
blocks_per_rank: usize,
}
impl TpE4m3HostBounce {
pub fn new(devices: &[usize]) -> Result<Self, Box<dyn std::error::Error>> {
Self::new_inner(devices, false, false, false, false)
}
pub fn new_native_p2p(devices: &[usize]) -> Result<Self, Box<dyn std::error::Error>> {
Self::new_inner(devices, false, true, false, false)
}
pub fn new_native_p2p_device_arithmetic(
devices: &[usize],
) -> Result<Self, Box<dyn std::error::Error>> {
Self::new_inner(devices, false, true, true, false)
}
pub(crate) fn new_configured(
devices: &[usize],
native_p2p: bool,
ep_device_arithmetic: bool,
bulk_p2p: bool,
) -> Result<Self, Box<dyn std::error::Error>> {
Self::new_inner(devices, false, native_p2p, ep_device_arithmetic, bulk_p2p)
}
pub fn new_single_rank_oracle(device: usize) -> Result<Self, Box<dyn std::error::Error>> {
Self::new_inner(&[device], true, false, false, false)
}
fn new_inner(
devices: &[usize],
allow_single_rank: bool,
native_p2p: bool,
ep_device_arithmetic: bool,
bulk_p2p: bool,
) -> Result<Self, Box<dyn std::error::Error>> {
if ep_device_arithmetic && !native_p2p {
return Err("device-resident EP arithmetic requires native P2P".into());
}
if bulk_p2p && !native_p2p {
return Err("bulk TP transport requires native P2P".into());
}
let minimum = if allow_single_rank { 1 } else { 2 };
if !(minimum..=8).contains(&devices.len()) {
return Err(format!(
"TP reference requires {minimum}..=8 devices, got {}",
devices.len()
)
.into());
}
let mut unique = devices.to_vec();
unique.sort_unstable();
unique.dedup();
if unique.len() != devices.len() {
return Err(format!("TP devices must be distinct, got {devices:?}").into());
}
let ranks = devices
.iter()
.map(|&device| Engine::new(device))
.collect::<Result<Vec<_>, _>>()?;
if native_p2p {
configure_native_p2p(&ranks, devices)?;
}
if allow_single_rank {
eprintln!(
"[tp] canonical oracle transport=local device={} performance_claim=false",
devices[0]
);
} else if native_p2p {
if ep_device_arithmetic {
eprintln!(
"[tp] correctness transport=native-p2p devices={devices:?} \
native_p2p=true activation=device-host-exact \
accumulation=device-host-exact output=root-readback \
bulk_p2p={bulk_p2p} performance_claim=false"
);
} else {
eprintln!(
"[tp] correctness transport=native-p2p devices={devices:?} \
native_p2p=true activation=host-canonical bulk_p2p={bulk_p2p} \
performance_claim=false"
);
}
} else {
eprintln!(
"[tp] correctness transport=host-bounce devices={devices:?} \
native_p2p=false performance_claim=false"
);
}
Ok(Self {
devices: devices.to_vec(),
ranks,
native_p2p,
ep_device_arithmetic,
bulk_p2p,
decode_v2: std::sync::Mutex::new(Vec::new()),
})
}
pub fn devices(&self) -> &[usize] {
&self.devices
}
pub fn native_p2p(&self) -> bool {
self.native_p2p
}
pub fn bulk_p2p(&self) -> bool {
self.bulk_p2p
}
pub fn expert_activation_label(&self) -> &'static str {
if self.ep_device_arithmetic {
"device-host-exact"
} else {
"host-canonical"
}
}
pub fn expert_accumulation_label(&self) -> &'static str {
self.expert_activation_label()
}
pub fn expert_output_label(&self) -> &'static str {
if self.ep_device_arithmetic {
"root-readback"
} else {
"host-accumulated"
}
}
pub fn transport_label(&self) -> &'static str {
if self.devices.len() == 1 {
"local"
} else if self.native_p2p {
"native-p2p"
} else {
"host-bounce"
}
}
pub fn device_names(&self) -> Result<Vec<String>, Box<dyn std::error::Error>> {
self.ranks
.iter()
.map(|rank| rank.ctx().name().map_err(Into::into))
.collect()
}
pub fn rank_engine(&self, rank: usize) -> Option<&Engine> {
self.ranks.get(rank)
}
pub fn allocate_tp_kv_cache(
&self,
kv_dim_k: usize,
kv_dim_v: usize,
capacity: usize,
) -> Result<ResidentTpKvCache, Box<dyn std::error::Error>> {
self.allocate_tp_kv_cache_inner(kv_dim_k, kv_dim_v, capacity, None)
}
pub fn allocate_tp_swa_kv_cache(
&self,
kv_dim_k: usize,
kv_dim_v: usize,
capacity: usize,
window: usize,
) -> Result<ResidentTpKvCache, Box<dyn std::error::Error>> {
if window == 0 {
return Err("TP SWA KV window must be nonzero".into());
}
self.allocate_tp_kv_cache_inner(kv_dim_k, kv_dim_v, capacity, Some(window))
}
fn allocate_tp_kv_cache_inner(
&self,
kv_dim_k: usize,
kv_dim_v: usize,
capacity: usize,
window: Option<usize>,
) -> Result<ResidentTpKvCache, Box<dyn std::error::Error>> {
if capacity == 0 || capacity > i32::MAX as usize {
return Err(
format!("TP KV capacity must be in 1..={}, got {capacity}", i32::MAX).into(),
);
}
let tp = self.ranks.len();
let shape = crate::cache::tp_kv_rank_allocation_shape(kv_dim_k, kv_dim_v, tp)?;
let physical_rows = window
.map(|window| crate::cache::swa_ring_rows(window, capacity))
.unwrap_or(capacity);
let k_plane_bytes = physical_rows
.checked_mul(shape.k_token_bytes)
.and_then(|bytes| bytes.checked_add(8))
.ok_or("TP KV K plane-byte overflow")?;
let v_plane_bytes = physical_rows
.checked_mul(shape.v_token_bytes)
.and_then(|bytes| bytes.checked_add(8))
.ok_or("TP KV V plane-byte overflow")?;
let mut ranks = Vec::with_capacity(tp);
for engine in &self.ranks {
let _main = engine.gpu.enter_main()?;
ranks.push(ResidentTpKvCacheRank::new(
engine.alloc_u8(k_plane_bytes)?,
engine.alloc_u8(v_plane_bytes)?,
engine.htod_i32(&[0])?,
));
}
Ok(match window {
Some(window) => ResidentTpKvCache::new_swa(
ranks,
shape.kv_dim_k,
shape.kv_dim_v,
shape.k_token_bytes,
shape.v_token_bytes,
capacity,
window,
),
None => ResidentTpKvCache::new(
ranks,
shape.kv_dim_k,
shape.kv_dim_v,
shape.k_token_bytes,
shape.v_token_bytes,
capacity,
),
})
}
pub fn grow_tp_kv_cache(
&self,
source: &ResidentTpKvCache,
target_capacity: usize,
rows: usize,
) -> Result<ResidentTpKvCache, Box<dyn std::error::Error>> {
self.validate_tp_kv_cache(source)?;
let plan = source.prepare_grow(target_capacity, rows)?;
let ranks = self.ranks.len();
let global_k = source
.kv_dim_k()
.checked_mul(ranks)
.ok_or("TP KV grow global K dimension overflow")?;
let global_v = source
.kv_dim_v()
.checked_mul(ranks)
.ok_or("TP KV grow global V dimension overflow")?;
let mut target = match source.ring_window() {
Some(window) => {
self.allocate_tp_swa_kv_cache(global_k, global_v, target_capacity, window)?
}
None => self.allocate_tp_kv_cache(global_k, global_v, target_capacity)?,
};
self.validate_tp_kv_cache(&target)?;
for (rank, engine) in self.ranks.iter().enumerate() {
let _main = engine.gpu.enter_main()?;
let src = source
.rank(rank)
.ok_or_else(|| format!("TP KV grow source has no rank {rank}"))?;
let dst = target
.rank_mut(rank)
.ok_or_else(|| format!("TP KV grow target has no rank {rank}"))?;
if plan.k_bytes() > 0 {
engine.copy_u8_range_into(
dst.k_mut(),
0,
src.k(),
plan.source_row() * source.k_tok_bytes(),
plan.k_bytes(),
)?;
}
if plan.v_bytes() > 0 {
engine.copy_u8_range_into(
dst.v_mut(),
0,
src.v(),
plan.source_row() * source.v_tok_bytes(),
plan.v_bytes(),
)?;
}
}
self.set_tp_kv_len_mirrors(&mut target, plan.rows())?;
for engine in &self.ranks {
let _main = engine.gpu.enter_main()?;
engine.stream().synchronize()?;
}
let physical_copy_rows = plan.copy_rows();
target.publish_grow(plan)?;
eprintln!(
"[step-tp-kv-grow] rows={} source_capacity={} target_capacity={} ranks={} \
physical_copy_rows={} ring_window={:?} copy=rank-local-dtod \
rank_streams_synchronized=true generation_preserved=true",
rows,
source.capacity(),
target_capacity,
ranks,
physical_copy_rows,
source.ring_window(),
);
Ok(target)
}
pub fn hydrate_tp_kv_cache(
&self,
cache: &mut ResidentTpKvCache,
rows: usize,
k_rows: &[u8],
v_rows: &[u8],
) -> Result<(), Box<dyn std::error::Error>> {
self.hydrate_tp_kv_cache_from(cache, rows, 0, k_rows, v_rows)
}
pub fn hydrate_tp_kv_cache_from(
&self,
cache: &mut ResidentTpKvCache,
logical_len: usize,
resident_start: usize,
k_rows: &[u8],
v_rows: &[u8],
) -> Result<(), Box<dyn std::error::Error>> {
self.validate_tp_kv_cache(cache)?;
if cache.committed_len() != 0 || cache.staged_len() != 0 {
return Err(format!(
"TP KV hydration requires an empty cache, got committed/staged={}/{}",
cache.committed_len(),
cache.staged_len()
)
.into());
}
if resident_start > logical_len || logical_len > cache.capacity() {
return Err(format!(
"TP KV hydration range [{resident_start},{logical_len}) exceeds capacity {}",
cache.capacity(),
)
.into());
}
let rows = logical_len - resident_start;
if rows > cache.physical_capacity() {
return Err(format!(
"TP KV hydration rows {rows} exceed physical capacity {}",
cache.physical_capacity()
)
.into());
}
for rank in 0..self.ranks.len() {
let k_rank =
cache_rank_rows(k_rows, rows, cache.k_tok_bytes(), self.ranks.len(), rank)?;
let v_rank =
cache_rank_rows(v_rows, rows, cache.v_tok_bytes(), self.ranks.len(), rank)?;
let engine = &self.ranks[rank];
let _main = engine.gpu.enter_main()?;
let rank_cache = cache
.rank_mut(rank)
.ok_or_else(|| format!("TP KV cache has no rank {rank}"))?;
engine.htod_u8_into(rank_cache.k_mut(), 0, &k_rank)?;
engine.htod_u8_into(rank_cache.v_mut(), 0, &v_rank)?;
}
cache.publish_hydration(logical_len, resident_start)?;
Ok(())
}
pub fn append_tp_kv_transaction(
&self,
cache: &mut ResidentTpKvCache,
transaction: TpKvTransaction,
k_shards: &[CudaSlice<f32>],
v_shards: &[CudaSlice<f32>],
rows: usize,
) -> Result<(), Box<dyn std::error::Error>> {
self.append_tp_kv_transaction_inner(cache, transaction, k_shards, v_shards, rows, false)
}
#[allow(clippy::too_many_arguments)]
pub fn append_tp_kv_transaction_inner(
&self,
cache: &mut ResidentTpKvCache,
transaction: TpKvTransaction,
k_shards: &[CudaSlice<f32>],
v_shards: &[CudaSlice<f32>],
rows: usize,
external_rank_appends: bool,
) -> Result<(), Box<dyn std::error::Error>> {
self.validate_tp_kv_cache(cache)?;
let plan = cache.prepare_append(transaction, rows)?;
let target = plan.target();
let expected_k = rows
.checked_mul(cache.kv_dim_k())
.ok_or("TP KV K append size overflow")?;
let expected_v = rows
.checked_mul(cache.kv_dim_v())
.ok_or("TP KV V append size overflow")?;
if !external_rank_appends
&& (k_shards.len() != self.ranks.len() || v_shards.len() != self.ranks.len())
{
return Err(format!(
"TP KV append shard counts k={} v={} != ranks {}",
k_shards.len(),
v_shards.len(),
self.ranks.len()
)
.into());
}
let kv_dim_k = cache.kv_dim_k();
let kv_dim_v = cache.kv_dim_v();
let k_tok_bytes = cache.k_tok_bytes();
let v_tok_bytes = cache.v_tok_bytes();
if let Some(KvRingAppend::Rebase {
src_row,
keep_rows,
new_base,
..
}) = plan.ring_append()
{
for rank in 0..self.ranks.len() {
let engine = &self.ranks[rank];
let _main = engine.gpu.enter_main()?;
let rank_cache = cache
.rank_mut(rank)
.ok_or_else(|| format!("TP KV cache has no rank {rank}"))?;
if keep_rows > 0 {
let k_len = keep_rows
.checked_mul(k_tok_bytes)
.ok_or("TP KV K rebase-byte overflow")?;
let v_len = keep_rows
.checked_mul(v_tok_bytes)
.ok_or("TP KV V rebase-byte overflow")?;
let mut k_tmp = engine.alloc_u8_uninit(k_len)?;
let mut v_tmp = engine.alloc_u8_uninit(v_len)?;
engine.copy_u8_range_into(
&mut k_tmp,
0,
rank_cache.k(),
src_row * k_tok_bytes,
k_len,
)?;
engine.copy_u8_range_into(
&mut v_tmp,
0,
rank_cache.v(),
src_row * v_tok_bytes,
v_len,
)?;
engine.copy_u8_into(rank_cache.k_mut(), 0, &k_tmp, k_len)?;
engine.copy_u8_into(rank_cache.v_mut(), 0, &v_tmp, v_len)?;
}
if rank_cache.base_d().is_some() {
let value = new_base as i32;
let rank_cache = cache
.rank_mut(rank)
.ok_or_else(|| format!("TP KV cache has no rank {rank}"))?;
if let Some(base_d) = rank_cache.base_d_mut() {
engine.set_i32_one(base_d, value)?;
}
}
}
}
cache.publish_append_rebase(plan)?;
let write_row = plan.write_row();
for rank in 0..self.ranks.len() {
if external_rank_appends {
break;
}
let engine = &self.ranks[rank];
let _main = engine.gpu.enter_main()?;
if k_shards[rank].len() != expected_k
|| v_shards[rank].len() != expected_v
|| k_shards[rank].ordinal() != engine.ctx().ordinal()
|| v_shards[rank].ordinal() != engine.ctx().ordinal()
{
return Err(format!(
"TP KV rank {rank} shard geometry/device k={}/{} v={}/{} \
!= expected {expected_k}/{expected_v} on device {}",
k_shards[rank].len(),
k_shards[rank].ordinal(),
v_shards[rank].len(),
v_shards[rank].ordinal(),
engine.ctx().ordinal(),
)
.into());
}
let rank_cache = cache
.rank_mut(rank)
.ok_or_else(|| format!("TP KV cache has no rank {rank}"))?;
let (rank_k, rank_v) = rank_cache.planes_mut();
engine.append_kv_quantized_rows(
&k_shards[rank],
&v_shards[rank],
rank_k,
rank_v,
write_row,
rows,
kv_dim_k,
kv_dim_v,
k_tok_bytes,
v_tok_bytes,
Engine::kv_fp8_on(),
)?;
}
if !external_rank_appends {
self.set_tp_kv_len_mirrors(cache, target)?;
}
cache.publish_append_plan(plan)?;
Ok(())
}
pub fn commit_tp_kv_transaction(
&self,
cache: &mut ResidentTpKvCache,
transaction: TpKvTransaction,
accepted_rows: usize,
) -> Result<(), Box<dyn std::error::Error>> {
self.validate_tp_kv_cache(cache)?;
let target = cache.commit_target(transaction, accepted_rows)?;
self.set_tp_kv_len_mirrors(cache, target)?;
cache.publish_finalize(transaction, target)?;
Ok(())
}
pub fn commit_tp_kv_transaction_external(
&self,
cache: &mut ResidentTpKvCache,
transaction: TpKvTransaction,
accepted_rows: usize,
) -> Result<(), Box<dyn std::error::Error>> {
self.validate_tp_kv_cache(cache)?;
let target = cache.commit_target(transaction, accepted_rows)?;
cache.publish_finalize(transaction, target)?;
Ok(())
}
pub fn rollback_tp_kv_transaction(
&self,
cache: &mut ResidentTpKvCache,
transaction: TpKvTransaction,
) -> Result<(), Box<dyn std::error::Error>> {
self.validate_tp_kv_cache(cache)?;
cache.validate_transaction(transaction)?;
let target = transaction.base_len();
self.set_tp_kv_len_mirrors(cache, target)?;
cache.publish_finalize(transaction, target)?;
Ok(())
}
pub fn tp_kv_device_lengths(
&self,
cache: &ResidentTpKvCache,
) -> Result<Vec<i32>, Box<dyn std::error::Error>> {
self.validate_tp_kv_cache(cache)?;
let mut lengths = Vec::with_capacity(self.ranks.len());
for (engine, rank_cache) in self.ranks.iter().zip(cache.ranks()) {
let _main = engine.gpu.enter_main()?;
lengths.push(engine.dtoh_i32_one(rank_cache.len_d())?);
}
Ok(lengths)
}
fn set_tp_kv_len_mirrors(
&self,
cache: &mut ResidentTpKvCache,
len: usize,
) -> Result<(), Box<dyn std::error::Error>> {
let len = i32::try_from(len).map_err(|_| "TP KV length exceeds i32 device mirror")?;
for (engine, rank_cache) in self.ranks.iter().zip(cache.ranks_mut()) {
let _main = engine.gpu.enter_main()?;
engine.set_i32_one(rank_cache.len_d_mut(), len)?;
}
Ok(())
}
fn validate_tp_kv_cache(
&self,
cache: &ResidentTpKvCache,
) -> Result<(), Box<dyn std::error::Error>> {
if cache.ranks_len() != self.ranks.len() {
return Err(format!(
"TP KV cache ranks {} != runtime ranks {}",
cache.ranks_len(),
self.ranks.len()
)
.into());
}
let expected_k = cache
.physical_capacity()
.checked_mul(cache.k_tok_bytes())
.and_then(|bytes| bytes.checked_add(8))
.ok_or("TP KV K plane validation overflow")?;
let expected_v = cache
.physical_capacity()
.checked_mul(cache.v_tok_bytes())
.and_then(|bytes| bytes.checked_add(8))
.ok_or("TP KV V plane validation overflow")?;
for (rank, (engine, rank_cache)) in self.ranks.iter().zip(cache.ranks()).enumerate() {
let device = engine.ctx().ordinal();
if rank_cache.k().len() != expected_k
|| rank_cache.v().len() != expected_v
|| rank_cache.len_d().len() != 1
|| rank_cache.k().ordinal() != device
|| rank_cache.v().ordinal() != device
|| rank_cache.len_d().ordinal() != device
{
return Err(format!(
"TP KV rank {rank} residency does not match device {device} or plane geometry"
)
.into());
}
}
Ok(())
}
pub fn full(
&self,
matrix: E4m3BlockMatrix<'_>,
activations: &[f32],
tokens: usize,
) -> Result<Vec<f32>, Box<dyn std::error::Error>> {
matrix.validate()?;
validate_activations(activations, tokens, matrix.in_features)?;
run_rank(&self.ranks[0], matrix, activations, tokens)
}
pub fn column_parallel(
&self,
matrix: E4m3BlockMatrix<'_>,
activations: &[f32],
tokens: usize,
) -> Result<ColumnParallelResult, Box<dyn std::error::Error>> {
matrix.validate()?;
validate_activations(activations, tokens, matrix.in_features)?;
let tp = self.ranks.len();
if matrix.out_features % tp != 0 {
return Err(format!(
"column-parallel out_features {} is not divisible by TP={tp}",
matrix.out_features
)
.into());
}
let local_out = matrix.out_features / tp;
if local_out % FP8_BLOCK != 0 {
return Err(format!(
"column-parallel output shard {local_out} cuts through a {FP8_BLOCK}-row \
E4M3 scale block"
)
.into());
}
let mut gathered = vec![0.0f32; tokens * matrix.out_features];
let mut rank_outputs = Vec::with_capacity(tp);
for (rank_index, rank) in self.ranks.iter().enumerate() {
let shard = column_shard(matrix, tp, rank_index)?;
let output = run_rank(rank, shard, activations, tokens)?;
let row_start = rank_index * local_out;
for token in 0..tokens {
gathered[token * matrix.out_features + row_start
..token * matrix.out_features + row_start + local_out]
.copy_from_slice(&output[token * local_out..(token + 1) * local_out]);
}
rank_outputs.push(output);
}
Ok(ColumnParallelResult {
gathered,
rank_outputs,
})
}
pub fn upload_column_parallel(
&self,
matrix: E4m3BlockMatrix<'_>,
) -> Result<ResidentColumnParallel, Box<dyn std::error::Error>> {
matrix.validate()?;
let tp = self.ranks.len();
validate_column_shape(matrix, tp)?;
let mut ranks = Vec::with_capacity(tp);
for (rank_index, engine) in self.ranks.iter().enumerate() {
ranks.push(upload_rank(engine, column_shard(matrix, tp, rank_index)?)?);
}
Ok(ResidentColumnParallel {
ranks,
out_features: matrix.out_features,
in_features: matrix.in_features,
})
}
pub fn column_parallel_resident(
&self,
matrix: &ResidentColumnParallel,
activations: &[f32],
tokens: usize,
) -> Result<ColumnParallelResult, Box<dyn std::error::Error>> {
validate_resident_ranks(&self.ranks, &matrix.ranks)?;
validate_activations(activations, tokens, matrix.in_features)?;
let local_out = matrix.out_features / self.ranks.len();
let mut gathered = vec![0.0f32; tokens * matrix.out_features];
let mut rank_outputs = Vec::with_capacity(self.ranks.len());
for (rank_index, (engine, shard)) in self.ranks.iter().zip(&matrix.ranks).enumerate() {
let output = run_resident_rank(engine, shard, activations, tokens)?;
let row_start = rank_index * local_out;
for token in 0..tokens {
gathered[token * matrix.out_features + row_start
..token * matrix.out_features + row_start + local_out]
.copy_from_slice(&output[token * local_out..(token + 1) * local_out]);
}
rank_outputs.push(output);
}
Ok(ColumnParallelResult {
gathered,
rank_outputs,
})
}
pub fn row_parallel(
&self,
matrix: E4m3BlockMatrix<'_>,
activations: &[f32],
tokens: usize,
) -> Result<RowParallelResult, Box<dyn std::error::Error>> {
matrix.validate()?;
validate_activations(activations, tokens, matrix.in_features)?;
let tp = self.ranks.len();
if matrix.in_features % tp != 0 {
return Err(format!(
"row-parallel in_features {} is not divisible by TP={tp}",
matrix.in_features
)
.into());
}
let local_in = matrix.in_features / tp;
if local_in % FP8_BLOCK != 0 {
return Err(format!(
"row-parallel input shard {local_in} cuts through a {FP8_BLOCK}-column \
E4M3 scale block"
)
.into());
}
let mut reduced = vec![0.0f32; tokens * matrix.out_features];
let mut rank_partials = Vec::with_capacity(tp);
for (rank_index, rank) in self.ranks.iter().enumerate() {
let (codes, scales) = row_shard(matrix, tp, rank_index)?;
let local_activations =
activation_shard(activations, tokens, matrix.in_features, tp, rank_index);
let shard = E4m3BlockMatrix {
codes: &codes,
scales: &scales,
out_features: matrix.out_features,
in_features: local_in,
};
let partial = run_rank(rank, shard, &local_activations, tokens)?;
for (sum, value) in reduced.iter_mut().zip(&partial) {
*sum += *value;
}
rank_partials.push(partial);
}
Ok(RowParallelResult {
reduced,
rank_partials,
})
}
pub fn upload_row_parallel(
&self,
matrix: E4m3BlockMatrix<'_>,
) -> Result<ResidentRowParallel, Box<dyn std::error::Error>> {
matrix.validate()?;
let tp = self.ranks.len();
validate_row_shape(matrix, tp)?;
let local_in = matrix.in_features / tp;
let mut ranks = Vec::with_capacity(tp);
for (rank_index, engine) in self.ranks.iter().enumerate() {
let (codes, scales) = row_shard(matrix, tp, rank_index)?;
ranks.push(upload_rank(
engine,
E4m3BlockMatrix {
codes: &codes,
scales: &scales,
out_features: matrix.out_features,
in_features: local_in,
},
)?);
}
Ok(ResidentRowParallel {
ranks,
out_features: matrix.out_features,
in_features: matrix.in_features,
})
}
pub fn row_parallel_resident(
&self,
matrix: &ResidentRowParallel,
activations: &[f32],
tokens: usize,
) -> Result<RowParallelResult, Box<dyn std::error::Error>> {
validate_resident_ranks(&self.ranks, &matrix.ranks)?;
validate_activations(activations, tokens, matrix.in_features)?;
let tp = self.ranks.len();
let mut reduced = vec![0.0f32; tokens * matrix.out_features];
let mut rank_partials = Vec::with_capacity(tp);
for (rank_index, (engine, shard)) in self.ranks.iter().zip(&matrix.ranks).enumerate() {
let local_activations =
activation_shard(activations, tokens, matrix.in_features, tp, rank_index);
let partial = run_resident_rank(engine, shard, &local_activations, tokens)?;
for (sum, value) in reduced.iter_mut().zip(&partial) {
*sum += *value;
}
rank_partials.push(partial);
}
Ok(RowParallelResult {
reduced,
rank_partials,
})
}
pub fn upload_bf16_column_parallel(
&self,
matrix: Bf16Matrix<'_>,
) -> Result<ResidentBf16ColumnParallel, Box<dyn std::error::Error>> {
self.upload_bf16_column_parallel_inner(matrix, None, false)
}
pub fn upload_step_bf16_column_parallel(
&self,
matrix: Bf16Matrix<'_>,
) -> Result<ResidentBf16ColumnParallel, Box<dyn std::error::Error>> {
self.upload_step_bf16_column_parallel_inner(matrix, false)
}
pub fn upload_step_bf16_column_parallel_f32_mirror(
&self,
matrix: Bf16Matrix<'_>,
) -> Result<ResidentBf16ColumnParallel, Box<dyn std::error::Error>> {
self.upload_step_bf16_column_parallel_inner(matrix, true)
}
fn upload_step_bf16_column_parallel_inner(
&self,
matrix: Bf16Matrix<'_>,
f32_mirror: bool,
) -> Result<ResidentBf16ColumnParallel, Box<dyn std::error::Error>> {
let canonical_chunk_rows =
step_bf16_canonical_chunk_rows(matrix.out_features, self.ranks.len())?;
self.upload_bf16_column_parallel_inner(matrix, Some(canonical_chunk_rows), f32_mirror)
}
fn upload_bf16_column_parallel_inner(
&self,
matrix: Bf16Matrix<'_>,
canonical_chunk_rows: Option<usize>,
f32_mirror: bool,
) -> Result<ResidentBf16ColumnParallel, Box<dyn std::error::Error>> {
matrix.validate()?;
let tp = self.ranks.len();
if matrix.out_features % tp != 0 {
return Err(format!(
"BF16 column-parallel out_features {} is not divisible by TP={tp}",
matrix.out_features
)
.into());
}
let mut ranks = Vec::with_capacity(tp);
for (rank, engine) in self.ranks.iter().enumerate() {
ranks.push(upload_bf16_rank(
engine,
bf16_column_shard(matrix, tp, rank)?,
f32_mirror,
)?);
}
Ok(ResidentBf16ColumnParallel {
ranks,
out_features: matrix.out_features,
in_features: matrix.in_features,
canonical_chunk_rows,
})
}
pub fn bf16_column_parallel_resident(
&self,
matrix: &ResidentBf16ColumnParallel,
activations: &[f32],
tokens: usize,
) -> Result<ColumnParallelResult, Box<dyn std::error::Error>> {
validate_resident_bf16_ranks(&self.ranks, &matrix.ranks)?;
validate_activations(activations, tokens, matrix.in_features)?;
let local_out = matrix.out_features / self.ranks.len();
let mut gathered = vec![0.0f32; tokens * matrix.out_features];
let mut rank_outputs = Vec::with_capacity(self.ranks.len());
for (rank, (engine, shard)) in self.ranks.iter().zip(&matrix.ranks).enumerate() {
let output = run_resident_bf16_rank(
engine,
shard,
activations,
tokens,
matrix.canonical_chunk_rows,
)?;
for token in 0..tokens {
let src = &output[token * local_out..(token + 1) * local_out];
let dst_start = token * matrix.out_features + rank * local_out;
gathered[dst_start..dst_start + local_out].copy_from_slice(src);
}
rank_outputs.push(output);
}
Ok(ColumnParallelResult {
gathered,
rank_outputs,
})
}
pub fn bf16_column_parallel_resident_native(
&self,
matrix: &ResidentBf16ColumnParallel,
activations: &[f32],
tokens: usize,
) -> Result<Vec<f32>, Box<dyn std::error::Error>> {
let rank_outputs =
self.bf16_column_parallel_resident_device_shards(matrix, activations, tokens)?;
let local_out = matrix.out_features / self.ranks.len();
self.gather_native_column_shards(&rank_outputs, tokens, local_out)
}
pub fn bf16_column_parallel_resident_device_shards(
&self,
matrix: &ResidentBf16ColumnParallel,
activations: &[f32],
tokens: usize,
) -> Result<Vec<CudaSlice<f32>>, Box<dyn std::error::Error>> {
if self.ranks.len() > 1 && !self.native_p2p {
return Err("device-resident BF16 column parallelism requires native P2P ranks".into());
}
validate_resident_bf16_ranks(&self.ranks, &matrix.ranks)?;
validate_activations(activations, tokens, matrix.in_features)?;
let mut rank_inputs = Vec::with_capacity(self.ranks.len());
let root_input = {
let root = &self.ranks[0];
let _main = root.gpu.enter_main()?;
root.htod(activations)?
};
{
let root = &self.ranks[0];
let _main = root.gpu.enter_main()?;
root.stream().synchronize()?;
}
rank_inputs.push(root_input);
for engine in &self.ranks[1..] {
let peer_input = {
let _main = engine.gpu.enter_main()?;
let mut peer_input = engine.uninit(activations.len())?;
engine
.stream()
.memcpy_dtod(&rank_inputs[0], &mut peer_input)?;
peer_input
};
rank_inputs.push(peer_input);
}
let mut rank_outputs = Vec::with_capacity(self.ranks.len());
for rank in 0..self.ranks.len() {
rank_outputs.push(run_resident_bf16_rank_device(
&self.ranks[rank],
&matrix.ranks[rank],
&rank_inputs[rank],
tokens,
matrix.canonical_chunk_rows,
self.bulk_p2p,
)?);
}
Ok(rank_outputs)
}
pub fn allocate_replicated_device_rows(
&self,
tokens: usize,
width: usize,
) -> Result<ResidentReplicatedDeviceRows, Box<dyn std::error::Error>> {
if self.ranks.len() > 1 && !self.native_p2p {
return Err("replicated device rows require native P2P ranks".into());
}
let values = tokens
.checked_mul(width)
.ok_or("replicated device row size overflow")?;
let rank_lengths = vec![values; self.ranks.len()];
replicated_device_row_values(tokens, width, self.ranks.len(), &rank_lengths)?;
let mut ranks = Vec::with_capacity(self.ranks.len());
for engine in &self.ranks {
let _main = engine.gpu.enter_main()?;
ranks.push(engine.uninit(values)?);
}
Ok(ResidentReplicatedDeviceRows {
ranks,
tokens,
width,
})
}
pub fn refresh_replicated_device_rows_from_root(
&self,
rows: &mut ResidentReplicatedDeviceRows,
source: &CudaSlice<f32>,
) -> Result<(), Box<dyn std::error::Error>> {
if self.ranks.len() > 1 && !self.native_p2p {
return Err("replicated device rows require native P2P ranks".into());
}
validate_replicated_device_rows(&self.ranks, rows)?;
let root = self
.ranks
.first()
.ok_or("replicated rows have no root rank")?;
let values = replicated_device_row_source_values(
rows.tokens,
rows.width,
source.len(),
source.ordinal(),
root.ctx().ordinal(),
)?;
let (root_rows, peer_rows) = rows
.ranks
.split_first_mut()
.ok_or("replicated rows have no root allocation")?;
{
let _main = root.gpu.enter_main()?;
let mut destination = root_rows.slice_mut(0..values);
root.stream()
.memcpy_dtod(&source.slice(0..values), &mut destination)?;
root.stream().synchronize()?;
}
for (engine, peer_rows) in self.ranks.iter().skip(1).zip(peer_rows) {
let _main = engine.gpu.enter_main()?;
let mut destination = peer_rows.slice_mut(0..values);
engine
.stream()
.memcpy_dtod(&root_rows.slice(0..values), &mut destination)?;
}
Ok(())
}
pub fn upload_replicated_device_rows(
&self,
rows: &[f32],
tokens: usize,
width: usize,
) -> Result<ResidentReplicatedDeviceRows, Box<dyn std::error::Error>> {
if self.ranks.len() > 1 && !self.native_p2p {
return Err("replicated device rows require native P2P ranks".into());
}
validate_activations(rows, tokens, width)?;
let root = self
.ranks
.first()
.ok_or("replicated rows have no root rank")?;
let root_rows = {
let _main = root.gpu.enter_main()?;
root.htod(rows)?
};
{
let _main = root.gpu.enter_main()?;
root.stream().synchronize()?;
}
let mut ranks = Vec::with_capacity(self.ranks.len());
ranks.push(root_rows);
for engine in self.ranks.iter().skip(1) {
let _main = engine.gpu.enter_main()?;
let mut peer_rows = engine.uninit(rows.len())?;
engine.stream().memcpy_dtod(&ranks[0], &mut peer_rows)?;
ranks.push(peer_rows);
}
Ok(ResidentReplicatedDeviceRows {
ranks,
tokens,
width,
})
}
pub fn bf16_column_parallel_resident_replicated_device_shards(
&self,
matrix: &ResidentBf16ColumnParallel,
activations: &ResidentReplicatedDeviceRows,
) -> Result<Vec<CudaSlice<f32>>, Box<dyn std::error::Error>> {
validate_resident_bf16_ranks(&self.ranks, &matrix.ranks)?;
validate_replicated_device_rows(&self.ranks, activations)?;
if activations.width != matrix.in_features {
return Err(format!(
"replicated BF16 column input width {} != matrix width {}",
activations.width, matrix.in_features
)
.into());
}
let mut outputs = Vec::with_capacity(self.ranks.len());
for rank in 0..self.ranks.len() {
outputs.push(run_resident_bf16_rank_device(
&self.ranks[rank],
&matrix.ranks[rank],
&activations.ranks[rank],
activations.tokens,
matrix.canonical_chunk_rows,
self.bulk_p2p,
)?);
}
Ok(outputs)
}
#[allow(clippy::too_many_arguments)]
pub fn upload_sigmoid_topk_router(
&self,
weight: Bf16Matrix<'_>,
correction_bias: &[f32],
active: Option<&[bool]>,
experts_per_token: usize,
scaling_factor: f32,
route_norm: bool,
) -> Result<ResidentSigmoidTopKRouter, Box<dyn std::error::Error>> {
weight.validate()?;
if correction_bias.len() != weight.out_features
|| experts_per_token == 0
|| experts_per_token > weight.out_features
|| !correction_bias.iter().all(|value| value.is_finite())
|| !scaling_factor.is_finite()
|| scaling_factor <= 0.0
{
return Err(format!(
"sigmoid router geometry weight={}x{} bias={} top_k={} scale={scaling_factor}",
weight.out_features,
weight.in_features,
correction_bias.len(),
experts_per_token,
)
.into());
}
let active_row = active
.map(|mask| {
if mask.len() != weight.out_features {
return Err(format!(
"sigmoid router active mask {} != experts {}",
mask.len(),
weight.out_features
));
}
Ok(mask
.iter()
.map(|&enabled| u8::from(enabled))
.collect::<Vec<_>>())
})
.transpose()?
.unwrap_or_else(|| vec![1; weight.out_features]);
let active_count = active_row.iter().filter(|&&enabled| enabled != 0).count();
crate::sigrouter_contract::validate_active_count(experts_per_token, active_count)?;
let root = self
.ranks
.first()
.ok_or("sigmoid router runtime has no root rank")?;
let _main = root.gpu.enter_main()?;
let bf16 = root.htod_bytes(weight.bytes)?;
let weight_f32 = root.bf16_to_f32(
&bf16.slice(0..bf16.len()),
weight.out_features * weight.in_features,
)?;
Ok(ResidentSigmoidTopKRouter {
weight: weight_f32,
correction_bias: root.htod(correction_bias)?,
active: root.htod_bytes(&active_row)?,
root_device: root.ctx().ordinal(),
input_width: weight.in_features,
expert_count: weight.out_features,
experts_per_token,
active_count,
scaling_factor,
route_norm,
})
}
pub fn sigmoid_topk_replicated_device_rows_host(
&self,
router: &ResidentSigmoidTopKRouter,
input: &ResidentReplicatedDeviceRows,
) -> Result<SigmoidTopKHostOutput, Box<dyn std::error::Error>> {
validate_replicated_device_rows(&self.ranks, input)?;
if input.width != router.input_width {
return Err(format!(
"sigmoid router input width {} != resident width {}",
input.width, router.input_width
)
.into());
}
let root = self
.ranks
.first()
.ok_or("sigmoid router runtime has no root rank")?;
let _main = root.gpu.enter_main()?;
if root.ctx().ordinal() != router.root_device
|| router.weight.ordinal() != router.root_device
|| router.correction_bias.ordinal() != router.root_device
|| router.active.ordinal() != router.root_device
{
return Err("sigmoid router root residency changed".into());
}
let logits = root.router_gemv(
&router.weight,
&input.ranks[0],
router.input_width,
router.expert_count,
input.tokens,
)?;
let (selected, weights) = root.moe_router_sigmoid_topk_host(
&logits,
input.tokens,
router.expert_count,
router.experts_per_token,
router.active_count,
&router.correction_bias,
&router.active,
router.scaling_factor,
router.route_norm,
)?;
Ok(SigmoidTopKHostOutput {
logits: root.dtoh(&logits)?,
selected,
weights,
})
}
pub fn upload_replicated_bf16_swiglu(
&self,
gate: Bf16Matrix<'_>,
up: Bf16Matrix<'_>,
down: Bf16Matrix<'_>,
) -> Result<ResidentReplicatedBf16SwiGlu, Box<dyn std::error::Error>> {
gate.validate()?;
up.validate()?;
down.validate()?;
if gate.in_features != up.in_features
|| gate.out_features != up.out_features
|| down.in_features != gate.out_features
|| down.out_features != gate.in_features
{
return Err(format!(
"replicated BF16 SwiGLU geometry gate={}x{} up={}x{} down={}x{}",
gate.out_features,
gate.in_features,
up.out_features,
up.in_features,
down.out_features,
down.in_features,
)
.into());
}
let mut gate_ranks = Vec::with_capacity(self.ranks.len());
let mut up_ranks = Vec::with_capacity(self.ranks.len());
let mut down_ranks = Vec::with_capacity(self.ranks.len());
for engine in &self.ranks {
gate_ranks.push(upload_bf16_rank(engine, gate, false)?);
up_ranks.push(upload_bf16_rank(engine, up, false)?);
down_ranks.push(upload_bf16_rank(engine, down, false)?);
}
Ok(ResidentReplicatedBf16SwiGlu {
gate: gate_ranks,
up: up_ranks,
down: down_ranks,
input_width: gate.in_features,
intermediate_width: gate.out_features,
})
}
pub fn replicated_bf16_swiglu_resident_device(
&self,
mlp: &ResidentReplicatedBf16SwiGlu,
input: &ResidentReplicatedDeviceRows,
activation_limit: Option<f32>,
) -> Result<ResidentReplicatedDeviceRows, Box<dyn std::error::Error>> {
validate_step_expert_activation_limit(activation_limit)?;
validate_replicated_device_rows(&self.ranks, input)?;
validate_resident_bf16_ranks(&self.ranks, &mlp.gate)?;
validate_resident_bf16_ranks(&self.ranks, &mlp.up)?;
validate_resident_bf16_ranks(&self.ranks, &mlp.down)?;
if input.width != mlp.input_width
|| mlp.gate.len() != self.ranks.len()
|| mlp.up.len() != self.ranks.len()
|| mlp.down.len() != self.ranks.len()
{
return Err("replicated BF16 SwiGLU residency or input width changed".into());
}
let mut outputs = Vec::with_capacity(self.ranks.len());
for rank in 0..self.ranks.len() {
let engine = &self.ranks[rank];
let gate = run_resident_bf16_rank_device(
engine,
&mlp.gate[rank],
&input.ranks[rank],
input.tokens,
None,
self.bulk_p2p,
)?;
let up = run_resident_bf16_rank_device(
engine,
&mlp.up[rank],
&input.ranks[rank],
input.tokens,
None,
self.bulk_p2p,
)?;
let _main = engine.gpu.enter_main()?;
let values = input
.tokens
.checked_mul(mlp.intermediate_width)
.ok_or("replicated BF16 SwiGLU activation size overflow")?;
let mut activation = engine.uninit(values)?;
if let Some(limit) = activation_limit {
engine.silu_clamped_mul_host_expf(&gate, &up, limit, &mut activation, values)?;
} else {
engine.silu_mul_host_expf(&gate, &up, &mut activation, values)?;
}
outputs.push(run_resident_bf16_rank_device(
engine,
&mlp.down[rank],
&activation,
input.tokens,
None,
self.bulk_p2p,
)?);
}
Ok(ResidentReplicatedDeviceRows {
ranks: outputs,
tokens: input.tokens,
width: mlp.input_width,
})
}
pub fn rms_norm_replicated_device_rows(
&self,
input: &ResidentReplicatedDeviceRows,
weight: &[f32],
eps: f32,
) -> Result<ResidentReplicatedDeviceRows, Box<dyn std::error::Error>> {
validate_replicated_device_rows(&self.ranks, input)?;
if weight.len() != input.width || !eps.is_finite() || eps <= 0.0 {
return Err(format!(
"replicated RMS norm weight/eps {}/{} != width {}",
weight.len(),
eps,
input.width
)
.into());
}
let mut ranks = Vec::with_capacity(self.ranks.len());
for (rank, engine) in self.ranks.iter().enumerate() {
let _main = engine.gpu.enter_main()?;
let weight = engine.htod(weight)?;
let mut output = engine.uninit(input.tokens * input.width)?;
engine.rms_norm(
&input.ranks[rank],
&weight,
&mut output,
input.width,
input.tokens,
eps,
)?;
ranks.push(output);
}
Ok(ResidentReplicatedDeviceRows {
ranks,
tokens: input.tokens,
width: input.width,
})
}
pub fn add_rms_norm_replicated_device_rows(
&self,
input: &ResidentReplicatedDeviceRows,
update: &ResidentReplicatedDeviceRows,
weight: &[f32],
eps: f32,
) -> Result<
(ResidentReplicatedDeviceRows, ResidentReplicatedDeviceRows),
Box<dyn std::error::Error>,
> {
validate_replicated_device_rows(&self.ranks, input)?;
validate_replicated_device_rows(&self.ranks, update)?;
if input.tokens != update.tokens
|| input.width != update.width
|| weight.len() != input.width
|| !eps.is_finite()
|| eps <= 0.0
{
return Err(format!(
"replicated add/RMS geometry input={}x{} update={}x{} weight={} eps={eps}",
input.tokens,
input.width,
update.tokens,
update.width,
weight.len(),
)
.into());
}
let values = input.tokens * input.width;
let mut residual_ranks = Vec::with_capacity(self.ranks.len());
let mut normalized_ranks = Vec::with_capacity(self.ranks.len());
for (rank, engine) in self.ranks.iter().enumerate() {
let _main = engine.gpu.enter_main()?;
let weight = engine.htod(weight)?;
let mut residual = engine.uninit(values)?;
let mut normalized = engine.uninit(values)?;
engine.add_rms_norm(
&input.ranks[rank],
&update.ranks[rank],
&weight,
&mut residual,
&mut normalized,
input.width,
input.tokens,
eps,
)?;
residual_ranks.push(residual);
normalized_ranks.push(normalized);
}
Ok((
ResidentReplicatedDeviceRows {
ranks: residual_ranks,
tokens: input.tokens,
width: input.width,
},
ResidentReplicatedDeviceRows {
ranks: normalized_ranks,
tokens: input.tokens,
width: input.width,
},
))
}
pub fn collect_replicated_device_rows(
&self,
rows: &ResidentReplicatedDeviceRows,
) -> Result<Vec<Vec<f32>>, Box<dyn std::error::Error>> {
validate_replicated_device_rows(&self.ranks, rows)?;
let mut outputs = Vec::with_capacity(self.ranks.len());
for (rank, engine) in self.ranks.iter().enumerate() {
let _main = engine.gpu.enter_main()?;
outputs.push(engine.dtoh(&rows.ranks[rank])?);
}
Ok(outputs)
}
pub fn upload_bf16_row_parallel(
&self,
matrix: Bf16Matrix<'_>,
) -> Result<ResidentBf16RowParallel, Box<dyn std::error::Error>> {
matrix.validate()?;
let tp = self.ranks.len();
if matrix.in_features % tp != 0 {
return Err(format!(
"BF16 row-parallel in_features {} is not divisible by TP={tp}",
matrix.in_features
)
.into());
}
let mut ranks = Vec::with_capacity(tp);
for (rank, engine) in self.ranks.iter().enumerate() {
let shard = bf16_row_shard(matrix, tp, rank)?;
ranks.push(upload_bf16_rank(
engine,
Bf16Matrix {
bytes: &shard,
out_features: matrix.out_features,
in_features: matrix.in_features / tp,
},
false,
)?);
}
Ok(ResidentBf16RowParallel {
ranks,
out_features: matrix.out_features,
in_features: matrix.in_features,
})
}
pub fn bf16_row_parallel_resident(
&self,
matrix: &ResidentBf16RowParallel,
activations: &[f32],
tokens: usize,
) -> Result<RowParallelResult, Box<dyn std::error::Error>> {
validate_resident_bf16_ranks(&self.ranks, &matrix.ranks)?;
validate_activations(activations, tokens, matrix.in_features)?;
let tp = self.ranks.len();
let mut reduced = vec![0.0f32; tokens * matrix.out_features];
let mut rank_partials = Vec::with_capacity(tp);
for (rank, (engine, shard)) in self.ranks.iter().zip(&matrix.ranks).enumerate() {
let local_activations =
activation_shard(activations, tokens, matrix.in_features, tp, rank);
let partial = run_resident_bf16_rank(engine, shard, &local_activations, tokens, None)?;
for (sum, value) in reduced.iter_mut().zip(&partial) {
*sum += value;
}
rank_partials.push(partial);
}
Ok(RowParallelResult {
reduced,
rank_partials,
})
}
pub fn upload_step_bf16_row_parallel(
&self,
matrix: Bf16Matrix<'_>,
) -> Result<ResidentStepBf16RowParallel, Box<dyn std::error::Error>> {
self.upload_step_bf16_row_parallel_inner(matrix, false)
}
pub fn upload_step_bf16_row_parallel_f32_mirror(
&self,
matrix: Bf16Matrix<'_>,
) -> Result<ResidentStepBf16RowParallel, Box<dyn std::error::Error>> {
self.upload_step_bf16_row_parallel_inner(matrix, true)
}
fn upload_step_bf16_row_parallel_inner(
&self,
matrix: Bf16Matrix<'_>,
f32_mirror: bool,
) -> Result<ResidentStepBf16RowParallel, Box<dyn std::error::Error>> {
matrix.validate()?;
let tp = self.ranks.len();
let canonical_chunk_cols = step_bf16_canonical_chunk_cols(matrix.in_features, tp)?;
let local_in = matrix.in_features / tp;
let blocks_per_rank = local_in / canonical_chunk_cols;
let mut ranks = Vec::with_capacity(tp);
for (rank, engine) in self.ranks.iter().enumerate() {
let mut blocks = Vec::with_capacity(blocks_per_rank);
for block in 0..blocks_per_rank {
let global_block = rank * blocks_per_rank + block;
let col_start = global_block * canonical_chunk_cols;
let bytes = bf16_row_block(matrix, col_start, canonical_chunk_cols)?;
blocks.push(upload_bf16_rank(
engine,
Bf16Matrix {
bytes: &bytes,
out_features: matrix.out_features,
in_features: canonical_chunk_cols,
},
f32_mirror,
)?);
}
ranks.push(blocks);
}
Ok(ResidentStepBf16RowParallel {
ranks,
out_features: matrix.out_features,
in_features: matrix.in_features,
canonical_chunk_cols,
})
}
pub fn step_bf16_row_parallel_resident(
&self,
matrix: &ResidentStepBf16RowParallel,
activations: &[f32],
tokens: usize,
) -> Result<Vec<f32>, Box<dyn std::error::Error>> {
validate_step_bf16_row_residency(&self.ranks, matrix)?;
validate_activations(activations, tokens, matrix.in_features)?;
let root = &self.ranks[0];
let output_len = tokens
.checked_mul(matrix.out_features)
.ok_or("Step BF16 row output size overflow")?;
let mut reduced = {
let _main = root.gpu.enter_main()?;
root.htod(&vec![0.0f32; output_len])?
};
let blocks_per_rank = PRODUCT_MAX_CARDS / self.ranks.len();
for (rank, blocks) in matrix.ranks.iter().enumerate() {
for (block, resident) in blocks.iter().enumerate() {
let global_block = rank * blocks_per_rank + block;
let input = activation_shard(
activations,
tokens,
matrix.in_features,
PRODUCT_MAX_CARDS,
global_block,
);
let partial =
run_resident_bf16_rank(&self.ranks[rank], resident, &input, tokens, None)?;
let next = {
let _main = root.gpu.enter_main()?;
let partial = root.htod(&partial)?;
let mut next = root.uninit(output_len)?;
root.add(&reduced, &partial, &mut next, output_len)?;
next
};
reduced = next;
}
}
let _main = root.gpu.enter_main()?;
root.dtoh(&reduced)
}
pub fn step_bf16_row_parallel_resident_native(
&self,
matrix: &ResidentStepBf16RowParallel,
activations: &[f32],
tokens: usize,
) -> Result<Vec<f32>, Box<dyn std::error::Error>> {
if self.ranks.len() > 1 && !self.native_p2p {
return Err("native Step BF16 row parallelism requires P2P ranks".into());
}
validate_step_bf16_row_residency(&self.ranks, matrix)?;
validate_activations(activations, tokens, matrix.in_features)?;
let root = &self.ranks[0];
let root_input = {
let _main = root.gpu.enter_main()?;
root.htod(activations)?
};
let output_len = tokens
.checked_mul(matrix.out_features)
.ok_or("native Step BF16 row output size overflow")?;
let mut reduced = {
let _main = root.gpu.enter_main()?;
root.htod(&vec![0.0f32; output_len])?
};
{
let _main = root.gpu.enter_main()?;
root.stream().synchronize()?;
}
let blocks_per_rank = PRODUCT_MAX_CARDS / self.ranks.len();
let mut block_input_keepalive = Vec::with_capacity(PRODUCT_MAX_CARDS);
let mut root_packed_keepalive = Vec::with_capacity(PRODUCT_MAX_CARDS);
let mut remote_partial_keepalive = Vec::new();
for (rank, blocks) in matrix.ranks.iter().enumerate() {
for (block, resident) in blocks.iter().enumerate() {
let global_block = rank * blocks_per_rank + block;
let col_start = global_block * matrix.canonical_chunk_cols;
let block_len = tokens
.checked_mul(matrix.canonical_chunk_cols)
.ok_or("native Step BF16 row block size overflow")?;
let block_input = if self.bulk_p2p {
let root_packed = {
let _main = root.gpu.enter_main()?;
let mut root_packed = root.uninit(block_len)?;
root.copy_rows_strided(
&root_input,
&mut root_packed,
matrix.canonical_chunk_cols,
tokens,
matrix.in_features,
col_start,
)?;
root_packed
};
if rank == 0 {
root_packed
} else {
{
let _main = root.gpu.enter_main()?;
root.stream().synchronize()?;
}
let engine = &self.ranks[rank];
let _main = engine.gpu.enter_main()?;
let mut block_input = engine.uninit(block_len)?;
engine
.stream()
.memcpy_dtod(&root_packed, &mut block_input)?;
root_packed_keepalive.push(root_packed);
block_input
}
} else {
let engine = &self.ranks[rank];
let _main = engine.gpu.enter_main()?;
let mut block_input = engine.uninit(block_len)?;
for token in 0..tokens {
let source_start = token * matrix.in_features + col_start;
let source = root_input
.slice(source_start..source_start + matrix.canonical_chunk_cols);
let destination_start = token * matrix.canonical_chunk_cols;
let mut destination = block_input.slice_mut(
destination_start..destination_start + matrix.canonical_chunk_cols,
);
engine.stream().memcpy_dtod(&source, &mut destination)?;
}
block_input
};
let partial = run_resident_bf16_rank_device(
&self.ranks[rank],
resident,
&block_input,
tokens,
None,
self.bulk_p2p,
)?;
block_input_keepalive.push(block_input);
let root_partial = if rank == 0 {
partial
} else {
{
let engine = &self.ranks[rank];
let _main = engine.gpu.enter_main()?;
engine.stream().synchronize()?;
}
let _main = root.gpu.enter_main()?;
let mut peer_partial = root.uninit(output_len)?;
root.stream().memcpy_dtod(&partial, &mut peer_partial)?;
remote_partial_keepalive.push(partial);
peer_partial
};
let next = {
let _main = root.gpu.enter_main()?;
let mut next = root.uninit(output_len)?;
root.add(&reduced, &root_partial, &mut next, output_len)?;
next
};
reduced = next;
}
}
let output = {
let _main = root.gpu.enter_main()?;
root.dtoh(&reduced)?
};
drop(remote_partial_keepalive);
drop(root_packed_keepalive);
drop(block_input_keepalive);
Ok(output)
}
pub fn step_bf16_row_parallel_resident_root_device(
&self,
matrix: &ResidentStepBf16RowParallel,
rank_activations: &[CudaSlice<f32>],
tokens: usize,
) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
if self.ranks.len() > 1 && !self.native_p2p {
return Err(
"device-resident Step BF16 row parallelism requires native P2P ranks".into(),
);
}
validate_step_bf16_row_residency(&self.ranks, matrix)?;
let local_width = matrix.in_features / self.ranks.len();
let shard_len = tokens
.checked_mul(local_width)
.ok_or("device Step BF16 row shard size overflow")?;
if tokens == 0
|| rank_activations.len() != self.ranks.len()
|| rank_activations
.iter()
.zip(&self.ranks)
.any(|(rows, engine)| {
rows.len() != shard_len || rows.ordinal() != engine.ctx().ordinal()
})
{
return Err("device Step BF16 row activation shard geometry changed".into());
}
let blocks_per_rank = PRODUCT_MAX_CARDS / self.ranks.len();
let mut block_inputs = Vec::with_capacity(self.ranks.len());
let mut partials = Vec::with_capacity(self.ranks.len());
for (rank, blocks) in matrix.ranks.iter().enumerate() {
if blocks.len() != blocks_per_rank {
return Err(format!(
"device Step BF16 row rank {rank} blocks {} != {blocks_per_rank}",
blocks.len()
)
.into());
}
let engine = &self.ranks[rank];
let _main = engine.gpu.enter_main()?;
let mut rank_inputs = Vec::with_capacity(blocks_per_rank);
let mut rank_partials = Vec::with_capacity(blocks_per_rank);
for (block, resident) in blocks.iter().enumerate() {
let block_len = tokens
.checked_mul(matrix.canonical_chunk_cols)
.ok_or("device Step BF16 row block size overflow")?;
let mut block_input = engine.uninit(block_len)?;
let local_col_start = block * matrix.canonical_chunk_cols;
if self.bulk_p2p {
engine.copy_rows_strided(
&rank_activations[rank],
&mut block_input,
matrix.canonical_chunk_cols,
tokens,
local_width,
local_col_start,
)?;
} else {
for token in 0..tokens {
let source_start = token * local_width + local_col_start;
let source = rank_activations[rank]
.slice(source_start..source_start + matrix.canonical_chunk_cols);
let destination_start = token * matrix.canonical_chunk_cols;
let mut destination = block_input.slice_mut(
destination_start..destination_start + matrix.canonical_chunk_cols,
);
engine.stream().memcpy_dtod(&source, &mut destination)?;
}
}
let partial = run_resident_bf16_rank_device(
engine,
resident,
&block_input,
tokens,
None,
self.bulk_p2p,
)?;
rank_inputs.push(block_input);
rank_partials.push(partial);
}
block_inputs.push(rank_inputs);
partials.push(rank_partials);
}
for engine in self.ranks.iter().skip(1) {
let _main = engine.gpu.enter_main()?;
engine.stream().synchronize()?;
}
let output_len = tokens
.checked_mul(matrix.out_features)
.ok_or("device Step BF16 row output size overflow")?;
let root = &self.ranks[0];
let _main = root.gpu.enter_main()?;
let mut reduced = root.htod(&vec![0.0f32; output_len])?;
let mut remote_partials = Vec::new();
for (rank, rank_partials) in partials.into_iter().enumerate() {
for partial in rank_partials {
let root_partial = if rank == 0 {
partial
} else {
let mut peer_partial = root.uninit(output_len)?;
root.stream().memcpy_dtod(&partial, &mut peer_partial)?;
remote_partials.push(partial);
peer_partial
};
let mut next = root.uninit(output_len)?;
root.add(&reduced, &root_partial, &mut next, output_len)?;
reduced = next;
}
}
root.stream().synchronize()?;
drop(remote_partials);
drop(block_inputs);
Ok(reduced)
}
pub fn step_bf16_row_parallel_resident_replicated_device(
&self,
matrix: &ResidentStepBf16RowParallel,
rank_activations: &[CudaSlice<f32>],
tokens: usize,
) -> Result<ResidentReplicatedDeviceRows, Box<dyn std::error::Error>> {
let reduced =
self.step_bf16_row_parallel_resident_root_device(matrix, rank_activations, tokens)?;
let output_len = tokens
.checked_mul(matrix.out_features)
.ok_or("device Step BF16 row output size overflow")?;
let mut ranks = Vec::with_capacity(self.ranks.len());
ranks.push(reduced);
for engine in self.ranks.iter().skip(1) {
let _main = engine.gpu.enter_main()?;
let mut peer_output = engine.uninit(output_len)?;
engine.stream().memcpy_dtod(&ranks[0], &mut peer_output)?;
ranks.push(peer_output);
}
Ok(ResidentReplicatedDeviceRows {
ranks,
tokens,
width: matrix.out_features,
})
}
pub fn upload_expert(
&self,
gate: E4m3BlockMatrix<'_>,
up: E4m3BlockMatrix<'_>,
down: E4m3BlockMatrix<'_>,
) -> Result<ResidentTpExpert, Box<dyn std::error::Error>> {
if gate.in_features != up.in_features || gate.out_features != up.out_features {
return Err("TP expert gate/up dimensions differ".into());
}
if down.in_features != gate.out_features || down.out_features != gate.in_features {
return Err(format!(
"TP expert down {}x{} does not invert gate/up {}x{}",
down.out_features, down.in_features, gate.out_features, gate.in_features
)
.into());
}
Ok(ResidentTpExpert {
gate: self.upload_column_parallel(gate)?,
up: self.upload_column_parallel(up)?,
down: self.upload_row_parallel(down)?,
input_width: gate.in_features,
expert_width: gate.out_features,
})
}
pub fn run_expert(
&self,
expert: &ResidentTpExpert,
input: &[f32],
tokens: usize,
) -> Result<Vec<f32>, Box<dyn std::error::Error>> {
validate_activations(input, tokens, expert.input_width)?;
let gate = self.column_parallel_resident(&expert.gate, input, tokens)?;
let up = self.column_parallel_resident(&expert.up, input, tokens)?;
let activated: Vec<f32> = gate
.gathered
.iter()
.zip(&up.gathered)
.map(|(&gate, &up)| gate / (1.0 + (-gate).exp()) * up)
.collect();
debug_assert_eq!(activated.len(), tokens * expert.expert_width);
Ok(self
.row_parallel_resident(&expert.down, &activated, tokens)?
.reduced)
}
pub fn upload_expert_parallel(
&self,
gate: E4m3ExpertBank<'_>,
up: E4m3ExpertBank<'_>,
down: E4m3ExpertBank<'_>,
) -> Result<ResidentExpertParallel, Box<dyn std::error::Error>> {
gate.validate()?;
up.validate()?;
down.validate()?;
if gate.expert_count != up.expert_count || gate.expert_count != down.expert_count {
return Err("EP gate/up/down expert counts differ".into());
}
if gate.in_features != up.in_features || gate.out_features != up.out_features {
return Err("EP gate/up dimensions differ".into());
}
if down.in_features != gate.out_features || down.out_features != gate.in_features {
return Err(format!(
"EP down {}x{} does not invert gate/up {}x{}",
down.out_features, down.in_features, gate.out_features, gate.in_features
)
.into());
}
if gate.expert_count % self.ranks.len() != 0 {
return Err(format!(
"EP expert count {} is not divisible by {} ranks",
gate.expert_count,
self.ranks.len()
)
.into());
}
let per_rank = gate.expert_count / self.ranks.len();
let mut ranks = Vec::with_capacity(self.ranks.len());
for (rank, engine) in self.ranks.iter().enumerate() {
let expert_range = rank * per_rank..(rank + 1) * per_rank;
ranks.push(ResidentEpRank {
gate: upload_expert_bank_rank(engine, gate, expert_range.clone())?,
up: upload_expert_bank_rank(engine, up, expert_range.clone())?,
down: upload_expert_bank_rank(engine, down, expert_range)?,
});
}
Ok(ResidentExpertParallel {
ranks,
expert_count: gate.expert_count,
input_width: gate.in_features,
expert_width: gate.out_features,
})
}
#[allow(clippy::too_many_arguments)]
pub fn prepare_step_grouped_fp8_gate(
&self,
gate: E4m3ExpertBank<'_>,
up: E4m3ExpertBank<'_>,
down: E4m3ExpertBank<'_>,
input: &[f32],
tokens: usize,
selected: &[usize],
activation_limit: Option<f32>,
) -> Result<PreparedStepGroupedFp8Gate, Box<dyn std::error::Error>> {
gate.validate()?;
up.validate()?;
down.validate()?;
validate_step_expert_activation_limit(activation_limit)?;
if gate.expert_count != STEP_GROUPED_FP8_EXPERTS
|| up.expert_count != STEP_GROUPED_FP8_EXPERTS
|| down.expert_count != STEP_GROUPED_FP8_EXPERTS
{
return Err(format!(
"official Step grouped FP8 gate requires {STEP_GROUPED_FP8_EXPERTS} experts, \
got gate/up/down={}/{}/{}",
gate.expert_count, up.expert_count, down.expert_count,
)
.into());
}
if gate.in_features != up.in_features
|| gate.out_features != STEP_GROUPED_FP8_WIDTH
|| up.out_features != STEP_GROUPED_FP8_WIDTH
|| down.in_features != STEP_GROUPED_FP8_WIDTH
|| down.out_features != gate.in_features
{
return Err(format!(
"official Step grouped FP8 geometry gate={}x{} up={}x{} down={}x{}",
gate.out_features,
gate.in_features,
up.out_features,
up.in_features,
down.out_features,
down.in_features,
)
.into());
}
validate_activations(input, tokens, gate.in_features)?;
let pairs = tokens
.checked_mul(STEP_GROUPED_FP8_TOP_K)
.ok_or("official Step grouped FP8 route count overflow")?;
if selected.len() != pairs {
return Err(format!(
"official Step grouped FP8 routes {} != {tokens}x{STEP_GROUPED_FP8_TOP_K} \
({pairs})",
selected.len()
)
.into());
}
for (token, routes) in selected.chunks_exact(STEP_GROUPED_FP8_TOP_K).enumerate() {
let mut unique = routes.to_vec();
unique.sort_unstable();
unique.dedup();
if unique.len() != STEP_GROUPED_FP8_TOP_K {
return Err(format!(
"official Step grouped FP8 token {token} routes are not top-8 unique: \
{routes:?}"
)
.into());
}
}
let engine = self
.ranks
.first()
.ok_or("official Step grouped FP8 gate has no rank-zero engine")?;
let _main = engine.gpu.enter_main()?;
let expert_range = 0..STEP_GROUPED_FP8_EXPERTS;
let gate = upload_expert_bank_rank(engine, gate, expert_range.clone())?;
let up = upload_expert_bank_rank(engine, up, expert_range.clone())?;
let down = upload_expert_bank_rank(engine, down, expert_range)?;
let input = engine.htod(input)?;
let route_csr = ExpertCsr::from_token_routes(
STEP_GROUPED_FP8_EXPERTS,
tokens,
STEP_GROUPED_FP8_TOP_K,
selected,
)?
.upload(engine)?;
let pair_rows = (0..pairs).collect::<Vec<_>>();
let down_csr =
ExpertCsr::from_pair_rows(STEP_GROUPED_FP8_EXPERTS, pairs, selected, &pair_rows)?
.upload(engine)?;
let gate_workspace =
Fp8GroupedWorkspace::new(engine, gate.in_features, gate.out_features, tokens, pairs)?;
let up_workspace =
Fp8GroupedWorkspace::new(engine, up.in_features, up.out_features, tokens, pairs)?;
let down_workspace =
Fp8GroupedWorkspace::new(engine, down.in_features, down.out_features, pairs, pairs)?;
let activation = engine.uninit(pairs * STEP_GROUPED_FP8_WIDTH)?;
Ok(PreparedStepGroupedFp8Gate {
device: engine.ctx().ordinal(),
gate,
up,
down,
input,
route_csr,
down_csr,
gate_workspace,
up_workspace,
down_workspace,
activation,
activation_limit,
tokens,
pairs,
})
}
pub fn run_step_grouped_fp8_gate(
&self,
plan: &mut PreparedStepGroupedFp8Gate,
) -> Result<StepGroupedFp8ProjectionOutput, Box<dyn std::error::Error>> {
let engine = self
.ranks
.first()
.ok_or("official Step grouped FP8 gate has no rank-zero engine")?;
if engine.ctx().ordinal() != plan.device {
return Err(format!(
"official Step grouped FP8 plan device {} != rank-zero device {}",
plan.device,
engine.ctx().ordinal()
)
.into());
}
let _main = engine.gpu.enter_main()?;
plan.gate_workspace.quantize(engine, &plan.input)?;
plan.gate_workspace.project(
engine,
&plan.gate.codes,
&plan.gate.scales,
&plan.route_csr,
plan.gate.code_stride,
plan.gate.scale_stride,
1.0,
)?;
plan.up_workspace.quantize(engine, &plan.input)?;
plan.up_workspace.project(
engine,
&plan.up.codes,
&plan.up.scales,
&plan.route_csr,
plan.up.code_stride,
plan.up.scale_stride,
1.0,
)?;
if let Some(limit) = plan.activation_limit {
engine.silu_clamped_mul_host_expf(
plan.gate_workspace.output(),
plan.up_workspace.output(),
limit,
&mut plan.activation,
plan.pairs * STEP_GROUPED_FP8_WIDTH,
)?;
} else {
engine.silu_mul_host_expf(
plan.gate_workspace.output(),
plan.up_workspace.output(),
&mut plan.activation,
plan.pairs * STEP_GROUPED_FP8_WIDTH,
)?;
}
plan.down_workspace.quantize(engine, &plan.activation)?;
plan.down_workspace.project(
engine,
&plan.down.codes,
&plan.down.scales,
&plan.down_csr,
plan.down.code_stride,
plan.down.scale_stride,
1.0,
)?;
Ok(StepGroupedFp8ProjectionOutput {
gate: engine.dtoh(plan.gate_workspace.output())?,
up: engine.dtoh(plan.up_workspace.output())?,
down: engine.dtoh(plan.down_workspace.output())?,
})
}
pub fn prepare_step_grouped_expert_parallel_gate(
&self,
experts: &ResidentExpertParallel,
input: &[f32],
tokens: usize,
selected: &[usize],
activation_limit: Option<f32>,
) -> Result<PreparedStepGroupedExpertParallelGate, Box<dyn std::error::Error>> {
self.prepare_step_grouped_expert_parallel_gate_with_capacity(
experts,
input,
tokens,
selected,
activation_limit,
tokens,
)
}
#[allow(clippy::too_many_arguments)]
pub fn prepare_step_grouped_expert_parallel_gate_with_capacity(
&self,
experts: &ResidentExpertParallel,
input: &[f32],
tokens: usize,
selected: &[usize],
activation_limit: Option<f32>,
max_tokens: usize,
) -> Result<PreparedStepGroupedExpertParallelGate, Box<dyn std::error::Error>> {
if !self.native_p2p || !self.ep_device_arithmetic {
return Err(
"Step owner-grouped FP8 requires native P2P and device-resident arithmetic".into(),
);
}
validate_step_expert_activation_limit(activation_limit)?;
validate_ep_residency(&self.ranks, experts)?;
validate_activations(input, tokens, experts.input_width)?;
if max_tokens < tokens || max_tokens > i32::MAX as usize {
return Err(format!(
"official Step owner-grouped FP8 tokens {tokens} exceed capacity {max_tokens}"
)
.into());
}
if experts.expert_count != STEP_GROUPED_FP8_EXPERTS
|| experts.expert_width != STEP_GROUPED_FP8_WIDTH
{
return Err(format!(
"official Step owner-grouped FP8 requires {} experts at width {}, got {} at {}",
STEP_GROUPED_FP8_EXPERTS,
STEP_GROUPED_FP8_WIDTH,
experts.expert_count,
experts.expert_width,
)
.into());
}
validate_step_grouped_owner_routes(experts.expert_count, tokens, selected)?;
let max_pairs = max_tokens
.checked_mul(STEP_GROUPED_FP8_TOP_K)
.ok_or("official Step owner-grouped FP8 capacity route count overflow")?;
let input_capacity = max_tokens
.checked_mul(experts.input_width)
.ok_or("official Step owner-grouped FP8 input capacity overflow")?;
let mut rank_inputs = Vec::with_capacity(self.ranks.len());
for engine in &self.ranks {
let _main = engine.gpu.enter_main()?;
rank_inputs.push(engine.uninit(input_capacity)?);
}
let mut owners = Vec::with_capacity(self.ranks.len());
for (owner_rank, rank) in experts.ranks.iter().enumerate() {
if rank.gate.expert_range != rank.up.expert_range
|| rank.gate.expert_range != rank.down.expert_range
{
return Err(format!(
"owner-grouped FP8 rank {} gate/up/down expert ranges differ",
owner_rank
)
.into());
}
let local_experts = rank.gate.expert_range.len();
let engine = &self.ranks[owner_rank];
let _main = engine.gpu.enter_main()?;
let route_csr =
DeviceExpertCsr::with_capacity(engine, local_experts, max_tokens, max_pairs)?;
let down_csr =
DeviceExpertCsr::with_capacity(engine, local_experts, max_pairs, max_pairs)?;
let gate_workspace = Fp8GroupedWorkspace::new(
engine,
experts.input_width,
experts.expert_width,
max_tokens,
max_pairs,
)?;
let up_workspace = Fp8GroupedWorkspace::new(
engine,
experts.input_width,
experts.expert_width,
max_tokens,
max_pairs,
)?;
let down_workspace = Fp8GroupedWorkspace::new(
engine,
experts.expert_width,
experts.input_width,
max_pairs,
max_pairs,
)?;
let activation = engine.uninit(
max_pairs
.checked_mul(experts.expert_width)
.ok_or("official Step owner-grouped FP8 activation capacity overflow")?,
)?;
owners.push(PreparedStepGroupedExpertOwner {
rank: owner_rank,
global_pairs: Vec::new(),
route_csr,
down_csr,
gate_workspace,
up_workspace,
down_workspace,
activation,
});
}
let mut plan = PreparedStepGroupedExpertParallelGate {
rank_inputs,
owners,
activation_limit,
tokens: 0,
pairs: 0,
max_tokens,
max_pairs,
input_width: experts.input_width,
expert_width: experts.expert_width,
generation: 0,
executed_generation: None,
ready: false,
};
self.refresh_step_grouped_expert_parallel_gate(
experts, &mut plan, input, tokens, selected,
)?;
Ok(plan)
}
fn prepare_step_grouped_expert_parallel_refresh(
&self,
experts: &ResidentExpertParallel,
plan: &PreparedStepGroupedExpertParallelGate,
tokens: usize,
selected: &[usize],
) -> Result<(usize, u64, Vec<Option<StepGroupedExpertOwnerSchedule>>), Box<dyn std::error::Error>>
{
validate_ep_residency(&self.ranks, experts)?;
if plan.rank_inputs.len() != self.ranks.len()
|| plan.owners.len() != self.ranks.len()
|| plan.input_width != experts.input_width
|| plan.expert_width != experts.expert_width
|| tokens > plan.max_tokens
{
return Err(format!(
"Step owner-grouped FP8 refresh geometry changed ranks={}/{} owners={}/{} \
input={}/{} expert={}/{} tokens={}/{}",
plan.rank_inputs.len(),
self.ranks.len(),
plan.owners.len(),
self.ranks.len(),
plan.input_width,
experts.input_width,
plan.expert_width,
experts.expert_width,
tokens,
plan.max_tokens,
)
.into());
}
let pairs = validate_step_grouped_owner_routes(experts.expert_count, tokens, selected)?;
if pairs > plan.max_pairs {
return Err(format!(
"Step owner-grouped FP8 route count {pairs} exceeds capacity {}",
plan.max_pairs
)
.into());
}
let next_generation = plan
.generation
.checked_add(1)
.ok_or("Step owner-grouped FP8 plan generation overflow")?;
let owner_routes = partition_expert_owner_routes(
experts.expert_count,
self.ranks.len(),
tokens,
STEP_GROUPED_FP8_TOP_K,
selected,
)?;
let mut schedules = Vec::with_capacity(self.ranks.len());
for routes in owner_routes {
if routes.selected.is_empty() {
schedules.push(None);
continue;
}
let local_experts = experts.ranks[routes.rank].gate.expert_range.len();
let local_pairs = routes.selected.len();
let route_csr = ExpertCsr::from_pair_rows(
local_experts,
tokens,
&routes.selected,
&routes.token_rows,
)?;
let down_rows = (0..local_pairs).collect::<Vec<_>>();
let down_csr = ExpertCsr::from_pair_rows(
local_experts,
local_pairs,
&routes.selected,
&down_rows,
)?;
schedules.push(Some(StepGroupedExpertOwnerSchedule {
global_pairs: routes.global_pairs,
route_csr,
down_csr,
}));
}
Ok((pairs, next_generation, schedules))
}
fn commit_step_grouped_expert_parallel_refresh(
&self,
plan: &mut PreparedStepGroupedExpertParallelGate,
tokens: usize,
pairs: usize,
next_generation: u64,
schedules: Vec<Option<StepGroupedExpertOwnerSchedule>>,
) -> Result<(), Box<dyn std::error::Error>> {
for (owner, schedule) in plan.owners.iter_mut().zip(schedules) {
let engine = &self.ranks[owner.rank];
let _main = engine.gpu.enter_main()?;
if let Some(schedule) = schedule {
owner.route_csr.refresh(engine, &schedule.route_csr)?;
owner.down_csr.refresh(engine, &schedule.down_csr)?;
owner.global_pairs = schedule.global_pairs;
} else {
owner.route_csr.clear();
owner.down_csr.clear();
owner.global_pairs.clear();
}
}
plan.tokens = tokens;
plan.pairs = pairs;
plan.generation = next_generation;
plan.ready = true;
Ok(())
}
pub fn refresh_step_grouped_expert_parallel_gate(
&self,
experts: &ResidentExpertParallel,
plan: &mut PreparedStepGroupedExpertParallelGate,
input: &[f32],
tokens: usize,
selected: &[usize],
) -> Result<(), Box<dyn std::error::Error>> {
validate_activations(input, tokens, experts.input_width)?;
let (pairs, next_generation, schedules) =
self.prepare_step_grouped_expert_parallel_refresh(experts, plan, tokens, selected)?;
plan.ready = false;
plan.executed_generation = None;
{
let root = &self.ranks[0];
let _main = root.gpu.enter_main()?;
let mut destination = plan.rank_inputs[0].slice_mut(0..input.len());
root.stream().memcpy_htod(input, &mut destination)?;
root.stream().synchronize()?;
}
let (root_inputs, peer_inputs) = plan.rank_inputs.split_at_mut(1);
let root_input = &root_inputs[0];
for (rank, peer_input) in peer_inputs.iter_mut().enumerate() {
let engine = &self.ranks[rank + 1];
let _main = engine.gpu.enter_main()?;
let mut destination = peer_input.slice_mut(0..input.len());
engine
.stream()
.memcpy_dtod(&root_input.slice(0..input.len()), &mut destination)?;
}
self.commit_step_grouped_expert_parallel_refresh(
plan,
tokens,
pairs,
next_generation,
schedules,
)
}
pub fn refresh_step_grouped_expert_parallel_gate_from_root_device(
&self,
experts: &ResidentExpertParallel,
plan: &mut PreparedStepGroupedExpertParallelGate,
input: &CudaSlice<f32>,
tokens: usize,
selected: &[usize],
) -> Result<(), Box<dyn std::error::Error>> {
let input_values = tokens
.checked_mul(experts.input_width)
.ok_or("Step owner-grouped FP8 input size overflow")?;
let root = self
.ranks
.first()
.ok_or("Step owner-grouped FP8 runtime has no root rank")?;
if input.len() < input_values || input.ordinal() != root.ctx().ordinal() {
return Err(format!(
"Step owner-grouped FP8 root input len/device {}/{} does not cover {} values on \
device {}",
input.len(),
input.ordinal(),
input_values,
root.ctx().ordinal(),
)
.into());
}
let (pairs, next_generation, schedules) =
self.prepare_step_grouped_expert_parallel_refresh(experts, plan, tokens, selected)?;
plan.ready = false;
plan.executed_generation = None;
{
let _main = root.gpu.enter_main()?;
let mut destination = plan.rank_inputs[0].slice_mut(0..input_values);
root.stream()
.memcpy_dtod(&input.slice(0..input_values), &mut destination)?;
root.stream().synchronize()?;
}
let (root_inputs, peer_inputs) = plan.rank_inputs.split_at_mut(1);
let root_input = &root_inputs[0];
for (rank, peer_input) in peer_inputs.iter_mut().enumerate() {
let engine = &self.ranks[rank + 1];
let _main = engine.gpu.enter_main()?;
let mut destination = peer_input.slice_mut(0..input_values);
engine
.stream()
.memcpy_dtod(&root_input.slice(0..input_values), &mut destination)?;
}
self.commit_step_grouped_expert_parallel_refresh(
plan,
tokens,
pairs,
next_generation,
schedules,
)
}
pub fn refresh_step_grouped_expert_parallel_inputs_from_replicated(
&self,
experts: &ResidentExpertParallel,
plan: &mut PreparedStepGroupedExpertParallelGate,
input: &ResidentReplicatedDeviceRows,
) -> Result<(), Box<dyn std::error::Error>> {
validate_ep_residency(&self.ranks, experts)?;
validate_replicated_device_rows(&self.ranks, input)?;
if !plan.ready
|| input.tokens != plan.tokens
|| input.width != plan.input_width
|| input.tokens > plan.max_tokens
|| plan.rank_inputs.len() != self.ranks.len()
|| plan.owners.len() != self.ranks.len()
|| plan.input_width != experts.input_width
|| plan.expert_width != experts.expert_width
{
return Err("Step owner-grouped replicated input geometry changed".into());
}
let values = input
.tokens
.checked_mul(input.width)
.ok_or("Step owner-grouped replicated input size overflow")?;
let next_generation = plan
.generation
.checked_add(1)
.ok_or("Step owner-grouped FP8 plan generation overflow")?;
plan.ready = false;
plan.executed_generation = None;
for (rank, engine) in self.ranks.iter().enumerate() {
let _main = engine.gpu.enter_main()?;
let mut destination = plan.rank_inputs[rank].slice_mut(0..values);
engine
.stream()
.memcpy_dtod(&input.ranks[rank], &mut destination)?;
}
plan.generation = next_generation;
plan.ready = true;
Ok(())
}
pub fn execute_step_grouped_expert_parallel_gate(
&self,
experts: &ResidentExpertParallel,
plan: &mut PreparedStepGroupedExpertParallelGate,
) -> Result<(), Box<dyn std::error::Error>> {
validate_ep_residency(&self.ranks, experts)?;
if !plan.ready
|| plan.rank_inputs.len() != self.ranks.len()
|| plan.owners.len() != self.ranks.len()
|| plan.input_width != experts.input_width
|| plan.expert_width != experts.expert_width
{
return Err("Step owner-grouped FP8 plan is not ready or its geometry changed".into());
}
plan.executed_generation = None;
for owner in &mut plan.owners {
if owner.global_pairs.is_empty() {
continue;
}
let engine = &self.ranks[owner.rank];
let bank = &experts.ranks[owner.rank];
let _main = engine.gpu.enter_main()?;
let local_pairs = owner.global_pairs.len();
owner.gate_workspace.quantize_for_shape(
engine,
&plan.rank_inputs[owner.rank],
plan.tokens,
local_pairs,
)?;
owner.gate_workspace.project(
engine,
&bank.gate.codes,
&bank.gate.scales,
&owner.route_csr,
bank.gate.code_stride,
bank.gate.scale_stride,
1.0,
)?;
owner.up_workspace.quantize_for_shape(
engine,
&plan.rank_inputs[owner.rank],
plan.tokens,
local_pairs,
)?;
owner.up_workspace.project(
engine,
&bank.up.codes,
&bank.up.scales,
&owner.route_csr,
bank.up.code_stride,
bank.up.scale_stride,
1.0,
)?;
}
for owner in &mut plan.owners {
if owner.global_pairs.is_empty() {
continue;
}
let engine = &self.ranks[owner.rank];
let _main = engine.gpu.enter_main()?;
let values = owner.global_pairs.len() * plan.expert_width;
if let Some(limit) = plan.activation_limit {
engine.silu_clamped_mul_host_expf(
owner.gate_workspace.output(),
owner.up_workspace.output(),
limit,
&mut owner.activation,
values,
)?;
} else {
engine.silu_mul_host_expf(
owner.gate_workspace.output(),
owner.up_workspace.output(),
&mut owner.activation,
values,
)?;
}
}
for owner in &mut plan.owners {
if owner.global_pairs.is_empty() {
continue;
}
let engine = &self.ranks[owner.rank];
let bank = &experts.ranks[owner.rank];
let _main = engine.gpu.enter_main()?;
let local_pairs = owner.global_pairs.len();
owner.down_workspace.quantize_for_shape(
engine,
&owner.activation,
local_pairs,
local_pairs,
)?;
owner.down_workspace.project(
engine,
&bank.down.codes,
&bank.down.scales,
&owner.down_csr,
bank.down.code_stride,
bank.down.scale_stride,
1.0,
)?;
}
plan.executed_generation = Some(plan.generation);
Ok(())
}
pub fn collect_step_grouped_expert_parallel_gate(
&self,
plan: &PreparedStepGroupedExpertParallelGate,
) -> Result<StepGroupedFp8ProjectionOutput, Box<dyn std::error::Error>> {
if !plan.ready || plan.executed_generation != Some(plan.generation) {
return Err("Step owner-grouped FP8 projection is stale or has not executed".into());
}
let mut gate = vec![0.0f32; plan.pairs * plan.expert_width];
let mut up = vec![0.0f32; plan.pairs * plan.expert_width];
let mut down = vec![0.0f32; plan.pairs * plan.input_width];
for owner in &plan.owners {
if owner.global_pairs.is_empty() {
continue;
}
let engine = &self.ranks[owner.rank];
let _main = engine.gpu.enter_main()?;
let owner_gate = engine.dtoh_view(
&owner
.gate_workspace
.output()
.slice(0..owner.gate_workspace.output_len()),
)?;
let owner_up = engine.dtoh_view(
&owner
.up_workspace
.output()
.slice(0..owner.up_workspace.output_len()),
)?;
let owner_down = engine.dtoh_view(
&owner
.down_workspace
.output()
.slice(0..owner.down_workspace.output_len()),
)?;
for (local_pair, &global_pair) in owner.global_pairs.iter().enumerate() {
let local_expert = local_pair * plan.expert_width;
let global_expert = global_pair * plan.expert_width;
gate[global_expert..global_expert + plan.expert_width]
.copy_from_slice(&owner_gate[local_expert..local_expert + plan.expert_width]);
up[global_expert..global_expert + plan.expert_width]
.copy_from_slice(&owner_up[local_expert..local_expert + plan.expert_width]);
let local_hidden = local_pair * plan.input_width;
let global_hidden = global_pair * plan.input_width;
down[global_hidden..global_hidden + plan.input_width]
.copy_from_slice(&owner_down[local_hidden..local_hidden + plan.input_width]);
}
}
Ok(StepGroupedFp8ProjectionOutput { gate, up, down })
}
pub fn run_step_grouped_expert_parallel_gate(
&self,
experts: &ResidentExpertParallel,
plan: &mut PreparedStepGroupedExpertParallelGate,
) -> Result<StepGroupedFp8ProjectionOutput, Box<dyn std::error::Error>> {
self.execute_step_grouped_expert_parallel_gate(experts, plan)?;
self.collect_step_grouped_expert_parallel_gate(plan)
}
pub fn prepare_step_grouped_expert_parallel_combine(
&self,
plan: &PreparedStepGroupedExpertParallelGate,
route_weights: &[f32],
) -> Result<PreparedPeerWeightedRouteCombine, Box<dyn std::error::Error>> {
if !self.native_p2p || !self.ep_device_arithmetic || !plan.ready {
return Err(
"Step owner-grouped combine requires a ready native-P2P device plan".into(),
);
}
let owner_pairs = plan
.owners
.iter()
.map(|owner| owner.global_pairs.as_slice())
.collect::<Vec<_>>();
let shape = validate_weighted_route_combine(
plan.input_width,
STEP_GROUPED_FP8_TOP_K,
plan.max_tokens,
plan.tokens,
&owner_pairs,
route_weights,
)?;
if shape.max_pairs != plan.max_pairs {
return Err(format!(
"Step owner-grouped combine capacity {} != projection capacity {}",
shape.max_pairs, plan.max_pairs
)
.into());
}
let root = self
.ranks
.first()
.ok_or("Step owner-grouped combine has no root rank")?;
let slot_values = shape
.max_pairs
.checked_mul(plan.input_width)
.ok_or("Step owner-grouped combine slot capacity overflow")?;
let output_values = plan
.max_tokens
.checked_mul(plan.input_width)
.ok_or("Step owner-grouped combine output capacity overflow")?;
let (root_device, owners, peer_staging, slots, weights, output) = {
let _main = root.gpu.enter_main()?;
let mut owners = Vec::with_capacity(plan.owners.len());
for _ in &plan.owners {
owners.push(PreparedPeerWeightedRouteOwner {
token_rows: root.htod_i32(&vec![0; shape.max_pairs])?,
slots: root.htod_i32(&vec![0; shape.max_pairs])?,
weights: root.htod(&vec![0.0; shape.max_pairs])?,
active_pairs: 0,
});
}
(
root.ctx().ordinal(),
owners,
root.uninit(slot_values)?,
root.uninit(slot_values)?,
root.uninit(shape.max_pairs)?,
root.uninit(output_values)?,
)
};
let mut peer_devices = Vec::with_capacity(self.ranks.len().saturating_sub(1));
let mut peer_outputs = Vec::with_capacity(self.ranks.len().saturating_sub(1));
for engine in self.ranks.iter().skip(1) {
let _main = engine.gpu.enter_main()?;
peer_devices.push(engine.ctx().ordinal());
peer_outputs.push(engine.uninit(output_values)?);
}
let mut combine = PreparedPeerWeightedRouteCombine {
root_device,
owners,
peer_staging,
slots,
weights,
output,
peer_devices,
peer_outputs,
width: plan.input_width,
experts_per_token: STEP_GROUPED_FP8_TOP_K,
max_tokens: plan.max_tokens,
max_pairs: shape.max_pairs,
tokens: 0,
pairs: 0,
projection_generation: 0,
output_generation: None,
broadcast_generation: None,
ready: false,
};
self.refresh_step_grouped_expert_parallel_combine(plan, &mut combine, route_weights)?;
Ok(combine)
}
pub fn refresh_step_grouped_expert_parallel_combine(
&self,
plan: &PreparedStepGroupedExpertParallelGate,
combine: &mut PreparedPeerWeightedRouteCombine,
route_weights: &[f32],
) -> Result<(), Box<dyn std::error::Error>> {
let output_capacity = combine
.max_tokens
.checked_mul(combine.width)
.ok_or("Step owner-grouped combine output capacity overflow")?;
if !plan.ready
|| combine.owners.len() != plan.owners.len()
|| combine.peer_devices.len() + 1 != self.ranks.len()
|| combine.peer_outputs.len() + 1 != self.ranks.len()
|| combine.width != plan.input_width
|| combine.experts_per_token != STEP_GROUPED_FP8_TOP_K
|| combine.max_tokens != plan.max_tokens
|| combine.max_pairs != plan.max_pairs
|| combine.output.len() < output_capacity
|| combine
.peer_outputs
.iter()
.any(|output| output.len() < output_capacity)
{
return Err("Step owner-grouped combine/projection geometry changed".into());
}
if self
.ranks
.iter()
.skip(1)
.zip(&combine.peer_devices)
.any(|(engine, &device)| engine.ctx().ordinal() != device)
{
return Err("Step owner-grouped combine peer devices changed".into());
}
let owner_pairs = plan
.owners
.iter()
.map(|owner| owner.global_pairs.as_slice())
.collect::<Vec<_>>();
let shape = validate_weighted_route_combine(
combine.width,
combine.experts_per_token,
combine.max_tokens,
plan.tokens,
&owner_pairs,
route_weights,
)?;
if shape.max_pairs != combine.max_pairs {
return Err("Step owner-grouped combine capacity changed during refresh".into());
}
let metadata = owner_pairs
.iter()
.map(|pairs| {
let token_rows = pairs
.iter()
.map(|&pair| (pair / combine.experts_per_token) as i32)
.collect::<Vec<_>>();
let slots = pairs
.iter()
.map(|&pair| (pair % combine.experts_per_token) as i32)
.collect::<Vec<_>>();
let weights = pairs
.iter()
.map(|&pair| route_weights[pair])
.collect::<Vec<_>>();
(token_rows, slots, weights)
})
.collect::<Vec<_>>();
combine.ready = false;
combine.output_generation = None;
combine.broadcast_generation = None;
let root = self
.ranks
.first()
.ok_or("Step owner-grouped combine has no root rank")?;
let _main = root.gpu.enter_main()?;
if root.ctx().ordinal() != combine.root_device {
return Err(format!(
"Step owner-grouped combine root device changed {} != {}",
root.ctx().ordinal(),
combine.root_device
)
.into());
}
for (owner, (token_rows, slots, weights)) in combine.owners.iter_mut().zip(metadata) {
if token_rows.is_empty() {
owner.active_pairs = 0;
continue;
}
root.htod_i32_into(&mut owner.token_rows, &token_rows)?;
root.htod_i32_into(&mut owner.slots, &slots)?;
let mut weight_prefix = owner.weights.slice_mut(0..weights.len());
root.stream().memcpy_htod(&weights, &mut weight_prefix)?;
owner.active_pairs = token_rows.len();
}
combine.tokens = plan.tokens;
combine.pairs = shape.pairs;
combine.projection_generation = plan.generation;
combine.ready = true;
Ok(())
}
pub fn execute_step_grouped_expert_parallel_combine(
&self,
plan: &PreparedStepGroupedExpertParallelGate,
combine: &mut PreparedPeerWeightedRouteCombine,
) -> Result<(), Box<dyn std::error::Error>> {
if !plan.ready
|| plan.executed_generation != Some(plan.generation)
|| !combine.ready
|| combine.tokens != plan.tokens
|| combine.pairs != plan.pairs
|| combine.width != plan.input_width
|| combine.owners.len() != plan.owners.len()
|| combine.projection_generation != plan.generation
{
return Err("Step owner-grouped combine is stale or its geometry changed".into());
}
combine.output_generation = None;
combine.broadcast_generation = None;
for owner in &plan.owners {
if owner.rank == 0 || owner.global_pairs.is_empty() {
continue;
}
let engine = &self.ranks[owner.rank];
let _main = engine.gpu.enter_main()?;
engine.stream().synchronize()?;
}
let root = self
.ranks
.first()
.ok_or("Step owner-grouped combine has no root rank")?;
let _main = root.gpu.enter_main()?;
if root.ctx().ordinal() != combine.root_device {
return Err("Step owner-grouped combine is not resident on the root device".into());
}
for (index, owner) in plan.owners.iter().enumerate() {
let metadata = &combine.owners[index];
if owner.global_pairs.len() != metadata.active_pairs {
return Err(format!(
"Step owner-grouped combine owner {index} rows {} != metadata {}",
owner.global_pairs.len(),
metadata.active_pairs
)
.into());
}
if metadata.active_pairs == 0 {
continue;
}
let values = metadata
.active_pairs
.checked_mul(combine.width)
.ok_or("Step owner-grouped combine peer value count overflow")?;
if owner.rank == 0 {
root.scatter_slot(
owner.down_workspace.output(),
&metadata.token_rows,
&metadata.slots,
&metadata.weights,
&mut combine.slots,
&mut combine.weights,
combine.width,
combine.experts_per_token,
metadata.active_pairs,
)?;
} else {
let source = owner.down_workspace.output().slice(0..values);
let mut destination = combine.peer_staging.slice_mut(0..values);
root.stream().memcpy_dtod(&source, &mut destination)?;
root.scatter_slot(
&combine.peer_staging,
&metadata.token_rows,
&metadata.slots,
&metadata.weights,
&mut combine.slots,
&mut combine.weights,
combine.width,
combine.experts_per_token,
metadata.active_pairs,
)?;
}
}
root.reduce_slots_host(
&combine.slots,
&combine.weights,
&mut combine.output,
combine.width,
combine.experts_per_token,
combine.tokens,
)?;
combine.output_generation = Some(plan.generation);
Ok(())
}
pub fn collect_step_grouped_expert_parallel_combine(
&self,
plan: &PreparedStepGroupedExpertParallelGate,
combine: &PreparedPeerWeightedRouteCombine,
) -> Result<Vec<f32>, Box<dyn std::error::Error>> {
if !plan.ready
|| combine.output_generation != Some(plan.generation)
|| combine.projection_generation != plan.generation
{
return Err("Step owner-grouped combine output is stale or has not executed".into());
}
let root = self
.ranks
.first()
.ok_or("Step owner-grouped combine has no root rank")?;
let _main = root.gpu.enter_main()?;
if root.ctx().ordinal() != combine.root_device {
return Err("Step owner-grouped combine is not resident on the root device".into());
}
root.dtoh_view(&combine.output.slice(0..combine.tokens * combine.width))
}
pub fn copy_step_grouped_expert_parallel_combine_root(
&self,
plan: &PreparedStepGroupedExpertParallelGate,
combine: &PreparedPeerWeightedRouteCombine,
destination: &Engine,
) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
if !plan.ready
|| combine.output_generation != Some(plan.generation)
|| combine.projection_generation != plan.generation
{
return Err("Step owner-grouped combine output is stale or has not executed".into());
}
let root = self
.ranks
.first()
.ok_or("Step owner-grouped combine has no root rank")?;
if root.ctx().ordinal() != combine.root_device
|| destination.ctx().ordinal() != combine.root_device
{
return Err(format!(
"Step owner-grouped combine root/destination devices {}/{} != {}",
root.ctx().ordinal(),
destination.ctx().ordinal(),
combine.root_device,
)
.into());
}
let values = combine
.tokens
.checked_mul(combine.width)
.ok_or("Step owner-grouped combine copy size overflow")?;
{
let _main = root.gpu.enter_main()?;
root.stream().synchronize()?;
}
let _main = destination.gpu.enter_main()?;
let mut output = destination.uninit(values)?;
destination
.stream()
.memcpy_dtod(&combine.output.slice(0..values), &mut output)?;
Ok(output)
}
pub fn broadcast_step_grouped_expert_parallel_combine(
&self,
plan: &PreparedStepGroupedExpertParallelGate,
combine: &mut PreparedPeerWeightedRouteCombine,
) -> Result<(), Box<dyn std::error::Error>> {
if !plan.ready
|| combine.output_generation != Some(plan.generation)
|| combine.projection_generation != plan.generation
|| combine.peer_devices.len() + 1 != self.ranks.len()
|| combine.peer_outputs.len() + 1 != self.ranks.len()
{
return Err("Step owner-grouped combine output cannot be broadcast".into());
}
combine.broadcast_generation = None;
let values = combine
.tokens
.checked_mul(combine.width)
.ok_or("Step owner-grouped combine broadcast size overflow")?;
{
let root = self
.ranks
.first()
.ok_or("Step owner-grouped combine has no root rank")?;
let _main = root.gpu.enter_main()?;
if root.ctx().ordinal() != combine.root_device {
return Err("Step owner-grouped combine root device changed".into());
}
root.stream().synchronize()?;
}
let source = &combine.output;
for (index, destination_buffer) in combine.peer_outputs.iter_mut().enumerate() {
let engine = &self.ranks[index + 1];
let _main = engine.gpu.enter_main()?;
if engine.ctx().ordinal() != combine.peer_devices[index] {
return Err(format!(
"Step owner-grouped combine peer {} device changed",
index + 1
)
.into());
}
let mut destination = destination_buffer.slice_mut(0..values);
engine
.stream()
.memcpy_dtod(&source.slice(0..values), &mut destination)?;
}
combine.broadcast_generation = Some(plan.generation);
Ok(())
}
pub fn collect_step_grouped_expert_parallel_broadcast(
&self,
plan: &PreparedStepGroupedExpertParallelGate,
combine: &PreparedPeerWeightedRouteCombine,
) -> Result<Vec<Vec<f32>>, Box<dyn std::error::Error>> {
if !plan.ready
|| combine.output_generation != Some(plan.generation)
|| combine.broadcast_generation != Some(plan.generation)
|| combine.peer_outputs.len() + 1 != self.ranks.len()
{
return Err("Step owner-grouped combine broadcast is stale or incomplete".into());
}
let values = combine
.tokens
.checked_mul(combine.width)
.ok_or("Step owner-grouped combine collection size overflow")?;
let mut outputs = Vec::with_capacity(self.ranks.len());
{
let root = &self.ranks[0];
let _main = root.gpu.enter_main()?;
outputs.push(root.dtoh_view(&combine.output.slice(0..values))?);
}
for (index, output) in combine.peer_outputs.iter().enumerate() {
let engine = &self.ranks[index + 1];
let _main = engine.gpu.enter_main()?;
outputs.push(engine.dtoh_view(&output.slice(0..values))?);
}
Ok(outputs)
}
pub fn finish_step_grouped_expert_parallel_layer(
&self,
plan: &PreparedStepGroupedExpertParallelGate,
combine: &PreparedPeerWeightedRouteCombine,
shared: &ResidentReplicatedDeviceRows,
residual: &ResidentReplicatedDeviceRows,
) -> Result<ResidentReplicatedDeviceRows, Box<dyn std::error::Error>> {
validate_replicated_device_rows(&self.ranks, shared)?;
validate_replicated_device_rows(&self.ranks, residual)?;
if !plan.ready
|| plan.executed_generation != Some(plan.generation)
|| combine.output_generation != Some(plan.generation)
|| combine.broadcast_generation != Some(plan.generation)
|| combine.projection_generation != plan.generation
|| combine.peer_outputs.len() + 1 != self.ranks.len()
|| shared.tokens != combine.tokens
|| residual.tokens != combine.tokens
|| shared.width != combine.width
|| residual.width != combine.width
{
return Err("Step full-layer finish inputs are stale or their geometry changed".into());
}
let values = combine
.tokens
.checked_mul(combine.width)
.ok_or("Step full-layer output size overflow")?;
let mut ranks = Vec::with_capacity(self.ranks.len());
for rank in 0..self.ranks.len() {
let engine = &self.ranks[rank];
let _main = engine.gpu.enter_main()?;
let routed = if rank == 0 {
&combine.output
} else {
&combine.peer_outputs[rank - 1]
};
let mut ffn = engine.uninit(values)?;
engine.add(routed, &shared.ranks[rank], &mut ffn, values)?;
let mut output = engine.uninit(values)?;
engine.add(&residual.ranks[rank], &ffn, &mut output, values)?;
ranks.push(output);
}
Ok(ResidentReplicatedDeviceRows {
ranks,
tokens: combine.tokens,
width: combine.width,
})
}
pub fn run_step_grouped_expert_parallel_combine(
&self,
plan: &PreparedStepGroupedExpertParallelGate,
combine: &mut PreparedPeerWeightedRouteCombine,
) -> Result<Vec<f32>, Box<dyn std::error::Error>> {
self.execute_step_grouped_expert_parallel_combine(plan, combine)?;
self.collect_step_grouped_expert_parallel_combine(plan, combine)
}
pub fn upload_tensor_parallel(
&self,
gate: E4m3ExpertBank<'_>,
up: E4m3ExpertBank<'_>,
down: E4m3ExpertBank<'_>,
) -> Result<ResidentTensorParallel, Box<dyn std::error::Error>> {
gate.validate()?;
up.validate()?;
down.validate()?;
if gate.expert_count != up.expert_count || gate.expert_count != down.expert_count {
return Err("TP gate/up/down expert counts differ".into());
}
if gate.in_features != up.in_features || gate.out_features != up.out_features {
return Err("TP gate/up dimensions differ".into());
}
if down.in_features != gate.out_features || down.out_features != gate.in_features {
return Err(format!(
"TP down {}x{} does not invert gate/up {}x{}",
down.out_features, down.in_features, gate.out_features, gate.in_features
)
.into());
}
let tp = self.ranks.len();
validate_column_bank_shape(gate, tp)?;
validate_column_bank_shape(up, tp)?;
validate_row_bank_shape(down, tp)?;
let mut gate_ranks = Vec::with_capacity(tp);
let mut up_ranks = Vec::with_capacity(tp);
let mut down_ranks = Vec::with_capacity(tp);
for (rank, engine) in self.ranks.iter().enumerate() {
gate_ranks.push(upload_column_bank_rank(engine, gate, tp, rank)?);
up_ranks.push(upload_column_bank_rank(engine, up, tp, rank)?);
down_ranks.push(upload_row_bank_rank(engine, down, tp, rank)?);
}
Ok(ResidentTensorParallel {
bank: ResidentTpExpertBank {
gate: gate_ranks,
up: up_ranks,
down: down_ranks,
expert_count: gate.expert_count,
input_width: gate.in_features,
expert_width: gate.out_features,
},
})
}
pub fn run_tensor_parallel_routes(
&self,
experts: &ResidentTensorParallel,
input: &[f32],
tokens: usize,
selected: &[usize],
route_weights: &[f32],
experts_per_token: usize,
) -> Result<Vec<f32>, Box<dyn std::error::Error>> {
validate_tp_bank_residency(&self.ranks, &experts.bank)?;
validate_activations(input, tokens, experts.bank.input_width)?;
let pairs = tokens
.checked_mul(experts_per_token)
.ok_or("TP route count overflow")?;
if selected.len() != pairs || route_weights.len() != pairs {
return Err(format!(
"TP routes selected={} weights={} != tokens {tokens} x experts/token \
{experts_per_token} ({pairs})",
selected.len(),
route_weights.len(),
)
.into());
}
if !route_weights.iter().all(|weight| weight.is_finite()) {
return Err("TP route weights contain a non-finite value".into());
}
let mut output = vec![0.0f32; tokens * experts.bank.input_width];
for token in 0..tokens {
let input_row =
&input[token * experts.bank.input_width..(token + 1) * experts.bank.input_width];
for slot in 0..experts_per_token {
let pair = token * experts_per_token + slot;
let expert = selected[pair];
if expert >= experts.bank.expert_count {
return Err(format!(
"TP selected expert {expert} outside 0..{}",
experts.bank.expert_count
)
.into());
}
let down = if self.native_p2p {
self.run_tensor_parallel_expert_native(&experts.bank, expert, input_row)?
} else {
let gate =
self.run_column_bank_expert(&experts.bank.gate, expert, input_row)?;
let up = self.run_column_bank_expert(&experts.bank.up, expert, input_row)?;
let activated: Vec<f32> = gate
.iter()
.zip(&up)
.map(|(&gate, &up)| gate / (1.0 + (-gate).exp()) * up)
.collect();
debug_assert_eq!(activated.len(), experts.bank.expert_width);
self.run_row_bank_expert(&experts.bank.down, expert, &activated)?
};
let weight = route_weights[pair];
for (sum, value) in output
[token * experts.bank.input_width..(token + 1) * experts.bank.input_width]
.iter_mut()
.zip(down)
{
*sum += weight * value;
}
}
}
Ok(output)
}
fn run_column_bank_expert(
&self,
ranks: &[ResidentE4m3ExpertBankRank],
expert: usize,
input: &[f32],
) -> Result<Vec<f32>, Box<dyn std::error::Error>> {
let local_out = ranks
.first()
.ok_or("TP column bank has no ranks")?
.out_features;
let mut gathered = vec![0.0f32; local_out * ranks.len()];
for (rank, (engine, bank)) in self.ranks.iter().zip(ranks).enumerate() {
let shard = run_resident_bank_expert(engine, bank, expert, input, 1)?;
gathered[rank * local_out..(rank + 1) * local_out].copy_from_slice(&shard);
}
Ok(gathered)
}
fn run_row_bank_expert(
&self,
ranks: &[ResidentE4m3ExpertBankRank],
expert: usize,
input: &[f32],
) -> Result<Vec<f32>, Box<dyn std::error::Error>> {
let local_in = ranks.first().ok_or("TP row bank has no ranks")?.in_features;
if input.len() != local_in * ranks.len() {
return Err(format!(
"TP row input {} != {} ranks x {local_in}",
input.len(),
ranks.len()
)
.into());
}
let out_features = ranks[0].out_features;
let mut reduced = vec![0.0f32; out_features];
for (rank, (engine, bank)) in self.ranks.iter().zip(ranks).enumerate() {
let blocks = bank
.k_blocks
.ok_or("TP row bank is not packed in native K-block order")?;
if blocks * FP8_BLOCK != local_in {
return Err(format!(
"TP row bank has {blocks} blocks but local input width is {local_in}"
)
.into());
}
for block in 0..blocks {
let global_start = rank * local_in + block * FP8_BLOCK;
let partial = run_resident_bank_expert_block(
engine,
bank,
expert,
block,
&input[global_start..global_start + FP8_BLOCK],
)?;
for (sum, value) in reduced.iter_mut().zip(partial) {
*sum += value;
}
}
}
Ok(reduced)
}
fn run_tensor_parallel_expert_native(
&self,
bank: &ResidentTpExpertBank,
expert: usize,
input: &[f32],
) -> Result<Vec<f32>, Box<dyn std::error::Error>> {
if !self.native_p2p || self.ranks.len() < 2 {
return Err("native TP expert execution requires at least two P2P ranks".into());
}
let local_out = bank
.gate
.first()
.ok_or("native TP gate bank has no ranks")?
.out_features;
if local_out * self.ranks.len() != bank.expert_width {
return Err(format!(
"native TP gate shards {}x{local_out} != expert width {}",
self.ranks.len(),
bank.expert_width
)
.into());
}
let mut rank_inputs = Vec::with_capacity(self.ranks.len());
let root_input = {
let root = &self.ranks[0];
let _main = root.gpu.enter_main()?;
root.htod(input)?
};
rank_inputs.push(root_input);
for engine in &self.ranks[1..] {
let peer_input = {
let _main = engine.gpu.enter_main()?;
let mut peer_input = engine.uninit(input.len())?;
engine
.stream()
.memcpy_dtod(&rank_inputs[0], &mut peer_input)?;
peer_input
};
rank_inputs.push(peer_input);
}
let mut gate_shards = Vec::with_capacity(self.ranks.len());
let mut up_shards = Vec::with_capacity(self.ranks.len());
for rank in 0..self.ranks.len() {
gate_shards.push(run_resident_bank_expert_device(
&self.ranks[rank],
&bank.gate[rank],
expert,
&rank_inputs[rank],
1,
)?);
up_shards.push(run_resident_bank_expert_device(
&self.ranks[rank],
&bank.up[rank],
expert,
&rank_inputs[rank],
1,
)?);
}
let gate = self.gather_native_column_shards(&gate_shards, 1, local_out)?;
let up = self.gather_native_column_shards(&up_shards, 1, local_out)?;
let activated = gate
.iter()
.zip(&up)
.map(|(&gate, &up)| gate / (1.0 + (-gate).exp()) * up)
.collect::<Vec<_>>();
debug_assert_eq!(activated.len(), bank.expert_width);
let root_activated = {
let root = &self.ranks[0];
let _main = root.gpu.enter_main()?;
root.htod(&activated)?
};
let mut rank_activated = Vec::with_capacity(self.ranks.len());
for (rank, engine) in self.ranks.iter().enumerate() {
let start = rank * local_out;
let source = root_activated.slice(start..start + local_out);
let local = {
let _main = engine.gpu.enter_main()?;
let mut local = engine.uninit(local_out)?;
engine.stream().memcpy_dtod(&source, &mut local)?;
local
};
rank_activated.push(local);
}
let out_features = bank
.down
.first()
.ok_or("native TP down bank has no ranks")?
.out_features;
let mut reduced = {
let root = &self.ranks[0];
let _main = root.gpu.enter_main()?;
root.htod(&vec![0.0f32; out_features])?
};
let mut remote_partial_keepalive = Vec::new();
for rank in 0..self.ranks.len() {
let down = &bank.down[rank];
let blocks = down
.k_blocks
.ok_or("native TP row bank is not packed in checkpoint-block order")?;
if blocks * FP8_BLOCK != local_out {
return Err(format!(
"native TP rank {rank} has {blocks} blocks but local activation width is \
{local_out}"
)
.into());
}
for block in 0..blocks {
let start = block * FP8_BLOCK;
let input_block = rank_activated[rank].slice(start..start + FP8_BLOCK);
let partial = run_resident_bank_expert_block_device(
&self.ranks[rank],
down,
expert,
block,
&input_block,
)?;
let root_partial = if rank == 0 {
partial
} else {
let root = &self.ranks[0];
let _main = root.gpu.enter_main()?;
let mut peer_partial = root.uninit(out_features)?;
root.stream().memcpy_dtod(&partial, &mut peer_partial)?;
remote_partial_keepalive.push(partial);
peer_partial
};
let next = {
let root = &self.ranks[0];
let _main = root.gpu.enter_main()?;
let mut next = root.uninit(out_features)?;
root.add(&reduced, &root_partial, &mut next, out_features)?;
next
};
reduced = next;
}
}
let output = {
let root = &self.ranks[0];
let _main = root.gpu.enter_main()?;
root.dtoh(&reduced)?
};
drop(remote_partial_keepalive);
Ok(output)
}
pub fn gather_native_column_shards_device(
&self,
shards: &[CudaSlice<f32>],
tokens: usize,
local_out: usize,
) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
let shard_len = tokens
.checked_mul(local_out)
.ok_or("native TP gather shard size overflow")?;
if shards.len() != self.ranks.len() || shards.iter().any(|shard| shard.len() != shard_len) {
return Err("native TP gather shard geometry mismatch".into());
}
for engine in &self.ranks[1..] {
let _main = engine.gpu.enter_main()?;
engine.stream().synchronize()?;
}
let root = &self.ranks[0];
let _main = root.gpu.enter_main()?;
let global_out = shards
.len()
.checked_mul(local_out)
.ok_or("native TP gather output width overflow")?;
let gathered_len = tokens
.checked_mul(global_out)
.ok_or("native TP gather output size overflow")?;
let mut gathered = root.uninit(gathered_len)?;
if self.bulk_p2p {
root.place_rows_strided(&shards[0], &mut gathered, local_out, tokens, global_out, 0)?;
if shards.len() > 1 {
let mut staging = root.uninit(shard_len)?;
for (rank, shard) in shards.iter().enumerate().skip(1) {
root.stream().memcpy_dtod(shard, &mut staging)?;
root.place_rows_strided(
&staging,
&mut gathered,
local_out,
tokens,
global_out,
rank * local_out,
)?;
}
}
} else {
for token in 0..tokens {
for (rank, shard) in shards.iter().enumerate() {
let source = shard.slice(token * local_out..(token + 1) * local_out);
let start = token * global_out + rank * local_out;
let mut destination = gathered.slice_mut(start..start + local_out);
root.stream().memcpy_dtod(&source, &mut destination)?;
}
}
}
Ok(gathered)
}
pub fn gather_native_column_shards(
&self,
shards: &[CudaSlice<f32>],
tokens: usize,
local_out: usize,
) -> Result<Vec<f32>, Box<dyn std::error::Error>> {
let gathered = self.gather_native_column_shards_device(shards, tokens, local_out)?;
let root = &self.ranks[0];
let _main = root.gpu.enter_main()?;
root.dtoh(&gathered)
}
pub(crate) fn decode_v2_workspace(&self) -> &std::sync::Mutex<Vec<StepTpDecodeV2Ws>> {
&self.decode_v2
}
pub(crate) fn decode_v2_ensure(
&self,
e: &Engine,
q_m: &ResidentBf16ColumnParallel,
k_m: &ResidentBf16ColumnParallel,
v_m: &ResidentBf16ColumnParallel,
o_m: &ResidentStepBf16RowParallel,
heads: usize,
) -> Result<usize, Box<dyn std::error::Error>> {
if self.ranks.len() > 1 && !self.native_p2p {
return Err("step TP decode v2 requires native P2P ranks".into());
}
let ranks = self.ranks.len();
let fused_door = step_tp_qkv_fused_enabled()?;
let arm_ok = |weight: &ResidentBf16Weight| match weight {
ResidentBf16Weight::F32(_) => true,
ResidentBf16Weight::Bf16(_) => fused_door,
};
for matrix in [q_m, k_m, v_m] {
validate_resident_bf16_ranks(&self.ranks, &matrix.ranks)?;
if matrix.out_features % ranks != 0 || matrix.in_features != q_m.in_features {
return Err("step TP decode v2 QKV geometry mismatch".into());
}
for rank in &matrix.ranks {
if !arm_ok(&rank.weight) {
return Err("step TP decode v2 requires MEMRA_STEP_TP_F32_MIRROR=1 or \
MEMRA_STEP_TP_QKV_FUSED=1 (bf16-resident fused kernels)"
.into());
}
}
}
validate_step_bf16_row_residency(&self.ranks, o_m)?;
for blocks in &o_m.ranks {
for block in blocks {
if !arm_ok(&block.weight) {
return Err("step TP decode v2 requires MEMRA_STEP_TP_F32_MIRROR=1 or \
MEMRA_STEP_TP_QKV_FUSED=1 (bf16-resident fused kernels)"
.into());
}
}
}
if v_m.out_features != k_m.out_features
|| o_m.in_features != q_m.out_features
|| heads == 0
|| heads % ranks != 0
{
return Err("step TP decode v2 K/V/O geometry mismatch".into());
}
let local_q_dim = q_m.out_features / ranks;
let local_kv_dim = k_m.out_features / ranks;
let o_out = o_m.out_features;
let o_block_cols = o_m.canonical_chunk_cols;
let blocks_per_rank = o_m.ranks.first().map(Vec::len).unwrap_or(0);
if blocks_per_rank == 0
|| o_m
.ranks
.iter()
.any(|blocks| blocks.len() != blocks_per_rank)
|| blocks_per_rank * o_block_cols * ranks != o_m.in_features
{
return Err("step TP decode v2 O canonical block grid mismatch".into());
}
let mut guard = self
.decode_v2
.lock()
.map_err(|_| "step TP decode v2 workspace lock is poisoned")?;
if let Some(index) = guard.iter().position(|ws| {
ws.local_q_dim == local_q_dim
&& ws.local_kv_dim == local_kv_dim
&& ws.heads == heads
&& ws.o_out == o_out
&& ws.o_block_cols == o_block_cols
&& ws.blocks_per_rank == blocks_per_rank
&& ws.e_device == e.ctx().ordinal()
&& ws.q.len() == ranks
}) {
return Ok(index);
}
let mut q_raw = Vec::with_capacity(ranks);
let mut k_raw = Vec::with_capacity(ranks);
let mut v_raw = Vec::with_capacity(ranks);
let mut q = Vec::with_capacity(ranks);
let mut k = Vec::with_capacity(ranks);
let mut pos = Vec::with_capacity(ranks);
let mut gate = Vec::with_capacity(ranks);
let mut attn_out = Vec::with_capacity(ranks);
let mut gated = Vec::with_capacity(ranks);
let mut fuse_ctr = Vec::with_capacity(ranks);
let mut o_partials = Vec::with_capacity(ranks);
let mut ev_rank = Vec::with_capacity(ranks);
let direct_join = oproj_direct_on();
for (rank, engine) in self.ranks.iter().enumerate() {
let _main = engine.gpu.enter_main()?;
q_raw.push(engine.uninit(local_q_dim)?);
k_raw.push(engine.uninit(local_kv_dim)?);
v_raw.push(engine.uninit(local_kv_dim)?);
q.push(engine.uninit(local_q_dim)?);
k.push(engine.uninit(local_kv_dim)?);
pos.push(engine.htod_i32(&[0])?);
fuse_ctr.push(engine.stream().clone_htod(&[0u32])?);
gate.push(engine.uninit(heads / ranks)?);
attn_out.push(engine.uninit(local_q_dim)?);
gated.push(engine.uninit(local_q_dim)?);
let mut rank_partials = Vec::with_capacity(blocks_per_rank);
for _ in 0..blocks_per_rank {
if direct_join && rank != 0 {
let root = &self.ranks[0];
let _root_main = root.gpu.enter_main()?;
rank_partials.push(root.uninit(o_out)?);
} else {
rank_partials.push(engine.uninit(o_out)?);
}
}
o_partials.push(rank_partials);
ev_rank.push(engine.ctx().new_event(None)?);
}
let root = &self.ranks[0];
let (peer_partial, reduce_a, reduce_b, zeros, k_shadow, v_shadow, ev_refresh, ev_oproj) = {
let _main = root.gpu.enter_main()?;
(
root.uninit(o_out)?,
root.uninit(o_out)?,
root.uninit(o_out)?,
root.htod(&vec![0.0f32; o_out])?,
root.uninit(ranks * local_kv_dim)?,
root.uninit(ranks * local_kv_dim)?,
root.ctx().new_event(None)?,
root.ctx().new_event(None)?,
)
};
let (gate_e, ev_entry) = {
let _main = e.gpu.enter_main()?;
(e.uninit(heads)?, e.ctx().new_event(None)?)
};
let raw_attn_in = Vec::new();
let raw_pos = Vec::new();
guard.push(StepTpDecodeV2Ws {
tcol_q: Vec::new(),
tcol_k: Vec::new(),
tcol_v: Vec::new(),
tcol_g: Vec::new(),
tcol_in: Vec::new(),
tcol_cap: 0,
tcol_gated: Vec::new(),
tcol_opart: Vec::new(),
tcol_opeer: None,
tcol_omix: None,
tcol_ocap: 0,
q_raw,
k_raw,
v_raw,
q,
k,
pos,
fuse_ctr,
gate,
attn_out,
gated,
o_partials,
ev_rank,
peer_partial,
reduce_a,
reduce_b,
zeros,
k_shadow,
v_shadow,
ev_refresh,
ev_oproj,
gate_e,
attn_in: Vec::new(),
h_stage: None,
pos_stage: None,
raw_h_stage: 0,
raw_pos_stage: 0,
raw_attn_in,
raw_pos,
raw_o_partial1: 0,
raw_peer_partial: 0,
raw_k1: 0,
raw_v1: 0,
raw_k_shadow: 0,
raw_v_shadow: 0,
raw_mixed_stage_e: 0,
raw_reduce_a: 0,
raw_shadow_stage_e: (0, 0),
ev_entry,
e_device: e.ctx().ordinal(),
local_q_dim,
local_kv_dim,
heads,
o_out,
o_block_cols,
blocks_per_rank,
});
eprintln!(
"[step-tp-decode-v2] workspace ranks={ranks} local_q={local_q_dim} \
local_kv={local_kv_dim} heads={heads} o_blocks={blocks_per_rank}x{o_block_cols} \
residency=persistent ordering=evented performance_claim=false"
);
Ok(guard.len() - 1)
}
#[allow(clippy::too_many_arguments)]
#[allow(clippy::too_many_arguments)]
pub fn decode_v2_input_qkv_tcol(
&self,
ws_index: usize,
e: &Engine,
h_t: &CudaSlice<f32>,
t: usize,
q_m: &ResidentBf16ColumnParallel,
k_m: &ResidentBf16ColumnParallel,
v_m: &ResidentBf16ColumnParallel,
gate_shards: Option<StepTpGateShards<'_>>,
) -> Result<(), Box<dyn std::error::Error>> {
let ranks = self.ranks.len();
let mut guard = self
.decode_v2
.lock()
.map_err(|_| "step TP decode v2 workspace lock is poisoned")?;
let ws = guard
.get_mut(ws_index)
.ok_or("step TP decode v2 workspace index out of range")?;
let in_f = q_m.in_features;
if h_t.len() < t * in_f || t == 0 || t > 8 {
return Err("decode_v2_input_qkv_tcol geometry".into());
}
if ws.tcol_cap < t || ws.tcol_q.len() != ranks {
ws.tcol_q.clear();
ws.tcol_k.clear();
ws.tcol_v.clear();
ws.tcol_g.clear();
ws.tcol_in.clear();
for engine in &self.ranks {
let _m = engine.gpu.enter_main()?;
ws.tcol_q.push(engine.uninit(8 * ws.local_q_dim)?);
ws.tcol_k.push(engine.uninit(8 * ws.local_kv_dim)?);
ws.tcol_v.push(engine.uninit(8 * ws.local_kv_dim)?);
ws.tcol_g
.push(engine.uninit(8 * (ws.heads / ranks).max(1))?);
ws.tcol_in.push(engine.uninit(8 * in_f)?);
}
ws.tcol_cap = 8;
}
use cudarc::driver::DevicePtr;
let raw_src = {
let _main = e.gpu.enter_main()?;
let stream = e.stream();
let (p, _g) = h_t.device_ptr(&stream);
ws.ev_entry.record(&stream)?;
p as u64
};
for rank in 0..ranks {
let engine = &self.ranks[rank];
let _main = engine.gpu.enter_main()?;
engine.stream().wait(&ws.ev_entry)?;
let raw_dst = {
let stream = engine.stream();
let (p, _g) = ws.tcol_in[rank].device_ptr(&stream);
p as u64
};
raw_copy_bytes(raw_dst, raw_src, t * in_f * 4, engine)?;
let out_g = match &gate_shards {
Some(_) => ws.heads / ranks,
None => 0,
};
match (
&q_m.ranks[rank].weight,
&k_m.ranks[rank].weight,
&v_m.ranks[rank].weight,
) {
(
ResidentBf16Weight::Bf16(wq),
ResidentBf16Weight::Bf16(wk),
ResidentBf16Weight::Bf16(wv),
) => {
let wg = match &gate_shards {
Some(StepTpGateShards::Bf16(shards)) => &shards[rank],
Some(StepTpGateShards::F32(_)) => {
return Err(
"tcol verify: gate shard class does not match bf16 QKV".into()
);
}
None => wq,
};
let StepTpDecodeV2Ws {
tcol_q,
tcol_k,
tcol_v,
tcol_g,
tcol_in,
local_q_dim,
local_kv_dim,
..
} = &mut *ws;
static REFK: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
let refk = *REFK
.get_or_init(|| std::env::var("MEMRA_TCOL_REFKERN").as_deref() == Ok("1"));
if refk {
let lq = *local_q_dim;
let lkv = *local_kv_dim;
let mut hrow = engine.uninit(in_f)?;
let mut qr = engine.uninit(lq)?;
let mut kr = engine.uninit(lkv)?;
let mut vr = engine.uninit(lkv)?;
let mut gr = engine.uninit(out_g.max(1))?;
for c in 0..t {
{
let mut dst = hrow.slice_mut(0..in_f);
engine.stream().memcpy_dtod(
&tcol_in[rank].slice(c * in_f..(c + 1) * in_f),
&mut dst,
)?;
}
engine.matvec_bf16_qkvg_into(
wq, wk, wv, wg, &hrow, &mut qr, &mut kr, &mut vr, &mut gr, in_f,
lq, lkv, out_g,
)?;
let stream = engine.stream();
{
let mut dst = tcol_q[rank].slice_mut(c * lq..(c + 1) * lq);
stream.memcpy_dtod(&qr.slice(0..lq), &mut dst)?;
}
{
let mut dst = tcol_k[rank].slice_mut(c * lkv..(c + 1) * lkv);
stream.memcpy_dtod(&kr.slice(0..lkv), &mut dst)?;
}
{
let mut dst = tcol_v[rank].slice_mut(c * lkv..(c + 1) * lkv);
stream.memcpy_dtod(&vr.slice(0..lkv), &mut dst)?;
}
if out_g > 0 {
let mut dst = tcol_g[rank].slice_mut(c * out_g..(c + 1) * out_g);
stream.memcpy_dtod(&gr.slice(0..out_g), &mut dst)?;
}
}
} else {
engine.matvec_bf16_qkvg_tcol_into(
wq,
wk,
wv,
wg,
&tcol_in[rank],
&mut tcol_q[rank],
&mut tcol_k[rank],
&mut tcol_v[rank],
&mut tcol_g[rank],
in_f,
*local_q_dim,
*local_kv_dim,
out_g,
t,
)?;
}
}
_ => return Err("tcol verify requires bf16-resident fused QKV".into()),
}
}
Ok(())
}
pub(crate) fn decode_v2_oproj_tcol_eligible(
&self,
ws: &StepTpDecodeV2Ws,
o_m: &ResidentStepBf16RowParallel,
) -> bool {
self.ranks.len() == 2
&& ws.blocks_per_rank == 4
&& step_tp_qkv_fused_enabled().unwrap_or(false)
&& no_local_shadow_on()
&& std::env::var("MEMRA_B4_X2").as_deref() != Ok("1")
&& o_m
.ranks
.iter()
.flatten()
.all(|block| matches!(block.weight, ResidentBf16Weight::Bf16(_)))
}
pub(crate) fn decode_v2_stash_gated(
&self,
ws: &mut StepTpDecodeV2Ws,
e: &Engine,
col: usize,
) -> Result<(), Box<dyn std::error::Error>> {
let ranks = self.ranks.len();
if col >= 8 {
return Err("decode_v2_stash_gated column out of range".into());
}
let lq = ws.local_q_dim;
if ws.tcol_ocap == 0 || ws.tcol_gated.len() != ranks {
ws.tcol_gated.clear();
ws.tcol_opart.clear();
for engine in &self.ranks {
let _m = engine.gpu.enter_main()?;
ws.tcol_gated.push(engine.uninit(8 * lq)?);
ws.tcol_opart.push(engine.uninit(8 * ws.o_out)?);
}
let root = &self.ranks[0];
let _m = root.gpu.enter_main()?;
ws.tcol_opeer = Some(root.uninit(8 * ws.o_out)?);
ws.tcol_omix = Some(root.uninit(8 * ws.o_out)?);
ws.tcol_ocap = 8;
}
for rank in 0..ranks {
let engine = &self.ranks[rank];
let _main = engine.gpu.enter_main()?;
let mut dst = ws.tcol_gated[rank].slice_mut(col * lq..(col + 1) * lq);
engine
.stream()
.memcpy_dtod(&ws.gated[rank].slice(0..lq), &mut dst)?;
ws.ev_rank[rank].record(&engine.stream())?;
}
{
let _main = e.gpu.enter_main()?;
for ev in ws.ev_rank.iter() {
e.stream().wait(ev)?;
}
}
Ok(())
}
pub(crate) fn decode_v2_oproj_tcol(
&self,
ws_index: usize,
e: &Engine,
o_m: &ResidentStepBf16RowParallel,
t: usize,
) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
let ranks = self.ranks.len();
let mut guard = self
.decode_v2
.lock()
.map_err(|_| "step TP decode v2 workspace lock is poisoned")?;
let ws = guard
.get_mut(ws_index)
.ok_or("step TP decode v2 workspace index out of range")?;
if ranks != 2 || ws.blocks_per_rank != 4 || t == 0 || t > 8 || ws.tcol_ocap < t {
return Err("decode_v2_oproj_tcol geometry".into());
}
for rank in 0..ranks {
let engine = &self.ranks[rank];
let _main = engine.gpu.enter_main()?;
let mut weights = Vec::with_capacity(4);
for block in 0..4 {
let ResidentBf16Weight::Bf16(weight) = &o_m.ranks[rank][block].weight else {
return Err("tcol o_proj requires bf16-resident O blocks".into());
};
weights.push(weight);
}
{
let StepTpDecodeV2Ws {
tcol_gated,
tcol_opart,
local_q_dim,
o_block_cols,
o_out,
..
} = &mut *ws;
static REFK: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
let refk = *REFK
.get_or_init(|| std::env::var("MEMRA_TCOL_OPROJ_REF").as_deref() == Ok("1"));
if refk {
let lq = *local_q_dim;
let mut xr = engine.uninit(lq)?;
let mut yr = engine.uninit(*o_out)?;
for c in 0..t {
{
let mut dst = xr.slice_mut(0..lq);
engine.stream().memcpy_dtod(
&tcol_gated[rank].slice(c * lq..(c + 1) * lq),
&mut dst,
)?;
}
engine.matvec_bf16_b4_into(
[weights[0], weights[1], weights[2], weights[3]],
&xr,
&mut yr,
*o_block_cols,
*o_out,
)?;
let mut dst = tcol_opart[rank].slice_mut(c * *o_out..(c + 1) * *o_out);
engine
.stream()
.memcpy_dtod(&yr.slice(0..*o_out), &mut dst)?;
}
} else {
engine.matvec_bf16_b4_tcol_into(
[weights[0], weights[1], weights[2], weights[3]],
&tcol_gated[rank],
&mut tcol_opart[rank],
*o_block_cols,
*o_out,
t,
)?;
}
}
if rank != 0 {
ws.ev_rank[rank].record(&engine.stream())?;
}
}
let root = &self.ranks[0];
{
let _main = root.gpu.enter_main()?;
for ev in ws.ev_rank.iter().skip(1) {
root.stream().wait(ev)?;
}
{
let StepTpDecodeV2Ws {
tcol_opart,
tcol_opeer,
tcol_omix,
o_out,
..
} = &mut *ws;
let opeer = tcol_opeer.as_mut().ok_or("tcol o_proj slabs not armed")?;
let omix = tcol_omix.as_mut().ok_or("tcol o_proj slabs not armed")?;
{
let mut dst = opeer.slice_mut(0..t * *o_out);
root.stream()
.memcpy_dtod(&tcol_opart[1].slice(0..t * *o_out), &mut dst)?;
}
root.add(&tcol_opart[0], opeer, omix, t * *o_out)?;
}
ws.ev_oproj.record(&root.stream())?;
}
let _main = e.gpu.enter_main()?;
e.stream().wait(&ws.ev_oproj)?;
let mut out = e.uninit(t * ws.o_out)?;
let omix = ws.tcol_omix.as_ref().ok_or("tcol o_proj slabs not armed")?;
e.stream().memcpy_dtod(
&omix.slice(0..t * ws.o_out),
&mut out.slice_mut(0..t * ws.o_out),
)?;
Ok(out)
}
pub(crate) fn decode_v2_input_qkv(
&self,
ws: &mut StepTpDecodeV2Ws,
e: &Engine,
h: &CudaSlice<f32>,
pos_d: &CudaSlice<i32>,
gate_raw: Option<&CudaSlice<f32>>,
gate_shards: Option<StepTpGateShards<'_>>,
decode_input: &mut ResidentReplicatedDeviceRows,
q_m: &ResidentBf16ColumnParallel,
k_m: &ResidentBf16ColumnParallel,
v_m: &ResidentBf16ColumnParallel,
q_norm: &[CudaSlice<f32>],
k_norm: &[CudaSlice<f32>],
head_dim: usize,
n_rot: usize,
rope_base: f32,
rope_freqs: &[Option<&CudaSlice<f32>>],
rms_eps: f32,
defer_norm_rope: bool,
tcol_col: Option<usize>,
) -> Result<(), Box<dyn std::error::Error>> {
let ranks = self.ranks.len();
validate_replicated_device_rows(&self.ranks, decode_input)?;
if decode_input.tokens != 1
|| decode_input.width != q_m.in_features
|| pos_d.len() != 1
|| gate_raw.is_some_and(|gate| gate.len() != ws.heads)
|| gate_raw.is_none() != gate_shards.is_some()
|| gate_shards.as_ref().is_some_and(|shards| match shards {
StepTpGateShards::F32(shards) => shards.len() != ranks,
StepTpGateShards::Bf16(shards) => shards.len() != ranks,
})
|| q_norm.len() != ranks
|| k_norm.len() != ranks
|| rope_freqs.len() != ranks
|| e.ctx().ordinal() != ws.e_device
{
return Err("step TP decode v2 input geometry mismatch".into());
}
let qkv_fused = step_tp_qkv_fused_enabled()?;
if gate_shards.is_some() && !qkv_fused {
return Err("step TP decode v2 gate shards require MEMRA_STEP_TP_QKV_FUSED=1".into());
}
let values = decode_input.width;
if h.len() != values {
return Err(format!(
"step TP decode v2 hidden width {} != replicated width {values}",
h.len()
)
.into());
}
if qkv_fused {
if ws.h_stage.is_none() {
use cudarc::driver::DevicePtr;
let _main = e.gpu.enter_main()?;
let h_stage = e.uninit(values)?;
let pos_stage = e.htod_i32(&[0])?;
{
let stream = e.stream();
let (hp, _g0) = h_stage.device_ptr(&stream);
let (pp, _g1) = pos_stage.device_ptr(&stream);
ws.raw_h_stage = hp as u64;
ws.raw_pos_stage = pp as u64;
}
ws.h_stage = Some(h_stage);
ws.pos_stage = Some(pos_stage);
for rank in 0..ranks {
use cudarc::driver::DevicePtr;
let engine = &self.ranks[rank];
let _rmain = engine.gpu.enter_main()?;
let attn_in = engine.uninit(values)?;
let (dp, pp) = {
let stream = engine.stream();
let (dp, _g2) = attn_in.device_ptr(&stream);
let (pp, _g3) = ws.pos[rank].device_ptr(&stream);
(dp as u64, pp as u64)
};
ws.raw_attn_in.push(dp);
ws.raw_pos.push(pp);
ws.attn_in.push(attn_in);
}
{
use cudarc::driver::DevicePtr;
let root = &self.ranks[0];
let _rmain = root.gpu.enter_main()?;
let stream = root.stream();
let (a, _g) = ws.peer_partial.device_ptr(&stream);
let (b, _g) = ws.k_shadow.device_ptr(&stream);
let (c, _g) = ws.v_shadow.device_ptr(&stream);
ws.raw_peer_partial = a as u64;
ws.raw_k_shadow = b as u64;
ws.raw_v_shadow = c as u64;
}
{
use cudarc::driver::DevicePtr;
let rank1 = &self.ranks[1];
let _rmain = rank1.gpu.enter_main()?;
let stream = rank1.stream();
let (a, _g) = ws.o_partials[1][0].device_ptr(&stream);
let (b, _g) = ws.k[1].device_ptr(&stream);
let (c, _g) = ws.v_raw[1].device_ptr(&stream);
ws.raw_o_partial1 = a as u64;
ws.raw_k1 = b as u64;
ws.raw_v1 = c as u64;
}
}
{
let _main = e.gpu.enter_main()?;
{
let h_stage = ws.h_stage.as_mut().expect("stage armed above");
let mut dst = h_stage.slice_mut(0..values);
e.stream().memcpy_dtod(&h.slice(0..values), &mut dst)?;
}
{
let pos_stage = ws.pos_stage.as_mut().expect("stage armed above");
let mut dst = pos_stage.slice_mut(0..1);
e.stream().memcpy_dtod(&pos_d.slice(0..1), &mut dst)?;
}
ws.ev_entry.record(&e.stream())?;
}
for rank in 0..ranks {
let engine = &self.ranks[rank];
let _main = engine.gpu.enter_main()?;
engine.stream().wait(&ws.ev_entry)?;
}
} else {
{
let _main = e.gpu.enter_main()?;
if let Some(gate_raw) = gate_raw {
let mut gate_dst = ws.gate_e.slice_mut(0..ws.heads);
e.stream()
.memcpy_dtod(&gate_raw.slice(0..ws.heads), &mut gate_dst)?;
}
ws.ev_entry.record(&e.stream())?;
}
{
let root = &self.ranks[0];
let _main = root.gpu.enter_main()?;
root.stream().wait(&ws.ev_entry)?;
let mut destination = decode_input.ranks[0].slice_mut(0..values);
root.stream()
.memcpy_dtod(&h.slice(0..values), &mut destination)?;
ws.ev_refresh.record(&root.stream())?;
}
for rank in 1..ranks {
let engine = &self.ranks[rank];
let _main = engine.gpu.enter_main()?;
engine.stream().wait(&ws.ev_refresh)?;
let (root_rows, peer_rows) = decode_input.ranks.split_at_mut(rank);
let mut destination = peer_rows[0].slice_mut(0..values);
engine
.stream()
.memcpy_dtod(&root_rows[0].slice(0..values), &mut destination)?;
}
}
for rank in 0..ranks {
self.decode_v2_input_qkv_rank(
ws,
pos_d,
decode_input,
q_m,
k_m,
v_m,
q_norm,
k_norm,
head_dim,
n_rot,
rope_base,
rope_freqs,
rms_eps,
gate_shards.as_ref(),
qkv_fused,
defer_norm_rope,
rank,
tcol_col,
)?;
}
Ok(())
}
#[allow(clippy::too_many_arguments)]
pub(crate) fn decode_v2_input_qkv_rank(
&self,
ws: &mut StepTpDecodeV2Ws,
pos_d: &CudaSlice<i32>,
decode_input: &mut ResidentReplicatedDeviceRows,
q_m: &ResidentBf16ColumnParallel,
k_m: &ResidentBf16ColumnParallel,
v_m: &ResidentBf16ColumnParallel,
q_norm: &[CudaSlice<f32>],
k_norm: &[CudaSlice<f32>],
head_dim: usize,
n_rot: usize,
rope_base: f32,
rope_freqs: &[Option<&CudaSlice<f32>>],
rms_eps: f32,
gate_shards: Option<&StepTpGateShards<'_>>,
qkv_fused: bool,
defer_norm_rope: bool,
rank: usize,
tcol_col: Option<usize>,
) -> Result<(), Box<dyn std::error::Error>> {
let ranks = self.ranks.len();
let local_heads = ws.local_q_dim / head_dim;
let local_kv_heads = ws.local_kv_dim / head_dim;
let engine = &self.ranks[rank];
let _main = engine.gpu.enter_main()?;
let ws_e_device = ws.e_device;
if qkv_fused && tcol_col.is_some() {
let c = tcol_col.expect("checked");
if ws.tcol_cap == 0 || ws.tcol_q.len() != ranks {
return Err("tcol select without precompute".into());
}
if engine.ctx().ordinal() != ws_e_device {
raw_copy_bytes(ws.raw_pos[rank], ws.raw_pos_stage, 4, engine)?;
}
let StepTpDecodeV2Ws {
tcol_q,
tcol_k,
tcol_v,
tcol_g,
q_raw,
k_raw,
v_raw,
gate,
local_q_dim,
local_kv_dim,
heads,
..
} = &mut *ws;
let lg = *heads / ranks;
let stream = engine.stream();
{
let mut dst = q_raw[rank].slice_mut(0..*local_q_dim);
stream.memcpy_dtod(
&tcol_q[rank].slice(c * *local_q_dim..(c + 1) * *local_q_dim),
&mut dst,
)?;
}
{
let mut dst = k_raw[rank].slice_mut(0..*local_kv_dim);
stream.memcpy_dtod(
&tcol_k[rank].slice(c * *local_kv_dim..(c + 1) * *local_kv_dim),
&mut dst,
)?;
}
{
let mut dst = v_raw[rank].slice_mut(0..*local_kv_dim);
stream.memcpy_dtod(
&tcol_v[rank].slice(c * *local_kv_dim..(c + 1) * *local_kv_dim),
&mut dst,
)?;
}
if lg > 0 {
let mut dst = gate[rank].slice_mut(0..lg);
stream.memcpy_dtod(&tcol_g[rank].slice(c * lg..(c + 1) * lg), &mut dst)?;
}
if !defer_norm_rope {
} else {
return Ok(());
}
}
if qkv_fused {
let same_dev = engine.ctx().ordinal() == ws.e_device;
if !same_dev {
raw_copy_bytes(
ws.raw_attn_in[rank],
ws.raw_h_stage,
q_m.in_features * 4,
engine,
)?;
raw_copy_bytes(ws.raw_pos[rank], ws.raw_pos_stage, 4, engine)?;
}
let StepTpDecodeV2Ws {
q_raw,
k_raw,
v_raw,
gate,
gate_e,
attn_in,
h_stage,
heads,
local_q_dim,
local_kv_dim,
..
} = &mut *ws;
let input_ref: &CudaSlice<f32> = if same_dev {
h_stage
.as_ref()
.ok_or("step TP decode v2 stage not armed")?
} else {
&attn_in[rank]
};
match (
&q_m.ranks[rank].weight,
&k_m.ranks[rank].weight,
&v_m.ranks[rank].weight,
) {
(
ResidentBf16Weight::F32(wq),
ResidentBf16Weight::F32(wk),
ResidentBf16Weight::F32(wv),
) => {
let (wg, out_g) = match &gate_shards {
Some(StepTpGateShards::F32(shards)) => (&shards[rank], *heads / ranks),
Some(StepTpGateShards::Bf16(_)) => {
return Err("step TP decode v2 gate shard class does not \
match the F32 projections"
.into());
}
None => (&*gate_e, 0),
};
engine.matvec_f32_qkv_into(
wq,
wk,
wv,
wg,
input_ref,
&mut q_raw[rank],
&mut k_raw[rank],
&mut v_raw[rank],
&mut gate[rank],
q_m.in_features,
*local_q_dim,
*local_kv_dim,
out_g,
)?;
}
(
ResidentBf16Weight::Bf16(wq),
ResidentBf16Weight::Bf16(wk),
ResidentBf16Weight::Bf16(wv),
) => {
let (wg, out_g) = match &gate_shards {
Some(StepTpGateShards::Bf16(shards)) => (&shards[rank], *heads / ranks),
Some(StepTpGateShards::F32(_)) => {
return Err("step TP decode v2 gate shard class does not \
match the bf16 projections"
.into());
}
None => (wq, 0),
};
engine.matvec_bf16_qkvg_into(
wq,
wk,
wv,
wg,
input_ref,
&mut q_raw[rank],
&mut k_raw[rank],
&mut v_raw[rank],
&mut gate[rank],
q_m.in_features,
*local_q_dim,
*local_kv_dim,
out_g,
)?;
}
_ => {
return Err("step TP decode v2 QKV projections mix residency classes".into());
}
}
} else {
for (matrix, local_out, raw) in [
(q_m, ws.local_q_dim, &mut ws.q_raw),
(k_m, ws.local_kv_dim, &mut ws.k_raw),
(v_m, ws.local_kv_dim, &mut ws.v_raw),
] {
let ResidentBf16Weight::F32(values_w) = &matrix.ranks[rank].weight else {
return Err("step TP decode v2 lost its F32 projection residency".into());
};
let chunk_rows = matrix.canonical_chunk_rows.unwrap_or(local_out);
engine.linear_f32_resident_canonical_rows_t1_into(
&decode_input.ranks[rank],
values_w,
&mut raw[rank],
matrix.in_features,
local_out,
chunk_rows,
)?;
}
}
if qkv_fused && defer_norm_rope {
} else if qkv_fused {
let StepTpDecodeV2Ws {
q_raw,
k_raw,
q,
k,
pos,
pos_stage,
..
} = &mut *ws;
let same_dev = engine.ctx().ordinal() == ws_e_device;
let pos_ref: &CudaSlice<i32> = if same_dev {
pos_stage
.as_ref()
.ok_or("step TP decode v2 pos stage not armed")?
} else {
&pos[rank]
};
engine.qk_norm_rope_into(
&q_raw[rank],
&k_raw[rank],
&q_norm[rank],
&k_norm[rank],
&mut q[rank],
&mut k[rank],
pos_ref,
head_dim,
n_rot,
local_heads,
local_kv_heads,
rms_eps,
rope_base,
1.0,
rope_freqs[rank],
)?;
} else {
engine.rms_norm(
&ws.q_raw[rank],
&q_norm[rank],
&mut ws.q[rank],
head_dim,
local_heads,
rms_eps,
)?;
engine.rms_norm(
&ws.k_raw[rank],
&k_norm[rank],
&mut ws.k[rank],
head_dim,
local_kv_heads,
rms_eps,
)?;
{
let mut pos_dst = ws.pos[rank].slice_mut(0..1);
engine
.stream()
.memcpy_dtod(&pos_d.slice(0..1), &mut pos_dst)?;
}
engine.rope_neox2(
&mut ws.q[rank],
&mut ws.k[rank],
&ws.pos[rank],
head_dim,
n_rot,
local_heads,
local_kv_heads,
1,
rope_base,
1.0,
rope_freqs[rank],
)?;
}
if gate_shards.is_none() {
let gate_start = rank * (ws.heads / ranks);
let mut gate_dst = ws.gate[rank].slice_mut(0..ws.heads / ranks);
engine.stream().memcpy_dtod(
&ws.gate_e.slice(gate_start..gate_start + ws.heads / ranks),
&mut gate_dst,
)?;
}
Ok(())
}
pub(crate) fn decode_v2_finish_rank_partial(
&self,
ws: &mut StepTpDecodeV2Ws,
o_m: &ResidentStepBf16RowParallel,
o_fused: bool,
rank: usize,
) -> Result<(), Box<dyn std::error::Error>> {
let engine = &self.ranks[rank];
let _main = engine.gpu.enter_main()?;
if o_fused {
let StepTpDecodeV2Ws {
gated,
o_partials,
o_block_cols,
o_out,
..
} = &mut *ws;
let all_f32 = o_m.ranks[rank]
.iter()
.all(|block| matches!(block.weight, ResidentBf16Weight::F32(_)));
if all_f32 {
let mut weights = Vec::with_capacity(4);
for block in 0..4 {
let ResidentBf16Weight::F32(weight) = &o_m.ranks[rank][block].weight else {
unreachable!("all_f32 checked above");
};
weights.push(weight);
}
engine.matvec_f32_b4_into(
[weights[0], weights[1], weights[2], weights[3]],
&gated[rank],
&mut o_partials[rank][0],
*o_block_cols,
*o_out,
)?;
} else {
let mut weights = Vec::with_capacity(4);
for block in 0..4 {
let ResidentBf16Weight::Bf16(weight) = &o_m.ranks[rank][block].weight else {
return Err("step TP decode v2 O projections mix residency classes".into());
};
weights.push(weight);
}
engine.matvec_bf16_b4_into(
[weights[0], weights[1], weights[2], weights[3]],
&gated[rank],
&mut o_partials[rank][0],
*o_block_cols,
*o_out,
)?;
}
} else {
for block in 0..ws.blocks_per_rank {
let ResidentBf16Weight::F32(weight) = &o_m.ranks[rank][block].weight else {
return Err("step TP decode v2 lost its F32 O residency".into());
};
let x =
ws.gated[rank].slice(block * ws.o_block_cols..(block + 1) * ws.o_block_cols);
let w = weight.slice(0..weight.len());
let mut y = ws.o_partials[rank][block].slice_mut(0..ws.o_out);
engine.linear_t1_into(&x, &w, &mut y, ws.o_block_cols, ws.o_out)?;
}
}
Ok(())
}
pub(crate) fn decode_v2_finish(
&self,
ws: &mut StepTpDecodeV2Ws,
e: &Engine,
o_m: &ResidentStepBf16RowParallel,
) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
let ranks = self.ranks.len();
if e.ctx().ordinal() != ws.e_device {
return Err("step TP decode v2 finish engine changed".into());
}
let o_fused = step_tp_qkv_fused_enabled()? && ws.blocks_per_rank == 4 && ranks == 2;
for rank in 0..ranks {
self.decode_v2_finish_rank_partial(ws, o_m, o_fused, rank)?;
if rank == 0 {
continue;
}
let engine = &self.ranks[rank];
let _main = engine.gpu.enter_main()?;
ws.ev_rank[rank].record(&engine.stream())?;
}
let root = &self.ranks[0];
#[allow(unused_assignments)]
let mut final_in_a = false;
{
let _main = root.gpu.enter_main()?;
for ev in ws.ev_rank.iter().skip(1) {
root.stream().wait(ev)?;
}
if o_fused && oproj_direct_on() && ranks == 2 && no_local_shadow_on() {
ws.ev_oproj.record(&root.stream())?;
let _main = e.gpu.enter_main()?;
e.stream().wait(&ws.ev_oproj)?;
let mut output = e.uninit(ws.o_out)?;
if oproj_tail_on() && oproj_tail_eligible() {
use cudarc::driver::DevicePtr;
let stream = e.stream();
let (p0, _g0) = ws.o_partials[0][0].device_ptr(&stream);
let (p1, _g1) = ws.o_partials[1][0].device_ptr(&stream);
set_oproj_tail((p0 as u64, p1 as u64));
return Ok(output);
}
e.add(
&ws.o_partials[0][0],
&ws.o_partials[1][0],
&mut output,
ws.o_out,
)?;
return Ok(output);
}
if o_fused {
self.decode_v2_finish_root_fused(ws)?;
ws.ev_oproj.record(&root.stream())?;
let _main = e.gpu.enter_main()?;
e.stream().wait(&ws.ev_oproj)?;
let mut output = e.uninit(ws.o_out)?;
e.stream().memcpy_dtod(
&ws.reduce_a.slice(0..ws.o_out),
&mut output.slice_mut(0..ws.o_out),
)?;
return Ok(output);
}
let mut first = true;
let mut current_is_a = false;
for rank in 0..ranks {
for block in 0..ws.blocks_per_rank {
let use_peer = rank != 0;
if use_peer {
root.stream()
.memcpy_dtod(&ws.o_partials[rank][block], &mut ws.peer_partial)?;
}
match (first, current_is_a, use_peer) {
(true, _, true) => {
root.add(&ws.zeros, &ws.peer_partial, &mut ws.reduce_a, ws.o_out)?
}
(true, _, false) => root.add(
&ws.zeros,
&ws.o_partials[0][block],
&mut ws.reduce_a,
ws.o_out,
)?,
(false, true, true) => {
root.add(&ws.reduce_a, &ws.peer_partial, &mut ws.reduce_b, ws.o_out)?
}
(false, true, false) => root.add(
&ws.reduce_a,
&ws.o_partials[0][block],
&mut ws.reduce_b,
ws.o_out,
)?,
(false, false, true) => {
root.add(&ws.reduce_b, &ws.peer_partial, &mut ws.reduce_a, ws.o_out)?
}
(false, false, false) => root.add(
&ws.reduce_b,
&ws.o_partials[0][block],
&mut ws.reduce_a,
ws.o_out,
)?,
}
current_is_a = first || !current_is_a;
first = false;
}
}
final_in_a = current_is_a;
for rank in 0..ranks {
let start = rank * ws.local_kv_dim;
let mut k_dst = ws.k_shadow.slice_mut(start..start + ws.local_kv_dim);
root.stream().memcpy_dtod(&ws.k[rank], &mut k_dst)?;
let mut v_dst = ws.v_shadow.slice_mut(start..start + ws.local_kv_dim);
root.stream().memcpy_dtod(&ws.v_raw[rank], &mut v_dst)?;
}
ws.ev_oproj.record(&root.stream())?;
}
let _main = e.gpu.enter_main()?;
e.stream().wait(&ws.ev_oproj)?;
let mut output = e.uninit(ws.o_out)?;
let source = if final_in_a {
&ws.reduce_a
} else {
&ws.reduce_b
};
e.stream().memcpy_dtod(
&source.slice(0..ws.o_out),
&mut output.slice_mut(0..ws.o_out),
)?;
Ok(output)
}
pub fn run_routed_experts(
&self,
experts: &ResidentExpertParallel,
input: &[f32],
tokens: usize,
selected: &[usize],
route_weights: &[f32],
experts_per_token: usize,
activation_limit: Option<f32>,
) -> Result<Vec<f32>, Box<dyn std::error::Error>> {
validate_step_expert_activation_limit(activation_limit)?;
validate_ep_residency(&self.ranks, experts)?;
validate_activations(input, tokens, experts.input_width)?;
let pairs = tokens
.checked_mul(experts_per_token)
.ok_or("EP route count overflow")?;
if selected.len() != pairs || route_weights.len() != pairs {
return Err(format!(
"EP routes selected={} weights={} != tokens {tokens} x experts/token \
{experts_per_token} ({pairs})",
selected.len(),
route_weights.len(),
)
.into());
}
if !route_weights.iter().all(|weight| weight.is_finite()) {
return Err("EP route weights contain a non-finite value".into());
}
if self.native_p2p {
return self.run_routed_experts_native(
experts,
input,
tokens,
selected,
route_weights,
experts_per_token,
activation_limit,
);
}
let mut output = vec![0.0f32; tokens * experts.input_width];
let per_rank = experts.expert_count / experts.ranks.len();
for token in 0..tokens {
let input_row = &input[token * experts.input_width..(token + 1) * experts.input_width];
for slot in 0..experts_per_token {
let pair = token * experts_per_token + slot;
let expert = selected[pair];
if expert >= experts.expert_count {
return Err(format!(
"EP selected expert {expert} outside 0..{}",
experts.expert_count
)
.into());
}
let owner = expert / per_rank;
let local_expert = expert - experts.ranks[owner].gate.expert_range.start;
let rank = &experts.ranks[owner];
let engine = &self.ranks[owner];
let gate =
run_resident_bank_expert(engine, &rank.gate, local_expert, input_row, 1)?;
let up = run_resident_bank_expert(engine, &rank.up, local_expert, input_row, 1)?;
let activated: Vec<f32> = gate
.iter()
.zip(&up)
.map(|(&gate, &up)| step_expert_activation_host(gate, up, activation_limit))
.collect();
debug_assert_eq!(activated.len(), experts.expert_width);
let down =
run_resident_bank_expert(engine, &rank.down, local_expert, &activated, 1)?;
let weight = route_weights[pair];
for (sum, value) in output
[token * experts.input_width..(token + 1) * experts.input_width]
.iter_mut()
.zip(down)
{
*sum += weight * value;
}
}
}
Ok(output)
}
fn run_routed_experts_native(
&self,
experts: &ResidentExpertParallel,
input: &[f32],
tokens: usize,
selected: &[usize],
route_weights: &[f32],
experts_per_token: usize,
activation_limit: Option<f32>,
) -> Result<Vec<f32>, Box<dyn std::error::Error>> {
if !self.native_p2p || self.ranks.len() < 2 {
return Err("native EP execution requires at least two P2P ranks".into());
}
if self.ep_device_arithmetic {
return self.run_routed_experts_native_device(
experts,
input,
tokens,
selected,
route_weights,
experts_per_token,
activation_limit,
);
}
let mut output = vec![0.0f32; tokens * experts.input_width];
let per_rank = experts.expert_count / experts.ranks.len();
for token in 0..tokens {
let input_row = &input[token * experts.input_width..(token + 1) * experts.input_width];
let mut rank_inputs = (0..self.ranks.len())
.map(|_| None)
.collect::<Vec<Option<CudaSlice<f32>>>>();
rank_inputs[0] = Some({
let root = &self.ranks[0];
let _main = root.gpu.enter_main()?;
root.htod(input_row)?
});
for slot in 0..experts_per_token {
let pair = token * experts_per_token + slot;
let expert = selected[pair];
if expert >= experts.expert_count {
return Err(format!(
"EP selected expert {expert} outside 0..{}",
experts.expert_count
)
.into());
}
let owner = expert / per_rank;
let local_expert = expert - experts.ranks[owner].gate.expert_range.start;
if rank_inputs[owner].is_none() {
let peer_input = {
let root_input = rank_inputs[0]
.as_ref()
.ok_or("native EP lost its root input")?;
let engine = &self.ranks[owner];
let _main = engine.gpu.enter_main()?;
let mut peer_input = engine.uninit(experts.input_width)?;
engine.stream().memcpy_dtod(root_input, &mut peer_input)?;
peer_input
};
rank_inputs[owner] = Some(peer_input);
}
let rank = &experts.ranks[owner];
let engine = &self.ranks[owner];
let owner_input = rank_inputs[owner]
.as_ref()
.ok_or("native EP owner input is absent after dispatch")?;
let gate = run_resident_bank_expert_device(
engine,
&rank.gate,
local_expert,
owner_input,
1,
)?;
let up = run_resident_bank_expert_device(
engine,
&rank.up,
local_expert,
owner_input,
1,
)?;
let (gate, up) = {
let _main = engine.gpu.enter_main()?;
(engine.dtoh(&gate)?, engine.dtoh(&up)?)
};
let activated = gate
.iter()
.zip(&up)
.map(|(&gate, &up)| step_expert_activation_host(gate, up, activation_limit))
.collect::<Vec<_>>();
debug_assert_eq!(activated.len(), experts.expert_width);
let activated = {
let _main = engine.gpu.enter_main()?;
engine.htod(&activated)?
};
let down = run_resident_bank_expert_device(
engine,
&rank.down,
local_expert,
&activated,
1,
)?;
let down = if owner == 0 {
let _main = engine.gpu.enter_main()?;
engine.dtoh(&down)?
} else {
let root = &self.ranks[0];
let _main = root.gpu.enter_main()?;
let mut root_down = root.uninit(experts.input_width)?;
root.stream().memcpy_dtod(&down, &mut root_down)?;
root.dtoh(&root_down)?
};
let weight = route_weights[pair];
for (sum, value) in output
[token * experts.input_width..(token + 1) * experts.input_width]
.iter_mut()
.zip(down)
{
*sum += weight * value;
}
}
}
Ok(output)
}
fn run_routed_experts_native_device(
&self,
experts: &ResidentExpertParallel,
input: &[f32],
tokens: usize,
selected: &[usize],
route_weights: &[f32],
experts_per_token: usize,
activation_limit: Option<f32>,
) -> Result<Vec<f32>, Box<dyn std::error::Error>> {
if !self.native_p2p || !self.ep_device_arithmetic || self.ranks.len() < 2 {
return Err(
"device-resident EP arithmetic requires at least two native P2P ranks".into(),
);
}
let mut output = Vec::with_capacity(tokens * experts.input_width);
let per_rank = experts.expert_count / experts.ranks.len();
let root = &self.ranks[0];
for token in 0..tokens {
let input_row = &input[token * experts.input_width..(token + 1) * experts.input_width];
let mut rank_inputs = (0..self.ranks.len())
.map(|_| None)
.collect::<Vec<Option<CudaSlice<f32>>>>();
rank_inputs[0] = Some({
let _main = root.gpu.enter_main()?;
root.htod(input_row)?
});
let mut root_output = {
let _main = root.gpu.enter_main()?;
root.zeros(experts.input_width)?
};
let mut remote_down_keepalive = Vec::new();
for slot in 0..experts_per_token {
let pair = token * experts_per_token + slot;
let expert = selected[pair];
if expert >= experts.expert_count {
return Err(format!(
"EP selected expert {expert} outside 0..{}",
experts.expert_count
)
.into());
}
let owner = expert / per_rank;
let local_expert = expert - experts.ranks[owner].gate.expert_range.start;
if rank_inputs[owner].is_none() {
let peer_input = {
let root_input = rank_inputs[0]
.as_ref()
.ok_or("native EP lost its root input")?;
let engine = &self.ranks[owner];
let _main = engine.gpu.enter_main()?;
let mut peer_input = engine.uninit(experts.input_width)?;
engine.stream().memcpy_dtod(root_input, &mut peer_input)?;
peer_input
};
rank_inputs[owner] = Some(peer_input);
}
let rank = &experts.ranks[owner];
let engine = &self.ranks[owner];
let owner_input = rank_inputs[owner]
.as_ref()
.ok_or("native EP owner input is absent after dispatch")?;
let gate = run_resident_bank_expert_device(
engine,
&rank.gate,
local_expert,
owner_input,
1,
)?;
let up = run_resident_bank_expert_device(
engine,
&rank.up,
local_expert,
owner_input,
1,
)?;
let activated = {
let _main = engine.gpu.enter_main()?;
let mut activated = engine.uninit(experts.expert_width)?;
if let Some(limit) = activation_limit {
engine.silu_clamped_mul_host_expf(
&gate,
&up,
limit,
&mut activated,
experts.expert_width,
)?;
} else {
engine.silu_mul_host_expf(
&gate,
&up,
&mut activated,
experts.expert_width,
)?;
}
activated
};
let down = run_resident_bank_expert_device(
engine,
&rank.down,
local_expert,
&activated,
1,
)?;
let root_down = if owner == 0 {
down
} else {
let _main = root.gpu.enter_main()?;
let mut root_down = root.uninit(experts.input_width)?;
root.stream().memcpy_dtod(&down, &mut root_down)?;
remote_down_keepalive.push(down);
root_down
};
let _main = root.gpu.enter_main()?;
let mut destination = root_output.slice_mut(0..experts.input_width);
root.axpy_host_into(
&root_down.slice(0..root_down.len()),
route_weights[pair],
&mut destination,
experts.input_width,
)?;
}
let _main = root.gpu.enter_main()?;
let root_output = root.dtoh(&root_output)?;
drop(remote_down_keepalive);
output.extend(root_output);
}
Ok(output)
}
}
fn validate_column_shape(matrix: E4m3BlockMatrix<'_>, tp: usize) -> Result<(), String> {
if matrix.out_features % tp != 0 {
return Err(format!(
"column-parallel out_features {} is not divisible by TP={tp}",
matrix.out_features
));
}
let local_out = matrix.out_features / tp;
if local_out % FP8_BLOCK != 0 {
return Err(format!(
"column-parallel output shard {local_out} cuts through a {FP8_BLOCK}-row \
E4M3 scale block"
));
}
Ok(())
}
fn step_bf16_canonical_chunk_rows(out_features: usize, tp: usize) -> Result<usize, String> {
if !matches!(tp, 1 | 2 | 4 | 8) {
return Err(format!(
"Step BF16 canonical projection requires TP1/TP2/TP4/TP8, got TP={tp}"
));
}
if out_features == 0 || out_features % PRODUCT_MAX_CARDS != 0 {
return Err(format!(
"Step BF16 output width {out_features} is not divisible by the TP8 product envelope"
));
}
let canonical_rows = out_features / PRODUCT_MAX_CARDS;
let local_out = out_features / tp;
if local_out % canonical_rows != 0 {
return Err(format!(
"Step BF16 TP={tp} output shard {local_out} is not divisible by canonical \
{canonical_rows}-row chunks"
));
}
Ok(canonical_rows)
}
fn step_bf16_canonical_chunk_cols(in_features: usize, tp: usize) -> Result<usize, String> {
if !matches!(tp, 1 | 2 | 4 | 8) {
return Err(format!(
"Step BF16 canonical row projection requires TP1/TP2/TP4/TP8, got TP={tp}"
));
}
if in_features == 0 || in_features % PRODUCT_MAX_CARDS != 0 {
return Err(format!(
"Step BF16 input width {in_features} is not divisible by the TP8 product envelope"
));
}
let canonical_cols = in_features / PRODUCT_MAX_CARDS;
let local_in = in_features / tp;
if local_in % canonical_cols != 0 {
return Err(format!(
"Step BF16 TP={tp} input shard {local_in} is not divisible by canonical \
{canonical_cols}-column chunks"
));
}
Ok(canonical_cols)
}
fn validate_row_shape(matrix: E4m3BlockMatrix<'_>, tp: usize) -> Result<(), String> {
if matrix.in_features % tp != 0 {
return Err(format!(
"row-parallel in_features {} is not divisible by TP={tp}",
matrix.in_features
));
}
let local_in = matrix.in_features / tp;
if local_in % FP8_BLOCK != 0 {
return Err(format!(
"row-parallel input shard {local_in} cuts through a {FP8_BLOCK}-column \
E4M3 scale block"
));
}
Ok(())
}
fn upload_rank(
engine: &Engine,
matrix: E4m3BlockMatrix<'_>,
) -> Result<ResidentE4m3Rank, Box<dyn std::error::Error>> {
let _main = engine.gpu.enter_main()?;
matrix.validate()?;
Ok(ResidentE4m3Rank {
codes: engine.htod_bytes(matrix.codes)?,
scales: engine.htod(matrix.scales)?,
out_features: matrix.out_features,
in_features: matrix.in_features,
})
}
fn upload_bf16_rank(
engine: &Engine,
matrix: Bf16Matrix<'_>,
f32_mirror: bool,
) -> Result<ResidentBf16Rank, Box<dyn std::error::Error>> {
let _main = engine.gpu.enter_main()?;
matrix.validate()?;
let bytes = engine.htod_bytes(matrix.bytes)?;
let weight = if f32_mirror {
let values = matrix
.out_features
.checked_mul(matrix.in_features)
.ok_or("resident BF16 mirror element count overflow")?;
ResidentBf16Weight::F32(engine.bf16_to_f32(&bytes.slice(0..bytes.len()), values)?)
} else {
ResidentBf16Weight::Bf16(bytes)
};
Ok(ResidentBf16Rank {
weight,
out_features: matrix.out_features,
in_features: matrix.in_features,
})
}
fn upload_expert_bank_rank(
engine: &Engine,
bank: E4m3ExpertBank<'_>,
expert_range: Range<usize>,
) -> Result<ResidentE4m3ExpertBankRank, Box<dyn std::error::Error>> {
let _main = engine.gpu.enter_main()?;
bank.validate()?;
if expert_range.start >= expert_range.end || expert_range.end > bank.expert_count {
return Err(format!(
"invalid EP expert range {expert_range:?} for {} experts",
bank.expert_count
)
.into());
}
let code_stride = bank.out_features * bank.in_features;
let scale_stride = bank.out_features.div_ceil(FP8_BLOCK) * bank.in_features.div_ceil(FP8_BLOCK);
Ok(ResidentE4m3ExpertBankRank {
codes: engine.htod_bytes(
&bank.codes[expert_range.start * code_stride..expert_range.end * code_stride],
)?,
scales: engine.htod(
&bank.scales[expert_range.start * scale_stride..expert_range.end * scale_stride],
)?,
expert_range,
out_features: bank.out_features,
in_features: bank.in_features,
code_stride,
scale_stride,
k_blocks: None,
})
}
fn validate_column_bank_shape(bank: E4m3ExpertBank<'_>, tp: usize) -> Result<(), String> {
if bank.out_features % tp != 0 {
return Err(format!(
"TP expert output width {} is not divisible by TP={tp}",
bank.out_features
));
}
let local_out = bank.out_features / tp;
if local_out % FP8_BLOCK != 0 {
return Err(format!(
"TP expert output shard {local_out} cuts through a {FP8_BLOCK}-row E4M3 scale block"
));
}
Ok(())
}
fn validate_row_bank_shape(bank: E4m3ExpertBank<'_>, tp: usize) -> Result<(), String> {
if bank.in_features % tp != 0 {
return Err(format!(
"TP expert input width {} is not divisible by TP={tp}",
bank.in_features
));
}
let local_in = bank.in_features / tp;
if local_in % FP8_BLOCK != 0 {
return Err(format!(
"TP expert input shard {local_in} cuts through a {FP8_BLOCK}-column E4M3 scale block"
));
}
Ok(())
}
fn upload_column_bank_rank(
engine: &Engine,
bank: E4m3ExpertBank<'_>,
tp: usize,
rank: usize,
) -> Result<ResidentE4m3ExpertBankRank, Box<dyn std::error::Error>> {
let _main = engine.gpu.enter_main()?;
let packed = pack_column_bank_rank(bank, tp, rank)?;
Ok(ResidentE4m3ExpertBankRank {
codes: engine.htod_bytes(&packed.codes)?,
scales: engine.htod(&packed.scales)?,
expert_range: packed.expert_range,
out_features: packed.out_features,
in_features: packed.in_features,
code_stride: packed.code_stride,
scale_stride: packed.scale_stride,
k_blocks: packed.k_blocks,
})
}
fn pack_column_bank_rank(
bank: E4m3ExpertBank<'_>,
tp: usize,
rank: usize,
) -> Result<PackedE4m3ExpertBankRank, String> {
bank.validate()?;
validate_column_bank_shape(bank, tp)?;
if rank >= tp {
return Err(format!("TP rank {rank} outside 0..{tp}"));
}
let local_out = bank.out_features / tp;
let full_code_stride = bank.out_features * bank.in_features;
let local_code_stride = local_out * bank.in_features;
let scale_cols = bank.in_features.div_ceil(FP8_BLOCK);
let full_scale_stride = bank.out_features.div_ceil(FP8_BLOCK) * scale_cols;
let local_scale_rows = local_out / FP8_BLOCK;
let local_scale_stride = local_scale_rows * scale_cols;
let mut codes = Vec::with_capacity(bank.expert_count * local_code_stride);
let mut scales = Vec::with_capacity(bank.expert_count * local_scale_stride);
let row_start = rank * local_out;
let scale_row_start = rank * local_scale_rows;
for expert in 0..bank.expert_count {
let code_start = expert * full_code_stride + row_start * bank.in_features;
codes.extend_from_slice(&bank.codes[code_start..code_start + local_code_stride]);
let scale_start = expert * full_scale_stride + scale_row_start * scale_cols;
scales.extend_from_slice(&bank.scales[scale_start..scale_start + local_scale_stride]);
}
Ok(PackedE4m3ExpertBankRank {
codes,
scales,
expert_range: 0..bank.expert_count,
out_features: local_out,
in_features: bank.in_features,
code_stride: local_code_stride,
scale_stride: local_scale_stride,
k_blocks: None,
})
}
fn upload_row_bank_rank(
engine: &Engine,
bank: E4m3ExpertBank<'_>,
tp: usize,
rank: usize,
) -> Result<ResidentE4m3ExpertBankRank, Box<dyn std::error::Error>> {
let _main = engine.gpu.enter_main()?;
let packed = pack_row_bank_rank(bank, tp, rank)?;
Ok(ResidentE4m3ExpertBankRank {
codes: engine.htod_bytes(&packed.codes)?,
scales: engine.htod(&packed.scales)?,
expert_range: packed.expert_range,
out_features: packed.out_features,
in_features: packed.in_features,
code_stride: packed.code_stride,
scale_stride: packed.scale_stride,
k_blocks: packed.k_blocks,
})
}
fn pack_row_bank_rank(
bank: E4m3ExpertBank<'_>,
tp: usize,
rank: usize,
) -> Result<PackedE4m3ExpertBankRank, String> {
bank.validate()?;
validate_row_bank_shape(bank, tp)?;
if rank >= tp {
return Err(format!("TP rank {rank} outside 0..{tp}"));
}
let local_in = bank.in_features / tp;
let full_code_stride = bank.out_features * bank.in_features;
let local_code_stride = bank.out_features * local_in;
let full_scale_cols = bank.in_features.div_ceil(FP8_BLOCK);
let local_scale_cols = local_in / FP8_BLOCK;
let scale_rows = bank.out_features.div_ceil(FP8_BLOCK);
let full_scale_stride = scale_rows * full_scale_cols;
let local_scale_stride = scale_rows * local_scale_cols;
let global_block_start = rank * local_scale_cols;
let mut codes = Vec::with_capacity(bank.expert_count * local_code_stride);
let mut scales = Vec::with_capacity(bank.expert_count * local_scale_stride);
for expert in 0..bank.expert_count {
let expert_code_start = expert * full_code_stride;
let expert_scale_start = expert * full_scale_stride;
for local_block in 0..local_scale_cols {
let global_block = global_block_start + local_block;
let column_start = global_block * FP8_BLOCK;
for row in 0..bank.out_features {
let start = expert_code_start + row * bank.in_features + column_start;
codes.extend_from_slice(&bank.codes[start..start + FP8_BLOCK]);
}
for row in 0..scale_rows {
scales.push(bank.scales[expert_scale_start + row * full_scale_cols + global_block]);
}
}
}
Ok(PackedE4m3ExpertBankRank {
codes,
scales,
expert_range: 0..bank.expert_count,
out_features: bank.out_features,
in_features: local_in,
code_stride: local_code_stride,
scale_stride: local_scale_stride,
k_blocks: Some(local_scale_cols),
})
}
fn validate_resident_ranks(engines: &[Engine], ranks: &[ResidentE4m3Rank]) -> Result<(), String> {
if engines.len() != ranks.len() {
return Err(format!(
"resident TP rank count {} != runtime rank count {}",
ranks.len(),
engines.len()
));
}
for (rank, (engine, matrix)) in engines.iter().zip(ranks).enumerate() {
let device = engine.ctx().ordinal();
if matrix.codes.ordinal() != device || matrix.scales.ordinal() != device {
return Err(format!(
"resident TP rank {rank} is not owned by runtime device {device}"
));
}
}
Ok(())
}
fn validate_tp_bank_residency(
engines: &[Engine],
experts: &ResidentTpExpertBank,
) -> Result<(), String> {
if engines.len() != experts.gate.len()
|| engines.len() != experts.up.len()
|| engines.len() != experts.down.len()
{
return Err(format!(
"resident TP expert-bank rank counts gate={} up={} down={} != runtime {}",
experts.gate.len(),
experts.up.len(),
experts.down.len(),
engines.len()
));
}
for (rank, engine) in engines.iter().enumerate() {
let device = engine.ctx().ordinal();
for (projection, bank) in [
("gate", &experts.gate[rank]),
("up", &experts.up[rank]),
("down", &experts.down[rank]),
] {
if bank.codes.ordinal() != device || bank.scales.ordinal() != device {
return Err(format!(
"resident TP rank {rank} {projection} bank is not owned by runtime device \
{device}"
));
}
}
}
Ok(())
}
fn validate_ep_residency(
engines: &[Engine],
experts: &ResidentExpertParallel,
) -> Result<(), String> {
if engines.len() != experts.ranks.len() {
return Err(format!(
"resident EP rank count {} != runtime rank count {}",
experts.ranks.len(),
engines.len()
));
}
for (rank, (engine, resident)) in engines.iter().zip(&experts.ranks).enumerate() {
let device = engine.ctx().ordinal();
for (projection, bank) in [
("gate", &resident.gate),
("up", &resident.up),
("down", &resident.down),
] {
if bank.codes.ordinal() != device || bank.scales.ordinal() != device {
return Err(format!(
"resident EP rank {rank} {projection} bank is not owned by runtime device \
{device}"
));
}
}
}
Ok(())
}
fn run_rank(
engine: &Engine,
matrix: E4m3BlockMatrix<'_>,
activations: &[f32],
tokens: usize,
) -> Result<Vec<f32>, Box<dyn std::error::Error>> {
let _main = engine.gpu.enter_main()?;
let codes = engine.htod_bytes(matrix.codes)?;
let scales = engine.htod(matrix.scales)?;
let activations = engine.htod(activations)?;
let output = engine.qmatvec_mmq_fp8_blk(
&codes,
&scales,
&activations,
tokens,
matrix.in_features,
matrix.out_features,
)?;
engine.dtoh(&output)
}
fn run_resident_rank(
engine: &Engine,
matrix: &ResidentE4m3Rank,
activations: &[f32],
tokens: usize,
) -> Result<Vec<f32>, Box<dyn std::error::Error>> {
let _main = engine.gpu.enter_main()?;
let activations = engine.htod(activations)?;
let output = engine.qmatvec_mmq_fp8_blk(
&matrix.codes,
&matrix.scales,
&activations,
tokens,
matrix.in_features,
matrix.out_features,
)?;
engine.dtoh(&output)
}
fn run_resident_bf16_rank(
engine: &Engine,
matrix: &ResidentBf16Rank,
activations: &[f32],
tokens: usize,
canonical_chunk_rows: Option<usize>,
) -> Result<Vec<f32>, Box<dyn std::error::Error>> {
let _main = engine.gpu.enter_main()?;
let activations = engine.htod(activations)?;
let output = run_resident_bf16_rank_device(
engine,
matrix,
&activations,
tokens,
canonical_chunk_rows,
false,
)?;
engine.dtoh(&output)
}
fn run_resident_bf16_rank_device(
engine: &Engine,
matrix: &ResidentBf16Rank,
activations: &CudaSlice<f32>,
tokens: usize,
canonical_chunk_rows: Option<usize>,
strided_chunk_output: bool,
) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
let _main = engine.gpu.enter_main()?;
if activations.ordinal() != engine.ctx().ordinal() {
return Err(format!(
"resident BF16 activation device {} != rank device {}",
activations.ordinal(),
engine.ctx().ordinal()
)
.into());
}
if activations.len() != tokens * matrix.in_features {
return Err(format!(
"resident BF16 activation count {} != {tokens}x{}",
activations.len(),
matrix.in_features
)
.into());
}
match (&matrix.weight, canonical_chunk_rows) {
(ResidentBf16Weight::Bf16(bytes), Some(rows)) => engine
.linear_bf16_resident_canonical_rows(
activations,
bytes,
tokens,
matrix.in_features,
matrix.out_features,
rows,
),
(ResidentBf16Weight::Bf16(bytes), None) => engine.linear_bf16_resident(
activations,
bytes,
tokens,
matrix.in_features,
matrix.out_features,
),
(ResidentBf16Weight::F32(values), Some(rows)) if strided_chunk_output => engine
.linear_f32_resident_canonical_rows_strided(
activations,
values,
tokens,
matrix.in_features,
matrix.out_features,
rows,
),
(ResidentBf16Weight::F32(values), Some(rows)) => engine.linear_f32_resident_canonical_rows(
activations,
values,
tokens,
matrix.in_features,
matrix.out_features,
rows,
),
(ResidentBf16Weight::F32(values), None) => engine.linear(
activations,
values,
tokens,
matrix.in_features,
matrix.out_features,
),
}
}
fn validate_resident_bf16_ranks(
engines: &[Engine],
ranks: &[ResidentBf16Rank],
) -> Result<(), String> {
if engines.len() != ranks.len() {
return Err(format!(
"resident BF16 TP rank count {} != runtime rank count {}",
ranks.len(),
engines.len(),
));
}
for (rank, (engine, matrix)) in engines.iter().zip(ranks).enumerate() {
let device = engine.ctx().ordinal();
if matrix.weight.ordinal() != device {
return Err(format!(
"resident BF16 TP rank {rank} is not owned by runtime device {device}"
));
}
}
Ok(())
}
fn validate_step_bf16_row_residency(
engines: &[Engine],
matrix: &ResidentStepBf16RowParallel,
) -> Result<(), String> {
if engines.len() != matrix.ranks.len() {
return Err(format!(
"resident Step BF16 row rank count {} != runtime rank count {}",
matrix.ranks.len(),
engines.len(),
));
}
let canonical_cols = step_bf16_canonical_chunk_cols(matrix.in_features, engines.len())?;
if matrix.canonical_chunk_cols != canonical_cols {
return Err(format!(
"resident Step BF16 row canonical columns {} != registered {canonical_cols}",
matrix.canonical_chunk_cols
));
}
let blocks_per_rank = PRODUCT_MAX_CARDS / engines.len();
for (rank, (engine, blocks)) in engines.iter().zip(&matrix.ranks).enumerate() {
if blocks.len() != blocks_per_rank {
return Err(format!(
"resident Step BF16 row rank {rank} has {} blocks, expected {blocks_per_rank}",
blocks.len()
));
}
let device = engine.ctx().ordinal();
for (block, resident) in blocks.iter().enumerate() {
if resident.weight.ordinal() != device
|| resident.in_features != canonical_cols
|| resident.out_features != matrix.out_features
{
return Err(format!(
"resident Step BF16 row rank {rank} block {block} has inconsistent \
device or geometry"
));
}
}
}
Ok(())
}
fn validate_replicated_device_rows(
engines: &[Engine],
rows: &ResidentReplicatedDeviceRows,
) -> Result<(), String> {
let rank_lengths = rows
.ranks
.iter()
.map(|rank_rows| rank_rows.len())
.collect::<Vec<_>>();
replicated_device_row_values(rows.tokens, rows.width, engines.len(), &rank_lengths)?;
if rows
.ranks
.iter()
.zip(engines)
.any(|(rank_rows, engine)| rank_rows.ordinal() != engine.ctx().ordinal())
{
return Err("replicated device rows are owned by the wrong CUDA contexts".into());
}
Ok(())
}
fn replicated_device_row_values(
tokens: usize,
width: usize,
expected_ranks: usize,
rank_lengths: &[usize],
) -> Result<usize, String> {
let values = tokens
.checked_mul(width)
.ok_or("replicated device row size overflow")?;
if tokens == 0
|| width == 0
|| expected_ranks == 0
|| rank_lengths.len() != expected_ranks
|| rank_lengths.iter().any(|&rank_len| rank_len != values)
{
return Err(format!(
"replicated device rows have inconsistent geometry tokens={} width={} ranks={}/{}",
tokens,
width,
rank_lengths.len(),
expected_ranks
));
}
Ok(values)
}
fn replicated_device_row_source_values(
tokens: usize,
width: usize,
source_len: usize,
source_device: usize,
root_device: usize,
) -> Result<usize, String> {
let values = tokens
.checked_mul(width)
.ok_or("replicated device row size overflow")?;
if tokens == 0 || width == 0 || source_len != values || source_device != root_device {
return Err(format!(
"replicated device row source has inconsistent geometry/device \
tokens={tokens} width={width} source={source_len}@{source_device} root={root_device}"
));
}
Ok(values)
}
fn bf16_column_shard(
matrix: Bf16Matrix<'_>,
tp: usize,
rank: usize,
) -> Result<Bf16Matrix<'_>, String> {
matrix.validate()?;
if tp == 0 || rank >= tp || matrix.out_features % tp != 0 {
return Err(format!(
"invalid BF16 column shard out={} TP={tp} rank={rank}",
matrix.out_features
));
}
let local_out = matrix.out_features / tp;
let row_bytes = matrix.in_features * 2;
let start = rank * local_out * row_bytes;
Ok(Bf16Matrix {
bytes: &matrix.bytes[start..start + local_out * row_bytes],
out_features: local_out,
in_features: matrix.in_features,
})
}
fn bf16_row_shard(matrix: Bf16Matrix<'_>, tp: usize, rank: usize) -> Result<Vec<u8>, String> {
matrix.validate()?;
if tp == 0 || rank >= tp || matrix.in_features % tp != 0 {
return Err(format!(
"invalid BF16 row shard in={} TP={tp} rank={rank}",
matrix.in_features
));
}
let local_in = matrix.in_features / tp;
let mut bytes = Vec::with_capacity(matrix.out_features * local_in * 2);
for row in 0..matrix.out_features {
let start = (row * matrix.in_features + rank * local_in) * 2;
bytes.extend_from_slice(&matrix.bytes[start..start + local_in * 2]);
}
Ok(bytes)
}
fn bf16_row_block(
matrix: Bf16Matrix<'_>,
col_start: usize,
block_cols: usize,
) -> Result<Vec<u8>, String> {
matrix.validate()?;
let col_end = col_start
.checked_add(block_cols)
.ok_or("BF16 row block column overflow")?;
if block_cols == 0 || col_end > matrix.in_features {
return Err(format!(
"invalid BF16 row block columns {col_start}..{col_end} for input width {}",
matrix.in_features
));
}
let mut bytes = Vec::with_capacity(matrix.out_features * block_cols * 2);
for row in 0..matrix.out_features {
let start = (row * matrix.in_features + col_start) * 2;
bytes.extend_from_slice(&matrix.bytes[start..start + block_cols * 2]);
}
Ok(bytes)
}
fn run_resident_bank_expert(
engine: &Engine,
bank: &ResidentE4m3ExpertBankRank,
local_expert: usize,
activations: &[f32],
tokens: usize,
) -> Result<Vec<f32>, Box<dyn std::error::Error>> {
let _main = engine.gpu.enter_main()?;
if bank.k_blocks.is_some() {
return Err("block-major TP row bank requires canonical block execution".into());
}
let local_count = bank.expert_range.end - bank.expert_range.start;
if local_expert >= local_count {
return Err(format!(
"local EP expert {local_expert} outside 0..{local_count} for range {:?}",
bank.expert_range
)
.into());
}
validate_activations(activations, tokens, bank.in_features)?;
let activations = engine.htod(activations)?;
let weight = bank
.codes
.slice(local_expert * bank.code_stride..(local_expert + 1) * bank.code_stride);
let scales = bank
.scales
.slice(local_expert * bank.scale_stride..(local_expert + 1) * bank.scale_stride);
let input = activations.slice(0..activations.len());
let output = engine.qmatvec_mmq_fp8_blk_view(
&weight,
&scales,
&input,
tokens,
bank.in_features,
bank.out_features,
)?;
engine.dtoh(&output)
}
fn run_resident_bank_expert_block(
engine: &Engine,
bank: &ResidentE4m3ExpertBankRank,
local_expert: usize,
block: usize,
activations: &[f32],
) -> Result<Vec<f32>, Box<dyn std::error::Error>> {
let _main = engine.gpu.enter_main()?;
let local_count = bank.expert_range.end - bank.expert_range.start;
if local_expert >= local_count {
return Err(format!(
"local TP expert {local_expert} outside 0..{local_count} for range {:?}",
bank.expert_range
)
.into());
}
let blocks = bank
.k_blocks
.ok_or("TP row bank is not packed in native K-block order")?;
if block >= blocks {
return Err(format!("TP row block {block} outside 0..{blocks}").into());
}
validate_activations(activations, 1, FP8_BLOCK)?;
let block_code_stride = bank.out_features * FP8_BLOCK;
let block_scale_stride = bank.out_features.div_ceil(FP8_BLOCK);
if bank.in_features != blocks * FP8_BLOCK
|| bank.code_stride != blocks * block_code_stride
|| bank.scale_stride != blocks * block_scale_stride
{
return Err("TP row bank block-major geometry is inconsistent".into());
}
let expert_code_start = local_expert * bank.code_stride;
let expert_scale_start = local_expert * bank.scale_stride;
let weight = bank.codes.slice(
expert_code_start + block * block_code_stride
..expert_code_start + (block + 1) * block_code_stride,
);
let scales = bank.scales.slice(
expert_scale_start + block * block_scale_stride
..expert_scale_start + (block + 1) * block_scale_stride,
);
let activations = engine.htod(activations)?;
let input = activations.slice(0..activations.len());
let output = engine.qmatvec_mmq_fp8_blk_view(
&weight,
&scales,
&input,
1,
FP8_BLOCK,
bank.out_features,
)?;
engine.dtoh(&output)
}
fn run_resident_bank_expert_device(
engine: &Engine,
bank: &ResidentE4m3ExpertBankRank,
local_expert: usize,
activations: &CudaSlice<f32>,
tokens: usize,
) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
let _main = engine.gpu.enter_main()?;
if bank.k_blocks.is_some() {
return Err("block-major TP row bank requires canonical block execution".into());
}
let local_count = bank.expert_range.end - bank.expert_range.start;
if local_expert >= local_count {
return Err(format!(
"local TP expert {local_expert} outside 0..{local_count} for range {:?}",
bank.expert_range
)
.into());
}
let expected = tokens
.checked_mul(bank.in_features)
.ok_or("native TP activation size overflow")?;
if activations.len() != expected || activations.ordinal() != engine.ctx().ordinal() {
return Err(format!(
"native TP activation len/device {}/{} != expected {expected}/{}",
activations.len(),
activations.ordinal(),
engine.ctx().ordinal()
)
.into());
}
let weight = bank
.codes
.slice(local_expert * bank.code_stride..(local_expert + 1) * bank.code_stride);
let scales = bank
.scales
.slice(local_expert * bank.scale_stride..(local_expert + 1) * bank.scale_stride);
let input = activations.slice(0..activations.len());
engine.qmatvec_mmq_fp8_blk_view(
&weight,
&scales,
&input,
tokens,
bank.in_features,
bank.out_features,
)
}
fn run_resident_bank_expert_block_device(
engine: &Engine,
bank: &ResidentE4m3ExpertBankRank,
local_expert: usize,
block: usize,
activations: &cudarc::driver::CudaView<'_, f32>,
) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
let _main = engine.gpu.enter_main()?;
let local_count = bank.expert_range.end - bank.expert_range.start;
if local_expert >= local_count {
return Err(format!(
"local TP expert {local_expert} outside 0..{local_count} for range {:?}",
bank.expert_range
)
.into());
}
let blocks = bank
.k_blocks
.ok_or("native TP row bank is not packed in checkpoint-block order")?;
if block >= blocks {
return Err(format!("native TP row block {block} outside 0..{blocks}").into());
}
let activation_device = activations.stream().context().ordinal();
if activations.len() != FP8_BLOCK || activation_device != engine.ctx().ordinal() {
return Err(format!(
"native TP block activation len/device {}/{} != expected {FP8_BLOCK}/{}",
activations.len(),
activation_device,
engine.ctx().ordinal()
)
.into());
}
let block_code_stride = bank.out_features * FP8_BLOCK;
let block_scale_stride = bank.out_features.div_ceil(FP8_BLOCK);
if bank.in_features != blocks * FP8_BLOCK
|| bank.code_stride != blocks * block_code_stride
|| bank.scale_stride != blocks * block_scale_stride
{
return Err("native TP row bank block-major geometry is inconsistent".into());
}
let expert_code_start = local_expert * bank.code_stride;
let expert_scale_start = local_expert * bank.scale_stride;
let weight = bank.codes.slice(
expert_code_start + block * block_code_stride
..expert_code_start + (block + 1) * block_code_stride,
);
let scales = bank.scales.slice(
expert_scale_start + block * block_scale_stride
..expert_scale_start + (block + 1) * block_scale_stride,
);
engine.qmatvec_mmq_fp8_blk_view(
&weight,
&scales,
activations,
1,
FP8_BLOCK,
bank.out_features,
)
}
fn configure_native_p2p(
ranks: &[Engine],
devices: &[usize],
) -> Result<(), Box<dyn std::error::Error>> {
if ranks.len() != devices.len() || ranks.len() < 2 {
return Err("native TP P2P setup requires matching multi-rank devices".into());
}
for (rank, (&device, engine)) in devices.iter().zip(ranks).enumerate() {
if engine.ctx().ordinal() != device {
return Err(format!(
"native TP rank {rank} context device {} != requested device {device}",
engine.ctx().ordinal()
)
.into());
}
}
for src in 0..ranks.len() {
for dst in 0..ranks.len() {
if src == dst {
continue;
}
let mut can_access = 0;
unsafe {
cudarc::driver::sys::cuDeviceCanAccessPeer(
&mut can_access,
ranks[src].ctx().cu_device(),
ranks[dst].ctx().cu_device(),
)
.result()?;
}
if can_access == 0 {
return Err(format!(
"native TP requires P2P, but dev{} cannot access dev{}",
devices[src], devices[dst]
)
.into());
}
ranks[src].ctx().bind_to_thread()?;
let rc =
unsafe { cudarc::driver::sys::cuCtxEnablePeerAccess(ranks[dst].ctx().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!(
"native TP cuCtxEnablePeerAccess(dev{} -> dev{}) failed: {rc:?}",
devices[src], devices[dst]
)
.into());
}
}
}
for &owner in devices {
for &accessor in devices {
if owner == accessor {
continue;
}
let device = 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, device).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!(
"native TP cuMemPoolSetAccess(dev{owner} pool -> dev{accessor}) failed: \
{rc:?}"
)
.into());
}
}
}
for src in 0..ranks.len() {
for dst in 0..ranks.len() {
if src == dst {
continue;
}
let expected = (0..NATIVE_P2P_PROBE_WORDS)
.map(|index| {
(index as u32)
.wrapping_mul(0x9e37_79b9)
.wrapping_add(((src as u32) << 16) | dst as u32)
})
.collect::<Vec<_>>();
let poison = expected.iter().map(|value| !value).collect::<Vec<_>>();
let source = ranks[src].htod_u32_v(&expected)?;
let mut destination = ranks[dst].htod_u32_v(&poison)?;
ranks[dst].stream().memcpy_dtod(&source, &mut destination)?;
let actual = ranks[dst].dtoh_u32(&destination)?;
if actual != expected {
let mismatches = actual
.iter()
.zip(&expected)
.filter(|(actual, expected)| actual != expected)
.count();
return Err(format!(
"native TP peer probe dev{}->dev{} failed: {mismatches}/{} words differ",
devices[src],
devices[dst],
expected.len()
)
.into());
}
}
}
ranks[0].ctx().bind_to_thread()?;
eprintln!(
"[tp] native peer byte-integrity probe PASS: devices={devices:?} \
directions={} bytes={} mismatches=0",
ranks.len() * (ranks.len() - 1),
NATIVE_P2P_PROBE_WORDS * std::mem::size_of::<u32>(),
);
Ok(())
}
fn validate_activations(
activations: &[f32],
tokens: usize,
in_features: usize,
) -> Result<(), String> {
let expected = tokens
.checked_mul(in_features)
.ok_or_else(|| "activation size overflow".to_string())?;
if activations.len() != expected {
return Err(format!(
"activation count {} != {tokens}x{in_features} ({expected})",
activations.len()
));
}
if !activations.iter().all(|value| value.is_finite()) {
return Err("activations contain a non-finite value".to_string());
}
Ok(())
}
fn column_shard(
matrix: E4m3BlockMatrix<'_>,
tp: usize,
rank: usize,
) -> Result<E4m3BlockMatrix<'_>, String> {
let local_out = matrix.out_features / tp;
let row_start = rank * local_out;
let code_start = row_start * matrix.in_features;
let code_end = code_start + local_out * matrix.in_features;
let scale_cols = matrix.in_features.div_ceil(FP8_BLOCK);
let local_scale_rows = local_out / FP8_BLOCK;
let scale_start = rank * local_scale_rows * scale_cols;
let scale_end = scale_start + local_scale_rows * scale_cols;
Ok(E4m3BlockMatrix {
codes: &matrix.codes[code_start..code_end],
scales: &matrix.scales[scale_start..scale_end],
out_features: local_out,
in_features: matrix.in_features,
})
}
fn row_shard(
matrix: E4m3BlockMatrix<'_>,
tp: usize,
rank: usize,
) -> Result<(Vec<u8>, Vec<f32>), String> {
let local_in = matrix.in_features / tp;
let col_start = rank * local_in;
let mut codes = Vec::with_capacity(matrix.out_features * local_in);
for row in 0..matrix.out_features {
let start = row * matrix.in_features + col_start;
codes.extend_from_slice(&matrix.codes[start..start + local_in]);
}
let scale_rows = matrix.out_features.div_ceil(FP8_BLOCK);
let scale_cols = matrix.in_features.div_ceil(FP8_BLOCK);
let local_scale_cols = local_in / FP8_BLOCK;
let scale_col_start = rank * local_scale_cols;
let mut scales = Vec::with_capacity(scale_rows * local_scale_cols);
for row in 0..scale_rows {
let start = row * scale_cols + scale_col_start;
scales.extend_from_slice(&matrix.scales[start..start + local_scale_cols]);
}
Ok((codes, scales))
}
fn activation_shard(
activations: &[f32],
tokens: usize,
in_features: usize,
tp: usize,
rank: usize,
) -> Vec<f32> {
let local_in = in_features / tp;
let col_start = rank * local_in;
let mut shard = Vec::with_capacity(tokens * local_in);
for token in 0..tokens {
let start = token * in_features + col_start;
shard.extend_from_slice(&activations[start..start + local_in]);
}
shard
}
#[derive(Clone, Copy)]
pub struct Nvfp4BlockMatrix<'a> {
pub codes: &'a [u8], pub scales: &'a [u8], pub macro_scale: f32, pub out_features: usize,
pub in_features: usize,
}
impl Nvfp4BlockMatrix<'_> {
pub fn validate(&self) -> Result<(), String> {
if self.in_features == 0 || self.out_features == 0 {
return Err("NVFP4 matrix has a zero dimension".to_string());
}
if self.in_features % 64 != 0 {
return Err(format!(
"NVFP4 in_features {} is not 64-aligned (memra block_nvfp4 superblock)",
self.in_features
));
}
if self.codes.len() != self.out_features * self.in_features / 2 {
return Err(format!(
"NVFP4 code bytes {} != {}x{}/2",
self.codes.len(),
self.out_features,
self.in_features
));
}
if self.scales.len() != self.out_features * self.in_features / 16 {
return Err(format!(
"NVFP4 scale bytes {} != {}x{}/16",
self.scales.len(),
self.out_features,
self.in_features
));
}
if !self.macro_scale.is_finite() || self.macro_scale <= 0.0 {
return Err(format!(
"NVFP4 macro scale {} is not finite-positive",
self.macro_scale
));
}
Ok(())
}
}
#[derive(Clone, Copy)]
pub struct Nvfp4ExpertBank<'a> {
pub codes: &'a [u8], pub scales: &'a [u8], pub macros: &'a [f32], pub expert_count: usize,
pub out_features: usize,
pub in_features: usize,
}
impl Nvfp4ExpertBank<'_> {
pub fn validate(&self) -> Result<(), String> {
if self.expert_count == 0 {
return Err("NVFP4 expert bank is empty".to_string());
}
if self.macros.len() != self.expert_count {
return Err(format!(
"NVFP4 bank macros {} != expert count {}",
self.macros.len(),
self.expert_count
));
}
self.expert(0).map(|_| ())
}
pub fn expert(&self, expert: usize) -> Result<Nvfp4BlockMatrix<'_>, String> {
if expert >= self.expert_count {
return Err(format!("expert {expert} outside 0..{}", self.expert_count));
}
let code_stride = self.out_features * self.in_features / 2;
let scale_stride = self.out_features * self.in_features / 16;
if self.codes.len() != self.expert_count * code_stride
|| self.scales.len() != self.expert_count * scale_stride
{
return Err("NVFP4 bank byte extents do not match the declared geometry".to_string());
}
let matrix = Nvfp4BlockMatrix {
codes: &self.codes[expert * code_stride..(expert + 1) * code_stride],
scales: &self.scales[expert * scale_stride..(expert + 1) * scale_stride],
macro_scale: self.macros[expert],
out_features: self.out_features,
in_features: self.in_features,
};
matrix.validate()?;
Ok(matrix)
}
}
pub struct ResidentNvfp4Rank {
blocks: crate::CudaSlice<u8>,
macro_scale: f32,
out_features: usize,
in_features: usize,
row_bytes: usize,
}
pub struct ResidentNvfp4ColumnParallel {
ranks: Vec<ResidentNvfp4Rank>,
pub out_features: usize,
pub in_features: usize,
}
pub struct ResidentNvfp4RowParallel {
ranks: Vec<ResidentNvfp4Rank>,
pub out_features: usize,
pub in_features: usize,
}
pub struct ResidentTpNvfp4Expert {
gate: ResidentNvfp4ColumnParallel,
up: ResidentNvfp4ColumnParallel,
down: ResidentNvfp4RowParallel,
pub input_width: usize,
pub expert_width: usize,
}
pub struct ResidentNvfp4ColumnBankRank {
bank: crate::CudaSlice<u8>,
expert_bytes: usize,
local_out: usize,
in_features: usize,
row_bytes: usize,
}
impl ResidentNvfp4ColumnBankRank {
fn expert(&self, index: usize) -> cudarc::driver::CudaView<'_, u8> {
self.bank
.slice(index * self.expert_bytes..(index + 1) * self.expert_bytes)
}
}
pub const NVFP4_CANONICAL_ROW_SHARDS: usize = 2;
pub struct ResidentNvfp4RowBankRank {
bank: crate::CudaSlice<u8>,
expert_bytes: usize,
device_rank: usize, out_features: usize,
local_in: usize,
row_bytes: usize,
}
impl ResidentNvfp4RowBankRank {
fn expert(&self, index: usize) -> cudarc::driver::CudaView<'_, u8> {
self.bank
.slice(index * self.expert_bytes..(index + 1) * self.expert_bytes)
}
}
impl ResidentNvfp4TensorParallel {
pub(crate) fn device_workspace_handle(
&self,
) -> &std::sync::Mutex<Option<Nvfp4DeviceRoutesWorkspace>> {
&self.device_workspace
}
}
pub struct ResidentNvfp4TensorParallel {
gate: Vec<ResidentNvfp4ColumnBankRank>,
up: Vec<ResidentNvfp4ColumnBankRank>,
down: Vec<ResidentNvfp4RowBankRank>,
macros_gate: Vec<f32>,
macros_up: Vec<f32>,
macros_down: Vec<f32>,
macros_gate_dev: Vec<crate::CudaSlice<f32>>,
macros_up_dev: Vec<crate::CudaSlice<f32>>,
macros_down_dev: Vec<crate::CudaSlice<f32>>,
pub expert_count: usize,
pub input_width: usize,
pub expert_width: usize,
device_workspace: std::sync::Mutex<Option<Nvfp4DeviceRoutesWorkspace>>,
t2_workspace: std::sync::Mutex<Option<Nvfp4T2Workspace>>,
}
pub struct Nvfp4T2Workspace {
input2: Vec<crate::CudaSlice<f32>>,
in_q2: Vec<crate::CudaSlice<i8>>,
in_d2: Vec<crate::CudaSlice<f32>>,
sel2: Vec<crate::CudaSlice<i32>>,
route_w2: Vec<crate::CudaSlice<f32>>,
gate_out2: Vec<crate::CudaSlice<f32>>,
up_out2: Vec<crate::CudaSlice<f32>>,
act_q2: Vec<crate::CudaSlice<i8>>,
act_d2: Vec<crate::CudaSlice<f32>>,
partial2: Vec<crate::CudaSlice<f32>>,
acc_a: Vec<crate::CudaSlice<f32>>,
acc_b: Vec<crate::CudaSlice<f32>>,
peer_a: crate::CudaSlice<f32>,
peer_b: crate::CudaSlice<f32>,
omix_a: crate::CudaSlice<f32>,
omix_b: crate::CudaSlice<f32>,
ev_entry: CudaEvent,
ev_rank: Vec<CudaEvent>,
ev_root: CudaEvent,
n_sel: usize,
e_device: usize,
}
struct RoutesGraph {
exec: cudarc::driver::sys::CUgraphExec,
parent: cudarc::driver::sys::CUgraph,
_children: Vec<cudarc::driver::CudaGraph>,
}
unsafe impl Send for RoutesGraph {}
impl Drop for RoutesGraph {
fn drop(&mut self) {
unsafe {
let _ = cudarc::driver::sys::cuGraphExecDestroy(self.exec);
let _ = cudarc::driver::sys::cuGraphDestroy(self.parent);
}
}
}
impl Nvfp4DeviceRoutesWorkspace {
pub(crate) fn in_stage_handle(&self) -> Option<&crate::CudaSlice<f32>> {
self.in_stage_e.as_ref()
}
pub(crate) fn in_stage_mut(&mut self) -> Option<&mut crate::CudaSlice<f32>> {
self.in_stage_e.as_mut()
}
pub(crate) fn out_stage_mut(&mut self) -> Option<&mut crate::CudaSlice<f32>> {
self.out_stage_e.as_mut()
}
pub(crate) fn arm_stages(
&mut self,
e: &Engine,
width: usize,
n_sel: usize,
) -> Result<(), Box<dyn std::error::Error>> {
let _main = e.gpu.enter_main()?;
if self.in_stage_e.is_none() {
self.in_stage_e = Some(e.htod(&vec![0.0f32; width])?);
self.out_stage_e = Some(e.htod(&vec![0.0f32; width])?);
}
if self.dev_route_e.is_none() {
self.dev_route_e = Some((
e.htod_i32(&vec![0i32; n_sel])?,
e.htod(&vec![0.0f32; n_sel])?,
));
}
Ok(())
}
pub(crate) fn in_and_out_stages_mut(
&mut self,
) -> Option<(&crate::CudaSlice<f32>, &mut crate::CudaSlice<f32>)> {
match (self.in_stage_e.as_ref(), self.out_stage_e.as_mut()) {
(Some(input), Some(output)) => Some((input, output)),
_ => None,
}
}
pub(crate) fn dev_route_e_mut(
&mut self,
) -> Option<(&mut crate::CudaSlice<i32>, &mut crate::CudaSlice<f32>)> {
self.dev_route_e.as_mut().map(|(a, b)| (a, b))
}
}
pub struct Nvfp4DeviceRoutesWorkspace {
gate_out: Vec<crate::CudaSlice<f32>>,
up_out: Vec<crate::CudaSlice<f32>>,
act_q: Vec<crate::CudaSlice<i8>>,
act_d: Vec<crate::CudaSlice<f32>>,
sel: Vec<crate::CudaSlice<i32>>,
partial: Vec<crate::CudaSlice<f32>>,
accumulator: Vec<crate::CudaSlice<f32>>,
combine_w: Vec<crate::CudaSlice<f32>>,
route_w: Vec<crate::CudaSlice<f32>>,
in_q: Vec<crate::CudaSlice<i8>>,
in_d: Vec<crate::CudaSlice<f32>>,
dev_route_e: Option<(crate::CudaSlice<i32>, crate::CudaSlice<f32>)>,
prestaged: bool,
rank1_routed: bool,
fence_flags_raw: u64,
fence_ticket: u32,
ev_input: Option<(CudaEvent, usize)>,
in_stage_e: Option<crate::CudaSlice<f32>>,
out_stage_e: Option<crate::CudaSlice<f32>>,
routes_graph: Option<RoutesGraph>,
raw_dev_route_e: Option<(u64, u64)>,
raw_combine: Option<(u64, u64, u64, u64)>,
raw_input: Vec<u64>,
raw_sel: Vec<u64>,
raw_route_w: Vec<u64>,
remote: crate::CudaSlice<f32>,
combined: crate::CudaSlice<f32>,
n_sel: usize,
input: Vec<crate::CudaSlice<f32>>,
ev_rank: Vec<CudaEvent>,
ev_done: Option<CudaEvent>,
ev_entry: Option<(CudaEvent, usize)>,
}
struct ResidentNvfp4EpRank {
gate: Vec<crate::CudaSlice<u8>>,
up: Vec<crate::CudaSlice<u8>>,
down: Vec<crate::CudaSlice<u8>>,
#[allow(dead_code)]
expert_range: Range<usize>,
}
pub struct ResidentNvfp4ExpertParallel {
ranks: Vec<ResidentNvfp4EpRank>,
macros_gate: Vec<f32>,
macros_up: Vec<f32>,
macros_down: Vec<f32>,
pub expert_count: usize,
pub input_width: usize,
pub expert_width: usize,
gate_row_bytes: usize,
down_row_bytes: usize,
}
fn nvfp4_repack_matrix(matrix: Nvfp4BlockMatrix<'_>) -> Vec<u8> {
memra_gguf::nvfp4_repack::repack_modelopt_to_gguf(
matrix.codes,
matrix.scales,
matrix.out_features,
matrix.in_features,
)
}
fn nvfp4_row_bytes(in_features: usize) -> usize {
in_features / 64 * 36 }
pub(crate) fn fuse_rope_append_on() -> bool {
static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
*ON.get_or_init(|| std::env::var("MEMRA_FUSE_ROPE_APPEND").as_deref() == Ok("1"))
}
pub(crate) fn no_local_shadow_on() -> bool {
static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
*ON.get_or_init(|| std::env::var("MEMRA_NO_LOCAL_SHADOW").as_deref() == Ok("1"))
}
pub(crate) fn nvfp4_bank_v2_on() -> bool {
static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
*ON.get_or_init(|| std::env::var("MEMRA_NVFP4_BANK_V2").as_deref() == Ok("1"))
}
fn nvfp4_matrix_v2_permute(v1: &[u8], out_features: usize, in_features: usize) -> Vec<u8> {
let row_bytes = nvfp4_row_bytes(in_features);
assert_eq!(v1.len(), out_features * row_bytes, "v2 permute geometry");
let n_slots = in_features / 32;
let mut out = Vec::with_capacity(v1.len());
for row in 0..out_features {
let r = &v1[row * row_bytes..(row + 1) * row_bytes];
for g in 0..n_slots {
let (sblk, h) = (g / 2, g % 2);
let b = &r[sblk * 36..sblk * 36 + 36];
out.extend_from_slice(&b[4 + 16 * h..4 + 16 * h + 16]);
}
for g in 0..n_slots {
let (sblk, h) = (g / 2, g % 2);
let b = &r[sblk * 36..sblk * 36 + 36];
out.push(b[2 * h]);
out.push(b[2 * h + 1]);
}
}
out
}
fn nvfp4_repack_bank_matrix(matrix: Nvfp4BlockMatrix<'_>) -> Vec<u8> {
let (out_features, in_features) = (matrix.out_features, matrix.in_features);
let v1 = nvfp4_repack_matrix(matrix);
if nvfp4_bank_v2_on() {
nvfp4_matrix_v2_permute(&v1, out_features, in_features)
} else {
v1
}
}
fn nvfp4_column_shard<'a>(
matrix: Nvfp4BlockMatrix<'a>,
tp: usize,
rank: usize,
) -> Result<Nvfp4BlockMatrix<'a>, String> {
if matrix.out_features % tp != 0 {
return Err(format!(
"NVFP4 column-parallel out_features {} is not divisible by TP={tp}",
matrix.out_features
));
}
let local_out = matrix.out_features / tp;
let code_row = matrix.in_features / 2;
let scale_row = matrix.in_features / 16;
Ok(Nvfp4BlockMatrix {
codes: &matrix.codes[rank * local_out * code_row..(rank + 1) * local_out * code_row],
scales: &matrix.scales[rank * local_out * scale_row..(rank + 1) * local_out * scale_row],
macro_scale: matrix.macro_scale,
out_features: local_out,
in_features: matrix.in_features,
})
}
fn nvfp4_row_shard(
matrix: Nvfp4BlockMatrix<'_>,
tp: usize,
rank: usize,
) -> Result<(Vec<u8>, Vec<u8>, usize), String> {
if matrix.in_features % tp != 0 {
return Err(format!(
"NVFP4 row-parallel in_features {} is not divisible by TP={tp}",
matrix.in_features
));
}
let local_in = matrix.in_features / tp;
if local_in % 64 != 0 {
return Err(format!(
"NVFP4 row-parallel input shard {local_in} cuts through a 64-element superblock"
));
}
let code_row = matrix.in_features / 2;
let scale_row = matrix.in_features / 16;
let local_code = local_in / 2;
let local_scale = local_in / 16;
let mut codes = Vec::with_capacity(matrix.out_features * local_code);
let mut scales = Vec::with_capacity(matrix.out_features * local_scale);
for row in 0..matrix.out_features {
let code_start = row * code_row + rank * local_code;
codes.extend_from_slice(&matrix.codes[code_start..code_start + local_code]);
let scale_start = row * scale_row + rank * local_scale;
scales.extend_from_slice(&matrix.scales[scale_start..scale_start + local_scale]);
}
Ok((codes, scales, local_in))
}
fn run_rank_nvfp4(
engine: &Engine,
matrix: Nvfp4BlockMatrix<'_>,
activations: &[f32],
tokens: usize,
) -> Result<Vec<f32>, Box<dyn std::error::Error>> {
matrix.validate()?;
validate_activations(activations, tokens, matrix.in_features)?;
let _main = engine.gpu.enter_main()?;
let blocks = engine.htod_bytes(&nvfp4_repack_matrix(matrix))?;
let activations = engine.htod(activations)?;
let output = engine.qmatvec_nvfp4_fast(
&blocks.slice(0..blocks.len()),
&activations,
tokens,
matrix.in_features,
matrix.out_features,
nvfp4_row_bytes(matrix.in_features),
)?;
engine.dtoh(&output)
}
fn upload_rank_nvfp4(
engine: &Engine,
matrix: Nvfp4BlockMatrix<'_>,
) -> Result<ResidentNvfp4Rank, Box<dyn std::error::Error>> {
matrix.validate()?;
let _main = engine.gpu.enter_main()?;
Ok(ResidentNvfp4Rank {
blocks: engine.htod_bytes(&nvfp4_repack_matrix(matrix))?,
macro_scale: matrix.macro_scale,
out_features: matrix.out_features,
in_features: matrix.in_features,
row_bytes: nvfp4_row_bytes(matrix.in_features),
})
}
fn run_resident_rank_nvfp4(
engine: &Engine,
rank: &ResidentNvfp4Rank,
activations: &[f32],
tokens: usize,
) -> Result<Vec<f32>, Box<dyn std::error::Error>> {
validate_activations(activations, tokens, rank.in_features)?;
let _main = engine.gpu.enter_main()?;
let activations = engine.htod(activations)?;
let output = engine.qmatvec_nvfp4_fast(
&rank.blocks.slice(0..rank.blocks.len()),
&activations,
tokens,
rank.in_features,
rank.out_features,
rank.row_bytes,
)?;
engine.dtoh(&output)
}
fn apply_macro(values: &mut [f32], macro_scale: f32) {
for value in values.iter_mut() {
*value *= macro_scale;
}
}
impl TpE4m3HostBounce {
pub fn full_nvfp4(
&self,
matrix: Nvfp4BlockMatrix<'_>,
activations: &[f32],
tokens: usize,
) -> Result<Vec<f32>, Box<dyn std::error::Error>> {
let mut output = run_rank_nvfp4(&self.ranks[0], matrix, activations, tokens)?;
apply_macro(&mut output, matrix.macro_scale);
Ok(output)
}
pub fn column_parallel_nvfp4(
&self,
matrix: Nvfp4BlockMatrix<'_>,
activations: &[f32],
tokens: usize,
) -> Result<ColumnParallelResult, Box<dyn std::error::Error>> {
matrix.validate()?;
validate_activations(activations, tokens, matrix.in_features)?;
let tp = self.ranks.len();
let local_out = matrix.out_features / tp;
let mut gathered = vec![0.0f32; tokens * matrix.out_features];
let mut rank_outputs = Vec::with_capacity(tp);
for (rank_index, rank) in self.ranks.iter().enumerate() {
let shard = nvfp4_column_shard(matrix, tp, rank_index)?;
let output = run_rank_nvfp4(rank, shard, activations, tokens)?;
let row_start = rank_index * local_out;
for token in 0..tokens {
gathered[token * matrix.out_features + row_start
..token * matrix.out_features + row_start + local_out]
.copy_from_slice(&output[token * local_out..(token + 1) * local_out]);
}
rank_outputs.push(output);
}
apply_macro(&mut gathered, matrix.macro_scale);
Ok(ColumnParallelResult {
gathered,
rank_outputs,
})
}
pub fn row_parallel_nvfp4(
&self,
matrix: Nvfp4BlockMatrix<'_>,
activations: &[f32],
tokens: usize,
) -> Result<RowParallelResult, Box<dyn std::error::Error>> {
matrix.validate()?;
validate_activations(activations, tokens, matrix.in_features)?;
let tp = self.ranks.len();
let mut reduced = vec![0.0f32; tokens * matrix.out_features];
let mut rank_partials = Vec::with_capacity(tp);
for (rank_index, rank) in self.ranks.iter().enumerate() {
let (codes, scales, local_in) = nvfp4_row_shard(matrix, tp, rank_index)?;
let local_activations =
activation_shard(activations, tokens, matrix.in_features, tp, rank_index);
let shard = Nvfp4BlockMatrix {
codes: &codes,
scales: &scales,
macro_scale: matrix.macro_scale,
out_features: matrix.out_features,
in_features: local_in,
};
let partial = run_rank_nvfp4(rank, shard, &local_activations, tokens)?;
for (sum, value) in reduced.iter_mut().zip(&partial) {
*sum += *value;
}
rank_partials.push(partial);
}
apply_macro(&mut reduced, matrix.macro_scale);
Ok(RowParallelResult {
reduced,
rank_partials,
})
}
pub fn upload_expert_nvfp4(
&self,
gate: Nvfp4BlockMatrix<'_>,
up: Nvfp4BlockMatrix<'_>,
down: Nvfp4BlockMatrix<'_>,
) -> Result<ResidentTpNvfp4Expert, Box<dyn std::error::Error>> {
if gate.in_features != up.in_features || gate.out_features != up.out_features {
return Err("NVFP4 TP expert gate/up dimensions differ".into());
}
if down.in_features != gate.out_features || down.out_features != gate.in_features {
return Err(format!(
"NVFP4 TP expert down {}x{} does not invert gate/up {}x{}",
down.out_features, down.in_features, gate.out_features, gate.in_features
)
.into());
}
let tp = self.ranks.len();
let mut gate_ranks = Vec::with_capacity(tp);
let mut up_ranks = Vec::with_capacity(tp);
let mut down_ranks = Vec::with_capacity(tp);
for (rank_index, engine) in self.ranks.iter().enumerate() {
gate_ranks.push(upload_rank_nvfp4(
engine,
nvfp4_column_shard(gate, tp, rank_index)?,
)?);
up_ranks.push(upload_rank_nvfp4(
engine,
nvfp4_column_shard(up, tp, rank_index)?,
)?);
let (codes, scales, local_in) = nvfp4_row_shard(down, tp, rank_index)?;
down_ranks.push(upload_rank_nvfp4(
engine,
Nvfp4BlockMatrix {
codes: &codes,
scales: &scales,
macro_scale: down.macro_scale,
out_features: down.out_features,
in_features: local_in,
},
)?);
}
Ok(ResidentTpNvfp4Expert {
gate: ResidentNvfp4ColumnParallel {
ranks: gate_ranks,
out_features: gate.out_features,
in_features: gate.in_features,
},
up: ResidentNvfp4ColumnParallel {
ranks: up_ranks,
out_features: up.out_features,
in_features: up.in_features,
},
down: ResidentNvfp4RowParallel {
ranks: down_ranks,
out_features: down.out_features,
in_features: down.in_features,
},
input_width: gate.in_features,
expert_width: gate.out_features,
})
}
fn column_parallel_resident_nvfp4(
&self,
matrix: &ResidentNvfp4ColumnParallel,
activations: &[f32],
tokens: usize,
) -> Result<Vec<f32>, Box<dyn std::error::Error>> {
validate_activations(activations, tokens, matrix.in_features)?;
let local_out = matrix.out_features / self.ranks.len();
let mut gathered = vec![0.0f32; tokens * matrix.out_features];
let mut macro_scale = None;
for (rank_index, (engine, shard)) in self.ranks.iter().zip(&matrix.ranks).enumerate() {
let output = run_resident_rank_nvfp4(engine, shard, activations, tokens)?;
let row_start = rank_index * local_out;
for token in 0..tokens {
gathered[token * matrix.out_features + row_start
..token * matrix.out_features + row_start + local_out]
.copy_from_slice(&output[token * local_out..(token + 1) * local_out]);
}
macro_scale = Some(shard.macro_scale);
}
apply_macro(
&mut gathered,
macro_scale.ok_or("NVFP4 column-parallel matrix has no ranks")?,
);
Ok(gathered)
}
fn row_parallel_resident_nvfp4(
&self,
matrix: &ResidentNvfp4RowParallel,
activations: &[f32],
tokens: usize,
) -> Result<Vec<f32>, Box<dyn std::error::Error>> {
validate_activations(activations, tokens, matrix.in_features)?;
let tp = self.ranks.len();
let local_in = matrix.in_features / tp;
let mut reduced = vec![0.0f32; tokens * matrix.out_features];
let mut macro_scale = None;
for (rank_index, (engine, shard)) in self.ranks.iter().zip(&matrix.ranks).enumerate() {
if shard.in_features != local_in {
return Err(format!(
"NVFP4 resident row shard in_features {} != expected {local_in}",
shard.in_features
)
.into());
}
let local_activations =
activation_shard(activations, tokens, matrix.in_features, tp, rank_index);
let partial = run_resident_rank_nvfp4(engine, shard, &local_activations, tokens)?;
for (sum, value) in reduced.iter_mut().zip(&partial) {
*sum += *value;
}
macro_scale = Some(shard.macro_scale);
}
apply_macro(
&mut reduced,
macro_scale.ok_or("NVFP4 row-parallel matrix has no ranks")?,
);
Ok(reduced)
}
pub fn run_expert_nvfp4(
&self,
expert: &ResidentTpNvfp4Expert,
input: &[f32],
tokens: usize,
) -> Result<Vec<f32>, Box<dyn std::error::Error>> {
validate_activations(input, tokens, expert.input_width)?;
let gate = self.column_parallel_resident_nvfp4(&expert.gate, input, tokens)?;
let up = self.column_parallel_resident_nvfp4(&expert.up, input, tokens)?;
let activated: Vec<f32> = gate
.iter()
.zip(&up)
.map(|(&gate, &up)| gate / (1.0 + (-gate).exp()) * up)
.collect();
debug_assert_eq!(activated.len(), tokens * expert.expert_width);
self.row_parallel_resident_nvfp4(&expert.down, &activated, tokens)
}
pub fn upload_tensor_parallel_nvfp4(
&self,
gate: Nvfp4ExpertBank<'_>,
up: Nvfp4ExpertBank<'_>,
down: Nvfp4ExpertBank<'_>,
) -> Result<ResidentNvfp4TensorParallel, Box<dyn std::error::Error>> {
gate.validate()?;
up.validate()?;
down.validate()?;
if gate.expert_count != up.expert_count || gate.expert_count != down.expert_count {
return Err("NVFP4 TP gate/up/down expert counts differ".into());
}
if gate.in_features != up.in_features || gate.out_features != up.out_features {
return Err("NVFP4 TP gate/up dimensions differ".into());
}
if down.in_features != gate.out_features || down.out_features != gate.in_features {
return Err(format!(
"NVFP4 TP down {}x{} does not invert gate/up {}x{}",
down.out_features, down.in_features, gate.out_features, gate.in_features
)
.into());
}
let tp = self.ranks.len();
if gate.out_features % tp != 0 {
return Err(format!(
"NVFP4 TP expert output width {} is not divisible by TP={tp}",
gate.out_features
)
.into());
}
if down.in_features % NVFP4_CANONICAL_ROW_SHARDS != 0
|| (down.in_features / NVFP4_CANONICAL_ROW_SHARDS) % 64 != 0
{
return Err(format!(
"NVFP4 TP expert input width {} does not split into 64-aligned canonical \
shards ({NVFP4_CANONICAL_ROW_SHARDS})",
down.in_features
)
.into());
}
if tp > NVFP4_CANONICAL_ROW_SHARDS {
return Err(format!(
"NVFP4 TP world {tp} exceeds the canonical row-shard grid \
({NVFP4_CANONICAL_ROW_SHARDS})"
)
.into());
}
let mut gate_ranks = Vec::with_capacity(tp);
let mut up_ranks = Vec::with_capacity(tp);
let mut macros_gate_dev = Vec::with_capacity(tp);
let mut macros_up_dev = Vec::with_capacity(tp);
let mut macros_down_dev = Vec::with_capacity(tp);
for (rank_index, engine) in self.ranks.iter().enumerate() {
let _main = engine.gpu.enter_main()?;
let mut gate_host: Vec<u8> = Vec::new();
let mut up_host: Vec<u8> = Vec::new();
for expert in 0..gate.expert_count {
let gate_shard = nvfp4_column_shard(gate.expert(expert)?, tp, rank_index)?;
gate_host.extend_from_slice(&nvfp4_repack_bank_matrix(gate_shard));
let up_shard = nvfp4_column_shard(up.expert(expert)?, tp, rank_index)?;
up_host.extend_from_slice(&nvfp4_repack_bank_matrix(up_shard));
}
let gate_expert_bytes = gate_host.len() / gate.expert_count;
let up_expert_bytes = up_host.len() / up.expert_count;
gate_ranks.push(ResidentNvfp4ColumnBankRank {
bank: engine.htod_bytes(&gate_host)?,
expert_bytes: gate_expert_bytes,
local_out: gate.out_features / tp,
in_features: gate.in_features,
row_bytes: nvfp4_row_bytes(gate.in_features),
});
up_ranks.push(ResidentNvfp4ColumnBankRank {
bank: engine.htod_bytes(&up_host)?,
expert_bytes: up_expert_bytes,
local_out: up.out_features / tp,
in_features: up.in_features,
row_bytes: nvfp4_row_bytes(up.in_features),
});
macros_gate_dev.push(engine.htod(gate.macros)?);
macros_up_dev.push(engine.htod(up.macros)?);
macros_down_dev.push(engine.htod(down.macros)?);
}
let mut down_ranks = Vec::with_capacity(NVFP4_CANONICAL_ROW_SHARDS);
for shard_index in 0..NVFP4_CANONICAL_ROW_SHARDS {
let device_rank = shard_index % tp;
let engine = &self.ranks[device_rank];
let _main = engine.gpu.enter_main()?;
let mut down_host: Vec<u8> = Vec::new();
for expert in 0..down.expert_count {
let down_matrix = down.expert(expert)?;
let (codes, scales, local_in) =
nvfp4_row_shard(down_matrix, NVFP4_CANONICAL_ROW_SHARDS, shard_index)?;
down_host.extend_from_slice(&nvfp4_repack_bank_matrix(Nvfp4BlockMatrix {
codes: &codes,
scales: &scales,
macro_scale: down_matrix.macro_scale,
out_features: down_matrix.out_features,
in_features: local_in,
}));
}
let down_expert_bytes = down_host.len() / down.expert_count;
down_ranks.push(ResidentNvfp4RowBankRank {
bank: engine.htod_bytes(&down_host)?,
expert_bytes: down_expert_bytes,
device_rank,
out_features: down.out_features,
local_in: down.in_features / NVFP4_CANONICAL_ROW_SHARDS,
row_bytes: nvfp4_row_bytes(down.in_features / NVFP4_CANONICAL_ROW_SHARDS),
});
}
Ok(ResidentNvfp4TensorParallel {
gate: gate_ranks,
up: up_ranks,
down: down_ranks,
macros_gate: gate.macros.to_vec(),
macros_up: up.macros.to_vec(),
macros_down: down.macros.to_vec(),
macros_gate_dev,
macros_up_dev,
macros_down_dev,
expert_count: gate.expert_count,
input_width: gate.in_features,
expert_width: gate.out_features,
device_workspace: std::sync::Mutex::new(None),
t2_workspace: std::sync::Mutex::new(None),
})
}
fn run_column_bank_expert_nvfp4(
&self,
ranks: &[ResidentNvfp4ColumnBankRank],
macros: &[f32],
expert: usize,
input: &[f32],
) -> Result<Vec<f32>, Box<dyn std::error::Error>> {
let local_out = ranks
.first()
.ok_or("NVFP4 TP column bank has no ranks")?
.local_out;
let mut gathered = vec![0.0f32; local_out * ranks.len()];
for (rank_index, (engine, bank)) in self.ranks.iter().zip(ranks).enumerate() {
let _main = engine.gpu.enter_main()?;
let activations = engine.htod(input)?;
let output = if nvfp4_bank_v2_on() {
engine.qmatvec_nvfp4_fast_v2(
&bank.expert(expert),
&activations,
1,
bank.in_features,
bank.local_out,
bank.row_bytes,
)?
} else {
engine.qmatvec_nvfp4_fast(
&bank.expert(expert),
&activations,
1,
bank.in_features,
bank.local_out,
bank.row_bytes,
)?
};
let output = engine.dtoh(&output)?;
gathered[rank_index * local_out..(rank_index + 1) * local_out].copy_from_slice(&output);
}
apply_macro(&mut gathered, macros[expert]);
Ok(gathered)
}
fn run_row_bank_expert_nvfp4(
&self,
shards: &[ResidentNvfp4RowBankRank],
macros: &[f32],
expert: usize,
input: &[f32],
) -> Result<Vec<f32>, Box<dyn std::error::Error>> {
let out_features = shards
.first()
.ok_or("NVFP4 TP row bank has no canonical shards")?
.out_features;
let in_features = shards.iter().map(|shard| shard.local_in).sum::<usize>();
let mut reduced = vec![0.0f32; out_features];
for (shard_index, shard) in shards.iter().enumerate() {
let engine = self
.ranks
.get(shard.device_rank)
.ok_or("NVFP4 canonical shard names a rank outside this runtime")?;
let _main = engine.gpu.enter_main()?;
let local_activations =
activation_shard(input, 1, in_features, shards.len(), shard_index);
let activations = engine.htod(&local_activations)?;
let output = if nvfp4_bank_v2_on() {
engine.qmatvec_nvfp4_fast_v2(
&shard.expert(expert),
&activations,
1,
shard.local_in,
shard.out_features,
shard.row_bytes,
)?
} else {
engine.qmatvec_nvfp4_fast(
&shard.expert(expert),
&activations,
1,
shard.local_in,
shard.out_features,
shard.row_bytes,
)?
};
let partial = engine.dtoh(&output)?;
for (sum, value) in reduced.iter_mut().zip(&partial) {
*sum += *value;
}
}
apply_macro(&mut reduced, macros[expert]);
Ok(reduced)
}
pub fn upload_expert_parallel_nvfp4(
&self,
gate: Nvfp4ExpertBank<'_>,
up: Nvfp4ExpertBank<'_>,
down: Nvfp4ExpertBank<'_>,
) -> Result<ResidentNvfp4ExpertParallel, Box<dyn std::error::Error>> {
gate.validate()?;
up.validate()?;
down.validate()?;
if gate.expert_count != up.expert_count || gate.expert_count != down.expert_count {
return Err("NVFP4 EP gate/up/down expert counts differ".into());
}
if gate.in_features != up.in_features || gate.out_features != up.out_features {
return Err("NVFP4 EP gate/up dimensions differ".into());
}
if down.in_features != gate.out_features || down.out_features != gate.in_features {
return Err(format!(
"NVFP4 EP down {}x{} does not invert gate/up {}x{}",
down.out_features, down.in_features, gate.out_features, gate.in_features
)
.into());
}
let world = self.ranks.len();
if gate.expert_count % world != 0 {
return Err(format!(
"NVFP4 EP expert count {} is not divisible by {world} ranks",
gate.expert_count
)
.into());
}
let experts_per_rank = gate.expert_count / world;
let mut ranks = Vec::with_capacity(world);
for (rank_index, engine) in self.ranks.iter().enumerate() {
let _main = engine.gpu.enter_main()?;
let expert_range = rank_index * experts_per_rank..(rank_index + 1) * experts_per_rank;
let mut gate_experts = Vec::with_capacity(experts_per_rank);
let mut up_experts = Vec::with_capacity(experts_per_rank);
let mut down_experts = Vec::with_capacity(experts_per_rank);
for expert in expert_range.clone() {
gate_experts.push(engine.htod_bytes(&nvfp4_repack_matrix(gate.expert(expert)?))?);
up_experts.push(engine.htod_bytes(&nvfp4_repack_matrix(up.expert(expert)?))?);
down_experts.push(engine.htod_bytes(&nvfp4_repack_matrix(down.expert(expert)?))?);
}
ranks.push(ResidentNvfp4EpRank {
gate: gate_experts,
up: up_experts,
down: down_experts,
expert_range,
});
}
Ok(ResidentNvfp4ExpertParallel {
ranks,
macros_gate: gate.macros.to_vec(),
macros_up: up.macros.to_vec(),
macros_down: down.macros.to_vec(),
expert_count: gate.expert_count,
input_width: gate.in_features,
expert_width: gate.out_features,
gate_row_bytes: nvfp4_row_bytes(gate.in_features),
down_row_bytes: nvfp4_row_bytes(down.in_features),
})
}
#[allow(clippy::too_many_arguments)]
pub fn run_routed_experts_nvfp4(
&self,
experts: &ResidentNvfp4ExpertParallel,
input: &[f32],
tokens: usize,
selected: &[usize],
route_weights: &[f32],
experts_per_token: usize,
activation_limit: Option<f32>,
) -> Result<Vec<f32>, Box<dyn std::error::Error>> {
validate_activations(input, tokens, experts.input_width)?;
let pairs = tokens
.checked_mul(experts_per_token)
.ok_or("NVFP4 EP route count overflow")?;
if selected.len() != pairs || route_weights.len() != pairs {
return Err(format!(
"NVFP4 EP routes selected={} weights={} != tokens {tokens} x experts/token \
{experts_per_token} ({pairs})",
selected.len(),
route_weights.len(),
)
.into());
}
if !route_weights.iter().all(|weight| weight.is_finite()) {
return Err("NVFP4 EP route weights contain a non-finite value".into());
}
let experts_per_rank = experts.expert_count / experts.ranks.len();
let mut output = vec![0.0f32; tokens * experts.input_width];
for token in 0..tokens {
let input_row = &input[token * experts.input_width..(token + 1) * experts.input_width];
for slot in 0..experts_per_token {
let pair = token * experts_per_token + slot;
let expert = selected[pair];
if expert >= experts.expert_count {
return Err(format!(
"NVFP4 EP selected expert {expert} outside 0..{}",
experts.expert_count
)
.into());
}
let owner = expert / experts_per_rank;
let local = expert - owner * experts_per_rank;
let rank = &experts.ranks[owner];
let engine = &self.ranks[owner];
let _main = engine.gpu.enter_main()?;
let device_input = engine.htod(input_row)?;
let gate_out = engine.qmatvec_nvfp4_fast(
&rank.gate[local].slice(0..rank.gate[local].len()),
&device_input,
1,
experts.input_width,
experts.expert_width,
experts.gate_row_bytes,
)?;
let up_out = engine.qmatvec_nvfp4_fast(
&rank.up[local].slice(0..rank.up[local].len()),
&device_input,
1,
experts.input_width,
experts.expert_width,
experts.gate_row_bytes,
)?;
let mut gate_host = engine.dtoh(&gate_out)?;
let mut up_host = engine.dtoh(&up_out)?;
apply_macro(&mut gate_host, experts.macros_gate[expert]);
apply_macro(&mut up_host, experts.macros_up[expert]);
let activated: Vec<f32> = gate_host
.iter()
.zip(&up_host)
.map(|(&gate, &up)| step_expert_activation_host(gate, up, activation_limit))
.collect();
let device_activated = engine.htod(&activated)?;
let down_out = engine.qmatvec_nvfp4_fast(
&rank.down[local].slice(0..rank.down[local].len()),
&device_activated,
1,
experts.expert_width,
experts.input_width,
experts.down_row_bytes,
)?;
let mut down_host = engine.dtoh(&down_out)?;
apply_macro(&mut down_host, experts.macros_down[expert]);
let weight = route_weights[pair];
for (sum, value) in output
[token * experts.input_width..(token + 1) * experts.input_width]
.iter_mut()
.zip(down_host)
{
*sum += weight * value;
}
}
}
Ok(output)
}
pub fn run_tensor_parallel_routes_nvfp4_device(
&self,
experts: &ResidentNvfp4TensorParallel,
input: &[f32],
selected: &[usize],
route_weights: &[f32],
experts_per_token: usize,
activation_limit: Option<f32>,
) -> Result<Vec<f32>, Box<dyn std::error::Error>> {
validate_activations(input, 1, experts.input_width)?;
if selected.len() != experts_per_token || route_weights.len() != experts_per_token {
return Err(format!(
"NVFP4 device routes selected={} weights={} != experts/token {experts_per_token}",
selected.len(),
route_weights.len(),
)
.into());
}
if !route_weights.iter().all(|weight| weight.is_finite()) {
return Err("NVFP4 device route weights contain a non-finite value".into());
}
let world = self.ranks.len();
if world != NVFP4_CANONICAL_ROW_SHARDS {
return Err(format!(
"NVFP4 device routes require world == canonical shard grid \
({NVFP4_CANONICAL_ROW_SHARDS}), got {world}"
)
.into());
}
let local_out = experts.expert_width / world;
static TIMING_NS: std::sync::atomic::AtomicU64 = std::sync::atomic::AtomicU64::new(0);
static TIMING_CALLS: std::sync::atomic::AtomicU64 = std::sync::atomic::AtomicU64::new(0);
let timing = std::env::var("MEMRA_STEP_TP_TIMING").as_deref() == Ok("1");
let started = timing.then(std::time::Instant::now);
let n_sel = experts_per_token;
let mut workspace_guard = experts
.device_workspace
.lock()
.map_err(|_| "NVFP4 device routes workspace lock is poisoned")?;
if workspace_guard.is_none() {
let mut gate_out = Vec::with_capacity(world);
let mut up_out = Vec::with_capacity(world);
let mut act_q = Vec::with_capacity(world);
let mut act_d = Vec::with_capacity(world);
let mut sel = Vec::with_capacity(world);
let mut partial = Vec::with_capacity(world);
let mut accumulator = Vec::with_capacity(world);
let mut combine_w = Vec::with_capacity(world);
let mut route_w = Vec::with_capacity(world);
let mut in_q = Vec::with_capacity(world);
let mut in_d = Vec::with_capacity(world);
let mut input = Vec::with_capacity(world);
let mut ev_rank = Vec::with_capacity(world);
let moe_direct = moe_direct_on();
for (rank, engine) in self.ranks.iter().enumerate() {
let _main = engine.gpu.enter_main()?;
gate_out.push(engine.uninit(n_sel * local_out)?);
up_out.push(engine.uninit(n_sel * local_out)?);
act_q.push(engine.uninit_i8(n_sel * local_out)?);
act_d.push(engine.uninit(n_sel * local_out / 32)?);
sel.push(engine.htod_i32(&vec![0i32; n_sel])?);
partial.push(engine.uninit(n_sel * experts.input_width)?);
if moe_direct && rank != 0 {
let root = &self.ranks[0];
let _root_main = root.gpu.enter_main()?;
accumulator.push(root.zeros(experts.input_width)?);
} else {
accumulator.push(engine.zeros(experts.input_width)?);
}
combine_w.push(engine.htod(&vec![0.0f32; n_sel])?);
route_w.push(engine.htod(&vec![0.0f32; n_sel])?);
in_q.push(engine.uninit_i8(experts.input_width)?);
in_d.push(engine.uninit(experts.input_width / 32)?);
input.push(engine.uninit(experts.input_width)?);
ev_rank.push(engine.ctx().new_event(None)?);
}
let root = &self.ranks[0];
let _main = root.gpu.enter_main()?;
*workspace_guard = Some(Nvfp4DeviceRoutesWorkspace {
prestaged: false,
rank1_routed: false,
ev_input: None,
fence_flags_raw: 0,
fence_ticket: 0,
gate_out,
up_out,
act_q,
act_d,
sel,
partial,
accumulator,
combine_w,
route_w,
in_q,
in_d,
dev_route_e: None,
in_stage_e: None,
out_stage_e: None,
routes_graph: None,
raw_dev_route_e: None,
raw_combine: None,
raw_input: Vec::new(),
raw_sel: Vec::new(),
raw_route_w: Vec::new(),
remote: root.uninit(experts.input_width)?,
combined: root.uninit(experts.input_width)?,
n_sel,
input,
ev_rank,
ev_done: Some(root.ctx().new_event(None)?),
ev_entry: None,
});
}
let workspace = workspace_guard
.as_mut()
.expect("NVFP4 device routes workspace initialized above");
if workspace.n_sel != n_sel {
return Err(format!(
"NVFP4 device routes experts/token changed: workspace {} != call {n_sel}",
workspace.n_sel
)
.into());
}
for &expert in selected {
if expert >= experts.expert_count {
return Err(format!(
"NVFP4 device selected expert {expert} outside 0..{}",
experts.expert_count
)
.into());
}
}
let sel_i32 = selected
.iter()
.map(|&expert| expert as i32)
.collect::<Vec<_>>();
for (rank_index, engine) in self.ranks.iter().enumerate() {
let _main = engine.gpu.enter_main()?;
let device_input = engine.htod(input)?;
let Nvfp4DeviceRoutesWorkspace { in_q, in_d, .. } = &mut *workspace;
engine.quantize_q8_1_into(
&device_input,
1,
experts.input_width,
&mut in_q[rank_index],
&mut in_d[rank_index],
)?;
}
self.nvfp4_routes_batched_sweeps(
experts,
workspace,
selected,
route_weights,
&sel_i32,
local_out,
n_sel,
activation_limit,
false,
)?;
let root = &self.ranks[0];
for engine in &self.ranks[1..] {
let _main = engine.gpu.enter_main()?;
engine.stream().synchronize()?;
}
let _main = root.gpu.enter_main()?;
root.stream()
.memcpy_dtod(&workspace.accumulator[1], &mut workspace.remote)?;
root.add(
&workspace.accumulator[0],
&workspace.remote,
&mut workspace.combined,
experts.input_width,
)?;
let output = root.dtoh(&workspace.combined)?;
if let Some(started) = started {
use std::sync::atomic::Ordering;
let ns = TIMING_NS.fetch_add(started.elapsed().as_nanos() as u64, Ordering::Relaxed)
+ started.elapsed().as_nanos() as u64;
let calls = TIMING_CALLS.fetch_add(1, Ordering::Relaxed) + 1;
if calls % 430 == 0 {
eprintln!(
"[nvfp4-dev-routes-timing] calls={calls} total_ms={:.1} avg_us={:.1}",
ns as f64 / 1.0e6,
ns as f64 / calls as f64 / 1.0e3,
);
}
}
Ok(output)
}
#[allow(clippy::too_many_arguments)]
fn nvfp4_routes_batched_sweeps(
&self,
experts: &ResidentNvfp4TensorParallel,
workspace: &mut Nvfp4DeviceRoutesWorkspace,
selected: &[usize],
route_weights: &[f32],
sel_i32: &[i32],
local_out: usize,
n_sel: usize,
activation_limit: Option<f32>,
device_routed: bool,
) -> Result<(), Box<dyn std::error::Error>> {
for rank_index in 0..self.ranks.len() {
self.nvfp4_routes_batched_sweeps_rank(
experts,
workspace,
selected,
route_weights,
sel_i32,
local_out,
n_sel,
activation_limit,
device_routed,
rank_index,
)?;
}
Ok(())
}
#[allow(clippy::too_many_arguments)]
fn nvfp4_routes_batched_sweeps_rank(
&self,
experts: &ResidentNvfp4TensorParallel,
workspace: &mut Nvfp4DeviceRoutesWorkspace,
selected: &[usize],
route_weights: &[f32],
sel_i32: &[i32],
local_out: usize,
n_sel: usize,
activation_limit: Option<f32>,
device_routed: bool,
rank_index: usize,
) -> Result<(), Box<dyn std::error::Error>> {
{
let engine = &self.ranks[rank_index];
let _main = engine.gpu.enter_main()?;
if !device_routed {
engine.htod_i32_into(&mut workspace.sel[rank_index], sel_i32)?;
let folded = (0..n_sel)
.map(|pair| route_weights[pair] * experts.macros_down[selected[pair]])
.collect::<Vec<_>>();
let mut view = workspace.combine_w[rank_index].slice_mut(0..n_sel);
engine.stream().memcpy_htod(&folded, &mut view)?;
}
let gate_bank = &experts.gate[rank_index];
let up_bank = &experts.up[rank_index];
let (aq, ad) = (&workspace.in_q[rank_index], &workspace.in_d[rank_index]);
let gu_fused = nvfp4_bank_v2_on()
&& gate_bank.in_features == up_bank.in_features
&& gate_bank.local_out == up_bank.local_out
&& gate_bank.row_bytes == up_bank.row_bytes
&& gate_bank.expert_bytes == up_bank.expert_bytes;
if gu_fused {
let Nvfp4DeviceRoutesWorkspace {
sel,
gate_out,
up_out,
in_q,
in_d,
..
} = &mut *workspace;
engine.qmatvec_nvfp4_sel_gu_into(
&gate_bank.bank,
&up_bank.bank,
&sel[rank_index],
&in_q[rank_index],
&in_d[rank_index],
&mut gate_out[rank_index],
&mut up_out[rank_index],
n_sel,
gate_bank.in_features,
gate_bank.local_out,
gate_bank.row_bytes,
gate_bank.expert_bytes,
)?;
} else {
engine.qmatvec_nvfp4_sel_into(
&gate_bank.bank,
&workspace.sel[rank_index],
aq,
ad,
&mut workspace.gate_out[rank_index],
n_sel,
gate_bank.in_features,
gate_bank.local_out,
gate_bank.row_bytes,
gate_bank.expert_bytes,
0,
0,
)?;
engine.qmatvec_nvfp4_sel_into(
&up_bank.bank,
&workspace.sel[rank_index],
aq,
ad,
&mut workspace.up_out[rank_index],
n_sel,
up_bank.in_features,
up_bank.local_out,
up_bank.row_bytes,
up_bank.expert_bytes,
0,
0,
)?;
}
{
let Nvfp4DeviceRoutesWorkspace {
gate_out,
up_out,
sel,
act_q,
act_d,
..
} = &mut *workspace;
engine.silu_mul_scaled_q8_1_sel_into(
&gate_out[rank_index],
&up_out[rank_index],
&experts.macros_gate_dev[rank_index],
&experts.macros_up_dev[rank_index],
&sel[rank_index],
activation_limit,
&mut act_q[rank_index],
&mut act_d[rank_index],
local_out,
n_sel,
)?;
}
let shard = &experts.down[rank_index];
if shard.device_rank != rank_index || shard.local_in != local_out {
return Err(
"NVFP4 device routes: down canonical shard placement drifted from \
the gate/up column split"
.into(),
);
}
{
let Nvfp4DeviceRoutesWorkspace {
sel,
act_q,
act_d,
partial,
..
} = &mut *workspace;
engine.qmatvec_nvfp4_sel_into(
&shard.bank,
&sel[rank_index],
&act_q[rank_index],
&act_d[rank_index],
&mut partial[rank_index],
n_sel,
shard.local_in,
shard.out_features,
shard.row_bytes,
shard.expert_bytes,
local_out,
local_out / 32,
)?;
}
{
let Nvfp4DeviceRoutesWorkspace {
partial,
combine_w,
route_w,
sel,
accumulator,
..
} = &mut *workspace;
if device_routed {
engine.axpy_rows_seq_md_into(
&partial[rank_index],
&route_w[rank_index],
&experts.macros_down_dev[rank_index],
&sel[rank_index],
&mut accumulator[rank_index],
experts.input_width,
n_sel,
)?;
} else {
engine.axpy_rows_seq_into(
&partial[rank_index],
&combine_w[rank_index],
&mut accumulator[rank_index],
experts.input_width,
n_sel,
)?;
}
}
}
Ok(())
}
pub fn run_tensor_parallel_routes_nvfp4_device_io(
&self,
experts: &ResidentNvfp4TensorParallel,
e: &Engine,
input_dev: &crate::CudaSlice<f32>,
selected: &[usize],
route_weights: &[f32],
experts_per_token: usize,
activation_limit: Option<f32>,
) -> Result<crate::CudaSlice<f32>, Box<dyn std::error::Error>> {
if input_dev.len() != experts.input_width {
return Err(format!(
"NVFP4 device-io routes input {} != width {}",
input_dev.len(),
experts.input_width
)
.into());
}
if selected.len() != experts_per_token || route_weights.len() != experts_per_token {
return Err(format!(
"NVFP4 device-io routes selected={} weights={} != experts/token {experts_per_token}",
selected.len(),
route_weights.len(),
)
.into());
}
if !route_weights.iter().all(|weight| weight.is_finite()) {
return Err("NVFP4 device route weights contain a non-finite value".into());
}
let world = self.ranks.len();
if world != NVFP4_CANONICAL_ROW_SHARDS {
return Err(format!(
"NVFP4 device routes require world == canonical shard grid \
({NVFP4_CANONICAL_ROW_SHARDS}), got {world}"
)
.into());
}
let local_out = experts.expert_width / world;
let n_sel = experts_per_token;
static TIMING_NS: std::sync::atomic::AtomicU64 = std::sync::atomic::AtomicU64::new(0);
static TIMING_CALLS: std::sync::atomic::AtomicU64 = std::sync::atomic::AtomicU64::new(0);
let timing = std::env::var("MEMRA_STEP_TP_TIMING").as_deref() == Ok("1");
let started = timing.then(std::time::Instant::now);
let mut workspace_guard = experts
.device_workspace
.lock()
.map_err(|_| "NVFP4 device routes workspace lock is poisoned")?;
if workspace_guard.is_none() {
drop(workspace_guard);
let zero = vec![0.0f32; experts.input_width];
let zero_sel = vec![0usize; n_sel];
let zero_w = vec![0.0f32; n_sel];
let _ = self.run_tensor_parallel_routes_nvfp4_device(
experts,
&zero,
&zero_sel,
&zero_w,
n_sel,
activation_limit,
)?;
workspace_guard = experts
.device_workspace
.lock()
.map_err(|_| "NVFP4 device routes workspace lock is poisoned")?;
}
let workspace = workspace_guard
.as_mut()
.expect("NVFP4 device routes workspace initialized above");
if workspace.n_sel != n_sel {
return Err(format!(
"NVFP4 device routes experts/token changed: workspace {} != call {n_sel}",
workspace.n_sel
)
.into());
}
for &expert in selected {
if expert >= experts.expert_count {
return Err(format!(
"NVFP4 device selected expert {expert} outside 0..{}",
experts.expert_count
)
.into());
}
}
let sel_i32 = selected
.iter()
.map(|&expert| expert as i32)
.collect::<Vec<_>>();
if let Some((_, device)) = workspace.ev_entry.as_ref() {
if *device != e.ctx().ordinal() {
return Err("NVFP4 device-io routes engine changed".into());
}
} else {
let _main = e.gpu.enter_main()?;
workspace.ev_entry = Some((e.ctx().new_event(None)?, e.ctx().ordinal()));
}
{
let _main = e.gpu.enter_main()?;
let (ev_entry, _) = workspace.ev_entry.as_ref().expect("entry event set above");
ev_entry.record(&e.stream())?;
}
for (rank_index, engine) in self.ranks.iter().enumerate() {
let _main = engine.gpu.enter_main()?;
let (ev_entry, _) = workspace.ev_entry.as_ref().expect("entry event set above");
engine.stream().wait(ev_entry)?;
{
let mut destination = workspace.input[rank_index].slice_mut(0..experts.input_width);
engine
.stream()
.memcpy_dtod(&input_dev.slice(0..experts.input_width), &mut destination)?;
}
{
let Nvfp4DeviceRoutesWorkspace {
input, in_q, in_d, ..
} = &mut *workspace;
engine.quantize_q8_1_into(
&input[rank_index],
1,
experts.input_width,
&mut in_q[rank_index],
&mut in_d[rank_index],
)?;
}
}
self.nvfp4_routes_batched_sweeps(
experts,
workspace,
selected,
route_weights,
&sel_i32,
local_out,
n_sel,
activation_limit,
false,
)?;
for (rank_index, engine) in self.ranks.iter().enumerate().skip(1) {
let _main = engine.gpu.enter_main()?;
workspace.ev_rank[rank_index].record(&engine.stream())?;
}
if moe_direct_on() && self.ranks.len() == 2 {
{
let root = &self.ranks[0];
let _main = root.gpu.enter_main()?;
workspace
.ev_done
.as_ref()
.expect("device routes done event")
.record(&root.stream())?;
}
let _main = e.gpu.enter_main()?;
e.stream().wait(
workspace
.ev_done
.as_ref()
.expect("device routes done event"),
)?;
for ev in workspace.ev_rank.iter().skip(1) {
e.stream().wait(ev)?;
}
let mut output = e.uninit(experts.input_width)?;
e.add(
&workspace.accumulator[0],
&workspace.accumulator[1],
&mut output,
experts.input_width,
)?;
let output = output;
if let Some(started) = started {
use std::sync::atomic::Ordering;
let ns = TIMING_NS
.fetch_add(started.elapsed().as_nanos() as u64, Ordering::Relaxed)
+ started.elapsed().as_nanos() as u64;
let calls = TIMING_CALLS.fetch_add(1, Ordering::Relaxed) + 1;
if calls % 430 == 0 {
eprintln!(
"[nvfp4-dev-routes-direct-timing] calls={calls} total_ms={:.1} avg_us={:.1}",
ns as f64 / 1.0e6,
ns as f64 / calls as f64 / 1.0e3,
);
}
}
return Ok(output);
}
{
let root = &self.ranks[0];
let _main = root.gpu.enter_main()?;
for ev in workspace.ev_rank.iter().skip(1) {
root.stream().wait(ev)?;
}
root.stream()
.memcpy_dtod(&workspace.accumulator[1], &mut workspace.remote)?;
{
let Nvfp4DeviceRoutesWorkspace {
accumulator,
remote,
combined,
..
} = &mut *workspace;
root.add(&accumulator[0], remote, combined, experts.input_width)?;
}
workspace
.ev_done
.as_ref()
.expect("device routes done event")
.record(&root.stream())?;
}
let output = {
let _main = e.gpu.enter_main()?;
e.stream().wait(
workspace
.ev_done
.as_ref()
.expect("device routes done event"),
)?;
let mut output = e.uninit(experts.input_width)?;
e.stream().memcpy_dtod(
&workspace.combined.slice(0..experts.input_width),
&mut output.slice_mut(0..experts.input_width),
)?;
output
};
if let Some(started) = started {
use std::sync::atomic::Ordering;
let ns = TIMING_NS.fetch_add(started.elapsed().as_nanos() as u64, Ordering::Relaxed)
+ started.elapsed().as_nanos() as u64;
let calls = TIMING_CALLS.fetch_add(1, Ordering::Relaxed) + 1;
if calls % 430 == 0 {
eprintln!(
"[nvfp4-dev-routes-io-timing] calls={calls} total_ms={:.1} avg_us={:.1}",
ns as f64 / 1.0e6,
ns as f64 / calls as f64 / 1.0e3,
);
}
}
Ok(output)
}
#[allow(clippy::too_many_arguments)]
pub fn nvfp4_routes_prestage(
&self,
experts: &ResidentNvfp4TensorParallel,
e: &Engine,
input_dev: &crate::CudaSlice<f32>,
) -> Result<bool, Box<dyn std::error::Error>> {
self.nvfp4_routes_prestage_with(experts, e, input_dev, |_, _, _, _| Ok(false))
}
pub fn nvfp4_routes_prestage_with(
&self,
experts: &ResidentNvfp4TensorParallel,
e: &Engine,
input_dev: &crate::CudaSlice<f32>,
rank1_router: impl FnOnce(
&Engine,
&crate::CudaSlice<f32>,
&mut crate::CudaSlice<i32>,
&mut crate::CudaSlice<f32>,
) -> Result<bool, Box<dyn std::error::Error>>,
) -> Result<bool, Box<dyn std::error::Error>> {
if !routes_prestage_on() || step_tp_graph_enabled()? {
return Ok(false);
}
if input_dev.len() != experts.input_width {
return Err("NVFP4 prestage input width mismatch".into());
}
let mut workspace_guard = experts
.device_workspace
.lock()
.map_err(|_| "NVFP4 device routes workspace lock is poisoned")?;
let Some(workspace) = workspace_guard.as_mut() else {
return Ok(false);
};
if workspace.ev_input.is_none() {
let _main = e.gpu.enter_main()?;
workspace.ev_input = Some((e.ctx().new_event(None)?, e.ctx().ordinal()));
} else if workspace.ev_input.as_ref().map(|(_, d)| *d) != Some(e.ctx().ordinal()) {
return Err("NVFP4 prestage engine changed".into());
}
{
let _main = e.gpu.enter_main()?;
let (ev, _) = workspace.ev_input.as_ref().expect("armed above");
ev.record(&e.stream())?;
}
for (rank_index, engine) in self.ranks.iter().enumerate() {
let _main = engine.gpu.enter_main()?;
let (ev, _) = workspace.ev_input.as_ref().expect("armed above");
engine.stream().wait(ev)?;
{
let mut destination = workspace.input[rank_index].slice_mut(0..experts.input_width);
engine
.stream()
.memcpy_dtod(&input_dev.slice(0..experts.input_width), &mut destination)?;
}
{
let Nvfp4DeviceRoutesWorkspace {
input, in_q, in_d, ..
} = &mut *workspace;
engine.quantize_q8_1_into(
&input[rank_index],
1,
experts.input_width,
&mut in_q[rank_index],
&mut in_d[rank_index],
)?;
}
}
if self.ranks.len() == 2 {
let rank1 = &self.ranks[1];
let _r1 = rank1.gpu.enter_main()?;
let Nvfp4DeviceRoutesWorkspace {
input,
sel,
route_w,
..
} = &mut *workspace;
let (in1, rest_sel) = (&input[1], &mut sel[1]);
if rank1_router(rank1, in1, rest_sel, &mut route_w[1])? {
workspace.rank1_routed = true;
}
}
workspace.prestaged = true;
Ok(true)
}
#[allow(clippy::too_many_arguments)]
pub fn run_tensor_parallel_routes_nvfp4_device_routed_t2(
&self,
experts: &ResidentNvfp4TensorParallel,
e: &Engine,
z2: &crate::CudaSlice<f32>,
sel_d: &crate::CudaSlice<i32>,
w_d: &crate::CudaSlice<f32>,
n_sel_col: usize,
activation_limit: Option<f32>,
) -> Result<crate::CudaSlice<f32>, Box<dyn std::error::Error>> {
let world = self.ranks.len();
if world != NVFP4_CANONICAL_ROW_SHARDS {
return Err("NVFP4 t2 routes require the canonical 2-shard grid".into());
}
let width = experts.input_width;
let n_sel = 2 * n_sel_col;
if z2.len() < 2 * width || sel_d.len() < n_sel || w_d.len() < n_sel {
return Err("NVFP4 t2 routes geometry".into());
}
if !nvfp4_bank_v2_on() {
return Err("NVFP4 t2 routes require the v2 banks (MEMRA_NVFP4_BANK_V2=1)".into());
}
let local_out = experts.expert_width / world;
let mut guard = experts
.t2_workspace
.lock()
.map_err(|_| "NVFP4 t2 workspace lock is poisoned")?;
if guard.as_ref().is_none_or(|ws| ws.n_sel != n_sel) {
let mut input2 = Vec::new();
let mut in_q2 = Vec::new();
let mut in_d2 = Vec::new();
let mut sel2 = Vec::new();
let mut route_w2 = Vec::new();
let mut gate_out2 = Vec::new();
let mut up_out2 = Vec::new();
let mut act_q2 = Vec::new();
let mut act_d2 = Vec::new();
let mut partial2 = Vec::new();
let mut acc_a = Vec::new();
let mut acc_b = Vec::new();
let mut ev_rank = Vec::new();
for engine in &self.ranks {
let _m = engine.gpu.enter_main()?;
input2.push(engine.uninit(2 * width)?);
in_q2.push(engine.alloc_i8_uninit(2 * width)?);
in_d2.push(engine.uninit(2 * (width / 32))?);
sel2.push(engine.htod_i32(&vec![0i32; n_sel])?);
route_w2.push(engine.uninit(n_sel)?);
gate_out2.push(engine.uninit(n_sel * local_out)?);
up_out2.push(engine.uninit(n_sel * local_out)?);
act_q2.push(engine.alloc_i8_uninit(n_sel * local_out)?);
act_d2.push(engine.uninit(n_sel * (local_out / 32))?);
partial2.push(engine.uninit(n_sel * width)?);
acc_a.push(engine.uninit(width)?);
acc_b.push(engine.uninit(width)?);
ev_rank.push(engine.ctx().new_event(None)?);
}
let root = &self.ranks[0];
let (peer_a, peer_b, omix_a, omix_b, ev_root) = {
let _m = root.gpu.enter_main()?;
(
root.uninit(width)?,
root.uninit(width)?,
root.uninit(width)?,
root.uninit(width)?,
root.ctx().new_event(None)?,
)
};
let ev_entry = {
let _m = e.gpu.enter_main()?;
e.ctx().new_event(None)?
};
*guard = Some(Nvfp4T2Workspace {
input2,
in_q2,
in_d2,
sel2,
route_w2,
gate_out2,
up_out2,
act_q2,
act_d2,
partial2,
acc_a,
acc_b,
peer_a,
peer_b,
omix_a,
omix_b,
ev_entry,
ev_rank,
ev_root,
n_sel,
e_device: e.ctx().ordinal(),
});
}
let ws = guard.as_mut().expect("armed above");
if ws.e_device != e.ctx().ordinal() {
return Err("NVFP4 t2 routes engine changed".into());
}
{
let _main = e.gpu.enter_main()?;
ws.ev_entry.record(&e.stream())?;
}
for rank in 0..world {
let engine = &self.ranks[rank];
let _main = engine.gpu.enter_main()?;
engine.stream().wait(&ws.ev_entry)?;
{
let mut dst = ws.input2[rank].slice_mut(0..2 * width);
engine
.stream()
.memcpy_dtod(&z2.slice(0..2 * width), &mut dst)?;
}
{
let mut dst = ws.sel2[rank].slice_mut(0..n_sel);
engine
.stream()
.memcpy_dtod(&sel_d.slice(0..n_sel), &mut dst)?;
}
{
let mut dst = ws.route_w2[rank].slice_mut(0..n_sel);
engine
.stream()
.memcpy_dtod(&w_d.slice(0..n_sel), &mut dst)?;
}
{
let Nvfp4T2Workspace {
input2,
in_q2,
in_d2,
..
} = &mut *ws;
engine.quantize_q8_1_into(
&input2[rank],
2,
width,
&mut in_q2[rank],
&mut in_d2[rank],
)?;
}
let gate_bank = &experts.gate[rank];
let up_bank = &experts.up[rank];
if gate_bank.in_features != up_bank.in_features
|| gate_bank.local_out != up_bank.local_out
|| gate_bank.row_bytes != up_bank.row_bytes
|| gate_bank.expert_bytes != up_bank.expert_bytes
{
return Err("NVFP4 t2 routes need matched gate/up bank geometry".into());
}
{
let Nvfp4T2Workspace {
sel2,
in_q2,
in_d2,
gate_out2,
up_out2,
..
} = &mut *ws;
engine.qmatvec_nvfp4_sel_gu_tcol_into(
&gate_bank.bank,
&up_bank.bank,
&sel2[rank],
&in_q2[rank],
&in_d2[rank],
&mut gate_out2[rank],
&mut up_out2[rank],
n_sel,
n_sel_col,
gate_bank.in_features,
gate_bank.local_out,
gate_bank.row_bytes,
gate_bank.expert_bytes,
width,
width / 32,
)?;
}
{
let Nvfp4T2Workspace {
gate_out2,
up_out2,
sel2,
act_q2,
act_d2,
..
} = &mut *ws;
engine.silu_mul_scaled_q8_1_sel_into(
&gate_out2[rank],
&up_out2[rank],
&experts.macros_gate_dev[rank],
&experts.macros_up_dev[rank],
&sel2[rank],
activation_limit,
&mut act_q2[rank],
&mut act_d2[rank],
local_out,
n_sel,
)?;
}
let shard = &experts.down[rank];
if shard.device_rank != rank || shard.local_in != local_out {
return Err("NVFP4 t2 routes: down shard placement drifted".into());
}
{
let Nvfp4T2Workspace {
sel2,
act_q2,
act_d2,
partial2,
..
} = &mut *ws;
engine.qmatvec_nvfp4_sel_into(
&shard.bank,
&sel2[rank],
&act_q2[rank],
&act_d2[rank],
&mut partial2[rank],
n_sel,
shard.local_in,
shard.out_features,
shard.row_bytes,
shard.expert_bytes,
local_out,
local_out / 32,
)?;
}
{
let Nvfp4T2Workspace {
partial2,
route_w2,
sel2,
acc_a,
acc_b,
..
} = &mut *ws;
engine.axpy_rows_seq_md_off_into(
&partial2[rank],
&route_w2[rank],
&experts.macros_down_dev[rank],
&sel2[rank],
&mut acc_a[rank],
width,
n_sel_col,
0,
)?;
engine.axpy_rows_seq_md_off_into(
&partial2[rank],
&route_w2[rank],
&experts.macros_down_dev[rank],
&sel2[rank],
&mut acc_b[rank],
width,
n_sel_col,
n_sel_col,
)?;
}
if rank != 0 {
ws.ev_rank[rank].record(&engine.stream())?;
}
}
let root = &self.ranks[0];
{
let _main = root.gpu.enter_main()?;
for ev in ws.ev_rank.iter().skip(1) {
root.stream().wait(ev)?;
}
{
let Nvfp4T2Workspace {
acc_a,
acc_b,
peer_a,
peer_b,
omix_a,
omix_b,
..
} = &mut *ws;
{
let mut dst = peer_a.slice_mut(0..width);
root.stream()
.memcpy_dtod(&acc_a[1].slice(0..width), &mut dst)?;
}
{
let mut dst = peer_b.slice_mut(0..width);
root.stream()
.memcpy_dtod(&acc_b[1].slice(0..width), &mut dst)?;
}
root.add(&acc_a[0], peer_a, omix_a, width)?;
root.add(&acc_b[0], peer_b, omix_b, width)?;
}
ws.ev_root.record(&root.stream())?;
}
let _main = e.gpu.enter_main()?;
e.stream().wait(&ws.ev_root)?;
let mut out = e.uninit(2 * width)?;
e.stream()
.memcpy_dtod(&ws.omix_a.slice(0..width), &mut out.slice_mut(0..width))?;
e.stream().memcpy_dtod(
&ws.omix_b.slice(0..width),
&mut out.slice_mut(width..2 * width),
)?;
Ok(out)
}
pub fn run_tensor_parallel_routes_nvfp4_device_routed(
&self,
experts: &ResidentNvfp4TensorParallel,
e: &Engine,
input_dev: &crate::CudaSlice<f32>,
sel_d: &crate::CudaSlice<i32>,
w_d: &crate::CudaSlice<f32>,
experts_per_token: usize,
activation_limit: Option<f32>,
) -> Result<crate::CudaSlice<f32>, Box<dyn std::error::Error>> {
self.run_tensor_parallel_routes_nvfp4_device_routed_prejoin(
experts,
e,
input_dev,
sel_d,
w_d,
experts_per_token,
activation_limit,
|| Ok(()),
)
}
#[allow(clippy::too_many_arguments)]
pub fn run_tensor_parallel_routes_nvfp4_device_routed_prejoin(
&self,
experts: &ResidentNvfp4TensorParallel,
e: &Engine,
input_dev: &crate::CudaSlice<f32>,
sel_d: &crate::CudaSlice<i32>,
w_d: &crate::CudaSlice<f32>,
experts_per_token: usize,
activation_limit: Option<f32>,
pre_join: impl FnOnce() -> Result<(), Box<dyn std::error::Error>>,
) -> Result<crate::CudaSlice<f32>, Box<dyn std::error::Error>> {
self.run_tensor_parallel_routes_nvfp4_device_routed_prejoin_add3(
experts,
e,
input_dev,
sel_d,
w_d,
experts_per_token,
activation_limit,
pre_join,
None,
)
}
#[allow(clippy::too_many_arguments)]
pub fn run_tensor_parallel_routes_nvfp4_device_routed_prejoin_add3(
&self,
experts: &ResidentNvfp4TensorParallel,
e: &Engine,
input_dev: &crate::CudaSlice<f32>,
sel_d: &crate::CudaSlice<i32>,
w_d: &crate::CudaSlice<f32>,
experts_per_token: usize,
activation_limit: Option<f32>,
pre_join: impl FnOnce() -> Result<(), Box<dyn std::error::Error>>,
post_add: Option<(u64, u64)>,
) -> Result<crate::CudaSlice<f32>, Box<dyn std::error::Error>> {
if input_dev.len() != experts.input_width {
return Err(format!(
"NVFP4 device-routed input {} != width {}",
input_dev.len(),
experts.input_width
)
.into());
}
let n_sel = experts_per_token;
if sel_d.len() < n_sel || w_d.len() < n_sel {
return Err(format!(
"NVFP4 device-routed routes sel={} w={} < experts/token {n_sel}",
sel_d.len(),
w_d.len()
)
.into());
}
let world = self.ranks.len();
if world != NVFP4_CANONICAL_ROW_SHARDS {
return Err(format!(
"NVFP4 device routes require world == canonical shard grid \
({NVFP4_CANONICAL_ROW_SHARDS}), got {world}"
)
.into());
}
let local_out = experts.expert_width / world;
static TIMING_NS: std::sync::atomic::AtomicU64 = std::sync::atomic::AtomicU64::new(0);
static TIMING_CALLS: std::sync::atomic::AtomicU64 = std::sync::atomic::AtomicU64::new(0);
let timing = std::env::var("MEMRA_STEP_TP_TIMING").as_deref() == Ok("1");
let started = timing.then(std::time::Instant::now);
let mut workspace_guard = experts
.device_workspace
.lock()
.map_err(|_| "NVFP4 device routes workspace lock is poisoned")?;
if workspace_guard.is_none() {
drop(workspace_guard);
let zero = vec![0.0f32; experts.input_width];
let zero_sel = vec![0usize; n_sel];
let zero_w = vec![0.0f32; n_sel];
let _ = self.run_tensor_parallel_routes_nvfp4_device(
experts,
&zero,
&zero_sel,
&zero_w,
n_sel,
activation_limit,
)?;
workspace_guard = experts
.device_workspace
.lock()
.map_err(|_| "NVFP4 device routes workspace lock is poisoned")?;
}
let workspace = workspace_guard
.as_mut()
.expect("NVFP4 device routes workspace initialized above");
if workspace.n_sel != n_sel {
return Err(format!(
"NVFP4 device routes experts/token changed: workspace {} != call {n_sel}",
workspace.n_sel
)
.into());
}
if step_tp_graph_enabled()? {
if workspace.dev_route_e.is_none() {
let _main = e.gpu.enter_main()?;
workspace.dev_route_e = Some((
e.htod_i32(&vec![0i32; n_sel])?,
e.htod(&vec![0.0f32; n_sel])?,
));
}
if workspace.in_stage_e.is_none() {
let _main = e.gpu.enter_main()?;
workspace.in_stage_e = Some(e.htod(&vec![0.0f32; experts.input_width])?);
workspace.out_stage_e = Some(e.htod(&vec![0.0f32; experts.input_width])?);
}
if workspace.routes_graph.is_none() {
let graph = self.nvfp4_routes_build_graph(
experts,
workspace,
local_out,
n_sel,
activation_limit,
)?;
workspace.routes_graph = Some(graph);
eprintln!(
"[step-tp-graph] routes segment captured: ranks={world} n_sel={n_sel} \
children=3 updates=none performance_claim=false"
);
}
let output = {
let _main = e.gpu.enter_main()?;
{
let (sel_e, w_e) = workspace
.dev_route_e
.as_mut()
.expect("device route staging set above");
{
let mut dst = sel_e.slice_mut(0..n_sel);
e.stream().memcpy_dtod(&sel_d.slice(0..n_sel), &mut dst)?;
}
{
let mut dst = w_e.slice_mut(0..n_sel);
e.stream().memcpy_dtod(&w_d.slice(0..n_sel), &mut dst)?;
}
}
{
let in_stage = workspace
.in_stage_e
.as_mut()
.expect("graph staging set above");
let mut dst = in_stage.slice_mut(0..experts.input_width);
e.stream()
.memcpy_dtod(&input_dev.slice(0..experts.input_width), &mut dst)?;
}
unsafe {
let r = cudarc::driver::sys::cuGraphLaunch(
workspace
.routes_graph
.as_ref()
.expect("routes graph built above")
.exec,
e.stream().cu_stream() as cudarc::driver::sys::CUstream,
);
if r != cudarc::driver::sys::CUresult::CUDA_SUCCESS {
return Err(format!("routes graph launch: {r:?}").into());
}
}
let mut output = e.uninit(experts.input_width)?;
{
let out_stage = workspace
.out_stage_e
.as_ref()
.expect("graph staging set above");
e.stream().memcpy_dtod(
&out_stage.slice(0..experts.input_width),
&mut output.slice_mut(0..experts.input_width),
)?;
}
output
};
if let Some(started) = started {
use std::sync::atomic::Ordering;
let ns = TIMING_NS
.fetch_add(started.elapsed().as_nanos() as u64, Ordering::Relaxed)
+ started.elapsed().as_nanos() as u64;
let calls = TIMING_CALLS.fetch_add(1, Ordering::Relaxed) + 1;
if calls % 430 == 0 {
eprintln!(
"[nvfp4-dev-routed-timing] calls={calls} total_ms={:.1} avg_us={:.1}",
ns as f64 / 1.0e6,
ns as f64 / calls as f64 / 1.0e3,
);
}
}
return Ok(output);
}
if let Some((_, device)) = workspace.ev_entry.as_ref() {
if *device != e.ctx().ordinal() {
return Err("NVFP4 device-routed routes engine changed".into());
}
} else {
let _main = e.gpu.enter_main()?;
workspace.ev_entry = Some((e.ctx().new_event(None)?, e.ctx().ordinal()));
}
if workspace.dev_route_e.is_none() {
let _main = e.gpu.enter_main()?;
workspace.dev_route_e = Some((
e.htod_i32(&vec![0i32; n_sel])?,
e.htod(&vec![0.0f32; n_sel])?,
));
}
{
let _main = e.gpu.enter_main()?;
let (sel_e, w_e) = workspace
.dev_route_e
.as_mut()
.expect("device route staging set above");
{
let mut dst = sel_e.slice_mut(0..n_sel);
e.stream().memcpy_dtod(&sel_d.slice(0..n_sel), &mut dst)?;
}
{
let mut dst = w_e.slice_mut(0..n_sel);
e.stream().memcpy_dtod(&w_d.slice(0..n_sel), &mut dst)?;
}
let (ev_entry, _) = workspace.ev_entry.as_ref().expect("entry event set above");
ev_entry.record(&e.stream())?;
}
let prestaged = std::mem::take(&mut workspace.prestaged);
let rank1_routed = std::mem::take(&mut workspace.rank1_routed);
for (rank_index, engine) in self.ranks.iter().enumerate() {
let _main = engine.gpu.enter_main()?;
let (ev_entry, _) = workspace.ev_entry.as_ref().expect("entry event set above");
engine.stream().wait(ev_entry)?;
if !prestaged {
let mut destination = workspace.input[rank_index].slice_mut(0..experts.input_width);
engine
.stream()
.memcpy_dtod(&input_dev.slice(0..experts.input_width), &mut destination)?;
}
if !(rank1_routed && rank_index == 1) {
let (sel_e, w_e) = workspace
.dev_route_e
.as_ref()
.expect("device route staging set above");
{
let mut dst = workspace.sel[rank_index].slice_mut(0..n_sel);
engine
.stream()
.memcpy_dtod(&sel_e.slice(0..n_sel), &mut dst)?;
}
{
let mut dst = workspace.route_w[rank_index].slice_mut(0..n_sel);
engine
.stream()
.memcpy_dtod(&w_e.slice(0..n_sel), &mut dst)?;
}
}
if !prestaged {
let Nvfp4DeviceRoutesWorkspace {
input, in_q, in_d, ..
} = &mut *workspace;
engine.quantize_q8_1_into(
&input[rank_index],
1,
experts.input_width,
&mut in_q[rank_index],
&mut in_d[rank_index],
)?;
}
}
self.nvfp4_routes_batched_sweeps(
experts,
workspace,
&[],
&[],
&[],
local_out,
n_sel,
activation_limit,
true,
)?;
for (rank_index, engine) in self.ranks.iter().enumerate().skip(1) {
let _main = engine.gpu.enter_main()?;
workspace.ev_rank[rank_index].record(&engine.stream())?;
}
let memops = fence_memops_on() && moe_direct_on() && self.ranks.len() == 2;
let mut ticket = 0u32;
if memops {
use cudarc::driver::sys;
if workspace.fence_flags_raw == 0 {
let root = &self.ranks[0];
let _main = root.gpu.enter_main()?;
let mut ptr: sys::CUdeviceptr = 0;
let r = unsafe { sys::cuMemAlloc_v2(&mut ptr, 8) };
if r != sys::CUresult::CUDA_SUCCESS {
return Err(format!("fence flag alloc: {r:?}").into());
}
let r = unsafe { sys::cuMemsetD8_v2(ptr, 0, 8) };
if r != sys::CUresult::CUDA_SUCCESS {
return Err(format!("fence flag memset: {r:?}").into());
}
workspace.fence_flags_raw = ptr as u64;
}
workspace.fence_ticket = workspace.fence_ticket.wrapping_add(1).max(1);
ticket = workspace.fence_ticket;
let base = workspace.fence_flags_raw;
{
let root = &self.ranks[0];
let _main = root.gpu.enter_main()?;
let r = unsafe {
sys::cuStreamWriteValue32_v2(
root.stream().cu_stream() as sys::CUstream,
(base + 4) as sys::CUdeviceptr,
ticket,
0,
)
};
if r != sys::CUresult::CUDA_SUCCESS {
return Err(format!("fence write root: {r:?}").into());
}
}
}
pre_join()?;
if moe_direct_on() && self.ranks.len() == 2 {
let _main = e.gpu.enter_main()?;
if memops {
use cudarc::driver::sys;
let base = workspace.fence_flags_raw;
let r = unsafe {
sys::cuStreamWaitValue32_v2(
e.stream().cu_stream() as sys::CUstream,
(base + 4) as sys::CUdeviceptr,
ticket,
sys::CUstreamWaitValue_flags::CU_STREAM_WAIT_VALUE_GEQ as u32,
)
};
if r != sys::CUresult::CUDA_SUCCESS {
return Err(format!("fence wait: {r:?}").into());
}
for ev in workspace.ev_rank.iter().skip(1) {
e.stream().wait(ev)?;
}
} else {
{
let root = &self.ranks[0];
let _rmain = root.gpu.enter_main()?;
workspace
.ev_done
.as_ref()
.expect("device routes done event")
.record(&root.stream())?;
}
e.stream().wait(
workspace
.ev_done
.as_ref()
.expect("device routes done event"),
)?;
for ev in workspace.ev_rank.iter().skip(1) {
e.stream().wait(ev)?;
}
}
let mut output = e.uninit(experts.input_width)?;
if let Some((sh_raw, scale_raw)) = post_add {
e.add3_raw(
&workspace.accumulator[0],
&workspace.accumulator[1],
sh_raw,
scale_raw,
&mut output,
experts.input_width,
)?;
} else {
e.add(
&workspace.accumulator[0],
&workspace.accumulator[1],
&mut output,
experts.input_width,
)?;
}
let output = output;
if let Some(started) = started {
use std::sync::atomic::Ordering;
let ns = TIMING_NS
.fetch_add(started.elapsed().as_nanos() as u64, Ordering::Relaxed)
+ started.elapsed().as_nanos() as u64;
let calls = TIMING_CALLS.fetch_add(1, Ordering::Relaxed) + 1;
if calls % 430 == 0 {
eprintln!(
"[nvfp4-dev-routes-direct-timing] calls={calls} total_ms={:.1} avg_us={:.1}",
ns as f64 / 1.0e6,
ns as f64 / calls as f64 / 1.0e3,
);
}
}
return Ok(output);
}
{
let root = &self.ranks[0];
let _main = root.gpu.enter_main()?;
for ev in workspace.ev_rank.iter().skip(1) {
root.stream().wait(ev)?;
}
root.stream()
.memcpy_dtod(&workspace.accumulator[1], &mut workspace.remote)?;
{
let Nvfp4DeviceRoutesWorkspace {
accumulator,
remote,
combined,
..
} = &mut *workspace;
root.add(&accumulator[0], remote, combined, experts.input_width)?;
}
workspace
.ev_done
.as_ref()
.expect("device routes done event")
.record(&root.stream())?;
}
let output = {
let _main = e.gpu.enter_main()?;
e.stream().wait(
workspace
.ev_done
.as_ref()
.expect("device routes done event"),
)?;
let mut output = e.uninit(experts.input_width)?;
e.stream().memcpy_dtod(
&workspace.combined.slice(0..experts.input_width),
&mut output.slice_mut(0..experts.input_width),
)?;
output
};
if let Some(started) = started {
use std::sync::atomic::Ordering;
let ns = TIMING_NS.fetch_add(started.elapsed().as_nanos() as u64, Ordering::Relaxed)
+ started.elapsed().as_nanos() as u64;
let calls = TIMING_CALLS.fetch_add(1, Ordering::Relaxed) + 1;
if calls % 430 == 0 {
eprintln!(
"[nvfp4-dev-routed-timing] calls={calls} total_ms={:.1} avg_us={:.1}",
ns as f64 / 1.0e6,
ns as f64 / calls as f64 / 1.0e3,
);
}
}
Ok(output)
}
pub(crate) fn decode_v2_finish_root_fused(
&self,
ws: &mut StepTpDecodeV2Ws,
) -> Result<(), Box<dyn std::error::Error>> {
let root = &self.ranks[0];
let _main = root.gpu.enter_main()?;
if ws.raw_peer_partial != 0 {
raw_copy_bytes(ws.raw_peer_partial, ws.raw_o_partial1, ws.o_out * 4, root)?;
} else {
root.stream()
.memcpy_dtod(&ws.o_partials[1][0], &mut ws.peer_partial)?;
}
{
let StepTpDecodeV2Ws {
o_partials,
peer_partial,
reduce_a,
o_out,
..
} = &mut *ws;
root.add(&o_partials[0][0], peer_partial, reduce_a, *o_out)?;
}
let shadows = !no_local_shadow_on() || ws.raw_mixed_stage_e != 0;
if shadows {
let mut k_dst = ws.k_shadow.slice_mut(0..ws.local_kv_dim);
root.stream().memcpy_dtod(&ws.k[0], &mut k_dst)?;
let mut v_dst = ws.v_shadow.slice_mut(0..ws.local_kv_dim);
root.stream().memcpy_dtod(&ws.v_raw[0], &mut v_dst)?;
}
if shadows && ws.raw_peer_partial != 0 {
raw_copy_bytes(
ws.raw_k_shadow + (ws.local_kv_dim * 4) as u64,
ws.raw_k1,
ws.local_kv_dim * 4,
root,
)?;
raw_copy_bytes(
ws.raw_v_shadow + (ws.local_kv_dim * 4) as u64,
ws.raw_v1,
ws.local_kv_dim * 4,
root,
)?;
} else if shadows {
let start = ws.local_kv_dim;
let mut k_dst = ws.k_shadow.slice_mut(start..start + ws.local_kv_dim);
root.stream().memcpy_dtod(&ws.k[1], &mut k_dst)?;
let mut v_dst = ws.v_shadow.slice_mut(start..start + ws.local_kv_dim);
root.stream().memcpy_dtod(&ws.v_raw[1], &mut v_dst)?;
}
if ws.raw_mixed_stage_e != 0 {
raw_copy_bytes(ws.raw_mixed_stage_e, ws.raw_reduce_a, ws.o_out * 4, root)?;
let (k_stage, v_stage) = ws.raw_shadow_stage_e;
raw_copy_bytes(k_stage, ws.raw_k_shadow, 2 * ws.local_kv_dim * 4, root)?;
raw_copy_bytes(v_stage, ws.raw_v_shadow, 2 * ws.local_kv_dim * 4, root)?;
}
Ok(())
}
pub(crate) fn decode_v2_arm_token_mirrors(
&self,
ws: &mut StepTpDecodeV2Ws,
mixed_stage_e: u64,
shadow_stage_e: (u64, u64),
) -> Result<(), Box<dyn std::error::Error>> {
use cudarc::driver::DevicePtr;
let root = &self.ranks[0];
let _main = root.gpu.enter_main()?;
let stream = root.stream();
let (a, _g) = ws.reduce_a.device_ptr(&stream);
ws.raw_reduce_a = a as u64;
ws.raw_mixed_stage_e = mixed_stage_e;
ws.raw_shadow_stage_e = shadow_stage_e;
Ok(())
}
fn nvfp4_routes_build_graph(
&self,
experts: &ResidentNvfp4TensorParallel,
workspace: &mut Nvfp4DeviceRoutesWorkspace,
local_out: usize,
n_sel: usize,
activation_limit: Option<f32>,
) -> Result<RoutesGraph, Box<dyn std::error::Error>> {
use cudarc::driver::DevicePtr;
use cudarc::driver::sys;
fn cu_try(r: sys::CUresult, what: &str) -> Result<(), Box<dyn std::error::Error>> {
if r == sys::CUresult::CUDA_SUCCESS {
Ok(())
} else {
Err(format!("{what}: {r:?}").into())
}
}
let world = self.ranks.len();
if world != 2 {
return Err("routes graph door is built for the TP2 pair".into());
}
let width = experts.input_width;
let ptr_f32 = |buf: &crate::CudaSlice<f32>, engine: &Engine| -> u64 {
let stream = engine.stream();
let (ptr, _g) = buf.device_ptr(&stream);
ptr as u64
};
let ptr_i32 = |buf: &crate::CudaSlice<i32>, engine: &Engine| -> u64 {
let stream = engine.stream();
let (ptr, _g) = buf.device_ptr(&stream);
ptr as u64
};
let (sel_e, w_e) = workspace
.dev_route_e
.as_ref()
.expect("device route staging set before graph build");
let root_engine = &self.ranks[0];
let p_in_stage = ptr_f32(
workspace.in_stage_e.as_ref().expect("graph staging"),
root_engine,
);
let p_out_stage = ptr_f32(
workspace.out_stage_e.as_ref().expect("graph staging"),
root_engine,
);
let p_sel_e = ptr_i32(sel_e, root_engine);
let p_w_e = ptr_f32(w_e, root_engine);
let p_input: Vec<u64> = (0..world)
.map(|r| ptr_f32(&workspace.input[r], &self.ranks[r]))
.collect();
let p_sel: Vec<u64> = (0..world)
.map(|r| ptr_i32(&workspace.sel[r], &self.ranks[r]))
.collect();
let p_route_w: Vec<u64> = (0..world)
.map(|r| ptr_f32(&workspace.route_w[r], &self.ranks[r]))
.collect();
let p_acc1 = ptr_f32(&workspace.accumulator[1], &self.ranks[1]);
let p_remote = ptr_f32(&workspace.remote, root_engine);
let p_combined = ptr_f32(&workspace.combined, root_engine);
let raw_copy = |dst: u64,
src: u64,
bytes: usize,
engine: &Engine|
-> Result<(), Box<dyn std::error::Error>> {
unsafe {
cu_try(
sys::cuMemcpyAsync(
dst as sys::CUdeviceptr,
src as sys::CUdeviceptr,
bytes,
engine.stream().cu_stream() as sys::CUstream,
),
"routes graph cuMemcpyAsync",
)
}
};
let mut children = Vec::with_capacity(3);
for rank in 0..world {
let engine = &self.ranks[rank];
let _main = engine.gpu.enter_main()?;
let (child, _retained) = engine.capture_graph_retained(|_| {
raw_copy(p_input[rank], p_in_stage, width * 4, engine)?;
raw_copy(p_sel[rank], p_sel_e, n_sel * 4, engine)?;
raw_copy(p_route_w[rank], p_w_e, n_sel * 4, engine)?;
{
let Nvfp4DeviceRoutesWorkspace {
input, in_q, in_d, ..
} = &mut *workspace;
engine.quantize_q8_1_into(
&input[rank],
1,
width,
&mut in_q[rank],
&mut in_d[rank],
)?;
}
self.nvfp4_routes_batched_sweeps_rank(
experts,
workspace,
&[],
&[],
&[],
local_out,
n_sel,
activation_limit,
true,
rank,
)?;
Ok(())
})?;
children.push(child);
}
{
let root = &self.ranks[0];
let _main = root.gpu.enter_main()?;
let (child, _retained) = root.capture_graph_retained(|_| {
raw_copy(p_remote, p_acc1, width * 4, root)?;
{
let Nvfp4DeviceRoutesWorkspace {
accumulator,
remote,
combined,
..
} = &mut *workspace;
root.add(&accumulator[0], remote, combined, width)?;
}
raw_copy(p_out_stage, p_combined, width * 4, root)?;
Ok(())
})?;
children.push(child);
}
let mut parent: sys::CUgraph = std::ptr::null_mut();
unsafe {
cu_try(sys::cuGraphCreate(&mut parent, 0), "routes cuGraphCreate")?;
}
let mut n0: sys::CUgraphNode = std::ptr::null_mut();
let mut n1: sys::CUgraphNode = std::ptr::null_mut();
let mut n2: sys::CUgraphNode = std::ptr::null_mut();
unsafe {
cu_try(
sys::cuGraphAddChildGraphNode(
&mut n0,
parent,
std::ptr::null(),
0,
children[0].cu_graph(),
),
"routes child r0",
)?;
cu_try(
sys::cuGraphAddChildGraphNode(
&mut n1,
parent,
std::ptr::null(),
0,
children[1].cu_graph(),
),
"routes child r1",
)?;
let deps = [n0, n1];
cu_try(
sys::cuGraphAddChildGraphNode(
&mut n2,
parent,
deps.as_ptr(),
2,
children[2].cu_graph(),
),
"routes child root",
)?;
}
let mut exec: sys::CUgraphExec = std::ptr::null_mut();
unsafe {
cu_try(
sys::cuGraphInstantiateWithFlags(&mut exec, parent, 0),
"routes instantiate",
)?;
}
Ok(RoutesGraph {
exec,
parent,
_children: children,
})
}
#[allow(clippy::too_many_arguments)]
pub(crate) fn routes_rank_section(
&self,
experts: &ResidentNvfp4TensorParallel,
workspace: &mut Nvfp4DeviceRoutesWorkspace,
raw_input_src: u64,
local_out: usize,
n_sel: usize,
activation_limit: Option<f32>,
rank_index: usize,
) -> Result<(), Box<dyn std::error::Error>> {
let engine = &self.ranks[rank_index];
{
let _main = engine.gpu.enter_main()?;
let (sel_e_ptr, w_e_ptr) = workspace
.raw_dev_route_e
.ok_or("routes rank section requires armed staging pointers")?;
raw_copy_bytes(
workspace.raw_input[rank_index],
raw_input_src,
experts.input_width * 4,
engine,
)?;
raw_copy_bytes(workspace.raw_sel[rank_index], sel_e_ptr, n_sel * 4, engine)?;
raw_copy_bytes(
workspace.raw_route_w[rank_index],
w_e_ptr,
n_sel * 4,
engine,
)?;
{
let Nvfp4DeviceRoutesWorkspace {
input, in_q, in_d, ..
} = &mut *workspace;
engine.quantize_q8_1_into(
&input[rank_index],
1,
experts.input_width,
&mut in_q[rank_index],
&mut in_d[rank_index],
)?;
}
}
self.nvfp4_routes_batched_sweeps_rank(
experts,
workspace,
&[],
&[],
&[],
local_out,
n_sel,
activation_limit,
true,
rank_index,
)
}
pub(crate) fn routes_root_section(
&self,
experts: &ResidentNvfp4TensorParallel,
workspace: &mut Nvfp4DeviceRoutesWorkspace,
) -> Result<(), Box<dyn std::error::Error>> {
let root = &self.ranks[0];
let _main = root.gpu.enter_main()?;
let (acc1_ptr, remote_ptr, combined_ptr, out_stage_ptr) = workspace
.raw_combine
.ok_or("routes root section requires armed combine pointers")?;
raw_copy_bytes(remote_ptr, acc1_ptr, experts.input_width * 4, root)?;
{
let Nvfp4DeviceRoutesWorkspace {
accumulator,
remote,
combined,
..
} = &mut *workspace;
root.add(&accumulator[0], remote, combined, experts.input_width)?;
}
raw_copy_bytes(out_stage_ptr, combined_ptr, experts.input_width * 4, root)?;
Ok(())
}
pub(crate) fn routes_arm_raw(
&self,
experts: &ResidentNvfp4TensorParallel,
workspace: &mut Nvfp4DeviceRoutesWorkspace,
) -> Result<(), Box<dyn std::error::Error>> {
use cudarc::driver::DevicePtr;
if workspace.raw_dev_route_e.is_some() {
return Ok(());
}
let _ = experts;
let (sel_e, w_e) = workspace
.dev_route_e
.as_ref()
.ok_or("routes staging not armed")?;
let root = &self.ranks[0];
{
let _main = root.gpu.enter_main()?;
let stream = root.stream();
let (a, _g) = sel_e.device_ptr(&stream);
let (b, _g) = w_e.device_ptr(&stream);
workspace.raw_dev_route_e = Some((a as u64, b as u64));
let (c, _g) = workspace.accumulator[1].device_ptr(&stream);
let (d, _g) = workspace.remote.device_ptr(&stream);
let (f, _g) = workspace.combined.device_ptr(&stream);
let out_stage = workspace
.out_stage_e
.as_ref()
.ok_or("routes out stage not armed")?;
let (g_, _g) = out_stage.device_ptr(&stream);
workspace.raw_combine = Some((c as u64, d as u64, f as u64, g_ as u64));
}
for rank in 0..self.ranks.len() {
let engine = &self.ranks[rank];
let _main = engine.gpu.enter_main()?;
let stream = engine.stream();
let (a, _g) = workspace.input[rank].device_ptr(&stream);
let (b, _g) = workspace.sel[rank].device_ptr(&stream);
let (c, _g) = workspace.route_w[rank].device_ptr(&stream);
workspace.raw_input.push(a as u64);
workspace.raw_sel.push(b as u64);
workspace.raw_route_w.push(c as u64);
}
Ok(())
}
pub fn run_tensor_parallel_routes_nvfp4(
&self,
experts: &ResidentNvfp4TensorParallel,
input: &[f32],
tokens: usize,
selected: &[usize],
route_weights: &[f32],
experts_per_token: usize,
activation_limit: Option<f32>,
) -> Result<Vec<f32>, Box<dyn std::error::Error>> {
validate_activations(input, tokens, experts.input_width)?;
let pairs = tokens
.checked_mul(experts_per_token)
.ok_or("NVFP4 TP route count overflow")?;
if selected.len() != pairs || route_weights.len() != pairs {
return Err(format!(
"NVFP4 TP routes selected={} weights={} != tokens {tokens} x experts/token \
{experts_per_token} ({pairs})",
selected.len(),
route_weights.len(),
)
.into());
}
if !route_weights.iter().all(|weight| weight.is_finite()) {
return Err("NVFP4 TP route weights contain a non-finite value".into());
}
let mut output = vec![0.0f32; tokens * experts.input_width];
for token in 0..tokens {
let input_row = &input[token * experts.input_width..(token + 1) * experts.input_width];
for slot in 0..experts_per_token {
let pair = token * experts_per_token + slot;
let expert = selected[pair];
if expert >= experts.expert_count {
return Err(format!(
"NVFP4 TP selected expert {expert} outside 0..{}",
experts.expert_count
)
.into());
}
let gate = self.run_column_bank_expert_nvfp4(
&experts.gate,
&experts.macros_gate,
expert,
input_row,
)?;
let up = self.run_column_bank_expert_nvfp4(
&experts.up,
&experts.macros_up,
expert,
input_row,
)?;
let activated: Vec<f32> = gate
.iter()
.zip(&up)
.map(|(&gate, &up)| step_expert_activation_host(gate, up, activation_limit))
.collect();
debug_assert_eq!(activated.len(), experts.expert_width);
let down = self.run_row_bank_expert_nvfp4(
&experts.down,
&experts.macros_down,
expert,
&activated,
)?;
let weight = route_weights[pair];
for (sum, value) in output
[token * experts.input_width..(token + 1) * experts.input_width]
.iter_mut()
.zip(down)
{
*sum += weight * value;
}
}
}
Ok(output)
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn step_expert_activation_clamps_each_arm_by_the_official_contract() {
let limit = Some(7.0);
assert_eq!(step_expert_activation_host(20.0, 9.0, limit), 49.0);
assert_eq!(step_expert_activation_host(20.0, -9.0, limit), -49.0);
assert!(
step_expert_activation_host(-20.0, 9.0, limit).abs()
< step_expert_activation_host(-20.0, 9.0, None).abs()
);
assert!(validate_step_expert_activation_limit(Some(f32::NAN)).is_err());
assert!(validate_step_expert_activation_limit(Some(0.0)).is_err());
assert!(validate_step_expert_activation_limit(limit).is_ok());
}
#[test]
fn moe_residual_host_preserves_official_add_order() {
let output = moe_residual_host(&[1.0e20], &[-1.0e20], &[1.0]).unwrap();
assert_eq!(output, [0.0]);
assert_eq!(
moe_residual_host(&[0.0], &[0.0, 1.0], &[0.0]).unwrap_err(),
"MoE residual lengths residual=1 routed=2 shared=1"
);
}
#[test]
fn expert_owner_routes_preserve_global_pair_order_with_local_expert_ids() {
let selected = [0, 36, 72, 108, 144, 180, 216, 252];
let owners = partition_expert_owner_routes(288, 4, 1, 8, &selected).unwrap();
assert_eq!(owners.len(), 4);
for (rank, owner) in owners.iter().enumerate() {
assert_eq!(owner.rank, rank);
assert_eq!(owner.selected, vec![0, 36]);
assert_eq!(owner.token_rows, vec![0, 0]);
assert_eq!(owner.global_pairs, vec![rank * 2, rank * 2 + 1]);
}
}
#[test]
fn expert_owner_routes_validate_geometry_and_selected_experts() {
assert!(partition_expert_owner_routes(288, 5, 1, 8, &[0; 8]).is_err());
assert!(partition_expert_owner_routes(288, 4, 2, 8, &[0; 8]).is_err());
let error = partition_expert_owner_routes(288, 4, 1, 8, &[288; 8]).unwrap_err();
assert!(error.contains("outside 0..288"));
}
#[test]
fn step_grouped_owner_routes_validate_dynamic_top8_shapes() {
let selected = [
1, 73, 80, 145, 152, 159, 217, 224, 12, 84, 91, 156, 163, 170, 228, 235,
];
assert_eq!(
validate_step_grouped_owner_routes(288, 2, &selected).unwrap(),
16
);
let owners = partition_expert_owner_routes(288, 4, 2, 8, &selected).unwrap();
assert_eq!(
owners
.iter()
.map(|owner| owner.selected.len())
.collect::<Vec<_>>(),
vec![2, 4, 6, 4]
);
assert!(validate_step_grouped_owner_routes(288, 2, &selected[..8]).is_err());
assert!(validate_step_grouped_owner_routes(288, 1, &[0; 8]).is_err());
assert!(validate_step_grouped_owner_routes(287, 2, &selected).is_err());
}
#[test]
fn weighted_route_combine_requires_a_canonical_pair_permutation() {
let owner0 = [0usize, 3];
let owner1 = [1usize, 2];
let owners = [owner0.as_slice(), owner1.as_slice()];
assert_eq!(
validate_weighted_route_combine(4096, 4, 3, 1, &owners, &[0.1, 0.2, 0.3, 0.4],)
.unwrap(),
WeightedRouteCombineShape {
pairs: 4,
max_pairs: 12,
}
);
let duplicate = [owner0.as_slice(), &[1usize, 1][..]];
assert!(
validate_weighted_route_combine(4096, 4, 3, 1, &duplicate, &[0.1, 0.2, 0.3, 0.4],)
.is_err()
);
assert!(
validate_weighted_route_combine(4096, 4, 3, 1, &owners, &[0.1, f32::NAN, 0.3, 0.4],)
.is_err()
);
assert!(
validate_weighted_route_combine(4096, 4, 1, 2, &owners, &[0.1, 0.2, 0.3, 0.4],)
.is_err()
);
}
#[test]
fn native_p2p_door_is_strict_and_default_off() {
assert!(!parse_step_tp_native_p2p(None).unwrap());
assert!(!parse_step_tp_native_p2p(Some("")).unwrap());
assert!(!parse_step_tp_native_p2p(Some("0")).unwrap());
assert!(parse_step_tp_native_p2p(Some("1")).unwrap());
assert!(parse_step_tp_native_p2p(Some("true")).is_err());
assert!(parse_step_tp_native_p2p(Some("2")).is_err());
}
#[test]
fn bulk_p2p_door_is_strict_and_default_off() {
assert!(!parse_step_tp_bulk_p2p(None).unwrap());
assert!(!parse_step_tp_bulk_p2p(Some("")).unwrap());
assert!(!parse_step_tp_bulk_p2p(Some("0")).unwrap());
assert!(parse_step_tp_bulk_p2p(Some("1")).unwrap());
assert!(parse_step_tp_bulk_p2p(Some("true")).is_err());
assert!(parse_step_tp_bulk_p2p(Some("2")).is_err());
}
#[test]
fn ep_device_arithmetic_door_is_strict_and_default_off() {
assert!(!parse_step_ep_device_arithmetic(None).unwrap());
assert!(!parse_step_ep_device_arithmetic(Some("")).unwrap());
assert!(!parse_step_ep_device_arithmetic(Some("0")).unwrap());
assert!(parse_step_ep_device_arithmetic(Some("1")).unwrap());
assert!(parse_step_ep_device_arithmetic(Some("true")).is_err());
assert!(parse_step_ep_device_arithmetic(Some("2")).is_err());
}
#[test]
fn f32_mirror_door_is_strict_and_default_off() {
assert!(!parse_step_tp_f32_mirror(None).unwrap());
assert!(!parse_step_tp_f32_mirror(Some("")).unwrap());
assert!(!parse_step_tp_f32_mirror(Some("0")).unwrap());
assert!(parse_step_tp_f32_mirror(Some("1")).unwrap());
assert!(parse_step_tp_f32_mirror(Some("true")).is_err());
assert!(parse_step_tp_f32_mirror(Some("2")).is_err());
}
fn matrix(out_features: usize, in_features: usize) -> (Vec<u8>, Vec<f32>) {
let codes = (0..out_features * in_features)
.map(|index| (index % 251) as u8)
.collect();
let scales = (0..out_features.div_ceil(FP8_BLOCK) * in_features.div_ceil(FP8_BLOCK))
.map(|index| index as f32 + 1.0)
.collect();
(codes, scales)
}
fn bf16_matrix_bytes(out_features: usize, in_features: usize) -> Vec<u8> {
(0..out_features * in_features)
.flat_map(|value| (value as u16).to_le_bytes())
.collect()
}
fn decode_u16(bytes: &[u8]) -> Vec<u16> {
bytes
.chunks_exact(2)
.map(|bytes| u16::from_le_bytes([bytes[0], bytes[1]]))
.collect()
}
#[test]
fn bf16_matrix_rejects_wrong_byte_count() {
let bytes = vec![0u8; 4 * 4 * 2 - 1];
let matrix = Bf16Matrix {
bytes: &bytes,
out_features: 4,
in_features: 4,
};
assert!(matrix.validate().unwrap_err().contains("4x4x2"));
}
#[test]
fn replicated_device_rows_require_exact_rank_local_shapes() {
assert_eq!(
replicated_device_row_values(3, 4096, 4, &[12_288; 4]).unwrap(),
12_288
);
assert!(replicated_device_row_values(0, 4096, 4, &[0; 4]).is_err());
assert!(replicated_device_row_values(3, 0, 4, &[0; 4]).is_err());
assert!(replicated_device_row_values(3, 4096, 4, &[12_288; 3]).is_err());
assert!(
replicated_device_row_values(3, 4096, 4, &[12_288, 12_288, 12_287, 12_288]).is_err()
);
assert!(replicated_device_row_values(usize::MAX, 2, 1, &[0]).is_err());
}
#[test]
fn replicated_device_row_refresh_requires_exact_root_source() {
assert_eq!(
replicated_device_row_source_values(1, 12_288, 12_288, 3, 3).unwrap(),
12_288
);
assert!(replicated_device_row_source_values(0, 12_288, 0, 3, 3).is_err());
assert!(replicated_device_row_source_values(1, 0, 0, 3, 3).is_err());
assert!(replicated_device_row_source_values(1, 12_288, 12_287, 3, 3).is_err());
assert!(replicated_device_row_source_values(1, 12_288, 12_288, 2, 3).is_err());
assert!(replicated_device_row_source_values(usize::MAX, 2, 0, 3, 3).is_err());
}
#[test]
fn step_bf16_canonical_rows_are_topology_invariant_through_tp8() {
for tp in [1, 2, 4, 8] {
assert_eq!(step_bf16_canonical_chunk_rows(8_192, tp).unwrap(), 1_024);
assert_eq!(step_bf16_canonical_chunk_rows(12_288, tp).unwrap(), 1_536);
assert_eq!(step_bf16_canonical_chunk_rows(1_024, tp).unwrap(), 128);
assert_eq!(step_bf16_canonical_chunk_cols(8_192, tp).unwrap(), 1_024);
assert_eq!(step_bf16_canonical_chunk_cols(12_288, tp).unwrap(), 1_536);
}
assert!(step_bf16_canonical_chunk_rows(12_288, 3).is_err());
assert!(step_bf16_canonical_chunk_rows(1_001, 2).is_err());
assert!(step_bf16_canonical_chunk_cols(12_288, 3).is_err());
assert!(step_bf16_canonical_chunk_cols(1_001, 2).is_err());
}
#[test]
fn cache_rows_split_by_token_then_rank() {
let rows = (0u8..24).collect::<Vec<_>>();
assert_eq!(
cache_rank_rows(&rows, 3, 4, 2, 0).unwrap(),
vec![0, 1, 2, 3, 8, 9, 10, 11, 16, 17, 18, 19]
);
assert_eq!(
cache_rank_rows(&rows, 3, 4, 2, 1).unwrap(),
vec![4, 5, 6, 7, 12, 13, 14, 15, 20, 21, 22, 23]
);
assert!(cache_rank_rows(&rows[..23], 3, 4, 2, 0).is_err());
assert!(cache_rank_rows(&rows, 3, 4, 2, 2).is_err());
}
#[test]
fn bf16_column_shard_preserves_contiguous_output_rows() {
let bytes = bf16_matrix_bytes(4, 4);
let matrix = Bf16Matrix {
bytes: &bytes,
out_features: 4,
in_features: 4,
};
let shard = bf16_column_shard(matrix, 2, 1).unwrap();
assert_eq!(shard.out_features, 2);
assert_eq!(shard.in_features, 4);
assert_eq!(decode_u16(shard.bytes), (8..16).collect::<Vec<_>>());
}
#[test]
fn bf16_row_shard_preserves_each_input_column_window() {
let bytes = bf16_matrix_bytes(3, 4);
let matrix = Bf16Matrix {
bytes: &bytes,
out_features: 3,
in_features: 4,
};
let shard = bf16_row_shard(matrix, 2, 1).unwrap();
assert_eq!(decode_u16(&shard), vec![2, 3, 6, 7, 10, 11]);
}
#[test]
fn bf16_row_block_preserves_global_column_order() {
let bytes = bf16_matrix_bytes(3, 8);
let matrix = Bf16Matrix {
bytes: &bytes,
out_features: 3,
in_features: 8,
};
let block = bf16_row_block(matrix, 2, 3).unwrap();
assert_eq!(decode_u16(&block), vec![2, 3, 4, 10, 11, 12, 18, 19, 20]);
}
#[test]
fn column_shard_preserves_contiguous_weight_and_scale_rows() {
let (codes, scales) = matrix(1280, 4096);
let matrix = E4m3BlockMatrix {
codes: &codes,
scales: &scales,
out_features: 1280,
in_features: 4096,
};
let shard = column_shard(matrix, 2, 1).unwrap();
assert_eq!(shard.out_features, 640);
assert_eq!(shard.codes, &codes[640 * 4096..]);
assert_eq!(shard.scales, &scales[5 * 32..]);
}
#[test]
fn row_shard_preserves_each_weight_and_scale_column_window() {
let (codes, scales) = matrix(4096, 1280);
let matrix = E4m3BlockMatrix {
codes: &codes,
scales: &scales,
out_features: 4096,
in_features: 1280,
};
let (shard_codes, shard_scales) = row_shard(matrix, 2, 1).unwrap();
assert_eq!(shard_codes.len(), 4096 * 640);
assert_eq!(&shard_codes[..640], &codes[640..1280]);
assert_eq!(&shard_codes[640..1280], &codes[1280 + 640..2560]);
assert_eq!(shard_scales.len(), 32 * 5);
assert_eq!(&shard_scales[..5], &scales[5..10]);
assert_eq!(&shard_scales[5..10], &scales[15..20]);
}
#[test]
fn activation_shards_keep_token_rows_separate() {
let activations: Vec<f32> = (0..2 * 8).map(|value| value as f32).collect();
assert_eq!(
activation_shard(&activations, 2, 8, 2, 1),
vec![4.0, 5.0, 6.0, 7.0, 12.0, 13.0, 14.0, 15.0],
);
}
#[test]
fn expert_bank_selects_expert_major_code_and_scale_planes() {
let expert_count = 2;
let out_features = 128;
let in_features = 128;
let code_stride = out_features * in_features;
let codes: Vec<u8> = (0..expert_count * code_stride)
.map(|index| (index % 251) as u8)
.collect();
let scales = vec![1.0f32, 2.0];
let bank = E4m3ExpertBank {
codes: &codes,
scales: &scales,
expert_count,
out_features,
in_features,
};
bank.validate().unwrap();
let expert = bank.expert(1).unwrap();
assert_eq!(expert.codes, &codes[code_stride..]);
assert_eq!(expert.scales, &[2.0]);
}
#[test]
fn expert_bank_rejects_non_positive_scale() {
let codes = vec![0u8; 128 * 128];
let scales = vec![0.0f32];
let bank = E4m3ExpertBank {
codes: &codes,
scales: &scales,
expert_count: 1,
out_features: 128,
in_features: 128,
};
assert!(bank.validate().unwrap_err().contains("non-positive"));
}
#[test]
fn tensor_parallel_column_bank_keeps_each_expert_scale_plane_separate() {
let expert_count = 2;
let out_features = 256;
let in_features = 128;
let code_stride = out_features * in_features;
let scale_stride = 2;
let codes = (0..expert_count * code_stride)
.map(|index| (index % 251) as u8)
.collect::<Vec<_>>();
let scales = vec![10.0f32, 11.0, 20.0, 21.0];
let bank = E4m3ExpertBank {
codes: &codes,
scales: &scales,
expert_count,
out_features,
in_features,
};
let rank = pack_column_bank_rank(bank, 2, 1).unwrap();
assert_eq!(rank.out_features, 128);
assert_eq!(rank.in_features, 128);
assert_eq!(rank.codes.len(), expert_count * 128 * 128);
assert_eq!(rank.scales, vec![11.0, 21.0]);
assert_eq!(&rank.codes[..128 * 128], &codes[128 * 128..256 * 128]);
assert_eq!(
&rank.codes[128 * 128..],
&codes[code_stride + 128 * 128..2 * code_stride]
);
assert_eq!(scale_stride, scales.len() / expert_count);
}
#[test]
fn tensor_parallel_row_bank_keeps_each_expert_scale_plane_separate() {
let expert_count = 2;
let out_features = 128;
let in_features = 256;
let code_stride = out_features * in_features;
let codes = (0..expert_count * code_stride)
.map(|index| (index % 251) as u8)
.collect::<Vec<_>>();
let scales = vec![10.0f32, 11.0, 20.0, 21.0];
let bank = E4m3ExpertBank {
codes: &codes,
scales: &scales,
expert_count,
out_features,
in_features,
};
let rank = pack_row_bank_rank(bank, 2, 1).unwrap();
assert_eq!(rank.out_features, 128);
assert_eq!(rank.in_features, 128);
assert_eq!(rank.k_blocks, Some(1));
assert_eq!(rank.codes.len(), expert_count * 128 * 128);
assert_eq!(rank.scales, vec![11.0, 21.0]);
assert_eq!(&rank.codes[..128], &codes[128..256]);
assert_eq!(
&rank.codes[128 * 128..128 * 128 + 128],
&codes[code_stride + 128..code_stride + 256]
);
}
#[test]
fn tensor_parallel_row_bank_preserves_global_k_block_order() {
let expert_count = 2;
let out_features = 256;
let in_features = 512;
let code_stride = out_features * in_features;
let mut codes = vec![0u8; expert_count * code_stride];
for expert in 0..expert_count {
for row in 0..out_features {
for block in 0..4 {
let value = (expert * 80 + block * 16 + row % 16) as u8;
let start = expert * code_stride + row * in_features + block * FP8_BLOCK;
codes[start..start + FP8_BLOCK].fill(value);
}
}
}
let scales = vec![
1.0f32, 2.0, 3.0, 4.0, 11.0, 12.0, 13.0, 14.0, 101.0, 102.0, 103.0, 104.0, 111.0,
112.0, 113.0, 114.0,
];
let bank = E4m3ExpertBank {
codes: &codes,
scales: &scales,
expert_count,
out_features,
in_features,
};
let rank = pack_row_bank_rank(bank, 2, 1).unwrap();
assert_eq!(rank.out_features, out_features);
assert_eq!(rank.in_features, 256);
assert_eq!(rank.k_blocks, Some(2));
assert_eq!(rank.code_stride, out_features * 256);
assert_eq!(rank.scale_stride, 4);
assert_eq!(&rank.scales[..4], &[3.0, 13.0, 4.0, 14.0]);
assert_eq!(&rank.scales[4..], &[103.0, 113.0, 104.0, 114.0]);
let block_stride = out_features * FP8_BLOCK;
assert!(rank.codes[..FP8_BLOCK].iter().all(|&code| code == 32));
assert!(
rank.codes[block_stride..block_stride + FP8_BLOCK]
.iter()
.all(|&code| code == 48)
);
assert!(
rank.codes[rank.code_stride..rank.code_stride + FP8_BLOCK]
.iter()
.all(|&code| code == 112)
);
assert!(
rank.codes
[rank.code_stride + block_stride..rank.code_stride + block_stride + FP8_BLOCK]
.iter()
.all(|&code| code == 128)
);
}
#[test]
fn step_ep_layer_specs_are_literal_and_fail_closed() {
assert!(parse_step_ep_layer_specs(None).unwrap().is_empty());
assert!(parse_step_ep_layer_specs(Some("0")).unwrap().is_empty());
assert_eq!(
parse_step_ep_layer_specs(Some("24@1,2")).unwrap(),
vec![StepEpLayerSpec {
layer: 24,
devices: vec![1, 2],
}]
);
assert_eq!(
parse_step_ep_layer_specs(Some("24-25@1,2;31@0,2")).unwrap(),
vec![
StepEpLayerSpec {
layer: 24,
devices: vec![1, 2],
},
StepEpLayerSpec {
layer: 25,
devices: vec![1, 2],
},
StepEpLayerSpec {
layer: 31,
devices: vec![0, 2],
},
]
);
assert!(parse_step_ep_layer_specs(Some("24@1")).is_err());
assert!(parse_step_ep_layer_specs(Some("24@1,1")).is_err());
assert!(parse_step_ep_layer_specs(Some("layer@1,2")).is_err());
assert!(parse_step_ep_layer_specs(Some("25-24@1,2")).is_err());
assert!(parse_step_ep_layer_specs(Some("0-128@1,2")).is_err());
assert!(parse_step_ep_layer_specs(Some("24-25@1,2;25@0,2")).is_err());
assert!(parse_step_ep_layer_specs(Some("all@0,1")).is_err());
}
#[test]
fn step_tp_layer_specs_share_the_fail_closed_layer_contract() {
assert!(parse_step_tp_layer_specs(None).unwrap().is_empty());
assert!(parse_step_tp_layer_specs(Some("0")).unwrap().is_empty());
assert_eq!(
parse_step_tp_layer_specs(Some("24-25@1,2")).unwrap(),
vec![
StepTpLayerSpec {
layer: 24,
devices: vec![1, 2],
},
StepTpLayerSpec {
layer: 25,
devices: vec![1, 2],
},
]
);
let error = parse_step_tp_layer_specs(Some("24@1")).unwrap_err();
assert!(error.contains("MEMRA_STEP_TP"));
assert!(parse_step_tp_layer_specs(Some("24@1,1")).is_err());
assert!(parse_step_tp_layer_specs(Some("24-25@1,2;25@0,2")).is_err());
let all = parse_step_tp_layer_specs(Some("all@0,1,2,3,4,5,6,7")).unwrap();
assert_eq!(all.len(), STEP37_TRUNK_LAYERS);
assert_eq!(all.first().unwrap().layer, 0);
assert_eq!(all.last().unwrap().layer, STEP37_TRUNK_LAYERS - 1);
let devices = (0..8).collect::<Vec<_>>();
assert!(all.iter().all(|spec| spec.devices == devices));
assert!(parse_step_tp_layer_specs(Some("all@0,1;44@0,1")).is_err());
}
}
struct TokenGraphChild {
graph: cudarc::driver::CudaGraph,
node: cudarc::driver::sys::CUgraphNode,
ctx: cudarc::driver::sys::CUcontext,
}
struct TokenGraphFaSite {
ctx: cudarc::driver::sys::CUcontext,
memset_o: cudarc::driver::sys::CUgraphNode,
memset_m: [cudarc::driver::sys::CUgraphNode; 2],
fa: cudarc::driver::sys::CUgraphNode,
combine: cudarc::driver::sys::CUgraphNode,
window: usize,
n_head: usize,
n_head_kv: usize,
head_dim: usize,
}
pub struct TokenGraphBuilder {
parent: cudarc::driver::sys::CUgraph,
children: Vec<TokenGraphChild>,
frontier: Vec<cudarc::driver::sys::CUgraphNode>,
pending_detached: Vec<cudarc::driver::sys::CUgraphNode>,
group: Option<(
u32,
Vec<cudarc::driver::sys::CUgraphNode>,
Vec<cudarc::driver::sys::CUgraphNode>,
)>,
}
unsafe impl Send for TokenGraphBuilder {}
impl TokenGraphBuilder {
pub fn new() -> Result<Self, Box<dyn std::error::Error>> {
use cudarc::driver::sys;
let mut parent: sys::CUgraph = std::ptr::null_mut();
let r = unsafe { sys::cuGraphCreate(&mut parent, 0) };
if r != sys::CUresult::CUDA_SUCCESS {
return Err(format!("token graph create: {r:?}").into());
}
Ok(Self {
parent,
children: Vec::new(),
frontier: Vec::new(),
pending_detached: Vec::new(),
group: None,
})
}
fn push_child(
&mut self,
graph: cudarc::driver::CudaGraph,
parallel_group: Option<u32>,
detached: bool,
absorb: bool,
ctx: cudarc::driver::sys::CUcontext,
) -> Result<(), Box<dyn std::error::Error>> {
use cudarc::driver::sys;
let deps: Vec<sys::CUgraphNode> = match (&mut self.group, parallel_group) {
(Some((open, base, _)), Some(group)) if *open == group => base.clone(),
(state, Some(group)) => {
if let Some((_, _, members)) = state.take() {
self.frontier = members;
}
let base = self.frontier.clone();
*state = Some((group, base.clone(), Vec::new()));
base
}
(state, None) if detached => match state.as_ref() {
Some((_, base, _)) => base.clone(),
None => self.frontier.clone(),
},
(state, None) => {
if let Some((_, _, members)) = state.take() {
self.frontier = members;
}
let mut deps = self.frontier.clone();
if absorb {
deps.append(&mut self.pending_detached);
}
deps
}
};
let mut node: sys::CUgraphNode = std::ptr::null_mut();
let r = unsafe {
sys::cuGraphAddChildGraphNode(
&mut node,
self.parent,
if deps.is_empty() {
std::ptr::null()
} else {
deps.as_ptr()
},
deps.len(),
graph.cu_graph(),
)
};
if r != sys::CUresult::CUDA_SUCCESS {
return Err(format!("token graph child: {r:?}").into());
}
match (&mut self.group, parallel_group, detached) {
(_, None, true) => self.pending_detached.push(node),
(Some((_, _, members)), Some(_), _) => members.push(node),
_ => self.frontier = vec![node],
}
self.children.push(TokenGraphChild { graph, node, ctx });
Ok(())
}
pub fn finish(mut self) -> Result<TokenGraph, Box<dyn std::error::Error>> {
use cudarc::driver::sys;
if let Some((_, _, members)) = self.group.take() {
self.frontier = members;
}
let mut fa_sites = Vec::new();
for child in &self.children {
if let Some(site) = discover_fa_site(child.node, child.ctx)? {
fa_sites.push(site);
}
}
let mut exec: sys::CUgraphExec = std::ptr::null_mut();
let r = unsafe { sys::cuGraphInstantiateWithFlags(&mut exec, self.parent, 0) };
if r != sys::CUresult::CUDA_SUCCESS {
return Err(format!("token graph instantiate: {r:?}").into());
}
Ok(TokenGraph {
exec,
parent: self.parent,
_children: self.children,
fa_sites,
})
}
}
fn discover_fa_site(
child_node: cudarc::driver::sys::CUgraphNode,
ctx: cudarc::driver::sys::CUcontext,
) -> Result<Option<TokenGraphFaSite>, Box<dyn std::error::Error>> {
use cudarc::driver::sys;
fn cu_try(r: sys::CUresult, what: &str) -> Result<(), Box<dyn std::error::Error>> {
if r == sys::CUresult::CUDA_SUCCESS {
Ok(())
} else {
Err(format!("{what}: {r:?}").into())
}
}
let mut graph: sys::CUgraph = std::ptr::null_mut();
unsafe {
cu_try(
sys::cuGraphChildGraphNodeGetGraph(child_node, &mut graph),
"fa-site child GetGraph",
)?;
}
let mut count: usize = 0;
unsafe {
cu_try(
sys::cuGraphGetNodes(graph, std::ptr::null_mut(), &mut count),
"fa-site GetNodes(count)",
)?;
}
let mut nodes: Vec<sys::CUgraphNode> = vec![std::ptr::null_mut(); count];
unsafe {
cu_try(
sys::cuGraphGetNodes(graph, nodes.as_mut_ptr(), &mut count),
"fa-site GetNodes",
)?;
}
nodes.truncate(count);
let node_type =
|node: sys::CUgraphNode| -> Result<sys::CUgraphNodeType, Box<dyn std::error::Error>> {
let mut ty = sys::CUgraphNodeType::CU_GRAPH_NODE_TYPE_EMPTY;
unsafe {
cu_try(
sys::cuGraphNodeGetType(node, &mut ty),
"fa-site NodeGetType",
)?;
}
Ok(ty)
};
let memsets: Vec<sys::CUgraphNode> = {
let mut v = Vec::new();
for &node in &nodes {
if node_type(node)? == sys::CUgraphNodeType::CU_GRAPH_NODE_TYPE_MEMSET {
v.push(node);
}
}
v
};
if memsets.len() != 3 {
return Ok(None);
}
let dependents =
|node: sys::CUgraphNode| -> Result<Vec<sys::CUgraphNode>, Box<dyn std::error::Error>> {
let mut n: usize = 0;
unsafe {
cu_try(
sys::cuGraphNodeGetDependentNodes_v2(
node,
std::ptr::null_mut(),
std::ptr::null_mut(),
&mut n,
),
"fa-site GetDependentNodes(count)",
)?;
}
let mut v: Vec<sys::CUgraphNode> = vec![std::ptr::null_mut(); n];
unsafe {
cu_try(
sys::cuGraphNodeGetDependentNodes_v2(
node,
v.as_mut_ptr(),
std::ptr::null_mut(),
&mut n,
),
"fa-site GetDependentNodes",
)?;
}
v.truncate(n);
Ok(v)
};
let mut fa: Option<sys::CUgraphNode> = None;
let mut last_memset: Option<sys::CUgraphNode> = None;
for &ms in &memsets {
for dep in dependents(ms)? {
if node_type(dep)? == sys::CUgraphNodeType::CU_GRAPH_NODE_TYPE_KERNEL {
fa = Some(dep);
last_memset = Some(ms);
}
}
}
let (Some(fa), Some(_last)) = (fa, last_memset) else {
return Ok(None);
};
let mut combine: Option<sys::CUgraphNode> = None;
for dep in dependents(fa)? {
if node_type(dep)? == sys::CUgraphNodeType::CU_GRAPH_NODE_TYPE_KERNEL {
combine = Some(dep);
}
}
let Some(combine) = combine else {
return Ok(None);
};
let mut params: sys::CUDA_KERNEL_NODE_PARAMS = unsafe { std::mem::zeroed() };
unsafe {
cu_try(
sys::cuGraphKernelNodeGetParams_v2(fa, &mut params),
"fa-site KernelNodeGetParams",
)?;
}
let arg_i32 =
|slot: usize| -> i32 { unsafe { *(*params.kernelParams.add(slot) as *const i32) } };
let (hd, nh, nhkv, win) = (arg_i32(6), arg_i32(7), arg_i32(8), arg_i32(11));
let width_of = |node: sys::CUgraphNode| -> Result<usize, Box<dyn std::error::Error>> {
let mut mp: sys::CUDA_MEMSET_NODE_PARAMS = unsafe { std::mem::zeroed() };
unsafe {
cu_try(
sys::cuGraphMemsetNodeGetParams(node, &mut mp),
"fa-site MemsetNodeGetParams",
)?;
}
Ok(mp.width)
};
let mut widest = memsets[0];
for &ms in &memsets[1..] {
if width_of(ms)? > width_of(widest)? {
widest = ms;
}
}
let memset_m: Vec<sys::CUgraphNode> =
memsets.iter().copied().filter(|&m| m != widest).collect();
Ok(Some(TokenGraphFaSite {
ctx,
memset_o: widest,
memset_m: [memset_m[0], memset_m[1]],
fa,
combine,
window: win as usize,
n_head: nh as usize,
n_head_kv: nhkv as usize,
head_dim: hd as usize,
}))
}
pub struct TokenGraph {
exec: cudarc::driver::sys::CUgraphExec,
parent: cudarc::driver::sys::CUgraph,
_children: Vec<TokenGraphChild>,
fa_sites: Vec<TokenGraphFaSite>,
}
unsafe impl Send for TokenGraph {}
impl TokenGraph {
pub fn retarget_bucket(&mut self, bucket: usize) -> Result<(), Box<dyn std::error::Error>> {
use cudarc::driver::sys;
fn cu_try(r: sys::CUresult, what: &str) -> Result<(), Box<dyn std::error::Error>> {
if r == sys::CUresult::CUDA_SUCCESS {
Ok(())
} else {
Err(format!("{what}: {r:?}").into())
}
}
for site in &self.fa_sites {
let layer_bucket = if site.window > 0 {
bucket.min(site.window)
} else {
bucket
};
let sp = crate::fa_split_keys(layer_bucket, site.n_head_kv);
let nsp = layer_bucket.div_ceil(sp).max(1);
let mut params: sys::CUDA_KERNEL_NODE_PARAMS = unsafe { std::mem::zeroed() };
unsafe {
cu_try(
sys::cuGraphKernelNodeGetParams_v2(site.fa, &mut params),
"retarget fa GetParams",
)?;
*(*params.kernelParams.add(13) as *mut i32) = nsp as i32;
*(*params.kernelParams.add(14) as *mut i32) = sp as i32;
params.gridDimY = nsp as u32;
cu_try(
sys::cuGraphExecKernelNodeSetParams_v2(self.exec, site.fa, ¶ms),
"retarget fa SetParams",
)?;
}
let mut cparams: sys::CUDA_KERNEL_NODE_PARAMS = unsafe { std::mem::zeroed() };
unsafe {
cu_try(
sys::cuGraphKernelNodeGetParams_v2(site.combine, &mut cparams),
"retarget combine GetParams",
)?;
*(*cparams.kernelParams.add(6) as *mut i32) = nsp as i32;
cu_try(
sys::cuGraphExecKernelNodeSetParams_v2(self.exec, site.combine, &cparams),
"retarget combine SetParams",
)?;
}
let set_width =
|node: sys::CUgraphNode, width: usize| -> Result<(), Box<dyn std::error::Error>> {
let mut mp: sys::CUDA_MEMSET_NODE_PARAMS = unsafe { std::mem::zeroed() };
unsafe {
cu_try(
sys::cuGraphMemsetNodeGetParams(node, &mut mp),
"retarget memset GetParams",
)?;
}
mp.width = width;
unsafe {
cu_try(
sys::cuGraphExecMemsetNodeSetParams(self.exec, node, &mp, site.ctx),
"retarget memset SetParams",
)?;
}
Ok(())
};
set_width(site.memset_o, site.n_head * nsp * site.head_dim)?;
set_width(site.memset_m[0], site.n_head * nsp)?;
set_width(site.memset_m[1], site.n_head * nsp)?;
}
Ok(())
}
pub fn launch(&self, e: &Engine) -> Result<(), Box<dyn std::error::Error>> {
use cudarc::driver::sys;
let _main = e.gpu.enter_main()?;
let r = unsafe { sys::cuGraphLaunch(self.exec, e.stream().cu_stream() as sys::CUstream) };
if r != sys::CUresult::CUDA_SUCCESS {
return Err(format!("token graph launch: {r:?}").into());
}
Ok(())
}
}
impl Drop for TokenGraph {
fn drop(&mut self) {
unsafe {
let _ = cudarc::driver::sys::cuGraphExecDestroy(self.exec);
let _ = cudarc::driver::sys::cuGraphDestroy(self.parent);
}
}
}
std::thread_local! {
static TOKEN_GRAPH_BUILDER: std::cell::RefCell<Option<TokenGraphBuilder>> =
const { std::cell::RefCell::new(None) };
}
pub fn token_graph_build_begin() -> Result<(), Box<dyn std::error::Error>> {
let builder = TokenGraphBuilder::new()?;
TOKEN_GRAPH_BUILDER.with(|cell| *cell.borrow_mut() = Some(builder));
Ok(())
}
pub fn token_graph_build_finish() -> Result<TokenGraph, Box<dyn std::error::Error>> {
let builder = TOKEN_GRAPH_BUILDER
.with(|cell| cell.borrow_mut().take())
.ok_or("token graph build was not begun")?;
builder.finish()
}
pub fn token_graph_building() -> bool {
TOKEN_GRAPH_BUILDER.with(|cell| cell.borrow().is_some())
}
pub fn graph_section<F>(
engine: &Engine,
parallel_group: Option<u32>,
f: F,
) -> Result<(), Box<dyn std::error::Error>>
where
F: FnMut() -> Result<(), Box<dyn std::error::Error>>,
{
graph_section_opts(engine, parallel_group, false, false, f)
}
pub fn graph_section_absorbing<F>(engine: &Engine, f: F) -> Result<(), Box<dyn std::error::Error>>
where
F: FnMut() -> Result<(), Box<dyn std::error::Error>>,
{
graph_section_opts(engine, None, false, true, f)
}
pub fn graph_section_detached<F>(engine: &Engine, f: F) -> Result<(), Box<dyn std::error::Error>>
where
F: FnMut() -> Result<(), Box<dyn std::error::Error>>,
{
graph_section_opts(engine, None, true, false, f)
}
pub fn graph_section_opts<F>(
engine: &Engine,
parallel_group: Option<u32>,
detached: bool,
absorb: bool,
f: F,
) -> Result<(), Box<dyn std::error::Error>>
where
F: FnMut() -> Result<(), Box<dyn std::error::Error>>,
{
let building = token_graph_building();
if !building {
let mut f = f;
return f();
}
let (child, ctx) = {
let _main = engine.gpu.enter_main()?;
let mut ctx: cudarc::driver::sys::CUcontext = std::ptr::null_mut();
let r = unsafe { cudarc::driver::sys::cuCtxGetCurrent(&mut ctx) };
if r != cudarc::driver::sys::CUresult::CUDA_SUCCESS {
return Err(format!("graph section ctx query: {r:?}").into());
}
let mut f = f;
let (child, _retained) = engine.capture_graph_retained_nowarm(|_| f())?;
(child, ctx)
};
TOKEN_GRAPH_BUILDER.with(|cell| {
cell.borrow_mut()
.as_mut()
.expect("builder checked above")
.push_child(child, parallel_group, detached, absorb, ctx)
})
}