1use serde::{Deserialize, Serialize};
4
5#[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#[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}