use burn_core as burn;
use burn::config::Config;
use burn::module::{Module, Param};
use burn::tensor::Device;
use burn::tensor::Tensor;
use burn::tensor::activation::{relu, softplus};
use burn_nn::{Gelu, Linear, LinearConfig};
const CLIP_MEAN: [f32; 3] = [0.481_454_66, 0.457_827_5, 0.408_210_73];
const CLIP_STD: [f32; 3] = [0.268_629_54, 0.261_302_58, 0.275_777_11];
const CHNS: [usize; 13] = [
3, 768, 768, 768, 768, 768, 768, 768, 768, 768, 768, 768, 768,
];
const NUM_LEVELS: usize = 13;
const Q_HIDDEN_DIM: usize = 128;
const Q_PROJ_HEAD_OUT: usize = 768;
const Q_PROJ_HEAD_IN: usize = 6 + Q_HIDDEN_DIM * 12;
const D_CHNS_SUM: usize = 9219;
const EPS: f64 = 1e-10;
fn build_clip_mean_std(device: &Device) -> (Tensor<4>, Tensor<4>) {
let mean = Tensor::from_floats(
[[[[CLIP_MEAN[0]]], [[CLIP_MEAN[1]]], [[CLIP_MEAN[2]]]]],
device,
);
let std = Tensor::from_floats(
[[[[CLIP_STD[0]]], [[CLIP_STD[1]]], [[CLIP_STD[2]]]]],
device,
);
(mean, std)
}
fn level_mean_var(feat: Tensor<3>) -> (Tensor<3>, Tensor<3>, Tensor<2>) {
let mean = feat.clone().mean_dim(1);
let centered = feat - mean.clone();
let var = centered.powi_scalar(2).mean_dim(1);
let descriptor = Tensor::cat(
vec![
mean.clone().flatten::<2>(1, 2),
var.clone().flatten::<2>(1, 2),
],
1,
);
(mean, var, descriptor)
}
#[derive(Config, Debug)]
pub struct AfineQHeadConfig {}
impl AfineQHeadConfig {
pub fn init(&self, device: &Device) -> AfineQHead {
let (mean, std) = build_clip_mean_std(device);
AfineQHead {
mean,
std,
proj_feat: LinearConfig::new(2 * 768, Q_HIDDEN_DIM)
.with_bias(true)
.init(device),
proj_head_fc1: LinearConfig::new(Q_PROJ_HEAD_IN, Q_PROJ_HEAD_OUT)
.with_bias(true)
.init(device),
proj_head_fc2: LinearConfig::new(Q_PROJ_HEAD_OUT, 1)
.with_bias(true)
.init(device),
activation: Gelu::new(),
}
}
}
#[derive(Module, Debug)]
pub struct AfineQHead {
pub(crate) mean: Tensor<4>,
pub(crate) std: Tensor<4>,
pub(crate) proj_feat: Linear,
pub(crate) proj_head_fc1: Linear,
pub(crate) proj_head_fc2: Linear,
pub(crate) activation: Gelu,
}
impl AfineQHead {
pub fn forward(&self, image: Tensor<4>, clip_features: &[Tensor<3>]) -> Tensor<2> {
assert_eq!(
clip_features.len(),
12,
"AfineQHead expects 12 CLIP feature maps, got {}",
clip_features.len()
);
let [batch, channels, height, width] = image.dims();
let img = image * self.std.clone() + self.mean.clone();
let img_feat = img
.reshape([batch, channels, height * width])
.swap_dims(1, 2);
let mut level_descriptors: Vec<Tensor<2>> = Vec::with_capacity(NUM_LEVELS);
let (_, _, raw_descriptor) = level_mean_var(img_feat);
level_descriptors.push(raw_descriptor);
for h in clip_features {
let activated = relu(h.clone());
let (_, _, descriptor) = level_mean_var(activated);
level_descriptors.push(self.proj_feat.forward(descriptor));
}
let concat_all = Tensor::cat(level_descriptors, 1);
let hidden = self
.activation
.forward(self.proj_head_fc1.forward(concat_all));
self.proj_head_fc2.forward(hidden)
}
}
#[derive(Config, Debug)]
pub struct AfineDHeadConfig {}
impl AfineDHeadConfig {
pub fn init(&self, device: &Device) -> AfineDHead {
let (mean, std) = build_clip_mean_std(device);
let alpha = Tensor::random(
[1, 1, D_CHNS_SUM],
burn::tensor::Distribution::Normal(0.1, 0.01),
device,
);
let beta = Tensor::random(
[1, 1, D_CHNS_SUM],
burn::tensor::Distribution::Normal(0.1, 0.01),
device,
);
AfineDHead {
mean,
std,
alpha: Param::from_tensor(alpha),
beta: Param::from_tensor(beta),
}
}
}
#[derive(Module, Debug)]
pub struct AfineDHead {
pub(crate) mean: Tensor<4>,
pub(crate) std: Tensor<4>,
pub(crate) alpha: Param<Tensor<3>>,
pub(crate) beta: Param<Tensor<3>>,
}
impl AfineDHead {
pub fn forward(
&self,
distorted: Tensor<4>,
reference: Tensor<4>,
feat_dis: &[Tensor<3>],
feat_ref: &[Tensor<3>],
) -> Tensor<2> {
assert_eq!(feat_dis.len(), 12);
assert_eq!(feat_ref.len(), 12);
let [batch, channels, height, width] = distorted.dims();
let raw_x = (distorted * self.std.clone() + self.mean.clone())
.reshape([batch, channels, height * width])
.swap_dims(1, 2);
let raw_y = (reference * self.std.clone() + self.mean.clone())
.reshape([batch, channels, height * width])
.swap_dims(1, 2);
let mut feat_x: Vec<Tensor<3>> = Vec::with_capacity(NUM_LEVELS);
let mut feat_y: Vec<Tensor<3>> = Vec::with_capacity(NUM_LEVELS);
feat_x.push(raw_x);
feat_y.push(raw_y);
for h in feat_dis {
feat_x.push(relu(h.clone()));
}
for h in feat_ref {
feat_y.push(relu(h.clone()));
}
let alpha_sp = softplus(self.alpha.val(), 1.0);
let beta_sp = softplus(self.beta.val(), 1.0);
let w_sum = (alpha_sp.clone().sum() + beta_sp.clone().sum())
.add_scalar(EPS)
.reshape([1, 1, 1]);
let alpha_norm = alpha_sp / w_sum.clone();
let beta_norm = beta_sp / w_sum;
let mut dist1: Option<Tensor<3>> = None;
let mut dist2: Option<Tensor<3>> = None;
let mut offset: usize = 0;
for k in 0..NUM_LEVELS {
let cn = CHNS[k];
let alpha_k = alpha_norm.clone().slice([0..1, 0..1, offset..offset + cn]);
let beta_k = beta_norm.clone().slice([0..1, 0..1, offset..offset + cn]);
let xm = feat_x[k].clone().mean_dim(1); let ym = feat_y[k].clone().mean_dim(1);
let s1_num = (xm.clone() * ym.clone()).mul_scalar(2.0).add_scalar(EPS);
let s1_den = xm.clone().powi_scalar(2) + ym.clone().powi_scalar(2);
let s1 = s1_num / s1_den.add_scalar(EPS);
let term1 = (alpha_k * s1).sum_dim(2); dist1 = Some(match dist1 {
None => term1,
Some(d) => d + term1,
});
let xc = feat_x[k].clone() - xm.clone();
let yc = feat_y[k].clone() - ym.clone();
let x_var = xc.powi_scalar(2).mean_dim(1);
let y_var = yc.powi_scalar(2).mean_dim(1);
let xy_mean = (feat_x[k].clone() * feat_y[k].clone()).mean_dim(1);
let xy_cov = xy_mean - xm * ym;
let s2_num = xy_cov.mul_scalar(2.0).add_scalar(EPS);
let s2_den = (x_var + y_var).add_scalar(EPS);
let s2 = s2_num / s2_den;
let term2 = (beta_k * s2).sum_dim(2);
dist2 = Some(match dist2 {
None => term2,
Some(d) => d + term2,
});
offset += cn;
}
let total = dist1.unwrap() + dist2.unwrap(); let score = total.ones_like() - total; score.squeeze_dim::<2>(2) }
}
#[cfg(test)]
mod tests {
use super::*;
use burn::tensor::Distribution;
#[test]
fn afine_q_head_forward_shape() {
let device = Default::default();
let head = AfineQHeadConfig::new().init(&device);
let image = Tensor::<4>::random([2, 3, 64, 64], Distribution::Default, &device);
let features: Vec<_> = (0..12)
.map(|_| Tensor::<3>::random([2, 4, 768], Distribution::Default, &device))
.collect();
let out = head.forward(image, &features);
assert_eq!(out.dims(), [2, 1]);
}
#[test]
fn afine_d_head_forward_shape() {
let device = Default::default();
let head = AfineDHeadConfig::new().init(&device);
let dis = Tensor::<4>::random([2, 3, 64, 64], Distribution::Default, &device);
let reference = Tensor::<4>::random([2, 3, 64, 64], Distribution::Default, &device);
let feat_dis: Vec<_> = (0..12)
.map(|_| Tensor::<3>::random([2, 4, 768], Distribution::Default, &device))
.collect();
let feat_ref: Vec<_> = (0..12)
.map(|_| Tensor::<3>::random([2, 4, 768], Distribution::Default, &device))
.collect();
let out = head.forward(dis, reference, &feat_dis, &feat_ref);
assert_eq!(out.dims(), [2, 1]);
}
#[test]
fn afine_d_head_identical_inputs_unit_score() {
let device = Default::default();
let head = AfineDHeadConfig::new().init(&device);
let image = Tensor::<4>::random([1, 3, 32, 32], Distribution::Default, &device);
let features: Vec<_> = (0..12)
.map(|_| Tensor::<3>::random([1, 1, 768], Distribution::Default, &device))
.collect();
let out = head.forward(image.clone(), image, &features.clone(), &features);
let value = out.into_data().to_vec::<f32>().unwrap()[0];
assert!(
value.abs() < 0.1,
"fidelity head on identical inputs should yield ~0, got {value}"
);
}
}