use serde::{Deserialize, Serialize};
use crate::trace::memory::{MemoryCategory, MemoryProfile, MemorySummary};
pub const SCHEMA: &str = "candle-graph/graph/3";
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
pub struct ExecutionGraph {
pub schema: String,
pub spans: Vec<GraphNode>,
pub edges: Vec<GraphEdge>,
pub tensors: Vec<TensorRecord>,
pub gradients: Vec<GradientRecord>,
pub summary: GraphSummary,
pub memory: MemoryProfile,
}
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
pub struct TensorRecord {
pub span_id: String,
pub tensor_id: String,
pub shape: Vec<usize>,
pub dtype: String,
pub device: String,
pub requires_grad: bool,
pub storage_bytes: u64,
pub category: MemoryCategory,
}
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
pub struct GraphNode {
pub id: String,
pub parent_id: Option<String>,
pub name: String,
pub kind: GraphNodeKind,
pub start_ns: u64,
pub self_time_ns: u64,
pub total_time_ns: u64,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub shape: Option<Vec<usize>>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub dtype: Option<String>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub device: Option<String>,
#[serde(default)]
pub bytes: u64,
#[serde(default)]
pub peak_bytes: u64,
#[serde(default)]
pub residual_bytes: u64,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub storage_bytes: Option<u64>,
}
#[derive(Debug, Clone, PartialEq, Eq, Hash, Serialize, Deserialize)]
#[serde(rename_all = "snake_case")]
pub enum GraphNodeKind {
Root,
Function,
Module,
Op,
#[serde(other)]
Other,
}
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
pub struct GraphEdge {
pub from: String,
pub to: String,
pub kind: GraphEdgeKind,
pub duration_ns: u64,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub label: Option<String>,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, Serialize, Deserialize)]
#[serde(rename_all = "snake_case")]
pub enum GraphEdgeKind {
Call,
Data,
}
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
pub struct GradientRecord {
pub root: String,
pub key: String,
pub state: GradientRecordState,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub norm: Option<f64>,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, Serialize, Deserialize)]
#[serde(rename_all = "snake_case")]
pub enum GradientRecordState {
Present,
Missing,
Zero,
NonFinite,
}
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
pub struct GraphSummary {
pub entrypoint: String,
pub total_ms: f64,
pub slowest_spans: Vec<SlowSpan>,
pub heaviest_spans: Vec<HeavySpan>,
pub memory: MemorySummary,
}
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
pub struct SlowSpan {
pub id: String,
pub name: String,
pub self_time_ns: u64,
}
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
pub struct HeavySpan {
pub id: String,
pub name: String,
pub bytes: u64,
pub peak_bytes: u64,
}
impl ExecutionGraph {
pub fn node(&self, id: &str) -> Option<&GraphNode> {
self.spans.iter().find(|n| n.id == id)
}
pub fn children(&self, parent_id: &str) -> Vec<&GraphNode> {
self.spans
.iter()
.filter(|n| n.parent_id.as_deref() == Some(parent_id))
.collect()
}
}