use std::collections::HashMap;
use std::sync::atomic::{AtomicU64, Ordering};
use std::sync::{Arc, Mutex};
use onnx_runtime_ep_api::EpError;
use onnx_runtime_ir::Node;
use onnx_runtime_memory_governor::{
HolderId, LeaseLedger, LedgerGovernor, MemoryError, MemoryGovernor, MemoryLease, MemoryRole,
Tier,
};
fn next_holder_id() -> HolderId {
static NEXT_HOLDER: AtomicU64 = AtomicU64::new(1);
HolderId::new(NEXT_HOLDER.fetch_add(1, Ordering::Relaxed))
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub(crate) enum CsaCacheFormat {
F32,
Fp8E4m3Block64,
Fp4E2m1Block32,
}
impl CsaCacheFormat {
fn parse(raw: &str) -> Option<Self> {
match raw {
"f32" => Some(Self::F32),
"fp8_e4m3_block64" => Some(Self::Fp8E4m3Block64),
"fp4_e2m1_block32" => Some(Self::Fp4E2m1Block32),
_ => None,
}
}
fn as_str(self) -> &'static str {
match self {
Self::F32 => "f32",
Self::Fp8E4m3Block64 => "fp8_e4m3_block64",
Self::Fp4E2m1Block32 => "fp4_e2m1_block32",
}
}
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub(crate) enum CsaStateGroupRefusal {
UnsupportedRatio { ratio: usize },
UnknownCacheFormat { raw: String },
Ratio4RequiresFp8Cache { cache_format: &'static str },
Ratio4RequiresIndexHeadDim128 { index_head_dim: usize },
Ratio128RejectsFp4,
InvalidHeadGeometry { num_heads: usize, head_dim: usize },
RopeExceedsHeadDim {
qk_rope_head_dim: usize,
head_dim: usize,
},
MultiDeviceAmbiguity { device_count: u32 },
MissingStateEdge { which: &'static str },
UnsupportedC1Ratio { ratio: usize },
OutOfMemory {
request: u64,
device_ordinal: u32,
requested: u64,
resident: u64,
limit: u64,
},
}
impl std::fmt::Display for CsaStateGroupRefusal {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
Self::UnsupportedRatio { ratio } => {
write!(
f,
"compression_ratio {ratio} is neither CSA (4) nor HCA (128)"
)
}
Self::UnknownCacheFormat { raw } => write!(f, "unknown cache_format {raw:?}"),
Self::Ratio4RequiresFp8Cache { cache_format } => {
write!(f, "ratio-4 records are hybrid FP8/BF16, not {cache_format}")
}
Self::Ratio4RequiresIndexHeadDim128 { index_head_dim } => write!(
f,
"ratio-4 selection needs a 128-wide index key, not {index_head_dim}"
),
Self::Ratio128RejectsFp4 => {
write!(
f,
"ratio-128 attention-compressor records are f32 or FP8/BF16, not FP4"
)
}
Self::InvalidHeadGeometry {
num_heads,
head_dim,
} => write!(
f,
"invalid head geometry num_heads={num_heads} head_dim={head_dim}"
),
Self::RopeExceedsHeadDim {
qk_rope_head_dim,
head_dim,
} => write!(
f,
"qk_rope_head_dim {qk_rope_head_dim} exceeds head_dim {head_dim}"
),
Self::MultiDeviceAmbiguity { device_count } => write!(
f,
"a CSA state group spans {device_count} devices; v1 requires exactly one"
),
Self::MissingStateEdge { which } => {
write!(f, "required state edge {which} is missing from the node")
}
Self::UnsupportedC1Ratio { ratio } => write!(
f,
"the C1 runtime slice threads only ratio-128 (HCA); ratio-{ratio} is a follow-up slice"
),
Self::OutOfMemory {
request,
device_ordinal,
requested,
resident,
limit,
} => write!(
f,
"CSA state group (request {request}, device {device_ordinal}) needs {requested} B \
but {resident} B are resident against a {limit} B limit"
),
}
}
}
impl From<CsaStateGroupRefusal> for EpError {
fn from(refusal: CsaStateGroupRefusal) -> Self {
match refusal {
CsaStateGroupRefusal::OutOfMemory {
requested,
resident,
limit,
..
} => EpError::OutOfMemory {
requested: usize::try_from(requested).unwrap_or(usize::MAX),
available: usize::try_from(limit.saturating_sub(resident)).unwrap_or(usize::MAX),
},
other => EpError::KernelFailed(format!(
"CompressedSparseAttention state group refused: {other}"
)),
}
}
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub(crate) struct CsaStateGroupDescriptor {
pub ratio: usize,
pub cache_format: CsaCacheFormat,
pub num_heads: usize,
pub head_dim: usize,
pub qk_rope_head_dim: usize,
pub index_head_dim: usize,
pub device_ordinal: u32,
pub device_count: u32,
pub request_id: u64,
pub has_compressed_state: bool,
pub has_carry_state: bool,
}
impl CsaStateGroupDescriptor {
pub(crate) fn from_node(
node: &Node,
device_ordinal: u32,
device_count: u32,
request_id: u64,
) -> Result<Self, CsaStateGroupRefusal> {
let ratio = attr_usize(node, "compression_ratio").unwrap_or(0);
let raw_format = node
.attr("cache_format")
.and_then(|attribute| attribute.as_str())
.unwrap_or("f32");
let cache_format = CsaCacheFormat::parse(raw_format).ok_or_else(|| {
CsaStateGroupRefusal::UnknownCacheFormat {
raw: raw_format.to_string(),
}
})?;
Ok(Self {
ratio,
cache_format,
num_heads: attr_usize(node, "num_heads").unwrap_or(0),
head_dim: attr_usize(node, "head_dim").unwrap_or(0),
qk_rope_head_dim: attr_usize(node, "qk_rope_head_dim").unwrap_or(0),
index_head_dim: attr_usize(node, "index_head_dim").unwrap_or(0),
device_ordinal,
device_count,
request_id,
has_compressed_state: node.inputs.get(6).map(Option::is_some).unwrap_or(false),
has_carry_state: node.inputs.get(7).map(Option::is_some).unwrap_or(false),
})
}
pub(crate) fn validate(&self) -> Result<(), CsaStateGroupRefusal> {
if self.ratio != 4 && self.ratio != 128 {
return Err(CsaStateGroupRefusal::UnsupportedRatio { ratio: self.ratio });
}
if self.num_heads == 0 || self.head_dim == 0 {
return Err(CsaStateGroupRefusal::InvalidHeadGeometry {
num_heads: self.num_heads,
head_dim: self.head_dim,
});
}
if self.qk_rope_head_dim > self.head_dim {
return Err(CsaStateGroupRefusal::RopeExceedsHeadDim {
qk_rope_head_dim: self.qk_rope_head_dim,
head_dim: self.head_dim,
});
}
if self.device_count != 1 {
return Err(CsaStateGroupRefusal::MultiDeviceAmbiguity {
device_count: self.device_count,
});
}
if !self.has_compressed_state {
return Err(CsaStateGroupRefusal::MissingStateEdge {
which: "past_compressed_kv",
});
}
if !self.has_carry_state {
return Err(CsaStateGroupRefusal::MissingStateEdge {
which: "past_compression_carry",
});
}
match self.ratio {
4 => {
if self.cache_format != CsaCacheFormat::Fp8E4m3Block64 {
return Err(CsaStateGroupRefusal::Ratio4RequiresFp8Cache {
cache_format: self.cache_format.as_str(),
});
}
if self.index_head_dim != 128 {
return Err(CsaStateGroupRefusal::Ratio4RequiresIndexHeadDim128 {
index_head_dim: self.index_head_dim,
});
}
}
128 => {
if self.cache_format == CsaCacheFormat::Fp4E2m1Block32 {
return Err(CsaStateGroupRefusal::Ratio128RejectsFp4);
}
}
_ => unreachable!("ratio guarded above"),
}
Ok(())
}
pub(crate) fn charge_key(&self) -> (u64, u32) {
(self.request_id, self.device_ordinal)
}
#[cfg_attr(not(test), allow(dead_code))]
pub(crate) fn validate_c1_runtime(&self) -> Result<(), CsaStateGroupRefusal> {
self.validate()?;
if self.ratio != 128 {
return Err(CsaStateGroupRefusal::UnsupportedC1Ratio { ratio: self.ratio });
}
Ok(())
}
}
#[derive(Clone, Copy, Debug, Default, PartialEq, Eq)]
pub(crate) struct CsaStateGroupBytes {
pub compressed: u64,
pub carry: u64,
pub dense_ring: u64,
pub index: u64,
pub index_carry: u64,
pub scratch: u64,
}
impl CsaStateGroupBytes {
pub(crate) fn workspace_only(workspace_bytes: &[usize]) -> Self {
Self {
scratch: workspace_bytes.iter().map(|&size| size.max(1) as u64).sum(),
..Self::default()
}
}
pub(crate) fn total(&self) -> u64 {
self.compressed
.saturating_add(self.carry)
.saturating_add(self.dense_ring)
.saturating_add(self.index)
.saturating_add(self.index_carry)
.saturating_add(self.scratch)
}
}
pub(crate) struct CsaStateGroupLedger {
governor: Arc<dyn MemoryGovernor + Send + Sync>,
holder: HolderId,
resident: AtomicU64,
peak: AtomicU64,
compressed: AtomicU64,
carry: AtomicU64,
dense_ring: AtomicU64,
index: AtomicU64,
index_carry: AtomicU64,
scratch: AtomicU64,
charge_failures: AtomicU64,
request_seq: AtomicU64,
charges: Mutex<HashMap<(u64, u32), u64>>,
}
impl std::fmt::Debug for CsaStateGroupLedger {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("CsaStateGroupLedger")
.field("holder", &self.holder)
.field("resident", &self.resident.load(Ordering::Relaxed))
.field("peak", &self.peak.load(Ordering::Relaxed))
.field(
"charge_failures",
&self.charge_failures.load(Ordering::Relaxed),
)
.finish_non_exhaustive()
}
}
impl Default for CsaStateGroupLedger {
fn default() -> Self {
Self::new(Arc::new(LedgerGovernor::new(LeaseLedger::new(
u64::MAX,
0,
0,
))))
}
}
impl CsaStateGroupLedger {
pub(crate) fn new(governor: Arc<dyn MemoryGovernor + Send + Sync>) -> Self {
Self {
governor,
holder: next_holder_id(),
resident: AtomicU64::new(0),
peak: AtomicU64::new(0),
compressed: AtomicU64::new(0),
carry: AtomicU64::new(0),
dense_ring: AtomicU64::new(0),
index: AtomicU64::new(0),
index_carry: AtomicU64::new(0),
scratch: AtomicU64::new(0),
charge_failures: AtomicU64::new(0),
request_seq: AtomicU64::new(0),
charges: Mutex::new(HashMap::new()),
}
}
#[cfg(test)]
pub(crate) fn with_device_limit(device_bytes: u64) -> Self {
Self::new(Arc::new(LedgerGovernor::new(LeaseLedger::new(
device_bytes,
0,
0,
))))
}
pub(crate) fn device_available_bytes(&self) -> u64 {
self.governor.available(Tier::Device)
}
pub(crate) fn governor_device_used(&self) -> u64 {
self.governor.used(Tier::Device)
}
pub(crate) fn next_request_id(&self) -> u64 {
self.request_seq.fetch_add(1, Ordering::Relaxed)
}
pub(crate) fn try_charge(
self: &Arc<Self>,
key: (u64, u32),
bytes: CsaStateGroupBytes,
) -> Result<CsaStateGroupCharge, CsaStateGroupRefusal> {
let total = bytes.total();
let lease = self
.governor
.reserve(
Tier::Device,
total,
MemoryRole::Workspace { step_scoped: false },
self.holder,
)
.map_err(|error| {
self.charge_failures.fetch_add(1, Ordering::Relaxed);
self.refusal_from(key, total, &error)
})?;
let mut charges = self
.charges
.lock()
.expect("CSA state-group ledger poisoned");
let next = self.resident.load(Ordering::Relaxed).saturating_add(total);
self.resident.store(next, Ordering::Relaxed);
bump_peak(&self.peak, next);
self.compressed
.fetch_add(bytes.compressed, Ordering::Relaxed);
self.carry.fetch_add(bytes.carry, Ordering::Relaxed);
self.dense_ring
.fetch_add(bytes.dense_ring, Ordering::Relaxed);
self.index.fetch_add(bytes.index, Ordering::Relaxed);
self.index_carry
.fetch_add(bytes.index_carry, Ordering::Relaxed);
self.scratch.fetch_add(bytes.scratch, Ordering::Relaxed);
*charges.entry(key).or_insert(0) += total;
drop(charges);
Ok(CsaStateGroupCharge {
ledger: Arc::clone(self),
key,
bytes,
lease,
})
}
fn refusal_from(
&self,
key: (u64, u32),
requested: u64,
error: &MemoryError,
) -> CsaStateGroupRefusal {
let (resident, limit) = match error {
MemoryError::TierExhausted { used, limit, .. } => (*used, *limit),
_ => {
let used = self.governor.used(Tier::Device);
(
used,
used.saturating_add(self.governor.available(Tier::Device)),
)
}
};
CsaStateGroupRefusal::OutOfMemory {
request: key.0,
device_ordinal: key.1,
requested,
resident,
limit,
}
}
fn release(&self, key: (u64, u32), bytes: CsaStateGroupBytes) {
let total = bytes.total();
let mut charges = self
.charges
.lock()
.expect("CSA state-group ledger poisoned");
self.resident.fetch_sub(
total.min(self.resident.load(Ordering::Relaxed)),
Ordering::Relaxed,
);
sub_saturating(&self.compressed, bytes.compressed);
sub_saturating(&self.carry, bytes.carry);
sub_saturating(&self.dense_ring, bytes.dense_ring);
sub_saturating(&self.index, bytes.index);
sub_saturating(&self.index_carry, bytes.index_carry);
sub_saturating(&self.scratch, bytes.scratch);
if let Some(entry) = charges.get_mut(&key) {
*entry = entry.saturating_sub(total);
if *entry == 0 {
charges.remove(&key);
}
}
}
pub(crate) fn resident_bytes(&self) -> u64 {
self.resident.load(Ordering::Relaxed)
}
pub(crate) fn peak_bytes(&self) -> u64 {
self.peak.load(Ordering::Relaxed)
}
pub(crate) fn compressed_bytes(&self) -> u64 {
self.compressed.load(Ordering::Relaxed)
}
pub(crate) fn dense_ring_bytes(&self) -> u64 {
self.dense_ring.load(Ordering::Relaxed)
}
pub(crate) fn charge_failures(&self) -> u64 {
self.charge_failures.load(Ordering::Relaxed)
}
pub(crate) fn resident_for(&self, request: u64, device_ordinal: u32) -> u64 {
self.charges
.lock()
.expect("CSA state-group ledger poisoned")
.get(&(request, device_ordinal))
.copied()
.unwrap_or(0)
}
pub(crate) fn active_group_count(&self) -> usize {
self.charges
.lock()
.expect("CSA state-group ledger poisoned")
.len()
}
}
pub(crate) struct CsaStateGroupCharge {
ledger: Arc<CsaStateGroupLedger>,
key: (u64, u32),
bytes: CsaStateGroupBytes,
lease: MemoryLease,
}
impl CsaStateGroupCharge {
#[cfg(test)]
pub(crate) fn total(&self) -> u64 {
self.bytes.total()
}
#[cfg(test)]
pub(crate) fn lease_bytes(&self) -> u64 {
self.lease.bytes()
}
}
impl std::fmt::Debug for CsaStateGroupCharge {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("CsaStateGroupCharge")
.field("key", &self.key)
.field("total", &self.bytes.total())
.field("lease_bytes", &self.lease.bytes())
.finish()
}
}
impl Drop for CsaStateGroupCharge {
fn drop(&mut self) {
self.ledger.release(self.key, self.bytes);
}
}
fn attr_usize(node: &Node, name: &str) -> Option<usize> {
node.attr(name)
.and_then(|attribute| attribute.as_int())
.and_then(|value| usize::try_from(value).ok())
}
fn bump_peak(peak: &AtomicU64, candidate: u64) {
let mut current = peak.load(Ordering::Relaxed);
while candidate > current {
match peak.compare_exchange_weak(current, candidate, Ordering::Relaxed, Ordering::Relaxed) {
Ok(_) => break,
Err(observed) => current = observed,
}
}
}
fn sub_saturating(counter: &AtomicU64, delta: u64) {
let mut current = counter.load(Ordering::Relaxed);
loop {
let next = current.saturating_sub(delta);
match counter.compare_exchange_weak(current, next, Ordering::Relaxed, Ordering::Relaxed) {
Ok(_) => break,
Err(observed) => current = observed,
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use onnx_runtime_ir::{Attribute, Node, NodeId};
fn hca_node(cache_format: &str, with_state: bool) -> Node {
let input_count = if with_state { 8 } else { 6 };
let inputs = (0..input_count)
.map(|_| Some(onnx_runtime_ir::ValueId(0)))
.collect();
let mut node = Node::new(NodeId(0), "CompressedSparseAttention", inputs, vec![]);
node.domain = "pkg.nxrt".into();
node.attributes
.insert("compression_ratio".into(), Attribute::Int(128));
node.attributes
.insert("num_heads".into(), Attribute::Int(64));
node.attributes
.insert("head_dim".into(), Attribute::Int(512));
node.attributes
.insert("qk_rope_head_dim".into(), Attribute::Int(64));
node.attributes.insert(
"cache_format".into(),
Attribute::String(cache_format.into()),
);
node
}
#[test]
fn hca_descriptor_accepts_property_compatible_ratio128() {
let node = hca_node("f32", true);
let descriptor = CsaStateGroupDescriptor::from_node(&node, 0, 1, 7).unwrap();
assert_eq!(descriptor.ratio, 128);
assert_eq!(descriptor.cache_format, CsaCacheFormat::F32);
assert!(descriptor.validate().is_ok());
}
#[test]
fn ratio128_rejects_fp4_by_property() {
let node = hca_node("fp4_e2m1_block32", true);
let descriptor = CsaStateGroupDescriptor::from_node(&node, 0, 1, 7).unwrap();
assert_eq!(
descriptor.validate(),
Err(CsaStateGroupRefusal::Ratio128RejectsFp4)
);
}
#[test]
fn ratio4_requires_fp8_and_index_head_dim_128() {
let mut node = hca_node("f32", true);
node.attributes
.insert("compression_ratio".into(), Attribute::Int(4));
let descriptor = CsaStateGroupDescriptor::from_node(&node, 0, 1, 7).unwrap();
assert_eq!(
descriptor.validate(),
Err(CsaStateGroupRefusal::Ratio4RequiresFp8Cache {
cache_format: "f32"
})
);
node.attributes.insert(
"cache_format".into(),
Attribute::String("fp8_e4m3_block64".into()),
);
let descriptor = CsaStateGroupDescriptor::from_node(&node, 0, 1, 7).unwrap();
assert_eq!(
descriptor.validate(),
Err(CsaStateGroupRefusal::Ratio4RequiresIndexHeadDim128 { index_head_dim: 0 })
);
node.attributes
.insert("index_head_dim".into(), Attribute::Int(128));
let descriptor = CsaStateGroupDescriptor::from_node(&node, 0, 1, 7).unwrap();
assert!(descriptor.validate().is_ok());
}
#[test]
fn unsupported_ratio_and_geometry_and_multidevice_and_missing_edges_refuse() {
let mut node = hca_node("f32", true);
node.attributes
.insert("compression_ratio".into(), Attribute::Int(2));
assert_eq!(
CsaStateGroupDescriptor::from_node(&node, 0, 1, 7)
.unwrap()
.validate(),
Err(CsaStateGroupRefusal::UnsupportedRatio { ratio: 2 })
);
let mut node = hca_node("f32", true);
node.attributes.insert("head_dim".into(), Attribute::Int(0));
assert_eq!(
CsaStateGroupDescriptor::from_node(&node, 0, 1, 7)
.unwrap()
.validate(),
Err(CsaStateGroupRefusal::InvalidHeadGeometry {
num_heads: 64,
head_dim: 0
})
);
let node = hca_node("f32", true);
let descriptor = CsaStateGroupDescriptor::from_node(&node, 0, 2, 7).unwrap();
assert_eq!(
descriptor.validate(),
Err(CsaStateGroupRefusal::MultiDeviceAmbiguity { device_count: 2 })
);
let node = hca_node("f32", false);
let descriptor = CsaStateGroupDescriptor::from_node(&node, 0, 1, 7).unwrap();
assert_eq!(
descriptor.validate(),
Err(CsaStateGroupRefusal::MissingStateEdge {
which: "past_compressed_kv"
})
);
}
#[test]
fn unknown_cache_format_is_typed() {
let node = hca_node("bf16_block16", true);
assert_eq!(
CsaStateGroupDescriptor::from_node(&node, 0, 1, 7),
Err(CsaStateGroupRefusal::UnknownCacheFormat {
raw: "bf16_block16".into()
})
);
}
fn bytes(total_split: [u64; 6]) -> CsaStateGroupBytes {
CsaStateGroupBytes {
compressed: total_split[0],
carry: total_split[1],
dense_ring: total_split[2],
index: total_split[3],
index_carry: total_split[4],
scratch: total_split[5],
}
}
#[test]
fn ledger_charges_release_and_return_to_baseline() {
let ledger = Arc::new(CsaStateGroupLedger::default());
assert_eq!(ledger.resident_bytes(), 0);
let charge = ledger
.try_charge((1, 0), bytes([100, 20, 40, 0, 0, 8]))
.unwrap();
assert_eq!(ledger.resident_bytes(), 168);
assert_eq!(ledger.compressed_bytes(), 100);
assert_eq!(ledger.dense_ring_bytes(), 40);
assert_eq!(ledger.peak_bytes(), 168);
assert_eq!(ledger.active_group_count(), 1);
drop(charge);
assert_eq!(ledger.resident_bytes(), 0);
assert_eq!(ledger.active_group_count(), 0);
assert_eq!(ledger.peak_bytes(), 168);
}
#[test]
fn ledger_fails_closed_over_limit_without_mutating() {
let ledger = Arc::new(CsaStateGroupLedger::with_device_limit(150));
let _first = ledger
.try_charge((1, 0), bytes([100, 0, 0, 0, 0, 0]))
.unwrap();
let refusal = ledger
.try_charge((2, 0), bytes([100, 0, 0, 0, 0, 0]))
.unwrap_err();
assert!(matches!(refusal, CsaStateGroupRefusal::OutOfMemory { .. }));
assert_eq!(ledger.resident_bytes(), 100);
assert_eq!(ledger.resident_for(2, 0), 0);
assert_eq!(ledger.charge_failures(), 1);
assert_eq!(ledger.governor_device_used(), 100);
}
#[test]
fn ledger_isolates_requests_and_devices() {
let ledger = Arc::new(CsaStateGroupLedger::default());
let a = ledger
.try_charge((1, 0), bytes([100, 0, 0, 0, 0, 0]))
.unwrap();
let b = ledger
.try_charge((2, 0), bytes([50, 0, 0, 0, 0, 0]))
.unwrap();
let c = ledger
.try_charge((1, 1), bytes([25, 0, 0, 0, 0, 0]))
.unwrap();
assert_eq!(ledger.resident_for(1, 0), 100);
assert_eq!(ledger.resident_for(2, 0), 50);
assert_eq!(ledger.resident_for(1, 1), 25);
assert_eq!(ledger.resident_bytes(), 175);
assert_eq!(ledger.active_group_count(), 3);
drop(b);
assert_eq!(ledger.resident_for(1, 0), 100);
assert_eq!(ledger.resident_for(2, 0), 0);
assert_eq!(ledger.resident_for(1, 1), 25);
assert_eq!(ledger.resident_bytes(), 125);
drop(a);
drop(c);
assert_eq!(ledger.resident_bytes(), 0);
}
#[test]
fn next_request_id_is_monotonic() {
let ledger = Arc::new(CsaStateGroupLedger::default());
let first = ledger.next_request_id();
let second = ledger.next_request_id();
assert!(second > first);
}
#[test]
fn c1_runtime_admits_only_ratio128() {
let node = hca_node("f32", true);
let descriptor = CsaStateGroupDescriptor::from_node(&node, 0, 1, 7).unwrap();
assert!(descriptor.validate_c1_runtime().is_ok());
let mut node = hca_node("fp8_e4m3_block64", true);
node.attributes
.insert("compression_ratio".into(), Attribute::Int(4));
node.attributes
.insert("index_head_dim".into(), Attribute::Int(128));
let descriptor = CsaStateGroupDescriptor::from_node(&node, 0, 1, 7).unwrap();
assert!(
descriptor.validate().is_ok(),
"ratio-4 is a valid op config"
);
assert_eq!(
descriptor.validate_c1_runtime(),
Err(CsaStateGroupRefusal::UnsupportedC1Ratio { ratio: 4 }),
"ratio-4 is out of C1 scope"
);
}
#[test]
fn ledger_reports_governor_device_ceiling() {
let unlimited = Arc::new(CsaStateGroupLedger::default());
assert_eq!(
unlimited.device_available_bytes(),
u64::MAX,
"disarmed default is an unlimited reference governor"
);
let capped = Arc::new(CsaStateGroupLedger::with_device_limit(4096));
assert_eq!(capped.device_available_bytes(), 4096);
let _charge = capped
.try_charge((1, 0), bytes([1000, 0, 0, 0, 0, 0]))
.unwrap();
assert_eq!(capped.device_available_bytes(), 3096);
assert_eq!(capped.governor_device_used(), 1000);
}
#[test]
fn governor_sees_csa_reservation_in_shared_books() {
let ledger_books = Arc::new(LeaseLedger::new(8 << 20, 0, 0));
let governor = Arc::new(LedgerGovernor::new(Arc::clone(&ledger_books)));
let other = governor
.reserve(Tier::Device, 4096, MemoryRole::KvCache, HolderId::new(999))
.unwrap();
let ledger = Arc::new(CsaStateGroupLedger::new(
Arc::clone(&governor) as Arc<dyn MemoryGovernor + Send + Sync>
));
let charge = ledger
.try_charge((1, 0), bytes([100, 20, 40, 0, 0, 8]))
.unwrap();
assert_eq!(governor.used(Tier::Device), 4096 + 168);
assert_eq!(ledger.governor_device_used(), 4096 + 168);
assert_eq!(charge.lease_bytes(), 168);
drop(charge);
assert_eq!(governor.used(Tier::Device), 4096);
drop(other);
assert_eq!(governor.used(Tier::Device), 0);
}
#[derive(Debug, Default)]
struct CountingBooks {
device_used: Mutex<u64>,
reserves: AtomicU64,
releases: AtomicU64,
}
impl onnx_runtime_memory_governor::LeaseAccounting for CountingBooks {
fn try_claim(
&self,
_tier: Tier,
_bytes: u64,
_role: MemoryRole,
) -> Result<(), MemoryError> {
Ok(())
}
fn release(&self, _tier: Tier, bytes: u64) {
self.releases.fetch_add(1, Ordering::Relaxed);
let mut used = self.device_used.lock().unwrap();
*used = used.saturating_sub(bytes);
}
}
#[derive(Debug)]
struct CountingGovernor {
books: Arc<CountingBooks>,
authority: onnx_runtime_memory_governor::MemoryAuthorityId,
}
impl MemoryGovernor for CountingGovernor {
fn authority_id(&self) -> onnx_runtime_memory_governor::MemoryAuthorityId {
self.authority
}
fn reserve(
&self,
tier: Tier,
bytes: u64,
role: MemoryRole,
holder: HolderId,
) -> Result<MemoryLease, MemoryError> {
self.books.reserves.fetch_add(1, Ordering::Relaxed);
{
let mut used = self.books.device_used.lock().unwrap();
*used += bytes;
}
Ok(MemoryLease::new(
tier,
bytes,
role,
holder,
Arc::clone(&self.books) as Arc<dyn onnx_runtime_memory_governor::LeaseAccounting>,
))
}
fn available(&self, _tier: Tier) -> u64 {
u64::MAX
}
fn used(&self, _tier: Tier) -> u64 {
*self.books.device_used.lock().unwrap()
}
}
#[test]
fn charge_reserves_and_releases_governor_exactly_once() {
let books = Arc::new(CountingBooks::default());
let governor = Arc::new(CountingGovernor {
books: Arc::clone(&books),
authority: onnx_runtime_memory_governor::MemoryAuthorityId::new(
onnx_runtime_memory_governor::DeviceKey::device(0),
),
});
let ledger = Arc::new(CsaStateGroupLedger::new(
governor as Arc<dyn MemoryGovernor + Send + Sync>,
));
let charge = ledger
.try_charge((7, 0), bytes([100, 20, 40, 0, 0, 8]))
.unwrap();
assert_eq!(books.reserves.load(Ordering::Relaxed), 1);
assert_eq!(books.releases.load(Ordering::Relaxed), 0);
assert_eq!(*books.device_used.lock().unwrap(), 168);
drop(charge);
assert_eq!(books.reserves.load(Ordering::Relaxed), 1);
assert_eq!(books.releases.load(Ordering::Relaxed), 1);
assert_eq!(*books.device_used.lock().unwrap(), 0);
}
}