use tch::nn;
use tch::nn::Module;
use tch::Tensor;
pub const REG_MAX: i64 = 16;
pub struct ClassifyHead {
fc: nn::Linear,
pub num_classes: i64,
}
impl ClassifyHead {
pub fn new(p: &nn::Path, in_c: i64, num_classes: i64) -> Self {
Self {
fc: nn::linear(p / "fc", in_c, num_classes, Default::default()),
num_classes,
}
}
pub fn logits(&self, pooled: &Tensor) -> Tensor {
self.fc.forward(pooled)
}
}
struct LevelHead {
cls1: nn::Conv2D,
cls2: nn::Conv2D,
cls_out: nn::Conv2D,
box1: nn::Conv2D,
box2: nn::Conv2D,
box_out: nn::Conv2D,
theta_out: Option<nn::Conv2D>,
}
impl LevelHead {
fn new(p: &nn::Path, in_c: i64, mid: i64, num_classes: i64, obb: bool) -> Self {
let cc = nn::ConvConfig {
padding: 1,
..Default::default()
};
Self {
cls1: nn::conv2d(p / "cls1", in_c, mid, 3, cc),
cls2: nn::conv2d(p / "cls2", mid, mid, 3, cc),
cls_out: nn::conv2d(p / "cls3", mid, num_classes, 1, Default::default()),
box1: nn::conv2d(p / "box1", in_c, mid, 3, cc),
box2: nn::conv2d(p / "box2", mid, mid, 3, cc),
box_out: nn::conv2d(p / "box3", mid, 4 * REG_MAX, 1, Default::default()),
theta_out: obb.then(|| nn::conv2d(p / "theta", mid, 1, 1, Default::default())),
}
}
fn forward(&self, feat: &Tensor) -> (Tensor, Tensor, Option<Tensor>) {
let cls = self
.cls_out
.forward(&self.cls2.forward(&self.cls1.forward(feat).relu()).relu());
let box_hidden = self.box2.forward(&self.box1.forward(feat).relu()).relu();
let box_dist = self.box_out.forward(&box_hidden);
let theta = self.theta_out.as_ref().map(|t| t.forward(&box_hidden));
(cls, box_dist, theta)
}
}
pub struct DetectHead {
levels: Vec<LevelHead>,
pub strides: Vec<u32>,
pub num_classes: i64,
}
impl DetectHead {
pub fn new(p: &nn::Path, strides: &[u32], channels: &[i64], num_classes: i64) -> Self {
Self::with_mode(p, strides, channels, num_classes, false)
}
pub fn new_obb(p: &nn::Path, strides: &[u32], channels: &[i64], num_classes: i64) -> Self {
Self::with_mode(p, strides, channels, num_classes, true)
}
pub fn with_mode(
p: &nn::Path,
strides: &[u32],
channels: &[i64],
num_classes: i64,
obb: bool,
) -> Self {
assert_eq!(
strides.len(),
channels.len(),
"head_levels 与 channels 数量必须一致"
);
let levels = strides
.iter()
.zip(channels.iter())
.map(|(&s, &c)| LevelHead::new(&(p / format!("s{s}")), c, 64, num_classes, obb))
.collect();
Self {
levels,
strides: strides.to_vec(),
num_classes,
}
}
pub fn forward(&self, feats: &[&Tensor]) -> Vec<(Tensor, Tensor, Option<Tensor>)> {
assert_eq!(
feats.len(),
self.levels.len(),
"特征层数与检测头层数必须一致"
);
self.levels
.iter()
.zip(feats.iter())
.map(|(head, feat)| head.forward(feat))
.collect()
}
}