use std::cell::Cell;
use tch::nn;
use tch::nn::{Module, ModuleT};
use tch::Tensor;
use av_core::config::BackboneCfg;
use av_core::error::{AvError, AvResult};
use av_core::traits::{BackboneSpec, BaseBackbone, FeatureMap, FeaturePyramid, LevelSpec};
pub const FAMILY_NAME: &str = "resnet18";
pub const STAGE_CHANNELS: [i64; 4] = [64, 128, 256, 512];
const STAGE_STRIDES: [i64; 4] = [1, 2, 2, 2];
const BLOCKS_PER_STAGE: usize = 2;
pub const POOLED_CHANNELS: i64 = STAGE_CHANNELS[3];
struct BasicBlock {
conv1: nn::Conv2D,
bn1: nn::BatchNorm,
conv2: nn::Conv2D,
bn2: nn::BatchNorm,
downsample: Option<(nn::Conv2D, nn::BatchNorm)>,
}
impl BasicBlock {
fn new(p: &nn::Path, in_ch: i64, out_ch: i64, stride: i64) -> Self {
let conv1 = nn::conv2d(
p / "conv1",
in_ch,
out_ch,
3,
nn::ConvConfig {
stride,
padding: 1,
bias: false,
..Default::default()
},
);
let bn1 = nn::batch_norm2d(p / "bn1", out_ch, Default::default());
let conv2 = nn::conv2d(
p / "conv2",
out_ch,
out_ch,
3,
nn::ConvConfig {
padding: 1,
bias: false,
..Default::default()
},
);
let bn2 = nn::batch_norm2d(p / "bn2", out_ch, Default::default());
let downsample = if stride != 1 || in_ch != out_ch {
let ds = p / "downsample";
let conv = nn::conv2d(
&ds / "0",
in_ch,
out_ch,
1,
nn::ConvConfig {
stride,
bias: false,
..Default::default()
},
);
let bn = nn::batch_norm2d(&ds / "1", out_ch, Default::default());
Some((conv, bn))
} else {
None
};
Self {
conv1,
bn1,
conv2,
bn2,
downsample,
}
}
fn forward(&self, x: &Tensor, train: bool) -> Tensor {
let mut out = self.conv1.forward(x);
out = self.bn1.forward_t(&out, train).relu();
out = self.conv2.forward(&out);
out = self.bn2.forward_t(&out, train);
let identity = match &self.downsample {
Some((conv, bn)) => bn.forward_t(&conv.forward(x), train),
None => x.shallow_clone(),
};
(out + identity).relu()
}
}
pub struct ResNetBackbone {
conv1: nn::Conv2D,
bn1: nn::BatchNorm,
layers: [Vec<BasicBlock>; 4],
train: Cell<bool>,
}
impl ResNetBackbone {
pub fn new(p: &nn::Path, cfg: &BackboneCfg) -> AvResult<Self> {
if (cfg.width - 1.0).abs() > 1e-6 {
return Err(AvError::config(format!(
"{FAMILY_NAME} 固定 torchvision 宽度(64/128/256/512),width = {} 不支持:\
ImageNet 预训练权重形状兼容优先,改宽请用 simple-cnn 或后续 resnet 变宽档",
cfg.width
)));
}
let stem = nn::ConvConfig {
stride: 2,
padding: 3,
bias: false,
..Default::default()
};
let conv1 = nn::conv2d(p / "conv1", 3, STAGE_CHANNELS[0], 7, stem);
let bn1 = nn::batch_norm2d(p / "bn1", STAGE_CHANNELS[0], Default::default());
let mut layers: [Vec<BasicBlock>; 4] = Default::default();
let mut in_ch = STAGE_CHANNELS[0];
for (stage, &out_ch) in STAGE_CHANNELS.iter().enumerate() {
let lp = p / &format!("layer{}", stage + 1);
let stride = STAGE_STRIDES[stage];
let mut blocks = Vec::with_capacity(BLOCKS_PER_STAGE);
for b in 0..BLOCKS_PER_STAGE {
let bp = &lp / &b.to_string();
blocks.push(BasicBlock::new(
&bp,
in_ch,
out_ch,
if b == 0 { stride } else { 1 },
));
in_ch = out_ch;
}
layers[stage] = blocks;
}
Ok(Self {
conv1,
bn1,
layers,
train: Cell::new(false),
})
}
pub fn set_train(&self, train: bool) {
self.train.set(train);
}
pub fn pooled_channels(&self) -> i64 {
POOLED_CHANNELS
}
pub fn stride_channels(&self, stride: u32) -> AvResult<i64> {
match stride {
4 => Ok(STAGE_CHANNELS[0]), 8 => Ok(STAGE_CHANNELS[1]), 16 => Ok(STAGE_CHANNELS[2]), other => Err(AvError::shape(format!(
"{FAMILY_NAME} 不存在 stride {other} 的特征层(金字塔 = layer1/2/3,\
layer4 专供 forward_pooled)"
))),
}
}
fn forward_all(&self, x: &Tensor, train: bool) -> (Tensor, Tensor, Tensor, Tensor) {
let mut x = self.conv1.forward(x);
x = self.bn1.forward_t(&x, train).relu();
x = x.max_pool2d([3i64, 3], [2i64, 2], [1i64, 1], [1i64, 1], false);
let mut feats = Vec::with_capacity(4);
for stage in &self.layers {
let mut out = feats.last().unwrap_or(&x).shallow_clone();
for block in stage {
out = block.forward(&out, train);
}
feats.push(out);
}
let mut it = feats.into_iter();
(
it.next().expect("layer1"),
it.next().expect("layer2"),
it.next().expect("layer3"),
it.next().expect("layer4"),
)
}
}
impl BaseBackbone for ResNetBackbone {
fn forward_features(&self, x: &Tensor) -> AvResult<FeaturePyramid> {
let (l1, l2, l3, _l4) = self.forward_all(x, self.train.get());
let mut pyramid = FeaturePyramid::default();
pyramid.levels.push(FeatureMap::new(l1, 4)?);
pyramid.levels.push(FeatureMap::new(l2, 8)?);
pyramid.levels.push(FeatureMap::new(l3, 16)?);
pyramid.validate_ascending()?;
Ok(pyramid)
}
fn forward_pooled(&self, x: &Tensor) -> AvResult<Tensor> {
let (_l1, _l2, _l3, l4) = self.forward_all(x, self.train.get());
let pooled = l4
.adaptive_avg_pool2d([1, 1])
.reshape([-1, POOLED_CHANNELS]);
Ok(pooled)
}
fn spec(&self) -> BackboneSpec {
BackboneSpec {
levels: vec![
LevelSpec {
stride: 4,
channels: STAGE_CHANNELS[0] as usize,
},
LevelSpec {
stride: 8,
channels: STAGE_CHANNELS[1] as usize,
},
LevelSpec {
stride: 16,
channels: STAGE_CHANNELS[2] as usize,
},
],
}
}
}
#[cfg(all(test, feature = "torch"))]
mod tests {
use super::*;
use tch::{Device, Kind};
fn resnet18_backbone(vs: &nn::VarStore) -> ResNetBackbone {
ResNetBackbone::new(&(vs.root() / "backbone"), &BackboneCfg::default())
.expect("resnet18 默认配置(width=1.0)应可装配")
}
#[test]
fn rejects_width_scaling() {
let vs = nn::VarStore::new(Device::Cpu);
let cfg = BackboneCfg {
width: 0.5,
..Default::default()
};
let err = match ResNetBackbone::new(&(vs.root() / "backbone"), &cfg) {
Err(e) => e,
Ok(_) => panic!("width = 0.5 应被拒绝"),
};
assert!(err.to_string().contains("width"), "{err}");
}
#[test]
fn feature_pyramid_contract() {
let vs = nn::VarStore::new(Device::Cpu);
let backbone = resnet18_backbone(&vs);
let x = Tensor::randn([2, 3, 64, 64], (Kind::Float, Device::Cpu));
let pyramid = backbone.forward_features(&x).unwrap();
let strides: Vec<u32> = pyramid.levels.iter().map(|l| l.stride).collect();
assert_eq!(strides, vec![4, 8, 16]);
let channels: Vec<usize> = pyramid.levels.iter().map(|l| l.channels).collect();
assert_eq!(channels, vec![64, 128, 256], "ResNet18 layer1/2/3 真实宽度");
for (level, expect) in pyramid.levels.iter().zip([16i64, 8, 4]) {
assert_eq!(
level.tensor.size()[2],
expect,
"stride {} 空间尺寸",
level.stride
);
}
let pooled = backbone.forward_pooled(&x).unwrap();
assert_eq!(pooled.size(), vec![2, backbone.pooled_channels()]);
assert_eq!(backbone.pooled_channels(), POOLED_CHANNELS);
for lv in backbone.spec().levels {
assert_eq!(
backbone.stride_channels(lv.stride).unwrap(),
lv.channels as i64
);
}
}
#[test]
fn feature_pyramid_spatial_contract_at_320() {
let vs = nn::VarStore::new(Device::Cpu);
let backbone = resnet18_backbone(&vs);
let x = Tensor::randn([2, 3, 320, 320], (Kind::Float, Device::Cpu));
let pyramid = backbone.forward_features(&x).unwrap();
for (level, (ch, hw)) in pyramid
.levels
.iter()
.zip([(64i64, 80i64), (128, 40), (256, 20)])
{
assert_eq!(
level.tensor.size(),
vec![2, ch, hw, hw],
"stride {}",
level.stride
);
}
}
#[test]
fn torchvision_layer_names_inventory() {
let vs = nn::VarStore::new(Device::Cpu);
let _ = resnet18_backbone(&vs);
let mut expected: Vec<String> = vec!["backbone.conv1.weight".into()];
let bn = ["weight", "bias", "running_mean", "running_var"];
expected.extend(bn.iter().map(|s| format!("backbone.bn1.{s}")));
for stage in 1..=4usize {
for b in 0..BLOCKS_PER_STAGE {
let base = format!("backbone.layer{stage}.{b}");
expected.push(format!("{base}.conv1.weight"));
expected.extend(bn.iter().map(|s| format!("{base}.bn1.{s}")));
expected.push(format!("{base}.conv2.weight"));
expected.extend(bn.iter().map(|s| format!("{base}.bn2.{s}")));
if stage > 1 && b == 0 {
expected.push(format!("{base}.downsample.0.weight"));
expected.extend(bn.iter().map(|s| format!("{base}.downsample.1.{s}")));
}
}
}
assert_eq!(expected.len(), 100, "20 conv + 40 BN 参数 + 40 BN 统计量");
let mut got: Vec<String> = vs
.variables()
.into_keys()
.filter(|n| n.starts_with("backbone."))
.collect();
got.sort();
expected.sort();
assert_eq!(
got, expected,
"变量名必须与 torchvision 层名逐一同名(+前缀)"
);
}
#[test]
fn bn_buffers_live_in_varstore_and_drive_eval() {
let vs = nn::VarStore::new(Device::Cpu);
let backbone = resnet18_backbone(&vs);
assert!(vs.variables().contains_key("backbone.bn1.running_mean"));
assert!(vs
.variables()
.contains_key("backbone.layer2.0.downsample.1.running_var"));
let x = Tensor::randn([2, 3, 64, 64], (Kind::Float, Device::Cpu));
let out1 = backbone.forward_pooled(&x).unwrap();
let shift = Tensor::ones([64i64], (Kind::Float, Device::Cpu)) * 5.0;
tch::no_grad(|| {
let mut vars = vs.variables();
let mut rm = vars
.remove("backbone.bn1.running_mean")
.expect("统计量应存在");
rm.copy_(&shift);
});
let out2 = backbone.forward_pooled(&x).unwrap();
let diff = (&out2 - &out1).abs().max().double_value(&[]);
assert!(
diff > 1e-3,
"改写 running_mean 必须改变推理输出(diff={diff})"
);
let out3 = backbone.forward_pooled(&x).unwrap();
assert_eq!((&out3 - &out2).abs().max().double_value(&[]), 0.0);
}
#[test]
fn bn_train_flag_switches_batch_stats_and_updates_running() {
let vs = nn::VarStore::new(Device::Cpu);
let backbone = resnet18_backbone(&vs);
let x = Tensor::randn([4, 3, 64, 64], (Kind::Float, Device::Cpu));
let e1 = tch::no_grad(|| backbone.forward_pooled(&x).unwrap());
let e2 = tch::no_grad(|| backbone.forward_pooled(&x).unwrap());
assert_eq!(
(&e1 - &e2).abs().max().double_value(&[]),
0.0,
"默认 FrozenBN(eval)两次前向应逐位一致"
);
let rm_sum_before = vs.variables()["backbone.bn1.running_mean"]
.sum(Kind::Float)
.double_value(&[]);
backbone.set_train(true);
let t1 = tch::no_grad(|| backbone.forward_pooled(&x).unwrap());
let t2 = tch::no_grad(|| backbone.forward_pooled(&x).unwrap());
assert_eq!(
(&t1 - &t2).abs().max().double_value(&[]),
0.0,
"train 态输出 = 批统计归一化,同批两次前向应一致(running 不进训练态输出)"
);
let t_vs_e = tch::no_grad(|| (&t1 - &e1).abs().max().double_value(&[]));
assert!(t_vs_e > 1e-4, "train 态应用批统计(与 eval 差 {t_vs_e})");
let rm_sum_after = vs.variables()["backbone.bn1.running_mean"]
.sum(Kind::Float)
.double_value(&[]);
assert!(
(rm_sum_after - rm_sum_before).abs() > 1e-6,
"train 前向应更新 running_mean(sum {rm_sum_before} → {rm_sum_after})"
);
backbone.set_train(false);
let e3 = tch::no_grad(|| backbone.forward_pooled(&x).unwrap());
let e4 = tch::no_grad(|| backbone.forward_pooled(&x).unwrap());
assert_eq!(
(&e3 - &e4).abs().max().double_value(&[]),
0.0,
"回 eval 应恢复确定性"
);
let e3_vs_e1 = tch::no_grad(|| (&e3 - &e1).abs().max().double_value(&[]));
assert!(
e3_vs_e1 > 1e-6,
"eval 前向应消费更新后的 running 统计量(diff={e3_vs_e1})"
);
}
#[test]
fn bn_eval_matches_manual_formula() {
let vs = nn::VarStore::new(Device::Cpu);
let mut bn = nn::batch_norm2d(&(vs.root() / "bn"), 2, Default::default());
tch::no_grad(|| {
bn.running_mean.copy_(&Tensor::from_slice(&[1.0f32, -2.0]));
bn.running_var.copy_(&Tensor::from_slice(&[4.0f32, 0.25]));
bn.ws
.as_mut()
.unwrap()
.copy_(&Tensor::from_slice(&[2.0f32, 3.0]));
bn.bs
.as_mut()
.unwrap()
.copy_(&Tensor::from_slice(&[0.5f32, -1.0]));
});
let x = Tensor::randn([2, 2, 3, 3], (Kind::Float, Device::Cpu));
let y = bn.forward_t(&x, false);
let eps = 1e-5f64;
let mean = Tensor::from_slice(&[1.0f32, -2.0]).reshape([1i64, 2, 1, 1]);
let var = Tensor::from_slice(&[4.0f32, 0.25]).reshape([1i64, 2, 1, 1]);
let gamma = Tensor::from_slice(&[2.0f32, 3.0]).reshape([1i64, 2, 1, 1]);
let beta = Tensor::from_slice(&[0.5f32, -1.0]).reshape([1i64, 2, 1, 1]);
let manual = &(&(&x - &mean) / &(&var + eps).sqrt() * &gamma) + β
let diff = (&y - &manual).abs().max().double_value(&[]);
assert!(diff < 1e-5, "推理 BN 应等于手工公式(diff={diff})");
}
}