use tch::nn;
use tch::nn::Module;
use tch::Tensor;
pub const KP_HEAD_STRIDE: u32 = 8;
pub const KP_CLS_CH: i64 = 0;
pub const KP_BOX_CH: i64 = 1;
pub const KP_OFF_CH: i64 = 5;
pub fn kp_out_channels(num_keypoints: i64) -> i64 {
5 + 3 * num_keypoints
}
pub struct KeypointHead {
c1: nn::Conv2D,
c2: nn::Conv2D,
out: nn::Conv2D,
pub num_keypoints: i64,
}
impl KeypointHead {
pub fn new(p: &nn::Path, in_c: i64, num_keypoints: i64) -> Self {
let cc = nn::ConvConfig {
padding: 1,
..Default::default()
};
Self {
c1: nn::conv2d(p / "c1", in_c, 64, 3, cc),
c2: nn::conv2d(p / "c2", 64, 64, 3, cc),
out: nn::conv2d(
p / "out",
64,
kp_out_channels(num_keypoints),
1,
Default::default(),
),
num_keypoints,
}
}
pub fn forward(&self, f8: &Tensor) -> Tensor {
self.out
.forward(&self.c2.forward(&self.c1.forward(f8).relu()).relu())
}
}
#[cfg(all(test, feature = "torch"))]
mod tests {
use super::*;
use tch::Device;
#[test]
fn head_output_channel_layout() {
let vs = tch::nn::VarStore::new(Device::Cpu);
let head = KeypointHead::new(&vs.root(), 64, 3);
assert_eq!(kp_out_channels(3), 14);
let x = Tensor::randn([2, 64, 8, 8], (tch::Kind::Float, Device::Cpu));
let out = head.forward(&x);
assert_eq!(out.size(), vec![2, kp_out_channels(3), 8, 8]);
assert_eq!(head.num_keypoints, 3);
}
}