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