use std::path::Path;
use arcstr::ArcStr;
use crate::error::YoloError;
#[derive(Debug)]
pub struct YoloModelSession {
pub session: ort::session::Session,
pub labels: Vec<ArcStr>,
pub probability_threshold: Option<f32>, pub iou_threshold: Option<f32>, }
impl YoloModelSession {
pub fn new(
session: ort::session::Session,
labels: impl Iterator<Item = impl Into<ArcStr>>,
) -> Self {
Self {
session,
labels: labels.map(Into::into).collect(),
probability_threshold: None,
iou_threshold: None,
}
}
pub fn new_v8(session: ort::session::Session) -> Self {
const LABELS: &[ArcStr] = &[
arcstr::literal!("person"),
arcstr::literal!("bicycle"),
arcstr::literal!("car"),
arcstr::literal!("motorcycle"),
arcstr::literal!("airplane"),
arcstr::literal!("bus"),
arcstr::literal!("train"),
arcstr::literal!("truck"),
arcstr::literal!("boat"),
arcstr::literal!("traffic light"),
arcstr::literal!("fire hydrant"),
arcstr::literal!("stop sign"),
arcstr::literal!("parking meter"),
arcstr::literal!("bench"),
arcstr::literal!("bird"),
arcstr::literal!("cat"),
arcstr::literal!("dog"),
arcstr::literal!("horse"),
arcstr::literal!("sheep"),
arcstr::literal!("cow"),
arcstr::literal!("elephant"),
arcstr::literal!("bear"),
arcstr::literal!("zebra"),
arcstr::literal!("giraffe"),
arcstr::literal!("backpack"),
arcstr::literal!("umbrella"),
arcstr::literal!("handbag"),
arcstr::literal!("tie"),
arcstr::literal!("suitcase"),
arcstr::literal!("frisbee"),
arcstr::literal!("skis"),
arcstr::literal!("snowboard"),
arcstr::literal!("sports ball"),
arcstr::literal!("kite"),
arcstr::literal!("baseball bat"),
arcstr::literal!("baseball glove"),
arcstr::literal!("skateboard"),
arcstr::literal!("surfboard"),
arcstr::literal!("tennis racket"),
arcstr::literal!("bottle"),
arcstr::literal!("wine glass"),
arcstr::literal!("cup"),
arcstr::literal!("fork"),
arcstr::literal!("knife"),
arcstr::literal!("spoon"),
arcstr::literal!("bowl"),
arcstr::literal!("banana"),
arcstr::literal!("apple"),
arcstr::literal!("sandwich"),
arcstr::literal!("orange"),
arcstr::literal!("broccoli"),
arcstr::literal!("carrot"),
arcstr::literal!("hot dog"),
arcstr::literal!("pizza"),
arcstr::literal!("donut"),
arcstr::literal!("cake"),
arcstr::literal!("chair"),
arcstr::literal!("couch"),
arcstr::literal!("potted plant"),
arcstr::literal!("bed"),
arcstr::literal!("dining table"),
arcstr::literal!("toilet"),
arcstr::literal!("tv"),
arcstr::literal!("laptop"),
arcstr::literal!("mouse"),
arcstr::literal!("remote"),
arcstr::literal!("keyboard"),
arcstr::literal!("cell phone"),
arcstr::literal!("microwave"),
arcstr::literal!("oven"),
arcstr::literal!("toaster"),
arcstr::literal!("sink"),
arcstr::literal!("refrigerator"),
arcstr::literal!("book"),
arcstr::literal!("clock"),
arcstr::literal!("vase"),
arcstr::literal!("scissors"),
arcstr::literal!("teddy bear"),
arcstr::literal!("hair drier"),
arcstr::literal!("toothbrush"),
];
Self {
session,
labels: LABELS.to_vec(),
probability_threshold: None,
iou_threshold: None,
}
}
pub fn from_filename_v8(filename: impl AsRef<Path>) -> Result<Self, YoloError> {
let session = ort::session::Session::builder()
.map_err(YoloError::OrtSessionBuildError)?
.commit_from_file(filename)
.map_err(YoloError::OrtSessionLoadError)?;
Ok(Self::new_v8(session))
}
pub fn get_labels(&self) -> &[ArcStr] {
&self.labels
}
pub fn get_probability_threshold(&self) -> f32 {
self.probability_threshold.unwrap_or(0.5)
}
pub fn get_iou_threshold(&self) -> f32 {
self.iou_threshold.unwrap_or(0.7)
}
}
impl AsRef<ort::session::Session> for YoloModelSession {
fn as_ref(&self) -> &ort::session::Session {
&self.session
}
}
impl AsMut<ort::session::Session> for YoloModelSession {
fn as_mut(&mut self) -> &mut ort::session::Session {
&mut self.session
}
}