use crate::id::KernelId;
use crate::memory_management::ManagedMemoryId;
use crate::server::{IoError, ServerError};
use alloc::boxed::Box;
use alloc::vec::Vec;
use core::num::NonZeroU64;
use cubecl_environment::backtrace::BackTrace;
use cubecl_environment::collections::HashMap;
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, PartialOrd, Ord)]
pub struct FailureId(NonZeroU64);
impl core::fmt::Display for FailureId {
fn fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
write!(f, "#{}", self.0)
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum Claim {
Failed(FailureId, ManagedMemoryId),
Unallocated(ManagedMemoryId),
}
#[derive(Debug, Default)]
pub struct ErrorGraph {
nodes: HashMap<FailureId, Failure>,
minted: u64,
}
#[derive(Debug)]
struct Failure {
error: ServerError,
tagged: u32,
skipped: Vec<Skipped>,
skipped_total: u64,
}
#[derive(Debug, Clone)]
pub struct Skipped {
pub kernel: KernelId,
pub needed: ManagedMemoryId,
pub produced: Vec<ManagedMemoryId>,
}
impl ErrorGraph {
pub const MAX_SKIPPED: usize = 16;
pub fn insert(&mut self, error: ServerError) -> FailureId {
self.minted = self
.minted
.checked_add(1)
.expect("a failure id was minted for every u64");
let id = FailureId(NonZeroU64::new(self.minted).expect("minted starts above zero"));
self.nodes.insert(
id,
Failure {
error,
tagged: 0,
skipped: Vec::new(),
skipped_total: 0,
},
);
id
}
pub fn skipped(&mut self, failure: FailureId, record: Skipped) {
let Some(node) = self.nodes.get_mut(&failure) else {
return;
};
node.skipped_total += 1;
if node.skipped.len() == Self::MAX_SKIPPED {
node.skipped.remove(0);
}
node.skipped.push(record);
}
pub fn report(&self, failure: FailureId, memory: ManagedMemoryId) -> Option<ServerError> {
let node = self.nodes.get(&failure)?;
let mut chain = Vec::new();
let mut target = memory;
let mut upper = node.skipped.len();
while let Some(found) = node.skipped[..upper]
.iter()
.rposition(|record| record.produced.contains(&target))
{
let record = &node.skipped[found];
chain.push(alloc::format!(
"skipped `{}`: it needed memory {:?}, which carried the failure",
record.kernel.short_name(),
record.needed,
));
target = record.needed;
upper = found;
}
let dropped = node.skipped_total.saturating_sub(node.skipped.len() as u64);
if !chain.is_empty() && dropped > 0 {
chain.push(alloc::format!(
"({dropped} older skip record(s) were dropped; the walk may stop before the root)"
));
}
Some(ServerError::Unwritten {
failure: failure.0.get(),
claimed: node.tagged,
chain,
root: Box::new(node.error.clone()),
backtrace: BackTrace::capture(),
})
}
pub fn reports(&self, claims: impl Iterator<Item = Claim>) -> Result<(), ServerError> {
let mut seen: Vec<FailureId> = Vec::new();
let mut errors = Vec::new();
for claim in claims {
let (failure, memory) = match claim {
Claim::Failed(failure, memory) => (failure, memory),
Claim::Unallocated(memory) => {
errors.push(Self::unallocated(memory));
continue;
}
};
if seen.contains(&failure) {
continue;
}
seen.push(failure);
if let Some(error) = self.report(failure, memory) {
errors.push(error);
}
}
match errors.is_empty() {
true => Ok(()),
false => Err(ServerError::Several {
errors,
backtrace: BackTrace::capture(),
}),
}
}
fn unallocated(memory: ManagedMemoryId) -> ServerError {
IoError::NotFound {
backtrace: BackTrace::capture(),
reason: alloc::format!(
"memory {memory:?} was never allocated: the reservation behind it failed"
)
.into(),
}
.into()
}
pub(crate) fn tag(&mut self, failure: FailureId) {
self.node_mut(failure).tagged += 1;
}
pub fn untag(&mut self, failure: Option<FailureId>) {
let Some(failure) = failure else {
return;
};
let node = self.node_mut(failure);
debug_assert!(node.tagged > 0, "{failure} was shed more often than tagged");
node.tagged = node.tagged.saturating_sub(1);
if node.tagged == 0 {
self.nodes.remove(&failure);
}
}
pub fn replace(&mut self, failure: FailureId, error: ServerError) {
if let Some(node) = self.nodes.get_mut(&failure) {
node.error = error;
}
}
pub fn prune(&mut self, failure: FailureId) {
if let Some(node) = self.nodes.get(&failure)
&& node.tagged == 0
{
self.nodes.remove(&failure);
}
}
pub fn error(&self, failure: FailureId) -> Option<&ServerError> {
self.nodes.get(&failure).map(|node| &node.error)
}
pub fn len(&self) -> usize {
self.nodes.len()
}
pub fn is_empty(&self) -> bool {
self.nodes.is_empty()
}
fn node_mut(&mut self, failure: FailureId) -> &mut Failure {
self.nodes
.get_mut(&failure)
.expect("a carried failure id always has its node")
}
}
#[cfg(test)]
mod tests {
use super::*;
use alloc::string::ToString;
fn error(reason: &str) -> ServerError {
ServerError::Generic {
reason: reason.to_string(),
backtrace: Default::default(),
}
}
#[test]
fn a_node_lives_while_something_carries_it_and_no_longer() {
let mut graph = ErrorGraph::default();
let failure = graph.insert(error("launch"));
graph.tag(failure);
graph.tag(failure);
assert_eq!(graph.len(), 1);
graph.untag(Some(failure));
assert!(graph.error(failure).is_some(), "one carrier remains");
graph.untag(Some(failure));
assert!(
graph.error(failure).is_none(),
"nothing carries it, so it is gone"
);
assert!(graph.is_empty());
}
#[test]
fn a_failure_that_tainted_nothing_is_pruned() {
let mut graph = ErrorGraph::default();
let failure = graph.insert(error("dry-run"));
graph.prune(failure);
assert!(graph.is_empty());
}
#[test]
fn replacing_an_error_leaves_the_carriers_pointing_at_the_new_one() {
let mut graph = ErrorGraph::default();
let failure = graph.insert(error("torn down"));
graph.tag(failure);
graph.replace(failure, error("launch"));
match graph.error(failure) {
Some(ServerError::Generic { reason, .. }) => assert_eq!(reason, "launch"),
other => panic!("expected the replaced error, got {other:?}"),
}
graph.untag(Some(failure));
assert!(graph.is_empty());
}
fn skip(
kernel_name: KernelId,
needed: ManagedMemoryId,
produced: &[ManagedMemoryId],
) -> Skipped {
Skipped {
kernel: kernel_name,
needed,
produced: produced.to_vec(),
}
}
fn memory_id(value: usize) -> ManagedMemoryId {
ManagedMemoryId { value }
}
struct Fill;
struct Matmul;
struct Gelu;
#[test]
fn a_report_walks_the_skip_chain_back_to_the_root() {
let mut graph = ErrorGraph::default();
let failure = graph.insert(error("fill_f32 failed to compile"));
graph.tag(failure);
let (root_out, mid, last) = (memory_id(77), memory_id(91), memory_id(103));
graph.skipped(failure, skip(KernelId::new::<Matmul>(), root_out, &[mid]));
graph.skipped(failure, skip(KernelId::new::<Gelu>(), mid, &[last]));
let report = graph.report(failure, last).unwrap();
let text = alloc::format!("{report}");
let gelu = text.find("Gelu").expect("the newest hop comes first");
let matmul = text.find("Matmul").expect("then the one it needed");
assert!(gelu < matmul, "newest skip first, root last: {text}");
assert!(
text.contains("fill_f32 failed to compile"),
"the root is always in the report: {text}"
);
assert!(
text.contains(&alloc::format!("#{}", failure.0.get())),
"the failure id ties reads of the same failure together: {text}"
);
let report = graph.report(failure, memory_id(555)).unwrap();
let text = alloc::format!("{report}");
assert!(!text.contains("Gelu") && text.contains("fill_f32 failed to compile"));
}
#[test]
fn the_skip_cap_keeps_the_newest_records() {
let mut graph = ErrorGraph::default();
let failure = graph.insert(error("root"));
graph.tag(failure);
for i in 0..(ErrorGraph::MAX_SKIPPED + 4) {
graph.skipped(
failure,
skip(KernelId::new::<Fill>(), memory_id(i), &[memory_id(i + 1)]),
);
}
let newest = memory_id(ErrorGraph::MAX_SKIPPED + 4);
let report = graph.report(failure, newest).unwrap();
let text = alloc::format!("{report}");
assert!(
text.contains("Fill"),
"the newest buffer still has an entry to walk from: {text}"
);
assert!(
text.contains("4 older skip record(s) were dropped"),
"and the report says the walk may stop early: {text}"
);
}
#[test]
fn pruning_leaves_a_carried_failure_alone() {
let mut graph = ErrorGraph::default();
let failure = graph.insert(error("launch"));
graph.tag(failure);
graph.prune(failure);
assert!(graph.error(failure).is_some());
}
}