use serde::{Deserialize, Serialize};
use super::memory::{MemoryAction, MemoryCategory};
use super::schema::{GradientState, SpanKind, TraceRunMeta, SCHEMA};
use crate::phase::ExecutionStep;
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
#[serde(tag = "kind", rename_all = "snake_case")]
pub enum TraceEvent {
Meta {
schema: String,
#[serde(flatten)]
run: TraceRunMeta,
},
SpanStart(SpanStartEvent),
SpanEnd(SpanEndEvent),
Op(OpEvent),
Tensor(TensorEvent),
Memory(MemoryEvent),
DeviceMemory(DeviceMemoryEvent),
Gradient(GradientEvent),
Edge(EdgeEvent),
}
impl TraceEvent {
pub fn meta(run: TraceRunMeta) -> Self {
Self::Meta {
schema: SCHEMA.to_string(),
run,
}
}
}
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
pub struct SpanStartEvent {
pub id: String,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub parent_id: Option<String>,
pub name: String,
pub start_ns: u64,
#[serde(rename = "span_kind")]
pub kind: SpanKind,
#[serde(default)]
pub measured: bool,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub step: Option<ExecutionStep>,
}
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
pub struct SpanEndEvent {
pub id: String,
#[serde(default)]
pub duration_ns: u64,
}
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
pub struct OpEvent {
pub span_id: String,
pub op_name: String,
#[serde(default)]
pub inputs: Vec<String>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub output: Option<String>,
#[serde(default)]
pub shape: Vec<usize>,
pub dtype: String,
pub device: String,
pub duration_ns: u64,
#[serde(default)]
pub timestamp_ns: u64,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub storage_bytes: Option<u64>,
#[serde(default)]
pub input_storage_bytes: u64,
}
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
pub struct TensorEvent {
pub span_id: String,
pub tensor_id: String,
#[serde(default)]
pub shape: Vec<usize>,
pub dtype: String,
pub device: String,
#[serde(default)]
pub requires_grad: bool,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub storage_bytes: Option<u64>,
#[serde(default)]
pub category: MemoryCategory,
}
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
pub struct MemoryEvent {
pub timestamp_ns: u64,
pub tensor_id: String,
pub span_id: String,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub op_name: Option<String>,
pub device: String,
pub bytes: u64,
pub action: MemoryAction,
#[serde(default)]
pub shape: Vec<usize>,
pub dtype: String,
#[serde(default)]
pub category: MemoryCategory,
}
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
pub struct DeviceMemoryEvent {
pub timestamp_ns: u64,
pub device: String,
pub used_bytes: u64,
pub free_bytes: u64,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub reserved_bytes: Option<u64>,
}
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
pub struct GradientEvent {
pub event_id: String,
pub root: String,
pub key: String,
pub state: GradientState,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub norm: Option<f64>,
}
impl GradientEvent {
pub fn param_key(&self) -> &str {
&self.key
}
}
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
pub struct EdgeEvent {
pub from_span: String,
pub to_span: String,
pub duration_ns: u64,
}