use std::fs;
use std::path::{Path, PathBuf};
use std::sync::Arc;
use rayon::prelude::*;
use serde::Serialize;
use tch::nn::OptimizerConfig;
use tch::nn::VarStore;
use tch::{Device, Kind, Tensor};
use av_core::config::{DataPipeline, RunConfig, TaskCfg, TaskKind};
use av_core::error::{AvError, AvResult};
use av_core::geometry::Aabb;
use av_pretrain::av_weight as av_weight_store;
use av_pretrain::weight_adapter::{self, LayerMap};
use av_tasks::augment::{self, AugmentPlan};
use av_tasks::mask::MaskSummary;
use av_tasks::models::{
build_model, KeypointModel, PredictOutput, SegModel, TaskModel, TrainBatch,
};
use av_tasks::rng::XorShift;
use crate::dataset::{self, ClassifySample, KeypointSample, SampleTensor, SegSample};
use crate::eval_map::{CocoEvaluator, GtBox};
const EVAL_BATCH: i64 = 256;
const STEPS_PER_EPOCH: usize = 16;
const CKPT_DIR: &str = "best.ckpt";
struct EpochLossAcc(Option<Tensor>);
impl EpochLossAcc {
fn new() -> Self {
Self(None)
}
fn add(&mut self, loss: &Tensor) {
let l = loss.detach();
self.0 = Some(match &self.0 {
Some(t) => t + &l,
None => l,
});
}
fn mean(&self, steps: usize) -> f32 {
match &self.0 {
Some(t) if steps > 0 => (t / steps as f64).double_value(&[]) as f32,
_ => 0.0,
}
}
}
struct GradScaler {
scale: f64,
growth_factor: f64,
backoff_factor: f64,
growth_interval: u32,
steps_since_growth: u32,
}
impl GradScaler {
fn new() -> Self {
Self {
scale: 65536.0,
growth_factor: 2.0,
backoff_factor: 0.5,
growth_interval: 2000,
steps_since_growth: 0,
}
}
fn scaled(&self, loss: &Tensor) -> Tensor {
loss.to_kind(Kind::Float) * self.scale
}
fn unscale_and_check(&mut self, vars: &[Tensor]) -> bool {
let mut max_t: Option<Tensor> = None;
for v in vars {
let g = v.grad();
if !g.defined() {
continue;
}
let m = g.abs().max();
max_t = Some(match &max_t {
Some(acc) => acc.maximum(&m),
None => m,
});
}
let worst = match &max_t {
Some(t) => t.double_value(&[]),
None => 0.0,
};
if worst.is_nan() || worst.is_infinite() {
return false;
}
tch::no_grad(|| {
for v in vars {
let mut g = v.grad();
if g.defined() {
g /= self.scale;
}
}
});
true
}
fn update(&mut self, stepped: bool) {
if stepped {
self.steps_since_growth += 1;
if self.steps_since_growth >= self.growth_interval {
self.scale *= self.growth_factor;
self.steps_since_growth = 0;
}
} else {
self.scale *= self.backoff_factor;
self.steps_since_growth = 0;
}
}
}
struct WeightEma {
shadow: Vec<(String, Tensor)>,
decay: f32,
updates: u64,
}
impl WeightEma {
fn new(vs: &VarStore, decay: f32) -> Self {
let mut shadow: Vec<(String, Tensor)> = vs.variables().into_iter().collect();
shadow.sort_by(|a, b| a.0.cmp(&b.0));
Self {
shadow,
decay,
updates: 0,
}
}
fn update(&mut self, vs: &VarStore) {
self.updates += 1;
let t = self.updates as f64;
let d = (self.decay as f64).min((1.0 + t) / (10.0 + t));
let mut vars = vs.variables();
tch::no_grad(|| {
for (name, s) in self.shadow.iter_mut() {
if let Some(v) = vars.get_mut(name.as_str()) {
*s = &*s * d + &*v * (1.0 - d);
}
}
});
}
fn apply_to(&self, vs: &VarStore) -> Vec<(String, Tensor)> {
let mut vars = vs.variables();
let mut saved = Vec::new();
for (name, s) in &self.shadow {
if let Some(v) = vars.get_mut(name.as_str()) {
saved.push((name.clone(), v.copy()));
v.set_data(s);
}
}
saved
}
fn restore(&self, vs: &VarStore, saved: &[(String, Tensor)]) {
let mut vars = vs.variables();
for (name, old) in saved {
if let Some(v) = vars.get_mut(name.as_str()) {
v.set_data(old);
}
}
}
fn snapshot_cpu(&self) -> Vec<(String, Tensor)> {
self.shadow
.iter()
.map(|(n, t)| (n.clone(), t.to_device(Device::Cpu)))
.collect()
}
}
fn save_named_variables(vars: &[(String, Tensor)], dir: &Path) -> AvResult<()> {
fs::create_dir_all(dir)?;
for (name, t) in vars {
t.save(dir.join(ckpt_file_name(name)))
.map_err(|e| AvError::train(format!("保存张量 {name} 失败: {e}")))?;
}
Ok(())
}
fn optimizer_step(
amp: bool,
scaler: &mut GradScaler,
vs: &VarStore,
opt: &mut tch::nn::Optimizer,
ema: &mut WeightEma,
grad_clip: f32,
) -> bool {
if !amp {
if grad_clip > 0.0 {
opt.clip_grad_norm(grad_clip as f64);
}
opt.step();
opt.zero_grad();
ema.update(vs);
return true;
}
let vars = vs.trainable_variables();
if scaler.unscale_and_check(&vars) {
if grad_clip > 0.0 {
opt.clip_grad_norm(grad_clip as f64);
}
opt.step();
opt.zero_grad();
ema.update(vs);
scaler.update(true);
true
} else {
opt.zero_grad();
scaler.update(false);
false
}
}
fn save_checkpoint(vs: &VarStore, dir: &Path) -> AvResult<()> {
fs::create_dir_all(dir)?;
for (name, t) in vs.variables() {
let f = dir.join(ckpt_file_name(&name));
t.save(&f)
.map_err(|e| AvError::train(format!("保存张量 {name} 失败: {e}")))?;
}
Ok(())
}
fn save_checkpoint_epoch(vs: &VarStore, dir: &Path, epoch: u32) -> AvResult<()> {
save_checkpoint(vs, dir)?;
fs::write(dir.join("meta.json"), format!("{{\"epoch\":{epoch}}}"))
.map_err(|e| AvError::train(format!("写 checkpoint 元数据失败: {e}")))
}
fn read_checkpoint_epoch(dir: &Path) -> Option<u32> {
let raw = fs::read_to_string(dir.join("meta.json")).ok()?;
let v: serde_json::Value = serde_json::from_str(&raw).ok()?;
v.get("epoch")?.as_u64().map(|e| e as u32)
}
fn load_checkpoint(vs: &mut VarStore, dir: &Path) -> AvResult<()> {
tch::no_grad(|| {
for (name, mut t) in vs.variables() {
let f = dir.join(ckpt_file_name(&name));
let loaded = Tensor::load(&f)
.map_err(|e| AvError::train(format!("读取张量 {name} 失败: {e}")))?;
if loaded.size() != t.size() {
return Err(AvError::train(format!(
"checkpoint 与当前模型结构不匹配: {name} 期望 {:?} 得到 {:?}(模型结构变更后请重训)",
t.size(),
loaded.size()
)));
}
t.copy_(&loaded);
}
Ok(())
})
}
pub fn load_checkpoint_dir(vs: &mut VarStore, dir: &Path) -> AvResult<()> {
load_checkpoint(vs, dir)
}
pub fn export_checkpoint(vs: &VarStore, dir: &Path, fmt: &str) -> AvResult<()> {
match fmt {
"safetensors" | "st" => {
if !dir.to_string_lossy().ends_with(".safetensors") {
return Err(AvError::train(
"safetensors 导出的 out 必须以 .safetensors 结尾",
));
}
let named: Vec<(String, Tensor)> = vs.variables().into_iter().collect();
let refs: Vec<(&str, &Tensor)> = named.iter().map(|(n, t)| (n.as_str(), t)).collect();
tch::Tensor::write_safetensors(&refs, dir)
.map_err(|e| AvError::train(format!("safetensors 导出失败: {e}")))?;
Ok(())
}
"torch" | "ckpt" => save_checkpoint(vs, dir),
other => Err(AvError::train(format!(
"不支持的导出格式: {other}(safetensors | torch)"
))),
}
}
fn ckpt_file_name(name: &str) -> String {
name.replace(['/', '\\'], "__").replace('.', "_")
}
#[derive(Debug, Clone, Serialize)]
pub struct TrainReport {
pub run_id: String,
pub task: String,
pub epochs: u32,
pub final_loss: f32,
pub metric: String,
pub metric_value: f32,
pub secondary: Option<(String, f32)>,
pub run_dir: String,
}
pub fn train(cfg: &RunConfig) -> AvResult<TrainReport> {
train_impl(cfg, false)
}
pub fn train_resumed(cfg: &RunConfig) -> AvResult<TrainReport> {
match cfg.model.tasks.first() {
Some(TaskCfg::Seg(_)) => train_impl(cfg, true),
_ => Err(AvError::config(
"--resume 目前仅 seg 训练路径实现(其余任务的周期存盘尚未接入)",
)),
}
}
fn train_impl(cfg: &RunConfig, resume: bool) -> AvResult<TrainReport> {
let run_id = cfg.effective_run_id();
let run_dir = cfg.output_dir.join(&run_id);
fs::create_dir_all(&run_dir)?;
if resume {
println!("[resume] 续训模式:保留 runs/{run_id}/ 既有产物");
} else {
let _ = fs::remove_dir_all(run_dir.join(CKPT_DIR));
let _ = fs::remove_file(run_dir.join(crate::metrics::METRICS_FILE));
}
fs::write(run_dir.join("config.snapshot.toml"), cfg.snapshot_toml()?)?;
if cfg.train.deterministic {
tch::manual_seed(cfg.seed as i64);
println!(
"[deterministic] 已播种 torch 全局 RNG(seed={}),权重初始化可复现",
cfg.seed
);
}
let report = match cfg.model.tasks.first() {
Some(TaskCfg::Classify(_)) => train_classify(cfg, &run_id, &run_dir)?,
Some(TaskCfg::Detect(d)) if !d.obb_mode => match cfg.data.pipeline {
DataPipeline::Synthetic => train_detect_synthetic(cfg, &run_id, &run_dir)?,
DataPipeline::Dir | DataPipeline::AvPack => train_detect_yolo(cfg, &run_id, &run_dir)?,
},
Some(TaskCfg::Detect(d)) if d.obb_mode => match cfg.data.pipeline {
DataPipeline::Dir => train_detect_obb(cfg, &run_id, &run_dir)?,
_ => {
return Err(AvError::config(
"OBB 训练需要 data.pipeline = \"dir\"(DOTA 格式,见 configs/detect_obb_dota8.toml)",
))
}
},
Some(TaskCfg::Seg(_)) => match cfg.data.pipeline {
DataPipeline::Dir => train_seg(cfg, &run_id, &run_dir, resume)?,
_ => {
return Err(AvError::config(
"Seg 训练需要 data.pipeline = \"dir\"(COCO 分割格式,见 configs/seg_coco8.toml)",
))
}
},
Some(TaskCfg::Keypoint(_)) => match cfg.data.pipeline {
DataPipeline::Dir => train_keypoint(cfg, &run_id, &run_dir)?,
_ => {
return Err(AvError::config(
"Keypoint 训练需要 data.pipeline = \"dir\"(COCO 姿态格式,见 configs/keypoint_coco8.toml)",
))
}
},
_ => {
return Err(AvError::config(
"该任务类型在 v0.1 引擎未支持(见 PLAN §8 里程碑)",
))
}
};
fs::write(
run_dir.join("report.json"),
serde_json::to_string_pretty(&report)
.map_err(|e| AvError::train(format!("报告序列化失败: {e}")))?,
)?;
Ok(report)
}
pub fn infer(cfg: &RunConfig, weights: &Path, input: &Path) -> AvResult<serde_json::Value> {
let (model, device) = load_model(cfg, weights)?;
let (x, lb, orig_w, orig_h) = dataset::decode_image_with_meta(
input,
model.img_size(),
Device::Cpu,
dataset::ResizeMode::Letterbox,
imagenet_norm(cfg),
)?;
let x = x.to_device(device).unsqueeze(0);
match model.predict(&x, 0.25, 0.5)? {
PredictOutput::Classify { labels, confs } => {
let preds: Vec<serde_json::Value> = labels
.iter()
.zip(&confs)
.map(|(c, p)| serde_json::json!({ "class_id": c, "prob": p }))
.collect();
Ok(serde_json::json!({ "task": "classify", "predictions": preds }))
}
PredictOutput::Detect { per_image } => {
let mut dets = per_image.first().cloned().unwrap_or_default();
if let Some(lb) = lb {
for d in &mut dets {
d.bbox = lb.restore_box(d.bbox, orig_w, orig_h);
if let Some(kps) = &mut d.keypoints {
for kp in kps.iter_mut() {
kp[0] = (kp[0] - lb.pad_left) / lb.scale;
kp[1] = (kp[1] - lb.pad_top) / lb.scale;
}
}
}
}
Ok(serde_json::json!({
"task": "detect",
"image": { "width": orig_w, "height": orig_h },
"detections": dets,
}))
}
PredictOutput::Seg { per_image } => {
let insts = per_image.first().cloned().unwrap_or_default();
let mask_size = model.img_size() / 4;
let mut preds = Vec::new();
for it in insts {
let (mut x0, mut y0, mut x1, mut y1, mut area) =
(i64::MAX, i64::MAX, 0i64, 0i64, 0usize);
for (pi, &v) in it.mask.iter().enumerate() {
if v == 1 {
let (py, px) = (
(pi / mask_size as usize) as i64,
(pi % mask_size as usize) as i64,
);
x0 = x0.min(px);
y0 = y0.min(py);
x1 = x1.max(px);
y1 = y1.max(py);
area += 1;
}
}
preds.push(serde_json::json!({
"class_id": it.label,
"score": it.score,
"mask_bbox": [x0, y0, x1, y1],
"mask_area": area,
}));
}
Ok(serde_json::json!({ "task": "seg", "predictions": preds }))
}
PredictOutput::Keypoint { per_image } => {
let mut dets = per_image.first().cloned().unwrap_or_default();
if let Some(lb) = lb {
for d in &mut dets {
d.restore(&lb);
}
}
Ok(serde_json::json!({
"task": "keypoint",
"image": { "width": orig_w, "height": orig_h },
"detections": dets,
}))
}
}
}
pub fn infer_sliced(
cfg: &RunConfig,
weights: &Path,
input: &Path,
window: u32,
overlap_frac: f32,
) -> AvResult<serde_json::Value> {
let (model, device) = load_model(cfg, weights)?;
let img = image::open(input).map_err(|e| AvError::data(format!("读图失败 {input:?}: {e}")))?;
let rgb = img.to_rgb8();
let (w, h) = (rgb.width(), rgb.height());
let s_model = model.img_size();
let (tw, th) = (window.min(w).max(1), window.min(h).max(1));
let stride = (((1.0 - overlap_frac.clamp(0.0, 0.8)) * tw as f32).round() as u32).max(1);
fn starts(total: u32, win: u32, stride: u32) -> Vec<u32> {
if total <= win {
return vec![0];
}
let mut v = Vec::new();
let mut p = 0u32;
while p + win < total {
v.push(p);
p += stride;
}
v.push(total - win);
v
}
let y_starts = starts(h, th, stride);
let x_starts = starts(w, tw, stride);
let mut tiles: Vec<(u32, u32, Tensor, Option<av_core::geometry::Letterbox>)> = Vec::new();
for &y0 in &y_starts {
for &x0 in &x_starts {
let crop = image::imageops::crop_imm(&rgb, x0, y0, tw, th).to_image();
let (x, lb) = dataset::decode_rgb_with_meta(
&crop,
s_model,
Device::Cpu,
dataset::ResizeMode::Letterbox,
imagenet_norm(cfg),
)?;
tiles.push((x0, y0, x, lb));
}
}
let mut all: Vec<av_core::types::Detection> = Vec::new();
for batch in tiles.chunks(8) {
let xs: Vec<Tensor> = batch.iter().map(|(_, _, x, _)| x.copy()).collect();
let x = Tensor::stack(&xs, 0).to_device(device);
let per = model.predict(&x, 0.25, 0.5)?;
let per_image = match per {
PredictOutput::Detect { per_image } => per_image,
PredictOutput::Classify { .. } => continue,
PredictOutput::Seg { .. } => continue,
PredictOutput::Keypoint { .. } => continue,
};
for (bi, dets) in per_image.into_iter().enumerate() {
let (x0, y0, _, lb) = &batch[bi];
for mut d in dets {
match lb {
Some(lb) => {
d.bbox.x1 = (d.bbox.x1 - lb.pad_left) / lb.scale + *x0 as f32;
d.bbox.y1 = (d.bbox.y1 - lb.pad_top) / lb.scale + *y0 as f32;
d.bbox.x2 = (d.bbox.x2 - lb.pad_left) / lb.scale + *x0 as f32;
d.bbox.y2 = (d.bbox.y2 - lb.pad_top) / lb.scale + *y0 as f32;
if let Some(kps) = &mut d.keypoints {
for kp in kps.iter_mut() {
kp[0] = (kp[0] - lb.pad_left) / lb.scale + *x0 as f32;
kp[1] = (kp[1] - lb.pad_top) / lb.scale + *y0 as f32;
}
}
}
None => {
let sx = *x0 as f32;
let sy = *y0 as f32;
d.bbox.x1 += sx;
d.bbox.y1 += sy;
d.bbox.x2 += sx;
d.bbox.y2 += sy;
}
}
d.bbox.x1 = d.bbox.x1.clamp(0.0, w as f32);
d.bbox.y1 = d.bbox.y1.clamp(0.0, h as f32);
d.bbox.x2 = d.bbox.x2.clamp(0.0, w as f32);
d.bbox.y2 = d.bbox.y2.clamp(0.0, h as f32);
if d.bbox.x2 > d.bbox.x1 && d.bbox.y2 > d.bbox.y1 {
all.push(d);
}
}
}
}
let max_center_dist = 0.25 * tw.min(th) as f32;
let all = merge_tile_fragments(all, 0.3, max_center_dist);
let all = av_core::types::nms(all, 0.5);
let tiles_cnt = (y_starts.len() * x_starts.len()) as u32;
Ok(serde_json::json!({
"task": "detect",
"mode": "sliced",
"image": { "width": w, "height": h },
"tiles": tiles_cnt,
"detections": all,
}))
}
fn log_epoch_metrics(
run_dir: &Path,
epoch: u32,
loss: f32,
metric: &str,
value: f32,
secondary: Option<(&str, f32)>,
) {
log_epoch_metrics_with_extra(run_dir, epoch, loss, metric, value, secondary, None);
}
fn log_epoch_metrics_with_extra(
run_dir: &Path,
epoch: u32,
loss: f32,
metric: &str,
value: f32,
secondary: Option<(&str, f32)>,
extra: Option<(&str, f32)>,
) {
let row = serde_json::json!({
"epoch": epoch,
"loss": loss,
"metric": metric,
"metric_value": value,
"secondary": secondary.map(|(n, v)| serde_json::json!([n, v])),
"extra": extra.map(|(n, v)| serde_json::json!([n, v])),
"ts": crate::metrics::now_unix(),
});
if let Err(e) = crate::metrics::append(run_dir, &row) {
tracing::warn!("metrics.jsonl 追加失败(忽略): {e}");
}
}
pub fn merge_tile_fragments(
mut dets: Vec<av_core::types::Detection>,
iou_thr: f32,
max_center_dist: f32,
) -> Vec<av_core::types::Detection> {
dets.sort_by(|a, b| {
b.score
.partial_cmp(&a.score)
.unwrap_or(std::cmp::Ordering::Equal)
});
let mut out: Vec<av_core::types::Detection> = Vec::new();
for d in dets {
let (dcx, dcy) = ((d.bbox.x1 + d.bbox.x2) * 0.5, (d.bbox.y1 + d.bbox.y2) * 0.5);
let mut merged = false;
for k in out.iter_mut() {
if k.class_id != d.class_id {
continue;
}
let (kcx, kcy) = ((k.bbox.x1 + k.bbox.x2) * 0.5, (k.bbox.y1 + k.bbox.y2) * 0.5);
let dist = ((kcx - dcx).powi(2) + (kcy - dcy).powi(2)).sqrt();
if k.bbox.iou(&d.bbox) >= iou_thr && dist < max_center_dist {
k.bbox.x1 = k.bbox.x1.min(d.bbox.x1);
k.bbox.y1 = k.bbox.y1.min(d.bbox.y1);
k.bbox.x2 = k.bbox.x2.max(d.bbox.x2);
k.bbox.y2 = k.bbox.y2.max(d.bbox.y2);
merged = true;
break;
}
}
if !merged {
out.push(d);
}
}
out
}
const GPU_CACHE_BUDGET_BYTES: u64 = 4 * 1024 * 1024 * 1024;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
enum SegCacheMode {
Gpu,
Ram,
Off,
}
fn resolve_seg_cache_mode(
cfg_value: &str,
n: usize,
img_size: u32,
device: Device,
) -> SegCacheMode {
let bytes = n as u64 * 3 * img_size as u64 * img_size as u64 * 4;
let is_cuda = matches!(device, Device::Cuda(_));
match cfg_value {
"gpu" if is_cuda => SegCacheMode::Gpu,
"gpu" => {
tracing::warn!("data.cache = \"gpu\" 但设备无 CUDA,回落 ram");
SegCacheMode::Ram
}
"ram" => SegCacheMode::Ram,
"off" => SegCacheMode::Off,
_ => {
if is_cuda && bytes <= GPU_CACHE_BUDGET_BYTES {
SegCacheMode::Gpu
} else {
SegCacheMode::Ram
}
}
}
}
struct SyncTensor(tch::Tensor);
unsafe impl Send for SyncTensor {}
unsafe impl Sync for SyncTensor {}
enum SegEncoder {
Off {
raw: std::sync::Arc<Vec<dataset::RawSegSample>>,
img_size: u32,
inorm: bool,
},
Ram {
cache: std::sync::Arc<Vec<dataset::CachedSegSample>>,
img_size: u32,
inorm: bool,
},
Gpu {
stack: Arc<SyncTensor>,
meta: Arc<Vec<dataset::CachedSegSample>>,
img_size: u32,
inorm: bool,
},
}
impl SegEncoder {
fn encode(&self, idx: &[usize], plans: &[AugmentPlan]) -> AvResult<Vec<SegSample>> {
match self {
SegEncoder::Off {
raw,
img_size,
inorm,
} => idx
.par_iter()
.zip(plans)
.map(|(&i, plan)| {
dataset::encode_seg_sample(&raw[i], *img_size, Device::Cpu, plan, *inorm)
})
.collect(),
SegEncoder::Ram {
cache,
img_size,
inorm,
} => {
dataset::encode_seg_batch_cached(cache, idx, plans, *img_size, Device::Cpu, *inorm)
}
SegEncoder::Gpu {
stack,
meta,
img_size,
inorm,
} => idx
.par_iter()
.zip(plans)
.map(|(&i, plan)| {
dataset::encode_seg_sample_gpu(
&stack.0, i as u32, &meta[i], *img_size, plan, *inorm,
)
})
.collect(),
}
}
fn describe(&self, n: usize) -> String {
match self {
SegEncoder::Off { .. } => "off(历史全分辨率路径)".into(),
SegEncoder::Ram { .. } => {
format!("ram({n} 张内容贴片缓存 + rayon 并行 + 双缓冲预取)")
}
SegEncoder::Gpu { .. } => {
format!("gpu({n} 张画布显存驻留 + GPU 张量增强)")
}
}
}
}
fn train_seg(cfg: &RunConfig, run_id: &str, run_dir: &Path, resume: bool) -> AvResult<TrainReport> {
let (num_classes, img_size) = match cfg.model.tasks.first() {
Some(TaskCfg::Seg(s)) => (s.num_classes as u32, s.img_size),
_ => unreachable!("train_seg 只处理分割任务"),
};
let device = resolve_device(cfg);
let root = cfg
.data
.sources
.train
.dir
.clone()
.ok_or_else(|| AvError::config("dir 数据源缺 data.sources.train.dir"))?;
let split = cfg.data.sources.train.split.as_deref().unwrap_or("train");
let aug_cfg = train_augment_cfg(cfg, TaskKind::Seg);
let mut train_raw = if let Some(a) = &aug_cfg {
println!(
"[augment] seg 训练增强生效: flip={} hsv={:?} scale_jitter={:?} close_last_epochs={}",
a.flip, a.hsv, a.scale_jitter, a.close_last_epochs
);
Some(dataset::load_cocoseg_dir_raw(&root, split)?)
} else {
None
};
let train = if train_raw.is_some() {
Vec::new()
} else {
dataset::load_cocoseg_dir(&root, split, img_size, Device::Cpu, imagenet_norm(cfg))?
};
let val = match cfg.data.sources.val.dir.as_ref() {
Some(d) => dataset::load_cocoseg_dir(
d,
cfg.data.sources.val.split.as_deref().unwrap_or("val"),
img_size,
Device::Cpu,
imagenet_norm(cfg),
)?,
None => dataset::load_cocoseg_dir(&root, "val", img_size, Device::Cpu, imagenet_norm(cfg))?,
};
let train_n = train_raw.as_ref().map_or(train.len(), |r| r.len());
let train_insts: usize = match &train_raw {
Some(raw) => raw.iter().map(|s| s.polys.len()).sum(),
None => train.iter().map(|s| s.masks.len()).sum(),
};
println!(
"[seg] root={} train={}图/{}实例 val={}图/{}实例 classes={} img_size={} mask画布={}×{}",
root.display(),
train_n,
train_insts,
val.len(),
val.iter().map(|s| s.masks.len()).sum::<usize>(),
num_classes,
img_size,
img_size / 4,
img_size / 4
);
let seg_encoder: Option<Arc<SegEncoder>> = match (&train_raw, &aug_cfg) {
(Some(raw), Some(_)) => {
let mode = resolve_seg_cache_mode(&cfg.data.cache, raw.len(), img_size, device);
match mode {
SegCacheMode::Off => Some(Arc::new(SegEncoder::Off {
raw: Arc::new(std::mem::take(&mut train_raw).unwrap()),
img_size,
inorm: imagenet_norm(cfg),
})),
SegCacheMode::Ram => {
let (cache, bytes) = dataset::build_seg_cache(raw, img_size)?;
println!(
"[seg] 数据缓存: ram({} 张内容贴片,{:.0}MB;rayon 并行 + 双缓冲预取)",
cache.len(),
bytes as f64 / 1e6
);
drop(std::mem::take(&mut train_raw));
Some(Arc::new(SegEncoder::Ram {
cache: Arc::new(cache),
img_size,
inorm: imagenet_norm(cfg),
}))
}
SegCacheMode::Gpu => {
let (cache, bytes) = dataset::build_seg_cache(raw, img_size)?;
let stack = {
let res = std::panic::catch_unwind(std::panic::AssertUnwindSafe(|| {
dataset::build_seg_canvas_stack(&cache, img_size, device)
}));
match res {
Ok(Ok(t)) => Some(t),
Ok(Err(e)) => {
tracing::warn!("显存驻留上传失败({e}),降级 ram 缓存");
None
}
Err(_) => {
tracing::warn!("显存驻留上传 OOM,降级 ram 缓存");
None
}
}
};
match stack {
Some(t) => {
println!(
"[seg] 数据缓存: gpu({} 张画布显存驻留 {:.0}MB;GPU 张量增强)",
cache.len(),
bytes as f64 * 4.0 / 1e6
);
drop(std::mem::take(&mut train_raw));
Some(Arc::new(SegEncoder::Gpu {
stack: Arc::new(SyncTensor(t)),
meta: Arc::new(cache),
img_size,
inorm: imagenet_norm(cfg),
}))
}
None => {
println!(
"[seg] 数据缓存: ram(降级;{} 张内容贴片,{:.0}MB)",
cache.len(),
bytes as f64 / 1e6
);
drop(std::mem::take(&mut train_raw));
Some(Arc::new(SegEncoder::Ram {
cache: Arc::new(cache),
img_size,
inorm: imagenet_norm(cfg),
}))
}
}
}
}
}
_ => None,
};
let encoder_desc = seg_encoder
.as_ref()
.map(|e| e.describe(train_n))
.unwrap_or_else(|| "off(无增强,整集预解码)".into());
if seg_encoder.is_some() {
println!("[seg] 数据管线: {encoder_desc}");
}
let mut vs = VarStore::new(device);
let model = build_model(&vs.root(), cfg)?;
apply_pretrain(&vs, cfg)?;
let mut start_epoch = 1u32;
if resume {
let last = run_dir.join("last.ckpt");
match read_checkpoint_epoch(&last) {
Some(done) => {
load_checkpoint(&mut vs, &last)?;
start_epoch = done + 1;
println!(
"[resume] 已从 last.ckpt 恢复(完成 {done} epoch),从 epoch {start_epoch} 续训;EMA 重新累计"
);
}
None => println!("[resume] 未找到 last.ckpt(runs/{run_id}/),从头训练"),
}
}
let mut opt = make_opt(&vs, cfg)?;
let mut rng = XorShift::new(cfg.seed);
let bs = (cfg.train.batch_size as usize).max(1);
let mut final_loss = 0f32;
#[allow(unused_assignments)] let (mut val_miou, mut val_r50, mut val_p50) = (0f32, 0f32, 0f32);
let bn_train = bn_train_mode(cfg);
for epoch in start_epoch..=cfg.train.epochs {
model.set_train(bn_train);
opt.set_lr(schedule_lr(cfg, epoch));
let mut order: Vec<usize> = (0..train_n).collect();
for i in (1..order.len()).rev() {
let j = rng.next_usize(i + 1);
order.swap(i, j);
}
let strong_on = aug_cfg
.as_ref()
.map(|a| strong_aug_on(a, cfg, epoch))
.unwrap_or(false);
let mut aug_rng = epoch_aug_rng(cfg.seed, epoch);
let mut epoch_loss = EpochLossAcc::new();
let mut steps = 0usize;
let acc_steps = cfg.train.accumulate_steps.max(1) as usize;
opt.zero_grad();
let strong_plans: Vec<AugmentPlan> = match (&seg_encoder, &aug_cfg) {
(Some(_), Some(a)) if strong_on => (0..order.len())
.map(|_| augment::draw_plan(a, &mut aug_rng))
.collect(),
_ => Vec::new(),
};
let chunk_starts: Vec<usize> = (0..order.len()).step_by(bs).collect();
let mut in_flight: Option<std::thread::JoinHandle<AvResult<Vec<SegSample>>>> = None;
for (ci, &start) in chunk_starts.iter().enumerate() {
let end = (start + bs).min(order.len());
let (x, masks, labels) = match in_flight.take() {
Some(h) => {
let batch_samples: Vec<SegSample> =
h.join().map_err(|_| AvError::train("数据编码线程崩溃"))??;
let x = dataset::stack_seg_samples(&batch_samples)?.to_device(device);
(
x,
batch_samples.iter().map(|s| s.masks.clone()).collect(),
batch_samples.iter().map(|s| s.labels.clone()).collect(),
)
}
None => {
let idx = &order[start..end];
match &seg_encoder {
Some(enc) => {
let plans: Vec<AugmentPlan> = if strong_on {
strong_plans[start..end].to_vec()
} else {
vec![AugmentPlan::none(); end - start]
};
let batch_samples = enc.encode(idx, &plans)?;
let x = dataset::stack_seg_samples(&batch_samples)?.to_device(device);
(
x,
batch_samples.iter().map(|s| s.masks.clone()).collect(),
batch_samples.iter().map(|s| s.labels.clone()).collect(),
)
}
None => {
let xs: Vec<&Tensor> = idx.iter().map(|&i| &train[i].x).collect();
(
Tensor::stack(&xs, 0).to_device(device),
idx.iter().map(|&i| train[i].masks.clone()).collect(),
idx.iter().map(|&i| train[i].labels.clone()).collect(),
)
}
}
}
};
if let (Some(enc), Some(&next_start)) = (&seg_encoder, chunk_starts.get(ci + 1)) {
let next_end = (next_start + bs).min(order.len());
let idx = order[next_start..next_end].to_vec();
let plans: Vec<AugmentPlan> = if strong_on {
strong_plans[next_start..next_end].to_vec()
} else {
vec![AugmentPlan::none(); next_end - next_start]
};
let enc = Arc::clone(enc);
in_flight = Some(std::thread::spawn(move || enc.encode(&idx, &plans)));
}
let batch = TrainBatch::Seg { masks, labels };
let loss = model.loss(&x, &batch)?;
loss.backward();
epoch_loss.add(&loss);
steps += 1;
if steps.is_multiple_of(acc_steps) {
if cfg.train.grad_clip > 0.0 {
opt.clip_grad_norm(cfg.train.grad_clip as f64);
}
opt.step();
opt.zero_grad();
}
}
if !steps.is_multiple_of(acc_steps) {
if cfg.train.grad_clip > 0.0 {
opt.clip_grad_norm(cfg.train.grad_clip as f64);
}
opt.step();
opt.zero_grad();
}
final_loss = epoch_loss.mean(steps);
let eval_due = epoch == 1
|| epoch == cfg.train.epochs
|| cfg.eval.interval_epochs == 0
|| epoch % cfg.eval.interval_epochs == 0;
if eval_due {
model.set_train(false); let TaskModel::Seg(m) = &model else {
unreachable!("分割任务模型类型")
};
(val_miou, val_r50, val_p50, _, _) = eval_seg_samples(m, &val, device)?;
model.set_train(bn_train); println!(
"[seg] run={run_id} epoch={epoch}/{} loss={final_loss:.4} val掩码mIoU={val_miou:.3} val R@0.5={val_r50:.3} val P@0.5={val_p50:.3}",
cfg.train.epochs
);
match save_checkpoint_epoch(&vs, &run_dir.join("last.ckpt"), epoch) {
Ok(()) => println!("[ckpt] last.ckpt 已更新(epoch {epoch})"),
Err(e) => tracing::warn!("last.ckpt 保存失败(忽略,不影响训练): {e}"),
}
}
log_epoch_metrics_with_extra(
run_dir,
epoch,
final_loss,
"mask_miou",
val_miou,
Some(("recall@0.5", val_r50)),
Some(("precision@0.5", val_p50)),
);
}
model.set_train(false); let TaskModel::Seg(m) = &model else {
unreachable!("分割任务模型类型")
};
let (train_miou, train_r50, train_p50, n_inst, _) = match &seg_encoder {
Some(enc) => {
let idx: Vec<usize> = (0..train_n).collect();
let plans = vec![AugmentPlan::none(); train_n];
let clean = enc.encode(&idx, &plans)?;
eval_seg_samples(m, &clean, device)?
}
None => eval_seg_samples(m, &train, device)?,
};
save_checkpoint(&vs, &run_dir.join(CKPT_DIR))?;
Ok(TrainReport {
run_id: run_id.to_string(),
task: "seg".into(),
epochs: cfg.train.epochs,
final_loss,
metric: "train_mask_miou".into(),
metric_value: train_miou,
secondary: Some(("val_mask_miou".into(), val_miou)),
run_dir: run_dir.display().to_string(),
})
.inspect(|_| {
println!(
"[seg] 训练集验收:掩码 mIoU={train_miou:.3} R@0.5={train_r50:.3} P@0.5={train_p50:.3}({n_inst} 个 gt 实例)\
| val mIoU={val_miou:.3} val R@0.5={val_r50:.3} val P@0.5={val_p50:.3}",
);
})
}
type SegEval = (f32, f32, f32, usize, Vec<(u32, f32, usize)>);
fn eval_seg_samples(m: &SegModel, val: &[SegSample], device: Device) -> AvResult<SegEval> {
if val.is_empty() {
return Err(AvError::data("验证集为空"));
}
let mut per_image: Vec<Vec<av_tasks::models::SegInstance>> = Vec::new();
for chunk in val.chunks(16) {
let x = dataset::stack_seg_samples(chunk)?.to_device(device);
per_image.extend(m.predict(&x, 0.1, 0.5)?);
}
let mut sum_best = 0f32;
let mut n_gt = 0usize;
let mut hits = 0usize;
let mut tp = 0usize;
let mut n_pred = 0usize;
let mut per_class_acc: std::collections::BTreeMap<u32, (f64, usize)> =
std::collections::BTreeMap::new();
for (gi, s) in val.iter().enumerate() {
let preds = &per_image[gi];
n_gt += s.masks.len();
n_pred += preds.len();
let gt_sum: Vec<MaskSummary> = s.masks.iter().map(|m| MaskSummary::of(m)).collect();
let ious: Vec<Vec<f32>> = preds
.iter()
.map(|d| {
let ds = MaskSummary::of(&d.mask);
s.masks
.iter()
.zip(>_sum)
.map(|(gt, gs)| ds.iou(&d.mask, gs, gt))
.collect()
})
.collect();
for (g, _gt_mask) in s.masks.iter().enumerate() {
let label = s.labels[g];
let best = ious.iter().map(|row| row[g]).fold(0f32, f32::max);
sum_best += best;
let acc = per_class_acc.entry(label).or_insert((0.0, 0));
acc.0 += best as f64;
acc.1 += 1;
if preds
.iter()
.enumerate()
.any(|(pi, d)| d.label == label && ious[pi][g] >= 0.5)
{
hits += 1;
}
}
tp += greedy_mask_match(
preds.iter().map(|d| (d.label, d.score)).collect(),
&s.labels,
&ious,
0.5,
);
}
if n_gt == 0 {
return Err(AvError::data("验证集无 gt 实例"));
}
let p50 = if n_pred == 0 {
0.0
} else {
tp as f32 / n_pred as f32
};
let per_class = per_class_acc
.into_iter()
.map(|(cls, (sum, n))| (cls, (sum / n as f64) as f32, n))
.collect();
Ok((
sum_best / n_gt as f32,
hits as f32 / n_gt as f32,
p50,
n_gt,
per_class,
))
}
fn greedy_mask_match(
preds: Vec<(u32, f32)>,
gt_labels: &[u32],
ious: &[Vec<f32>],
thr: f32,
) -> usize {
let mut order: Vec<usize> = (0..preds.len()).collect();
order.sort_by(|&a, &b| preds[b].1.total_cmp(&preds[a].1));
let mut used = vec![false; gt_labels.len()];
let mut tp = 0usize;
for pi in order {
let mut best_g = None;
let mut best_v = thr;
for (g, &used_g) in used.iter().enumerate() {
if used_g || preds[pi].0 != gt_labels[g] {
continue;
}
let v = ious[pi][g];
if v >= best_v {
best_v = v;
best_g = Some(g);
}
}
if let Some(g) = best_g {
used[g] = true;
tp += 1;
}
}
tp
}
fn train_keypoint(cfg: &RunConfig, run_id: &str, run_dir: &Path) -> AvResult<TrainReport> {
let (num_keypoints, img_size) = match cfg.model.tasks.first() {
Some(TaskCfg::Keypoint(k)) => (k.num_keypoints, k.img_size),
_ => unreachable!("train_keypoint 只处理关键点任务"),
};
let device = resolve_device(cfg);
let root = cfg
.data
.sources
.train
.dir
.clone()
.ok_or_else(|| AvError::config("dir 数据源缺 data.sources.train.dir"))?;
let split = cfg.data.sources.train.split.as_deref().unwrap_or("train");
let aug_cfg = train_augment_cfg(cfg, TaskKind::Keypoint);
let train_raw = if let Some(a) = &aug_cfg {
println!(
"[augment] keypoint 训练增强生效: flip={} hsv={:?} scale_jitter={:?} close_last_epochs={}",
a.flip, a.hsv, a.scale_jitter, a.close_last_epochs
);
Some(dataset::load_cocopose_dir_raw(&root, split)?)
} else {
None
};
let train = if train_raw.is_some() {
Vec::new()
} else {
dataset::load_cocopose_dir(&root, split, img_size, Device::Cpu, imagenet_norm(cfg))?
};
let val = match cfg.data.sources.val.dir.as_ref() {
Some(d) => dataset::load_cocopose_dir(
d,
cfg.data.sources.val.split.as_deref().unwrap_or("val"),
img_size,
Device::Cpu,
imagenet_norm(cfg),
)?,
None => {
dataset::load_cocopose_dir(&root, "val", img_size, Device::Cpu, imagenet_norm(cfg))?
}
};
let train_kpts: Vec<&Vec<[f32; 3]>> = match &train_raw {
Some(raw) => raw.iter().flat_map(|s| s.kpts.iter()).collect(),
None => train.iter().flat_map(|s| s.kpts.iter()).collect(),
};
for gk in train_kpts
.iter()
.copied()
.chain(val.iter().flat_map(|s| s.kpts.iter()))
{
if gk.len() != num_keypoints {
return Err(AvError::data(format!(
"标注关键点数 {} 与 keypoint.num_keypoints = {num_keypoints} 不一致\
(数据集与配置必须同一关键点模板)",
gk.len()
)));
}
}
let train_n = train_raw.as_ref().map_or(train.len(), |r| r.len());
let train_insts: usize = match &train_raw {
Some(raw) => raw.iter().map(|s| s.kpts.len()).sum(),
None => train.iter().map(|s| s.kpts.len()).sum(),
};
println!(
"[keypoint] root={} train={}图/{}实例 val={}图/{}实例 K={} img_size={} stride=8",
root.display(),
train_n,
train_insts,
val.len(),
val.iter().map(|s| s.kpts.len()).sum::<usize>(),
num_keypoints,
img_size
);
let vs = VarStore::new(device);
let model = build_model(&vs.root(), cfg)?;
apply_pretrain(&vs, cfg)?;
let mut opt = make_opt(&vs, cfg)?;
let mut rng = XorShift::new(cfg.seed);
let bs = (cfg.train.batch_size as usize).max(1);
let mut final_loss = 0f32;
#[allow(unused_assignments)] let (mut val_pck, mut val_oks) = (0f32, 0f32);
let bn_train = bn_train_mode(cfg);
for epoch in 1..=cfg.train.epochs {
model.set_train(bn_train);
opt.set_lr(schedule_lr(cfg, epoch));
let mut order: Vec<usize> = (0..train_n).collect();
for i in (1..order.len()).rev() {
let j = rng.next_usize(i + 1);
order.swap(i, j);
}
let strong_on = aug_cfg
.as_ref()
.map(|a| strong_aug_on(a, cfg, epoch))
.unwrap_or(false);
let mut aug_rng = epoch_aug_rng(cfg.seed, epoch);
let plans: Vec<AugmentPlan> = match (&train_raw, &aug_cfg) {
(Some(_), Some(a)) if strong_on => (0..order.len())
.map(|_| augment::draw_plan(a, &mut aug_rng))
.collect(),
(Some(_), Some(_)) => vec![AugmentPlan::none(); order.len()],
_ => Vec::new(),
};
let mut epoch_loss = EpochLossAcc::new();
let mut steps = 0usize;
let acc_steps = cfg.train.accumulate_steps.max(1) as usize;
opt.zero_grad();
for (ci, chunk) in order.chunks(bs).enumerate() {
let (x, boxes, kpts, labels) = match (&train_raw, &aug_cfg) {
(Some(raw), Some(_)) => {
let batch_samples: Vec<KeypointSample> = chunk
.par_iter()
.zip(&plans[ci * bs..ci * bs + chunk.len()])
.map(|(&i, plan)| {
dataset::encode_keypoint_sample(
&raw[i],
img_size,
Device::Cpu,
plan,
imagenet_norm(cfg),
)
})
.collect::<AvResult<Vec<_>>>()?;
let x = dataset::stack_kp_samples(&batch_samples)?.to_device(device);
(
x,
batch_samples.iter().map(|s| s.boxes.clone()).collect(),
batch_samples.iter().map(|s| s.kpts.clone()).collect(),
batch_samples.iter().map(|s| s.labels.clone()).collect(),
)
}
_ => {
let xs: Vec<&Tensor> = chunk.iter().map(|&i| &train[i].x).collect();
(
Tensor::stack(&xs, 0).to_device(device),
chunk.iter().map(|&i| train[i].boxes.clone()).collect(),
chunk.iter().map(|&i| train[i].kpts.clone()).collect(),
chunk.iter().map(|&i| train[i].labels.clone()).collect(),
)
}
};
let batch = TrainBatch::Keypoint {
boxes,
kpts,
labels,
};
let loss = model.loss(&x, &batch)?;
loss.backward();
epoch_loss.add(&loss);
steps += 1;
if steps.is_multiple_of(acc_steps) {
if cfg.train.grad_clip > 0.0 {
opt.clip_grad_norm(cfg.train.grad_clip as f64);
}
opt.step();
opt.zero_grad();
}
}
if !steps.is_multiple_of(acc_steps) {
if cfg.train.grad_clip > 0.0 {
opt.clip_grad_norm(cfg.train.grad_clip as f64);
}
opt.step();
opt.zero_grad();
}
final_loss = epoch_loss.mean(steps);
let eval_due = epoch == 1
|| epoch == cfg.train.epochs
|| cfg.eval.interval_epochs == 0
|| epoch % cfg.eval.interval_epochs == 0;
if eval_due {
model.set_train(false); let TaskModel::Keypoint(m) = &model else {
unreachable!("关键点任务模型类型")
};
(val_pck, val_oks, _, _) = eval_kp_samples(m, &val, device)?;
model.set_train(bn_train); println!(
"[keypoint] run={run_id} epoch={epoch}/{} loss={final_loss:.4} val PCK@0.5={val_pck:.3} val meanOKS={val_oks:.3}",
cfg.train.epochs
);
}
log_epoch_metrics(
run_dir,
epoch,
final_loss,
"pck@0.5",
val_pck,
Some(("mean_oks", val_oks)),
);
}
model.set_train(false); let TaskModel::Keypoint(m) = &model else {
unreachable!("关键点任务模型类型")
};
let (train_pck, train_oks, n_vis, n_inst) = match &train_raw {
Some(raw) => {
let clean: Vec<KeypointSample> = raw
.iter()
.map(|s| {
dataset::encode_keypoint_sample(
s,
img_size,
Device::Cpu,
&AugmentPlan::none(),
imagenet_norm(cfg),
)
})
.collect::<AvResult<Vec<_>>>()?;
eval_kp_samples(m, &clean, device)?
}
None => eval_kp_samples(m, &train, device)?,
};
save_checkpoint(&vs, &run_dir.join(CKPT_DIR))?;
Ok(TrainReport {
run_id: run_id.to_string(),
task: "keypoint".into(),
epochs: cfg.train.epochs,
final_loss,
metric: "train_pck@0.5".into(),
metric_value: train_pck,
secondary: Some(("val_pck@0.5".into(), val_pck)),
run_dir: run_dir.display().to_string(),
})
.inspect(|_| {
println!(
"[keypoint] 训练集验收:PCK@0.5={train_pck:.3} meanOKS={train_oks:.3}\
({n_inst} 个 gt 实例 / {n_vis} 个可见 gt 点)| val PCK@0.5={val_pck:.3} val meanOKS={val_oks:.3}",
);
})
}
fn eval_kp_samples(
m: &KeypointModel,
samples: &[KeypointSample],
device: Device,
) -> AvResult<(f32, f32, usize, usize)> {
if samples.is_empty() {
return Err(AvError::data("验证集为空"));
}
let thr = 0.1 * m.img_size() as f32;
let mut per_image: Vec<Vec<av_core::types::Detection>> = Vec::new();
for chunk in samples.chunks(16) {
let x = dataset::stack_kp_samples(chunk)?.to_device(device);
per_image.extend(m.predict(&x, 0.1, 0.5)?);
}
let (mut hits, mut visible, mut n_inst) = (0usize, 0usize, 0usize);
let (mut ok_sum, mut ok_cnt) = (0f64, 0usize);
for (gi, s) in samples.iter().enumerate() {
for (g, gk) in s.kpts.iter().enumerate() {
n_inst += 1;
let n_vis_here = gk.iter().filter(|p| p[2] > 0.0).count();
visible += n_vis_here;
let (cx, cy, bw, bh) = {
let b = s.boxes[g];
(b[0], b[1], b[2], b[3])
};
let g_aabb = Aabb::new(cx - bw / 2.0, cy - bh / 2.0, cx + bw / 2.0, cy + bh / 2.0);
let mut best: Option<(f32, &av_core::types::Detection)> = None;
for d in &per_image[gi] {
let v = g_aabb.iou(&d.bbox);
if best.map(|(bv, _)| v > bv).unwrap_or(true) {
best = Some((v, d));
}
}
let Some((best_iou, d)) = best else { continue };
if best_iou < 0.1 {
continue;
}
let Some(kps) = d.keypoints.as_ref() else {
continue;
};
if n_vis_here > 0 {
let scale = (bw * bh).sqrt().max(1e-3);
ok_sum += av_tasks::oks::oks_scalar(kps, gk, scale) as f64;
ok_cnt += 1;
}
for (j, gp) in gk.iter().enumerate() {
if gp[2] <= 0.0 {
continue;
}
if let Some(pk) = kps.get(j) {
let dx = pk[0] - gp[0];
let dy = pk[1] - gp[1];
if dx * dx + dy * dy < thr * thr {
hits += 1;
}
}
}
}
}
if visible == 0 {
return Err(AvError::data("评测集无可见 gt 关键点"));
}
Ok((
hits as f32 / visible as f32,
(ok_sum / ok_cnt.max(1) as f64) as f32,
visible,
n_inst,
))
}
fn train_detect_obb(cfg: &RunConfig, run_id: &str, run_dir: &Path) -> AvResult<TrainReport> {
let (num_classes, img_size) = match cfg.model.tasks.first() {
Some(TaskCfg::Detect(d)) => (d.num_classes as u32, d.img_size),
_ => unreachable!("train_detect_obb 只处理 OBB 任务"),
};
let device = resolve_device(cfg);
let root = cfg
.data
.sources
.train
.dir
.clone()
.ok_or_else(|| AvError::config("dir 数据源缺 data.sources.train.dir"))?;
let split = cfg.data.sources.train.split.as_deref().unwrap_or("train");
let aug_cfg = train_augment_cfg(cfg, TaskKind::Detect);
let train_raw = if let Some(a) = &aug_cfg {
println!(
"[augment] obb 训练增强生效: flip={} hsv={:?} scale_jitter={:?} close_last_epochs={}",
a.flip, a.hsv, a.scale_jitter, a.close_last_epochs
);
Some(dataset::load_dota_dir_raw(&root, split)?)
} else {
None
};
let train = if train_raw.is_some() {
Vec::new()
} else {
dataset::load_dota_dir(&root, split, img_size, Device::Cpu, imagenet_norm(cfg))?
};
let val = match cfg.data.sources.val.dir.as_ref() {
Some(d) => dataset::load_dota_dir(
d,
cfg.data.sources.val.split.as_deref().unwrap_or("val"),
img_size,
Device::Cpu,
imagenet_norm(cfg),
)?,
None => dataset::load_dota_dir(&root, "val", img_size, Device::Cpu, imagenet_norm(cfg))?,
};
let train_n = train_raw.as_ref().map_or(train.len(), |r| r.len());
println!(
"[obb] root={} train={} val={} classes={} img_size={}",
root.display(),
train_n,
val.len(),
num_classes,
img_size
);
let vs = VarStore::new(device);
let model = build_model(&vs.root(), cfg)?;
let mut opt = make_opt(&vs, cfg)?;
let mut rng = XorShift::new(cfg.seed);
let bs = (cfg.train.batch_size as usize).max(1);
let mut final_loss = 0f32;
let mut miou = 0f32;
let mut r50 = 0f32;
let bn_train = bn_train_mode(cfg);
for epoch in 1..=cfg.train.epochs {
model.set_train(bn_train);
opt.set_lr(schedule_lr(cfg, epoch));
let mut order: Vec<usize> = (0..train_n).collect();
for i in (1..order.len()).rev() {
let j = rng.next_usize(i + 1);
order.swap(i, j);
}
let strong_on = aug_cfg
.as_ref()
.map(|a| strong_aug_on(a, cfg, epoch))
.unwrap_or(false);
let mut aug_rng = epoch_aug_rng(cfg.seed, epoch);
let plans: Vec<AugmentPlan> = match (&train_raw, &aug_cfg) {
(Some(_), Some(a)) if strong_on => (0..order.len())
.map(|_| augment::draw_plan(a, &mut aug_rng))
.collect(),
(Some(_), Some(_)) => vec![AugmentPlan::none(); order.len()],
_ => Vec::new(),
};
let mut epoch_loss = EpochLossAcc::new();
let mut steps = 0usize;
let acc_steps = cfg.train.accumulate_steps.max(1) as usize;
opt.zero_grad();
for (ci, chunk) in order.chunks(bs).enumerate() {
let (x, boxes, labels) = match (&train_raw, &aug_cfg) {
(Some(raw), Some(_)) => {
let batch_samples: Vec<dataset::ObbSample> = chunk
.par_iter()
.zip(&plans[ci * bs..ci * bs + chunk.len()])
.map(|(&i, plan)| {
dataset::encode_obb_sample(
&raw[i],
img_size,
Device::Cpu,
plan,
imagenet_norm(cfg),
)
})
.collect::<AvResult<Vec<_>>>()?;
let x = dataset::stack_obb_samples(&batch_samples)?.to_device(device);
(
x,
batch_samples.iter().map(|s| s.boxes.clone()).collect(),
batch_samples.iter().map(|s| s.labels.clone()).collect(),
)
}
_ => {
let xs: Vec<&Tensor> = chunk.iter().map(|&i| &train[i].x).collect();
(
Tensor::stack(&xs, 0).to_device(device),
chunk.iter().map(|&i| train[i].boxes.clone()).collect(),
chunk.iter().map(|&i| train[i].labels.clone()).collect(),
)
}
};
let batch = TrainBatch::Obb { boxes, labels };
let loss = model.loss(&x, &batch)?;
loss.backward();
epoch_loss.add(&loss);
steps += 1;
if steps.is_multiple_of(acc_steps) {
if cfg.train.grad_clip > 0.0 {
opt.clip_grad_norm(cfg.train.grad_clip as f64);
}
opt.step();
opt.zero_grad();
}
}
if !steps.is_multiple_of(acc_steps) {
if cfg.train.grad_clip > 0.0 {
opt.clip_grad_norm(cfg.train.grad_clip as f64);
}
opt.step();
opt.zero_grad();
}
final_loss = epoch_loss.mean(steps);
model.set_train(false);
let TaskModel::Detect(m) = &model else {
unreachable!("OBB 模型类型")
};
(miou, r50) = eval_obb_samples(m, &val, device)?;
model.set_train(bn_train);
if epoch == 1 || epoch % 10 == 0 || epoch == cfg.train.epochs {
println!(
"[obb] run={run_id} epoch={epoch}/{} loss={final_loss:.4} 旋转mIoU={miou:.3} R@0.5={r50:.3}",
cfg.train.epochs
);
}
log_epoch_metrics(
run_dir,
epoch,
final_loss,
"obb_miou",
miou,
Some(("recall@0.5", r50)),
);
}
save_checkpoint(&vs, &run_dir.join(CKPT_DIR))?;
Ok(TrainReport {
run_id: run_id.to_string(),
task: "obb".into(),
epochs: cfg.train.epochs,
final_loss,
metric: "obb_miou".into(),
metric_value: miou,
secondary: Some(("recall@0.5".into(), r50)),
run_dir: run_dir.display().to_string(),
})
}
fn eval_obb_samples(
m: &av_tasks::models::DetectModel,
val: &[dataset::ObbSample],
device: Device,
) -> AvResult<(f32, f32)> {
use av_core::geometry::RotBox;
let n = val.len();
if n == 0 {
return Err(AvError::data("验证集为空"));
}
let mut per_image: Vec<Vec<av_core::types::Detection>> = Vec::new();
for chunk in val.chunks(32) {
let xs: Vec<&Tensor> = chunk.iter().map(|s| &s.x).collect();
let x = Tensor::stack(&xs, 0).to_device(device);
per_image.extend(m.predict(&x, 0.1, 0.5)?);
}
let mut best = vec![0f32; n];
let mut ok = vec![false; n];
for (gi, s) in val.iter().enumerate() {
for d in &per_image[gi] {
let dr = RotBox {
cx: (d.bbox.x1 + d.bbox.x2) / 2.0,
cy: (d.bbox.y1 + d.bbox.y2) / 2.0,
w: d.bbox.x2 - d.bbox.x1,
h: d.bbox.y2 - d.bbox.y1,
theta: d.angle.unwrap_or(0.0),
};
let mut hit = false;
for (b, &label) in s.boxes.iter().zip(&s.labels) {
let gr = RotBox {
cx: b[0],
cy: b[1],
w: b[2],
h: b[3],
theta: b[4],
};
let v = gr.iou(&dr);
if v > best[gi] {
best[gi] = v;
}
if d.class_id == label {
hit = true;
}
}
if hit {
ok[gi] = true;
}
}
}
let matched = best.iter().filter(|&&v| v >= 0.5).count();
Ok((
best.iter().sum::<f32>() / n as f32,
matched as f32 / n as f32,
))
}
pub fn eval(cfg: &RunConfig, weights: &Path) -> AvResult<serde_json::Value> {
let (model, device) = load_model(cfg, weights)?;
match cfg.model.tasks.first() {
Some(TaskCfg::Classify(c)) => {
let acc = if cfg.data.pipeline == DataPipeline::Dir {
let (root, split) = classify_val_source(cfg);
let (val, labels, class_map) = dataset::load_imagefolder(
&root,
&split,
c.img_size,
true,
Device::Cpu,
imagenet_norm(cfg),
)?;
if class_map.len() != c.num_classes {
return Err(AvError::config(format!(
"classify.num_classes = {} 与 ImageFolder 类别数 {} 不一致",
c.num_classes,
class_map.len()
)));
}
eval_classify_samples(&model, &val, &labels, device)?
} else {
eval_classify(&model, c.num_classes as u32, c.img_size, device)?
};
Ok(serde_json::json!({ "task": "classify", "top1": acc }))
}
Some(TaskCfg::Detect(d)) => {
let TaskModel::Detect(m) = &model else {
return Err(AvError::train("模型与任务不匹配"));
};
match cfg.data.pipeline {
DataPipeline::Dir => {
let dir = cfg
.data
.sources
.val
.dir
.as_ref()
.or(cfg.data.sources.train.dir.as_ref())
.ok_or_else(|| AvError::config("dir 数据源缺失"))?;
let split = cfg.data.sources.val.split.as_deref().unwrap_or("val");
let val = dataset::load_yolo_dir(
dir,
split,
d.img_size,
Device::Cpu,
imagenet_norm(cfg),
)?;
let evm = eval_detect_samples(m, &val, device, false)?;
Ok(serde_json::json!({
"task": "detect",
"mean_iou": evm.miou,
"recall@0.5": evm.r50,
"map50": evm.map50,
"map50_95": evm.map50_95,
}))
}
_ => {
let (miou, r50, _, _) =
eval_detect(&model, d.num_classes as u32, d.img_size, device)?;
Ok(serde_json::json!({ "task": "detect", "mean_iou": miou, "recall@0.5": r50 }))
}
}
}
Some(TaskCfg::Seg(_)) => {
let TaskModel::Seg(m) = &model else {
return Err(AvError::train("模型与任务不匹配"));
};
if cfg.data.pipeline != DataPipeline::Dir {
return Err(AvError::config(
"seg 评测需要 data.pipeline = \"dir\"(COCO 分割格式)",
));
}
let dir = cfg
.data
.sources
.val
.dir
.as_ref()
.or(cfg.data.sources.train.dir.as_ref())
.ok_or_else(|| AvError::config("dir 数据源缺失"))?;
let split = cfg.data.sources.val.split.as_deref().unwrap_or("val");
let val = dataset::load_cocoseg_dir(
dir,
split,
m.img_size(),
Device::Cpu,
imagenet_norm(cfg),
)?;
let (miou, r50, p50, n_gt, per_class) = eval_seg_samples(m, &val, device)?;
let per_class_json: serde_json::Map<String, serde_json::Value> = per_class
.into_iter()
.map(|(cls, v, n)| {
(
cls.to_string(),
serde_json::json!({ "mask_miou": v, "gt": n }),
)
})
.collect();
Ok(serde_json::json!({
"task": "seg",
"gt_instances": n_gt,
"mask_miou": miou,
"recall@0.5": r50,
"precision@0.5": p50,
"per_class_mask_miou": per_class_json,
}))
}
Some(TaskCfg::Keypoint(_)) => {
let TaskModel::Keypoint(m) = &model else {
return Err(AvError::train("模型与任务不匹配"));
};
if cfg.data.pipeline != DataPipeline::Dir {
return Err(AvError::config(
"keypoint 评测需要 data.pipeline = \"dir\"(COCO 姿态格式)",
));
}
let dir = cfg
.data
.sources
.val
.dir
.as_ref()
.or(cfg.data.sources.train.dir.as_ref())
.ok_or_else(|| AvError::config("dir 数据源缺失"))?;
let split = cfg.data.sources.val.split.as_deref().unwrap_or("val");
let val = dataset::load_cocopose_dir(
dir,
split,
m.img_size(),
Device::Cpu,
imagenet_norm(cfg),
)?;
let (pck, mean_oks, n_vis, n_inst) = eval_kp_samples(m, &val, device)?;
Ok(serde_json::json!({
"task": "keypoint",
"gt_instances": n_inst,
"visible_kpts": n_vis,
"pck@0.5": pck,
"mean_oks": mean_oks,
}))
}
_ => Err(AvError::config("该任务类型在 v0.1 引擎未支持")),
}
}
pub fn resolve_device(cfg: &RunConfig) -> Device {
let spec = cfg.device.trim().to_ascii_lowercase();
let requested = match spec.as_str() {
"cpu" => Device::Cpu,
"cuda" => Device::Cuda(0),
_ if spec.starts_with("cuda:") => match spec["cuda:".len()..].parse::<usize>() {
Ok(idx) => Device::Cuda(idx),
Err(_) => {
tracing::warn!(
"无法解析 device = {:?}(支持 cpu / cuda[:N]),回退 CPU",
cfg.device
);
return Device::Cpu;
}
},
_ => {
tracing::warn!(
"未知 device = {:?}(支持 cpu / cuda[:N]),回退 CPU",
cfg.device
);
return Device::Cpu;
}
};
if requested.is_cuda() {
crate::cuda_link::ensure_torch_cuda_loaded();
let available = Device::cuda_if_available();
if available == Device::Cpu {
tracing::warn!(
"device = {:?} 请求 CUDA,但当前环境不可用(CPU 版 libtorch 或无驱动/被禁用),回退 CPU",
cfg.device
);
return Device::Cpu;
}
if let Device::Cuda(idx) = requested {
if idx >= tch::Cuda::device_count().max(0) as usize {
tracing::warn!(
"device = {:?} 超出可用范围(device_count = {}),回退 {}",
cfg.device,
tch::Cuda::device_count(),
match available {
Device::Cuda(i) => format!("cuda:{i}"),
_ => "cpu".into(),
}
);
return available;
}
}
if !cuda_kernels_usable(requested) {
tracing::warn!(
"device = {:?} 的 CUDA 内核无法执行(本机 GPU 架构缺内核且无 PTX 可 JIT,\
报错详见 docs/gpu.md 实验 A),回退 CPU",
cfg.device
);
return Device::Cpu;
}
}
requested
}
fn cuda_kernels_usable(device: Device) -> bool {
let prev_hook = std::panic::take_hook();
std::panic::set_hook(Box::new(|_| {}));
let ok = std::panic::catch_unwind(|| {
let t = Tensor::from_slice(&[1.0f32, 2.0f32]).to_device(device);
t.sum(Kind::Float).double_value(&[]) == 3.0
})
.unwrap_or(false);
std::panic::set_hook(prev_hook);
ok
}
fn load_model(cfg: &RunConfig, weights: &Path) -> AvResult<(TaskModel, Device)> {
let device = resolve_device(cfg);
let mut vs = VarStore::new(device);
let model = build_model(&vs.root(), cfg)?;
load_checkpoint(&mut vs, weights)?;
Ok((model, device))
}
fn make_opt(vs: &VarStore, cfg: &RunConfig) -> AvResult<tch::nn::Optimizer> {
let lr = cfg.train.optimizer.lr as f64;
tch::nn::Adam::default()
.build(vs, lr)
.map_err(|e| AvError::train(format!("优化器构建失败: {e}")))
}
fn imagenet_norm(cfg: &RunConfig) -> bool {
cfg.model.backbone.imagenet_norm
}
fn bn_train_mode(cfg: &RunConfig) -> bool {
!cfg.pretrain.freeze_backbone
}
fn apply_pretrain(vs: &VarStore, cfg: &RunConfig) -> AvResult<()> {
let pre = &cfg.pretrain;
if !pre.enable {
return Ok(());
}
let weight_path = pre
.weight_path
.as_deref()
.ok_or_else(|| AvError::config("pretrain.enable = true 但未指定 pretrain.weight_path"))?;
let sources: Vec<(String, Tensor)> = if weight_path.is_dir() {
let manifest = av_weight_store::read_manifest(weight_path)?;
av_weight_store::verify_hashes(weight_path, &manifest)
.map_err(|e| AvError::train(format!("预训练权重校验失败: {e}")))?;
println!(
"[pretrain] avpretrain 目录 {}:{} 个张量,哈希校验通过(backbone={} task={} epoch={:?} created_at={})",
weight_path.display(),
manifest.tensors.len(),
manifest.backbone,
manifest.task,
manifest.epoch,
manifest.created_at,
);
av_weight_store::read_all_named(weight_path, &manifest)?
} else {
weight_adapter::read_safetensors_all(weight_path)?
};
let vars = vs.variables();
let targets: Vec<(String, Vec<i64>)> = vars
.iter()
.filter(|(n, _)| !pre.load_only_backbone || n.contains("backbone"))
.map(|(n, t)| (n.clone(), t.size()))
.collect();
let map = match &pre.layer_map {
Some(p) => LayerMap::from_toml_path(p)?,
None => LayerMap::default(),
};
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);
}
}
});
let mut frozen = 0usize;
if pre.freeze_backbone {
for (name, t) in vs.variables() {
if name.contains("backbone") {
let _ = t.set_requires_grad(false);
frozen += 1;
}
}
}
println!(
"[pretrain] 来源={} {}(frozen_backbone={frozen} / {} 个目标变量)",
weight_path.display(),
report.summary(),
targets.len()
);
for sm in &report.skipped_shape_mismatch {
tracing::debug!(
"[pretrain] 形状不匹配跳过: {} ← {} 期望 {:?} 得到 {:?}",
sm.target,
sm.source,
sm.expected,
sm.got
);
}
Ok(())
}
fn train_augment_cfg(cfg: &RunConfig, kind: TaskKind) -> Option<av_core::config::AugmentCfg> {
cfg.data
.sources
.train
.tasks
.iter()
.find(|t| t.kind == kind)
.map(|t| t.augment.clone())
.filter(augment::has_strength)
}
fn strong_aug_on(a: &av_core::config::AugmentCfg, cfg: &RunConfig, epoch: u32) -> bool {
augment::strong_aug_active(a, epoch, cfg.train.epochs)
}
fn epoch_aug_rng(seed: u64, epoch: u32) -> XorShift {
XorShift::new(seed ^ (epoch as u64).wrapping_mul(0x9E37_79B9_7F4A_7C15) ^ 0xA065_5EED)
}
fn schedule_lr(cfg: &RunConfig, epoch: u32) -> f64 {
let lr0 = cfg.train.optimizer.lr as f64;
let warm = cfg.train.warmup_epochs;
if warm > 0.0 && (epoch as f32) <= warm {
return lr0 * (epoch as f32 / warm).min(1.0) as f64;
}
let span = (cfg.train.epochs as f32 - warm).max(1.0);
let t = (((epoch as f32) - warm) / span).clamp(0.0, 1.0) as f64;
let min = lr0 * cfg.train.scheduler.lr_min_factor as f64;
min + 0.5 * (lr0 - min) * (1.0 + t.cos())
}
fn train_classify(cfg: &RunConfig, run_id: &str, run_dir: &Path) -> AvResult<TrainReport> {
match cfg.data.pipeline {
DataPipeline::Synthetic => train_classify_synthetic(cfg, run_id, run_dir),
DataPipeline::Dir => train_classify_imagenette(cfg, run_id, run_dir),
DataPipeline::AvPack => Err(AvError::config("avpack 数据源按 M2 落地(PLAN 附录 B)")),
}
}
fn classify_data_root(cfg: &RunConfig) -> AvResult<PathBuf> {
if let Some(TaskCfg::Classify(c)) = cfg.model.tasks.first() {
if let Some(d) = &c.data_dir {
return Ok(d.clone());
}
}
cfg.data.sources.train.dir.clone().ok_or_else(|| {
AvError::config("分类 dir 数据源缺失:需指定 classify.data_dir 或 data.sources.train.dir")
})
}
fn classify_val_source(cfg: &RunConfig) -> (PathBuf, String) {
let split = cfg
.data
.sources
.val
.split
.as_deref()
.unwrap_or("val")
.to_string();
match cfg.data.sources.val.dir.as_ref() {
Some(d) => (d.clone(), split),
None => (classify_data_root(cfg).unwrap_or_default(), split),
}
}
fn train_classify_imagenette(
cfg: &RunConfig,
run_id: &str,
run_dir: &Path,
) -> AvResult<TrainReport> {
let (num_classes, img_size) = match cfg.model.tasks.first() {
Some(TaskCfg::Classify(c)) => (c.num_classes as u32, c.img_size),
_ => unreachable!("train_classify 只处理分类任务"),
};
let device = resolve_device(cfg);
let root = classify_data_root(cfg)?;
let train_split = cfg.data.sources.train.split.as_deref().unwrap_or("train");
let (train, train_labels, class_map) = dataset::load_imagefolder(
&root,
train_split,
img_size,
true,
Device::Cpu,
imagenet_norm(cfg),
)?;
if class_map.len() != num_classes as usize {
return Err(AvError::config(format!(
"classify.num_classes = {num_classes} 与 ImageFolder 目录类别数 {} 不一致",
class_map.len()
)));
}
let (val_root, val_split) = classify_val_source(cfg);
let (val, val_labels, _) = dataset::load_imagefolder_with_classes(
&val_root,
&val_split,
img_size,
Some(&class_map),
Device::Cpu,
imagenet_norm(cfg),
)?;
println!(
"[classify-imagenette] root={} train={} val={} classes={} img_size={}",
root.display(),
train.len(),
val.len(),
class_map.len(),
img_size
);
let vs = VarStore::new(device);
let model = build_model(&vs.root(), cfg)?;
apply_pretrain(&vs, cfg)?;
let mut opt = make_opt(&vs, cfg)?;
let mut rng = XorShift::new(cfg.seed);
let bs = (cfg.train.batch_size as usize).max(1);
let mut final_loss = 0f32;
let mut acc = 0f32;
let bn_train = bn_train_mode(cfg);
for epoch in 1..=cfg.train.epochs {
model.set_train(bn_train);
opt.set_lr(schedule_lr(cfg, epoch));
let mut order: Vec<usize> = (0..train.len()).collect();
for i in (1..order.len()).rev() {
let j = rng.next_usize(i + 1);
order.swap(i, j);
}
let mut epoch_loss = EpochLossAcc::new();
let mut steps = 0usize;
let acc_steps = cfg.train.accumulate_steps.max(1) as usize;
opt.zero_grad();
for chunk in order.chunks(bs) {
let xs: Vec<&Tensor> = chunk.iter().map(|&i| &train[i].x).collect();
let x = Tensor::stack(&xs, 0).to_device(device);
let y: Vec<i64> = chunk.iter().map(|&i| train_labels[i] as i64).collect();
let labels = Tensor::from_slice(&y).to_device(device);
let loss = model.loss(&x, &TrainBatch::Classify { labels })?;
loss.backward();
epoch_loss.add(&loss);
steps += 1;
if steps.is_multiple_of(acc_steps) {
if cfg.train.grad_clip > 0.0 {
opt.clip_grad_norm(cfg.train.grad_clip as f64);
}
opt.step();
opt.zero_grad();
}
}
if !steps.is_multiple_of(acc_steps) {
if cfg.train.grad_clip > 0.0 {
opt.clip_grad_norm(cfg.train.grad_clip as f64);
}
opt.step();
opt.zero_grad();
}
final_loss = epoch_loss.mean(steps);
let eval_due = epoch == 1
|| epoch == cfg.train.epochs
|| cfg.eval.interval_epochs == 0
|| epoch % cfg.eval.interval_epochs == 0;
if eval_due {
model.set_train(false); acc = eval_classify_samples(&model, &val, &val_labels, device)?;
model.set_train(bn_train); println!(
"[classify-imagenette] run={run_id} epoch={epoch}/{} loss={final_loss:.4} top1={acc:.3}",
cfg.train.epochs
);
}
log_epoch_metrics(run_dir, epoch, final_loss, "top1", acc, None);
}
save_checkpoint(&vs, &run_dir.join(CKPT_DIR))?;
Ok(TrainReport {
run_id: run_id.to_string(),
task: "classify".into(),
epochs: cfg.train.epochs,
final_loss,
metric: "top1".into(),
metric_value: acc,
secondary: None,
run_dir: run_dir.display().to_string(),
})
}
fn train_classify_synthetic(
cfg: &RunConfig,
run_id: &str,
run_dir: &Path,
) -> AvResult<TrainReport> {
let task_cfg = match cfg.model.tasks.first() {
Some(TaskCfg::Classify(c)) => (c.num_classes as u32, c.img_size),
_ => unreachable!("train_classify 只处理分类任务"),
};
let (num_classes, img_size) = task_cfg;
let device = resolve_device(cfg);
let vs = VarStore::new(device);
let model = build_model(&vs.root(), cfg)?;
apply_pretrain(&vs, cfg)?;
let mut opt = make_opt(&vs, cfg)?;
let mut rng = XorShift::new(cfg.seed);
let bs = cfg.train.batch_size as i64;
let mut final_loss = 0f32;
let mut acc = 0f32;
let bn_train = bn_train_mode(cfg);
for epoch in 1..=cfg.train.epochs {
model.set_train(bn_train);
opt.set_lr(schedule_lr(cfg, epoch));
let mut epoch_loss = EpochLossAcc::new();
let acc_steps = cfg.train.accumulate_steps.max(1) as usize;
opt.zero_grad();
for si in 0..STEPS_PER_EPOCH {
let (x, y, _) = synthetic_classify(&mut rng, bs, num_classes, img_size, device);
let loss = model.loss(&x, &TrainBatch::Classify { labels: y })?;
loss.backward();
epoch_loss.add(&loss);
if (si + 1) % acc_steps == 0 {
if cfg.train.grad_clip > 0.0 {
opt.clip_grad_norm(cfg.train.grad_clip as f64);
}
opt.step();
opt.zero_grad();
}
}
if !STEPS_PER_EPOCH.is_multiple_of(acc_steps) {
if cfg.train.grad_clip > 0.0 {
opt.clip_grad_norm(cfg.train.grad_clip as f64);
}
opt.step();
opt.zero_grad();
}
final_loss = epoch_loss.mean(STEPS_PER_EPOCH);
model.set_train(false); acc = eval_classify(&model, num_classes, img_size, device)?;
model.set_train(bn_train); println!(
"[classify-smoke] run={run_id} epoch={epoch}/{} loss={final_loss:.4} top1={acc:.3}",
cfg.train.epochs
);
log_epoch_metrics(run_dir, epoch, final_loss, "top1", acc, None);
}
save_checkpoint(&vs, &run_dir.join(CKPT_DIR))?;
Ok(TrainReport {
run_id: run_id.to_string(),
task: "classify".into(),
epochs: cfg.train.epochs,
final_loss,
metric: "top1".into(),
metric_value: acc,
secondary: None,
run_dir: run_dir.display().to_string(),
})
}
fn train_detect_synthetic(cfg: &RunConfig, run_id: &str, run_dir: &Path) -> AvResult<TrainReport> {
let (num_classes, img_size) = match cfg.model.tasks.first() {
Some(TaskCfg::Detect(d)) => (d.num_classes as u32, d.img_size),
_ => unreachable!("train_detect 只处理检测任务"),
};
let device = resolve_device(cfg);
let vs = VarStore::new(device);
let model = build_model(&vs.root(), cfg)?;
apply_pretrain(&vs, cfg)?;
let mut opt = make_opt(&vs, cfg)?;
let mut rng = XorShift::new(cfg.seed);
let bs = cfg.train.batch_size as i64;
let mut final_loss = 0f32;
let mut miou = 0f32;
let mut r50 = 0f32;
#[allow(unused_assignments)]
let mut det_stats = 0f32;
#[allow(unused_assignments)]
let mut dbg;
let bn_train = bn_train_mode(cfg);
for epoch in 1..=cfg.train.epochs {
model.set_train(bn_train);
opt.set_lr(schedule_lr(cfg, epoch));
let mut epoch_loss = EpochLossAcc::new();
let acc_steps = cfg.train.accumulate_steps.max(1) as usize;
opt.zero_grad();
for si in 0..STEPS_PER_EPOCH {
let (x, boxes, labels) = synthetic_detect(&mut rng, bs, num_classes, img_size, device);
let batch = TrainBatch::Detect {
boxes: boxes.iter().map(|b| vec![*b]).collect(),
labels: labels.iter().map(|&l| vec![l]).collect(),
};
if epoch % 10 == 0 {
av_tasks::models::LOSS_DEBUG.with(|d| *d.borrow_mut() = Some(String::new()));
}
let loss = model.loss(&x, &batch)?;
loss.backward();
epoch_loss.add(&loss);
if (si + 1) % acc_steps == 0 {
if cfg.train.grad_clip > 0.0 {
opt.clip_grad_norm(cfg.train.grad_clip as f64);
}
opt.step();
opt.zero_grad();
}
}
final_loss = epoch_loss.mean(STEPS_PER_EPOCH);
model.set_train(false); (miou, r50, det_stats, dbg) = eval_detect(&model, num_classes, img_size, device)?;
model.set_train(bn_train); if epoch % 10 == 0 || epoch == cfg.train.epochs {
println!(
"[detect-smoke] run={run_id} epoch={epoch}/{} loss={final_loss:.4} mIoU={miou:.3} R@0.5={r50:.3} dets/图={det_stats:.1} {dbg}",
cfg.train.epochs
);
}
log_epoch_metrics(
run_dir,
epoch,
final_loss,
"mean_iou",
miou,
Some(("recall@0.5", r50)),
);
}
save_checkpoint(&vs, &run_dir.join(CKPT_DIR))?;
Ok(TrainReport {
run_id: run_id.to_string(),
task: "detect".into(),
epochs: cfg.train.epochs,
final_loss,
metric: "mean_iou".into(),
metric_value: miou,
secondary: Some(("recall@0.5".into(), r50)),
run_dir: run_dir.display().to_string(),
})
}
fn print_detect_cache_stats(tiles: &[dataset::RawDetectSample]) {
let mb: usize = tiles.iter().map(|t| t.rgb.len()).sum::<usize>() / (1024 * 1024);
println!(
"[detect] 数据缓存: ram({} 张内容贴片,{mb}MB;全分辨率重采样一次性完成,逐 epoch 零大图编码)",
tiles.len()
);
}
#[derive(Debug, Clone, Copy)]
struct DetectDraw {
comp: augment::CompositeDraw,
partners: [usize; 3],
mix_partner: usize,
plan: AugmentPlan,
}
fn encode_detect_batch(
raw: &[dataset::RawDetectSample],
idx: &[usize],
draws: &[DetectDraw],
img_size: u32,
inorm: bool,
) -> AvResult<Vec<SampleTensor>> {
idx.par_iter()
.zip(draws)
.map(|(&i, d)| -> AvResult<SampleTensor> {
let mosaic_img = if d.comp.mosaic {
Some(dataset::mosaic4_raw([
&raw[i],
&raw[d.partners[0]],
&raw[d.partners[1]],
&raw[d.partners[2]],
])?)
} else {
None
};
let mix_img = if d.comp.mixup {
let base = mosaic_img.as_ref().unwrap_or(&raw[i]);
Some(dataset::mixup_raw(
base,
&raw[d.mix_partner],
d.comp.mixup_lam,
)?)
} else {
None
};
let composed: &dataset::RawDetectSample =
mix_img.as_ref().or(mosaic_img.as_ref()).unwrap_or(&raw[i]);
dataset::encode_detect_sample(
composed,
img_size,
Device::Cpu,
dataset::ResizeMode::Letterbox,
&d.plan,
inorm,
)
})
.collect()
}
fn train_detect_yolo(cfg: &RunConfig, run_id: &str, run_dir: &Path) -> AvResult<TrainReport> {
let (num_classes, img_size) = match cfg.model.tasks.first() {
Some(TaskCfg::Detect(d)) => (d.num_classes as u32, d.img_size),
_ => unreachable!("train_detect_yolo 只处理检测任务"),
};
let device = resolve_device(cfg);
let (root, pack) =
match cfg.data.pipeline {
DataPipeline::AvPack => {
let p =
cfg.data.sources.train.avpack.clone().ok_or_else(|| {
AvError::config("avpack 数据源缺 data.sources.train.avpack")
})?;
(p.clone(), Some(p))
}
_ => {
let d = cfg
.data
.sources
.train
.dir
.clone()
.ok_or_else(|| AvError::config("dir 数据源缺 data.sources.train.dir"))?;
(d, None)
}
};
let train_split = cfg.data.sources.train.split.as_deref().unwrap_or("train");
let aug_cfg = train_augment_cfg(cfg, TaskKind::Detect);
let train_raw = if let Some(a) = &aug_cfg {
println!(
"[augment] detect 训练增强生效: mosaic={} mixup={} flip={} hsv={:?} scale_jitter={:?} close_last_epochs={}",
a.mosaic, a.mixup, a.flip, a.hsv, a.scale_jitter, a.close_last_epochs
);
println!(
"[augment] detect 串联顺序: mosaic → mixup → flip → hsv → scale(mixup 仅检测/分类语义,关键点不适用)"
);
let cache_on = cfg.data.cache != "off";
let raws: Vec<dataset::RawDetectSample> = match &pack {
Some(p) => {
let raws = dataset::load_yolo_avpack_raw(p, train_split)?;
if cache_on {
let tiles = dataset::build_detect_tile_cache(raws, img_size)?;
print_detect_cache_stats(&tiles);
tiles
} else {
println!("[detect] 数据缓存: off(raw 逐 epoch 全分辨率路径)");
raws
}
}
None => {
if cache_on {
let tiles = dataset::build_detect_cache_from_dir(&root, train_split, img_size)?;
print_detect_cache_stats(&tiles);
tiles
} else {
println!("[detect] 数据缓存: off(raw 逐 epoch 全分辨率路径)");
dataset::load_yolo_dir_raw(&root, train_split)?
}
}
};
Some(Arc::new(raws))
} else {
None
};
let train = if train_raw.is_some() {
Vec::new()
} else {
match &pack {
Some(p) => dataset::load_yolo_avpack(
p,
train_split,
img_size,
Device::Cpu,
imagenet_norm(cfg),
)?,
None => dataset::load_yolo_dir(
&root,
train_split,
img_size,
Device::Cpu,
imagenet_norm(cfg),
)?,
}
};
let val = if cfg.data.pipeline == DataPipeline::AvPack {
let val_split = cfg.data.sources.val.split.as_deref().unwrap_or("val");
match cfg.data.sources.val.avpack.as_ref() {
Some(p) => {
dataset::load_yolo_avpack(p, val_split, img_size, Device::Cpu, imagenet_norm(cfg))?
}
None => dataset::load_yolo_avpack(
&root,
val_split,
img_size,
Device::Cpu,
imagenet_norm(cfg),
)?,
}
} else {
match cfg.data.sources.val.dir.as_ref() {
Some(d) => dataset::load_yolo_dir(
d,
cfg.data.sources.val.split.as_deref().unwrap_or("val"),
img_size,
Device::Cpu,
imagenet_norm(cfg),
)?,
None => {
dataset::load_yolo_dir(&root, "val", img_size, Device::Cpu, imagenet_norm(cfg))?
}
}
};
let train_n = train_raw.as_ref().map_or(train.len(), |r| r.len());
println!(
"[yolo] {}={} train={} val={} classes={} img_size={}",
if pack.is_some() { "avpack" } else { "root" },
root.display(),
train_n,
val.len(),
num_classes,
img_size
);
let vs = VarStore::new(device);
let model = build_model(&vs.root(), cfg)?;
apply_pretrain(&vs, cfg)?;
let mut opt = make_opt(&vs, cfg)?;
let mut rng = XorShift::new(cfg.seed);
let bs = (cfg.train.batch_size as usize).max(1);
let mut final_loss = 0f32;
let mut miou = 0f32;
let mut r50 = 0f32;
let bn_train = bn_train_mode(cfg);
let loss_timing = std::env::var("AV_LOSS_TIMING").is_ok();
let step_clock = std::sync::atomic::AtomicUsize::new(0);
let amp = cfg.train.amp && device != Device::Cpu;
if amp {
println!("[amp] fp16 autocast + 动态梯度缩放启用");
}
let mut scaler = GradScaler::new();
let mut ema = WeightEma::new(&vs, cfg.train.ema_decay);
let mut best_fitness = f32::NEG_INFINITY;
let mut best_state: Option<Vec<(String, Tensor)>> = None;
for epoch in 1..=cfg.train.epochs {
model.set_train(bn_train);
opt.set_lr(schedule_lr(cfg, epoch));
let mut order: Vec<usize> = (0..train_n).collect();
for i in (1..order.len()).rev() {
let j = rng.next_usize(i + 1);
order.swap(i, j);
}
let strong_on = aug_cfg
.as_ref()
.map(|a| strong_aug_on(a, cfg, epoch))
.unwrap_or(false);
let mut aug_rng = epoch_aug_rng(cfg.seed, epoch);
let draws: Vec<DetectDraw> = match (&train_raw, &aug_cfg) {
(Some(raw), Some(a)) => (0..order.len())
.map(|_| -> DetectDraw {
if strong_on {
let comp = augment::draw_composite(a, &mut aug_rng);
let partners = if comp.mosaic {
[
aug_rng.next_usize(raw.len()),
aug_rng.next_usize(raw.len()),
aug_rng.next_usize(raw.len()),
]
} else {
[0; 3]
};
let mix_partner = if comp.mixup {
aug_rng.next_usize(raw.len())
} else {
0
};
let plan = augment::draw_plan(a, &mut aug_rng);
DetectDraw {
comp,
partners,
mix_partner,
plan,
}
} else {
DetectDraw {
comp: augment::CompositeDraw::none(),
partners: [0; 3],
mix_partner: 0,
plan: AugmentPlan::none(),
}
}
})
.collect(),
_ => Vec::new(),
};
let mut epoch_loss = EpochLossAcc::new();
let mut steps = 0usize;
let acc_steps = cfg.train.accumulate_steps.max(1) as usize;
opt.zero_grad();
let chunk_starts: Vec<usize> = (0..order.len()).step_by(bs).collect();
let mut in_flight: Option<std::thread::JoinHandle<AvResult<Vec<SampleTensor>>>> = None;
for (ci, &start) in chunk_starts.iter().enumerate() {
let end = (start + bs).min(order.len());
let (x, boxes, labels) = match in_flight.take() {
Some(h) => {
let batch_samples: Vec<SampleTensor> =
h.join().map_err(|_| AvError::train("数据编码线程崩溃"))??;
let x = dataset::stack_samples(&batch_samples)?.to_device(device);
(
x,
batch_samples.iter().map(|s| s.boxes.clone()).collect(),
batch_samples.iter().map(|s| s.labels.clone()).collect(),
)
}
None => match (&train_raw, &aug_cfg) {
(Some(raw), Some(_)) => {
let batch_samples = encode_detect_batch(
raw,
&order[start..end],
&draws[start..end],
img_size,
imagenet_norm(cfg),
)?;
let x = dataset::stack_samples(&batch_samples)?.to_device(device);
(
x,
batch_samples.iter().map(|s| s.boxes.clone()).collect(),
batch_samples.iter().map(|s| s.labels.clone()).collect(),
)
}
_ => {
let xs: Vec<&Tensor> =
order[start..end].iter().map(|&i| &train[i].x).collect();
(
Tensor::stack(&xs, 0).to_device(device),
order[start..end]
.iter()
.map(|&i| train[i].boxes.clone())
.collect(),
order[start..end]
.iter()
.map(|&i| train[i].labels.clone())
.collect(),
)
}
},
};
if let (Some(raw), Some(_)) = (&train_raw, &aug_cfg) {
if let Some(&next_start) = chunk_starts.get(ci + 1) {
let next_end = (next_start + bs).min(order.len());
let idx = order[next_start..next_end].to_vec();
let draws_next = draws[next_start..next_end].to_vec();
let raw = Arc::clone(raw);
let inorm = imagenet_norm(cfg);
in_flight = Some(std::thread::spawn(move || {
encode_detect_batch(&raw, &idx, &draws_next, img_size, inorm)
}));
}
}
let batch = TrainBatch::Detect { boxes, labels };
let timing = loss_timing;
let t0 = std::time::Instant::now();
let loss = if amp {
tch::autocast(true, || model.loss(&x, &batch))?.to_kind(Kind::Float)
} else {
model.loss(&x, &batch)?
};
let t_loss = t0.elapsed();
let backward_from = if amp {
scaler.scaled(&loss)
} else {
loss.copy()
};
backward_from.backward();
epoch_loss.add(&loss);
steps += 1;
if steps.is_multiple_of(acc_steps) {
optimizer_step(
amp,
&mut scaler,
&vs,
&mut opt,
&mut ema,
cfg.train.grad_clip,
);
}
if timing {
let t_all = t0.elapsed();
let k = step_clock.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
if k.is_multiple_of(20) {
let bwd = t_all - t_loss;
eprintln!("[timing] loss={t_loss:?} backward+step={bwd:?} total={t_all:?}");
}
}
}
if !steps.is_multiple_of(acc_steps) {
optimizer_step(
amp,
&mut scaler,
&vs,
&mut opt,
&mut ema,
cfg.train.grad_clip,
);
}
final_loss = epoch_loss.mean(steps);
let TaskModel::Detect(m) = &model else {
unreachable!("检测任务模型类型")
};
let eval_due = epoch == 1
|| epoch == cfg.train.epochs
|| cfg.eval.interval_epochs == 0
|| epoch % cfg.eval.interval_epochs == 0;
if eval_due {
model.set_train(false); let ema_saved = ema.apply_to(&vs);
let evm = eval_detect_samples(m, &val, device, amp)?;
ema.restore(&vs, &ema_saved);
model.set_train(bn_train); (miou, r50) = (evm.miou, evm.r50);
println!(
"[detect-yolo] run={run_id} epoch={epoch}/{} loss={final_loss:.4} mIoU={miou:.3} R@0.5={r50:.3} mAP50={:.3} mAP50:95={:.3}",
cfg.train.epochs, evm.map50, evm.map50_95
);
let fitness = 0.9 * evm.map50_95 + 0.1 * evm.map50;
if fitness > best_fitness {
best_fitness = fitness;
best_state = Some(ema.snapshot_cpu());
}
}
log_epoch_metrics(
run_dir,
epoch,
final_loss,
"mean_iou",
miou,
Some(("recall@0.5", r50)),
);
}
let final_state = best_state.unwrap_or_else(|| ema.snapshot_cpu());
save_named_variables(&final_state, &run_dir.join(CKPT_DIR))?;
Ok(TrainReport {
run_id: run_id.to_string(),
task: "detect".into(),
epochs: cfg.train.epochs,
final_loss,
metric: "mean_iou".into(),
metric_value: miou,
secondary: Some(("recall@0.5".into(), r50)),
run_dir: run_dir.display().to_string(),
})
}
#[derive(Debug, Clone)]
struct DetectEvalMetrics {
miou: f32,
r50: f32,
map50: f32,
map50_95: f32,
#[allow(dead_code)]
dets_per_img: f32,
#[allow(dead_code)]
dbg: String,
}
fn eval_detect_samples(
m: &av_tasks::models::DetectModel,
val: &[SampleTensor],
device: Device,
amp: bool,
) -> AvResult<DetectEvalMetrics> {
let n = val.len();
if n == 0 {
return Err(AvError::data("验证集为空"));
}
let mut per_image: Vec<Vec<av_core::types::Detection>> = Vec::new();
for chunk in val.chunks(32) {
let x = dataset::stack_samples(chunk)?.to_device(device);
let pred = if amp {
tch::autocast(true, || m.predict(&x, 0.1, 0.5))?
} else {
m.predict(&x, 0.1, 0.5)?
};
per_image.extend(pred);
}
let mut best_ious = vec![0f32; n];
let mut class_ok = vec![false; n];
for (gi, s) in val.iter().enumerate() {
for d in &per_image[gi] {
let mut matched = false;
for (b, &label) in s.boxes.iter().zip(&s.labels) {
let g = Aabb::new(b[0], b[1], b[2], b[3]);
best_ious[gi] = best_ious[gi].max(g.iou(&d.bbox));
if d.class_id == label {
matched = true;
}
}
if matched {
class_ok[gi] = true;
}
}
}
let matched_cnt = best_ious.iter().filter(|&&v| v >= 0.5).count();
let total_dets: usize = per_image.iter().map(|d| d.len()).sum();
let mut parts = vec![format!(
"类对率={:.2}",
class_ok.iter().filter(|&&v| v).count() as f32 / n as f32
)];
for gi in 0..n.min(3) {
parts.push(format!(
"图{gi} best_iou={:.2} dets={}",
best_ious[gi],
per_image[gi].len()
));
}
let mut ev = CocoEvaluator::new();
for (gi, s) in val.iter().enumerate() {
let gts: Vec<GtBox> = s
.boxes
.iter()
.zip(&s.labels)
.map(|(b, &l)| GtBox::new(Aabb::new(b[0], b[1], b[2], b[3]), l))
.collect();
ev.update(gi as u32, &per_image[gi], >s);
}
let map = ev.finalize();
Ok(DetectEvalMetrics {
miou: best_ious.iter().sum::<f32>() / n as f32,
r50: matched_cnt as f32 / n as f32,
map50: map.map50,
map50_95: map.map50_95,
dets_per_img: total_dets as f32 / n as f32,
dbg: parts.join(" | "),
})
}
fn eval_classify(
model: &TaskModel,
num_classes: u32,
img_size: u32,
device: Device,
) -> AvResult<f32> {
let mut rng = XorShift::new(0x5EED_0001);
let (x, _, labels) = synthetic_classify(&mut rng, EVAL_BATCH, num_classes, img_size, device);
let TaskModel::Classify(m) = model else {
return Err(AvError::train("模型与任务不匹配"));
};
let (pred, _) = m.predict(&x)?;
Ok(pred.iter().zip(&labels).filter(|(a, b)| a == b).count() as f32 / labels.len() as f32)
}
fn eval_classify_samples(
model: &TaskModel,
val: &[ClassifySample],
labels: &[u32],
device: Device,
) -> AvResult<f32> {
let n = val.len();
if n == 0 || labels.len() != n {
return Err(AvError::data("分类验证集为空或标签数不匹配"));
}
let TaskModel::Classify(m) = model else {
return Err(AvError::train("模型与任务不匹配"));
};
let mut correct = 0usize;
for (ci, chunk) in val.chunks(EVAL_BATCH as usize).enumerate() {
let start = ci * EVAL_BATCH as usize;
let x = dataset::stack_classify(chunk)?.to_device(device);
let (pred, _) = m.predict(&x)?;
correct += pred
.iter()
.zip(&labels[start..start + chunk.len()])
.filter(|(a, b)| a == b)
.count();
}
Ok(correct as f32 / n as f32)
}
fn eval_detect(
model: &TaskModel,
num_classes: u32,
img_size: u32,
device: Device,
) -> AvResult<(f32, f32, f32, String)> {
let mut rng = XorShift::new(0x5EED_0002);
let n = 64usize;
let (x, boxes, labels) = synthetic_detect(&mut rng, n as i64, num_classes, img_size, device);
let TaskModel::Detect(m) = model else {
return Err(AvError::train("模型与任务不匹配"));
};
let per_image = m.predict(&x, 0.25, 0.5)?;
let total_dets: usize = per_image.iter().map(|d| d.len()).sum();
let mut best_ious = vec![0f32; n];
let mut class_ok = vec![false; n];
let gt_boxes: Vec<Aabb> = boxes
.iter()
.map(|b| Aabb::new(b[0], b[1], b[2], b[3]))
.collect();
for (gi, g) in gt_boxes.iter().enumerate() {
for d in &per_image[gi] {
if d.class_id == labels[gi] {
class_ok[gi] = true;
}
if d.class_id != labels[gi] {
continue;
}
best_ious[gi] = best_ious[gi].max(g.iou(&d.bbox));
}
}
let matched = best_ious.iter().filter(|&&v| v >= 0.5).count();
let dbg = {
let mut parts = vec![format!(
"类对率={:.2} 无类过滤mIoU={:.2}",
class_ok.iter().filter(|&&v| v).count() as f32 / n as f32,
{
let mut best_all = vec![0f32; n];
for (gi, g) in gt_boxes.iter().enumerate() {
for d in &per_image[gi] {
best_all[gi] = best_all[gi].max(g.iou(&d.bbox));
}
}
best_all.iter().sum::<f32>() / n as f32
}
)];
for gi in 0..4 {
let top = per_image[gi].iter().max_by(|a, b| {
a.score
.partial_cmp(&b.score)
.unwrap_or(std::cmp::Ordering::Equal)
});
match top {
Some(d) => parts.push(format!(
"图{gi} pred=({:.0},{:.0},{:.0},{:.0}) c{} | gt=({:.0},{:.0},{:.0},{:.0}) c{} iou={:.2}",
d.bbox.x1,
d.bbox.y1,
d.bbox.x2,
d.bbox.y2,
d.class_id,
boxes[gi][0],
boxes[gi][1],
boxes[gi][2],
boxes[gi][3],
labels[gi],
gt_boxes[gi].iou(&d.bbox)
)),
None => parts.push(format!("图{gi} 无检出")),
}
}
parts.join(" | ")
};
Ok((
best_ious.iter().sum::<f32>() / n as f32,
matched as f32 / n as f32,
total_dets as f32 / n as f32,
dbg,
))
}
fn synthetic_classify(
rng: &mut XorShift,
n: i64,
num_classes: u32,
img: u32,
device: Device,
) -> (Tensor, Tensor, Vec<u32>) {
let s = img as usize;
let n = n as usize;
let mut buf = vec![0f32; n * 3 * s * s];
let mut labels = vec![0u32; n];
for (ni, lbl) in labels.iter_mut().enumerate() {
let c = rng.next_usize(num_classes as usize);
*lbl = c as u32;
let base = ni * 3 * s * s;
for v in buf[base..base + 3 * s * s].iter_mut() {
*v = rng.next_f32() * 0.3;
}
let ch = ni * 3 * s * s + (c % 3) * s * s;
let bx = (c * 13 + 7) % (s - 16);
let by = (c * 29 + 3) % (s - 16);
for yy in by..by + 16 {
for xx in bx..bx + 16 {
buf[ch + yy * s + xx] += 1.5;
}
}
}
let y: Vec<i64> = labels.iter().map(|&v| v as i64).collect();
let x = Tensor::from_slice(&buf)
.to_kind(Kind::Float)
.clamp(0.0, 1.0) .to_device(device)
.reshape([n as i64, 3, s as i64, s as i64]);
let y = Tensor::from_slice(&y).to_device(device);
(x, y, labels)
}
fn synthetic_detect(
rng: &mut XorShift,
n: i64,
num_classes: u32,
img: u32,
device: Device,
) -> (Tensor, Vec<[f32; 4]>, Vec<u32>) {
let s = img as usize;
let n = n as usize;
let mut buf = vec![0f32; n * 3 * s * s];
let mut boxes = Vec::with_capacity(n);
let mut labels = Vec::with_capacity(n);
for ni in 0..n {
for v in buf[ni * 3 * s * s..(ni + 1) * 3 * s * s].iter_mut() {
*v = rng.next_f32() * 0.3;
}
let sw = rng.next_range(0.18, 0.4) * s as f32;
let sh = rng.next_range(0.18, 0.4) * s as f32;
let cx = rng.next_range(sw / 2.0 + 1.0, s as f32 - sw / 2.0 - 1.0);
let cy = rng.next_range(sh / 2.0 + 1.0, s as f32 - sh / 2.0 - 1.0);
let (x1, y1, x2, y2) = (
(cx - sw / 2.0).round(),
(cy - sh / 2.0).round(),
(cx + sw / 2.0).round(),
(cy + sh / 2.0).round(),
);
labels.push(rng.next_usize(num_classes as usize) as u32);
boxes.push([x1, y1, x2, y2]);
let ch_off = ni * 3 * s * s + (labels[ni] as usize % 3) * s * s;
for yy in (y1 as usize)..(y2 as usize).min(s) {
for xx in (x1 as usize)..(x2 as usize).min(s) {
buf[ch_off + yy * s + xx] += 1.5;
}
}
}
let x = Tensor::from_slice(&buf)
.to_kind(Kind::Float)
.clamp(0.0, 1.0) .to_device(device)
.reshape([n as i64, 3, s as i64, s as i64]);
(x, boxes, labels)
}
pub(crate) fn testing_load_model(cfg: &RunConfig, weights: &Path) -> AvResult<TaskModel> {
load_model(cfg, weights).map(|(model, _)| model)
}
pub(crate) fn testing_synthetic_detect(
rng: &mut XorShift,
n: i64,
num_classes: u32,
img_size: u32,
) -> (Tensor, Vec<[f32; 4]>, Vec<u32>) {
synthetic_detect(rng, n, num_classes, img_size, Device::Cpu)
}
pub(crate) fn testing_eval_kp_samples(
m: &KeypointModel,
samples: &[KeypointSample],
) -> AvResult<(f32, f32, usize, usize)> {
eval_kp_samples(m, samples, Device::Cpu)
}
#[cfg(test)]
mod fragment_merge_tests {
use super::merge_tile_fragments;
use av_core::geometry::Aabb;
fn det(
x1: f32,
y1: f32,
x2: f32,
y2: f32,
score: f32,
class_id: u32,
) -> av_core::types::Detection {
av_core::types::Detection {
bbox: Aabb::new(x1, y1, x2, y2),
score,
class_id,
angle: None,
keypoints: None,
}
}
#[test]
fn two_overlapping_fragments_merge_into_one() {
let out = merge_tile_fragments(
vec![
det(0.0, 0.0, 40.0, 40.0, 0.9, 0),
det(20.0, 0.0, 60.0, 40.0, 0.6, 0),
],
0.3,
25.0,
);
assert_eq!(out.len(), 1, "两个重叠碎片必须合并为一个");
assert!((out[0].score - 0.9).abs() < 1e-6, "代表取分数最高者");
let b = &out[0].bbox;
assert!((b.x1 - 0.0).abs() < 1e-6 && (b.y1 - 0.0).abs() < 1e-6);
assert!(
(b.x2 - 60.0).abs() < 1e-6 && (b.y2 - 40.0).abs() < 1e-6,
"并集框 {:?}",
b
);
}
#[test]
fn merge_requires_iou_and_distance_and_same_class() {
let dets = vec![
det(20.0, 0.0, 60.0, 40.0, 0.6, 0), det(100.0, 100.0, 300.0, 300.0, 0.7, 0), det(0.0, 0.0, 40.0, 40.0, 0.9, 0), det(0.0, 0.0, 300.0, 300.0, 0.8, 0), det(1000.0, 1000.0, 1020.0, 1020.0, 0.5, 0), det(2000.0, 2000.0, 2040.0, 2040.0, 0.9, 1), det(2020.0, 2000.0, 2060.0, 2040.0, 0.6, 0), ];
let out = merge_tile_fragments(dets, 0.3, 25.0);
assert_eq!(out.len(), 6, "只有 a+b 合并,得 {:?}", out);
let a = out
.iter()
.find(|k| (k.score - 0.9).abs() < 1e-6 && k.class_id == 0)
.expect("代表框必须保留");
assert!((a.bbox.x2 - 60.0).abs() < 1e-6 && (a.bbox.y2 - 40.0).abs() < 1e-6);
assert!(out.iter().any(|k| (k.score - 0.8).abs() < 1e-6));
assert!(out.iter().any(|k| (k.score - 0.7).abs() < 1e-6));
assert_eq!(
out.iter().filter(|k| (k.score - 0.9).abs() < 1e-6).count(),
2
);
assert!(out
.iter()
.any(|k| (k.score - 0.6).abs() < 1e-6 && k.class_id == 0));
assert!(out.iter().any(|k| (k.score - 0.5).abs() < 1e-6));
}
#[test]
fn chain_merge_absorbs_via_grown_representative() {
let a = det(0.0, 0.0, 40.0, 40.0, 0.9, 0);
let b = det(20.0, 0.0, 60.0, 40.0, 0.6, 0); let c = det(40.0, 0.0, 80.0, 40.0, 0.5, 0); let out = merge_tile_fragments(vec![a, b, c], 0.3, 25.0);
assert_eq!(out.len(), 2, "c 中心距 30 超限,独立保留: {:?}", out);
assert!((out[0].bbox.x2 - 60.0).abs() < 1e-6);
}
}
#[cfg(test)]
mod seg_precision_tests {
use super::greedy_mask_match;
fn ious(rows: &[&[f32]]) -> Vec<Vec<f32>> {
rows.iter().map(|r| r.to_vec()).collect()
}
#[test]
fn threshold_and_class_constraints() {
let preds = vec![(0u32, 0.9f32), (1, 0.8), (0, 0.7)];
let m = ious(&[&[0.8, 0.0], &[0.0, 0.3], &[0.2, 0.0]]);
assert_eq!(greedy_mask_match(preds, &[0, 1], &m, 0.5), 1);
}
#[test]
fn one_to_one_no_double_count() {
let preds = vec![(0u32, 0.9f32), (0, 0.8)];
let m = ious(&[&[0.9], &[0.9]]);
assert_eq!(greedy_mask_match(preds, &[0], &m, 0.5), 1);
}
#[test]
fn score_order_claims_best_match_first() {
let preds = vec![(0u32, 0.9f32), (0, 0.8)];
let m = ious(&[&[0.55, 0.95], &[0.95, 0.0]]);
assert_eq!(greedy_mask_match(preds, &[0, 0], &m, 0.5), 2);
}
#[test]
fn no_predictions_zero_tp() {
assert_eq!(greedy_mask_match(vec![], &[0], &[], 0.5), 0);
}
}
#[cfg(test)]
mod seg_cache_mode_tests {
use super::{resolve_seg_cache_mode, SegCacheMode};
use tch::Device;
#[test]
fn auto_selects_gpu_only_within_budget_and_cuda() {
let cuda = Device::Cuda(0);
let cpu = Device::Cpu;
assert_eq!(
resolve_seg_cache_mode("auto", 526, 640, cuda),
SegCacheMode::Gpu
);
assert_eq!(
resolve_seg_cache_mode("auto", 200_000, 640, cuda),
SegCacheMode::Ram
);
assert_eq!(
resolve_seg_cache_mode("auto", 526, 640, cpu),
SegCacheMode::Ram
);
}
#[test]
fn explicit_values_override_auto() {
let cuda = Device::Cuda(0);
let cpu = Device::Cpu;
assert_eq!(
resolve_seg_cache_mode("off", 526, 640, cuda),
SegCacheMode::Off
);
assert_eq!(
resolve_seg_cache_mode("ram", 526, 640, cuda),
SegCacheMode::Ram
);
assert_eq!(
resolve_seg_cache_mode("gpu", 526, 640, cuda),
SegCacheMode::Gpu
);
assert_eq!(
resolve_seg_cache_mode("gpu", 526, 640, cpu),
SegCacheMode::Ram
);
}
}
#[cfg(test)]
mod ckpt_epoch_tests {
use super::{read_checkpoint_epoch, save_checkpoint_epoch};
use tch::nn::{self, VarStore};
use tch::Device;
#[test]
fn checkpoint_epoch_roundtrip() {
let dir = std::env::temp_dir().join(format!("av-ckpt-ep-{}", std::process::id()));
let _ = std::fs::remove_dir_all(&dir);
let vs = VarStore::new(Device::Cpu);
let _fc = nn::linear(&(vs.root() / "head") / "fc", 3, 2, Default::default());
save_checkpoint_epoch(&vs, &dir, 57).unwrap();
assert_eq!(read_checkpoint_epoch(&dir), Some(57));
std::fs::remove_dir_all(&dir).unwrap();
assert_eq!(read_checkpoint_epoch(&dir), None, "目录不存在应返回 None");
}
}
#[cfg(test)]
mod amp_ema_tests {
use super::*;
#[test]
fn grad_scaler_grows_on_stable_steps_and_backs_off_on_inf() {
let mut s = GradScaler::new();
assert_eq!(s.scale, 65536.0);
for _ in 0..(s.growth_interval - 1) {
s.update(true);
}
assert_eq!(s.scale, 65536.0, "未达 interval 不放大");
s.update(true);
assert_eq!(s.scale, 131072.0, "达到 interval 放大 2×");
s.update(false);
assert_eq!(s.scale, 65536.0, "异常步进回退 0.5× 并重置计数");
assert_eq!(s.steps_since_growth, 0);
}
#[test]
fn grad_scaler_unscale_divides_finite_grads() {
let vs = VarStore::new(Device::Cpu);
let a = vs.root().var("a", &[2_i64], tch::nn::Init::Const(1.0));
let loss = (&a * &a).sum(Kind::Float);
loss.backward();
let mut scaler = GradScaler::new();
scaler.scale = 2.0;
let vars = vs.trainable_variables();
assert!(scaler.unscale_and_check(&vars), "有限梯度应通过检测");
let g = vars[0].grad();
let (g1, g2) = (g.double_value(&[0]), g.double_value(&[1]));
assert!((g1 - 1.0).abs() < 1e-6 && (g2 - 1.0).abs() < 1e-6);
}
#[test]
fn grad_scaler_rejects_nan_grads() {
let t = Tensor::from_slice(&[f32::NAN]).set_requires_grad(true);
let loss = (&t * &t).sum(Kind::Float);
loss.backward();
let holder = VarStore::new(Device::Cpu);
let mut vars = holder.trainable_variables(); vars.push(t);
let mut scaler = GradScaler::new();
assert!(!scaler.unscale_and_check(&vars), "NaN 梯度必须被检出");
}
#[test]
fn weight_ema_scheduled_decay_update_and_swap_roundtrip() {
let vs = VarStore::new(Device::Cpu);
let v = vs.root().var("w", &[1_i64], tch::nn::Init::Const(1.0));
let mut ema = WeightEma::new(&vs, 0.999);
ema.shadow[0].1 = Tensor::from_slice(&[1.0f32]);
let mut v = v;
v.set_data(&Tensor::from_slice(&[3.0f32]));
ema.update(&vs);
let d = 2.0 / 11.0;
let expect = 3.0 - 2.0 * d;
let got = ema.shadow[0].1.double_value(&[0]);
assert!(
(got - expect).abs() < 1e-5,
"EMA 更新偏差: {got} vs {expect}"
);
let saved = ema.apply_to(&vs);
assert!(
(v.double_value(&[0]) - expect).abs() < 1e-6,
"apply 后应为影子值"
);
ema.restore(&vs, &saved);
assert!(
(v.double_value(&[0]) - 3.0).abs() < 1e-6,
"restore 后应还原原值"
);
}
#[test]
fn save_named_variables_roundtrip_matches_checkpoint_naming() {
let vs = VarStore::new(Device::Cpu);
let a = vs
.root()
.var("head/cls_weight", &[1_i64], tch::nn::Init::Const(5.0));
let dir = std::env::temp_dir().join(format!("av-named-ckpt-{}", std::process::id()));
let vars = vec![("head/cls_weight".to_string(), a.copy())];
save_named_variables(&vars, &dir).unwrap();
let mut vs2 = VarStore::new(Device::Cpu);
let b = vs2
.root()
.var("head/cls_weight", &[1_i64], tch::nn::Init::Const(0.0));
load_checkpoint(&mut vs2, &dir).unwrap();
assert!(
(b.double_value(&[0]) - a.double_value(&[0])).abs() < 1e-6,
"保存-加载应逐位一致"
);
std::fs::remove_dir_all(&dir).unwrap();
let _ = b;
}
}