#![cfg(feature = "torch")]
use av_core::config::RunConfig;
use av_tasks::rng::XorShift;
fn classify_cfg() -> RunConfig {
let mut cfg = RunConfig::from_toml_str(include_str!("../../../configs/classify_smoke.toml"))
.expect("classify smoke 配置必须合法");
cfg.output_dir = std::env::temp_dir().join("av-test-runs");
cfg
}
fn detect_cfg() -> RunConfig {
let mut cfg = RunConfig::from_toml_str(include_str!("../../../configs/detect_smoke.toml"))
.expect("detect smoke 配置必须合法");
cfg.output_dir = std::env::temp_dir().join("av-test-runs");
cfg
}
use av_core::config::PretrainCfg;
use av_pretrain::av_weight::WeightMeta;
use av_pretrain::weight_adapter::{self, LayerMap};
use tch::{Device, Kind, Tensor};
fn temp_dir(tag: &str) -> std::path::PathBuf {
static N: std::sync::atomic::AtomicUsize = std::sync::atomic::AtomicUsize::new(0);
let n = N.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
std::env::temp_dir().join(format!("av-test-{tag}-{}-{n}", std::process::id()))
}
#[test]
fn pretrain_native_roundtrip_restores_inference() {
let mut cfg = classify_cfg();
cfg.run_id = "classify-pretrain-roundtrip".into();
let report = av_runtime::engine::train(&cfg).expect("分类训练应成功");
let ckpt = std::path::Path::new(&report.run_dir).join("best.ckpt");
let export_dir = temp_dir("pretrain-export");
let vs = tch::nn::VarStore::new(Device::Cpu);
let model_ref = av_tasks::models::build_model(&vs.root(), &cfg).expect("模型构建应成功");
let mut vars: Vec<(String, Tensor)> = vs.variables().into_iter().collect();
av_pretrain::av_weight::load_named(&ckpt, &mut vars).expect("checkpoint 应可读回");
let meta = WeightMeta {
backbone: "simple-cnn".into(),
source_dataset: "synthetic".into(),
epoch: Some(report.epochs),
task: "classify".into(),
};
let exported: Vec<(String, Tensor)> = vs.variables().into_iter().collect();
let manifest = av_pretrain::av_weight::export_pretrain_dir(&export_dir, &exported, meta)
.expect("avpretrain 导出应成功");
assert_eq!(manifest.format, "avpretrain");
assert_eq!(manifest.tensors.len(), exported.len());
assert_eq!(manifest.backbone, "simple-cnn");
assert!(manifest.created_at.ends_with('Z'));
let vs2 = tch::nn::VarStore::new(Device::Cpu);
let model_new = av_tasks::models::build_model(&vs2.root(), &cfg).expect("模型构建应成功");
let manifest2 = av_pretrain::av_weight::read_manifest(&export_dir).expect("manifest 应可读");
av_pretrain::av_weight::verify_hashes(&export_dir, &manifest2)
.expect("哈希校验应通过(导出后无人篡改)");
let mut vars2: Vec<(String, Tensor)> = vs2.variables().into_iter().collect();
av_pretrain::av_weight::load_named(&export_dir, &mut vars2).expect("avpretrain 应可加载");
tch::manual_seed(2024);
let x = Tensor::randn([4, 3, 64, 64], (Kind::Float, Device::Cpu));
let out_ref = model_ref.predict(&x, 0.0, 0.0).expect("参考模型推理应成功");
let out_new = model_new.predict(&x, 0.0, 0.0).expect("加载模型推理应成功");
let (
av_tasks::models::PredictOutput::Classify {
labels: l1,
confs: c1,
},
av_tasks::models::PredictOutput::Classify {
labels: l2,
confs: c2,
},
) = (out_ref, out_new)
else {
panic!("应为分类输出");
};
assert_eq!(l1, l2, "标签应完全一致");
for (a, b) in c1.iter().zip(&c2) {
assert!((a - b).abs() < 1e-6, "概率应一致: {a} vs {b}");
}
let _ = std::fs::remove_dir_all(&export_dir);
}
#[test]
fn pretrain_freeze_backbone_keeps_weights_frozen_during_training() {
let cfg0 = classify_cfg();
let vs = tch::nn::VarStore::new(Device::Cpu);
let _ = av_tasks::models::build_model(&vs.root(), &cfg0).expect("模型构建应成功");
let src_st = temp_dir("freeze-src").join("init.safetensors");
std::fs::create_dir_all(src_st.parent().unwrap()).unwrap();
let vars: Vec<(String, Tensor)> = vs.variables().into_iter().collect();
let refs: Vec<(&str, &Tensor)> = vars.iter().map(|(n, t)| (n.as_str(), t)).collect();
Tensor::write_safetensors(&refs, &src_st).expect("safetensors 导出应成功");
let mut cfg = classify_cfg();
cfg.run_id = "classify-pretrain-freeze".into();
cfg.pretrain = PretrainCfg {
enable: true,
weight_path: Some(src_st.clone()),
load_only_backbone: true,
freeze_backbone: true,
layer_map: None, };
cfg.train.epochs = 1;
let report = av_runtime::engine::train(&cfg).expect("冻结训练应成功");
let src_named = Tensor::read_safetensors(&src_st).expect("源应可读");
let src_of = |name: &str| {
src_named
.iter()
.find(|(n, _)| n == name)
.map(|(_, t)| t)
.expect("源应含该变量")
.shallow_clone()
};
let read_ckpt = |name: &str| {
let f = std::path::Path::new(&report.run_dir)
.join("best.ckpt")
.join(name.replace('.', "_"));
Tensor::load(&f)
.expect("checkpoint 张量应可读")
.to_device(Device::Cpu)
};
let bb_src = src_of("backbone.c1.weight");
let bb_trained = read_ckpt("backbone.c1.weight");
let bb_diff = (bb_src - bb_trained).abs().max().double_value(&[]);
assert_eq!(
bb_diff, 0.0,
"冻结的 backbone 权重必须与导入源逐位一致(diff={bb_diff})"
);
let head_src = src_of("head.fc.weight");
let head_trained = read_ckpt("head.fc.weight");
let head_diff = (head_src - head_trained).abs().max().double_value(&[]);
assert!(
head_diff > 0.0,
"未冻结的头部权重应被训练更新(diff={head_diff})"
);
let _ = std::fs::remove_file(&src_st);
}
#[test]
fn pretrain_safetensors_import_report_chain() {
let src = std::path::Path::new("../../data/pretrain/yolov8n_backbone.safetensors");
if !src.is_file() {
println!("[skip] {src:?} 不存在,跳过外部 safetensors 链路测试");
return;
}
let n_sources;
let report = {
let sources = weight_adapter::read_safetensors_all(src).expect("yolov8n 导出应可读");
n_sources = sources.len();
assert_eq!(n_sources, 42, "导出侧车应产出 42 个张量");
let cfg = classify_cfg();
let vs = tch::nn::VarStore::new(Device::Cpu);
let _ = av_tasks::models::build_model(&vs.root(), &cfg).expect("模型构建应成功");
let targets: Vec<(String, Vec<i64>)> = vs
.variables()
.into_iter()
.filter(|(n, _)| n.contains("backbone"))
.map(|(n, t)| (n, t.size()))
.collect();
assert_eq!(
targets.len(),
8,
"simple-cnn 分类模型应有 8 个 backbone 变量"
);
let map_path = temp_dir("layer-map").join("layer_map.toml");
std::fs::create_dir_all(map_path.parent().unwrap()).unwrap();
std::fs::write(
&map_path,
r#"# yolov8n stem conv → simple-cnn backbone 演示映射
[[entries]]
from = '^model\.0\.conv\.'
to = 'backbone.c1.'
[[entries]]
from = '^model\.1\.conv\.'
to = 'backbone.c2.'
[[entries]]
from = '^model\.3\.conv\.'
to = 'backbone.c3.'
# 故意演示形状不匹配:C2f 内部 1x1/3x3 分支与主干 conv 不同构
[[entries]]
from = '^model\.2\.m\.0\.cv1\.conv\.weight$'
to = 'backbone.c3.weight'
"#,
)
.unwrap();
let map = LayerMap::from_toml_path(&map_path).expect("层映射应可加载");
weight_adapter::adapt(sources, &map, &targets)
};
println!("[pretrain] AdaptReport: {}", report.summary());
for l in &report.loaded {
println!(" loaded: {} ← {} {:?}", l.target, l.source, l.shape);
}
for sm in &report.skipped_shape_mismatch {
println!(
" skipped_shape_mismatch: {} ← {} 期望 {:?} 得到 {:?}",
sm.target, sm.source, sm.expected, sm.got
);
}
let loaded_of = |target: &str| {
report
.loaded
.iter()
.find(|l| l.target == target)
.unwrap_or_else(|| panic!("应载入 {target}: {report:?}"))
};
assert_eq!(report.loaded.len(), 3, "{}", report.summary());
assert_eq!(
loaded_of("backbone.c1.weight").source,
"model.0.conv.weight"
);
assert_eq!(loaded_of("backbone.c1.weight").shape, vec![16, 3, 3, 3]);
assert_eq!(
loaded_of("backbone.c2.weight").source,
"model.1.conv.weight"
);
assert_eq!(
loaded_of("backbone.c3.weight").source,
"model.3.conv.weight"
);
assert_eq!(loaded_of("backbone.c3.weight").shape, vec![64, 32, 3, 3]);
assert_eq!(
report.skipped_shape_mismatch.len(),
1,
"{}",
report.summary()
);
assert_eq!(report.skipped_shape_mismatch[0].got, vec![16, 16, 3, 3]);
assert_eq!(
report.skipped_shape_mismatch[0].expected,
vec![64, 32, 3, 3]
);
assert_eq!(report.unexpected.len(), n_sources - report.loaded.len() - 1);
assert_eq!(report.missing.len(), 8 - report.loaded.len());
assert_eq!(report.tensors.len(), report.loaded.len());
}
#[test]
fn resnet18_pretrain_import() {
let src = std::path::Path::new("../../data/pretrain/resnet18_imagenet.safetensors");
if !src.is_file() {
println!(
"[skip] {src:?} 不存在,跳过 ResNet18 导入验证(导出:tools/export/export_resnet18.py)"
);
return;
}
let mut cfg = RunConfig::from_toml_str(include_str!("../../../configs/classify_resnet18.toml"))
.expect("classify_resnet18 配置必须合法");
cfg.output_dir = std::env::temp_dir().join("av-test-runs");
let vs = tch::nn::VarStore::new(Device::Cpu);
let model = av_tasks::models::build_model(&vs.root(), &cfg).expect("resnet18 模型应可装配");
let targets: Vec<(String, Vec<i64>)> = vs
.variables()
.into_iter()
.filter(|(n, _)| n.contains("backbone"))
.map(|(n, t)| (n, t.size()))
.collect();
assert_eq!(
targets.len(),
100,
"resnet18 骨干应有 100 个变量(20 conv + 40 BN 参数 + 40 BN 统计量)"
);
assert!(
targets
.iter()
.any(|(n, _)| n == "backbone.bn1.running_mean"),
"BN 统计量必须是 VarStore 命名变量(导入链路的前提)"
);
let sources = weight_adapter::read_safetensors_all(src).expect("torchvision 导出应可读");
assert_eq!(
sources.len(),
122,
"torchvision resnet18 state_dict 应 122 张量(100 骨干 + 20 num_batches_tracked + fc×2)"
);
let map_path = temp_dir("resnet18-map").join("resnet18_map.toml");
std::fs::create_dir_all(map_path.parent().unwrap()).unwrap();
std::fs::write(
&map_path,
include_str!("../../../configs/resnet18_map.toml"),
)
.unwrap();
let map = LayerMap::from_toml_path(&map_path).expect("层映射应可加载");
let sources_n = sources.len(); let report = weight_adapter::adapt(sources, &map, &targets);
println!("[pretrain][resnet18] AdaptReport: {}", report.summary());
for l in &report.loaded {
println!(" loaded: {} ← {} {:?}", l.target, l.source, l.shape);
}
assert!(
report.skipped_shape_mismatch.is_empty(),
"同名同构导入不应有形状不匹配: {:#?}",
report.skipped_shape_mismatch
);
let loaded_ratio = report.loaded.len() as f64 / targets.len() as f64;
assert!(
loaded_ratio >= 0.9,
"loaded 比例应 ≥ 90%,实际 {:.1}%({})",
loaded_ratio * 100.0,
report.summary()
);
assert_eq!(report.loaded.len(), 100, "{}", report.summary());
assert!(
report.missing.is_empty(),
"骨干目标应全部命中: {:?}",
report.missing
);
assert_eq!(report.unexpected.len(), sources_n - report.loaded.len());
assert!(
report
.loaded
.iter()
.any(|l| l.target == "backbone.bn1.running_mean"),
"BN running 统计量应被导入"
);
assert!(
report
.loaded
.iter()
.any(|l| l.target == "backbone.layer4.1.bn2.bias"),
"深层 BN 参数应被导入"
);
tch::no_grad(|| {
for (name, mut t) in vs.variables() {
if let Some((_, src_t)) = report.tensors.iter().find(|(n, _)| *n == name) {
t.copy_(src_t);
}
}
});
let conv1 = vs
.variables()
.remove("backbone.conv1.weight")
.expect("conv1.weight 应存在");
let conv1_src = &report
.tensors
.iter()
.find(|(n, _)| n == "backbone.conv1.weight")
.expect("conv1.weight 应已载入")
.1;
let diff = (&conv1 - conv1_src).abs().max().double_value(&[]);
assert_eq!(
diff, 0.0,
"ImageNet conv1.weight 写回必须逐位一致(diff={diff})"
);
let av_tasks::models::TaskModel::Classify(m) = &model else {
panic!("应为分类模型");
};
let x = Tensor::randn([2, 3, 64, 64], (Kind::Float, Device::Cpu));
let logits = m.logits(&x).expect("resnet18 forward 应成功");
assert_eq!(logits.size(), vec![2, 4], "logits 应为 [N, num_classes]");
let nan_cnt = logits.isnan().sum(tch::Kind::Int).int64_value(&[]);
assert_eq!(nan_cnt, 0, "ImageNet 权重前向不应产生 NaN");
let _ = std::fs::remove_dir_all(map_path.parent().unwrap());
}
#[test]
fn classification_trains_and_saves() {
let cfg = classify_cfg();
let report = av_runtime::engine::train(&cfg).expect("分类训练应成功");
assert_eq!(report.task, "classify");
assert!(
report.metric_value >= 0.95,
"合成分类应过拟合到 top1>=0.95,实际 {}",
report.metric_value
);
let weights = std::path::Path::new(&report.run_dir).join("best.ckpt");
assert!(weights.is_dir(), "checkpoint 目录应已落盘");
}
#[test]
fn detection_trains_and_localizes() {
let cfg = detect_cfg();
let report = av_runtime::engine::train(&cfg).expect("检测训练应成功");
assert_eq!(report.task, "detect");
assert!(
report.metric_value >= 0.55,
"合成检测 mean_iou 应 >= 0.55,实际 {}",
report.metric_value
);
let (_, r50) = report.secondary.clone().expect("应有 recall@0.5");
assert!(r50 >= 0.6, "recall@0.5 应 >= 0.6,实际 {r50}");
let weights = std::path::Path::new(&report.run_dir).join("best.ckpt");
let snapshot = std::path::Path::new(&report.run_dir).join("config.snapshot.toml");
let cfg2 = RunConfig::from_path(&snapshot).expect("快照配置应可读");
let mut rng = XorShift::new(99);
let (x, _, _) = av_tasks_testing_detect_batch(&mut rng, &cfg2);
let model = av_runtime::testing::load_model(&cfg2, &weights).expect("权重应可加载");
let x = x.to_device(av_runtime::engine::resolve_device(&cfg2));
let out = model.predict(&x, 0.25, 0.5).expect("推理应成功");
let av_tasks::models::PredictOutput::Detect { per_image } = out else {
panic!("应为检测输出");
};
assert!(
per_image.iter().any(|dets| !dets.is_empty()),
"过拟合后的模型应能在合成样本上给出候选框"
);
}
fn av_tasks_testing_detect_batch(
rng: &mut XorShift,
cfg: &RunConfig,
) -> (tch::Tensor, Vec<[f32; 4]>, Vec<u32>) {
let (num_classes, img_size) = match cfg.model.tasks.first() {
Some(av_core::config::TaskCfg::Detect(d)) => (d.num_classes as u32, d.img_size),
_ => unreachable!(),
};
av_runtime::testing::synthetic_detect(rng, 4, num_classes, img_size)
}
#[test]
fn coco8_map_pipeline() {
let root = std::path::Path::new("../../data/coco8");
if !root.is_dir() {
println!("[skip] ../../data/coco8 不存在,跳过 mAP 管线测试");
return;
}
let cfg = RunConfig::from_toml_str(include_str!("../../../configs/detect_coco8.toml"))
.expect("coco8 配置必须合法");
let img_size = match cfg.model.tasks.first() {
Some(av_core::config::TaskCfg::Detect(d)) => d.img_size,
_ => unreachable!("coco8 配置应为检测任务"),
};
let vs = tch::nn::VarStore::new(tch::Device::Cpu);
let model = av_tasks::models::build_model(&vs.root(), &cfg).expect("模型构建应成功");
let av_tasks::models::TaskModel::Detect(m) = &model else {
panic!("应为检测模型");
};
let val = av_runtime::dataset::load_yolo_dir(root, "val", img_size, tch::Device::Cpu, false)
.expect("coco8 val 应可加载");
assert_eq!(val.len(), 4, "coco8 val 应有 4 张图");
let mut ev = av_runtime::eval_map::CocoEvaluator::new();
for (gi, chunk) in val.chunks(8).enumerate() {
let x = av_runtime::dataset::stack_samples(chunk).expect("样本堆叠应成功");
let dets_per_image = m.predict(&x, 0.25, 0.5).expect("推理应成功");
assert_eq!(dets_per_image.len(), chunk.len());
for (i, s) in chunk.iter().enumerate() {
let gts: Vec<av_runtime::eval_map::GtBox> = s
.boxes
.iter()
.zip(&s.labels)
.map(|(b, &l)| {
av_runtime::eval_map::GtBox::new(
av_core::geometry::Aabb::new(b[0], b[1], b[2], b[3]),
l,
)
})
.collect();
ev.update((gi * 8 + i) as u32, &dets_per_image[i], >s);
}
}
let rep = ev.finalize();
println!("coco8 mAP 报告: {rep:?}");
assert_eq!(rep.num_images, 4);
assert!(rep.num_gts > 0, "coco8 val 应有真值框");
for (name, v) in [("map50", rep.map50), ("map50_95", rep.map50_95)] {
assert!(!v.is_nan(), "{name} 不应为 NaN");
assert!((0.0..=1.0).contains(&v), "{name} 应在 [0,1],实际 {v}");
}
}
#[test]
fn imagenette_pipeline() {
let root = std::path::Path::new("../../data/imagenette2-160");
if !root.is_dir() {
println!("[skip] {root:?} 不存在,跳过 ImageNette 管线测试");
return;
}
let cfg = RunConfig::from_toml_str(include_str!("../../../configs/classify_imagenette.toml"))
.expect("classify_imagenette 配置必须合法");
let (num_classes, img_size) = match cfg.model.tasks.first() {
Some(av_core::config::TaskCfg::Classify(c)) => (c.num_classes as u32, c.img_size),
_ => panic!("classify_imagenette 应为分类任务"),
};
assert_eq!(num_classes, 10);
assert_eq!(img_size, 64);
let (val, labels, class_map) =
av_runtime::dataset::load_imagefolder(root, "val", img_size, true, tch::Device::Cpu, false)
.expect("ImageNette val 应可加载");
assert_eq!(class_map.len(), 10, "ImageNette 应有 10 个 wnid 类");
assert!(val.len() > 1000, "val 应有数千张图,实际 {}", val.len());
assert_eq!(val.len(), labels.len(), "样本与标签应一一对应");
assert!(labels.iter().all(|&l| l < num_classes));
let n = val.len().min(64);
let subset: Vec<_> = val[..n].to_vec();
let x = av_runtime::dataset::stack_classify(&subset).expect("样本堆叠应成功");
assert_eq!(
x.size(),
vec![n as i64, 3, img_size as i64, img_size as i64]
);
let (mn, mx) = (x.min().double_value(&[]), x.max().double_value(&[]));
assert!(
(0.0..=1.0).contains(&mn) && (0.0..=1.0).contains(&mx),
"min={mn} max={mx}"
);
let vs = tch::nn::VarStore::new(tch::Device::Cpu);
let model = av_tasks::models::build_model(&vs.root(), &cfg).expect("模型构建应成功");
let av_tasks::models::TaskModel::Classify(m) = &model else {
panic!("应为分类模型");
};
let logits = m.logits(&x).expect("forward 应成功");
assert_eq!(logits.size(), vec![n as i64, num_classes as i64]);
let nan_cnt = logits.isnan().sum(tch::Kind::Int).int64_value(&[]);
assert_eq!(nan_cnt, 0, "logits 不应含 NaN");
let (pred, confs) = m.predict(&x).expect("predict 应成功");
assert_eq!(pred.len(), n);
assert_eq!(confs.len(), n);
let correct = pred
.iter()
.zip(&labels[..n])
.filter(|(a, b)| a == b)
.count();
let top1 = correct as f32 / n as f32;
assert!((0.0..=1.0).contains(&top1), "top1 应在 [0,1],实际 {top1}");
}
#[test]
fn seg_coco8_pipeline() {
let root = std::path::Path::new("../../data/coco8-seg");
if !root.is_dir() {
println!("[skip] {root:?} 不存在,跳过 seg 管线测试");
return;
}
let cfg = RunConfig::from_toml_str(include_str!("../../../configs/seg_coco8.toml"))
.expect("seg_coco8 配置必须合法");
let img_size = match cfg.model.tasks.first() {
Some(av_core::config::TaskCfg::Seg(s)) => s.img_size,
_ => panic!("seg_coco8 应为分割任务"),
};
let mask_hw = (img_size / 4) as usize;
let val = av_runtime::dataset::load_cocoseg_dir(root, "val", img_size, tch::Device::Cpu, false)
.expect("coco8-seg val 应可加载");
assert_eq!(val.len(), 4, "coco8 val 应有 4 张图");
assert!(val.iter().all(|s| s.masks.len() == s.labels.len()));
let n_inst = val.iter().map(|s| s.masks.len()).sum::<usize>();
assert!(n_inst > 0, "coco8-seg val 应有实例");
assert!(val
.iter()
.flat_map(|s| s.masks.iter())
.all(|m| m.len() == mask_hw * mask_hw && m.iter().all(|&v| v <= 1)));
let vs = tch::nn::VarStore::new(tch::Device::Cpu);
let model = av_tasks::models::build_model(&vs.root(), &cfg).expect("模型构建应成功");
let av_tasks::models::TaskModel::Seg(m) = &model else {
panic!("应为分割模型");
};
let batch_samples: Vec<_> = val.iter().take(2).cloned().collect();
let x = av_runtime::dataset::stack_seg_samples(&batch_samples).expect("样本堆叠应成功");
assert_eq!(x.size(), vec![2, 3, img_size as i64, img_size as i64]);
let loss = model
.loss(
&x,
&av_tasks::models::TrainBatch::Seg {
masks: batch_samples.iter().map(|s| s.masks.clone()).collect(),
labels: batch_samples.iter().map(|s| s.labels.clone()).collect(),
},
)
.expect("Seg 损失应成功");
assert!(loss.double_value(&[]).is_finite(), "loss 不应为 NaN");
loss.backward();
let per_image = m.predict(&x, 0.05, 0.5).expect("predict 应成功");
assert_eq!(per_image.len(), batch_samples.len());
for insts in &per_image {
for it in insts {
assert_eq!(it.mask.len(), mask_hw * mask_hw);
assert!(it.mask.iter().all(|&v| v <= 1), "掩码应二值");
assert!((0.0..=1.0).contains(&it.score), "score 应在 [0,1]");
}
}
println!(
"coco8-seg 管线:val {} 图 / {n_inst} 实例,loss={:.4}",
val.len(),
loss.double_value(&[])
);
}
#[test]
fn keypoint_coco8_pipeline() {
use av_tasks::models::{PredictOutput, TaskModel, TrainBatch};
let root = std::path::Path::new("../../data/coco8-pose");
if !root.is_dir() {
println!("[skip] {root:?} 不存在,跳过 keypoint 管线测试");
return;
}
let cfg = RunConfig::from_toml_str(include_str!("../../../configs/keypoint_coco8.toml"))
.expect("keypoint_coco8 配置必须合法");
let (num_keypoints, img_size) = match cfg.model.tasks.first() {
Some(av_core::config::TaskCfg::Keypoint(k)) => (k.num_keypoints, k.img_size),
_ => panic!("keypoint_coco8 应为关键点任务"),
};
assert_eq!(num_keypoints, 17);
let val =
av_runtime::dataset::load_cocopose_dir(root, "val", img_size, tch::Device::Cpu, false)
.expect("coco8-pose val 应可加载");
assert_eq!(val.len(), 4, "coco8-pose val 应有 4 张图");
assert!(val.iter().all(|s| s.boxes.len() == s.kpts.len()));
let n_inst = val.iter().map(|s| s.kpts.len()).sum::<usize>();
assert!(n_inst > 0, "coco8-pose val 应有实例");
assert!(val
.iter()
.flat_map(|s| s.kpts.iter())
.all(|gk| gk.len() == num_keypoints));
for s in &val {
for gk in &s.kpts {
for kp in gk {
if kp[2] > 0.0 {
assert!(
(0.0..=img_size as f32).contains(&kp[0])
&& (0.0..=img_size as f32).contains(&kp[1]),
"可见关键点应落在画布内: {kp:?}"
);
}
}
}
}
let vs = tch::nn::VarStore::new(tch::Device::Cpu);
let model = av_tasks::models::build_model(&vs.root(), &cfg).expect("模型构建应成功");
let batch_samples: Vec<_> = val.iter().take(2).cloned().collect();
let x = av_runtime::dataset::stack_kp_samples(&batch_samples).expect("样本堆叠应成功");
assert_eq!(x.size(), vec![2, 3, img_size as i64, img_size as i64]);
let loss = model
.loss(
&x,
&TrainBatch::Keypoint {
boxes: batch_samples.iter().map(|s| s.boxes.clone()).collect(),
kpts: batch_samples.iter().map(|s| s.kpts.clone()).collect(),
labels: batch_samples.iter().map(|s| s.labels.clone()).collect(),
},
)
.expect("Keypoint 损失应成功");
assert!(loss.double_value(&[]).is_finite(), "loss 不应为 NaN");
loss.backward();
let PredictOutput::Keypoint { per_image } =
model.predict(&x, 0.05, 0.5).expect("predict 应成功")
else {
panic!("应为关键点输出");
};
assert_eq!(per_image.len(), batch_samples.len());
for dets in &per_image {
for d in dets {
let kps = d.keypoints.as_ref().expect("实例必须带关键点");
assert_eq!(kps.len(), num_keypoints);
for kp in kps {
assert!((0.0..=img_size as f32).contains(&kp[0]), "x 应在画布内");
assert!((0.0..=img_size as f32).contains(&kp[1]), "y 应在画布内");
assert!(kp[2] == 0.0 || kp[2] == 2.0, "v 应二值化 0/2");
}
}
}
let TaskModel::Keypoint(m) = &model else {
panic!("应为关键点模型");
};
let (pck, mean_oks, n_vis, n_inst_eval) =
av_runtime::testing::eval_kp_samples(m, &val).expect("PCK 评测应成功");
assert!((0.0..=1.0).contains(&pck), "PCK 应在 [0,1],实际 {pck}");
assert!(
(0.0..=1.0).contains(&mean_oks),
"mean OKS 应在 [0,1],实际 {mean_oks}"
);
assert!(n_vis > 0 && n_inst_eval == n_inst);
println!(
"coco8-pose 管线:val {} 图 / {n_inst} 实例 / {n_vis} 可见点,loss={:.4} 随机权重 PCK={pck:.3}",
val.len(),
loss.double_value(&[])
);
}
use av_tasks::augment::AugmentPlan;
#[test]
fn augment_raw_encode_consistency_on_real_data() {
use av_runtime::dataset;
let root = std::path::Path::new("../../data/coco8-pose");
if !root.is_dir() {
println!("[skip] {root:?} 不存在,跳过增强一致性测试(关键点)");
return;
}
let img_size =
match RunConfig::from_toml_str(include_str!("../../../configs/keypoint_coco8.toml"))
.expect("keypoint_coco8 配置必须合法")
.model
.tasks
.first()
{
Some(av_core::config::TaskCfg::Keypoint(k)) => k.img_size,
_ => panic!("keypoint_coco8 应为关键点任务"),
};
let plain = dataset::load_cocopose_dir(root, "train", img_size, tch::Device::Cpu, false)
.expect("coco8-pose train 应可加载");
let raw = dataset::load_cocopose_dir_raw(root, "train").expect("raw 加载应成功");
assert_eq!(plain.len(), raw.len());
assert!(!plain.is_empty());
for (p, r) in plain.iter().zip(&raw) {
let enc = dataset::encode_keypoint_sample(
r,
img_size,
tch::Device::Cpu,
&AugmentPlan::none(),
false,
)
.expect("none 编码应成功");
let diff = (&p.x - &enc.x).abs().max().double_value(&[]);
assert_eq!(diff, 0.0, "none() 编码张量应与 plain 加载器逐位一致");
assert_eq!(p.boxes, enc.boxes, "none() 框应一致");
assert_eq!(p.kpts, enc.kpts, "none() 关键点应一致");
}
for (si, s) in [0.9f32, 1.0, 1.1].iter().enumerate() {
let plan = AugmentPlan {
flip: si != 1, scale: *s,
..AugmentPlan::none()
};
for r in &raw {
let enc = dataset::encode_keypoint_sample(r, img_size, tch::Device::Cpu, &plan, false)
.expect("增强编码应成功");
assert_eq!(enc.x.size(), vec![3, img_size as i64, img_size as i64]);
for gk in &enc.kpts {
assert_eq!(gk.len(), 17);
for kp in gk {
if kp[2] > 0.0 {
assert!(
(0.0..=img_size as f32).contains(&kp[0])
&& (0.0..=img_size as f32).contains(&kp[1]),
"翻转/缩放后可见点应仍在画布内: {kp:?}"
);
}
}
}
for b in &enc.boxes {
assert!(b[2] > 0.0 && b[3] > 0.0, "框宽高应保持正: {b:?}");
assert!(
(0.0..=img_size as f32).contains(&b[0])
&& (0.0..=img_size as f32).contains(&b[1]),
"框中心应在画布内: {b:?}"
);
}
}
}
let root = std::path::Path::new("../../data/coco8");
if !root.is_dir() {
println!("[skip] {root:?} 不存在,跳过增强一致性测试(检测)");
return;
}
let img_size =
match RunConfig::from_toml_str(include_str!("../../../configs/detect_coco8.toml"))
.expect("detect_coco8 配置必须合法")
.model
.tasks
.first()
{
Some(av_core::config::TaskCfg::Detect(d)) => d.img_size,
_ => panic!("detect_coco8 应为检测任务"),
};
let plain = dataset::load_yolo_dir(root, "train", img_size, tch::Device::Cpu, false)
.expect("coco8 train 应可加载");
let raw = dataset::load_yolo_dir_raw(root, "train").expect("raw 加载应成功");
assert_eq!(plain.len(), raw.len());
for (p, r) in plain.iter().zip(&raw) {
let enc = dataset::encode_detect_sample(
r,
img_size,
tch::Device::Cpu,
dataset::ResizeMode::Letterbox,
&AugmentPlan::none(),
false,
)
.expect("none 编码应成功");
let diff = (&p.x - &enc.x).abs().max().double_value(&[]);
assert_eq!(diff, 0.0, "none() 编码张量应与 plain 加载器逐位一致");
assert_eq!(p.boxes, enc.boxes);
assert_eq!(p.labels, enc.labels);
}
println!(
"增强一致性(真实数据):关键点 {} 图 / 检测 {} 图 none() 逐位一致,flip+scale 坐标同步",
plain.len(),
plain.len()
);
}
#[test]
fn augment_ab_keypoint_configs_parse() {
let base =
RunConfig::from_toml_str(include_str!("../../../configs/keypoint_coco8_ab_base.toml"))
.expect("A/B 基线配置必须合法");
let aug = RunConfig::from_toml_str(include_str!("../../../configs/keypoint_coco8_ab_aug.toml"))
.expect("A/B 增强配置必须合法");
let base_aug = &base.data.sources.train.tasks[0].augment;
assert!(!av_tasks::augment::has_strength(base_aug), "基线臂应无增强");
let aug_aug = &aug.data.sources.train.tasks[0].augment;
assert!(av_tasks::augment::has_strength(aug_aug), "增强臂应有增强");
assert!((aug_aug.flip - 0.5).abs() < 1e-6);
assert_eq!(aug_aug.hsv, [0.1, 0.1, 0.1]);
assert_eq!(aug_aug.scale_jitter, Some([0.9, 1.1]));
assert_eq!(aug_aug.close_last_epochs, 20);
let mut base2 = base.clone();
let mut aug2 = aug.clone();
base2.run_id = String::new();
aug2.run_id = String::new();
base2.data.sources.train.tasks[0].augment = Default::default();
aug2.data.sources.train.tasks[0].augment = Default::default();
assert_eq!(base2, aug2, "A/B 两臂除增强外必须同配置");
assert_eq!(base2.train.epochs, aug.train.epochs);
}
#[test]
fn coco8_avpack_pipeline_matches_dir_loader() {
use av_runtime::dataset;
let root = std::path::Path::new("../../data/coco8");
if !root.is_dir() {
println!("[skip] ../../data/coco8 不存在,跳过 avpack 管线测试");
return;
}
let tmp = temp_dir("avpack-pipeline");
std::fs::create_dir_all(&tmp).unwrap();
let pack_path = tmp.join("coco8.avpack");
let (n, _) = av_runtime::avpack::pack_dir(root, &pack_path).expect("coco8 打包应成功");
assert!(n > 0, "打包应产出条目");
let img_size =
match RunConfig::from_toml_str(include_str!("../../../configs/detect_coco8_avpack.toml"))
.expect("detect_coco8_avpack 配置必须合法")
.model
.tasks
.first()
{
Some(av_core::config::TaskCfg::Detect(d)) => d.img_size,
_ => panic!("detect_coco8_avpack 应为检测任务"),
};
let packed = dataset::load_yolo_avpack(&pack_path, "val", img_size, tch::Device::Cpu, false)
.expect("avpack val split 应可加载");
assert!(!packed.is_empty(), "avpack val 不应为空");
let dir = dataset::load_yolo_dir(root, "val", img_size, tch::Device::Cpu, false)
.expect("目录加载对照");
assert_eq!(packed.len(), dir.len(), "容器与目录加载样本数应一致");
for (i, (p, d)) in packed.iter().zip(&dir).enumerate() {
assert_eq!(p.labels, d.labels, "样本 {i} 类别应一致");
assert_eq!(p.boxes.len(), d.boxes.len(), "样本 {i} 框数应一致");
for (b, bref) in p.boxes.iter().zip(&d.boxes) {
for (v, r) in b.iter().zip(bref) {
assert!((v - r).abs() < 1e-3, "样本 {i} 框值 {v} vs 目录 {r}");
}
assert!(*b.first().unwrap() >= -0.5 && *b.last().unwrap() <= img_size as f32 + 0.5);
assert!(
b[0] <= b[2] + 1e-3 && b[1] <= b[3] + 1e-3,
"样本 {i} 框退化: {b:?}"
);
}
}
let batch = dataset::stack_samples(&packed).expect("堆批应成功");
assert_eq!(
batch.size(),
vec![packed.len() as i64, 3, img_size as i64, img_size as i64]
);
let _ = std::fs::remove_dir_all(&tmp);
println!(
"coco8 avpack 管线:容器 {} 条目,val {} 样本与目录加载一致",
n,
packed.len()
);
}
#[test]
fn coco8_avpack_train_smoke() {
let root = std::path::Path::new("../../data/coco8");
if !root.is_dir() {
println!("[skip] ../../data/coco8 不存在,跳过 avpack 训练冒烟");
return;
}
let tmp = temp_dir("avpack-train");
std::fs::create_dir_all(&tmp).unwrap();
let pack_path = tmp.join("coco8.avpack");
av_runtime::avpack::pack_dir(root, &pack_path).expect("coco8 打包应成功");
let mut cfg =
RunConfig::from_toml_str(include_str!("../../../configs/detect_coco8_avpack.toml"))
.expect("detect_coco8_avpack 配置必须合法");
cfg.run_id = format!("avpack-smoke-{}", std::process::id());
cfg.data.sources.train.avpack = Some(pack_path.clone());
cfg.data.sources.val.avpack = Some(pack_path.clone());
cfg.train.epochs = 2; cfg.output_dir = tmp.clone();
let rep = av_runtime::engine::train(&cfg).expect("avpack 训练冒烟应成功");
assert_eq!(rep.epochs, 2);
assert!(
rep.final_loss.is_finite(),
"loss 应有限: {}",
rep.final_loss
);
let _ = std::fs::remove_dir_all(&tmp);
println!(
"avpack 训练冒烟:task={} loss={:.4} run={}",
rep.task, rep.final_loss, rep.run_id
);
}
use av_core::traits::BaseBackbone;
use av_tasks::backbone_dino::DinoV2Backbone;
fn workspace_root() -> std::path::PathBuf {
std::path::PathBuf::from("../..").canonicalize().unwrap()
}
fn dinov2_ckpt_path() -> std::path::PathBuf {
workspace_root().join("data/pretrain/dinov2_small.safetensors")
}
fn tensor_max_abs_diff(a: &Tensor, b: &Tensor) -> f64 {
(a - b).abs().max().double_value(&[])
}
fn dino_has_nan(t: &Tensor) -> bool {
t.isnan().sum(tch::Kind::Float).double_value(&[]) > 0.0
}
#[test]
fn dinov2_pretrained_import_and_forward() {
let ckpt = dinov2_ckpt_path();
if !ckpt.exists() {
println!(
"[skip] {} 不存在,跳过 DINOv2 导入验证(下载:tools/export/export_dinov2.py)",
ckpt.display()
);
return;
}
let img = 224u32; let vs = tch::nn::VarStore::new(tch::Device::Cpu);
let cfg = av_core::config::BackboneCfg {
family: "dinov2".into(),
..Default::default()
};
let mut backbone = DinoV2Backbone::new(&vs.root(), &cfg, img).expect("224 应满足 14/32 整除");
let x = Tensor::randn(
[2, 3, img as i64, img as i64],
(tch::Kind::Float, tch::Device::Cpu),
);
let pooled_before = tch::no_grad(|| backbone.forward_pooled(&x).unwrap());
let stats = backbone
.load_dinov2_weights(&ckpt)
.expect("官方权重导入应成功");
println!("[dinov2] 导入统计: {}", stats.summary());
println!(
"[dinov2] 非骨干跳过样例: {:?}",
&stats.skipped_unrecognized[..stats.skipped_unrecognized.len().min(3)]
);
assert!(
stats.load_ratio() > 0.9,
"骨干导入比例应 > 90%,实际 {:.3}({})",
stats.load_ratio(),
stats.summary()
);
assert!(
stats.loaded >= 170,
"ViT-S/14 骨干应有 174 个目标张量,实际 {}",
stats.loaded
);
assert!(
stats.pos_embed_interpolated,
"518 训练网格 37×37 → 224 网格 16×16 必然触发 pos_embed 插值"
);
let pooled_after = tch::no_grad(|| backbone.forward_pooled(&x).unwrap());
let diff = tensor_max_abs_diff(&pooled_before, &pooled_after);
println!("[dinov2] 加载前后 pooled 最大差 = {diff:.4}");
assert!(diff > 1e-3, "导入后输出应有可测差异,实际 {diff}");
assert_eq!(pooled_after.size(), vec![2, 384]);
assert!(!dino_has_nan(&pooled_after), "pooled 不应含 NaN");
let pyramid = tch::no_grad(|| backbone.forward_features(&x).unwrap());
let strides: Vec<u32> = pyramid.levels.iter().map(|l| l.stride).collect();
assert_eq!(strides, vec![8, 16, 32]);
let sizes: Vec<Vec<i64>> = pyramid.levels.iter().map(|l| l.tensor.size()).collect();
assert_eq!(sizes[0], vec![2, 256, 28, 28]);
assert_eq!(sizes[1], vec![2, 256, 14, 14]);
assert_eq!(sizes[2], vec![2, 256, 7, 7]);
for l in &pyramid.levels {
assert!(
!dino_has_nan(&l.tensor),
"stride {} 特征不应含 NaN",
l.stride
);
}
println!("[dinov2] 金字塔形状: {sizes:?}");
}
#[test]
fn dinov2_forward_multi_resolution() {
let vs = tch::nn::VarStore::new(tch::Device::Cpu);
let cfg = av_core::config::BackboneCfg::default();
let backbone = DinoV2Backbone::new(&vs.root(), &cfg, 224).unwrap();
for img in [224i64, 448] {
let x = Tensor::randn([1, 3, img, img], (tch::Kind::Float, tch::Device::Cpu));
let pooled = tch::no_grad(|| backbone.forward_pooled(&x).unwrap());
assert_eq!(pooled.size(), vec![1, 384], "img={img}");
assert!(!dino_has_nan(&pooled), "img={img}");
let g = img / 14;
assert_eq!(g * g + 1, 1 + g * g); }
}
#[test]
fn dinov2_latency_report() {
let device = tch::Device::cuda_if_available();
let img = 448u32;
let vs = tch::nn::VarStore::new(device);
let cfg = av_core::config::BackboneCfg::default();
let backbone = DinoV2Backbone::new(&vs.root(), &cfg, img).unwrap();
let x = Tensor::randn([1, 3, img as i64, img as i64], (tch::Kind::Float, device));
for _ in 0..2 {
let _ = tch::no_grad(|| backbone.forward_pooled(&x).unwrap());
}
if device == tch::Device::Cpu {
std::thread::sleep(std::time::Duration::from_millis(100));
}
let runs = 5;
let mut times = Vec::new();
for _ in 0..runs {
let t0 = std::time::Instant::now();
let _ = tch::no_grad(|| backbone.forward_pooled(&x).unwrap());
times.push(t0.elapsed().as_secs_f64() * 1000.0);
}
let min = times.iter().cloned().fold(f64::INFINITY, f64::min);
let avg = times.iter().sum::<f64>() / runs as f64;
println!(
"[dinov2-latency] device={device:?} img={img} 单图 pooled 前向:min={min:.1}ms avg={avg:.1}ms({runs} 次)"
);
assert!(times.iter().all(|t| t.is_finite() && *t > 0.0));
}
fn matrix_arm_cfg(toml: &'static str, arm: &str) -> RunConfig {
let root = workspace_root();
let mut cfg = RunConfig::from_toml_str(toml).expect("矩阵臂配置必须合法");
cfg.run_id = format!("matrix-{arm}-{}", std::process::id());
cfg.output_dir = temp_dir("matrix");
cfg.data.sources.train.dir = Some(root.join("data/coco8"));
cfg.data.sources.val.dir = Some(root.join("data/coco8"));
cfg
}
fn matrix_initial_loss(
cfg: &RunConfig,
x: &Tensor,
batch: &av_tasks::models::TrainBatch,
pretrain: bool,
) -> f64 {
tch::manual_seed(2026);
let vs = tch::nn::VarStore::new(Device::Cpu);
let model = av_tasks::models::build_model(&vs.root(), cfg).expect("矩阵臂装配应成功");
if pretrain {
let weight = workspace_root().join("data/pretrain/resnet18_imagenet.safetensors");
let map_path = workspace_root().join("configs/resnet18_map.toml");
let sources = weight_adapter::read_safetensors_all(&weight).expect("ImageNet 导出应可读");
let map = LayerMap::from_toml_path(&map_path).expect("resnet18_map 应可读");
let targets: Vec<(String, Vec<i64>)> = vs
.variables()
.into_iter()
.filter(|(n, _)| n.contains("backbone")) .map(|(n, t)| (n, t.size()))
.collect();
let report = weight_adapter::adapt(sources, &map, &targets);
println!(
"[matrix] ImageNet 导入(初始 loss 前): {}",
report.summary()
);
assert_eq!(
report.loaded.len(),
100,
"resnet18 骨干应 100/100 全量导入: {}",
report.summary()
);
assert!(report.skipped_shape_mismatch.is_empty());
tch::no_grad(|| {
for (name, mut t) in vs.variables() {
if let Some((_, src)) = report.tensors.iter().find(|(n, _)| *n == name) {
t.copy_(src);
}
}
});
}
model
.loss(x, batch)
.expect("初始 loss 应成功")
.double_value(&[])
}
#[test]
fn pretrain_ab_matrix_coco8_detect_three_arms() {
let root = std::path::Path::new("../../data/coco8");
if !root.is_dir() {
println!("[skip] ../../data/coco8 不存在,跳过预训练 A/B 矩阵");
return;
}
let weight = workspace_root().join("data/pretrain/resnet18_imagenet.safetensors");
if !weight.is_file() {
println!("[skip] {weight:?} 不存在,跳过预训练 A/B 矩阵(导出:tools/export/export_resnet18.py)");
return;
}
let train = av_runtime::dataset::load_yolo_dir(root, "train", 320, Device::Cpu, false)
.expect("coco8 train 应可加载");
assert_eq!(train.len(), 4, "coco8 train 应有 4 张图");
let x = av_runtime::dataset::stack_samples(&train).expect("堆批应成功");
assert_eq!(
x.size(),
vec![4, 3, 320, 320],
"letterbox 后应为 [N,3,320,320]"
);
let batch = av_tasks::models::TrainBatch::Detect {
boxes: train.iter().map(|s| s.boxes.clone()).collect(),
labels: train.iter().map(|s| s.labels.clone()).collect(),
};
let cfg_a1 = matrix_arm_cfg(
include_str!("../../../configs/detect_coco8_matrix_arm1_simplecnn_scratch.toml"),
"arm1",
);
let cfg_a2 = matrix_arm_cfg(
include_str!("../../../configs/detect_coco8_matrix_arm2_resnet_scratch.toml"),
"arm2",
);
let mut cfg_a3 = matrix_arm_cfg(
include_str!("../../../configs/detect_coco8_matrix_arm3_resnet_imagenet.toml"),
"arm3",
);
cfg_a3.pretrain.weight_path = Some(weight.clone());
cfg_a3.pretrain.layer_map = Some(workspace_root().join("configs/resnet18_map.toml"));
let loss_a1 = matrix_initial_loss(&cfg_a1, &x, &batch, false);
let loss_a2 = matrix_initial_loss(&cfg_a2, &x, &batch, false);
let loss_a3 = matrix_initial_loss(&cfg_a3, &x, &batch, true);
println!(
"[matrix] 初始 loss:A1 simple-cnn 从零={loss_a1:.4} A2 resnet18 从零={loss_a2:.4} A3 resnet18+ImageNet={loss_a3:.4}"
);
assert!(
(loss_a3 - loss_a2).abs() > 1e-3,
"预训练臂初始 loss 应相对从零臂可测地不同(ImageNet 权重真实生效的信号):\
A3={loss_a3:.4} vs A2={loss_a2:.4}"
);
{
let vs = tch::nn::VarStore::new(Device::Cpu);
let bb_cfg = av_core::config::BackboneCfg {
family: "resnet18".into(),
..Default::default()
};
let bb = av_tasks::backbone_resnet::ResNetBackbone::new(&(vs.root() / "backbone"), &bb_cfg)
.expect("resnet18 骨干装配应成功");
let py = tch::no_grad(|| bb.forward_features(&x).expect("金字塔前向应成功"));
let expect = [(4u32, 64i64, 80i64), (8, 128, 40), (16, 256, 20)];
assert_eq!(py.levels.len(), expect.len());
for (lv, (s, c, hw)) in py.levels.iter().zip(expect) {
assert_eq!(lv.stride, s);
assert_eq!(lv.tensor.size(), vec![4, c, hw, hw], "stride {s} 层形状");
}
}
let mut rows = Vec::new();
for (name, cfg) in [
("A1 simple-cnn 从零", cfg_a1.clone()),
("A2 resnet18 从零", cfg_a2.clone()),
("A3 resnet18+ImageNet", cfg_a3.clone()),
] {
let rep =
av_runtime::engine::train(&cfg).unwrap_or_else(|e| panic!("{name} 训练应成功: {e}"));
assert_eq!(rep.epochs, 60, "{name} 应跑满 60 epochs");
assert!(rep.final_loss.is_finite(), "{name} 最终 loss 应有限");
assert!(
(0.0..=1.0).contains(&rep.metric_value),
"{name} mIoU 应在 [0,1],实际 {}",
rep.metric_value
);
let r50 = rep.secondary.as_ref().map(|(_, v)| *v).unwrap_or(0.0);
println!(
"[matrix] {name}: final_loss={:.4} mIoU={:.4} R@0.5={r50:.4}",
rep.final_loss, rep.metric_value
);
rows.push((name, rep.final_loss, rep.metric_value));
}
println!("\n==== 预训练 A/B 矩阵(coco8 检测 60ep CPU img320)====");
println!("臂 | 初始 loss | 最终 loss | 最终 mIoU");
println!(
"A1 simple-cnn 从零 | {loss_a1:9.4} | {:9.4} | {:9.4}",
rows[0].1, rows[0].2
);
println!(
"A2 resnet18 从零 | {loss_a2:9.4} | {:9.4} | {:9.4}",
rows[1].1, rows[1].2
);
println!(
"A3 resnet18+ImageNet | {loss_a3:9.4} | {:9.4} | {:9.4}",
rows[2].1, rows[2].2
);
for (dir, _) in [
(cfg_a1.output_dir.clone(), ()),
(cfg_a2.output_dir.clone(), ()),
(cfg_a3.output_dir.clone(), ()),
] {
let _ = std::fs::remove_dir_all(&dir);
}
}
fn ab2_arm_initial_loss(
cfg: &RunConfig,
x: &Tensor,
batch: &av_tasks::models::TrainBatch,
pretrain: bool,
weight: &std::path::Path,
seed: u64,
) -> f64 {
tch::manual_seed(seed as i64);
let vs = tch::nn::VarStore::new(Device::Cpu);
let model = av_tasks::models::build_model(&vs.root(), cfg).expect("ab2 臂装配应成功");
if pretrain {
let sources = weight_adapter::read_safetensors_all(weight).expect("ImageNet 导出应可读");
let map = LayerMap::from_toml_path(&workspace_root().join("configs/resnet18_map.toml"))
.expect("resnet18_map 应可读");
let targets: Vec<(String, Vec<i64>)> = vs
.variables()
.into_iter()
.filter(|(n, _)| n.contains("backbone")) .map(|(n, t)| (n, t.size()))
.collect();
let report = weight_adapter::adapt(sources, &map, &targets);
assert_eq!(
report.loaded.len(),
100,
"ab2 预训练臂骨干应 100/100 全量导入: {}",
report.summary()
);
let matched = report.tensors;
tch::no_grad(|| {
for (name, mut t) in vs.variables() {
if let Some((_, src)) = matched.iter().find(|(n, _)| *n == name) {
t.copy_(src);
}
}
});
}
model
.loss(x, batch)
.expect("ab2 初始 loss 应成功")
.double_value(&[])
}
#[test]
fn resnet18_ab_norm_bn_phase2() {
let root = std::path::Path::new("../../data/coco8");
if !root.is_dir() {
println!("[skip] ../../data/coco8 不存在,跳过第二期 A/B 复验");
return;
}
let weight = workspace_root().join("data/pretrain/resnet18_imagenet.safetensors");
if !weight.is_file() {
println!("[skip] {weight:?} 不存在,跳过第二期 A/B 复验(导出:tools/export/export_resnet18.py)");
return;
}
let cfg_base = matrix_arm_cfg(
include_str!("../../../configs/resnet18_ab_base.toml"),
"ab2-base",
);
let mut cfg_pre = matrix_arm_cfg(
include_str!("../../../configs/resnet18_ab_pretrain.toml"),
"ab2-pre",
);
let mut cfg_frz = matrix_arm_cfg(
include_str!("../../../configs/resnet18_ab_pretrain_frozen.toml"),
"ab2-frozen",
);
cfg_pre.pretrain.weight_path = Some(weight.clone());
cfg_pre.pretrain.layer_map = Some(workspace_root().join("configs/resnet18_map.toml"));
cfg_frz.pretrain.weight_path = Some(weight.clone());
cfg_frz.pretrain.layer_map = Some(workspace_root().join("configs/resnet18_map.toml"));
let train_base = av_runtime::dataset::load_yolo_dir(root, "train", 320, Device::Cpu, false)
.expect("coco8 train([0,1] 域)应可加载");
let train_pre = av_runtime::dataset::load_yolo_dir(root, "train", 320, Device::Cpu, true)
.expect("coco8 train(ImageNet 域)应可加载");
assert_eq!(train_base.len(), 4);
let x_base = av_runtime::dataset::stack_samples(&train_base).expect("堆批应成功");
let x_pre = av_runtime::dataset::stack_samples(&train_pre).expect("堆批应成功");
let domain_shift = tch::no_grad(|| (&x_pre - &x_base).abs().max().double_value(&[]));
println!("[ab2] 两域输入最大绝对差 = {domain_shift:.4}(ImageNet 归一化生效信号)");
assert!(
domain_shift > 0.5,
"imagenet_norm = true 必须改变输入域(差值 {domain_shift})"
);
let batch_base = av_tasks::models::TrainBatch::Detect {
boxes: train_base.iter().map(|s| s.boxes.clone()).collect(),
labels: train_base.iter().map(|s| s.labels.clone()).collect(),
};
let batch_pre = av_tasks::models::TrainBatch::Detect {
boxes: train_pre.iter().map(|s| s.boxes.clone()).collect(),
labels: train_pre.iter().map(|s| s.labels.clone()).collect(),
};
let seeds: [u64; 5] = [2026, 3, 7, 42, 99];
let (mut sum_base, mut sum_pre) = (0.0f64, 0.0f64);
println!("[ab2] 逐 seed 初始 loss(同 seed 同头初始化配对;两臂只差骨干权重与输入域):");
for s in seeds {
let loss_b = ab2_arm_initial_loss(&cfg_base, &x_base, &batch_base, false, &weight, s);
let loss_p = ab2_arm_initial_loss(&cfg_pre, &x_pre, &batch_pre, true, &weight, s);
println!(
"[ab2] seed={s:<5} base={loss_b:9.4} pretrain={loss_p:9.4} Δ={:+.4}",
loss_p - loss_b
);
sum_base += loss_b;
sum_pre += loss_p;
}
let loss_base = sum_base / seeds.len() as f64;
let loss_pre = sum_pre / seeds.len() as f64;
let delta_pct = (loss_pre - loss_base) / loss_base * 100.0;
println!(
"[ab2] 初始 loss 均值:base(resnet18 从零 [0,1])={loss_base:.4} pretrain(resnet18+ImageNet 归一化域)={loss_pre:.4} 差={delta_pct:+.2}%"
);
if loss_pre >= loss_base {
println!("[ab2] 预训练臂初始 loss 均值未低于从零臂,dump 两臂前三步 loss 分解:");
for (name, cfg, x, batch, pre) in [
("base", &cfg_base, &x_base, &batch_base, false),
("pretrain", &cfg_pre, &x_pre, &batch_pre, true),
] {
tch::manual_seed(2026);
let vs = tch::nn::VarStore::new(Device::Cpu);
let model = av_tasks::models::build_model(&vs.root(), cfg).expect("装配应成功");
if pre {
let sources = weight_adapter::read_safetensors_all(&weight).expect("权重应可读");
let map =
LayerMap::from_toml_path(&workspace_root().join("configs/resnet18_map.toml"))
.expect("resnet18_map 应可读");
let targets: Vec<(String, Vec<i64>)> = vs
.variables()
.into_iter()
.filter(|(n, _)| n.contains("backbone"))
.map(|(n, t)| (n, t.size()))
.collect();
let report = weight_adapter::adapt(sources, &map, &targets);
let matched = report.tensors;
tch::no_grad(|| {
for (name, mut t) in vs.variables() {
if let Some((_, src)) = matched.iter().find(|(n, _)| *n == name) {
t.copy_(src);
}
}
});
}
for step in 0..3 {
av_tasks::models::LOSS_DEBUG
.with(|d| *d.borrow_mut() = Some(format!("[ab2 {name} step{step}]")));
let l = model.loss(x, batch).expect("loss 应成功").double_value(&[]);
println!("[ab2] {name} step{step} total={l:.4}");
}
}
}
let loss_base_2026 =
ab2_arm_initial_loss(&cfg_base, &x_base, &batch_base, false, &weight, 2026);
let loss_pre_2026 = ab2_arm_initial_loss(&cfg_pre, &x_pre, &batch_pre, true, &weight, 2026);
assert!(
(loss_pre_2026 - loss_base_2026).abs() > 1e-3,
"预训练臂初始 loss 应相对从零臂可测地不同(100/100 导入权重真实生效信号):\
pretrain={loss_pre_2026:.4} vs base={loss_base_2026:.4}"
);
let mut rows = Vec::new();
for (name, cfg) in [
("base resnet18 从零", cfg_base.clone()),
("pretrain resnet18+ImageNet", cfg_pre.clone()),
("pretrain 冻结骨干(BN eval)", cfg_frz.clone()),
] {
let rep =
av_runtime::engine::train(&cfg).unwrap_or_else(|e| panic!("{name} 训练应成功: {e}"));
assert_eq!(rep.epochs, 120, "{name} 应跑满 120 epochs");
assert!(rep.final_loss.is_finite(), "{name} 最终 loss 应有限");
assert!(
(0.0..=1.0).contains(&rep.metric_value),
"{name} mIoU 应在 [0,1],实际 {}",
rep.metric_value
);
let r50 = rep.secondary.as_ref().map(|(_, v)| *v).unwrap_or(0.0);
rows.push((name, rep.final_loss, rep.metric_value, r50));
}
println!("\n==== 预训练 A/B 第二期(归一化+BN 修复后复验,coco8 检测 120ep CPU img320)====");
println!("臂 | 初始 loss(seed2026) | 最终 loss | 最终 mIoU | R@0.5");
println!(
"base resnet18 从零 | {loss_base_2026:9.4} | {:9.4} | {:9.4} | {:.3}",
rows[0].1, rows[0].2, rows[0].3
);
println!(
"pretrain resnet18+ImageNet | {loss_pre_2026:9.4} | {:9.4} | {:9.4} | {:.3}",
rows[1].1, rows[1].2, rows[1].3
);
println!(
"pretrain 冻结骨干(BN eval) | - | {:9.4} | {:9.4} | {:.3}",
rows[2].1, rows[2].2, rows[2].3
);
for dir in [
cfg_base.output_dir.clone(),
cfg_pre.output_dir.clone(),
cfg_frz.output_dir.clone(),
] {
let _ = std::fs::remove_dir_all(&dir);
}
}