Skip to main content

candle_graph/
phase.rs

1//! Train vs inference execution phase for trace metadata.
2
3use serde::{Deserialize, Serialize};
4
5/// Whether 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/// Training-step slice inside a probe run (PyTorch profiler timeline phases).
31#[derive(Debug, Clone, Copy, PartialEq, Eq, PartialOrd, Ord, Hash, Serialize, Deserialize)]
32#[serde(rename_all = "snake_case")]
33pub enum ExecutionStep {
34    Forward,
35    Backward,
36    Optimizer,
37}
38
39impl ExecutionStep {
40    pub fn as_str(self) -> &'static str {
41        match self {
42            Self::Forward => "forward",
43            Self::Backward => "backward",
44            Self::Optimizer => "optimizer",
45        }
46    }
47
48    pub fn parse(value: &str) -> Option<Self> {
49        match value.replace('-', "_").to_ascii_lowercase().as_str() {
50            "forward" | "fwd" => Some(Self::Forward),
51            "backward" | "backward_pass" | "bwd" | "backprop" => Some(Self::Backward),
52            "optimizer" | "optim" | "step" => Some(Self::Optimizer),
53            _ => None,
54        }
55    }
56}