Skip to main content

candle_graph/
phase.rs

1//! Train vs inference execution phase for static graphs and runtime traces.
2
3use serde::{Deserialize, Serialize};
4
5/// Whether analysis or profiling targets training (autograd) or inference (no-grad).
6#[derive(Debug, Clone, Copy, PartialEq, Eq, PartialOrd, Ord, Hash, Serialize, Deserialize)]
7#[serde(rename_all = "snake_case")]
8pub enum ExecutionPhase {
9    Train,
10    Infer,
11}
12
13impl ExecutionPhase {
14    pub fn as_str(self) -> &'static str {
15        match self {
16            Self::Train => "train",
17            Self::Infer => "infer",
18        }
19    }
20
21    pub fn parse(value: &str) -> Option<Self> {
22        match value.replace('-', "_").to_ascii_lowercase().as_str() {
23            "train" | "training" => Some(Self::Train),
24            "infer" | "inference" | "eval" | "evaluate" => Some(Self::Infer),
25            _ => None,
26        }
27    }
28}
29
30/// Classify which static graphs to build for an entrypoint.
31pub fn entrypoint_phases(name: &str, qualified_name: &str, is_loss: bool) -> Vec<ExecutionPhase> {
32    if is_loss {
33        return vec![ExecutionPhase::Train];
34    }
35    let lower = format!("{name} {qualified_name}").to_ascii_lowercase();
36    let infer_hint = lower.contains("eval")
37        || lower.contains("infer")
38        || lower.contains("predict")
39        || lower.contains("inference");
40    let train_hint = lower.contains("train")
41        || lower.contains("loss")
42        || lower.contains("backward")
43        || lower.contains("optim");
44    if name == "forward" || name == "forward_t" {
45        return vec![ExecutionPhase::Train, ExecutionPhase::Infer];
46    }
47    if infer_hint && train_hint {
48        return vec![ExecutionPhase::Train, ExecutionPhase::Infer];
49    }
50    if infer_hint {
51        return vec![ExecutionPhase::Infer];
52    }
53    if train_hint {
54        return vec![ExecutionPhase::Train];
55    }
56    vec![ExecutionPhase::Train]
57}