Skip to main content

candle_graph/instrument/
session.rs

1//! Representative-run profiler session — emits `candle-graph/trace/6` JSONL.
2
3use 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
24/// Streaming trace session writing TensorFlow-Profiler-style span JSONL.
25pub struct TraceSession {
26    path: PathBuf,
27    inner: RefCell<SessionInner>,
28}
29
30/// Required provenance for one representative profile run.
31#[derive(Debug, Clone, PartialEq, Eq)]
32pub struct ProfileRun {
33    pub entrypoint: String,
34    pub correlation_id: String,
35    pub phase: ExecutionPhase,
36    /// One-based selected update or inference invocation.
37    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    /// Mark only the caller-controlled measured region as device-synchronized.
82    /// Nested semantic spans remain host-timed.
83    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    /// Open a trace and own its single root span until [`Self::finish`].
102    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    /// Begin a nested span; parent is the top of the session span stack (TF Profiler call tree).
163    pub fn begin_span(&self, name: impl Into<String>, kind: SpanKind) -> SpanGuard<'_> {
164        self.begin_span_inner(name, kind, None, false)
165    }
166
167    /// Begin the single caller-controlled region used for total-time comparisons.
168    pub fn begin_measurement(&self, name: impl Into<String>) -> SpanGuard<'_> {
169        self.begin_span_inner(name, SpanKind::Function, None, true)
170    }
171
172    /// Begin a span tagged with a PyTorch-style training step (`forward` / `backward` / `optimizer`).
173    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    /// Record a timed op observation attached to `span_id`.
274    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    /// Record tensor metadata and an allocation event.
309    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    /// Record an explicit tensor allocation (TensorFlow memory timeline).
329    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    /// Record an explicit tensor deallocation.
350    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    /// Record a device-level memory checkpoint (cudaMemGetInfo-style).
371    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    /// Record one parameter gradient fact from a probe run.
393    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    /// Close the owned root span, flush, and return the trace path.
424    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
486/// Convert Unix seconds to UTC calendar components (Gregorian).
487fn 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}