1use std::collections::BTreeMap;
8
9use serde::{Deserialize, Serialize};
10
11pub use crate::phase::ExecutionPhase;
12
13pub 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 pub name: Option<String>,
134 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 #[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 pub cfg_predicates: Vec<String>,
238 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 #[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 #[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 #[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}