Skip to main content

candle_graph/trace/
document.rs

1//! Aggregate trace document and JSONL I/O for `candle-graph/trace/10`.
2
3use std::collections::{BTreeMap, HashMap, HashSet};
4use std::fs::File;
5use std::io::{BufRead, BufReader, Write};
6use std::path::Path;
7
8use anyhow::{bail, Context, Result};
9use serde::{Deserialize, Serialize};
10
11use super::events::{
12    DeviceIntervalEvent, DeviceMemoryEvent, EdgeEvent, GradientEvent, MemoryEvent, OpEvent,
13    SpanEndEvent, SpanStartEvent, TensorEvent, TensorStatsEvent, TerminalEvent, TraceEvent,
14};
15use super::memory::{resolve_dense_tensor_bytes, MemoryAction};
16use super::schema::{RunOutcome, SpanRecord, TraceRunMeta, TraceSummary, PREVIOUS_SCHEMA, SCHEMA};
17
18/// Full trace document assembled from JSONL events.
19#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
20pub struct TraceDocument {
21    pub schema: String,
22    pub run: TraceRunMeta,
23    #[serde(default)]
24    pub spans: Vec<SpanRecord>,
25    #[serde(default)]
26    pub ops: Vec<OpEvent>,
27    #[serde(default)]
28    pub tensors: Vec<TensorEvent>,
29    #[serde(default)]
30    pub tensor_stats: Vec<TensorStatsEvent>,
31    #[serde(default)]
32    pub memory: Vec<MemoryEvent>,
33    #[serde(default)]
34    pub device_memory: Vec<DeviceMemoryEvent>,
35    #[serde(default)]
36    pub device_intervals: Vec<DeviceIntervalEvent>,
37    #[serde(default)]
38    pub gradients: Vec<GradientEvent>,
39    #[serde(default)]
40    pub edges: Vec<EdgeEvent>,
41    pub terminal: TerminalEvent,
42}
43
44impl TraceDocument {
45    /// Build a document from an ordered event stream (meta must be first).
46    pub fn from_events(events: impl IntoIterator<Item = TraceEvent>) -> Result<Self> {
47        let mut schema: Option<String> = None;
48        let mut run: Option<TraceRunMeta> = None;
49        let mut span_starts: BTreeMap<String, SpanStartEvent> = BTreeMap::new();
50        let mut span_durations: HashMap<String, u64> = HashMap::new();
51        let mut span_closed: HashSet<String> = HashSet::new();
52        let mut ops = Vec::new();
53        let mut tensors = Vec::new();
54        let mut tensor_stats = Vec::new();
55        let mut memory = Vec::new();
56        let mut device_memory = Vec::new();
57        let mut device_intervals = Vec::new();
58        let mut gradients = Vec::new();
59        let mut edges = Vec::new();
60        let mut terminal: Option<TerminalEvent> = None;
61
62        for (index, event) in events.into_iter().enumerate() {
63            if terminal.is_some() {
64                bail!("terminal event must be the final trace record; found another event at index {index}");
65            }
66            match event {
67                TraceEvent::Meta {
68                    schema: s,
69                    run: meta,
70                } => {
71                    if index != 0 {
72                        bail!("meta event must be the first non-empty trace record, found at index {index}");
73                    }
74                    if schema.is_some() || run.is_some() {
75                        bail!(
76                            "duplicate meta event at index {index}; only one meta record is allowed"
77                        );
78                    }
79                    schema = Some(s);
80                    run = Some(*meta);
81                }
82                TraceEvent::SpanStart(start) => {
83                    if span_starts.contains_key(&start.id) {
84                        bail!("duplicate span_start id `{}` at index {index}", start.id);
85                    }
86                    span_starts.insert(start.id.clone(), start);
87                }
88                TraceEvent::SpanEnd(SpanEndEvent { id, duration_ns }) => {
89                    if !span_starts.contains_key(&id) {
90                        bail!("span_end for unknown span `{id}` at index {index}");
91                    }
92                    if span_closed.contains(&id) {
93                        bail!("duplicate span_end for `{id}` at index {index}");
94                    }
95                    span_closed.insert(id.clone());
96                    span_durations.insert(id, duration_ns);
97                }
98                TraceEvent::Op(mut op) => {
99                    op.output_dense_bytes =
100                        resolve_dense_tensor_bytes(op.output_dense_bytes, &op.shape, &op.dtype);
101                    ops.push(op);
102                }
103                TraceEvent::Tensor(mut tensor) => {
104                    tensor.dense_bytes = resolve_dense_tensor_bytes(
105                        tensor.dense_bytes,
106                        &tensor.shape,
107                        &tensor.dtype,
108                    );
109                    tensors.push(tensor);
110                }
111                TraceEvent::TensorStats(stats) => {
112                    let shape_elements = stats
113                        .shape
114                        .iter()
115                        .try_fold(1u64, |total, &dim| total.checked_mul(dim as u64));
116                    if stats.label.trim().is_empty()
117                        || shape_elements != Some(stats.elements)
118                        || stats.non_finite > stats.elements
119                        || !stats.rms.is_finite()
120                        || stats.rms < 0.0
121                        || !stats.abs_max.is_finite()
122                        || stats.abs_max < 0.0
123                        || !stats.mean.is_finite()
124                    {
125                        bail!("invalid tensor_stats event at index {index}");
126                    }
127                    tensor_stats.push(stats);
128                }
129                TraceEvent::Memory(mem) => memory.push(mem),
130                TraceEvent::DeviceMemory(snapshot) => device_memory.push(snapshot),
131                TraceEvent::DeviceInterval(interval) => device_intervals.push(interval),
132                TraceEvent::Gradient(gradient) => gradients.push(gradient),
133                TraceEvent::Edge(edge) => edges.push(edge),
134                TraceEvent::Terminal(event) => {
135                    if terminal.replace(event).is_some() {
136                        bail!("duplicate terminal event at index {index}");
137                    }
138                }
139            }
140        }
141
142        let schema = schema.unwrap_or_else(|| SCHEMA.to_string());
143        let run = run.context("trace stream is missing a meta event with run metadata")?;
144        let terminal = terminal.context("trace stream is missing its terminal event")?;
145
146        if schema != SCHEMA && schema != PREVIOUS_SCHEMA {
147            bail!(
148                "unsupported trace schema {schema:?}; expected {SCHEMA:?} or {PREVIOUS_SCHEMA:?}"
149            );
150        }
151        if schema == PREVIOUS_SCHEMA && !tensor_stats.is_empty() {
152            bail!(
153                "trace schema {PREVIOUS_SCHEMA:?} does not define tensor_stats events; \
154                 producers emitting tensor statistics must declare {SCHEMA:?}"
155            );
156        }
157
158        let mut spans: Vec<SpanRecord> = span_starts
159            .into_iter()
160            .map(|(id, start)| SpanRecord {
161                id: id.clone(),
162                parent_id: start.parent_id,
163                name: start.name,
164                kind: start.kind,
165                measured: start.measured,
166                start_ns: start.start_ns,
167                closed: span_closed.contains(&id),
168                duration_ns: span_durations.get(&id).copied().unwrap_or(0),
169                step: start.step,
170            })
171            .collect();
172        spans.sort_by(|a, b| a.id.cmp(&b.id));
173
174        match terminal.outcome {
175            RunOutcome::Complete if terminal.reason.is_some() => {
176                bail!("complete terminal event cannot contain a failure reason")
177            }
178            RunOutcome::Failed
179                if terminal
180                    .reason
181                    .as_deref()
182                    .is_none_or(|reason| reason.trim().is_empty()) =>
183            {
184                bail!("failed terminal event requires a non-empty reason")
185            }
186            _ => {}
187        }
188        let latest_host_timestamp_ns = spans
189            .iter()
190            .map(|span| {
191                span.start_ns
192                    .saturating_add(if span.closed { span.duration_ns } else { 0 })
193            })
194            .chain(
195                ops.iter()
196                    .map(|op| op.timestamp_ns.saturating_add(op.duration_ns)),
197            )
198            .chain(memory.iter().map(|event| event.timestamp_ns))
199            .chain(device_memory.iter().map(|event| event.timestamp_ns))
200            .max()
201            .unwrap_or(0);
202        if terminal.timestamp_ns < latest_host_timestamp_ns {
203            bail!(
204                "terminal timestamp {} precedes host evidence ending at {latest_host_timestamp_ns}",
205                terminal.timestamp_ns
206            );
207        }
208
209        Ok(Self {
210            schema,
211            run,
212            spans,
213            ops,
214            tensors,
215            tensor_stats,
216            memory,
217            device_memory,
218            device_intervals,
219            gradients,
220            edges,
221            terminal,
222        })
223    }
224
225    /// Profiler summary: op count, total wall time, span tree shape, memory totals.
226    pub fn build_summary(&self) -> TraceSummary {
227        let op_count = self.ops.len();
228        let total_ns = self
229            .spans
230            .iter()
231            .filter(|span| span.measured)
232            .map(|span| span.duration_ns)
233            .sum();
234        let span_count = self.spans.len();
235        let root_span_count = self
236            .spans
237            .iter()
238            .filter(|span| span.parent_id.is_none())
239            .count();
240
241        let parent_by_id: HashMap<&str, Option<&str>> = self
242            .spans
243            .iter()
244            .map(|span| (span.id.as_str(), span.parent_id.as_deref()))
245            .collect();
246
247        let mut max_depth = 0usize;
248        for span in &self.spans {
249            let mut depth = 0usize;
250            let mut current_parent = span.parent_id.as_deref();
251            let mut seen = HashSet::new();
252            while let Some(parent_id) = current_parent {
253                if !seen.insert(parent_id) {
254                    break;
255                }
256                depth += 1;
257                current_parent = parent_by_id.get(parent_id).copied().flatten();
258            }
259            max_depth = max_depth.max(depth);
260        }
261
262        let alloc_count = self
263            .memory
264            .iter()
265            .filter(|event| event.action == MemoryAction::Alloc)
266            .count();
267        let free_count = self
268            .memory
269            .iter()
270            .filter(|event| event.action == MemoryAction::Free)
271            .count();
272
273        let logical_peak_bytes = super::memory::analyze_memory(self)
274            .logical
275            .and_then(|profile| profile.peak.map(|peak| peak.live_bytes));
276
277        TraceSummary {
278            op_count,
279            total_ns,
280            span_count,
281            root_span_count,
282            max_depth,
283            alloc_count,
284            free_count,
285            logical_peak_bytes,
286        }
287    }
288
289    /// Flatten the document back into JSONL events (meta first, then body in stable order).
290    pub fn to_events(&self) -> Vec<TraceEvent> {
291        let mut events = vec![TraceEvent::Meta {
292            schema: self.schema.clone(),
293            run: Box::new(self.run.clone()),
294        }];
295
296        let mut span_ids: Vec<_> = self.spans.iter().map(|span| span.id.as_str()).collect();
297        span_ids.sort_unstable();
298        for id in span_ids {
299            let span = self
300                .spans
301                .iter()
302                .find(|span| span.id == id)
303                .expect("sorted id must exist");
304            events.push(TraceEvent::SpanStart(SpanStartEvent {
305                id: span.id.clone(),
306                parent_id: span.parent_id.clone(),
307                name: span.name.clone(),
308                kind: span.kind,
309                measured: span.measured,
310                start_ns: span.start_ns,
311                step: span.step,
312            }));
313            if span.closed {
314                events.push(TraceEvent::SpanEnd(SpanEndEvent {
315                    id: span.id.clone(),
316                    duration_ns: span.duration_ns,
317                }));
318            }
319        }
320
321        events.extend(self.ops.iter().cloned().map(TraceEvent::Op));
322        events.extend(self.tensors.iter().cloned().map(TraceEvent::Tensor));
323        events.extend(
324            self.tensor_stats
325                .iter()
326                .cloned()
327                .map(TraceEvent::TensorStats),
328        );
329        events.extend(self.memory.iter().cloned().map(TraceEvent::Memory));
330        events.extend(
331            self.device_memory
332                .iter()
333                .cloned()
334                .map(TraceEvent::DeviceMemory),
335        );
336        events.extend(
337            self.device_intervals
338                .iter()
339                .cloned()
340                .map(TraceEvent::DeviceInterval),
341        );
342        events.extend(self.gradients.iter().cloned().map(TraceEvent::Gradient));
343        events.extend(self.edges.iter().cloned().map(TraceEvent::Edge));
344        events.push(TraceEvent::Terminal(self.terminal.clone()));
345        events
346    }
347}
348
349/// Parse a JSONL trace file into a [`TraceDocument`].
350pub fn parse_trace(path: impl AsRef<Path>) -> Result<TraceDocument> {
351    let path = path.as_ref();
352    let file = File::open(path).with_context(|| format!("open trace file {}", path.display()))?;
353    let reader = BufReader::new(file);
354    let mut events = Vec::new();
355
356    for (line_no, line) in reader.lines().enumerate() {
357        let line = line.with_context(|| {
358            format!(
359                "read trace JSONL line {} from {}",
360                line_no + 1,
361                path.display()
362            )
363        })?;
364        let trimmed = line.trim();
365        if trimmed.is_empty() {
366            continue;
367        }
368        let event: TraceEvent = serde_json::from_str(trimmed).with_context(|| {
369            format!(
370                "parse trace JSONL line {} from {}",
371                line_no + 1,
372                path.display()
373            )
374        })?;
375        events.push(event);
376    }
377
378    TraceDocument::from_events(events)
379}
380
381/// Write JSONL events to `path` (creates parent directories when needed).
382pub fn write_jsonl(path: impl AsRef<Path>, events: &[TraceEvent]) -> Result<()> {
383    let path = path.as_ref();
384    if let Some(parent) = path.parent() {
385        std::fs::create_dir_all(parent)
386            .with_context(|| format!("create trace dir {}", parent.display()))?;
387    }
388    let mut file =
389        File::create(path).with_context(|| format!("create trace file {}", path.display()))?;
390    for event in events {
391        let mut line = serde_json::to_vec(event).context("serialize trace JSONL event")?;
392        line.push(b'\n');
393        file.write_all(&line)
394            .with_context(|| format!("write trace JSONL to {}", path.display()))?;
395    }
396    file.flush()
397        .with_context(|| format!("flush trace JSONL to {}", path.display()))?;
398    Ok(())
399}
400
401#[cfg(test)]
402mod tests {
403    use super::*;
404    use crate::capability::CaptureContract;
405    use crate::trace::events::TraceEvent;
406    use crate::trace::memory::MemoryCategory;
407    use crate::trace::schema::{GradientState, RunOutcome, SpanKind};
408
409    fn sample_meta() -> TraceRunMeta {
410        TraceRunMeta {
411            run_id: "run-1".into(),
412            correlation_id: "demo/update-1".into(),
413            entrypoint: "demo::train::loss".into(),
414            phase: crate::phase::ExecutionPhase::Train,
415            timestamp: "2026-08-04T18:00:00Z".into(),
416            capture_step: 1,
417            warmup_steps: 0,
418            device: "cpu".into(),
419            measured_region_device_synchronized: false,
420            timing_mode: crate::trace::TimingMode::Host,
421            capture_contract: CaptureContract::default(),
422            comparison_identity: None,
423            tags: Default::default(),
424            candle_version: Some("0.8.0".into()),
425        }
426    }
427
428    fn sample_events() -> Vec<TraceEvent> {
429        vec![
430            TraceEvent::meta(sample_meta()),
431            TraceEvent::SpanStart(SpanStartEvent {
432                id: "span-root".into(),
433                parent_id: None,
434                name: "demo::train::loss".into(),
435                start_ns: 0,
436                kind: SpanKind::Function,
437                measured: true,
438                step: None,
439            }),
440            TraceEvent::SpanStart(SpanStartEvent {
441                id: "span-op".into(),
442                parent_id: Some("span-root".into()),
443                name: "matmul".into(),
444                start_ns: 10,
445                kind: SpanKind::Op,
446                measured: false,
447                step: None,
448            }),
449            TraceEvent::Op(OpEvent {
450                span_id: "span-op".into(),
451                op_name: "matmul".into(),
452                inputs: vec!["t0".into(), "t1".into()],
453                output: Some("t2".into()),
454                shape: vec![32, 32],
455                dtype: "f32".into(),
456                device: "cpu".into(),
457                duration_ns: 1200,
458                timestamp_ns: 10,
459                output_dense_bytes: None,
460                input_dense_bytes: 0,
461            }),
462            TraceEvent::Memory(super::super::events::MemoryEvent {
463                timestamp_ns: 1200,
464                storage_id: "storage-t2".into(),
465                tensor_id: "t2".into(),
466                span_id: "span-op".into(),
467                op_name: Some("matmul".into()),
468                device: "cpu".into(),
469                bytes: 32 * 32 * 4,
470                action: MemoryAction::Alloc,
471                shape: vec![32, 32],
472                dtype: "f32".into(),
473                category: MemoryCategory::Activation,
474            }),
475            TraceEvent::Edge(EdgeEvent::Call {
476                from_span: "span-root".into(),
477                to_span: "span-op".into(),
478                host_duration_ns: 1200,
479            }),
480            TraceEvent::Gradient(GradientEvent {
481                event_id: "grad-1".into(),
482                root: "vb".into(),
483                key: "encoder.weight".into(),
484                state: GradientState::Present,
485                norm: Some(0.42),
486            }),
487            TraceEvent::SpanEnd(SpanEndEvent {
488                id: "span-op".into(),
489                duration_ns: 1_200,
490            }),
491            TraceEvent::SpanEnd(SpanEndEvent {
492                id: "span-root".into(),
493                duration_ns: 2_500,
494            }),
495            TraceEvent::Terminal(TerminalEvent {
496                outcome: RunOutcome::Complete,
497                timestamp_ns: 2_500,
498                reason: None,
499            }),
500        ]
501    }
502
503    #[test]
504    fn from_events_builds_document_and_summary() {
505        let doc = TraceDocument::from_events(sample_events()).unwrap();
506        assert_eq!(doc.schema, SCHEMA);
507        assert_eq!(doc.run.entrypoint, "demo::train::loss");
508        assert_eq!(doc.spans.len(), 2);
509        assert!(doc.spans.iter().all(|span| span.closed));
510        assert_eq!(doc.ops.len(), 1);
511        assert_eq!(doc.ops[0].output_dense_bytes, Some(32 * 32 * 4));
512        assert_eq!(doc.memory.len(), 1);
513        assert_eq!(doc.edges.len(), 1);
514        assert_eq!(doc.gradients.len(), 1);
515        assert_eq!(doc.gradients[0].param_key(), "encoder.weight");
516
517        let summary = doc.build_summary();
518        assert_eq!(summary.op_count, 1);
519        assert_eq!(summary.total_ns, 2_500);
520        assert_eq!(summary.span_count, 2);
521        assert_eq!(summary.root_span_count, 1);
522        assert_eq!(summary.max_depth, 1);
523        assert_eq!(summary.alloc_count, 1);
524        assert_eq!(summary.logical_peak_bytes, Some(32 * 32 * 4));
525    }
526
527    #[test]
528    fn jsonl_roundtrip_via_temp_file() {
529        let dir = std::env::temp_dir().join(format!(
530            "candle-graph-trace7-{}-{}",
531            std::process::id(),
532            std::time::SystemTime::now()
533                .duration_since(std::time::UNIX_EPOCH)
534                .unwrap()
535                .as_nanos()
536        ));
537        std::fs::create_dir_all(&dir).unwrap();
538        let path = dir.join("trace.jsonl");
539
540        let events = sample_events();
541        write_jsonl(&path, &events).unwrap();
542        let parsed = parse_trace(&path).unwrap();
543        assert_eq!(parsed, TraceDocument::from_events(events.clone()).unwrap());
544
545        let _ = std::fs::remove_dir_all(dir);
546    }
547
548    fn sample_tensor_stats() -> TensorStatsEvent {
549        TensorStatsEvent {
550            span_id: "s1".into(),
551            label: "seam/out_y".into(),
552            shape: vec![2, 3],
553            dtype: "f32".into(),
554            elements: 6,
555            non_finite: 0,
556            rms: 1.5,
557            abs_max: 3.0,
558            mean: -0.25,
559        }
560    }
561
562    #[test]
563    fn tensor_stats_round_trip_in_current_schema() {
564        let stats = sample_tensor_stats();
565        let events = vec![
566            TraceEvent::meta(sample_meta()),
567            TraceEvent::TensorStats(stats.clone()),
568            TraceEvent::Terminal(TerminalEvent {
569                outcome: RunOutcome::Complete,
570                timestamp_ns: 0,
571                reason: None,
572            }),
573        ];
574        let document = TraceDocument::from_events(events).unwrap();
575        assert_eq!(document.schema, SCHEMA);
576        assert_eq!(document.tensor_stats, vec![stats]);
577        let rebuilt = TraceDocument::from_events(document.to_events()).unwrap();
578        assert_eq!(rebuilt, document);
579    }
580
581    #[test]
582    fn previous_schema_remains_readable_without_tensor_stats() {
583        let events = vec![
584            TraceEvent::Meta {
585                schema: PREVIOUS_SCHEMA.into(),
586                run: Box::new(sample_meta()),
587            },
588            TraceEvent::Terminal(TerminalEvent {
589                outcome: RunOutcome::Complete,
590                timestamp_ns: 0,
591                reason: None,
592            }),
593        ];
594        let document = TraceDocument::from_events(events).unwrap();
595        assert_eq!(document.schema, PREVIOUS_SCHEMA);
596        assert!(document.tensor_stats.is_empty());
597    }
598
599    #[test]
600    fn previous_schema_rejects_tensor_stats_events() {
601        let events = vec![
602            TraceEvent::Meta {
603                schema: PREVIOUS_SCHEMA.into(),
604                run: Box::new(sample_meta()),
605            },
606            TraceEvent::TensorStats(sample_tensor_stats()),
607            TraceEvent::Terminal(TerminalEvent {
608                outcome: RunOutcome::Complete,
609                timestamp_ns: 0,
610                reason: None,
611            }),
612        ];
613        let error = TraceDocument::from_events(events).unwrap_err();
614        assert!(error.to_string().contains("does not define tensor_stats"));
615    }
616
617    #[test]
618    fn gradient_rejects_removed_param_key_alias() {
619        let line = r#"{"kind":"gradient","event_id":"g1","root":"vb","param_key":"w","state":"present","norm":1.0}"#;
620        assert!(serde_json::from_str::<TraceEvent>(line).is_err());
621    }
622
623    #[test]
624    fn rejects_unknown_schema() {
625        let events = vec![
626            TraceEvent::Meta {
627                schema: "not-candle-graph".into(),
628                run: Box::new(sample_meta()),
629            },
630            TraceEvent::Terminal(TerminalEvent {
631                outcome: RunOutcome::Complete,
632                timestamp_ns: 0,
633                reason: None,
634            }),
635        ];
636        let err = TraceDocument::from_events(events).unwrap_err();
637        assert!(err.to_string().contains("unsupported trace schema"));
638    }
639
640    #[test]
641    fn rejects_span_end_without_start() {
642        let events = vec![
643            TraceEvent::meta(sample_meta()),
644            TraceEvent::SpanEnd(SpanEndEvent {
645                id: "missing".into(),
646                duration_ns: 0,
647            }),
648        ];
649        let err = TraceDocument::from_events(events).unwrap_err();
650        assert!(err.to_string().contains("unknown span"));
651    }
652
653    #[test]
654    fn rejects_records_after_terminal_and_invalid_outcomes() {
655        let after_terminal = vec![
656            TraceEvent::meta(sample_meta()),
657            TraceEvent::Terminal(TerminalEvent {
658                outcome: RunOutcome::Complete,
659                timestamp_ns: 0,
660                reason: None,
661            }),
662            TraceEvent::Gradient(GradientEvent {
663                event_id: "late".into(),
664                root: "vb".into(),
665                key: "w".into(),
666                state: GradientState::Present,
667                norm: None,
668            }),
669        ];
670        assert!(TraceDocument::from_events(after_terminal)
671            .unwrap_err()
672            .to_string()
673            .contains("must be the final"));
674
675        let failed_without_reason = vec![
676            TraceEvent::meta(sample_meta()),
677            TraceEvent::Terminal(TerminalEvent {
678                outcome: RunOutcome::Failed,
679                timestamp_ns: 0,
680                reason: Some("  ".into()),
681            }),
682        ];
683        assert!(TraceDocument::from_events(failed_without_reason)
684            .unwrap_err()
685            .to_string()
686            .contains("non-empty reason"));
687    }
688}