use std::collections::BTreeMap;
use std::fmt;
use crate::phase::ExecutionPhase;
use crate::phase::ExecutionStep;
use serde::{Deserialize, Serialize};
use crate::capability::CaptureContract;
pub const SCHEMA: &str = "candle-graph/trace/10";
pub const PREVIOUS_SCHEMA: &str = "candle-graph/trace/9";
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
pub struct ComparisonIdentity {
#[serde(default, skip_serializing_if = "Option::is_none")]
pub implementation_id: Option<String>,
pub workload_id: String,
pub model_id: String,
pub config_id: String,
pub data_id: String,
pub seed_policy: String,
pub physical_batch: u64,
pub accumulation_steps: u64,
pub precision: String,
pub device_state: String,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub pair_id: Option<String>,
}
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
pub struct TraceRunMeta {
pub run_id: String,
pub correlation_id: String,
pub entrypoint: String,
pub phase: ExecutionPhase,
pub timestamp: String,
pub capture_step: u64,
pub warmup_steps: u64,
pub device: String,
#[serde(default)]
pub measured_region_device_synchronized: bool,
pub timing_mode: TimingMode,
pub capture_contract: CaptureContract,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub comparison_identity: Option<ComparisonIdentity>,
pub tags: BTreeMap<String, String>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub candle_version: Option<String>,
}
impl TraceRunMeta {
pub fn validate(&self) -> anyhow::Result<()> {
for (label, value) in [
("run_id", &self.run_id),
("correlation_id", &self.correlation_id),
("entrypoint", &self.entrypoint),
("timestamp", &self.timestamp),
("device", &self.device),
] {
anyhow::ensure!(
!value.trim().is_empty(),
"run provenance {label} must not be empty"
);
}
anyhow::ensure!(
self.capture_step > 0,
"capture_step must be one-based and greater than zero"
);
anyhow::ensure!(
self.warmup_steps < self.capture_step,
"warmup_steps ({}) must be fewer than the one-based capture_step ({})",
self.warmup_steps,
self.capture_step
);
if let Some(identity) = &self.comparison_identity {
identity.validate()?;
}
Ok(())
}
}
impl ComparisonIdentity {
pub fn validate(&self) -> anyhow::Result<()> {
for (label, value) in [
("workload_id", &self.workload_id),
("model_id", &self.model_id),
("config_id", &self.config_id),
("data_id", &self.data_id),
("seed_policy", &self.seed_policy),
("precision", &self.precision),
("device_state", &self.device_state),
] {
anyhow::ensure!(
!value.trim().is_empty(),
"comparison identity {label} must not be empty"
);
}
for (label, value) in [
("implementation_id", &self.implementation_id),
("pair_id", &self.pair_id),
] {
if let Some(value) = value {
anyhow::ensure!(
!value.trim().is_empty(),
"comparison identity {label} must not be empty when declared"
);
}
}
anyhow::ensure!(
self.physical_batch > 0,
"comparison identity physical_batch must be greater than zero"
);
anyhow::ensure!(
self.accumulation_steps > 0,
"comparison identity accumulation_steps must be greater than zero"
);
Ok(())
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
#[serde(rename_all = "snake_case")]
pub enum TimingMode {
Host,
DeviceSynchronized,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
#[serde(rename_all = "snake_case")]
pub enum RunOutcome {
Complete,
Failed,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, PartialOrd, Ord, Hash, Serialize, Deserialize)]
#[serde(rename_all = "snake_case")]
pub enum SpanKind {
Function,
Op,
Module,
}
impl fmt::Display for SpanKind {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
match self {
Self::Function => write!(f, "function"),
Self::Op => write!(f, "op"),
Self::Module => write!(f, "module"),
}
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, Serialize, Deserialize)]
#[serde(rename_all = "snake_case")]
pub enum GradientState {
Present,
Missing,
Zero,
NonFinite,
}
impl GradientState {
pub(crate) fn norm_is_valid(self, norm: Option<f64>) -> bool {
match (self, norm) {
(Self::Present, Some(norm)) => norm.is_finite() && norm > 0.0,
(Self::Zero, Some(norm)) => norm == 0.0 && !norm.is_sign_negative(),
(Self::Missing | Self::NonFinite, None) => true,
_ => false,
}
}
}
impl fmt::Display for GradientState {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
match self {
Self::Present => write!(f, "present"),
Self::Missing => write!(f, "missing"),
Self::Zero => write!(f, "zero"),
Self::NonFinite => write!(f, "non_finite"),
}
}
}
#[derive(Debug, Clone, Default, PartialEq, Eq, Serialize, Deserialize)]
pub struct TraceSummary {
pub op_count: usize,
pub total_ns: u64,
pub span_count: usize,
pub root_span_count: usize,
pub max_depth: usize,
pub alloc_count: usize,
pub free_count: usize,
pub logical_peak_bytes: Option<u64>,
}
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
pub struct SpanRecord {
pub id: String,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub parent_id: Option<String>,
pub name: String,
pub kind: SpanKind,
pub measured: bool,
pub start_ns: u64,
#[serde(default)]
pub closed: bool,
#[serde(default)]
pub duration_ns: u64,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub step: Option<ExecutionStep>,
}