Skip to main content

candle_graph/
model_ir.rs

1//! Unified, agent-oriented representation of a Candle model crate.
2//!
3//! The structure and expression analyzers intentionally keep their own compact arenas while
4//! running. `ModelIr` is the durable interchange layer that joins those arenas with Cargo,
5//! pipeline, optimizer, artifact, tensor-contract, and runtime evidence.
6
7use std::collections::BTreeMap;
8
9use serde::{Deserialize, Serialize};
10
11pub use crate::phase::ExecutionPhase;
12
13/// Versioned schema emitted by scans and consumed by the query/runtime layers.
14pub const MODEL_IR_SCHEMA: &str = "candle-graph/model/1";
15
16#[derive(Debug, Clone, PartialEq, Eq, PartialOrd, Ord, Hash, Serialize, Deserialize)]
17#[serde(transparent)]
18pub struct StableId(pub String);
19
20impl StableId {
21    pub fn new(kind: &str, parts: impl IntoIterator<Item = impl AsRef<str>>) -> Self {
22        let mut value = String::from(kind);
23        for part in parts {
24            value.push(':');
25            escape_id_part(part.as_ref(), &mut value);
26        }
27        Self(value)
28    }
29}
30
31impl std::fmt::Display for StableId {
32    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
33        f.write_str(&self.0)
34    }
35}
36
37#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
38#[serde(rename_all = "snake_case")]
39pub enum EvidenceKind {
40    Source,
41    Cargo,
42    Checkpoint,
43    Runtime,
44    Inferred,
45}
46
47#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
48#[serde(rename_all = "snake_case")]
49pub enum Confidence {
50    Proven,
51    Conditional,
52    Heuristic,
53    Unknown,
54}
55
56#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
57pub struct Evidence {
58    pub kind: EvidenceKind,
59    pub confidence: Confidence,
60    pub source: Option<String>,
61    pub detail: String,
62}
63
64#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
65#[serde(rename_all = "snake_case")]
66pub enum Visibility {
67    Public,
68    Crate,
69    Restricted,
70    Private,
71    Unknown,
72}
73
74#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
75#[serde(rename_all = "snake_case")]
76pub enum BuilderRole {
77    Trainable,
78    Frozen,
79    State,
80    Conditional,
81    Unknown,
82}
83
84#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
85#[serde(rename_all = "snake_case")]
86pub enum ParameterRole {
87    Optimized,
88    Frozen,
89    RunningState,
90    Excluded,
91    Conditional,
92    Unknown,
93}
94
95#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
96#[serde(rename_all = "snake_case")]
97pub enum TensorRole {
98    Input,
99    Output,
100    Parameter,
101    Activation,
102    Target,
103    Mask,
104    Loss,
105    Cache,
106    State,
107    Unknown,
108}
109
110#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
111#[serde(rename_all = "snake_case")]
112pub enum DeviceFact {
113    Cpu,
114    Cuda { ordinal: Option<u32> },
115    Metal,
116    SameAs(String),
117    Unknown,
118}
119
120#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
121#[serde(rename_all = "snake_case")]
122pub enum LayoutFact {
123    Contiguous,
124    NonContiguous,
125    Strided,
126    SameAs(String),
127    Unknown,
128}
129
130#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
131pub struct Dimension {
132    /// Stable semantic label when one is known (`batch`, `tokens`, `hidden`, ...).
133    pub name: Option<String>,
134    /// Literal or symbolic Rust expression.
135    pub expr: String,
136}
137
138#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize, Default)]
139pub struct ShapeFact {
140    pub rank: Option<usize>,
141    pub dimensions: Vec<Dimension>,
142    pub source_expr: Option<String>,
143}
144
145#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
146pub struct TensorContract {
147    pub id: StableId,
148    pub name: String,
149    pub role: TensorRole,
150    pub owner_function: StableId,
151    pub parameter: Option<StableId>,
152    pub shape: ShapeFact,
153    pub dtype: String,
154    pub device: DeviceFact,
155    pub layout: LayoutFact,
156    pub requires_grad: Option<bool>,
157    /// Train vs inference graph this contract belongs to.
158    #[serde(default, skip_serializing_if = "Option::is_none")]
159    pub execution_phase: Option<ExecutionPhase>,
160    pub evidence: Vec<Evidence>,
161}
162
163#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
164pub struct BuilderNamespace {
165    pub name: String,
166    pub role: BuilderRole,
167    pub evidence: Vec<Evidence>,
168}
169
170#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
171pub struct Component {
172    pub id: StableId,
173    pub name: String,
174    pub qualified_name: String,
175    pub source: String,
176    pub constructor: StableId,
177    pub builders: Vec<BuilderNamespace>,
178    pub modules: Vec<StableId>,
179    pub parameters: Vec<StableId>,
180    pub entrypoints: Vec<StableId>,
181    pub evidence: Vec<Evidence>,
182}
183
184#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
185pub struct ArchitectureEdge {
186    pub id: StableId,
187    pub from: StableId,
188    pub to: StableId,
189    pub via_function: StableId,
190    pub source: String,
191    pub evidence: Vec<Evidence>,
192}
193
194#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
195pub struct Module {
196    pub id: StableId,
197    pub component: StableId,
198    pub parent: Option<StableId>,
199    pub type_name: String,
200    pub qualified_type: Option<String>,
201    pub field: Option<String>,
202    pub builder_root: String,
203    pub prefix: String,
204    pub repeat: Option<String>,
205    pub source: String,
206    pub confidence: Confidence,
207}
208
209#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
210pub struct Parameter {
211    pub id: StableId,
212    pub component: StableId,
213    pub module: StableId,
214    pub key: String,
215    pub builder_root: String,
216    pub role: ParameterRole,
217    pub kind: String,
218    pub symbolic_shape: Option<String>,
219    pub checkpoint_shape: Option<Vec<usize>>,
220    pub checkpoint_dtype: Option<String>,
221    pub source: String,
222    pub uses: Vec<StableId>,
223    pub optimizer_memberships: Vec<StableId>,
224    pub evidence: Vec<Evidence>,
225}
226
227#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
228pub struct Function {
229    pub id: StableId,
230    pub name: String,
231    pub qualified_name: String,
232    pub owner_type: Option<String>,
233    pub visibility: Visibility,
234    pub parameters: Vec<FunctionParameter>,
235    pub return_type: Option<String>,
236    /// Source-level `#[cfg(...)]` predicates inherited by this definition.
237    pub cfg_predicates: Vec<String>,
238    /// Whether those predicates match the selected Cargo feature/target context.
239    /// `None` means no Cargo context was available or a predicate was unsupported.
240    pub cfg_active: Option<bool>,
241    pub source: String,
242    pub calls: Vec<StableId>,
243    pub tensor_inputs: Vec<StableId>,
244    pub tensor_outputs: Vec<StableId>,
245    pub is_entrypoint: bool,
246    pub is_loss: bool,
247    /// Static graphs built for this entrypoint (train and/or infer).
248    #[serde(default)]
249    pub execution_phases: Vec<ExecutionPhase>,
250}
251
252#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
253pub struct FunctionParameter {
254    pub name: String,
255    pub type_name: String,
256}
257
258#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
259pub struct Operation {
260    pub id: StableId,
261    pub function: StableId,
262    pub name: String,
263    pub qualified_name: Option<String>,
264    pub inputs: Vec<StableId>,
265    pub output: StableId,
266    pub source: String,
267    pub dtype_rule: String,
268    pub gradient_rule: String,
269    pub device_rule: String,
270    pub shape_rule: String,
271    /// Float-range transfer after rounding (`real`, `non_negative`, `saturating_unit`, …).
272    #[serde(default)]
273    pub domain_rule: String,
274    #[serde(default, skip_serializing_if = "Option::is_none")]
275    pub execution_phase: Option<ExecutionPhase>,
276    #[serde(default, skip_serializing_if = "Option::is_none")]
277    pub avg_duration_ns: Option<u64>,
278    #[serde(default, skip_serializing_if = "Option::is_none")]
279    pub timing_samples: Option<u64>,
280    pub evidence: Vec<Evidence>,
281}
282
283#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
284#[serde(rename_all = "snake_case")]
285pub enum StageKind {
286    Prepare,
287    Train,
288    Evaluate,
289    Export,
290    Probe,
291    Unknown,
292}
293
294#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize, Default)]
295#[serde(rename_all = "snake_case")]
296pub enum StageDispatchKind {
297    #[default]
298    Unknown,
299    Inline,
300    Subprocess,
301}
302
303#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
304pub struct PipelineStage {
305    pub id: StableId,
306    pub name: String,
307    pub kind: StageKind,
308    pub function: StableId,
309    pub order: Option<usize>,
310    pub components: Vec<StableId>,
311    pub consumes: Vec<StableId>,
312    pub produces: Vec<StableId>,
313    pub depends_on: Vec<StableId>,
314    pub source: String,
315    pub evidence: Vec<Evidence>,
316    #[serde(default)]
317    pub dispatch: StageDispatchKind,
318    #[serde(default, skip_serializing_if = "Option::is_none")]
319    pub subprocess_key: Option<String>,
320    #[serde(default, skip_serializing_if = "Vec::is_empty")]
321    pub cli_flags: Vec<String>,
322    #[serde(default, skip_serializing_if = "Option::is_none")]
323    pub launcher: Option<String>,
324    #[serde(default, skip_serializing_if = "Option::is_none")]
325    pub orchestrator: Option<String>,
326}
327
328#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
329#[serde(rename_all = "snake_case")]
330pub enum ArtifactKind {
331    Checkpoint,
332    OptimizerState,
333    Dataset,
334    Vocabulary,
335    Cache,
336    EvaluationReport,
337    Configuration,
338    Unknown,
339}
340
341#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
342pub struct Artifact {
343    pub id: StableId,
344    pub name: String,
345    pub kind: ArtifactKind,
346    pub path_expr: String,
347    pub produced_by: Vec<StableId>,
348    pub consumed_by: Vec<StableId>,
349    pub source: String,
350    pub evidence: Vec<Evidence>,
351}
352
353#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
354#[serde(rename_all = "snake_case")]
355pub enum BuilderSourceKind {
356    VarMap,
357    MmapSafetensors,
358    FromTensors,
359    BufferedSafetensors,
360    Unknown,
361}
362
363#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
364pub struct AssemblySite {
365    pub id: StableId,
366    pub function: StableId,
367    pub function_name: String,
368    pub component: StableId,
369    pub component_name: String,
370    pub builder_root: String,
371    pub prefix_chain: Vec<String>,
372    pub varmap: Option<String>,
373    pub source_kind: BuilderSourceKind,
374    pub role: BuilderRole,
375    pub checkpoint_load: Option<String>,
376    pub source: String,
377    pub evidence: Vec<Evidence>,
378}
379
380#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
381pub struct OptimizerMembership {
382    pub id: StableId,
383    pub stage: StableId,
384    pub optimizer: String,
385    pub varmap: String,
386    pub components: Vec<StableId>,
387    pub builder_roots: Vec<String>,
388    pub include_patterns: Vec<String>,
389    pub exclude_patterns: Vec<String>,
390    pub conditional: Option<String>,
391    pub source: String,
392    pub evidence: Vec<Evidence>,
393}
394
395#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
396#[serde(rename_all = "snake_case")]
397pub enum FindingSeverity {
398    Error,
399    Warning,
400    Information,
401}
402
403#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
404pub struct Finding {
405    pub id: StableId,
406    pub rule: String,
407    pub severity: FindingSeverity,
408    pub confidence: Confidence,
409    pub message: String,
410    pub source: Option<String>,
411    pub related: Vec<StableId>,
412    pub evidence: Vec<Evidence>,
413}
414
415#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize, Default)]
416pub struct ModelCoverage {
417    pub components: usize,
418    pub architecture_edges: usize,
419    pub modules: usize,
420    pub parameters: usize,
421    pub functions: usize,
422    pub entrypoints: usize,
423    pub component_entrypoints: usize,
424    pub composition_edges: usize,
425    pub assembly_sites: usize,
426    pub subprocess_stages: usize,
427    pub tensors: usize,
428    pub operations: usize,
429    pub pipeline_stages: usize,
430    pub artifacts: usize,
431    pub optimizer_memberships: usize,
432    pub linked_parameter_uses: usize,
433    pub tensors_with_shape: usize,
434    pub tensors_with_dtype: usize,
435    pub tensors_with_device: usize,
436    pub runtime_observations: usize,
437    pub diagnostics: usize,
438}
439
440#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
441pub struct CargoSummary {
442    /// Deterministic identity of the exact Cargo configuration used for this scan.
443    #[serde(default)]
444    pub build_id: String,
445    pub workspace_root: String,
446    pub manifest_path: String,
447    pub package_name: String,
448    pub package_version: String,
449    #[serde(default, skip_serializing_if = "Option::is_none")]
450    pub selected_target: Option<String>,
451    pub active_features: Vec<String>,
452    pub active_cfg: Vec<String>,
453    pub candle_packages: BTreeMap<String, String>,
454}
455
456#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
457pub struct EdgeTimingSummary {
458    pub from: StableId,
459    pub to: StableId,
460    pub avg_duration_ns: u64,
461    pub samples: u64,
462}
463
464#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
465pub struct RuntimeSummary {
466    pub trace_schema: String,
467    pub entrypoint: Option<String>,
468    pub profile: Option<String>,
469    pub tensor_observations: usize,
470    #[serde(default)]
471    pub operation_observations: usize,
472    pub gradient_observations: usize,
473    pub missing_gradients: usize,
474    pub zero_gradients: usize,
475    pub non_finite_gradients: usize,
476    #[serde(default)]
477    pub tensor_conflicts: usize,
478    #[serde(default)]
479    pub gradient_conflicts: usize,
480    #[serde(default)]
481    pub identity_mismatches: usize,
482    #[serde(default, skip_serializing_if = "Option::is_none")]
483    pub first_non_finite_step: Option<u64>,
484    #[serde(default)]
485    pub saturating_activations: usize,
486    #[serde(default)]
487    pub value_observations: usize,
488    #[serde(default, skip_serializing_if = "Option::is_none")]
489    pub execution_phase: Option<ExecutionPhase>,
490    #[serde(default, skip_serializing_if = "Option::is_none")]
491    pub avg_operation_duration_ns: Option<u64>,
492    #[serde(default)]
493    pub edge_timings: Vec<EdgeTimingSummary>,
494}
495
496#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
497pub struct ModelIr {
498    pub schema: String,
499    pub analysis_id: StableId,
500    pub cargo: Option<CargoSummary>,
501    pub coverage: ModelCoverage,
502    pub components: Vec<Component>,
503    pub architecture_edges: Vec<ArchitectureEdge>,
504    pub modules: Vec<Module>,
505    pub parameters: Vec<Parameter>,
506    pub functions: Vec<Function>,
507    pub tensors: Vec<TensorContract>,
508    pub operations: Vec<Operation>,
509    pub stages: Vec<PipelineStage>,
510    pub artifacts: Vec<Artifact>,
511    pub optimizers: Vec<OptimizerMembership>,
512    pub assembly_sites: Vec<AssemblySite>,
513    pub findings: Vec<Finding>,
514    pub runtime: Option<RuntimeSummary>,
515}
516
517impl ModelIr {
518    pub fn empty(analysis_id: StableId) -> Self {
519        Self {
520            schema: MODEL_IR_SCHEMA.to_string(),
521            analysis_id,
522            cargo: None,
523            coverage: ModelCoverage::default(),
524            components: Vec::new(),
525            architecture_edges: Vec::new(),
526            modules: Vec::new(),
527            parameters: Vec::new(),
528            functions: Vec::new(),
529            tensors: Vec::new(),
530            operations: Vec::new(),
531            stages: Vec::new(),
532            artifacts: Vec::new(),
533            optimizers: Vec::new(),
534            assembly_sites: Vec::new(),
535            findings: Vec::new(),
536            runtime: None,
537        }
538    }
539
540    pub fn normalize(&mut self) {
541        self.components.sort_by(|a, b| a.id.cmp(&b.id));
542        self.architecture_edges.sort_by(|a, b| a.id.cmp(&b.id));
543        self.modules.sort_by(|a, b| a.id.cmp(&b.id));
544        self.parameters.sort_by(|a, b| a.id.cmp(&b.id));
545        self.functions.sort_by(|a, b| a.id.cmp(&b.id));
546        self.tensors.sort_by(|a, b| a.id.cmp(&b.id));
547        self.operations.sort_by(|a, b| a.id.cmp(&b.id));
548        self.stages.sort_by(|a, b| {
549            (a.order.unwrap_or(usize::MAX), &a.id).cmp(&(b.order.unwrap_or(usize::MAX), &b.id))
550        });
551        self.artifacts.sort_by(|a, b| a.id.cmp(&b.id));
552        self.optimizers.sort_by(|a, b| a.id.cmp(&b.id));
553        self.findings.sort_by(|a, b| a.id.cmp(&b.id));
554        self.components.dedup_by(|a, b| a.id == b.id);
555        self.architecture_edges.dedup_by(|a, b| a.id == b.id);
556        self.modules.dedup_by(|a, b| a.id == b.id);
557        self.parameters.dedup_by(|a, b| a.id == b.id);
558        self.functions.dedup_by(|a, b| a.id == b.id);
559        self.tensors.dedup_by(|a, b| a.id == b.id);
560        self.operations.dedup_by(|a, b| a.id == b.id);
561        self.artifacts.dedup_by(|a, b| a.id == b.id);
562        self.optimizers.dedup_by(|a, b| a.id == b.id);
563        let mut merged_findings: Vec<Finding> = Vec::with_capacity(self.findings.len());
564        for finding in self.findings.drain(..) {
565            if let Some(existing) = merged_findings
566                .last_mut()
567                .filter(|existing| existing.id == finding.id)
568            {
569                existing.related.extend(finding.related);
570                existing.evidence.extend(finding.evidence);
571            } else {
572                merged_findings.push(finding);
573            }
574        }
575        self.findings = merged_findings;
576        for component in &mut self.components {
577            component.modules.sort();
578            component.modules.dedup();
579            component.parameters.sort();
580            component.parameters.dedup();
581            component.entrypoints.sort();
582            component.entrypoints.dedup();
583            component.builders.sort_by(|a, b| a.name.cmp(&b.name));
584        }
585        for function in &mut self.functions {
586            function.calls.sort();
587            function.calls.dedup();
588            function.tensor_inputs.sort();
589            function.tensor_inputs.dedup();
590            function.tensor_outputs.sort();
591            function.tensor_outputs.dedup();
592        }
593        for parameter in &mut self.parameters {
594            parameter.uses.sort();
595            parameter.uses.dedup();
596            parameter.optimizer_memberships.sort();
597            parameter.optimizer_memberships.dedup();
598        }
599        for finding in &mut self.findings {
600            finding.related.sort();
601            finding.related.dedup();
602            finding
603                .evidence
604                .sort_by(|a, b| (&a.source, &a.detail).cmp(&(&b.source, &b.detail)));
605            finding
606                .evidence
607                .dedup_by(|a, b| a.source == b.source && a.detail == b.detail);
608        }
609        self.refresh_coverage();
610    }
611
612    pub fn refresh_coverage(&mut self) {
613        self.coverage.components = self.components.len();
614        self.coverage.architecture_edges = self.architecture_edges.len();
615        self.coverage.modules = self.modules.len();
616        self.coverage.parameters = self.parameters.len();
617        self.coverage.functions = self.functions.len();
618        self.coverage.entrypoints = self.functions.iter().filter(|f| f.is_entrypoint).count();
619        let component_types: BTreeMap<String, ()> = self
620            .components
621            .iter()
622            .flat_map(|component| [component.name.clone(), component.qualified_name.clone()])
623            .map(|name| (name, ()))
624            .collect();
625        self.coverage.component_entrypoints = self
626            .functions
627            .iter()
628            .filter(|function| {
629                function.is_entrypoint
630                    && function
631                        .owner_type
632                        .as_ref()
633                        .is_some_and(|owner| component_types.contains_key(owner))
634            })
635            .count();
636        self.coverage.composition_edges = self
637            .architecture_edges
638            .iter()
639            .filter(|edge| edge.id.0.starts_with("composition-edge:"))
640            .count();
641        self.coverage.assembly_sites = self.assembly_sites.len();
642        self.coverage.subprocess_stages = self
643            .stages
644            .iter()
645            .filter(|stage| stage.dispatch == StageDispatchKind::Subprocess)
646            .count();
647        self.coverage.tensors = self.tensors.len();
648        self.coverage.operations = self.operations.len();
649        self.coverage.pipeline_stages = self.stages.len();
650        self.coverage.artifacts = self.artifacts.len();
651        self.coverage.optimizer_memberships = self.optimizers.len();
652        self.coverage.linked_parameter_uses =
653            self.parameters.iter().map(|p| p.uses.len()).sum::<usize>();
654        self.coverage.tensors_with_shape = self
655            .tensors
656            .iter()
657            .filter(|t| t.shape.rank.is_some() || !t.shape.dimensions.is_empty())
658            .count();
659        self.coverage.tensors_with_dtype =
660            self.tensors.iter().filter(|t| t.dtype != "Unknown").count();
661        self.coverage.tensors_with_device = self
662            .tensors
663            .iter()
664            .filter(|t| !matches!(t.device, DeviceFact::Unknown))
665            .count();
666        self.coverage.diagnostics = self.findings.len();
667    }
668}
669
670fn escape_id_part(part: &str, out: &mut String) {
671    for byte in part.bytes() {
672        match byte {
673            b'%' | b':' | b'/' | b'\\' | b' ' => {
674                out.push('%');
675                out.push_str(&format!("{byte:02X}"));
676            }
677            _ => out.push(byte as char),
678        }
679    }
680}