use std::collections::HashMap;
use std::sync::Arc;
use crate::error::{InferenceError, Result};
use crate::inference::Quantization;
use crate::task::Task;
#[derive(Debug, Clone)]
pub struct ModelMetadata {
pub description: String,
pub author: String,
pub date: String,
pub version: String,
pub license: String,
pub docs: String,
pub task: Task,
pub stride: u32,
pub batch: usize,
pub imgsz: Option<(usize, usize)>,
pub channels: usize,
pub quantize: Option<Quantization>,
#[doc(hidden)]
pub half: bool,
pub names: Arc<HashMap<usize, String>>,
pub end2end: bool,
pub kpt_shape: Option<(usize, usize)>,
}
impl ModelMetadata {
pub fn from_onnx_metadata(metadata_map: &HashMap<String, String>) -> Result<Self> {
let yaml_str = metadata_map
.get("metadata")
.or_else(|| metadata_map.get("model_metadata"))
.or_else(|| {
metadata_map.values().find(|v| v.contains("task:"))
})
.ok_or_else(|| {
InferenceError::ModelLoadError(
"No metadata found in ONNX model. Ensure the model was exported with Ultralytics.".to_string()
)
})?;
Self::from_yaml_str(yaml_str)
}
pub fn from_yaml_str(yaml_str: &str) -> Result<Self> {
let mut metadata = Self::default();
let mut names: HashMap<usize, String> = HashMap::new();
let mut quantize_set = false;
let mut half = None;
for line in yaml_str.lines() {
let line = line.trim();
if line.is_empty() || line.starts_with('#') {
continue;
}
if let Some((key, value)) = line.split_once(':') {
let key = key.trim();
let value = value.trim().trim_matches('\'').trim_matches('"');
match key {
"description" => metadata.description = value.to_string(),
"author" => metadata.author = value.to_string(),
"date" => metadata.date = value.to_string(),
"version" => metadata.version = value.to_string(),
"license" => metadata.license = value.to_string(),
"docs" => metadata.docs = value.to_string(),
"task" => {
metadata.task = value.parse().map_err(|e| {
InferenceError::ModelLoadError(format!("Invalid task in metadata: {e}"))
})?;
}
"stride" => {
metadata.stride = value.parse().map_err(|_| {
InferenceError::ModelLoadError(format!("Invalid stride value: {value}"))
})?;
}
"batch" => {
metadata.batch = value.parse().map_err(|_| {
InferenceError::ModelLoadError(format!("Invalid batch value: {value}"))
})?;
}
"channels" => {
metadata.channels = value.parse().map_err(|_| {
InferenceError::ModelLoadError(format!(
"Invalid channels value: {value}"
))
})?;
}
"quantize" => {
metadata.quantize = Self::parse_quantize(value)?;
quantize_set = true;
}
"half" => {
half = Some(Self::parse_bool(value));
}
"end2end" => {
metadata.end2end = value == "true" || value == "True";
}
"args" => {
if !quantize_set
&& let Some(value) = Self::parse_args_value(value, "quantize")
{
metadata.quantize = Self::parse_quantize(value)?;
quantize_set = true;
}
if half.is_none()
&& let Some(value) = Self::parse_args_value(value, "half")
{
half = Some(Self::parse_bool(value));
}
}
_ => {
if let Ok(class_id) = key.trim().parse::<usize>() {
names.insert(class_id, value.to_string());
}
}
}
}
}
if !quantize_set {
metadata.quantize = half.unwrap_or(false).then_some(Quantization::Fp16);
}
metadata.half = metadata.quantize == Some(Quantization::Fp16);
metadata.imgsz = Self::parse_int_pair(yaml_str, "imgsz");
metadata.kpt_shape = Self::parse_int_pair(yaml_str, "kpt_shape");
if names.is_empty() {
names = Self::parse_names_block(yaml_str);
}
metadata.names = Arc::new(names);
Ok(metadata)
}
fn parse_quantize(value: &str) -> Result<Option<Quantization>> {
let value = value.trim().trim_matches(|c| c == '\'' || c == '"');
if value.eq_ignore_ascii_case("none")
|| value.eq_ignore_ascii_case("null")
|| value.is_empty()
{
return Ok(None);
}
value.parse().map(Some).map_err(|e| {
InferenceError::ModelLoadError(format!("Invalid quantize value in metadata: {e}"))
})
}
fn parse_args_value<'a>(args: &'a str, name: &str) -> Option<&'a str> {
let single_quoted = format!("'{name}'");
let double_quoted = format!("\"{name}\"");
let (_, value) = [single_quoted.as_str(), double_quoted.as_str()]
.iter()
.find_map(|key| args.split_once(key))?;
Some(
value
.trim_start()
.strip_prefix(':')
.unwrap_or(value)
.split([',', '}'])
.next()
.unwrap_or_default(),
)
}
fn parse_bool(value: &str) -> bool {
!matches!(
value
.trim()
.trim_matches(|c| c == '\'' || c == '"')
.to_ascii_lowercase()
.as_str(),
"none" | "false" | "0" | ""
)
}
fn parse_int_pair(yaml_str: &str, key: &str) -> Option<(usize, usize)> {
let prefix = format!("{key}:");
let lines: Vec<&str> = yaml_str.lines().collect();
for (i, line) in lines.iter().enumerate() {
let Some(rest) = line.trim_start().strip_prefix(&prefix) else {
continue;
};
let rest = rest
.trim()
.trim_matches(|c| matches!(c, '[' | ']' | '(' | ')'));
let inline: Vec<usize> = rest
.split(',')
.filter_map(|s| s.trim().parse().ok())
.collect();
if inline.len() >= 2 {
return Some((inline[0], inline[1]));
}
let mut vals = Vec::new();
for following in lines.iter().skip(i + 1) {
let t = following.trim();
if let Some(v) = t.strip_prefix('-') {
if let Ok(n) = v.trim().parse::<usize>() {
vals.push(n);
}
} else if !t.is_empty() && !t.starts_with('#') {
break;
}
if vals.len() >= 2 {
break;
}
}
if vals.len() >= 2 {
return Some((vals[0], vals[1]));
}
}
None
}
fn parse_names_block(yaml_str: &str) -> HashMap<usize, String> {
let mut names = HashMap::new();
if let Some(start) = yaml_str.find("names:") {
let after_names = &yaml_str[start + 6..];
let trimmed = after_names.trim();
if trimmed.starts_with('{')
&& let Some(end) = trimmed.find('}')
{
let dict_str = &trimmed[1..end];
return Self::parse_python_dict(dict_str);
}
}
let lines: Vec<&str> = yaml_str.lines().collect();
let mut in_names_block = false;
let mut names_indent = 0;
for line in &lines {
let trimmed = line.trim();
if trimmed.starts_with("names:") {
in_names_block = true;
names_indent = line.len() - line.trim_start().len();
continue;
}
if in_names_block {
let current_indent = line.len() - line.trim_start().len();
if !trimmed.is_empty()
&& !trimmed.starts_with('#')
&& current_indent <= names_indent
{
if !trimmed.chars().next().is_some_and(|c| c.is_ascii_digit()) {
break;
}
}
if let Some((key, value)) = trimmed.split_once(':')
&& let Ok(class_id) = key.trim().parse::<usize>()
{
let class_name = value.trim().trim_matches('\'').trim_matches('"');
names.insert(class_id, class_name.to_string());
}
}
}
names
}
fn parse_python_dict(dict_str: &str) -> HashMap<usize, String> {
let mut names = HashMap::new();
for entry in dict_str.split(',') {
let entry = entry.trim();
if let Some((key, value)) = entry.split_once(':') {
let key = key.trim();
let value = value.trim().trim_matches('\'').trim_matches('"');
if let Ok(class_id) = key.parse::<usize>() {
names.insert(class_id, value.to_string());
}
}
}
names
}
#[must_use]
pub fn num_classes(&self) -> usize {
self.names.len()
}
#[must_use]
pub fn class_name(&self, class_id: usize) -> Option<&str> {
self.names.get(&class_id).map(String::as_str)
}
#[must_use]
pub fn model_name(&self) -> String {
self.description
.split_whitespace()
.find(|&word| word.to_lowercase().starts_with("yolo"))
.unwrap_or("YOLO")
.to_string()
}
}
impl Default for ModelMetadata {
fn default() -> Self {
Self {
description: String::new(),
author: "Ultralytics".to_string(),
date: String::new(),
version: String::new(),
license: "AGPL-3.0".to_string(),
docs: "https://docs.ultralytics.com".to_string(),
task: Task::Detect,
stride: 32,
batch: 1,
imgsz: None,
channels: 3,
quantize: None,
half: false,
names: Arc::new(HashMap::new()),
end2end: false,
kpt_shape: None,
}
}
}
#[cfg(test)]
mod tests {
use super::*;
const SAMPLE_METADATA: &str = r"
description: Ultralytics YOLO11n model trained on /usr/src/ultralytics/ultralytics/cfg/datasets/coco.yaml
author: Ultralytics
date: '2025-12-11T20:19:45.464021'
version: 8.3.236
license: AGPL-3.0 License (https://ultralytics.com/license)
docs: https://docs.ultralytics.com
stride: 32
task: detect
batch: 1
imgsz:
- 640
- 640
names:
0: person
1: bicycle
2: car
3: motorcycle
channels: 3
";
#[test]
fn test_parse_metadata() {
let metadata = ModelMetadata::from_yaml_str(SAMPLE_METADATA).unwrap();
assert_eq!(metadata.task, Task::Detect);
assert_eq!(metadata.stride, 32);
assert_eq!(metadata.batch, 1);
assert_eq!(metadata.imgsz, Some((640, 640)));
assert_eq!(metadata.channels, 3);
assert_eq!(metadata.num_classes(), 4);
assert_eq!(metadata.class_name(0), Some("person"));
assert_eq!(metadata.class_name(1), Some("bicycle"));
assert_eq!(metadata.class_name(2), Some("car"));
assert_eq!(metadata.class_name(3), Some("motorcycle"));
}
#[test]
fn test_parse_inline_imgsz() {
let yaml = "task: detect\nimgsz: [640, 640]\nstride: 32";
let metadata = ModelMetadata::from_yaml_str(yaml).unwrap();
assert_eq!(metadata.imgsz, Some((640, 640)));
}
#[test]
fn test_default_metadata() {
let metadata = ModelMetadata::default();
assert_eq!(metadata.task, Task::Detect);
assert_eq!(metadata.stride, 32);
assert_eq!(metadata.imgsz, None);
}
#[test]
fn test_parse_multiline_kpt_shape_and_imgsz() {
let yaml = "task: pose\nstride: 32\nimgsz:\n- 640\n- 640\nkpt_shape:\n- 17\n- 3\n";
let metadata = ModelMetadata::from_yaml_str(yaml).unwrap();
assert_eq!(metadata.task, Task::Pose);
assert_eq!(metadata.imgsz, Some((640, 640)));
assert_eq!(metadata.kpt_shape, Some((17, 3)));
}
#[test]
fn test_parse_inline_kpt_shape() {
let yaml = "task: pose\nkpt_shape: [17, 3]\nstride: 32";
let metadata = ModelMetadata::from_yaml_str(yaml).unwrap();
assert_eq!(metadata.kpt_shape, Some((17, 3)));
}
#[test]
fn test_parse_int_pair_inline_and_block() {
assert_eq!(
ModelMetadata::parse_int_pair("kpt_shape: [17, 3]", "kpt_shape"),
Some((17, 3))
);
assert_eq!(
ModelMetadata::parse_int_pair("kpt_shape: (17, 2)", "kpt_shape"),
Some((17, 2))
);
assert_eq!(
ModelMetadata::parse_int_pair("foo: 5, 6", "foo"),
Some((5, 6))
);
assert_eq!(
ModelMetadata::parse_int_pair("kpt_shape:\n- 17\n- 3", "kpt_shape"),
Some((17, 3))
);
assert_eq!(
ModelMetadata::parse_int_pair("kpt_shape: [17]", "kpt_shape"),
None
);
assert_eq!(ModelMetadata::parse_int_pair("foo: 5", "foo"), None);
assert_eq!(ModelMetadata::parse_int_pair("other: 1, 2", "foo"), None);
}
#[test]
fn test_python_dict_names() {
let yaml = "task: detect\nnames: {0: 'person', 1: 'bicycle', 2: 'car'}";
let m = ModelMetadata::from_yaml_str(yaml).unwrap();
assert_eq!(m.num_classes(), 3);
assert_eq!(m.class_name(1), Some("bicycle"));
}
#[test]
fn test_precision_metadata_compatibility() {
let metadata = ModelMetadata::from_yaml_str("args: {'half': True}").unwrap();
assert_eq!(metadata.quantize, Some(Quantization::Fp16));
assert!(metadata.half);
let metadata = ModelMetadata::from_yaml_str(
"quantize: 32\nhalf: true\nargs: {'quantize': 16, 'half': True}",
)
.unwrap();
assert_eq!(metadata.quantize, Some(Quantization::Fp32));
assert!(!metadata.half);
}
#[test]
fn test_tflite_style_metadata_text() {
let text = "task: detect\nimgsz: [640, 640]\nstride: 32\nend2end: false\nnames:\n 0: person\n 1: traffic light\n 2: men's shoe";
let m = ModelMetadata::from_yaml_str(text).unwrap();
assert_eq!(m.task, Task::Detect);
assert_eq!(m.imgsz, Some((640, 640)));
assert_eq!(m.num_classes(), 3);
assert_eq!(m.class_name(1), Some("traffic light"));
assert_eq!(m.class_name(2), Some("men's shoe"));
}
#[test]
fn test_names_block_with_yaml_significant_chars() {
let text = "task: detect\nnames:\n 0: ratio 16:9\n 1: #1 pick";
let m = ModelMetadata::from_yaml_str(text).unwrap();
assert_eq!(m.class_name(0), Some("ratio 16:9"));
assert_eq!(m.class_name(1), Some("#1 pick"));
}
#[test]
fn test_pose_kpt_shape_and_quantize() {
let yaml = "task: pose\nkpt_shape: [17, 3]\nquantize: 16\nend2end: True";
let m = ModelMetadata::from_yaml_str(yaml).unwrap();
assert_eq!(m.task, Task::Pose);
assert_eq!(m.kpt_shape, Some((17, 3)));
assert_eq!(m.quantize, Some(Quantization::Fp16));
assert!(m.end2end);
let yaml2 = "task: detect\nargs: {'quantize': 'w8a16', 'imgsz': 640}";
let m2 = ModelMetadata::from_yaml_str(yaml2).unwrap();
assert_eq!(m2.quantize, Some(Quantization::W8a16));
let yaml3 = "task: detect\nargs: {'quantize': None, 'imgsz': 640}";
let m3 = ModelMetadata::from_yaml_str(yaml3).unwrap();
assert_eq!(m3.quantize, None);
}
#[test]
fn test_invalid_fields_error() {
assert!(ModelMetadata::from_yaml_str("task: notarealtask").is_err());
assert!(ModelMetadata::from_yaml_str("task: detect\nstride: abc").is_err());
assert!(ModelMetadata::from_yaml_str("task: detect\nbatch: xyz").is_err());
}
#[test]
fn test_from_onnx_metadata_keys_and_error() {
let map = HashMap::from([("metadata".to_string(), SAMPLE_METADATA.to_string())]);
let m = ModelMetadata::from_onnx_metadata(&map).unwrap();
assert_eq!(m.task, Task::Detect);
let map = HashMap::from([("whatever".to_string(), "task: pose\nstride: 32".to_string())]);
let m = ModelMetadata::from_onnx_metadata(&map).unwrap();
assert_eq!(m.task, Task::Pose);
let map = HashMap::from([("unrelated".to_string(), "no yaml here".to_string())]);
assert!(ModelMetadata::from_onnx_metadata(&map).is_err());
}
#[test]
fn test_model_name_extraction() {
let m = ModelMetadata::from_yaml_str(SAMPLE_METADATA).unwrap();
assert_eq!(m.model_name(), "YOLO11n");
let plain = ModelMetadata::default();
assert_eq!(plain.model_name(), "YOLO");
}
}