use core::cmp::{Ordering, Reverse};
use std::collections::BinaryHeap;
use crate::audio::ced::{
NUM_CLASSES,
error::{ClassCountMismatch, Error, InvalidConfidence, Result},
};
pub use soundevents_dataset::RatedSoundEvent;
pub use soundevents_dataset::SoundEventId;
#[cfg(test)]
mod tests;
pub type WindowConfidences = windit::windowed::Windowed<Confidences>;
#[derive(Debug, Clone, Copy)]
pub struct EventPrediction {
event: &'static RatedSoundEvent,
confidence: f32,
}
impl EventPrediction {
pub(crate) fn from_confidence(class_index: usize, confidence: f32) -> Result<Self> {
let event =
RatedSoundEvent::from_index(class_index).ok_or(Error::UnknownClassIndex(class_index))?;
Ok(Self { event, confidence })
}
#[inline]
pub const fn event(&self) -> &'static RatedSoundEvent {
self.event
}
#[inline]
pub const fn index(&self) -> usize {
self.event.index()
}
#[inline]
pub const fn name(&self) -> &'static str {
self.event.name()
}
#[inline]
pub const fn id(&self) -> SoundEventId {
self.event.id()
}
#[inline]
pub const fn mid(&self) -> &'static str {
self.event.mid()
}
#[inline]
pub const fn confidence(&self) -> f32 {
self.confidence
}
}
#[derive(Debug, Clone, PartialEq)]
pub struct Confidences {
values: Vec<f32>,
}
impl Confidences {
pub(crate) fn new(values: Vec<f32>) -> Self {
assert!(
values.len() == NUM_CLASSES,
"Confidences requires exactly NUM_CLASSES values, got {}",
values.len()
);
Self { values }
}
pub(crate) fn from_logits(logits: &[f32]) -> Self {
Self::new(logits.iter().copied().map(sigmoid).collect())
}
pub fn try_from_slice(values: &[f32]) -> Result<Self> {
if values.len() != NUM_CLASSES {
return Err(Error::ClassCountMismatch(ClassCountMismatch::new(
NUM_CLASSES,
values.len(),
)));
}
for (index, &value) in values.iter().enumerate() {
if !(0.0..=1.0).contains(&value) {
return Err(Error::InvalidConfidence(InvalidConfidence::new(
index, value,
)));
}
}
Ok(Self::new(values.to_vec()))
}
#[inline]
pub fn as_slice(&self) -> &[f32] {
&self.values
}
pub fn top_k(&self, k: usize) -> Result<Vec<EventPrediction>> {
top_k_from_scores(self.values.iter().copied().enumerate(), k, |c| c)
}
}
pub(crate) fn sigmoid(x: f32) -> f32 {
1.0 / (1.0 + (-x).exp())
}
#[derive(Debug, Clone, Copy)]
struct RankedScore {
class_index: usize,
score: f32,
}
impl PartialEq for RankedScore {
fn eq(&self, other: &Self) -> bool {
self.class_index == other.class_index && self.score.total_cmp(&other.score) == Ordering::Equal
}
}
impl Eq for RankedScore {}
impl PartialOrd for RankedScore {
fn partial_cmp(&self, other: &Self) -> Option<Ordering> {
Some(self.cmp(other))
}
}
impl Ord for RankedScore {
fn cmp(&self, other: &Self) -> Ordering {
self
.score
.total_cmp(&other.score)
.then_with(|| other.class_index.cmp(&self.class_index))
}
}
pub(crate) fn top_k_from_scores(
scores: impl IntoIterator<Item = (usize, f32)>,
k: usize,
map_score: impl Fn(f32) -> f32,
) -> Result<Vec<EventPrediction>> {
if k == 0 {
return Ok(Vec::new());
}
let mut heap = BinaryHeap::with_capacity(k.min(NUM_CLASSES));
for (class_index, score) in scores {
let candidate = Reverse(RankedScore { class_index, score });
if heap.len() < k {
heap.push(candidate);
continue;
}
if heap.peek().is_some_and(|smallest| candidate.0 > smallest.0) {
heap.pop();
heap.push(candidate);
}
}
let mut predictions = Vec::with_capacity(heap.len());
while let Some(entry) = heap.pop() {
let ranked = entry.0;
predictions.push(EventPrediction::from_confidence(
ranked.class_index,
map_score(ranked.score),
)?);
}
predictions.reverse();
Ok(predictions)
}