use std::path::Path;
use tch::nn;
use tch::nn::Module;
use tch::{IndexOp, Kind, Tensor};
use av_core::config::BackboneCfg;
use av_core::error::{AvError, AvResult};
use av_core::traits::{BackboneSpec, BaseBackbone, FeatureMap, FeaturePyramid, LevelSpec};
pub const FAMILY_NAME: &str = "dinov2";
const PATCH_SIZE: i64 = 14;
const EMBED_DIM: i64 = 384;
const DEPTH: usize = 12;
const NUM_HEADS: i64 = 6;
const PYRAMID_CHANNELS: i64 = 256;
const LN_EPS: f64 = 1e-6;
fn ln_config() -> nn::LayerNormConfig {
nn::LayerNormConfig {
eps: LN_EPS,
..Default::default()
}
}
struct PatchEmbed {
proj: nn::Conv2D,
}
impl PatchEmbed {
fn new(p: nn::Path) -> Self {
let config = nn::ConvConfig {
stride: PATCH_SIZE,
..Default::default()
};
Self {
proj: nn::conv2d(p / "proj", 3, EMBED_DIM, PATCH_SIZE, config),
}
}
fn forward(&self, xs: &Tensor) -> Tensor {
let xs = xs.apply(&self.proj); let (b, c, h, w) = xs.size4().unwrap();
xs.reshape([b, c, h * w]).transpose(1, 2) }
}
struct Attention {
qkv: nn::Linear,
proj: nn::Linear,
num_heads: i64,
scale: f64,
}
impl Attention {
fn new(p: nn::Path, dim: i64, num_heads: i64) -> Self {
Self {
qkv: nn::linear(&p / "qkv", dim, dim * 3, Default::default()),
proj: nn::linear(&p / "proj", dim, dim, Default::default()),
num_heads,
scale: 1.0 / (((dim / num_heads) as f64).sqrt()),
}
}
fn forward(&self, xs: &Tensor) -> Tensor {
let (b, n, c) = xs.size3().unwrap();
let qkv = self
.qkv
.forward(xs)
.reshape([b, n, 3, self.num_heads, c / self.num_heads])
.permute([2, 0, 3, 1, 4]);
let q = qkv.get(0) * self.scale;
let k = qkv.get(1);
let v = qkv.get(2);
let attn = q.matmul(&k.transpose(-2, -1)).softmax(-1, Kind::Float);
attn.matmul(&v)
.transpose(1, 2)
.reshape([b, n, c])
.apply(&self.proj)
}
}
struct Mlp {
fc1: nn::Linear,
fc2: nn::Linear,
}
impl Mlp {
fn new(p: nn::Path, dim: i64) -> Self {
Self {
fc1: nn::linear(&p / "fc1", dim, dim * 4, Default::default()),
fc2: nn::linear(&p / "fc2", dim * 4, dim, Default::default()),
}
}
fn forward(&self, xs: &Tensor) -> Tensor {
xs.apply(&self.fc1).gelu("none").apply(&self.fc2)
}
}
struct Block {
norm1: nn::LayerNorm,
attn: Attention,
ls1: Tensor,
norm2: nn::LayerNorm,
mlp: Mlp,
ls2: Tensor,
}
impl Block {
fn new(p: nn::Path, dim: i64, num_heads: i64) -> Self {
Self {
norm1: nn::layer_norm(&p / "norm1", vec![dim], ln_config()),
attn: Attention::new(&p / "attn", dim, num_heads),
ls1: (&p / "ls1").var("gamma", &[dim], nn::Init::Const(1.0)),
norm2: nn::layer_norm(&p / "norm2", vec![dim], ln_config()),
mlp: Mlp::new(&p / "mlp", dim),
ls2: (&p / "ls2").var("gamma", &[dim], nn::Init::Const(1.0)),
}
}
fn forward(&self, xs: &Tensor) -> Tensor {
let x = xs + &(self.attn.forward(&xs.apply(&self.norm1)) * &self.ls1);
&x + &(self.mlp.forward(&x.apply(&self.norm2)) * &self.ls2)
}
}
pub struct DinoV2Backbone {
patch_embed: PatchEmbed,
cls_token: Tensor,
pos_embed: Tensor,
blocks: Vec<Block>,
norm: nn::LayerNorm,
conv_s8: nn::Conv2D,
conv_s16: nn::Conv2D,
conv_s32: nn::Conv2D,
grid_side: i64,
}
impl DinoV2Backbone {
pub fn new(p: &nn::Path, _cfg: &BackboneCfg, img_size: u32) -> AvResult<Self> {
let img = img_size as i64;
if img % PATCH_SIZE != 0 {
return Err(AvError::config(format!(
"{FAMILY_NAME}: img_size 必须为 {PATCH_SIZE} 的倍数(patch 网格整数),得到 {img}"
)));
}
if img % 32 != 0 {
return Err(AvError::config(format!(
"{FAMILY_NAME}: img_size 必须为 32 的倍数(金字塔 stride 8/16/32 整除 + \
项目 32 对齐约定),得到 {img}"
)));
}
let grid_side = img / PATCH_SIZE;
let n = grid_side * grid_side;
let cc = nn::ConvConfig {
padding: 1,
..Default::default()
};
Ok(Self {
patch_embed: PatchEmbed::new(p / "patch_embed"),
cls_token: p.var("cls_token", &[1, 1, EMBED_DIM], nn::Init::Const(0.)),
pos_embed: p.var("pos_embed", &[1, 1 + n, EMBED_DIM], nn::Init::Const(0.)),
blocks: (0..DEPTH)
.map(|i| Block::new(p / "blocks" / i, EMBED_DIM, NUM_HEADS))
.collect(),
norm: nn::layer_norm(p / "norm", vec![EMBED_DIM], ln_config()),
conv_s8: nn::conv2d(
p / "pyramid" / "stride8",
EMBED_DIM,
PYRAMID_CHANNELS,
3,
cc,
),
conv_s16: nn::conv2d(
p / "pyramid" / "stride16",
EMBED_DIM,
PYRAMID_CHANNELS,
3,
cc,
),
conv_s32: nn::conv2d(
p / "pyramid" / "stride32",
EMBED_DIM,
PYRAMID_CHANNELS,
3,
cc,
),
grid_side,
})
}
pub fn pooled_channels(&self) -> i64 {
EMBED_DIM
}
pub fn stride_channels(&self, stride: u32) -> AvResult<i64> {
match stride {
8 | 16 | 32 => Ok(PYRAMID_CHANNELS),
other => Err(AvError::shape(format!(
"{FAMILY_NAME} 不存在 stride {other} 的特征层"
))),
}
}
fn prepare_tokens(&self, xs: &Tensor) -> Tensor {
let (b, _c, h, w) = xs.size4().unwrap();
let patch = self.patch_embed.forward(xs);
let cls = self.cls_token.expand([b, -1, -1], false);
let x = Tensor::cat(&[&cls, &patch], 1);
let pos = self.interpolate_pos_encoding(patch.size()[1], w, h);
&x + &pos
}
fn interpolate_pos_encoding(&self, npatch: i64, w: i64, h: i64) -> Tensor {
let n = self.pos_embed.size()[1] - 1;
if npatch == n && w == h {
return self.pos_embed.copy();
}
let class_pos = self.pos_embed.i((.., ..1));
let patch_pos = self.pos_embed.i((.., 1..));
let dim = EMBED_DIM;
let sqrt_n = (n as f64).sqrt();
let (w0, h0) = ((w / PATCH_SIZE) as f64 + 0.1, (h / PATCH_SIZE) as f64 + 0.1);
let patch_pos = patch_pos
.reshape([1, sqrt_n as i64, sqrt_n as i64, dim])
.permute([0, 3, 1, 2])
.upsample_bicubic2d([w0 as i64, h0 as i64], false, w0 / sqrt_n, h0 / sqrt_n)
.permute([0, 2, 3, 1])
.reshape([1, -1, dim]);
Tensor::cat(&[&class_pos, &patch_pos], 1)
}
fn encode_tokens(&self, xs: &Tensor) -> Tensor {
let mut t = self.prepare_tokens(xs);
for blk in &self.blocks {
t = blk.forward(&t);
}
t.apply(&self.norm)
}
fn patch_grid(&self, tokens: &Tensor) -> Tensor {
let b = tokens.size()[0];
tokens.i((.., 1..)).permute([0, 2, 1]).reshape([
b,
EMBED_DIM,
self.grid_side,
self.grid_side,
])
}
pub fn load_dinov2_weights(&mut self, path: &Path) -> AvResult<DinoLoadStats> {
let sources = Tensor::read_safetensors(path).map_err(|e| {
AvError::train(format!("读取 safetensors {} 失败: {e}", path.display()))
})?;
let mut stats = DinoLoadStats {
sources_total: sources.len(),
..Default::default()
};
let register_rows = sources
.iter()
.filter(|(n, _)| n == "register_tokens" || n == "dinov2.register_tokens")
.map(|(_, t)| t.size()[1])
.next()
.unwrap_or(0);
stats.register_tokens_stripped = register_rows > 0;
let mut applied: Vec<(String, Tensor)> = Vec::new();
let mut qkv_w: std::collections::HashMap<(usize, u8), Tensor> = Default::default();
let mut qkv_b: std::collections::HashMap<(usize, u8), Tensor> = Default::default();
for (src_name, src) in sources {
match map_source_name(&src_name) {
Mapped::Skipped(reason) => {
stats
.skipped_unrecognized
.push(format!("{src_name}({reason})"));
}
Mapped::Qkv { layer, part, bias } => {
let store = if bias { &mut qkv_b } else { &mut qkv_w };
store.insert((layer, part), src);
}
Mapped::Direct(name) => {
if name == "pos_embed" {
let adapted = self.adapt_pos_embed(&src, register_rows);
let n_src = src.size()[1] - 1 - register_rows;
stats.pos_embed_interpolated =
n_src != self.pos_embed.size()[1] - 1 || register_rows > 0;
applied.push((name, adapted));
} else {
applied.push((name, src));
}
}
}
}
let mut layer_ids: Vec<usize> = qkv_w.keys().map(|(l, _)| *l).collect();
layer_ids.extend(qkv_b.keys().map(|(l, _)| *l));
layer_ids.sort_unstable();
layer_ids.dedup();
for layer in layer_ids {
for (store, suffix) in [(&qkv_w, "weight"), (&qkv_b, "bias")] {
let (Some(q), Some(k), Some(v)) = (
store.get(&(layer, 0)),
store.get(&(layer, 1)),
store.get(&(layer, 2)),
) else {
continue;
};
let fused = Tensor::cat(&[q, k, v], 0);
applied.push((format!("blocks.{layer}.attn.qkv.{suffix}"), fused));
}
}
stats.expected = applied.len();
let mut mismatch: Vec<String> = Vec::new();
tch::no_grad(|| {
for (name, src) in &applied {
if self.paste(name, src) {
stats.loaded += 1;
} else {
mismatch.push(name.clone());
}
}
});
stats.skipped_shape_mismatch = mismatch;
Ok(stats)
}
fn adapt_pos_embed(&self, src: &Tensor, register_rows: i64) -> Tensor {
let class_pos = src.i((.., ..1));
let grid = src.i((.., 1 + register_rows..));
let target_n = self.pos_embed.size()[1] - 1;
let n_src = grid.size()[1];
if n_src == target_n && register_rows == 0 {
return grid.copy();
}
let dim = EMBED_DIM;
let s = (n_src as f64).sqrt().round() as i64;
let g = self.grid_side;
let grid = grid
.reshape([1, s, s, dim])
.permute([0, 3, 1, 2])
.upsample_bicubic2d([g, g], false, g as f64 / s as f64, g as f64 / s as f64)
.permute([0, 2, 3, 1])
.reshape([1, -1, dim]);
Tensor::cat(&[&class_pos, &grid], 1)
}
fn paste(&mut self, name: &str, src: &Tensor) -> bool {
fn copy_same(dst: &mut Tensor, src: &Tensor) -> bool {
if dst.size() == src.size() {
dst.copy_(src);
true
} else {
false
}
}
fn copy_opt(dst: Option<&mut Tensor>, src: &Tensor) -> bool {
match dst {
Some(t) => copy_same(t, src),
None => false,
}
}
match name {
"patch_embed.proj.weight" => copy_same(&mut self.patch_embed.proj.ws, src),
"patch_embed.proj.bias" => copy_opt(self.patch_embed.proj.bs.as_mut(), src),
"cls_token" => copy_same(&mut self.cls_token, src),
"pos_embed" => copy_same(&mut self.pos_embed, src),
"norm.weight" => copy_opt(self.norm.ws.as_mut(), src),
"norm.bias" => copy_opt(self.norm.bs.as_mut(), src),
other => {
let Some(rest) = other.strip_prefix("blocks.") else {
return false;
};
let Some((idx, sub)) = rest.split_once('.') else {
return false;
};
let Ok(i) = idx.parse::<usize>() else {
return false;
};
let Some(blk) = self.blocks.get_mut(i) else {
return false;
};
match sub {
"norm1.weight" => copy_opt(blk.norm1.ws.as_mut(), src),
"norm1.bias" => copy_opt(blk.norm1.bs.as_mut(), src),
"attn.qkv.weight" => copy_same(&mut blk.attn.qkv.ws, src),
"attn.qkv.bias" => copy_opt(blk.attn.qkv.bs.as_mut(), src),
"attn.proj.weight" => copy_same(&mut blk.attn.proj.ws, src),
"attn.proj.bias" => copy_opt(blk.attn.proj.bs.as_mut(), src),
"norm2.weight" => copy_opt(blk.norm2.ws.as_mut(), src),
"norm2.bias" => copy_opt(blk.norm2.bs.as_mut(), src),
"mlp.fc1.weight" => copy_same(&mut blk.mlp.fc1.ws, src),
"mlp.fc1.bias" => copy_opt(blk.mlp.fc1.bs.as_mut(), src),
"mlp.fc2.weight" => copy_same(&mut blk.mlp.fc2.ws, src),
"mlp.fc2.bias" => copy_opt(blk.mlp.fc2.bs.as_mut(), src),
"ls1.gamma" => copy_same(&mut blk.ls1, src),
"ls2.gamma" => copy_same(&mut blk.ls2, src),
_ => false,
}
}
}
}
}
#[derive(Debug, PartialEq)]
enum Mapped {
Direct(String),
Qkv { layer: usize, part: u8, bias: bool },
Skipped(&'static str),
}
fn map_source_name(src: &str) -> Mapped {
let n = src.strip_prefix("dinov2.").unwrap_or(src);
if n == "cls_token" || n == "embeddings.cls_token" {
return Mapped::Direct("cls_token".into());
}
if matches!(
n,
"pos_embed" | "position_embedding" | "embeddings.position_embeddings"
) {
return Mapped::Direct("pos_embed".into());
}
if n == "mask_token" || n == "embeddings.mask_token" {
return Mapped::Skipped("mask_token(本实现无 masked forward)");
}
if n == "register_tokens" {
return Mapped::Skipped("register_tokens(已识别并剥离 pos_embed 对应行)");
}
if n.starts_with("head.") {
return Mapped::Skipped("任务头(分类头按任务重训)");
}
if let Some(rest) = n.strip_prefix("embeddings.patch_embeddings.projection.") {
return Mapped::Direct(format!("patch_embed.proj.{rest}"));
}
if n.starts_with("patch_embed.") {
return Mapped::Direct(n.into());
}
if n == "norm.weight" || n == "norm.bias" {
return Mapped::Direct(n.into());
}
if n == "layernorm.weight" || n == "layernorm.bias" {
return Mapped::Direct(format!("norm.{}", &n["layernorm.".len()..]));
}
if let Some(rest) = n.strip_prefix("encoder.layer.") {
let Some((idx, sub)) = rest.split_once('.') else {
return Mapped::Skipped("无法解析的 encoder.layer 层名");
};
let Ok(layer) = idx.parse::<usize>() else {
return Mapped::Skipped("无法解析的层号");
};
let b = format!("blocks.{layer}.");
for (part, name) in [(0u8, "query"), (1u8, "key"), (2u8, "value")] {
if let Some(t) = sub.strip_prefix(&format!("attention.attention.{name}.")) {
match t {
"weight" => {
return Mapped::Qkv {
layer,
part,
bias: false,
}
}
"bias" => {
return Mapped::Qkv {
layer,
part,
bias: true,
}
}
_ => return Mapped::Skipped("无法识别的 QKV 子张量"),
}
}
}
let mapped = if let Some(t) = sub.strip_prefix("attention.attention.qkv.") {
format!("{b}attn.qkv.{t}")
} else if let Some(t) = sub.strip_prefix("attention.attention.proj.") {
format!("{b}attn.proj.{t}")
} else if let Some(t) = sub.strip_prefix("attention.output.dense.") {
format!("{b}attn.proj.{t}")
} else if let Some(t) = sub.strip_prefix("layernorm_before.") {
format!("{b}norm1.{t}")
} else if let Some(t) = sub.strip_prefix("norm1.") {
format!("{b}norm1.{t}")
} else if let Some(t) = sub.strip_prefix("layernorm_after.") {
format!("{b}norm2.{t}")
} else if let Some(t) = sub.strip_prefix("norm2.") {
format!("{b}norm2.{t}")
} else if let Some(t) = sub.strip_prefix("mlp.fc1.") {
format!("{b}mlp.fc1.{t}")
} else if let Some(t) = sub.strip_prefix("mlp.fc2.") {
format!("{b}mlp.fc2.{t}")
} else if sub == "layer_scale1.lambda1" || sub == "layer_scale1.lambda2" {
format!("{b}ls1.gamma")
} else if sub == "layer_scale2.lambda1" || sub == "layer_scale2.lambda2" {
format!("{b}ls2.gamma")
} else {
return Mapped::Skipped("无法识别的 encoder.layer 子层");
};
return Mapped::Direct(mapped);
}
if n.starts_with("blocks.") {
return Mapped::Direct(n.into());
}
Mapped::Skipped("无法识别的名字")
}
#[derive(Debug, Default, Clone)]
pub struct DinoLoadStats {
pub sources_total: usize,
pub loaded: usize,
pub expected: usize,
pub pos_embed_interpolated: bool,
pub register_tokens_stripped: bool,
pub skipped_shape_mismatch: Vec<String>,
pub skipped_unrecognized: Vec<String>,
}
impl DinoLoadStats {
pub fn load_ratio(&self) -> f32 {
if self.expected == 0 {
0.0
} else {
self.loaded as f32 / self.expected as f32
}
}
pub fn summary(&self) -> String {
format!(
"loaded={}/{}(ratio {:.3},源 {} 个张量)pos_embed 插值={} register 剥离={} \
形状不符={} 非骨干跳过={}",
self.loaded,
self.expected,
self.load_ratio(),
self.sources_total,
self.pos_embed_interpolated,
self.register_tokens_stripped,
self.skipped_shape_mismatch.len(),
self.skipped_unrecognized.len(),
)
}
}
impl BaseBackbone for DinoV2Backbone {
fn forward_features(&self, x: &Tensor) -> AvResult<FeaturePyramid> {
let (_b, _c, h, w) = x.size4().unwrap();
let tokens = self.encode_tokens(x);
let grid = self.patch_grid(&tokens); let (h8, w8) = (h / 8, w / 8);
let (h16, w16) = (h / 16, w / 16);
let (h32, w32) = (h / 32, w / 32);
let s8 = grid
.upsample_bicubic2d([h8, w8], false, None, None)
.apply(&self.conv_s8);
let s16 = grid
.upsample_bicubic2d([h16, w16], false, None, None)
.apply(&self.conv_s16);
let s32 = grid.adaptive_avg_pool2d([h32, w32]).apply(&self.conv_s32);
let mut pyramid = FeaturePyramid::default();
pyramid.levels.push(FeatureMap::new(s8, 8)?);
pyramid.levels.push(FeatureMap::new(s16, 16)?);
pyramid.levels.push(FeatureMap::new(s32, 32)?);
pyramid.validate_ascending()?;
Ok(pyramid)
}
fn forward_pooled(&self, x: &Tensor) -> AvResult<Tensor> {
let tokens = self.encode_tokens(x);
Ok(tokens.i((.., 0)))
}
fn spec(&self) -> BackboneSpec {
BackboneSpec {
levels: vec![
LevelSpec {
stride: 8,
channels: PYRAMID_CHANNELS as usize,
},
LevelSpec {
stride: 16,
channels: PYRAMID_CHANNELS as usize,
},
LevelSpec {
stride: 32,
channels: PYRAMID_CHANNELS as usize,
},
],
}
}
}
#[cfg(all(test, feature = "torch"))]
mod tests {
use super::*;
fn backbone_at(img: u32) -> DinoV2Backbone {
let vs = nn::VarStore::new(tch::Device::Cpu);
let cfg = BackboneCfg::default();
DinoV2Backbone::new(&vs.root(), &cfg, img).unwrap()
}
fn no_nan(t: &Tensor) -> bool {
t.isnan().sum(Kind::Float).double_value(&[]) == 0.0
}
#[test]
fn constructor_rejects_non_multiple_sizes() {
let vs = nn::VarStore::new(tch::Device::Cpu);
let cfg = BackboneCfg::default();
assert!(DinoV2Backbone::new(&vs.root(), &cfg, 224 + 14 + 3).is_err());
assert!(DinoV2Backbone::new(&vs.root(), &cfg, 154).is_err());
}
#[test]
fn pyramid_and_pooled_contract_224() {
let backbone = backbone_at(224);
let x = Tensor::randn([2, 3, 224, 224], (tch::Kind::Float, tch::Device::Cpu));
let pyramid = backbone.forward_features(&x).unwrap();
let strides: Vec<u32> = pyramid.levels.iter().map(|l| l.stride).collect();
assert_eq!(strides, vec![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]);
let pooled = backbone.forward_pooled(&x).unwrap();
assert_eq!(pooled.size(), vec![2, 384]);
assert!(no_nan(&pooled));
assert_eq!(backbone.spec().levels.len(), 3);
}
#[test]
fn pos_embed_interpolation_changes_grid_and_keeps_cls() {
let vs = nn::VarStore::new(tch::Device::Cpu);
let cfg = BackboneCfg::default();
let mut backbone = DinoV2Backbone::new(&vs.root(), &cfg, 224).unwrap(); let rand = Tensor::randn(
backbone.pos_embed.size(),
(tch::Kind::Float, tch::Device::Cpu),
);
tch::no_grad(|| backbone.pos_embed.copy_(&rand));
let npatch_224 = 16 * 16;
let same = backbone.interpolate_pos_encoding(npatch_224, 224, 224);
assert_eq!(
tensor_max_diff(&same, &backbone.pos_embed),
0.0,
"同网格应零拷贝逐位一致"
);
let npatch_448 = 32 * 32;
let interp = backbone.interpolate_pos_encoding(npatch_448, 448, 448);
assert_eq!(interp.size(), vec![1, 1 + npatch_448, EMBED_DIM]);
let cls_diff = tensor_max_diff(&backbone.pos_embed.i((.., ..1)), &interp.i((.., ..1)));
assert!(cls_diff < 1e-6, "cls 行不应被插值改动,diff={cls_diff}");
let grid_after = interp.i((.., 1..));
assert!(no_nan(&grid_after));
assert!(
grid_after.abs().max().double_value(&[]) > 1e-3,
"随机源的插值网格不应全零"
);
}
fn tensor_max_diff(a: &Tensor, b: &Tensor) -> f64 {
(a - b).abs().max().double_value(&[])
}
#[test]
fn source_name_maps_hf_and_hub() {
match map_source_name("encoder.layer.3.attention.attention.query.weight") {
Mapped::Qkv { layer, part, bias } => {
assert_eq!((layer, part, bias), (3, 0, false));
}
other => panic!("应为 Qkv 分量,得到 {other:?}"),
}
match map_source_name("encoder.layer.3.attention.attention.value.bias") {
Mapped::Qkv { layer, part, bias } => {
assert_eq!((layer, part, bias), (3, 2, true));
}
other => panic!("应为 Qkv 分量,得到 {other:?}"),
}
match map_source_name("dinov2.encoder.layer.0.attention.attention.key.bias") {
Mapped::Qkv { layer, part, bias } => {
assert_eq!((layer, part, bias), (0, 1, true));
}
other => panic!("应为 Qkv 分量,得到 {other:?}"),
}
assert_eq!(
map_source_name("encoder.layer.0.attention.output.dense.bias"),
Mapped::Direct("blocks.0.attn.proj.bias".into())
);
assert_eq!(
map_source_name("encoder.layer.11.norm1.weight"),
Mapped::Direct("blocks.11.norm1.weight".into())
);
assert_eq!(
map_source_name("encoder.layer.11.layer_scale2.lambda1"),
Mapped::Direct("blocks.11.ls2.gamma".into())
);
assert_eq!(
map_source_name("embeddings.patch_embeddings.projection.weight"),
Mapped::Direct("patch_embed.proj.weight".into())
);
assert_eq!(
map_source_name("embeddings.position_embeddings"),
Mapped::Direct("pos_embed".into())
);
assert_eq!(
map_source_name("embeddings.cls_token"),
Mapped::Direct("cls_token".into())
);
assert_eq!(
map_source_name("layernorm.bias"),
Mapped::Direct("norm.bias".into())
);
assert_eq!(
map_source_name("dinov2.encoder.layer.0.attention.attention.qkv.weight"),
Mapped::Direct("blocks.0.attn.qkv.weight".into())
);
assert_eq!(
map_source_name("dinov2.encoder.layer.0.layernorm_before.weight"),
Mapped::Direct("blocks.0.norm1.weight".into())
);
assert_eq!(
map_source_name("dinov2.position_embedding"),
Mapped::Direct("pos_embed".into())
);
assert_eq!(
map_source_name("blocks.5.attn.qkv.weight"),
Mapped::Direct("blocks.5.attn.qkv.weight".into())
);
assert_eq!(
map_source_name("blocks.5.ls1.gamma"),
Mapped::Direct("blocks.5.ls1.gamma".into())
);
assert_eq!(
map_source_name("patch_embed.proj.bias"),
Mapped::Direct("patch_embed.proj.bias".into())
);
assert!(matches!(
map_source_name("embeddings.mask_token"),
Mapped::Skipped(_)
));
assert!(matches!(
map_source_name("register_tokens"),
Mapped::Skipped(_)
));
assert!(matches!(
map_source_name("head.fc.weight"),
Mapped::Skipped(_)
));
}
#[test]
fn load_fuses_split_qkv_and_reports_stats() {
let dim = EMBED_DIM;
let vs = nn::VarStore::new(tch::Device::Cpu);
let cfg = BackboneCfg::default();
let mut backbone = DinoV2Backbone::new(&vs.root(), &cfg, 224).unwrap();
let mk =
|v: f64, shape: &[i64]| Tensor::ones(shape, (tch::Kind::Float, tch::Device::Cpu)) * v;
let dir =
std::env::temp_dir().join(format!("av-dino-test-{}.safetensors", std::process::id()));
let entries: Vec<(&str, Tensor)> = vec![
(
"encoder.layer.0.attention.attention.query.weight",
mk(1.0, &[dim, dim]),
),
(
"encoder.layer.0.attention.attention.key.weight",
mk(2.0, &[dim, dim]),
),
(
"encoder.layer.0.attention.attention.value.weight",
mk(3.0, &[dim, dim]),
),
(
"encoder.layer.0.attention.attention.query.bias",
mk(4.0, &[dim]),
),
(
"encoder.layer.0.attention.attention.key.bias",
mk(5.0, &[dim]),
),
(
"encoder.layer.0.attention.attention.value.bias",
mk(6.0, &[dim]),
),
("encoder.layer.1.mlp.fc1.weight", mk(7.0, &[dim, dim])),
];
Tensor::write_safetensors(&entries, &dir).unwrap();
let stats = backbone.load_dinov2_weights(&dir).unwrap();
let fused = backbone.blocks[0].attn.qkv.ws.copy();
for (seg, v) in [(0i64, 1.0), (1, 2.0), (2, 3.0)] {
let part = fused.narrow(0, seg * dim, dim);
assert!(
(part.max().double_value(&[]) - v).abs() < 1e-6
&& (part.min().double_value(&[]) - v).abs() < 1e-6,
"融合段 {seg} 应全为 {v}"
);
}
let fused_b = backbone.blocks[0].attn.qkv.bs.as_ref().unwrap().copy();
assert!((fused_b.double_value(&[0]) - 4.0).abs() < 1e-6);
assert!((fused_b.double_value(&[dim]) - 5.0).abs() < 1e-6);
assert_eq!(stats.loaded, 2); assert_eq!(stats.expected, 3);
assert_eq!(
stats.skipped_shape_mismatch,
vec!["blocks.1.mlp.fc1.weight"]
);
assert_eq!(stats.sources_total, entries.len() as usize);
assert!(stats.load_ratio() < 0.9);
let _ = std::fs::remove_file(&dir);
}
}