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
30pub 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}