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 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#[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 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 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 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
255pub 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
287pub 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}