use std::ops::{Add, Mul};
use tch::Tensor;
use crate::error::{AvError, AvResult};
use crate::types::ImageMeta;
pub type LossDict = Vec<(&'static str, f32)>;
#[derive(Debug)]
pub struct FeatureMap {
pub tensor: Tensor,
pub stride: u32,
pub channels: usize,
}
impl FeatureMap {
pub fn new(tensor: Tensor, stride: u32) -> AvResult<Self> {
let size = tensor.size();
if size.len() != 4 {
return Err(AvError::shape(format!(
"FeatureMap 需要 4 维 [N,C,H,W],实际 {size:?}(stride={stride})"
)));
}
Ok(Self {
tensor,
stride,
channels: size[1] as usize,
})
}
pub fn height(&self) -> i64 {
self.tensor.size()[2]
}
pub fn width(&self) -> i64 {
self.tensor.size()[3]
}
}
#[derive(Debug, Default)]
pub struct FeaturePyramid {
pub levels: Vec<FeatureMap>,
}
impl FeaturePyramid {
pub fn validate_ascending(&self) -> AvResult<()> {
for w in self.levels.windows(2) {
if w[0].stride >= w[1].stride {
return Err(AvError::shape(format!(
"金字塔 stride 必须严格升序:{} -> {}",
w[0].stride, w[1].stride
)));
}
}
Ok(())
}
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct BackboneSpec {
pub levels: Vec<LevelSpec>,
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct LevelSpec {
pub stride: u32,
pub channels: usize,
}
pub trait Configurable: Sized {
fn from_config(cfg: &toml::Value) -> AvResult<Self>;
fn validate(&self) -> AvResult<()> {
Ok(())
}
}
pub trait BaseBackbone: Send {
fn forward_features(&self, x: &Tensor) -> AvResult<FeaturePyramid>;
fn forward_pooled(&self, x: &Tensor) -> AvResult<Tensor>;
fn spec(&self) -> BackboneSpec;
}
pub trait FeatureNeck: Send {
fn forward(&self, feats: FeaturePyramid) -> AvResult<FeaturePyramid>;
}
pub trait TaskHead: Send {
type Target;
type RawPred;
fn forward(&self, feats: &FeaturePyramid) -> AvResult<Self::RawPred>;
}
pub trait LossAggregator: Default {
fn add(&mut self, name: &'static str, value: Tensor, weight: f64);
fn total(&self) -> Tensor;
fn snapshot(&self) -> LossDict;
}
#[derive(Default)]
pub struct WeightedSumAggregator {
parts: Vec<(&'static str, Tensor, f64)>,
}
impl LossAggregator for WeightedSumAggregator {
fn add(&mut self, name: &'static str, value: Tensor, weight: f64) {
self.parts.push((name, value, weight));
}
fn total(&self) -> Tensor {
let mut acc: Option<Tensor> = None;
for (_, v, w) in &self.parts {
let term = v.mul(&Tensor::from(*w as f32));
acc = Some(match acc {
None => term,
Some(a) => a.add(&term),
});
}
acc.unwrap_or_else(|| Tensor::from(0f32))
}
fn snapshot(&self) -> LossDict {
self.parts
.iter()
.map(|(name, v, w)| {
let scalar = if v.size().is_empty() {
v.double_value(&[])
} else {
v.mean_dim(&[0i64][..], true, v.kind()).double_value(&[])
};
(*name, (scalar * *w) as f32)
})
.collect()
}
}
pub trait PostProcessor {
type RawPred;
type Output;
fn process(&self, pred: &Self::RawPred, meta: &ImageMeta) -> AvResult<Self::Output>;
}
#[cfg(all(test, feature = "torch"))]
mod tests {
use super::*;
#[test]
fn weighted_sum_total_and_snapshot() {
let mut agg = WeightedSumAggregator::default();
agg.add("cls", Tensor::from(2f32), 1.0);
agg.add("box", Tensor::from(4f32), 0.5);
let total = agg.total();
assert!((total.double_value(&[]) - 4.0).abs() < 1e-6);
let snap = agg.snapshot();
assert_eq!(snap.len(), 2);
assert!((snap[1].1 - 2.0).abs() < 1e-5);
}
#[test]
fn empty_total_is_zero() {
let agg = WeightedSumAggregator::default();
assert!(agg.total().double_value(&[]).abs() < 1e-7);
}
#[test]
fn feature_map_shape_contract() {
let t = Tensor::randn([1, 8, 16, 16], (tch::Kind::Float, tch::Device::Cpu));
let fm = FeatureMap::new(t, 8).unwrap();
assert_eq!(fm.channels, 8);
assert_eq!(fm.height(), 16);
let bad = Tensor::randn([8, 16, 16], (tch::Kind::Float, tch::Device::Cpu));
assert!(FeatureMap::new(bad, 8).is_err());
}
#[test]
fn pyramid_stride_ascending_contract() {
let mut p = FeaturePyramid::default();
p.levels.push(
FeatureMap::new(
Tensor::randn([1, 8, 32, 32], (tch::Kind::Float, tch::Device::Cpu)),
8,
)
.unwrap(),
);
p.levels.push(
FeatureMap::new(
Tensor::randn([1, 8, 16, 16], (tch::Kind::Float, tch::Device::Cpu)),
16,
)
.unwrap(),
);
assert!(p.validate_ascending().is_ok());
p.levels[1].stride = 8;
assert!(p.validate_ascending().is_err());
}
}