use cudarc::driver::sys::CUdeviceptr;
use onnx_runtime_ep_api::{DeviceGraphResource, EpError, Result};
use std::sync::Arc;
use std::sync::atomic::{AtomicU32, Ordering};
use crate::runtime::{CudaRuntime, GraphDeviceAllocation};
pub const H_EPOCH: usize = 0;
pub const H_REQUEST: usize = 1;
pub const H_DEVICE: usize = 2;
pub const H_OVERFLOW: usize = 3;
pub const H_POISON: usize = 4;
pub const H_COUNT: usize = 5;
pub const HEADER_LEN: usize = 6;
pub const HEADER_BYTES: usize = HEADER_LEN * 4;
pub fn words_for(num_experts: usize) -> usize {
num_experts.div_ceil(32)
}
pub(crate) const MARK_DEVICE_SRC: &str = r#"
// ==== expert-route telemetry (issue #1810 Slice 7A, inert observability) ====
// Header word indices: 0 epoch, 1 request, 2 device, 3 overflow, 4 poison,
// 5 count. `route_telemetry_bitmap`/`route_telemetry_header` are null when the
// owning kernel is disarmed, in which case every helper is a no-op and the
// route kernel's outputs are byte-for-byte unchanged.
__device__ __forceinline__ unsigned int route_telemetry_mark_one(
unsigned int* route_telemetry_bitmap,
unsigned int* route_telemetry_header,
int expert,
int experts)
{
if (expert < 0 || expert >= experts) {
atomicOr(&route_telemetry_header[4], 1u); // poison: fail closed
return 0u;
}
atomicOr(&route_telemetry_bitmap[expert >> 5], 1u << (expert & 31));
return 1u;
}
// Fuse point: called once per routed row after `indices[0..top_k]` is final.
__device__ __forceinline__ void route_telemetry_mark_row(
unsigned int* route_telemetry_bitmap,
unsigned int* route_telemetry_header,
const int* indices,
int top_k,
int experts)
{
if (route_telemetry_bitmap == 0) { return; } // disarmed: inert
unsigned int valid = 0u;
for (int slot = 0; slot < top_k; ++slot) {
valid += route_telemetry_mark_one(
route_telemetry_bitmap, route_telemetry_header, indices[slot], experts);
}
if (valid != 0u) {
unsigned int* count = &route_telemetry_header[5];
unsigned int observed = atomicAdd(count, 0u);
for (;;) {
if (observed == 0xffffffffu) {
atomicOr(&route_telemetry_header[3], 1u);
break;
}
bool overflow = valid > 0xffffffffu - observed;
unsigned int desired = overflow ? 0xffffffffu : observed + valid;
unsigned int prior = atomicCAS(count, observed, desired);
if (prior == observed) {
if (overflow) {
atomicOr(&route_telemetry_header[3], 1u);
}
break;
}
observed = prior;
}
}
}
// ==== end expert-route telemetry helpers ====
"#;
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub struct RouteTelemetryConfig {
pub request_id: u32,
pub device_id: u32,
pub num_experts: usize,
pub routes_per_row: usize,
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub enum TelemetryUnsupported {
DeviceMismatch { config: u32, runtime: u32 },
ZeroExperts,
InvalidRoutesPerRow {
routes_per_row: usize,
num_experts: usize,
},
RouteWidthMismatch { config: usize, execution: usize },
Alloc(String),
GraphInstalled,
}
impl std::fmt::Display for TelemetryUnsupported {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
Self::DeviceMismatch { config, runtime } => write!(
f,
"route telemetry device mismatch: config device {config} != runtime device {runtime} (multi-device fails closed)"
),
Self::ZeroExperts => write!(f, "route telemetry requires num_experts > 0"),
Self::InvalidRoutesPerRow {
routes_per_row,
num_experts,
} => write!(
f,
"route telemetry requires 0 < routes_per_row <= num_experts, got \
routes_per_row={routes_per_row} and num_experts={num_experts}"
),
Self::RouteWidthMismatch { config, execution } => write!(
f,
"route telemetry routes_per_row={config} does not match the prepared execution \
contract {execution}; re-arm with the kernel's actual selected-expert width"
),
Self::Alloc(message) => write!(f, "route telemetry buffer alloc failed: {message}"),
Self::GraphInstalled => write!(
f,
"route telemetry must be configured before device-graph capture"
),
}
}
}
#[derive(Debug)]
pub(crate) struct ArmedTelemetry {
request_id: u32,
device_id: u32,
num_experts: usize,
routes_per_row: u32,
words: usize,
bitmap: Arc<GraphDeviceAllocation>,
header: Arc<GraphDeviceAllocation>,
epoch: AtomicU32,
bitmap_bytes: usize,
}
impl ArmedTelemetry {
pub(crate) fn arm(
runtime: &Arc<CudaRuntime>,
config: RouteTelemetryConfig,
) -> std::result::Result<Self, TelemetryUnsupported> {
let device = runtime.ordinal();
if config.device_id != device {
return Err(TelemetryUnsupported::DeviceMismatch {
config: config.device_id,
runtime: device,
});
}
if config.num_experts == 0 {
return Err(TelemetryUnsupported::ZeroExperts);
}
if config.routes_per_row == 0 || config.routes_per_row > config.num_experts {
return Err(TelemetryUnsupported::InvalidRoutesPerRow {
routes_per_row: config.routes_per_row,
num_experts: config.num_experts,
});
}
let routes_per_row = u32::try_from(config.routes_per_row).map_err(|_| {
TelemetryUnsupported::InvalidRoutesPerRow {
routes_per_row: config.routes_per_row,
num_experts: config.num_experts,
}
})?;
let words = words_for(config.num_experts);
let bitmap_bytes = words * 4;
let bitmap = GraphDeviceAllocation::allocate(runtime, bitmap_bytes.max(1))
.map_err(|error| TelemetryUnsupported::Alloc(error.to_string()))?;
let header = GraphDeviceAllocation::allocate(runtime, HEADER_BYTES)
.map_err(|error| TelemetryUnsupported::Alloc(error.to_string()))?;
let armed = Self {
request_id: config.request_id,
device_id: config.device_id,
num_experts: config.num_experts,
routes_per_row,
words,
bitmap,
header,
epoch: AtomicU32::new(1),
bitmap_bytes,
};
if let Err(error) = armed.open_window(runtime) {
return Err(TelemetryUnsupported::Alloc(error.to_string()));
}
Ok(armed)
}
fn open_window(&self, runtime: &CudaRuntime) -> Result<()> {
let header_words: [u32; HEADER_LEN] = [
self.epoch.load(Ordering::Relaxed),
self.request_id,
self.device_id,
0,
0,
0,
];
let mut header_bytes = [0u8; HEADER_BYTES];
for (word, chunk) in header_words.iter().zip(header_bytes.chunks_exact_mut(4)) {
chunk.copy_from_slice(&word.to_ne_bytes());
}
unsafe {
runtime.htod(&header_bytes, self.header.ptr())?;
if self.bitmap_bytes > 0 {
let zeros = vec![0u8; self.bitmap_bytes];
runtime.htod(&zeros, self.bitmap.ptr())?;
}
}
Ok(())
}
pub(crate) fn reset_boundary(&self, runtime: &CudaRuntime) -> Result<()> {
if runtime.is_capturing()? {
return Err(EpError::KernelFailed(
"cuda_ep: route telemetry boundary reset is illegal during graph capture/replay \
(the window advances only at a coarse safe boundary)"
.into(),
));
}
runtime.drain_for_unmap()?;
self.epoch
.fetch_update(Ordering::AcqRel, Ordering::Acquire, |epoch| {
epoch.checked_add(1)
})
.map_err(|_| {
EpError::KernelFailed(
"cuda_ep: route telemetry epoch overflow; re-arm the observer".into(),
)
})?;
self.open_window(runtime)
}
pub(crate) fn matches_experts(&self, experts: usize) -> bool {
self.num_experts == experts
}
pub(crate) fn bitmap_ptr(&self) -> CUdeviceptr {
self.bitmap.ptr()
}
pub(crate) fn header_ptr(&self) -> CUdeviceptr {
self.header.ptr()
}
pub(crate) fn footprint_bytes(&self) -> usize {
self.bitmap_bytes + HEADER_BYTES
}
pub(crate) fn bitmap_addr(&self) -> CUdeviceptr {
self.bitmap.ptr()
}
pub(crate) fn snapshot(&self, runtime: &CudaRuntime) -> Result<TelemetrySnapshot> {
let mut header = [0u32; HEADER_LEN];
let mut bitmap = vec![0u32; self.words];
unsafe {
let header_bytes =
std::slice::from_raw_parts_mut(header.as_mut_ptr() as *mut u8, HEADER_BYTES);
runtime.dtoh(header_bytes, self.header.ptr())?;
if self.words > 0 {
let bitmap_bytes = std::slice::from_raw_parts_mut(
bitmap.as_mut_ptr() as *mut u8,
self.bitmap_bytes,
);
runtime.dtoh(bitmap_bytes, self.bitmap.ptr())?;
}
}
Ok(TelemetrySnapshot {
header,
bitmap,
num_experts: self.num_experts,
routes_per_row: self.routes_per_row,
})
}
pub(crate) fn validated_snapshot(
&self,
runtime: &CudaRuntime,
) -> Result<ValidatedTelemetrySnapshot> {
let snapshot = self.snapshot(runtime)?;
match consume_and_validate(
&snapshot.header,
&snapshot.bitmap,
self.epoch.load(Ordering::Acquire),
self.request_id,
self.device_id,
self.num_experts,
usize::try_from(self.routes_per_row).expect("u32 routes-per-row fits usize"),
) {
RouteDecision::HotSet(_) => {
let unique_expert_count = snapshot
.bitmap
.iter()
.try_fold(0_u32, |total, word| total.checked_add(word.count_ones()))
.ok_or_else(|| {
EpError::KernelFailed(
"cuda_ep: validated route telemetry unique-expert count overflow"
.into(),
)
})?;
Ok(ValidatedTelemetrySnapshot {
selected_route_count: snapshot.count(),
unique_expert_count,
})
}
RouteDecision::WholeBank(reason) => Err(EpError::KernelFailed(format!(
"cuda_ep: invalid BlockQuantizedMoE traffic record: {reason}"
))),
}
}
#[cfg(feature = "gpu-tests")]
pub(crate) fn inject_header_word(
&self,
runtime: &CudaRuntime,
index: usize,
value: u32,
) -> Result<()> {
if index >= HEADER_LEN {
return Err(EpError::KernelFailed(format!(
"cuda_ep: telemetry test header index {index} is out of range"
)));
}
let byte_offset = index.checked_mul(4).ok_or_else(|| {
EpError::KernelFailed("cuda_ep: telemetry test header offset overflow".into())
})?;
let destination = self
.header
.ptr()
.checked_add(byte_offset as u64)
.ok_or_else(|| {
EpError::KernelFailed("cuda_ep: telemetry test header address overflow".into())
})?;
unsafe { runtime.htod(&value.to_ne_bytes(), destination) }
}
pub(crate) fn device_graph_resources(&self) -> [DeviceGraphResource; 2] {
[
GraphDeviceAllocation::device_graph_resource(&self.bitmap),
GraphDeviceAllocation::device_graph_resource(&self.header),
]
}
}
#[derive(Clone, Debug)]
pub struct TelemetrySnapshot {
pub header: [u32; HEADER_LEN],
pub bitmap: Vec<u32>,
pub num_experts: usize,
pub routes_per_row: u32,
}
impl TelemetrySnapshot {
pub fn routed_experts(&self) -> Vec<usize> {
(0..self.num_experts)
.filter(|&e| self.bitmap[e >> 5] & (1u32 << (e & 31)) != 0)
.collect()
}
pub fn epoch(&self) -> u32 {
self.header[H_EPOCH]
}
pub fn count(&self) -> u32 {
self.header[H_COUNT]
}
pub fn poison(&self) -> bool {
self.header[H_POISON] != 0
}
pub fn overflow(&self) -> bool {
self.header[H_OVERFLOW] != 0
}
}
pub(crate) struct ValidatedTelemetrySnapshot {
selected_route_count: u32,
unique_expert_count: u32,
}
impl ValidatedTelemetrySnapshot {
pub(crate) fn selected_route_count(&self) -> u32 {
self.selected_route_count
}
pub(crate) fn unique_expert_count(&self) -> u32 {
self.unique_expert_count
}
}
pub fn cpu_bitmap(routes: &[i32], num_experts: usize) -> (Vec<u32>, bool) {
let mut bits = vec![0u32; words_for(num_experts)];
let mut poison = false;
for &route in routes {
if route < 0 || route as usize >= num_experts {
poison = true;
continue;
}
let e = route as usize;
bits[e >> 5] |= 1u32 << (e & 31);
}
(bits, poison)
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub enum RouteDecision {
HotSet(Vec<u32>),
WholeBank(String),
}
pub fn consume_and_validate(
header: &[u32],
bitmap: &[u32],
expected_epoch: u32,
expected_request: u32,
expected_device: u32,
expected_num_experts: usize,
expected_routes_per_row: usize,
) -> RouteDecision {
if header.len() != HEADER_LEN {
return RouteDecision::WholeBank(format!(
"header length mismatch: record={} expected {HEADER_LEN}",
header.len()
));
}
let expected_words = words_for(expected_num_experts);
if bitmap.len() != expected_words {
return RouteDecision::WholeBank(format!(
"bitmap length mismatch: record={} expected {expected_words}",
bitmap.len()
));
}
if expected_routes_per_row == 0 || expected_routes_per_row > expected_num_experts {
return RouteDecision::WholeBank(format!(
"invalid route-width contract: routes_per_row={expected_routes_per_row}, \
num_experts={expected_num_experts}"
));
}
if header[H_POISON] != 0 {
return RouteDecision::WholeBank("poison: out-of-range expert id observed".into());
}
if header[H_OVERFLOW] != 0 {
return RouteDecision::WholeBank("overflow: bounded route counter saturated".into());
}
if header[H_DEVICE] != expected_device {
return RouteDecision::WholeBank(format!(
"device mismatch: record dev={} expected {expected_device}",
header[H_DEVICE]
));
}
if header[H_REQUEST] != expected_request {
return RouteDecision::WholeBank(format!(
"request mismatch: record req={} expected {expected_request}",
header[H_REQUEST]
));
}
if header[H_EPOCH] != expected_epoch {
return RouteDecision::WholeBank(format!(
"epoch mismatch: record epoch={} expected {expected_epoch}",
header[H_EPOCH]
));
}
if let Some(last) = bitmap.last() {
let valid_tail_bits = expected_num_experts % 32;
if valid_tail_bits != 0 && (*last >> valid_tail_bits) != 0 {
return RouteDecision::WholeBank(
"bitmap contains experts outside the armed capacity".into(),
);
}
}
let Some(unique) = bitmap
.iter()
.try_fold(0u32, |total, word| total.checked_add(word.count_ones()))
else {
return RouteDecision::WholeBank("unique expert count overflow".into());
};
let count = header[H_COUNT];
let Ok(routes_per_row) = u32::try_from(expected_routes_per_row) else {
return RouteDecision::WholeBank(format!(
"route-width contract {expected_routes_per_row} exceeds the telemetry counter domain"
));
};
if !count.is_multiple_of(routes_per_row) {
return RouteDecision::WholeBank(format!(
"route count {count} is impossible for routes_per_row={routes_per_row}; every clean \
routed row contributes exactly {routes_per_row} selections"
));
}
if count < unique || (count == 0) != (unique == 0) {
return RouteDecision::WholeBank(format!(
"route count {count} is inconsistent with {unique} unique selected experts"
));
}
RouteDecision::HotSet(bitmap.to_vec())
}
#[allow(dead_code)]
fn _assert_ep6_contract() {
const _: () = assert!(HEADER_LEN == 6);
}
impl From<TelemetryUnsupported> for EpError {
fn from(reason: TelemetryUnsupported) -> Self {
EpError::KernelFailed(format!("cuda_ep: {reason}"))
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn words_for_rounds_up() {
assert_eq!(words_for(0), 0);
assert_eq!(words_for(1), 1);
assert_eq!(words_for(32), 1);
assert_eq!(words_for(33), 2);
assert_eq!(words_for(64), 2);
assert_eq!(words_for(160), 5);
}
#[test]
fn cpu_bitmap_sets_routed_bits_and_no_poison() {
let (bits, poison) = cpu_bitmap(&[0, 5, 31, 32, 63], 64);
assert!(!poison);
assert_eq!(bits.len(), 2);
assert_eq!(bits[0], (1 << 0) | (1 << 5) | (1 << 31));
assert_eq!(bits[1], (1 << 0) | (1 << 31)); }
#[test]
fn cpu_bitmap_flags_out_of_range_as_poison() {
let (_bits, low) = cpu_bitmap(&[-1], 8);
assert!(low, "negative id must poison");
let (bits, high) = cpu_bitmap(&[8, 2], 8);
assert!(high, "id == num_experts must poison");
assert_eq!(bits[0], 1 << 2);
}
#[test]
fn validate_accepts_matching_clean_record() {
let mut header = [0u32; HEADER_LEN];
header[H_EPOCH] = 4;
header[H_REQUEST] = 9;
header[H_DEVICE] = 1;
header[H_COUNT] = 2;
let bitmap = vec![0b1010u32];
assert_eq!(
consume_and_validate(&header, &bitmap, 4, 9, 1, 32, 2),
RouteDecision::HotSet(vec![0b1010u32])
);
assert!(matches!(
consume_and_validate(&header, &bitmap, 3, 9, 1, 32, 2),
RouteDecision::WholeBank(_)
));
}
#[test]
fn validate_fails_closed_on_each_defect() {
let clean = |epoch: u32, request: u32, device: u32| {
let mut header = [0u32; HEADER_LEN];
header[H_EPOCH] = epoch;
header[H_REQUEST] = request;
header[H_DEVICE] = device;
header[H_COUNT] = 1;
header
};
let bitmap = vec![1u32];
let mut poisoned = clean(4, 9, 1);
poisoned[H_POISON] = 1;
assert!(matches!(
consume_and_validate(&poisoned, &bitmap, 4, 9, 1, 32, 1),
RouteDecision::WholeBank(_)
));
let mut overflowed = clean(4, 9, 1);
overflowed[H_OVERFLOW] = 1;
assert!(matches!(
consume_and_validate(&overflowed, &bitmap, 4, 9, 1, 32, 1),
RouteDecision::WholeBank(_)
));
assert!(matches!(
consume_and_validate(&clean(4, 9, 2), &bitmap, 4, 9, 1, 32, 1),
RouteDecision::WholeBank(_)
));
assert!(matches!(
consume_and_validate(&clean(4, 8, 1), &bitmap, 4, 9, 1, 32, 1),
RouteDecision::WholeBank(_)
));
assert!(matches!(
consume_and_validate(&clean(2, 9, 1), &bitmap, 3, 9, 1, 32, 1),
RouteDecision::WholeBank(_)
));
let mut inconsistent_count = clean(4, 9, 1);
inconsistent_count[H_COUNT] = 0;
assert!(matches!(
consume_and_validate(&inconsistent_count, &bitmap, 4, 9, 1, 32, 1),
RouteDecision::WholeBank(_)
));
assert!(matches!(
consume_and_validate(&clean(4, 9, 1), &[1, 0], 4, 9, 1, 32, 1),
RouteDecision::WholeBank(_)
));
assert!(matches!(
consume_and_validate(&clean(4, 9, 1), &[1 << 31], 4, 9, 1, 17, 1),
RouteDecision::WholeBank(_)
));
let mut non_multiple = clean(4, 9, 1);
non_multiple[H_COUNT] = 3;
assert!(matches!(
consume_and_validate(&non_multiple, &bitmap, 4, 9, 1, 32, 2),
RouteDecision::WholeBank(reason) if reason.contains("impossible")
));
}
#[test]
fn device_counter_uses_saturating_cas_not_wrapping_add() {
assert!(MARK_DEVICE_SRC.contains("atomicCAS(count, observed, desired)"));
assert!(MARK_DEVICE_SRC.contains("desired = overflow ? 0xffffffffu"));
assert!(
!MARK_DEVICE_SRC.contains("atomicAdd(&route_telemetry_header[5], valid)"),
"the production counter must never use a wrapping increment"
);
}
#[test]
fn snapshot_reports_header_fields_and_routed_experts() {
let mut header = [0u32; HEADER_LEN];
header[H_EPOCH] = 7;
header[H_COUNT] = 3;
header[H_POISON] = 1;
header[H_OVERFLOW] = 1;
let snapshot = TelemetrySnapshot {
header,
bitmap: vec![(1 << 1) | (1 << 4)],
num_experts: 8,
routes_per_row: 1,
};
assert_eq!(snapshot.epoch(), 7);
assert_eq!(snapshot.count(), 3);
assert!(snapshot.poison());
assert!(snapshot.overflow());
assert_eq!(snapshot.routed_experts(), vec![1, 4]);
}
#[test]
fn device_mismatch_display_is_descriptive() {
let reason = TelemetryUnsupported::DeviceMismatch {
config: 3,
runtime: 0,
};
let text = reason.to_string();
assert!(text.contains("device mismatch"));
assert!(text.contains("fails closed"));
}
}