1use std::collections::BTreeMap;
8
9use serde::{Deserialize, Serialize};
10
11pub 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 pub name: Option<String>,
132 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 pub cfg_predicates: Vec<String>,
233 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 #[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 #[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}