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::capability::CaptureContract;
14use crate::phase::ExecutionPhase;
15#[cfg(feature = "candle")]
16use crate::trace::events::TensorStatsEvent;
17use crate::trace::events::{
18 DeviceIntervalEvent, DeviceMemoryEvent, EdgeEvent, GradientEvent, MemoryEvent, OpEvent,
19 SpanEndEvent, SpanStartEvent, TensorEvent, TerminalEvent, TraceEvent,
20};
21use crate::trace::memory::{resolve_dense_tensor_bytes, MemoryAction};
22use crate::trace::schema::{
23 ComparisonIdentity, GradientState, RunOutcome, TimingMode, TraceRunMeta,
24};
25
26use super::span::{
27 DeviceIntervalRecord, DeviceMemoryRecord, MemoryRecord, OpRecord, SpanGuard, SpanId, SpanKind,
28 TensorRecord,
29};
30
31pub struct TraceSession {
33 path: PathBuf,
34 inner: RefCell<SessionInner>,
35}
36
37#[derive(Debug, Clone, PartialEq, Eq)]
39pub struct ProfileRun {
40 pub entrypoint: String,
41 pub correlation_id: String,
42 pub phase: ExecutionPhase,
43 pub capture_step: u64,
45 pub warmup_steps: u64,
46 pub device: String,
47 pub measured_region_device_synchronized: bool,
48 pub timing_mode: TimingMode,
49 pub capture_contract: CaptureContract,
50 pub comparison_identity: Option<ComparisonIdentity>,
51 pub tags: BTreeMap<String, String>,
52}
53
54impl ProfileRun {
55 pub fn training(
56 entrypoint: impl Into<String>,
57 capture_step: u64,
58 device: impl Into<String>,
59 ) -> Self {
60 let entrypoint = entrypoint.into();
61 Self {
62 correlation_id: format!("{entrypoint}/update-{capture_step}"),
63 entrypoint,
64 phase: ExecutionPhase::Train,
65 capture_step,
66 warmup_steps: capture_step.saturating_sub(1),
67 device: device.into(),
68 measured_region_device_synchronized: false,
69 timing_mode: TimingMode::Host,
70 capture_contract: CaptureContract::default(),
71 comparison_identity: None,
72 tags: BTreeMap::new(),
73 }
74 }
75
76 pub fn tag(mut self, key: impl Into<String>, value: impl Into<String>) -> Self {
77 self.tags.insert(key.into(), value.into());
78 self
79 }
80
81 pub fn correlation_id(mut self, value: impl Into<String>) -> Self {
82 self.correlation_id = value.into();
83 self
84 }
85
86 pub fn device_synchronized(mut self) -> Self {
87 self.timing_mode = TimingMode::DeviceSynchronized;
88 self.measured_region_device_synchronized = true;
89 self
90 }
91
92 pub fn measured_region_device_synchronized(mut self) -> Self {
95 self.measured_region_device_synchronized = true;
96 self
97 }
98
99 pub fn inference(
100 entrypoint: impl Into<String>,
101 capture_step: u64,
102 device: impl Into<String>,
103 ) -> Self {
104 let entrypoint = entrypoint.into();
105 Self {
106 correlation_id: format!("{entrypoint}/inference-{capture_step}"),
107 entrypoint,
108 phase: ExecutionPhase::Infer,
109 capture_step,
110 warmup_steps: capture_step.saturating_sub(1),
111 device: device.into(),
112 measured_region_device_synchronized: false,
113 timing_mode: TimingMode::Host,
114 capture_contract: CaptureContract::default(),
115 comparison_identity: None,
116 tags: BTreeMap::new(),
117 }
118 }
119
120 pub fn capture_contract(mut self, contract: CaptureContract) -> Self {
121 self.capture_contract = contract;
122 self
123 }
124
125 pub fn comparison_identity(mut self, identity: ComparisonIdentity) -> Self {
126 self.comparison_identity = Some(identity);
127 self
128 }
129}
130
131struct SessionInner {
132 writer: io::BufWriter<File>,
133 span_stack: Vec<u64>,
134 next_span_id: u64,
135 next_event_id: u64,
136 id_buf: String,
137 probe_started: Instant,
138 sticky_error: Option<String>,
139}
140
141impl SessionInner {
142 fn write(&mut self, event: &TraceEvent) -> Result<()> {
148 if let Some(error) = &self.sticky_error {
149 anyhow::bail!("trace session previously failed: {error}");
150 }
151 if let Err(error) = write_event(&mut self.writer, event) {
152 self.sticky_error.get_or_insert_with(|| error.to_string());
153 return Err(error);
154 }
155 Ok(())
156 }
157}
158
159impl TraceSession {
160 pub fn open(path: impl AsRef<Path>, run: ProfileRun) -> Result<Self> {
162 let entrypoint = run.entrypoint.clone();
163 let meta = TraceRunMeta {
164 run_id: new_run_id(),
165 correlation_id: run.correlation_id,
166 entrypoint: run.entrypoint,
167 phase: run.phase,
168 timestamp: utc_iso8601_now(),
169 capture_step: run.capture_step,
170 warmup_steps: run.warmup_steps,
171 device: run.device,
172 measured_region_device_synchronized: run.measured_region_device_synchronized,
173 timing_mode: run.timing_mode,
174 capture_contract: run.capture_contract,
175 comparison_identity: run.comparison_identity,
176 tags: run.tags,
177 candle_version: None,
178 };
179 meta.validate().context("validate run provenance")?;
180 meta.capture_contract
181 .validate()
182 .context("validate capture contract")?;
183 let path = path.as_ref().to_path_buf();
184 if let Some(parent) = path.parent() {
185 std::fs::create_dir_all(parent)
186 .with_context(|| format!("create trace dir {}", parent.display()))?;
187 }
188 let file = OpenOptions::new()
189 .create(true)
190 .write(true)
191 .truncate(true)
192 .open(&path)
193 .with_context(|| format!("open trace {}", path.display()))?;
194 let mut writer = io::BufWriter::new(file);
195 write_event(&mut writer, &TraceEvent::meta(meta))?;
196 write_event(
197 &mut writer,
198 &TraceEvent::SpanStart(SpanStartEvent {
199 id: span_id_string(1),
200 parent_id: None,
201 name: entrypoint,
202 start_ns: 0,
203 kind: SpanKind::Function,
204 measured: false,
205 step: None,
206 }),
207 )?;
208 Ok(Self {
209 path,
210 inner: RefCell::new(SessionInner {
211 writer,
212 span_stack: vec![1],
213 next_span_id: 1,
214 next_event_id: 0,
215 id_buf: String::with_capacity(24),
216 probe_started: Instant::now(),
217 sticky_error: None,
218 }),
219 })
220 }
221
222 pub fn begin_span(&self, name: impl Into<String>, kind: SpanKind) -> SpanGuard<'_> {
224 self.begin_span_inner(name, kind, None, false)
225 }
226
227 pub fn begin_measurement(&self, name: impl Into<String>) -> SpanGuard<'_> {
229 self.begin_span_inner(name, SpanKind::Function, None, true)
230 }
231
232 pub fn begin_step_span(
234 &self,
235 name: impl Into<String>,
236 step: crate::phase::ExecutionStep,
237 kind: SpanKind,
238 ) -> SpanGuard<'_> {
239 self.begin_span_inner(name, kind, Some(step), false)
240 }
241
242 pub fn record_completed_host_span(
249 &self,
250 name: impl Into<String>,
251 kind: SpanKind,
252 start_ns: u64,
253 duration_ns: u64,
254 ) -> Result<SpanId> {
255 anyhow::ensure!(
256 duration_ns > 0,
257 "completed host span duration must be positive"
258 );
259 let end_ns = start_ns
260 .checked_add(duration_ns)
261 .context("completed host span interval overflows u64 nanoseconds")?;
262 anyhow::ensure!(
263 end_ns <= self.elapsed_ns(),
264 "completed host span ends in the future relative to the trace session"
265 );
266
267 let mut inner = self.inner.borrow_mut();
268 if let Some(error) = &inner.sticky_error {
269 anyhow::bail!("trace session previously failed: {error}");
270 }
271 let root_id = inner
272 .span_stack
273 .first()
274 .copied()
275 .context("trace session root span is missing")?;
276 inner.next_span_id += 1;
277 let span_id = SpanId(inner.next_span_id);
278 let id = span_id_string(span_id.0);
279
280 inner.write(&TraceEvent::SpanStart(SpanStartEvent {
281 id: id.clone(),
282 parent_id: Some(span_id_string(root_id)),
283 name: name.into(),
284 start_ns,
285 kind,
286 measured: false,
287 step: None,
288 }))?;
289 inner.write(&TraceEvent::SpanEnd(SpanEndEvent { id, duration_ns }))?;
290 Ok(span_id)
291 }
292
293 fn begin_span_inner(
294 &self,
295 name: impl Into<String>,
296 kind: SpanKind,
297 step: Option<crate::phase::ExecutionStep>,
298 measured: bool,
299 ) -> SpanGuard<'_> {
300 let started = Instant::now();
301 let start_ns = self.elapsed_ns();
302 let mut inner = self.inner.borrow_mut();
303 inner.next_span_id += 1;
304 let span_id = inner.next_span_id;
305 let parent_id = inner.span_stack.last().copied().map(span_id_string);
306
307 format_span_id(&mut inner.id_buf, span_id);
308 let id_str = inner.id_buf.clone();
309
310 let _ = inner.write(&TraceEvent::SpanStart(SpanStartEvent {
313 id: id_str,
314 parent_id,
315 name: name.into(),
316 start_ns,
317 kind,
318 measured,
319 step,
320 }));
321
322 inner.span_stack.push(span_id);
323
324 SpanGuard {
325 session: self,
326 id: SpanId(span_id),
327 started,
328 }
329 }
330
331 pub(crate) fn end_span(&self, id: SpanId, duration_ns: u64) -> Result<()> {
332 let mut inner = self.inner.borrow_mut();
333 let expected = inner
334 .span_stack
335 .last()
336 .copied()
337 .with_context(|| format!("span stack underflow closing span {}", id.0))?;
338 anyhow::ensure!(
339 expected == id.0,
340 "span_end id `{}` does not match open span `{}`",
341 id.0,
342 expected
343 );
344 inner.span_stack.pop();
345
346 format_span_id(&mut inner.id_buf, id.0);
347 let span_id = inner.id_buf.clone();
348 inner.write(&TraceEvent::SpanEnd(SpanEndEvent {
349 id: span_id,
350 duration_ns,
351 }))
352 }
353
354 pub fn elapsed_ns(&self) -> u64 {
355 self.inner
356 .borrow()
357 .probe_started
358 .elapsed()
359 .as_nanos()
360 .min(u64::MAX as u128) as u64
361 }
362
363 pub fn host_timestamp_ns(&self, instant: Instant) -> Result<u64> {
366 let probe_started = self.inner.borrow().probe_started;
367 let elapsed = instant
368 .checked_duration_since(probe_started)
369 .context("host instant predates the trace session")?;
370 Ok(elapsed.as_nanos().min(u64::MAX as u128) as u64)
371 }
372
373 pub fn record_op(&self, span_id: SpanId, op: OpRecord<'_>) -> Result<()> {
375 let output_dense_bytes =
376 resolve_dense_tensor_bytes(op.output_dense_bytes, op.shape, op.dtype);
377 let timestamp_ns = if op.timestamp_ns > 0 {
378 op.timestamp_ns
379 } else {
380 self.elapsed_ns().saturating_sub(op.duration_ns)
381 };
382 {
383 let mut inner = self.inner.borrow_mut();
384 inner.write(&TraceEvent::Op(OpEvent {
385 span_id: span_id_string(span_id.0),
386 op_name: op.op_name.into(),
387 inputs: op.inputs.to_vec(),
388 output: op.output.map(str::to_string),
389 shape: op.shape.to_vec(),
390 dtype: op.dtype.into(),
391 device: op.device.into(),
392 duration_ns: op.duration_ns,
393 timestamp_ns,
394 output_dense_bytes,
395 input_dense_bytes: op.input_dense_bytes,
396 }))?;
397 }
398
399 Ok(())
400 }
401
402 pub fn record_tensor(&self, span_id: SpanId, tensor: TensorRecord<'_>) -> Result<()> {
404 let dense_bytes =
405 resolve_dense_tensor_bytes(tensor.dense_bytes, tensor.shape, tensor.dtype);
406 let mut inner = self.inner.borrow_mut();
407 inner.write(&TraceEvent::Tensor(TensorEvent {
408 span_id: span_id_string(span_id.0),
409 tensor_id: tensor.tensor_id.into(),
410 label: tensor.label.map(str::to_string),
411 shape: tensor.shape.to_vec(),
412 dtype: tensor.dtype.into(),
413 device: tensor.device.into(),
414 requires_grad: tensor.requires_grad,
415 dense_bytes,
416 category: tensor.category,
417 }))?;
418 Ok(())
419 }
420
421 #[cfg(feature = "candle")]
423 pub fn record_tensor_stats(
424 &self,
425 span_id: &str,
426 label: &str,
427 tensor: &candle_core::Tensor,
428 ) -> Result<()> {
429 use candle_core::DType;
430
431 let elements = tensor.elem_count() as u64;
432 let (non_finite, rms, abs_max, mean) = if elements == 0 {
433 (0, 0.0, 0.0, 0.0)
434 } else {
435 let x = tensor
436 .detach()
437 .to_dtype(DType::F32)
438 .with_context(|| format!("cast tensor stats {label:?} to f32"))?;
439 let finite = x
440 .sub(&x)
441 .and_then(|delta| delta.eq(0.0))
442 .and_then(|mask| mask.to_dtype(DType::U32))
443 .and_then(|mask| mask.sum_all())
444 .and_then(|count| count.to_scalar::<u32>())
445 .with_context(|| format!("count finite elements for tensor stats {label:?}"))?
446 as u64;
447 let read = |value: candle_core::Result<candle_core::Tensor>, name: &str| {
448 value
449 .and_then(|value| value.to_scalar::<f32>())
450 .map(f64::from)
451 .with_context(|| format!("reduce {name} for tensor stats {label:?}"))
452 };
453 let rms = read(x.sqr().and_then(|x2| x2.mean_all()?.sqrt()), "rms")?;
454 let abs_max = read(x.abs().and_then(|abs| abs.max_all()), "abs_max")?;
455 let mean = read(x.mean_all(), "mean")?;
456 let json_number = |value: f64| if value.is_finite() { value } else { 0.0 };
460 (
461 elements.saturating_sub(finite),
462 json_number(rms),
463 json_number(abs_max),
464 json_number(mean),
465 )
466 };
467 let event = TensorStatsEvent {
468 span_id: span_id.to_string(),
469 label: label.to_string(),
470 shape: tensor.dims().to_vec(),
471 dtype: format!("{:?}", tensor.dtype()).to_ascii_lowercase(),
472 elements,
473 non_finite,
474 rms,
475 abs_max,
476 mean,
477 };
478 let mut inner = self.inner.borrow_mut();
479 inner.write(&TraceEvent::TensorStats(event))
480 }
481
482 pub fn record_memory_alloc(&self, span_id: SpanId, mem: MemoryRecord<'_>) -> Result<()> {
484 let timestamp_ns = mem.timestamp_ns.unwrap_or_else(|| self.elapsed_ns());
485 let mut inner = self.inner.borrow_mut();
486 inner.write(&TraceEvent::Memory(MemoryEvent {
487 timestamp_ns,
488 storage_id: mem.storage_id.into(),
489 tensor_id: mem.tensor_id.into(),
490 span_id: span_id_string(span_id.0),
491 op_name: mem.op_name.map(str::to_string),
492 device: mem.device.into(),
493 bytes: mem.bytes,
494 action: MemoryAction::Alloc,
495 shape: mem.shape.to_vec(),
496 dtype: mem.dtype.into(),
497 category: mem.category,
498 }))
499 }
500
501 pub fn record_memory_free(&self, span_id: SpanId, mem: MemoryRecord<'_>) -> Result<()> {
503 let timestamp_ns = mem.timestamp_ns.unwrap_or_else(|| self.elapsed_ns());
504 let mut inner = self.inner.borrow_mut();
505 inner.write(&TraceEvent::Memory(MemoryEvent {
506 timestamp_ns,
507 storage_id: mem.storage_id.into(),
508 tensor_id: mem.tensor_id.into(),
509 span_id: span_id_string(span_id.0),
510 op_name: mem.op_name.map(str::to_string),
511 device: mem.device.into(),
512 bytes: mem.bytes,
513 action: MemoryAction::Free,
514 shape: mem.shape.to_vec(),
515 dtype: mem.dtype.into(),
516 category: mem.category,
517 }))
518 }
519
520 pub fn record_device_memory(&self, sample: DeviceMemoryRecord<'_>) -> Result<()> {
522 anyhow::ensure!(
523 sample.used_bytes.is_some()
524 || sample.free_bytes.is_some()
525 || sample.reserved_bytes.is_some()
526 || sample.capacity_bytes.is_some(),
527 "device-memory checkpoint must contain at least one observation"
528 );
529 let timestamp_ns = sample.timestamp_ns.unwrap_or_else(|| self.elapsed_ns());
530 let mut inner = self.inner.borrow_mut();
531 inner.write(&TraceEvent::DeviceMemory(DeviceMemoryEvent {
532 timestamp_ns,
533 device: sample.device.into(),
534 used_bytes: sample.used_bytes,
535 free_bytes: sample.free_bytes,
536 reserved_bytes: sample.reserved_bytes,
537 capacity_bytes: sample.capacity_bytes,
538 }))
539 }
540
541 pub fn record_device_interval(
542 &self,
543 span_id: SpanId,
544 interval: DeviceIntervalRecord<'_>,
545 ) -> Result<()> {
546 anyhow::ensure!(
547 interval.duration_ns > 0,
548 "device interval duration must be positive"
549 );
550 let mut inner = self.inner.borrow_mut();
551 inner.write(&TraceEvent::DeviceInterval(DeviceIntervalEvent {
552 span_id: span_id_string(span_id.0),
553 device: interval.device.into(),
554 stream_id: interval.stream_id.into(),
555 clock_id: interval.clock_id.into(),
556 backend: interval.backend.into(),
557 start_ns: interval.start_ns,
558 duration_ns: interval.duration_ns,
559 }))
560 }
561
562 pub fn record_call_edge(&self, from: SpanId, to: SpanId, duration_ns: u64) -> Result<()> {
563 let mut inner = self.inner.borrow_mut();
564 inner.write(&TraceEvent::Edge(EdgeEvent::Call {
565 from_span: span_id_string(from.0),
566 to_span: span_id_string(to.0),
567 host_duration_ns: duration_ns,
568 }))
569 }
570
571 pub fn record_data_edge(&self, from_tensor: &str, to_tensor: &str) -> Result<()> {
572 anyhow::ensure!(
573 !from_tensor.is_empty() && !to_tensor.is_empty(),
574 "data-edge tensor IDs cannot be empty"
575 );
576 let mut inner = self.inner.borrow_mut();
577 inner.write(&TraceEvent::Edge(EdgeEvent::Data {
578 from_tensor: from_tensor.into(),
579 to_tensor: to_tensor.into(),
580 }))
581 }
582
583 pub fn record_gradient(
588 &self,
589 root: impl Into<String>,
590 key: impl Into<String>,
591 state: GradientState,
592 norm: Option<f64>,
593 ) -> Result<()> {
594 let root = root.into();
595 let key = key.into();
596 anyhow::ensure!(
597 !root.trim().is_empty() && !key.trim().is_empty(),
598 "gradient roots and parameter keys must not be empty"
599 );
600 anyhow::ensure!(
601 state.norm_is_valid(norm),
602 "gradient state `{state}` is inconsistent with norm {norm:?}"
603 );
604 let mut inner = self.inner.borrow_mut();
605 inner.next_event_id += 1;
606 let event_id = format!("gradient-{}", inner.next_event_id);
607 inner.write(&TraceEvent::Gradient(GradientEvent {
608 event_id,
609 root,
610 key,
611 state,
612 norm,
613 }))
614 }
615
616 pub fn flush(&self) -> Result<()> {
617 let mut inner = self.inner.borrow_mut();
618 if let Some(error) = &inner.sticky_error {
619 anyhow::bail!("trace session previously failed: {error}");
620 }
621 inner.writer.flush().context("flushing trace JSONL")
622 }
623
624 pub fn finish(self) -> Result<PathBuf> {
626 let duration_ns = self.elapsed_ns();
627 {
628 let mut inner = self.inner.borrow_mut();
629 if let Some(error) = &inner.sticky_error {
630 anyhow::bail!("trace session previously failed: {error}");
631 }
632 anyhow::ensure!(
633 inner.span_stack.as_slice() == [1],
634 "cannot finish trace with {} nested spans still open",
635 inner.span_stack.len().saturating_sub(1)
636 );
637 inner.span_stack.pop();
638 inner.write(&TraceEvent::SpanEnd(SpanEndEvent {
639 id: span_id_string(1),
640 duration_ns,
641 }))?;
642 inner.write(&TraceEvent::Terminal(TerminalEvent {
643 outcome: RunOutcome::Complete,
644 timestamp_ns: duration_ns,
645 reason: None,
646 }))?;
647 inner.writer.flush().context("flushing trace JSONL")?;
648 }
649 Ok(self.path)
650 }
651
652 pub fn finish_failed(self, reason: impl Into<String>) -> Result<PathBuf> {
654 let reason = reason.into();
655 anyhow::ensure!(!reason.trim().is_empty(), "failure reason cannot be empty");
656 let timestamp_ns = self.elapsed_ns();
657 let mut inner = self.inner.borrow_mut();
658 write_event(
659 &mut inner.writer,
660 &TraceEvent::Terminal(TerminalEvent {
661 outcome: RunOutcome::Failed,
662 timestamp_ns,
663 reason: Some(reason),
664 }),
665 )?;
666 inner
667 .writer
668 .flush()
669 .context("flushing failed trace JSONL")?;
670 drop(inner);
671 Ok(self.path.clone())
672 }
673}
674
675fn write_event<W: Write, T: Serialize>(writer: &mut W, event: &T) -> Result<()> {
676 let mut line = serde_json::to_vec(event).context("serializing trace JSONL event")?;
677 line.push(b'\n');
678 writer.write_all(&line).context("writing trace JSONL event")
679}
680
681fn format_span_id(buf: &mut String, id: u64) {
682 buf.clear();
683 use std::fmt::Write as _;
684 let _ = write!(buf, "s{id}");
685}
686
687fn span_id_string(id: u64) -> String {
688 format!("s{id}")
689}
690
691fn new_run_id() -> String {
692 let pid = std::process::id();
693 let nanos = SystemTime::now()
694 .duration_since(UNIX_EPOCH)
695 .map(|d| d.as_nanos())
696 .unwrap_or(0);
697 format!("run-{pid}-{nanos}")
698}
699
700fn utc_iso8601_now() -> String {
701 let now = SystemTime::now()
702 .duration_since(UNIX_EPOCH)
703 .expect("system clock before UNIX epoch");
704 let secs = now.as_secs();
705 let millis = now.subsec_millis();
706 let (year, month, day, hour, minute, second) = unix_secs_to_utc(secs);
707 format!("{year:04}-{month:02}-{day:02}T{hour:02}:{minute:02}:{second:02}.{millis:03}Z")
708}
709
710fn unix_secs_to_utc(secs: u64) -> (u64, u64, u64, u64, u64, u64) {
712 const SECS_PER_DAY: u64 = 86_400;
713 let days = secs / SECS_PER_DAY;
714 let rem = secs % SECS_PER_DAY;
715 let hour = rem / 3600;
716 let minute = (rem % 3600) / 60;
717 let second = rem % 60;
718
719 let z = days + 719_468;
720 let era = z / 146_097;
721 let doe = z - era * 146_097;
722 let yoe = (doe - doe / 1460 + doe / 36524 - doe / 146096) / 365;
723 let y = yoe + era * 400;
724 let doy = doe - (365 * yoe + yoe / 4 - yoe / 100);
725 let mp = (5 * doy + 2) / 153;
726 let day = doy - (153 * mp + 2) / 5 + 1;
727 let month = if mp < 10 { mp + 3 } else { mp - 9 };
728 let year = if month <= 2 { y + 1 } else { y };
729 (year, month, day, hour, minute, second)
730}
731
732#[cfg(test)]
733mod tests {
734 use super::*;
735 use crate::trace::document::parse_trace;
736 use crate::trace::events::SpanStartEvent;
737 use crate::trace::health::analyze_health;
738 use serde_json::Value;
739 use std::io::{BufRead, BufReader};
740
741 fn temp_trace(name: &str) -> PathBuf {
742 std::env::temp_dir().join(format!(
743 "candle-graph-trace-{}-{}-{name}",
744 std::process::id(),
745 std::time::SystemTime::now()
746 .duration_since(UNIX_EPOCH)
747 .unwrap()
748 .as_nanos()
749 ))
750 }
751
752 fn span_end_durations(path: &Path) -> Vec<(String, u64)> {
753 let file = std::fs::File::open(path).unwrap();
754 let reader = BufReader::new(file);
755 reader
756 .lines()
757 .map(|line| line.unwrap())
758 .filter_map(|line| {
759 let value: Value = serde_json::from_str(&line).unwrap();
760 if value.get("kind")?.as_str()? != "span_end" {
761 return None;
762 }
763 Some((
764 value["id"].as_str().unwrap().to_string(),
765 value["duration_ns"].as_u64().unwrap(),
766 ))
767 })
768 .collect()
769 }
770
771 fn read_events(path: &Path) -> Vec<TraceEvent> {
772 let file = std::fs::File::open(path).unwrap();
773 let reader = BufReader::new(file);
774 reader
775 .lines()
776 .map(|line| {
777 let line = line.unwrap();
778 serde_json::from_str(&line).unwrap_or_else(|err| {
779 panic!("invalid JSONL line `{line}`: {err}");
780 })
781 })
782 .collect()
783 }
784
785 #[test]
786 fn nested_spans_emit_parent_hierarchy_and_durations() {
787 let path = temp_trace("nested");
788 let session =
789 TraceSession::open(&path, ProfileRun::training("model::forward", 1, "cpu")).unwrap();
790
791 let inner_id = {
792 let _outer = session.begin_measurement("Model::forward");
793 std::thread::sleep(std::time::Duration::from_micros(50));
794 let inner = session.begin_span("matmul", SpanKind::Op);
795 std::thread::sleep(std::time::Duration::from_micros(50));
796 inner.id
797 };
798
799 session
800 .record_op(
801 inner_id,
802 OpRecord {
803 op_name: "matmul",
804 inputs: &["a".into(), "b".into()],
805 output: Some("c"),
806 shape: &[8, 8],
807 dtype: "f32",
808 device: "cpu",
809 duration_ns: 1200,
810 timestamp_ns: 0,
811 output_dense_bytes: None,
812 input_dense_bytes: 0,
813 },
814 )
815 .unwrap();
816
817 session.finish().unwrap();
818
819 let events = read_events(&path);
820 assert!(matches!(events.first(), Some(TraceEvent::Meta { .. })));
821
822 let starts: Vec<&SpanStartEvent> = events
823 .iter()
824 .filter_map(|event| match event {
825 TraceEvent::SpanStart(start) => Some(start),
826 _ => None,
827 })
828 .collect();
829 assert_eq!(starts.len(), 3);
830 assert_eq!(starts[0].name, "model::forward");
831 assert!(starts[0].parent_id.is_none());
832 assert_eq!(starts[1].name, "Model::forward");
833 assert_eq!(starts[1].parent_id.as_deref(), Some("s1"));
834 assert_eq!(starts[2].name, "matmul");
835 assert_eq!(starts[2].parent_id.as_deref(), Some("s2"));
836
837 let ends = span_end_durations(&path);
838 assert_eq!(ends.len(), 3);
839 assert!(ends[0].1 > 0);
840
841 let doc = parse_trace(&path).unwrap();
842 assert_eq!(doc.run.entrypoint, "model::forward");
843 assert_eq!(doc.ops.len(), 1);
844 assert_eq!(doc.ops[0].output_dense_bytes, Some(8 * 8 * 4));
845 assert!(
846 doc.memory.is_empty(),
847 "op metadata must not fabricate tensor lifetime"
848 );
849 }
850
851 #[test]
852 fn measured_region_sync_does_not_overstate_nested_span_timing() {
853 let run = ProfileRun::training("train::update", 2, "cuda:0")
854 .measured_region_device_synchronized();
855
856 assert!(run.measured_region_device_synchronized);
857 assert_eq!(run.timing_mode, TimingMode::Host);
858 }
859
860 #[test]
861 fn record_gradient_round_trips_through_trace_parser() {
862 let path = temp_trace("gradient");
863 let session =
864 TraceSession::open(&path, ProfileRun::training("train::loss", 1, "cpu")).unwrap();
865 session
866 .record_gradient("vb", "encoder.weight", GradientState::Present, Some(0.42))
867 .unwrap();
868 session.finish().unwrap();
869
870 let doc = parse_trace(&path).unwrap();
871 assert_eq!(doc.gradients.len(), 1);
872 assert_eq!(doc.gradients[0].key, "encoder.weight");
873 }
874
875 #[cfg(feature = "candle")]
876 #[test]
877 fn tensor_stats_reduce_finite_non_finite_and_empty_tensors() {
878 use candle_core::{Device, Tensor};
879
880 let path = temp_trace("tensor-stats");
881 let session =
882 TraceSession::open(&path, ProfileRun::training("train::loss", 1, "cpu")).unwrap();
883 let span = session.begin_span("forward", SpanKind::Function);
884 let span_id = format!("s{}", span.id().raw());
885 let finite = Tensor::from_vec(vec![1.0f32, -2.0, 3.0, -4.0], (4,), &Device::Cpu).unwrap();
886 session
887 .record_tensor_stats(&span_id, "finite", &finite)
888 .unwrap();
889 let corrupt =
890 Tensor::from_vec(vec![1.0f32, f32::NAN, f32::INFINITY], (3,), &Device::Cpu).unwrap();
891 session
892 .record_tensor_stats(&span_id, "corrupt", &corrupt)
893 .unwrap();
894 let empty = Tensor::zeros((0,), candle_core::DType::F32, &Device::Cpu).unwrap();
895 session
896 .record_tensor_stats(&span_id, "empty", &empty)
897 .unwrap();
898 drop(span);
899 session.finish().unwrap();
900
901 let doc = parse_trace(&path).unwrap();
902 assert_eq!(doc.tensor_stats.len(), 3);
903 let finite = &doc.tensor_stats[0];
904 assert_eq!(finite.elements, 4);
905 assert_eq!(finite.non_finite, 0);
906 assert!((finite.rms - 7.5f64.sqrt()).abs() < 1e-6);
907 assert_eq!(finite.abs_max, 4.0);
908 assert_eq!(finite.mean, -0.5);
909 assert_eq!(doc.tensor_stats[1].non_finite, 2);
910 assert_eq!(doc.tensor_stats[2].elements, 0);
911 assert_eq!(doc.tensor_stats[2].non_finite, 0);
912 assert_eq!(doc.tensor_stats[2].rms, 0.0);
913 }
914
915 #[test]
916 fn completed_host_span_is_root_attached_stack_neutral_and_overlap_valid() {
917 let path = temp_trace("completed-host-span");
918 let session =
919 TraceSession::open(&path, ProfileRun::inference("serve::request", 1, "cpu")).unwrap();
920
921 let measured = session.begin_measurement("request");
922 let live = session.begin_span("main-thread", SpanKind::Function);
923 let stack_before = session.inner.borrow().span_stack.clone();
924 let start_ns = session.elapsed_ns();
925 let end_ns = loop {
926 let elapsed = session.elapsed_ns();
927 if elapsed > start_ns {
928 break elapsed;
929 }
930 std::hint::spin_loop();
931 };
932 let duration_ns = end_ns - start_ns;
933
934 let completed_id = session
935 .record_completed_host_span(
936 "background-work",
937 SpanKind::Function,
938 start_ns,
939 duration_ns,
940 )
941 .unwrap();
942 assert_eq!(session.inner.borrow().span_stack, stack_before);
943
944 drop(live);
945 drop(measured);
946 session.finish().unwrap();
947
948 let doc = parse_trace(&path).unwrap();
949 let completed = doc
950 .spans
951 .iter()
952 .find(|span| span.id == span_id_string(completed_id.0))
953 .unwrap();
954 let measured = doc.spans.iter().find(|span| span.measured).unwrap();
955 assert_eq!(completed.parent_id.as_deref(), Some("s1"));
956 assert_eq!(completed.start_ns, start_ns);
957 assert_eq!(completed.duration_ns, duration_ns);
958 assert!(completed.closed);
959 assert!(completed.start_ns < measured.start_ns.saturating_add(measured.duration_ns));
960 assert!(measured.start_ns < completed.start_ns.saturating_add(completed.duration_ns));
961
962 let health = analyze_health(&doc);
963 assert!(health.structurally_valid, "{:?}", health.issues);
964 assert!(health.capture_complete);
965 }
966
967 #[test]
968 fn completed_host_span_requires_positive_duration() {
969 let path = temp_trace("completed-host-span-zero");
970 let session =
971 TraceSession::open(&path, ProfileRun::inference("serve::request", 1, "cpu")).unwrap();
972 let error = session
973 .record_completed_host_span("empty", SpanKind::Function, 0, 0)
974 .unwrap_err();
975 assert!(error.to_string().contains("duration must be positive"));
976
977 let overflow = session
978 .record_completed_host_span("overflow", SpanKind::Function, u64::MAX, 1)
979 .unwrap_err();
980 assert!(overflow.to_string().contains("overflows"));
981
982 let future = session
983 .record_completed_host_span(
984 "future",
985 SpanKind::Function,
986 session.elapsed_ns().saturating_add(60_000_000_000),
987 1,
988 )
989 .unwrap_err();
990 assert!(future.to_string().contains("ends in the future"));
991
992 let before_session = session
993 .inner
994 .borrow()
995 .probe_started
996 .checked_sub(std::time::Duration::from_nanos(1))
997 .unwrap();
998 assert!(session.host_timestamp_ns(before_session).is_err());
999 let now = Instant::now();
1000 assert!(session.host_timestamp_ns(now).unwrap() <= session.elapsed_ns());
1001 session.finish().unwrap();
1002 }
1003
1004 #[test]
1005 fn sticky_error_poisons_event_writes_and_finish() {
1006 let path = temp_trace("sticky-error");
1007 let session =
1008 TraceSession::open(&path, ProfileRun::training("train::update", 1, "cpu")).unwrap();
1009 session.inner.borrow_mut().sticky_error = Some("boom".into());
1010
1011 let error = session
1012 .record_op(
1013 SpanId(1),
1014 OpRecord {
1015 op_name: "matmul",
1016 inputs: &[],
1017 output: None,
1018 shape: &[2, 2],
1019 dtype: "f32",
1020 device: "cpu",
1021 duration_ns: 100,
1022 timestamp_ns: 0,
1023 output_dense_bytes: None,
1024 input_dense_bytes: 0,
1025 },
1026 )
1027 .unwrap_err();
1028 assert!(error.to_string().contains("previously failed"));
1029
1030 let error = session.finish().unwrap_err();
1031 assert!(error.to_string().contains("previously failed"));
1032 }
1033}