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- 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.graph.summary.total_ms,
208 );
209 push_list(
210 &mut out,
211 &self.findings,
212 "No trusted findings were derived.",
213 );
214 out.push_str("\n## Evidence gaps\n\n");
215 push_list(&mut out, &self.gaps, "No known gaps.");
216 out.push_str("\n## Coverage\n\n```json\n");
217 out.push_str(&serde_json::to_string_pretty(&self.health.coverage).unwrap_or_default());
218 out.push_str("\n```\n");
219 out.push_str("\n## Tensor checkpoints\n\n");
220 if self.graph.tensors.is_empty() {
221 out.push_str("No tensor checkpoints were captured.\n");
222 } else {
223 out.push_str("| Tensor | Shape | Dtype | Device | Storage | Requires grad |\n| --- | --- | --- | --- | ---: | --- |\n");
224 for tensor in &self.graph.tensors {
225 out.push_str(&format!(
226 "| `{}` | `{:?}` | `{}` | `{}` | {} B | {} |\n",
227 tensor.tensor_id,
228 tensor.shape,
229 tensor.dtype,
230 tensor.device,
231 tensor.storage_bytes,
232 tensor.requires_grad
233 ));
234 }
235 }
236 out.push_str("\n## Gradient evidence\n\n");
237 out.push_str(&format!(
238 "{} parameter gradients captured; {} require attention.\n",
239 self.graph.gradients.len(),
240 self.graph
241 .gradients
242 .iter()
243 .filter(|gradient| !matches!(
244 gradient.state,
245 crate::graph::GradientRecordState::Present
246 ))
247 .count()
248 ));
249 if let Some(comparison) = &self.comparison {
250 out.push_str("\n## Baseline comparison\n\n");
251 match comparison.total_delta_percent {
252 Some(percent) => out.push_str(&format!("Total time changed by {percent:+.2}%.\n")),
253 None => out.push_str(
254 "Total-time percentage is unavailable because the baseline was zero.\n",
255 ),
256 }
257 for warning in &comparison.warnings {
258 out.push_str(&format!("- Warning: {warning}\n"));
259 }
260 }
261 out
262 }
263}
264
265pub fn compare_documents(baseline: &TraceDocument, candidate: &TraceDocument) -> Comparison {
266 let mut warnings = Vec::new();
267 let mut comparable = analyze_health(baseline).trusted && analyze_health(candidate).trusted;
268 for (name, left, right) in [
269 (
270 "entrypoint",
271 baseline.run.entrypoint.as_str(),
272 candidate.run.entrypoint.as_str(),
273 ),
274 (
275 "phase",
276 baseline.run.phase.as_str(),
277 candidate.run.phase.as_str(),
278 ),
279 (
280 "device",
281 baseline.run.device.as_str(),
282 candidate.run.device.as_str(),
283 ),
284 ] {
285 if left != right {
286 comparable = false;
287 warnings.push(format!("{name} differs: `{left}` vs `{right}`"));
288 }
289 }
290 if baseline.run.timing_mode != candidate.run.timing_mode {
291 comparable = false;
292 warnings.push("timing mode differs".into());
293 }
294 if baseline.run.warmup_steps != candidate.run.warmup_steps {
295 comparable = false;
296 warnings.push(format!(
297 "warmup differs: {} vs {}",
298 baseline.run.warmup_steps, candidate.run.warmup_steps
299 ));
300 }
301 let descriptive = ["source_revision", "source_commit", "build_id"];
302 let baseline_conditions = baseline
303 .run
304 .tags
305 .iter()
306 .filter(|(key, _)| !descriptive.contains(&key.as_str()))
307 .collect::<BTreeMap<_, _>>();
308 let candidate_conditions = candidate
309 .run
310 .tags
311 .iter()
312 .filter(|(key, _)| !descriptive.contains(&key.as_str()))
313 .collect::<BTreeMap<_, _>>();
314 if baseline_conditions != candidate_conditions {
315 comparable = false;
316 warnings.push("workload tags differ; batch/model/precision conditions must match".into());
317 }
318 for key in descriptive {
319 if baseline.run.tags.get(key) != candidate.run.tags.get(key) {
320 warnings.push(format!(
321 "descriptive `{key}` differs, as expected for a code-change comparison"
322 ));
323 }
324 }
325 let baseline_spans = aggregate_spans(baseline);
326 let candidate_spans = aggregate_spans(candidate);
327 let mut names = baseline_spans
328 .keys()
329 .chain(candidate_spans.keys())
330 .cloned()
331 .collect::<Vec<_>>();
332 names.sort();
333 names.dedup();
334 let mut spans = names
335 .into_iter()
336 .map(|name| {
337 let (baseline_count, baseline_total_ns) =
338 baseline_spans.get(&name).copied().unwrap_or_default();
339 let (candidate_count, candidate_total_ns) =
340 candidate_spans.get(&name).copied().unwrap_or_default();
341 let delta_ns = candidate_total_ns as i128 - baseline_total_ns as i128;
342 SpanComparison {
343 name,
344 baseline_count,
345 candidate_count,
346 baseline_total_ns,
347 candidate_total_ns,
348 baseline_mean_ns: mean(baseline_total_ns, baseline_count),
349 candidate_mean_ns: mean(candidate_total_ns, candidate_count),
350 delta_ns,
351 delta_percent: percent(baseline_total_ns, candidate_total_ns),
352 }
353 })
354 .collect::<Vec<_>>();
355 spans.sort_by_key(|span| std::cmp::Reverse(span.delta_ns.unsigned_abs()));
356 spans.truncate(50);
357 let baseline_total = root_total(baseline);
358 let candidate_total = root_total(candidate);
359 let baseline_peak = crate::trace::memory::analyze_memory(baseline)
360 .summary
361 .peak_bytes;
362 let candidate_peak = crate::trace::memory::analyze_memory(candidate)
363 .summary
364 .peak_bytes;
365 Comparison {
366 schema: COMPARISON_SCHEMA.into(),
367 baseline_run_id: baseline.run.run_id.clone(),
368 candidate_run_id: candidate.run.run_id.clone(),
369 comparable,
370 warnings,
371 total_delta_ns: candidate_total as i128 - baseline_total as i128,
372 total_delta_percent: percent(baseline_total, candidate_total),
373 baseline_peak_bytes: baseline_peak,
374 candidate_peak_bytes: candidate_peak,
375 peak_delta_bytes: candidate_peak as i128 - baseline_peak as i128,
376 spans,
377 gradients: compare_gradients(baseline, candidate),
378 }
379}
380
381fn aggregate_spans(doc: &TraceDocument) -> BTreeMap<String, (usize, u64)> {
382 let mut result = BTreeMap::new();
383 let measured = measured_span_ids(doc);
384 for span in doc
385 .spans
386 .iter()
387 .filter(|span| measured.contains(span.id.as_str()))
388 {
389 let mut path = vec![span.name.clone()];
390 let mut parent = span.parent_id.as_deref();
391 let mut seen = std::collections::HashSet::new();
392 while let Some(id) = parent {
393 if !seen.insert(id) {
394 break;
395 }
396 let Some(parent_span) = doc.spans.iter().find(|candidate| candidate.id == id) else {
397 break;
398 };
399 path.push(parent_span.name.clone());
400 parent = parent_span.parent_id.as_deref();
401 }
402 path.reverse();
403 let step = span
404 .step
405 .map(|step| format!("/{step:?}"))
406 .unwrap_or_default();
407 let key = format!("{} [{}]{step}", path.join("/"), span.kind);
408 let entry = result.entry(key).or_insert((0usize, 0u64));
409 entry.0 += 1;
410 entry.1 = entry.1.saturating_add(span.duration_ns);
411 }
412 result
413}
414
415fn measured_span_ids(doc: &TraceDocument) -> std::collections::HashSet<&str> {
416 let mut ids = doc
417 .spans
418 .iter()
419 .filter(|span| span.measured)
420 .map(|span| span.id.as_str())
421 .collect::<std::collections::HashSet<_>>();
422 loop {
423 let before = ids.len();
424 for span in &doc.spans {
425 if span
426 .parent_id
427 .as_deref()
428 .is_some_and(|parent| ids.contains(parent))
429 {
430 ids.insert(span.id.as_str());
431 }
432 }
433 if ids.len() == before {
434 return ids;
435 }
436 }
437}
438
439fn compare_gradients(
440 baseline: &TraceDocument,
441 candidate: &TraceDocument,
442) -> Vec<GradientComparison> {
443 let baseline = baseline
444 .gradients
445 .iter()
446 .map(|item| (format!("{}/{}", item.root, item.key), item))
447 .collect::<BTreeMap<_, _>>();
448 let candidate = candidate
449 .gradients
450 .iter()
451 .map(|item| (format!("{}/{}", item.root, item.key), item))
452 .collect::<BTreeMap<_, _>>();
453 let mut keys = baseline
454 .keys()
455 .chain(candidate.keys())
456 .cloned()
457 .collect::<Vec<_>>();
458 keys.sort();
459 keys.dedup();
460 keys.into_iter()
461 .filter_map(|parameter| {
462 let left = baseline.get(¶meter).copied();
463 let right = candidate.get(¶meter).copied();
464 let changed = left.map(|x| (x.state, x.norm)) != right.map(|x| (x.state, x.norm));
465 changed.then(|| GradientComparison {
466 parameter,
467 baseline_state: left.map(|x| x.state.to_string()),
468 candidate_state: right.map(|x| x.state.to_string()),
469 baseline_norm: left.and_then(|x| x.norm),
470 candidate_norm: right.and_then(|x| x.norm),
471 norm_delta: left
472 .and_then(|x| x.norm)
473 .zip(right.and_then(|x| x.norm))
474 .map(|(a, b)| b - a),
475 })
476 })
477 .take(100)
478 .collect()
479}
480
481fn mean(total: u64, count: usize) -> u64 {
482 if count == 0 {
483 0
484 } else {
485 total / count as u64
486 }
487}
488
489fn root_total(doc: &TraceDocument) -> u64 {
490 doc.spans
491 .iter()
492 .filter(|span| span.measured)
493 .map(|span| span.duration_ns)
494 .sum()
495}
496
497fn percent(baseline: u64, candidate: u64) -> Option<f64> {
498 (baseline > 0).then(|| (candidate as f64 - baseline as f64) * 100.0 / baseline as f64)
499}
500
501fn push_list(out: &mut String, values: &[String], empty: &str) {
502 if values.is_empty() {
503 out.push_str(&format!("- {empty}\n"));
504 } else {
505 for value in values {
506 out.push_str(&format!("- {value}\n"));
507 }
508 }
509}
510
511#[cfg(test)]
512mod tests {
513 use super::*;
514 use crate::phase::{ExecutionPhase, ExecutionStep};
515 use crate::trace::{
516 GradientEvent, GradientState, SpanKind, SpanRecord, TimingMode, TraceRunMeta, SCHEMA,
517 };
518
519 fn document(run_id: &str, measured_ns: u64, forward_ns: u64) -> TraceDocument {
520 TraceDocument {
521 schema: SCHEMA.into(),
522 run: TraceRunMeta {
523 run_id: run_id.into(),
524 correlation_id: format!("demo/{run_id}"),
525 entrypoint: "demo::update".into(),
526 phase: ExecutionPhase::Train,
527 timestamp: "2026-08-08T00:00:00Z".into(),
528 capture_step: 1,
529 warmup_steps: 0,
530 device: "cpu".into(),
531 timing_mode: TimingMode::Host,
532 tags: [("physical_batch".into(), "2".into())].into(),
533 candle_version: None,
534 },
535 spans: vec![
536 span("session", None, false, measured_ns + 100, None),
537 span("update", Some("session"), true, measured_ns, None),
538 span(
539 "forward",
540 Some("update"),
541 false,
542 forward_ns,
543 Some(ExecutionStep::Forward),
544 ),
545 span(
546 "backward",
547 Some("update"),
548 false,
549 measured_ns.saturating_sub(forward_ns + 10),
550 Some(ExecutionStep::Backward),
551 ),
552 span(
553 "optimizer",
554 Some("update"),
555 false,
556 10,
557 Some(ExecutionStep::Optimizer),
558 ),
559 ],
560 ops: vec![],
561 tensors: vec![],
562 memory: vec![],
563 device_memory: vec![],
564 gradients: vec![GradientEvent {
565 event_id: "g".into(),
566 root: "vb".into(),
567 key: "weight".into(),
568 state: GradientState::Present,
569 norm: Some(forward_ns as f64),
570 }],
571 edges: vec![],
572 }
573 }
574
575 fn span(
576 id: &str,
577 parent: Option<&str>,
578 measured: bool,
579 duration_ns: u64,
580 step: Option<ExecutionStep>,
581 ) -> SpanRecord {
582 SpanRecord {
583 id: id.into(),
584 parent_id: parent.map(str::to_string),
585 name: id.into(),
586 kind: SpanKind::Function,
587 measured,
588 start_ns: 0,
589 closed: true,
590 duration_ns,
591 step,
592 }
593 }
594
595 #[test]
596 fn compares_measured_region_and_semantic_paths() {
597 let comparison =
598 compare_documents(&document("base", 100, 60), &document("candidate", 80, 40));
599 assert!(comparison.comparable);
600 assert_eq!(comparison.total_delta_ns, -20);
601 assert_eq!(comparison.total_delta_percent, Some(-20.0));
602 assert!(comparison
603 .spans
604 .iter()
605 .any(|span| span.name.contains("session/update/forward")));
606 assert_eq!(comparison.gradients.len(), 1);
607 }
608
609 #[test]
610 fn packet_markdown_exposes_trust_gaps_and_coverage() {
611 let packet = EvidencePacket::from_document(
612 document("candidate", 80, 40),
613 NsightEvidence::unavailable("nsys not installed"),
614 None,
615 )
616 .unwrap();
617 let markdown = packet.markdown();
618 assert!(markdown.contains("Status: **TRUSTED**"));
619 assert!(markdown.contains("nsys not installed"));
620 assert!(markdown.contains("optimizer_spans"));
621 assert!(!packet.facts.is_empty());
622 }
623}