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- 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(&parameter).copied();
463            let right = candidate.get(&parameter).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}