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::{sigmoid, softplus};
const NR_YITA1: f64 = 2.0;
const NR_YITA2: f64 = -2.0;
const NR_YITA3_INIT: f32 = 4.9592;
const NR_YITA4_INIT: f32 = 21.5968;
const FR_YITA1: f64 = 2.0;
const FR_YITA2: f64 = -2.0;
const FR_YITA3_INIT: f32 = 0.5;
const FR_YITA4_INIT: f32 = 0.15;
const FR_YITA3_MIN: f32 = 0.05;
const FR_YITA3_MAX: f32 = 0.95;
const FR_YITA4_MIN: f32 = 0.01;
const FR_YITA4_MAX: f32 = 0.70;
const ADAPTER_K_INIT: f32 = 5.0;
const SCALE_YITA1: f64 = 100.0;
const SCALE_YITA2: f64 = 0.0;
const SCALE_YITA3: f64 = -1.971_0;
const SCALE_YITA4: f64 = -2.373_4;
const EPS: f64 = 1e-10;
fn logistic_calibrate(
x: Tensor<2>,
yita3: Tensor<1>,
yita4_abs: Tensor<1>,
yita1: f64,
yita2: f64,
) -> Tensor<2> {
let yita3 = yita3.reshape([1, 1]);
let denom = yita4_abs.reshape([1, 1]).add_scalar(EPS);
let inner = (x - yita3) / denom;
sigmoid(inner).mul_scalar(yita1 - yita2).add_scalar(yita2)
}
#[derive(Config, Debug)]
pub(crate) struct NrCalibratorConfig {}
impl NrCalibratorConfig {
pub(crate) fn init(&self, device: &Device) -> NrCalibrator {
NrCalibrator {
yita3: Param::from_tensor(Tensor::from_floats([NR_YITA3_INIT], device)),
yita4: Param::from_tensor(Tensor::from_floats([NR_YITA4_INIT], device)),
}
}
}
#[derive(Module, Debug)]
pub(crate) struct NrCalibrator {
pub(crate) yita3: Param<Tensor<1>>,
pub(crate) yita4: Param<Tensor<1>>,
}
impl NrCalibrator {
pub(crate) fn forward(&self, x: Tensor<2>) -> Tensor<2> {
logistic_calibrate(
x,
self.yita3.val(),
self.yita4.val().abs(),
NR_YITA1,
NR_YITA2,
)
}
}
#[derive(Config, Debug)]
pub(crate) struct FrCalibratorWithLimitConfig {}
impl FrCalibratorWithLimitConfig {
pub(crate) fn init(&self, device: &Device) -> FrCalibratorWithLimit {
FrCalibratorWithLimit {
yita3: Param::from_tensor(Tensor::from_floats([FR_YITA3_INIT], device)),
yita4: Param::from_tensor(Tensor::from_floats([FR_YITA4_INIT], device)),
}
}
}
#[derive(Module, Debug)]
pub(crate) struct FrCalibratorWithLimit {
pub(crate) yita3: Param<Tensor<1>>,
pub(crate) yita4: Param<Tensor<1>>,
}
impl FrCalibratorWithLimit {
pub(crate) fn forward(&self, x: Tensor<2>) -> Tensor<2> {
let yita3 = self.yita3.val().clamp(FR_YITA3_MIN, FR_YITA3_MAX);
let yita4 = self.yita4.val().clamp(FR_YITA4_MIN, FR_YITA4_MAX);
logistic_calibrate(x, yita3, yita4.abs(), FR_YITA1, FR_YITA2)
}
}
#[derive(Config, Debug)]
pub(crate) struct AfineAdapterConfig {}
impl AfineAdapterConfig {
pub(crate) fn init(&self, device: &Device) -> AfineAdapter {
AfineAdapter {
k: Param::from_tensor(Tensor::from_floats([ADAPTER_K_INIT], device)),
}
}
}
#[derive(Module, Debug)]
pub(crate) struct AfineAdapter {
pub(crate) k: Param<Tensor<1>>,
}
impl AfineAdapter {
pub(crate) fn forward(
&self,
x_nr: Tensor<2>,
ref_nr: Tensor<2>,
xref_fr: Tensor<2>,
) -> Tensor<2> {
let k_pos = softplus(self.k.val(), 1.0).reshape([1, 1]);
let weight = (k_pos * (ref_nr - x_nr.clone())).exp();
weight * x_nr + xref_fr
}
}
pub(crate) fn scale_finalscore(score: Tensor<2>) -> Tensor<2> {
let denom = SCALE_YITA4.abs() + EPS;
let inner = score.sub_scalar(SCALE_YITA3).div_scalar(denom);
sigmoid(inner)
.mul_scalar(SCALE_YITA1 - SCALE_YITA2)
.add_scalar(SCALE_YITA2)
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn nr_calibrator_maps_to_bounded_range() {
let device = Default::default();
let calibrator = NrCalibratorConfig::new().init(&device);
let extremes = Tensor::<2>::from_floats([[-1000.0], [0.0], [1000.0]], &device);
let out = calibrator.forward(extremes);
let values = out.into_data().to_vec::<f32>().unwrap();
for v in &values {
assert!(*v >= -2.0 && *v <= 2.0, "out-of-range value: {v}");
}
assert!(values[0] < values[1]);
assert!(values[1] < values[2]);
}
#[test]
fn fr_calibrator_clamp_does_not_panic() {
let device = Default::default();
let calibrator = FrCalibratorWithLimitConfig::new().init(&device);
let input = Tensor::<2>::from_floats([[0.5], [1.5], [-0.5]], &device);
let out = calibrator.forward(input);
assert_eq!(out.dims(), [3, 1]);
}
#[test]
fn adapter_forward_propagates_shape() {
let device = Default::default();
let adapter = AfineAdapterConfig::new().init(&device);
let nr_dis = Tensor::<2>::from_floats([[0.5], [-0.3]], &device);
let nr_ref = Tensor::<2>::from_floats([[0.7], [-0.1]], &device);
let fr = Tensor::<2>::from_floats([[0.2], [0.4]], &device);
let out = adapter.forward(nr_dis, nr_ref, fr);
assert_eq!(out.dims(), [2, 1]);
}
#[test]
fn scale_finalscore_maps_to_0_100_range() {
let device = Default::default();
let scores = Tensor::<2>::from_floats([[-1000.0], [-1.971], [1000.0]], &device);
let out = scale_finalscore(scores);
let values = out.into_data().to_vec::<f32>().unwrap();
assert!(values[0] >= 0.0 && values[0] <= 100.0);
assert!(values[2] >= 0.0 && values[2] <= 100.0);
assert!(
(values[1] - 50.0).abs() < 0.5,
"midpoint should be ~50, got {}",
values[1]
);
}
}