use tch::nn;
use tch::nn::Module;
use tch::{Kind, Tensor};
pub const NUM_PROTOS: i64 = 32;
pub fn dice_loss(pred: &Tensor, gt: &Tensor) -> Tensor {
let eps = 1e-5f64;
let inter = (pred * gt).sum(Kind::Float);
let denom = pred.sum(Kind::Float) + gt.sum(Kind::Float);
(&denom - &inter * 2.0 + eps) / (&denom + eps)
}
pub fn mask_iou(a: &[u8], b: &[u8]) -> f32 {
if a.len() != b.len() || a.is_empty() {
return 0.0;
}
let mut inter = 0usize;
let mut union = 0usize;
for (&x, &y) in a.iter().zip(b) {
inter += (x != 0 && y != 0) as usize;
union += (x != 0 || y != 0) as usize;
}
if union == 0 {
return 0.0;
}
inter as f32 / union as f32
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct MaskSummary {
width: Option<usize>,
bbox: Option<[usize; 4]>,
len: usize,
}
impl MaskSummary {
pub fn of(mask: &[u8]) -> Self {
let w = (mask.len() as f64).sqrt() as usize;
let width = (w > 0 && w * w == mask.len()).then_some(w);
let bbox = width.and_then(|w| {
let mut bbox: Option<[usize; 4]> = None;
for (i, &v) in mask.iter().enumerate() {
if v == 0 {
continue;
}
let (x, y) = (i % w, i / w);
bbox = Some(match bbox {
None => [x, y, x, y],
Some([x1, y1, x2, y2]) => [x1.min(x), y1.min(y), x2.max(x), y2.max(y)],
});
}
bbox
});
Self {
width,
bbox,
len: mask.len(),
}
}
pub fn iou(&self, a: &[u8], other: &Self, b: &[u8]) -> f32 {
if self.len != other.len {
return 0.0; }
match (self.width, other.width) {
(Some(_), Some(_)) => match (self.bbox, other.bbox) {
(Some(ba), Some(bb)) => {
if ba[2] < bb[0] || bb[2] < ba[0] || ba[3] < bb[1] || bb[3] < ba[1] {
0.0 } else {
mask_iou(a, b)
}
}
_ => 0.0,
},
_ => mask_iou(a, b),
}
}
}
pub struct MaskBranch {
proto1: nn::Conv2D,
proto2: nn::Conv2D,
proto_out: nn::Conv2D,
coef1: nn::Conv2D,
coef2: nn::Conv2D,
coef_out: nn::Conv2D,
pub num_classes: i64,
pub num_protos: i64,
}
impl MaskBranch {
pub fn new(p: &nn::Path, in_c: i64, num_classes: i64, num_protos: i64) -> Self {
let cc = nn::ConvConfig {
padding: 1,
..Default::default()
};
Self {
proto1: nn::conv2d(p / "proto1", in_c, 64, 3, cc),
proto2: nn::conv2d(p / "proto2", 64, 64, 3, cc),
proto_out: nn::conv2d(p / "proto_out", 64, num_protos, 1, Default::default()),
coef1: nn::conv2d(p / "coef1", in_c, 64, 3, cc),
coef2: nn::conv2d(p / "coef2", 64, 64, 3, cc),
coef_out: nn::conv2d(
p / "coef_out",
64,
num_classes + num_protos,
1,
Default::default(),
),
num_classes,
num_protos,
}
}
pub fn forward(&self, f16: &Tensor, mask_hw: (i64, i64)) -> (Tensor, Tensor) {
let hidden = self.proto1.forward(f16).relu();
let up = Tensor::upsample_bilinear2d(&hidden, [mask_hw.0, mask_hw.1], false, None, None);
let proto = self.proto_out.forward(&self.proto2.forward(&up).relu());
let coef_hidden = self.coef2.forward(&self.coef1.forward(f16).relu()).relu();
let coef_cls = self.coef_out.forward(&coef_hidden);
(proto, coef_cls)
}
}
#[cfg(all(test, feature = "torch"))]
mod tests {
use super::*;
use tch::Device;
#[test]
fn dice_loss_identical_masks_near_zero() {
let gt = Tensor::from_slice(&[1.0f32, 1.0, 0.0, 0.0]).reshape([2i64, 2]);
let pred = Tensor::from_slice(&[20.0f32, 20.0, -20.0, -20.0])
.reshape([2i64, 2])
.sigmoid();
let l = dice_loss(&pred, >);
assert!(l.double_value(&[]) < 1e-4, "got {}", l.double_value(&[]));
}
#[test]
fn dice_loss_hand_computed_constant_prediction() {
let gt = Tensor::from_slice(&[1.0f32, 0.0, 0.0, 0.0]).reshape([2i64, 2]);
let pred = Tensor::zeros([2i64, 2], (Kind::Float, Device::Cpu)).sigmoid();
let l = dice_loss(&pred, >);
assert!(
(l.double_value(&[]) - 2.0 / 3.0).abs() < 1e-4,
"got {}",
l.double_value(&[])
);
}
#[test]
fn mask_iou_basic_cases() {
let a = vec![1u8, 1, 0, 0];
let b = vec![1u8, 0, 1, 0];
assert!((mask_iou(&a, &b) - 1.0 / 3.0).abs() < 1e-6);
assert!((mask_iou(&a, &a) - 1.0).abs() < 1e-6);
assert_eq!(mask_iou(&a, &[0u8; 4]), 0.0);
assert_eq!(mask_iou(&a, &[1u8]), 0.0, "长度不等应返回 0");
assert_eq!(mask_iou(&[0u8; 4], &[0u8; 4]), 0.0, "双空掩码应返回 0");
}
fn mask4(cells: &[(usize, usize)]) -> Vec<u8> {
let mut m = vec![0u8; 16];
for &(x, y) in cells {
m[y * 4 + x] = 1;
}
m
}
#[test]
fn mask_summary_iou_matches_mask_iou() {
let a = mask4(&[(0, 0), (1, 0), (1, 1)]);
let b = mask4(&[(1, 0), (1, 1), (2, 1)]);
let (sa, sb) = (MaskSummary::of(&a), MaskSummary::of(&b));
assert_eq!(sa.iou(&a, &sb, &b), mask_iou(&a, &b));
assert_eq!(sa.iou(&a, &sa, &a), 1.0);
}
#[test]
fn mask_summary_short_circuits_separated_boxes() {
let left = mask4(&[(0, 1), (0, 2)]);
let right = mask4(&[(3, 1), (3, 2)]);
let adjacent = mask4(&[(1, 1), (1, 2)]);
let (sl, sr, sadj) = (
MaskSummary::of(&left),
MaskSummary::of(&right),
MaskSummary::of(&adjacent),
);
assert_eq!(sl.iou(&left, &sr, &right), 0.0);
assert_eq!(sl.iou(&left, &sadj, &adjacent), 0.0);
assert_eq!(sl.bbox, Some([0, 1, 0, 2]));
}
#[test]
fn mask_summary_empty_and_unequal_masks() {
let a = mask4(&[(0, 0)]);
let (sa, sempty) = (MaskSummary::of(&a), MaskSummary::of(&[0u8; 16]));
assert_eq!(sempty.bbox, None);
assert_eq!(sa.iou(&a, &sempty, &[0u8; 16]), 0.0, "空+非空 → 0");
assert_eq!(sempty.iou(&[0u8; 16], &sempty, &[0u8; 16]), 0.0, "双空 → 0");
assert_eq!(sa.iou(&a, &MaskSummary::of(&[1u8]), &[1u8]), 0.0);
let nonsquare = vec![1u8, 0, 1];
let sn = MaskSummary::of(&nonsquare);
assert_eq!(sn.width, None);
assert_eq!(
sn.iou(&nonsquare, &sn, &nonsquare),
mask_iou(&nonsquare, &nonsquare)
);
}
}