1use std::collections::BTreeMap;
4use std::path::Path;
5
6use anyhow::{Context, Result};
7use serde::{Deserialize, Serialize};
8
9use crate::graph::{build_from_trace, ExecutionGraph};
10use crate::nsight::{GpuEvidenceStatus, NsightEvidence};
11use crate::trace::{analyze_health, parse_trace, TraceDocument, TraceHealth, TraceRunMeta};
12
13pub const EVIDENCE_SCHEMA: &str = "candle-graph/evidence/1";
14pub const COMPARISON_SCHEMA: &str = "candle-graph/comparison/1";
15
16#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
17pub struct EvidencePacket {
18 pub schema: String,
19 pub provenance: TraceRunMeta,
20 pub health: TraceHealth,
21 pub findings: Vec<String>,
22 pub facts: Vec<EvidenceFact>,
23 pub gaps: Vec<String>,
24 pub graph: ExecutionGraph,
25 pub gpu: NsightEvidence,
26 #[serde(default, skip_serializing_if = "Option::is_none")]
27 pub comparison: Option<Comparison>,
28}
29
30#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
31pub struct EvidenceFact {
32 pub code: String,
33 pub label: String,
34 pub value: f64,
35 pub unit: String,
36 pub source: String,
37}
38
39#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
40pub struct Comparison {
41 pub schema: String,
42 pub baseline_run_id: String,
43 pub candidate_run_id: String,
44 pub comparable: bool,
45 pub warnings: Vec<String>,
46 pub total_delta_ns: i128,
47 pub total_delta_percent: Option<f64>,
48 pub baseline_peak_bytes: u64,
49 pub candidate_peak_bytes: u64,
50 pub peak_delta_bytes: i128,
51 pub spans: Vec<SpanComparison>,
52 pub gradients: Vec<GradientComparison>,
53}
54
55#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
56pub struct SpanComparison {
57 pub name: String,
58 pub baseline_count: usize,
59 pub candidate_count: usize,
60 pub baseline_total_ns: u64,
61 pub candidate_total_ns: u64,
62 pub baseline_mean_ns: u64,
63 pub candidate_mean_ns: u64,
64 pub delta_ns: i128,
65 pub delta_percent: Option<f64>,
66}
67
68#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
69pub struct GradientComparison {
70 pub parameter: String,
71 pub baseline_state: Option<String>,
72 pub candidate_state: Option<String>,
73 pub baseline_norm: Option<f64>,
74 pub candidate_norm: Option<f64>,
75 pub norm_delta: Option<f64>,
76}
77
78pub fn build_evidence(
79 trace: &Path,
80 baseline: Option<&Path>,
81 nsight_dir: Option<&Path>,
82) -> Result<EvidencePacket> {
83 let doc = parse_trace(trace).with_context(|| format!("parse trace {}", trace.display()))?;
84 let comparison = baseline
85 .map(|path| {
86 let baseline =
87 parse_trace(path).with_context(|| format!("parse baseline {}", path.display()))?;
88 Ok::<_, anyhow::Error>(compare_documents(&baseline, &doc))
89 })
90 .transpose()?;
91 let measured = measured_span_ids(&doc);
92 let expected_semantic_keys = doc
93 .spans
94 .iter()
95 .filter(|span| measured.contains(span.id.as_str()))
96 .map(|span| span.name.clone())
97 .collect::<Vec<_>>();
98 EvidencePacket::from_document(
99 doc,
100 NsightEvidence::load_optional(nsight_dir, &expected_semantic_keys),
101 comparison,
102 )
103}
104
105impl EvidencePacket {
106 pub fn from_document(
107 doc: TraceDocument,
108 gpu: NsightEvidence,
109 comparison: Option<Comparison>,
110 ) -> Result<Self> {
111 let health = analyze_health(&doc);
112 let graph = build_from_trace(&doc)?;
113 let mut findings = graph
114 .summary
115 .slowest_spans
116 .iter()
117 .take(5)
118 .map(|span| {
119 format!(
120 "{}: {:.2} ms self time",
121 span.name,
122 span.self_time_ns as f64 / 1_000_000.0
123 )
124 })
125 .collect::<Vec<_>>();
126 let mut facts = vec![EvidenceFact {
127 code: "measured_total".into(),
128 label: "Measured update total".into(),
129 value: graph.summary.total_ms,
130 unit: "ms".into(),
131 source: "measured_span".into(),
132 }];
133 facts.extend(
134 graph
135 .summary
136 .slowest_spans
137 .iter()
138 .take(8)
139 .map(|span| EvidenceFact {
140 code: "span_self_time".into(),
141 label: span.name.clone(),
142 value: span.self_time_ns as f64,
143 unit: "ns".into(),
144 source: span.id.clone(),
145 }),
146 );
147 if !graph.gradients.is_empty() {
148 let concerning = graph
149 .gradients
150 .iter()
151 .filter(|gradient| {
152 !matches!(gradient.state, crate::graph::GradientRecordState::Present)
153 })
154 .count();
155 findings.push(format!(
156 "{} gradient facts captured; {} require attention",
157 graph.gradients.len(),
158 concerning
159 ));
160 }
161 if gpu.status == GpuEvidenceStatus::Available {
162 if let Some(kernel) = gpu.kernels.first() {
163 findings.push(format!(
164 "top GPU kernel `{}`: {:.2} ms total",
165 kernel.name,
166 kernel.total_ns as f64 / 1_000_000.0
167 ));
168 }
169 }
170 let mut gaps = health
171 .gaps()
172 .map(|issue| issue.message.clone())
173 .collect::<Vec<_>>();
174 if gpu.status != GpuEvidenceStatus::Available {
175 gaps.push(
176 gpu.reason
177 .clone()
178 .unwrap_or_else(|| "GPU evidence is unavailable".into()),
179 );
180 }
181 Ok(Self {
182 schema: EVIDENCE_SCHEMA.into(),
183 provenance: doc.run,
184 health,
185 findings,
186 facts,
187 gaps,
188 graph,
189 gpu,
190 comparison,
191 })
192 }
193
194 pub fn markdown(&self) -> String {
195 let trust = if self.health.trusted {
196 "TRUSTED"
197 } else {
198 "UNTRUSTED"
199 };
200 let mut out = format!(
201 "# candle-graph evidence\n\n- Status: **{trust}**\n- Entrypoint: `{}`\n- Capture update: {} ({} warmup update{})\n- Device: `{}`\n- Timing mode: `{:?}`\n- Measured region device-synchronized: {}\n- Total: {:.2} ms\n\n## Findings\n\n",
202 self.provenance.entrypoint,
203 self.provenance.capture_step,
204 self.provenance.warmup_steps,
205 if self.provenance.warmup_steps == 1 { "" } else { "s" },
206 self.provenance.device,
207 self.provenance.timing_mode,
208 self.provenance.measured_region_device_synchronized,
209 self.graph.summary.total_ms,
210 );
211 push_list(
212 &mut out,
213 &self.findings,
214 "No trusted findings were derived.",
215 );
216 out.push_str("\n## Evidence gaps\n\n");
217 push_list(&mut out, &self.gaps, "No known gaps.");
218 out.push_str("\n## Coverage\n\n```json\n");
219 out.push_str(&serde_json::to_string_pretty(&self.health.coverage).unwrap_or_default());
220 out.push_str("\n```\n");
221 out.push_str("\n## Tensor checkpoints\n\n");
222 if self.graph.tensors.is_empty() {
223 out.push_str("No tensor checkpoints were captured.\n");
224 } else {
225 out.push_str("| Tensor | Shape | Dtype | Device | Storage | Requires grad |\n| --- | --- | --- | --- | ---: | --- |\n");
226 for tensor in &self.graph.tensors {
227 out.push_str(&format!(
228 "| `{}` | `{:?}` | `{}` | `{}` | {} B | {} |\n",
229 tensor.tensor_id,
230 tensor.shape,
231 tensor.dtype,
232 tensor.device,
233 tensor.storage_bytes,
234 tensor.requires_grad
235 ));
236 }
237 }
238 out.push_str("\n## Gradient evidence\n\n");
239 out.push_str(&format!(
240 "{} parameter gradients captured; {} require attention.\n",
241 self.graph.gradients.len(),
242 self.graph
243 .gradients
244 .iter()
245 .filter(|gradient| !matches!(
246 gradient.state,
247 crate::graph::GradientRecordState::Present
248 ))
249 .count()
250 ));
251 if let Some(comparison) = &self.comparison {
252 out.push_str("\n## Baseline comparison\n\n");
253 match comparison.total_delta_percent {
254 Some(percent) => out.push_str(&format!("Total time changed by {percent:+.2}%.\n")),
255 None => out.push_str(
256 "Total-time percentage is unavailable because the baseline was zero.\n",
257 ),
258 }
259 for warning in &comparison.warnings {
260 out.push_str(&format!("- Warning: {warning}\n"));
261 }
262 }
263 out
264 }
265}
266
267pub fn compare_documents(baseline: &TraceDocument, candidate: &TraceDocument) -> Comparison {
268 let mut warnings = Vec::new();
269 let mut comparable = analyze_health(baseline).trusted && analyze_health(candidate).trusted;
270 for (name, left, right) in [
271 (
272 "entrypoint",
273 baseline.run.entrypoint.as_str(),
274 candidate.run.entrypoint.as_str(),
275 ),
276 (
277 "phase",
278 baseline.run.phase.as_str(),
279 candidate.run.phase.as_str(),
280 ),
281 (
282 "device",
283 baseline.run.device.as_str(),
284 candidate.run.device.as_str(),
285 ),
286 ] {
287 if left != right {
288 comparable = false;
289 warnings.push(format!("{name} differs: `{left}` vs `{right}`"));
290 }
291 }
292 if baseline.run.timing_mode != candidate.run.timing_mode {
293 comparable = false;
294 warnings.push("timing mode differs".into());
295 }
296 if baseline.run.measured_region_device_synchronized
297 != candidate.run.measured_region_device_synchronized
298 {
299 comparable = false;
300 warnings.push("measured-region device synchronization differs".into());
301 }
302 if baseline.run.warmup_steps != candidate.run.warmup_steps {
303 comparable = false;
304 warnings.push(format!(
305 "warmup differs: {} vs {}",
306 baseline.run.warmup_steps, candidate.run.warmup_steps
307 ));
308 }
309 let descriptive = ["source_revision", "source_commit", "build_id"];
310 let baseline_conditions = baseline
311 .run
312 .tags
313 .iter()
314 .filter(|(key, _)| !descriptive.contains(&key.as_str()))
315 .collect::<BTreeMap<_, _>>();
316 let candidate_conditions = candidate
317 .run
318 .tags
319 .iter()
320 .filter(|(key, _)| !descriptive.contains(&key.as_str()))
321 .collect::<BTreeMap<_, _>>();
322 if baseline_conditions != candidate_conditions {
323 comparable = false;
324 warnings.push("workload tags differ; batch/model/precision conditions must match".into());
325 }
326 for key in descriptive {
327 if baseline.run.tags.get(key) != candidate.run.tags.get(key) {
328 warnings.push(format!(
329 "descriptive `{key}` differs, as expected for a code-change comparison"
330 ));
331 }
332 }
333 let baseline_spans = aggregate_spans(baseline);
334 let candidate_spans = aggregate_spans(candidate);
335 let mut names = baseline_spans
336 .keys()
337 .chain(candidate_spans.keys())
338 .cloned()
339 .collect::<Vec<_>>();
340 names.sort();
341 names.dedup();
342 let mut spans = names
343 .into_iter()
344 .map(|name| {
345 let (baseline_count, baseline_total_ns) =
346 baseline_spans.get(&name).copied().unwrap_or_default();
347 let (candidate_count, candidate_total_ns) =
348 candidate_spans.get(&name).copied().unwrap_or_default();
349 let delta_ns = candidate_total_ns as i128 - baseline_total_ns as i128;
350 SpanComparison {
351 name,
352 baseline_count,
353 candidate_count,
354 baseline_total_ns,
355 candidate_total_ns,
356 baseline_mean_ns: mean(baseline_total_ns, baseline_count),
357 candidate_mean_ns: mean(candidate_total_ns, candidate_count),
358 delta_ns,
359 delta_percent: percent(baseline_total_ns, candidate_total_ns),
360 }
361 })
362 .collect::<Vec<_>>();
363 spans.sort_by_key(|span| std::cmp::Reverse(span.delta_ns.unsigned_abs()));
364 spans.truncate(50);
365 let baseline_total = root_total(baseline);
366 let candidate_total = root_total(candidate);
367 let baseline_peak = crate::trace::memory::analyze_memory(baseline)
368 .summary
369 .peak_bytes;
370 let candidate_peak = crate::trace::memory::analyze_memory(candidate)
371 .summary
372 .peak_bytes;
373 Comparison {
374 schema: COMPARISON_SCHEMA.into(),
375 baseline_run_id: baseline.run.run_id.clone(),
376 candidate_run_id: candidate.run.run_id.clone(),
377 comparable,
378 warnings,
379 total_delta_ns: candidate_total as i128 - baseline_total as i128,
380 total_delta_percent: percent(baseline_total, candidate_total),
381 baseline_peak_bytes: baseline_peak,
382 candidate_peak_bytes: candidate_peak,
383 peak_delta_bytes: candidate_peak as i128 - baseline_peak as i128,
384 spans,
385 gradients: compare_gradients(baseline, candidate),
386 }
387}
388
389fn aggregate_spans(doc: &TraceDocument) -> BTreeMap<String, (usize, u64)> {
390 let mut result = BTreeMap::new();
391 let measured = measured_span_ids(doc);
392 for span in doc
393 .spans
394 .iter()
395 .filter(|span| measured.contains(span.id.as_str()))
396 {
397 let mut path = vec![span.name.clone()];
398 let mut parent = span.parent_id.as_deref();
399 let mut seen = std::collections::HashSet::new();
400 while let Some(id) = parent {
401 if !seen.insert(id) {
402 break;
403 }
404 let Some(parent_span) = doc.spans.iter().find(|candidate| candidate.id == id) else {
405 break;
406 };
407 path.push(parent_span.name.clone());
408 parent = parent_span.parent_id.as_deref();
409 }
410 path.reverse();
411 let step = span
412 .step
413 .map(|step| format!("/{step:?}"))
414 .unwrap_or_default();
415 let key = format!("{} [{}]{step}", path.join("/"), span.kind);
416 let entry = result.entry(key).or_insert((0usize, 0u64));
417 entry.0 += 1;
418 entry.1 = entry.1.saturating_add(span.duration_ns);
419 }
420 result
421}
422
423fn measured_span_ids(doc: &TraceDocument) -> std::collections::HashSet<&str> {
424 let mut ids = doc
425 .spans
426 .iter()
427 .filter(|span| span.measured)
428 .map(|span| span.id.as_str())
429 .collect::<std::collections::HashSet<_>>();
430 loop {
431 let before = ids.len();
432 for span in &doc.spans {
433 if span
434 .parent_id
435 .as_deref()
436 .is_some_and(|parent| ids.contains(parent))
437 {
438 ids.insert(span.id.as_str());
439 }
440 }
441 if ids.len() == before {
442 return ids;
443 }
444 }
445}
446
447fn compare_gradients(
448 baseline: &TraceDocument,
449 candidate: &TraceDocument,
450) -> Vec<GradientComparison> {
451 let baseline = baseline
452 .gradients
453 .iter()
454 .map(|item| (format!("{}/{}", item.root, item.key), item))
455 .collect::<BTreeMap<_, _>>();
456 let candidate = candidate
457 .gradients
458 .iter()
459 .map(|item| (format!("{}/{}", item.root, item.key), item))
460 .collect::<BTreeMap<_, _>>();
461 let mut keys = baseline
462 .keys()
463 .chain(candidate.keys())
464 .cloned()
465 .collect::<Vec<_>>();
466 keys.sort();
467 keys.dedup();
468 keys.into_iter()
469 .filter_map(|parameter| {
470 let left = baseline.get(¶meter).copied();
471 let right = candidate.get(¶meter).copied();
472 let changed = left.map(|x| (x.state, x.norm)) != right.map(|x| (x.state, x.norm));
473 changed.then(|| GradientComparison {
474 parameter,
475 baseline_state: left.map(|x| x.state.to_string()),
476 candidate_state: right.map(|x| x.state.to_string()),
477 baseline_norm: left.and_then(|x| x.norm),
478 candidate_norm: right.and_then(|x| x.norm),
479 norm_delta: left
480 .and_then(|x| x.norm)
481 .zip(right.and_then(|x| x.norm))
482 .map(|(a, b)| b - a),
483 })
484 })
485 .take(100)
486 .collect()
487}
488
489fn mean(total: u64, count: usize) -> u64 {
490 if count == 0 {
491 0
492 } else {
493 total / count as u64
494 }
495}
496
497fn root_total(doc: &TraceDocument) -> u64 {
498 doc.spans
499 .iter()
500 .filter(|span| span.measured)
501 .map(|span| span.duration_ns)
502 .sum()
503}
504
505fn percent(baseline: u64, candidate: u64) -> Option<f64> {
506 (baseline > 0).then(|| (candidate as f64 - baseline as f64) * 100.0 / baseline as f64)
507}
508
509fn push_list(out: &mut String, values: &[String], empty: &str) {
510 if values.is_empty() {
511 out.push_str(&format!("- {empty}\n"));
512 } else {
513 for value in values {
514 out.push_str(&format!("- {value}\n"));
515 }
516 }
517}
518
519#[cfg(test)]
520mod tests {
521 use super::*;
522 use crate::phase::{ExecutionPhase, ExecutionStep};
523 use crate::trace::{
524 GradientEvent, GradientState, SpanKind, SpanRecord, TimingMode, TraceRunMeta, SCHEMA,
525 };
526
527 fn document(run_id: &str, measured_ns: u64, forward_ns: u64) -> TraceDocument {
528 TraceDocument {
529 schema: SCHEMA.into(),
530 run: TraceRunMeta {
531 run_id: run_id.into(),
532 correlation_id: format!("demo/{run_id}"),
533 entrypoint: "demo::update".into(),
534 phase: ExecutionPhase::Train,
535 timestamp: "2026-08-08T00:00:00Z".into(),
536 capture_step: 1,
537 warmup_steps: 0,
538 device: "cpu".into(),
539 measured_region_device_synchronized: false,
540 timing_mode: TimingMode::Host,
541 tags: [("physical_batch".into(), "2".into())].into(),
542 candle_version: None,
543 },
544 spans: vec![
545 span("session", None, false, measured_ns + 100, None),
546 span("update", Some("session"), true, measured_ns, None),
547 span(
548 "forward",
549 Some("update"),
550 false,
551 forward_ns,
552 Some(ExecutionStep::Forward),
553 ),
554 span(
555 "backward",
556 Some("update"),
557 false,
558 measured_ns.saturating_sub(forward_ns + 10),
559 Some(ExecutionStep::Backward),
560 ),
561 span(
562 "optimizer",
563 Some("update"),
564 false,
565 10,
566 Some(ExecutionStep::Optimizer),
567 ),
568 ],
569 ops: vec![],
570 tensors: vec![],
571 memory: vec![],
572 device_memory: vec![],
573 gradients: vec![GradientEvent {
574 event_id: "g".into(),
575 root: "vb".into(),
576 key: "weight".into(),
577 state: GradientState::Present,
578 norm: Some(forward_ns as f64),
579 }],
580 edges: vec![],
581 }
582 }
583
584 fn span(
585 id: &str,
586 parent: Option<&str>,
587 measured: bool,
588 duration_ns: u64,
589 step: Option<ExecutionStep>,
590 ) -> SpanRecord {
591 SpanRecord {
592 id: id.into(),
593 parent_id: parent.map(str::to_string),
594 name: id.into(),
595 kind: SpanKind::Function,
596 measured,
597 start_ns: 0,
598 closed: true,
599 duration_ns,
600 step,
601 }
602 }
603
604 #[test]
605 fn compares_measured_region_and_semantic_paths() {
606 let comparison =
607 compare_documents(&document("base", 100, 60), &document("candidate", 80, 40));
608 assert!(comparison.comparable);
609 assert_eq!(comparison.total_delta_ns, -20);
610 assert_eq!(comparison.total_delta_percent, Some(-20.0));
611 assert!(comparison
612 .spans
613 .iter()
614 .any(|span| span.name.contains("session/update/forward")));
615 assert_eq!(comparison.gradients.len(), 1);
616 }
617
618 #[test]
619 fn rejects_different_measured_region_synchronization_contracts() {
620 let baseline = document("base", 100, 60);
621 let mut candidate = document("candidate", 80, 40);
622 candidate.run.measured_region_device_synchronized = true;
623
624 let comparison = compare_documents(&baseline, &candidate);
625
626 assert!(!comparison.comparable);
627 assert!(comparison
628 .warnings
629 .iter()
630 .any(|warning| warning.contains("measured-region device synchronization differs")));
631 }
632
633 #[test]
634 fn packet_markdown_exposes_trust_gaps_and_coverage() {
635 let packet = EvidencePacket::from_document(
636 document("candidate", 80, 40),
637 NsightEvidence::unavailable("nsys not installed"),
638 None,
639 )
640 .unwrap();
641 let markdown = packet.markdown();
642 assert!(markdown.contains("Status: **TRUSTED**"));
643 assert!(markdown.contains("Measured region device-synchronized: false"));
644 assert!(markdown.contains("nsys not installed"));
645 assert!(markdown.contains("optimizer_spans"));
646 assert!(!packet.facts.is_empty());
647 }
648}