1use std::cell::RefCell;
4use std::collections::BTreeMap;
5use std::fs::{File, OpenOptions};
6use std::io::{self, Write};
7use std::path::{Path, PathBuf};
8use std::time::{Instant, SystemTime, UNIX_EPOCH};
9
10use anyhow::{Context, Result};
11use serde::Serialize;
12
13use crate::phase::ExecutionPhase;
14use crate::trace::events::{
15 DeviceMemoryEvent, GradientEvent, MemoryEvent, OpEvent, SpanEndEvent, SpanStartEvent,
16 TensorEvent, TraceEvent,
17};
18use crate::trace::memory::category_for_step;
19use crate::trace::memory::{resolve_storage_bytes, MemoryAction};
20use crate::trace::schema::{GradientState, TimingMode, TraceRunMeta};
21
22use super::span::{MemoryRecord, OpRecord, SpanGuard, SpanId, SpanKind, TensorRecord};
23
24pub struct TraceSession {
26 path: PathBuf,
27 inner: RefCell<SessionInner>,
28}
29
30#[derive(Debug, Clone, PartialEq, Eq)]
32pub struct ProfileRun {
33 pub entrypoint: String,
34 pub correlation_id: String,
35 pub phase: ExecutionPhase,
36 pub capture_step: u64,
38 pub warmup_steps: u64,
39 pub device: String,
40 pub measured_region_device_synchronized: bool,
41 pub timing_mode: TimingMode,
42 pub tags: BTreeMap<String, String>,
43}
44
45impl ProfileRun {
46 pub fn training(
47 entrypoint: impl Into<String>,
48 capture_step: u64,
49 device: impl Into<String>,
50 ) -> Self {
51 let entrypoint = entrypoint.into();
52 Self {
53 correlation_id: format!("{entrypoint}/update-{capture_step}"),
54 entrypoint,
55 phase: ExecutionPhase::Train,
56 capture_step,
57 warmup_steps: capture_step.saturating_sub(1),
58 device: device.into(),
59 measured_region_device_synchronized: false,
60 timing_mode: TimingMode::Host,
61 tags: BTreeMap::new(),
62 }
63 }
64
65 pub fn tag(mut self, key: impl Into<String>, value: impl Into<String>) -> Self {
66 self.tags.insert(key.into(), value.into());
67 self
68 }
69
70 pub fn correlation_id(mut self, value: impl Into<String>) -> Self {
71 self.correlation_id = value.into();
72 self
73 }
74
75 pub fn device_synchronized(mut self) -> Self {
76 self.timing_mode = TimingMode::DeviceSynchronized;
77 self.measured_region_device_synchronized = true;
78 self
79 }
80
81 pub fn measured_region_device_synchronized(mut self) -> Self {
84 self.measured_region_device_synchronized = true;
85 self
86 }
87}
88
89struct SessionInner {
90 writer: io::BufWriter<File>,
91 span_stack: Vec<u64>,
92 span_steps: Vec<Option<crate::phase::ExecutionStep>>,
93 next_span_id: u64,
94 next_event_id: u64,
95 id_buf: String,
96 probe_started: Instant,
97 sticky_error: Option<String>,
98}
99
100impl TraceSession {
101 pub fn open(path: impl AsRef<Path>, run: ProfileRun) -> Result<Self> {
103 anyhow::ensure!(
104 run.capture_step > 0,
105 "capture_step must be one-based and greater than zero"
106 );
107 let path = path.as_ref().to_path_buf();
108 if let Some(parent) = path.parent() {
109 std::fs::create_dir_all(parent)
110 .with_context(|| format!("create trace dir {}", parent.display()))?;
111 }
112 let file = OpenOptions::new()
113 .create(true)
114 .write(true)
115 .truncate(true)
116 .open(&path)
117 .with_context(|| format!("open trace {}", path.display()))?;
118 let mut writer = io::BufWriter::new(file);
119 let run_id = new_run_id();
120 let meta = TraceRunMeta {
121 run_id: run_id.clone(),
122 correlation_id: run.correlation_id,
123 entrypoint: run.entrypoint.clone(),
124 phase: run.phase,
125 timestamp: utc_iso8601_now(),
126 capture_step: run.capture_step,
127 warmup_steps: run.warmup_steps,
128 device: run.device,
129 measured_region_device_synchronized: run.measured_region_device_synchronized,
130 timing_mode: run.timing_mode,
131 tags: run.tags,
132 candle_version: None,
133 };
134 write_event(&mut writer, &TraceEvent::meta(meta))?;
135 write_event(
136 &mut writer,
137 &TraceEvent::SpanStart(SpanStartEvent {
138 id: span_id_string(1),
139 parent_id: None,
140 name: run.entrypoint,
141 start_ns: 0,
142 kind: SpanKind::Function,
143 measured: false,
144 step: None,
145 }),
146 )?;
147 Ok(Self {
148 path,
149 inner: RefCell::new(SessionInner {
150 writer,
151 span_stack: vec![1],
152 span_steps: vec![None],
153 next_span_id: 1,
154 next_event_id: 0,
155 id_buf: String::with_capacity(24),
156 probe_started: Instant::now(),
157 sticky_error: None,
158 }),
159 })
160 }
161
162 pub fn begin_span(&self, name: impl Into<String>, kind: SpanKind) -> SpanGuard<'_> {
164 self.begin_span_inner(name, kind, None, false)
165 }
166
167 pub fn begin_measurement(&self, name: impl Into<String>) -> SpanGuard<'_> {
169 self.begin_span_inner(name, SpanKind::Function, None, true)
170 }
171
172 pub fn begin_step_span(
174 &self,
175 name: impl Into<String>,
176 step: crate::phase::ExecutionStep,
177 kind: SpanKind,
178 ) -> SpanGuard<'_> {
179 self.begin_span_inner(name, kind, Some(step), false)
180 }
181
182 fn begin_span_inner(
183 &self,
184 name: impl Into<String>,
185 kind: SpanKind,
186 step: Option<crate::phase::ExecutionStep>,
187 measured: bool,
188 ) -> SpanGuard<'_> {
189 let started = Instant::now();
190 let start_ns = self.elapsed_ns();
191 let mut inner = self.inner.borrow_mut();
192 inner.next_span_id += 1;
193 let span_id = inner.next_span_id;
194 let parent_id = inner.span_stack.last().copied().map(span_id_string);
195
196 format_span_id(&mut inner.id_buf, span_id);
197 let id_str = inner.id_buf.clone();
198
199 if let Err(error) = write_event(
200 &mut inner.writer,
201 &TraceEvent::SpanStart(SpanStartEvent {
202 id: id_str,
203 parent_id,
204 name: name.into(),
205 start_ns,
206 kind,
207 measured,
208 step,
209 }),
210 ) {
211 inner.sticky_error.get_or_insert_with(|| error.to_string());
212 }
213
214 inner.span_stack.push(span_id);
215 inner.span_steps.push(step);
216
217 SpanGuard {
218 session: self,
219 id: SpanId(span_id),
220 started,
221 }
222 }
223
224 fn current_step(&self) -> Option<crate::phase::ExecutionStep> {
225 self.inner
226 .borrow()
227 .span_steps
228 .iter()
229 .rev()
230 .find_map(|step| *step)
231 }
232
233 pub(crate) fn end_span(&self, id: SpanId, duration_ns: u64) -> Result<()> {
234 let mut inner = self.inner.borrow_mut();
235 let expected = inner
236 .span_stack
237 .last()
238 .copied()
239 .with_context(|| format!("span stack underflow closing span {}", id.0))?;
240 anyhow::ensure!(
241 expected == id.0,
242 "span_end id `{}` does not match open span `{}`",
243 id.0,
244 expected
245 );
246 inner.span_stack.pop();
247 inner.span_steps.pop();
248
249 format_span_id(&mut inner.id_buf, id.0);
250 let span_id = inner.id_buf.clone();
251 if let Err(error) = write_event(
252 &mut inner.writer,
253 &TraceEvent::SpanEnd(SpanEndEvent {
254 id: span_id,
255 duration_ns,
256 }),
257 ) {
258 inner.sticky_error.get_or_insert_with(|| error.to_string());
259 return Err(error);
260 }
261 Ok(())
262 }
263
264 pub fn elapsed_ns(&self) -> u64 {
265 self.inner
266 .borrow()
267 .probe_started
268 .elapsed()
269 .as_nanos()
270 .min(u64::MAX as u128) as u64
271 }
272
273 pub fn record_op(&self, span_id: SpanId, op: OpRecord<'_>) -> Result<()> {
275 let storage_bytes = resolve_storage_bytes(op.storage_bytes, op.shape, op.dtype);
276 let timestamp_ns = if op.timestamp_ns > 0 {
277 op.timestamp_ns
278 } else {
279 self.elapsed_ns()
280 };
281 let category = op
282 .category
283 .unwrap_or_else(|| category_for_step(self.current_step(), false));
284 {
285 let mut inner = self.inner.borrow_mut();
286 write_event(
287 &mut inner.writer,
288 &TraceEvent::Op(OpEvent {
289 span_id: span_id_string(span_id.0),
290 op_name: op.op_name.into(),
291 inputs: op.inputs.to_vec(),
292 output: op.output.map(str::to_string),
293 shape: op.shape.to_vec(),
294 dtype: op.dtype.into(),
295 device: op.device.into(),
296 duration_ns: op.duration_ns,
297 timestamp_ns,
298 storage_bytes: Some(storage_bytes),
299 input_storage_bytes: op.input_storage_bytes,
300 }),
301 )?;
302 }
303
304 let _ = category;
305 Ok(())
306 }
307
308 pub fn record_tensor(&self, span_id: SpanId, tensor: TensorRecord<'_>) -> Result<()> {
310 let storage_bytes = resolve_storage_bytes(tensor.storage_bytes, tensor.shape, tensor.dtype);
311 let mut inner = self.inner.borrow_mut();
312 write_event(
313 &mut inner.writer,
314 &TraceEvent::Tensor(TensorEvent {
315 span_id: span_id_string(span_id.0),
316 tensor_id: tensor.tensor_id.into(),
317 shape: tensor.shape.to_vec(),
318 dtype: tensor.dtype.into(),
319 device: tensor.device.into(),
320 requires_grad: tensor.requires_grad,
321 storage_bytes: Some(storage_bytes),
322 category: tensor.category,
323 }),
324 )?;
325 Ok(())
326 }
327
328 pub fn record_memory_alloc(&self, span_id: SpanId, mem: MemoryRecord<'_>) -> Result<()> {
330 let timestamp_ns = mem.timestamp_ns.unwrap_or_else(|| self.elapsed_ns());
331 let mut inner = self.inner.borrow_mut();
332 write_event(
333 &mut inner.writer,
334 &TraceEvent::Memory(MemoryEvent {
335 timestamp_ns,
336 tensor_id: mem.tensor_id.into(),
337 span_id: span_id_string(span_id.0),
338 op_name: mem.op_name.map(str::to_string),
339 device: mem.device.into(),
340 bytes: mem.bytes,
341 action: MemoryAction::Alloc,
342 shape: mem.shape.to_vec(),
343 dtype: mem.dtype.into(),
344 category: mem.category,
345 }),
346 )
347 }
348
349 pub fn record_memory_free(&self, span_id: SpanId, mem: MemoryRecord<'_>) -> Result<()> {
351 let timestamp_ns = mem.timestamp_ns.unwrap_or_else(|| self.elapsed_ns());
352 let mut inner = self.inner.borrow_mut();
353 write_event(
354 &mut inner.writer,
355 &TraceEvent::Memory(MemoryEvent {
356 timestamp_ns,
357 tensor_id: mem.tensor_id.into(),
358 span_id: span_id_string(span_id.0),
359 op_name: mem.op_name.map(str::to_string),
360 device: mem.device.into(),
361 bytes: mem.bytes,
362 action: MemoryAction::Free,
363 shape: mem.shape.to_vec(),
364 dtype: mem.dtype.into(),
365 category: mem.category,
366 }),
367 )
368 }
369
370 pub fn record_device_memory(
372 &self,
373 device: impl Into<String>,
374 used_bytes: u64,
375 free_bytes: u64,
376 timestamp_ns: Option<u64>,
377 ) -> Result<()> {
378 let timestamp_ns = timestamp_ns.unwrap_or_else(|| self.elapsed_ns());
379 let mut inner = self.inner.borrow_mut();
380 write_event(
381 &mut inner.writer,
382 &TraceEvent::DeviceMemory(DeviceMemoryEvent {
383 timestamp_ns,
384 device: device.into(),
385 used_bytes,
386 free_bytes,
387 reserved_bytes: None,
388 }),
389 )
390 }
391
392 pub fn record_gradient(
394 &self,
395 root: impl Into<String>,
396 key: impl Into<String>,
397 state: GradientState,
398 norm: Option<f64>,
399 ) -> Result<()> {
400 let mut inner = self.inner.borrow_mut();
401 inner.next_event_id += 1;
402 let event_id = format!("gradient-{}", inner.next_event_id);
403 write_event(
404 &mut inner.writer,
405 &TraceEvent::Gradient(GradientEvent {
406 event_id,
407 root: root.into(),
408 key: key.into(),
409 state,
410 norm,
411 }),
412 )
413 }
414
415 pub fn flush(&self) -> Result<()> {
416 let mut inner = self.inner.borrow_mut();
417 if let Some(error) = &inner.sticky_error {
418 anyhow::bail!("trace session previously failed: {error}");
419 }
420 inner.writer.flush().context("flushing trace JSONL")
421 }
422
423 pub fn finish(self) -> Result<PathBuf> {
425 let duration_ns = self.elapsed_ns();
426 {
427 let mut inner = self.inner.borrow_mut();
428 if let Some(error) = &inner.sticky_error {
429 anyhow::bail!("trace session previously failed: {error}");
430 }
431 anyhow::ensure!(
432 inner.span_stack.as_slice() == [1],
433 "cannot finish trace with {} nested spans still open",
434 inner.span_stack.len().saturating_sub(1)
435 );
436 inner.span_stack.pop();
437 inner.span_steps.pop();
438 write_event(
439 &mut inner.writer,
440 &TraceEvent::SpanEnd(SpanEndEvent {
441 id: span_id_string(1),
442 duration_ns,
443 }),
444 )?;
445 inner.writer.flush().context("flushing trace JSONL")?;
446 }
447 Ok(self.path)
448 }
449}
450
451fn write_event<W: Write, T: Serialize>(writer: &mut W, event: &T) -> Result<()> {
452 let mut line = serde_json::to_vec(event).context("serializing trace JSONL event")?;
453 line.push(b'\n');
454 writer.write_all(&line).context("writing trace JSONL event")
455}
456
457fn format_span_id(buf: &mut String, id: u64) {
458 buf.clear();
459 use std::fmt::Write as _;
460 let _ = write!(buf, "s{id}");
461}
462
463fn span_id_string(id: u64) -> String {
464 format!("s{id}")
465}
466
467fn new_run_id() -> String {
468 let pid = std::process::id();
469 let nanos = SystemTime::now()
470 .duration_since(UNIX_EPOCH)
471 .map(|d| d.as_nanos())
472 .unwrap_or(0);
473 format!("run-{pid}-{nanos}")
474}
475
476fn utc_iso8601_now() -> String {
477 let now = SystemTime::now()
478 .duration_since(UNIX_EPOCH)
479 .expect("system clock before UNIX epoch");
480 let secs = now.as_secs();
481 let millis = now.subsec_millis();
482 let (year, month, day, hour, minute, second) = unix_secs_to_utc(secs);
483 format!("{year:04}-{month:02}-{day:02}T{hour:02}:{minute:02}:{second:02}.{millis:03}Z")
484}
485
486fn unix_secs_to_utc(secs: u64) -> (u64, u64, u64, u64, u64, u64) {
488 const SECS_PER_DAY: u64 = 86_400;
489 let days = secs / SECS_PER_DAY;
490 let rem = secs % SECS_PER_DAY;
491 let hour = rem / 3600;
492 let minute = (rem % 3600) / 60;
493 let second = rem % 60;
494
495 let z = days + 719_468;
496 let era = z / 146_097;
497 let doe = z - era * 146_097;
498 let yoe = (doe - doe / 1460 + doe / 36524 - doe / 146096) / 365;
499 let y = yoe + era * 400;
500 let doy = doe - (365 * yoe + yoe / 4 - yoe / 100);
501 let mp = (5 * doy + 2) / 153;
502 let day = doy - (153 * mp + 2) / 5 + 1;
503 let month = if mp < 10 { mp + 3 } else { mp - 9 };
504 let year = if month <= 2 { y + 1 } else { y };
505 (year, month, day, hour, minute, second)
506}
507
508#[cfg(test)]
509mod tests {
510 use super::*;
511 use crate::trace::document::parse_trace;
512 use crate::trace::events::SpanStartEvent;
513 use serde_json::Value;
514 use std::io::{BufRead, BufReader};
515
516 fn temp_trace(name: &str) -> PathBuf {
517 std::env::temp_dir().join(format!(
518 "candle-graph-trace-{}-{}-{name}",
519 std::process::id(),
520 std::time::SystemTime::now()
521 .duration_since(UNIX_EPOCH)
522 .unwrap()
523 .as_nanos()
524 ))
525 }
526
527 fn span_end_durations(path: &Path) -> Vec<(String, u64)> {
528 let file = std::fs::File::open(path).unwrap();
529 let reader = BufReader::new(file);
530 reader
531 .lines()
532 .map(|line| line.unwrap())
533 .filter_map(|line| {
534 let value: Value = serde_json::from_str(&line).unwrap();
535 if value.get("kind")?.as_str()? != "span_end" {
536 return None;
537 }
538 Some((
539 value["id"].as_str().unwrap().to_string(),
540 value["duration_ns"].as_u64().unwrap(),
541 ))
542 })
543 .collect()
544 }
545
546 fn read_events(path: &Path) -> Vec<TraceEvent> {
547 let file = std::fs::File::open(path).unwrap();
548 let reader = BufReader::new(file);
549 reader
550 .lines()
551 .map(|line| {
552 let line = line.unwrap();
553 serde_json::from_str(&line).unwrap_or_else(|err| {
554 panic!("invalid JSONL line `{line}`: {err}");
555 })
556 })
557 .collect()
558 }
559
560 #[test]
561 fn nested_spans_emit_parent_hierarchy_and_durations() {
562 let path = temp_trace("nested");
563 let session =
564 TraceSession::open(&path, ProfileRun::training("model::forward", 1, "cpu")).unwrap();
565
566 let inner_id = {
567 let _outer = session.begin_measurement("Model::forward");
568 std::thread::sleep(std::time::Duration::from_micros(50));
569 let inner = session.begin_span("matmul", SpanKind::Op);
570 std::thread::sleep(std::time::Duration::from_micros(50));
571 inner.id
572 };
573
574 session
575 .record_op(
576 inner_id,
577 OpRecord {
578 op_name: "matmul",
579 inputs: &["a".into(), "b".into()],
580 output: Some("c"),
581 shape: &[8, 8],
582 dtype: "f32",
583 device: "cpu",
584 duration_ns: 1200,
585 timestamp_ns: 0,
586 storage_bytes: None,
587 input_storage_bytes: 0,
588 category: None,
589 },
590 )
591 .unwrap();
592
593 session.finish().unwrap();
594
595 let events = read_events(&path);
596 assert!(matches!(events.first(), Some(TraceEvent::Meta { .. })));
597
598 let starts: Vec<&SpanStartEvent> = events
599 .iter()
600 .filter_map(|event| match event {
601 TraceEvent::SpanStart(start) => Some(start),
602 _ => None,
603 })
604 .collect();
605 assert_eq!(starts.len(), 3);
606 assert_eq!(starts[0].name, "model::forward");
607 assert!(starts[0].parent_id.is_none());
608 assert_eq!(starts[1].name, "Model::forward");
609 assert_eq!(starts[1].parent_id.as_deref(), Some("s1"));
610 assert_eq!(starts[2].name, "matmul");
611 assert_eq!(starts[2].parent_id.as_deref(), Some("s2"));
612
613 let ends = span_end_durations(&path);
614 assert_eq!(ends.len(), 3);
615 assert!(ends[0].1 > 0);
616
617 let doc = parse_trace(&path).unwrap();
618 assert_eq!(doc.run.entrypoint, "model::forward");
619 assert_eq!(doc.ops.len(), 1);
620 assert_eq!(doc.ops[0].storage_bytes, Some(8 * 8 * 4));
621 assert!(
622 doc.memory.is_empty(),
623 "op metadata must not fabricate tensor lifetime"
624 );
625 }
626
627 #[test]
628 fn measured_region_sync_does_not_overstate_nested_span_timing() {
629 let run = ProfileRun::training("train::update", 2, "cuda:0")
630 .measured_region_device_synchronized();
631
632 assert!(run.measured_region_device_synchronized);
633 assert_eq!(run.timing_mode, TimingMode::Host);
634 }
635
636 #[test]
637 fn record_gradient_round_trips_through_trace_parser() {
638 let path = temp_trace("gradient");
639 let session =
640 TraceSession::open(&path, ProfileRun::training("train::loss", 1, "cpu")).unwrap();
641 session
642 .record_gradient("vb", "encoder.weight", GradientState::Present, Some(0.42))
643 .unwrap();
644 session.finish().unwrap();
645
646 let doc = parse_trace(&path).unwrap();
647 assert_eq!(doc.gradients.len(), 1);
648 assert_eq!(doc.gradients[0].key, "encoder.weight");
649 }
650}