Skip to main content

candle_graph/
evidence.rs

1//! One bounded evidence packet for agents, reports, comparisons, and the HTML viewer.
2
3use 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(&parameter).copied();
471            let right = candidate.get(&parameter).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}