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