Skip to main content

candle_graph/trace/
document.rs

1//! Aggregate trace document and JSONL I/O for `candle-graph/trace/6`.
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    DeviceMemoryEvent, EdgeEvent, GradientEvent, MemoryEvent, OpEvent, SpanEndEvent,
13    SpanStartEvent, TensorEvent, TraceEvent,
14};
15use super::memory::{resolve_storage_bytes, MemoryAction};
16use super::schema::{SpanRecord, TraceRunMeta, TraceSummary, 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 memory: Vec<MemoryEvent>,
31    #[serde(default)]
32    pub device_memory: Vec<DeviceMemoryEvent>,
33    #[serde(default)]
34    pub gradients: Vec<GradientEvent>,
35    #[serde(default)]
36    pub edges: Vec<EdgeEvent>,
37}
38
39impl TraceDocument {
40    /// Build a document from an ordered event stream (meta must be first).
41    pub fn from_events(events: impl IntoIterator<Item = TraceEvent>) -> Result<Self> {
42        let mut schema: Option<String> = None;
43        let mut run: Option<TraceRunMeta> = None;
44        let mut span_starts: BTreeMap<String, SpanStartEvent> = BTreeMap::new();
45        let mut span_durations: HashMap<String, u64> = HashMap::new();
46        let mut span_closed: HashSet<String> = HashSet::new();
47        let mut ops = Vec::new();
48        let mut tensors = Vec::new();
49        let mut memory = Vec::new();
50        let mut device_memory = Vec::new();
51        let mut gradients = Vec::new();
52        let mut edges = Vec::new();
53
54        for (index, event) in events.into_iter().enumerate() {
55            match event {
56                TraceEvent::Meta {
57                    schema: s,
58                    run: meta,
59                } => {
60                    if index != 0 {
61                        bail!("meta event must be the first non-empty trace record, found at index {index}");
62                    }
63                    if schema.is_some() || run.is_some() {
64                        bail!(
65                            "duplicate meta event at index {index}; only one meta record is allowed"
66                        );
67                    }
68                    schema = Some(s);
69                    run = Some(meta);
70                }
71                TraceEvent::SpanStart(start) => {
72                    if span_starts.contains_key(&start.id) {
73                        bail!("duplicate span_start id `{}` at index {index}", start.id);
74                    }
75                    span_starts.insert(start.id.clone(), start);
76                }
77                TraceEvent::SpanEnd(SpanEndEvent { id, duration_ns }) => {
78                    if !span_starts.contains_key(&id) {
79                        bail!("span_end for unknown span `{id}` at index {index}");
80                    }
81                    if span_closed.contains(&id) {
82                        bail!("duplicate span_end for `{id}` at index {index}");
83                    }
84                    span_closed.insert(id.clone());
85                    span_durations.insert(id, duration_ns);
86                }
87                TraceEvent::Op(mut op) => {
88                    op.storage_bytes = Some(resolve_storage_bytes(
89                        op.storage_bytes,
90                        &op.shape,
91                        &op.dtype,
92                    ));
93                    ops.push(op);
94                }
95                TraceEvent::Tensor(mut tensor) => {
96                    tensor.storage_bytes = Some(resolve_storage_bytes(
97                        tensor.storage_bytes,
98                        &tensor.shape,
99                        &tensor.dtype,
100                    ));
101                    tensors.push(tensor);
102                }
103                TraceEvent::Memory(mem) => memory.push(mem),
104                TraceEvent::DeviceMemory(snapshot) => device_memory.push(snapshot),
105                TraceEvent::Gradient(gradient) => gradients.push(gradient),
106                TraceEvent::Edge(edge) => edges.push(edge),
107            }
108        }
109
110        let schema = schema.unwrap_or_else(|| SCHEMA.to_string());
111        let run = run.context("trace stream is missing a meta event with run metadata")?;
112
113        if schema != SCHEMA {
114            bail!("unsupported trace schema {schema:?}; expected {SCHEMA:?}");
115        }
116
117        let mut spans: Vec<SpanRecord> = span_starts
118            .into_iter()
119            .map(|(id, start)| SpanRecord {
120                id: id.clone(),
121                parent_id: start.parent_id,
122                name: start.name,
123                kind: start.kind,
124                measured: start.measured,
125                start_ns: start.start_ns,
126                closed: span_closed.contains(&id),
127                duration_ns: span_durations.get(&id).copied().unwrap_or(0),
128                step: start.step,
129            })
130            .collect();
131        spans.sort_by(|a, b| a.id.cmp(&b.id));
132
133        Ok(Self {
134            schema,
135            run,
136            spans,
137            ops,
138            tensors,
139            memory,
140            device_memory,
141            gradients,
142            edges,
143        })
144    }
145
146    /// Profiler summary: op count, total wall time, span tree shape, memory totals.
147    pub fn build_summary(&self) -> TraceSummary {
148        let op_count = self.ops.len();
149        let total_ns = self
150            .spans
151            .iter()
152            .filter(|span| span.measured)
153            .map(|span| span.duration_ns)
154            .sum();
155        let span_count = self.spans.len();
156        let root_span_count = self
157            .spans
158            .iter()
159            .filter(|span| span.parent_id.is_none())
160            .count();
161
162        let parent_by_id: HashMap<&str, Option<&str>> = self
163            .spans
164            .iter()
165            .map(|span| (span.id.as_str(), span.parent_id.as_deref()))
166            .collect();
167
168        let mut max_depth = 0usize;
169        for span in &self.spans {
170            let mut depth = 0usize;
171            let mut current_parent = span.parent_id.as_deref();
172            let mut seen = HashSet::new();
173            while let Some(parent_id) = current_parent {
174                if !seen.insert(parent_id) {
175                    break;
176                }
177                depth += 1;
178                current_parent = parent_by_id.get(parent_id).copied().flatten();
179            }
180            max_depth = max_depth.max(depth);
181        }
182
183        let alloc_count = self
184            .memory
185            .iter()
186            .filter(|event| event.action == MemoryAction::Alloc)
187            .count();
188        let free_count = self
189            .memory
190            .iter()
191            .filter(|event| event.action == MemoryAction::Free)
192            .count();
193
194        let peak_bytes = super::memory::analyze_memory(self).summary.peak_bytes;
195
196        TraceSummary {
197            op_count,
198            total_ns,
199            span_count,
200            root_span_count,
201            max_depth,
202            alloc_count,
203            free_count,
204            peak_bytes,
205        }
206    }
207
208    /// Flatten the document back into JSONL events (meta first, then body in stable order).
209    pub fn to_events(&self) -> Vec<TraceEvent> {
210        let mut events = vec![TraceEvent::Meta {
211            schema: self.schema.clone(),
212            run: self.run.clone(),
213        }];
214
215        let mut span_ids: Vec<_> = self.spans.iter().map(|span| span.id.as_str()).collect();
216        span_ids.sort_unstable();
217        for id in span_ids {
218            let span = self
219                .spans
220                .iter()
221                .find(|span| span.id == id)
222                .expect("sorted id must exist");
223            events.push(TraceEvent::SpanStart(SpanStartEvent {
224                id: span.id.clone(),
225                parent_id: span.parent_id.clone(),
226                name: span.name.clone(),
227                kind: span.kind,
228                measured: span.measured,
229                start_ns: span.start_ns,
230                step: span.step,
231            }));
232            if span.closed {
233                events.push(TraceEvent::SpanEnd(SpanEndEvent {
234                    id: span.id.clone(),
235                    duration_ns: span.duration_ns,
236                }));
237            }
238        }
239
240        events.extend(self.ops.iter().cloned().map(TraceEvent::Op));
241        events.extend(self.tensors.iter().cloned().map(TraceEvent::Tensor));
242        events.extend(self.memory.iter().cloned().map(TraceEvent::Memory));
243        events.extend(
244            self.device_memory
245                .iter()
246                .cloned()
247                .map(TraceEvent::DeviceMemory),
248        );
249        events.extend(self.gradients.iter().cloned().map(TraceEvent::Gradient));
250        events.extend(self.edges.iter().cloned().map(TraceEvent::Edge));
251        events
252    }
253}
254
255/// Parse a JSONL trace file into a [`TraceDocument`].
256pub fn parse_trace(path: impl AsRef<Path>) -> Result<TraceDocument> {
257    let path = path.as_ref();
258    let file = File::open(path).with_context(|| format!("open trace file {}", path.display()))?;
259    let reader = BufReader::new(file);
260    let mut events = Vec::new();
261
262    for (line_no, line) in reader.lines().enumerate() {
263        let line = line.with_context(|| {
264            format!(
265                "read trace JSONL line {} from {}",
266                line_no + 1,
267                path.display()
268            )
269        })?;
270        let trimmed = line.trim();
271        if trimmed.is_empty() {
272            continue;
273        }
274        let event: TraceEvent = serde_json::from_str(trimmed).with_context(|| {
275            format!(
276                "parse trace JSONL line {} from {}",
277                line_no + 1,
278                path.display()
279            )
280        })?;
281        events.push(event);
282    }
283
284    TraceDocument::from_events(events)
285}
286
287/// Write JSONL events to `path` (creates parent directories when needed).
288pub fn write_jsonl(path: impl AsRef<Path>, events: &[TraceEvent]) -> Result<()> {
289    let path = path.as_ref();
290    if let Some(parent) = path.parent() {
291        std::fs::create_dir_all(parent)
292            .with_context(|| format!("create trace dir {}", parent.display()))?;
293    }
294    let mut file =
295        File::create(path).with_context(|| format!("create trace file {}", path.display()))?;
296    for event in events {
297        let mut line = serde_json::to_vec(event).context("serialize trace JSONL event")?;
298        line.push(b'\n');
299        file.write_all(&line)
300            .with_context(|| format!("write trace JSONL to {}", path.display()))?;
301    }
302    file.flush()
303        .with_context(|| format!("flush trace JSONL to {}", path.display()))?;
304    Ok(())
305}
306
307#[cfg(test)]
308mod tests {
309    use super::*;
310    use crate::trace::events::TraceEvent;
311    use crate::trace::memory::MemoryCategory;
312    use crate::trace::schema::{GradientState, SpanKind};
313
314    fn sample_meta() -> TraceRunMeta {
315        TraceRunMeta {
316            run_id: "run-1".into(),
317            correlation_id: "demo/update-1".into(),
318            entrypoint: "demo::train::loss".into(),
319            phase: crate::phase::ExecutionPhase::Train,
320            timestamp: "2026-08-04T18:00:00Z".into(),
321            capture_step: 1,
322            warmup_steps: 0,
323            device: "cpu".into(),
324            timing_mode: crate::trace::TimingMode::Host,
325            tags: Default::default(),
326            candle_version: Some("0.8.0".into()),
327        }
328    }
329
330    fn sample_events() -> Vec<TraceEvent> {
331        vec![
332            TraceEvent::meta(sample_meta()),
333            TraceEvent::SpanStart(SpanStartEvent {
334                id: "span-root".into(),
335                parent_id: None,
336                name: "demo::train::loss".into(),
337                start_ns: 0,
338                kind: SpanKind::Function,
339                measured: true,
340                step: None,
341            }),
342            TraceEvent::SpanStart(SpanStartEvent {
343                id: "span-op".into(),
344                parent_id: Some("span-root".into()),
345                name: "matmul".into(),
346                start_ns: 10,
347                kind: SpanKind::Op,
348                measured: false,
349                step: None,
350            }),
351            TraceEvent::Op(OpEvent {
352                span_id: "span-op".into(),
353                op_name: "matmul".into(),
354                inputs: vec!["t0".into(), "t1".into()],
355                output: Some("t2".into()),
356                shape: vec![32, 32],
357                dtype: "f32".into(),
358                device: "cpu".into(),
359                duration_ns: 1200,
360                timestamp_ns: 1200,
361                storage_bytes: None,
362                input_storage_bytes: 0,
363            }),
364            TraceEvent::Memory(super::super::events::MemoryEvent {
365                timestamp_ns: 1200,
366                tensor_id: "t2".into(),
367                span_id: "span-op".into(),
368                op_name: Some("matmul".into()),
369                device: "cpu".into(),
370                bytes: 32 * 32 * 4,
371                action: MemoryAction::Alloc,
372                shape: vec![32, 32],
373                dtype: "f32".into(),
374                category: MemoryCategory::Activation,
375            }),
376            TraceEvent::Edge(EdgeEvent {
377                from_span: "span-root".into(),
378                to_span: "span-op".into(),
379                duration_ns: 1200,
380            }),
381            TraceEvent::Gradient(GradientEvent {
382                event_id: "grad-1".into(),
383                root: "vb".into(),
384                key: "encoder.weight".into(),
385                state: GradientState::Present,
386                norm: Some(0.42),
387            }),
388            TraceEvent::SpanEnd(SpanEndEvent {
389                id: "span-op".into(),
390                duration_ns: 1_200,
391            }),
392            TraceEvent::SpanEnd(SpanEndEvent {
393                id: "span-root".into(),
394                duration_ns: 2_000,
395            }),
396        ]
397    }
398
399    #[test]
400    fn from_events_builds_document_and_summary() {
401        let doc = TraceDocument::from_events(sample_events()).unwrap();
402        assert_eq!(doc.schema, SCHEMA);
403        assert_eq!(doc.run.entrypoint, "demo::train::loss");
404        assert_eq!(doc.spans.len(), 2);
405        assert!(doc.spans.iter().all(|span| span.closed));
406        assert_eq!(doc.ops.len(), 1);
407        assert_eq!(doc.ops[0].storage_bytes, Some(32 * 32 * 4));
408        assert_eq!(doc.memory.len(), 1);
409        assert_eq!(doc.edges.len(), 1);
410        assert_eq!(doc.gradients.len(), 1);
411        assert_eq!(doc.gradients[0].param_key(), "encoder.weight");
412
413        let summary = doc.build_summary();
414        assert_eq!(summary.op_count, 1);
415        assert_eq!(summary.total_ns, 2_000);
416        assert_eq!(summary.span_count, 2);
417        assert_eq!(summary.root_span_count, 1);
418        assert_eq!(summary.max_depth, 1);
419        assert_eq!(summary.alloc_count, 1);
420        assert_eq!(summary.peak_bytes, 32 * 32 * 4);
421    }
422
423    #[test]
424    fn jsonl_roundtrip_via_temp_file() {
425        let dir = std::env::temp_dir().join(format!(
426            "candle-graph-trace6-{}-{}",
427            std::process::id(),
428            std::time::SystemTime::now()
429                .duration_since(std::time::UNIX_EPOCH)
430                .unwrap()
431                .as_nanos()
432        ));
433        std::fs::create_dir_all(&dir).unwrap();
434        let path = dir.join("trace.jsonl");
435
436        let events = sample_events();
437        write_jsonl(&path, &events).unwrap();
438        let parsed = parse_trace(&path).unwrap();
439        assert_eq!(parsed, TraceDocument::from_events(events.clone()).unwrap());
440
441        let _ = std::fs::remove_dir_all(dir);
442    }
443
444    #[test]
445    fn gradient_rejects_removed_param_key_alias() {
446        let line = r#"{"kind":"gradient","event_id":"g1","root":"vb","param_key":"w","state":"present","norm":1.0}"#;
447        assert!(serde_json::from_str::<TraceEvent>(line).is_err());
448    }
449
450    #[test]
451    fn rejects_unknown_schema() {
452        let events = vec![TraceEvent::Meta {
453            schema: "not-candle-graph".into(),
454            run: sample_meta(),
455        }];
456        let err = TraceDocument::from_events(events).unwrap_err();
457        assert!(err.to_string().contains("unsupported trace schema"));
458    }
459
460    #[test]
461    fn rejects_span_end_without_start() {
462        let events = vec![
463            TraceEvent::meta(sample_meta()),
464            TraceEvent::SpanEnd(SpanEndEvent {
465                id: "missing".into(),
466                duration_ns: 0,
467            }),
468        ];
469        let err = TraceDocument::from_events(events).unwrap_err();
470        assert!(err.to_string().contains("unknown span"));
471    }
472}