use serde::{Deserialize, Serialize};
use crate::trace::memory::MemoryCategory;
pub const SCHEMA: &str = "candle-graph/graph/5";
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
pub struct ExecutionGraph {
#[serde(deserialize_with = "deserialize_schema")]
pub schema: String,
pub spans: Vec<GraphNode>,
pub edges: Vec<GraphEdge>,
pub tensors: Vec<TensorRecord>,
pub gradients: Vec<GradientRecord>,
pub summary: GraphSummary,
}
fn deserialize_schema<'de, D>(deserializer: D) -> std::result::Result<String, D::Error>
where
D: serde::Deserializer<'de>,
{
let schema = String::deserialize(deserializer)?;
if schema != SCHEMA {
return Err(serde::de::Error::custom(format_args!(
"unsupported graph schema {schema:?}; expected {SCHEMA:?}"
)));
}
Ok(schema)
}
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
pub struct TensorRecord {
pub span_id: String,
pub tensor_id: String,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub label: Option<String>,
pub shape: Vec<usize>,
pub dtype: String,
pub device: String,
pub requires_grad: bool,
#[serde(skip_serializing_if = "Option::is_none")]
pub dense_bytes: Option<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 host_self_time_ns: u64,
pub host_total_time_ns: u64,
#[serde(default, skip_serializing_if = "Vec::is_empty")]
pub device_timings: Vec<DeviceNodeTiming>,
#[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, skip_serializing_if = "Option::is_none")]
pub allocated_bytes: Option<u64>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub peak_live_bytes: Option<u64>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub residual_bytes: Option<u64>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub dense_bytes: Option<u64>,
}
#[derive(Debug, Clone, PartialEq, Eq, Hash, Serialize, Deserialize)]
#[serde(rename_all = "snake_case")]
pub enum GraphNodeKind {
Root,
Function,
Module,
Op,
Tensor,
#[serde(other)]
Other,
}
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
pub struct DeviceNodeTiming {
pub device: String,
pub clock_id: String,
pub busy_ns: u64,
}
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
#[serde(tag = "kind", rename_all = "snake_case")]
pub enum GraphEdge {
Call {
from_span: String,
to_span: String,
host_duration_ns: u64,
#[serde(default, skip_serializing_if = "Option::is_none")]
label: Option<String>,
},
Data {
from_tensor: String,
to_tensor: String,
#[serde(default, skip_serializing_if = "Option::is_none")]
label: Option<String>,
},
}
#[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 outer_wall_time_ns: u64,
pub slowest_host_spans: Vec<HostSpanCost>,
pub slowest_device_spans: Vec<DeviceSpanCost>,
pub heaviest_spans: Vec<HeavySpan>,
}
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
pub struct HostSpanCost {
pub id: String,
pub name: String,
pub scope: MeasuredHostScope,
pub host_self_time_ns: u64,
pub measured_overlap_self_time_ns: u64,
pub full_duration_ns: u64,
pub measured_overlap_duration_ns: u64,
}
#[derive(Debug, Clone, Copy, Default, PartialEq, Eq, Hash, Serialize, Deserialize)]
#[serde(rename_all = "snake_case")]
pub enum MeasuredHostScope {
#[default]
MeasuredSubtree,
ConcurrentOverlap,
}
impl MeasuredHostScope {
pub const fn as_str(self) -> &'static str {
match self {
Self::MeasuredSubtree => "measured_subtree",
Self::ConcurrentOverlap => "concurrent_overlap",
}
}
}
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
pub struct DeviceSpanCost {
pub id: String,
pub name: String,
pub device: String,
pub clock_id: String,
pub device_busy_ns: u64,
}
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
pub struct HeavySpan {
pub id: String,
pub name: String,
pub allocated_bytes: u64,
pub peak_live_bytes: Option<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()
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn measured_scope_fields_are_required() {
let error = serde_json::from_value::<HostSpanCost>(serde_json::json!({
"id": "span",
"name": "work",
"host_self_time_ns": 17
}))
.unwrap_err();
assert!(error.to_string().contains("scope"));
}
#[test]
fn measured_host_scope_has_stable_wire_names() {
assert_eq!(
serde_json::to_value(MeasuredHostScope::MeasuredSubtree).unwrap(),
serde_json::json!("measured_subtree")
);
assert_eq!(
serde_json::to_value(MeasuredHostScope::ConcurrentOverlap).unwrap(),
serde_json::json!("concurrent_overlap")
);
}
}