1use 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#[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 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 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 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
349pub 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
381pub 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}