use std::fmt;
use std::str::FromStr;
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, Default)]
pub enum Task {
#[default]
Detect,
Segment,
Pose,
Classify,
Obb,
Semantic,
Depth,
}
impl Task {
#[must_use]
pub const fn as_str(&self) -> &'static str {
match self {
Self::Detect => "detect",
Self::Segment => "segment",
Self::Pose => "pose",
Self::Classify => "classify",
Self::Obb => "obb",
Self::Semantic => "semantic",
Self::Depth => "depth",
}
}
#[must_use]
pub const fn model_suffix(&self) -> &'static str {
match self {
Self::Detect => "",
Self::Segment => "-seg",
Self::Pose => "-pose",
Self::Classify => "-cls",
Self::Obb => "-obb",
Self::Semantic => "-sem",
Self::Depth => "-depth",
}
}
#[must_use]
pub fn default_model(&self) -> String {
format!("yolo26n{}.onnx", self.model_suffix())
}
#[must_use]
pub const fn has_boxes(&self) -> bool {
matches!(self, Self::Detect | Self::Segment | Self::Pose | Self::Obb)
}
#[must_use]
pub const fn has_masks(&self) -> bool {
matches!(self, Self::Segment)
}
#[must_use]
pub const fn has_keypoints(&self) -> bool {
matches!(self, Self::Pose)
}
#[must_use]
pub const fn has_probs(&self) -> bool {
matches!(self, Self::Classify)
}
#[must_use]
pub const fn has_obb(&self) -> bool {
matches!(self, Self::Obb)
}
#[must_use]
pub const fn has_semantic_mask(&self) -> bool {
matches!(self, Self::Semantic)
}
#[must_use]
pub const fn has_depth(&self) -> bool {
matches!(self, Self::Depth)
}
}
impl fmt::Display for Task {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.write_str(self.as_str())
}
}
impl FromStr for Task {
type Err = TaskParseError;
fn from_str(s: &str) -> Result<Self, Self::Err> {
match s.to_lowercase().as_str() {
"detect" | "detection" => Ok(Self::Detect),
"segment" | "segmentation" => Ok(Self::Segment),
"pose" | "keypoint" | "keypoints" => Ok(Self::Pose),
"classify" | "classification" | "cls" => Ok(Self::Classify),
"obb" | "oriented" => Ok(Self::Obb),
"semantic" | "semantic_segmentation" | "semseg" => Ok(Self::Semantic),
"depth" | "depth_estimation" => Ok(Self::Depth),
_ => Err(TaskParseError(s.to_string())),
}
}
}
#[derive(Debug, Clone)]
pub struct TaskParseError(String);
impl fmt::Display for TaskParseError {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
write!(
f,
"invalid task '{}', expected one of: detect, segment, pose, classify, obb, semantic, depth",
self.0
)
}
}
impl std::error::Error for TaskParseError {}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_task_from_str() {
assert_eq!("detect".parse::<Task>().unwrap(), Task::Detect);
assert_eq!("segment".parse::<Task>().unwrap(), Task::Segment);
assert_eq!("pose".parse::<Task>().unwrap(), Task::Pose);
assert_eq!("classify".parse::<Task>().unwrap(), Task::Classify);
assert_eq!("obb".parse::<Task>().unwrap(), Task::Obb);
assert_eq!("semantic".parse::<Task>().unwrap(), Task::Semantic);
assert_eq!("depth".parse::<Task>().unwrap(), Task::Depth);
assert_eq!("detection".parse::<Task>().unwrap(), Task::Detect);
assert_eq!("segmentation".parse::<Task>().unwrap(), Task::Segment);
assert_eq!("keypoints".parse::<Task>().unwrap(), Task::Pose);
assert_eq!("cls".parse::<Task>().unwrap(), Task::Classify);
assert_eq!(
"semantic_segmentation".parse::<Task>().unwrap(),
Task::Semantic
);
assert_eq!("depth_estimation".parse::<Task>().unwrap(), Task::Depth);
}
#[test]
fn test_task_display() {
assert_eq!(Task::Detect.to_string(), "detect");
assert_eq!(Task::Segment.to_string(), "segment");
assert_eq!(Task::Semantic.to_string(), "semantic");
}
#[test]
fn test_task_capabilities() {
assert!(Task::Detect.has_boxes());
assert!(!Task::Detect.has_masks());
assert!(Task::Segment.has_masks());
assert!(Task::Pose.has_keypoints());
assert!(Task::Classify.has_probs());
assert!(Task::Obb.has_obb());
assert!(Task::Semantic.has_semantic_mask());
assert!(!Task::Detect.has_semantic_mask());
assert!(Task::Depth.has_depth());
assert!(!Task::Semantic.has_depth());
assert!(!Task::Depth.has_semantic_mask());
}
#[test]
fn test_task_suffix_and_default_model() {
let cases = [
(Task::Detect, "detect", "", "yolo26n.onnx"),
(Task::Segment, "segment", "-seg", "yolo26n-seg.onnx"),
(Task::Pose, "pose", "-pose", "yolo26n-pose.onnx"),
(Task::Classify, "classify", "-cls", "yolo26n-cls.onnx"),
(Task::Obb, "obb", "-obb", "yolo26n-obb.onnx"),
(Task::Semantic, "semantic", "-sem", "yolo26n-sem.onnx"),
(Task::Depth, "depth", "-depth", "yolo26n-depth.onnx"),
];
for (task, name, suffix, model) in cases {
assert_eq!(task.as_str(), name);
assert_eq!(task.model_suffix(), suffix);
assert_eq!(task.default_model(), model);
}
}
#[test]
fn test_task_from_str_aliases_and_errors() {
assert_eq!("KEYPOINT".parse::<Task>().unwrap(), Task::Pose);
assert_eq!("oriented".parse::<Task>().unwrap(), Task::Obb);
assert_eq!("semseg".parse::<Task>().unwrap(), Task::Semantic);
assert_eq!("Classification".parse::<Task>().unwrap(), Task::Classify);
assert!("not_a_task".parse::<Task>().is_err());
}
}