use std::collections::HashMap;
use std::marker::PhantomData;
use serde::de::{self, Visitor};
use serde::{Deserialize, Deserializer, Serialize};
#[derive(Debug, Clone, Default, Deserialize, Serialize)]
pub struct Dataset {
#[serde(default)]
pub info: Option<Info>,
#[serde(default)]
pub images: Vec<Image>,
#[serde(default)]
pub annotations: Vec<Annotation>,
#[serde(default)]
pub categories: Vec<Category>,
#[serde(default)]
pub licenses: Vec<License>,
}
#[derive(Debug, Clone, Deserialize, Serialize)]
pub struct Info {
#[serde(default)]
pub year: Option<u32>,
#[serde(default)]
pub version: Option<String>,
#[serde(default)]
pub description: Option<String>,
#[serde(default)]
pub contributor: Option<String>,
#[serde(default)]
pub url: Option<String>,
#[serde(default)]
pub date_created: Option<String>,
}
#[derive(Debug, Clone, Default, Deserialize, Serialize)]
pub struct Image {
#[serde(deserialize_with = "deserialize_uint")]
pub id: u64,
#[serde(default)]
pub file_name: String,
#[serde(default, deserialize_with = "deserialize_uint")]
pub height: u32,
#[serde(default, deserialize_with = "deserialize_uint")]
pub width: u32,
#[serde(default, deserialize_with = "deserialize_opt_uint")]
pub license: Option<u64>,
#[serde(default)]
pub coco_url: Option<String>,
#[serde(default)]
pub flickr_url: Option<String>,
#[serde(default)]
pub date_captured: Option<String>,
#[serde(default)]
pub neg_category_ids: Vec<u64>,
#[serde(default)]
pub not_exhaustive_category_ids: Vec<u64>,
#[serde(flatten, skip_serializing_if = "serde_json::Map::is_empty")]
pub extra: serde_json::Map<String, serde_json::Value>,
}
#[derive(Debug, Clone, Default, Deserialize, Serialize)]
pub struct Annotation {
#[serde(default, deserialize_with = "deserialize_uint")]
pub id: u64,
#[serde(deserialize_with = "deserialize_uint")]
pub image_id: u64,
#[serde(deserialize_with = "deserialize_uint")]
pub category_id: u64,
#[serde(default)]
pub bbox: Option<[f64; 4]>,
#[serde(default)]
pub area: Option<f64>,
#[serde(default)]
pub segmentation: Option<Segmentation>,
#[serde(default, deserialize_with = "deserialize_flag")]
pub iscrowd: bool,
#[serde(default)]
pub keypoints: Option<Vec<f64>>,
#[serde(default, deserialize_with = "deserialize_opt_uint")]
pub num_keypoints: Option<u32>,
#[serde(default)]
pub obb: Option<[f64; 5]>,
#[serde(default)]
pub score: Option<f64>,
#[serde(default, deserialize_with = "deserialize_opt_flag")]
pub is_group_of: Option<bool>,
#[serde(flatten, skip_serializing_if = "serde_json::Map::is_empty")]
pub extra: serde_json::Map<String, serde_json::Value>,
}
impl Annotation {
pub fn num_visible_keypoints(&self) -> u32 {
match (self.num_keypoints, &self.keypoints) {
(Some(n), _) => n,
(None, Some(k)) => k.chunks_exact(3).filter(|t| t[2] > 0.0).count() as u32,
(None, None) => 0,
}
}
}
struct Flag(bool);
impl<'de> Deserialize<'de> for Flag {
fn deserialize<D: Deserializer<'de>>(deserializer: D) -> Result<Self, D::Error> {
struct FlagVisitor;
impl Visitor<'_> for FlagVisitor {
type Value = Flag;
fn expecting(&self, f: &mut std::fmt::Formatter) -> std::fmt::Result {
f.write_str("a bool or a 0/1 flag")
}
fn visit_bool<E: de::Error>(self, v: bool) -> Result<Flag, E> {
Ok(Flag(v))
}
fn visit_u64<E: de::Error>(self, v: u64) -> Result<Flag, E> {
Ok(Flag(v != 0))
}
fn visit_i64<E: de::Error>(self, v: i64) -> Result<Flag, E> {
Ok(Flag(v != 0))
}
fn visit_f64<E: de::Error>(self, v: f64) -> Result<Flag, E> {
if v.fract() == 0.0 {
Ok(Flag(v != 0.0))
} else {
Err(E::custom(format!("expected a bool or 0/1 flag, got {v}")))
}
}
}
deserializer.deserialize_any(FlagVisitor)
}
}
fn deserialize_flag<'de, D: Deserializer<'de>>(deserializer: D) -> Result<bool, D::Error> {
Flag::deserialize(deserializer).map(|f| f.0)
}
fn deserialize_opt_flag<'de, D: Deserializer<'de>>(
deserializer: D,
) -> Result<Option<bool>, D::Error> {
Option::<Flag>::deserialize(deserializer).map(|o| o.map(|f| f.0))
}
struct Uint<T>(T);
impl<'de, T: TryFrom<u64>> Deserialize<'de> for Uint<T> {
fn deserialize<D: Deserializer<'de>>(deserializer: D) -> Result<Self, D::Error> {
struct UintVisitor<T>(PhantomData<T>);
impl<T: TryFrom<u64>> Visitor<'_> for UintVisitor<T> {
type Value = Uint<T>;
fn expecting(&self, f: &mut std::fmt::Formatter) -> std::fmt::Result {
f.write_str("a non-negative integer")
}
fn visit_u64<E: de::Error>(self, v: u64) -> Result<Uint<T>, E> {
T::try_from(v)
.map(Uint)
.map_err(|_| E::custom(format!("integer {v} is out of range for this field")))
}
fn visit_i64<E: de::Error>(self, v: i64) -> Result<Uint<T>, E> {
u64::try_from(v)
.map_err(|_| E::custom(format!("expected a non-negative integer, got {v}")))
.and_then(|u| self.visit_u64(u))
}
fn visit_f64<E: de::Error>(self, v: f64) -> Result<Uint<T>, E> {
if v.fract() == 0.0 && v >= 0.0 && v <= u64::MAX as f64 {
self.visit_u64(v as u64)
} else {
Err(E::custom(format!(
"expected a non-negative integer, got {v}"
)))
}
}
}
deserializer.deserialize_any(UintVisitor(PhantomData))
}
}
fn deserialize_uint<'de, D, T>(deserializer: D) -> Result<T, D::Error>
where
D: Deserializer<'de>,
T: TryFrom<u64>,
{
Uint::deserialize(deserializer).map(|u| u.0)
}
fn deserialize_opt_uint<'de, D, T>(deserializer: D) -> Result<Option<T>, D::Error>
where
D: Deserializer<'de>,
T: TryFrom<u64>,
{
Option::<Uint<T>>::deserialize(deserializer).map(|o| o.map(|u| u.0))
}
#[derive(Debug, Clone, Serialize)]
#[serde(untagged)]
pub enum Segmentation {
Polygon(Vec<Vec<f64>>),
CompressedRle { size: [u32; 2], counts: String },
UncompressedRle { size: [u32; 2], counts: Vec<u32> },
}
impl<'de> Deserialize<'de> for Segmentation {
fn deserialize<D: Deserializer<'de>>(deserializer: D) -> Result<Self, D::Error> {
enum Counts {
Str(String),
Ints(Vec<u32>),
}
impl<'de> Deserialize<'de> for Counts {
fn deserialize<D: Deserializer<'de>>(deserializer: D) -> Result<Self, D::Error> {
struct CountsVisitor;
impl<'de> serde::de::Visitor<'de> for CountsVisitor {
type Value = Counts;
fn expecting(&self, f: &mut std::fmt::Formatter) -> std::fmt::Result {
f.write_str("an RLE counts string or an array of run lengths")
}
fn visit_str<E: serde::de::Error>(self, v: &str) -> Result<Counts, E> {
Ok(Counts::Str(v.to_owned()))
}
fn visit_string<E: serde::de::Error>(self, v: String) -> Result<Counts, E> {
Ok(Counts::Str(v))
}
fn visit_seq<A: serde::de::SeqAccess<'de>>(
self,
mut seq: A,
) -> Result<Counts, A::Error> {
let mut v = Vec::with_capacity(seq.size_hint().unwrap_or(0));
while let Some(c) = seq.next_element()? {
v.push(c);
}
Ok(Counts::Ints(v))
}
}
deserializer.deserialize_any(CountsVisitor)
}
}
struct SegVisitor;
impl<'de> serde::de::Visitor<'de> for SegVisitor {
type Value = Segmentation;
fn expecting(&self, f: &mut std::fmt::Formatter) -> std::fmt::Result {
f.write_str("a list of polygons or an RLE object with `size` and `counts`")
}
fn visit_seq<A: serde::de::SeqAccess<'de>>(
self,
mut seq: A,
) -> Result<Segmentation, A::Error> {
let mut polys = Vec::with_capacity(seq.size_hint().unwrap_or(0));
while let Some(p) = seq.next_element()? {
polys.push(p);
}
Ok(Segmentation::Polygon(polys))
}
fn visit_map<A: serde::de::MapAccess<'de>>(
self,
mut map: A,
) -> Result<Segmentation, A::Error> {
let mut size: Option<[u32; 2]> = None;
let mut counts: Option<Counts> = None;
while let Some(key) = map.next_key::<std::borrow::Cow<'_, str>>()? {
match key.as_ref() {
"size" => size = Some(map.next_value()?),
"counts" => counts = Some(map.next_value()?),
_ => {
map.next_value::<serde::de::IgnoredAny>()?;
}
}
}
let size = size.ok_or_else(|| serde::de::Error::missing_field("size"))?;
match counts.ok_or_else(|| serde::de::Error::missing_field("counts"))? {
Counts::Str(counts) => Ok(Segmentation::CompressedRle { size, counts }),
Counts::Ints(counts) => Ok(Segmentation::UncompressedRle { size, counts }),
}
}
}
deserializer.deserialize_any(SegVisitor)
}
}
#[derive(Debug, Clone, Default, Deserialize, Serialize)]
pub struct Category {
#[serde(deserialize_with = "deserialize_uint")]
pub id: u64,
#[serde(default)]
pub name: String,
#[serde(default)]
pub supercategory: Option<String>,
#[serde(default)]
pub skeleton: Option<Vec<[u32; 2]>>,
#[serde(default)]
pub keypoints: Option<Vec<String>>,
#[serde(default)]
pub frequency: Option<String>,
#[serde(flatten, skip_serializing_if = "serde_json::Map::is_empty")]
pub extra: serde_json::Map<String, serde_json::Value>,
}
pub(crate) fn cat_id_to_name(dataset: &Dataset) -> HashMap<u64, &str> {
dataset
.categories
.iter()
.map(|c| (c.id, c.name.as_str()))
.collect()
}
pub(crate) fn cat_name_to_id(categories: &[Category]) -> HashMap<&str, u64> {
categories.iter().map(|c| (c.name.as_str(), c.id)).collect()
}
#[derive(Debug, Clone, Deserialize, Serialize)]
pub struct License {
#[serde(default)]
pub id: u64,
#[serde(default)]
pub name: Option<String>,
#[serde(default)]
pub url: Option<String>,
}
#[derive(Debug, Clone, PartialEq)]
pub struct Rle {
pub h: u32,
pub w: u32,
pub counts: Vec<u32>,
}
impl Rle {
pub fn new(h: u32, w: u32, counts: Vec<u32>) -> crate::error::Result<Self> {
let sum: u64 = counts.iter().map(|&c| c as u64).sum();
let expected = h as u64 * w as u64;
if sum != expected {
return Err(
format!("RLE counts must sum to h*w ({h} * {w} = {expected}), got {sum}").into(),
);
}
Ok(Self { h, w, counts })
}
}