use crate::chiral::Center;
use crate::optimize::Objective;
use crate::smooth::Bounds;
pub const VOL_LO: f64 = 5.0;
pub const VOL_HI: f64 = 100.0;
pub const UMBRELLA_LO: f64 = 0.3;
pub const UMBRELLA_HI: f64 = 30.0;
pub const WEIGHT_UMBRELLA: f64 = 1.0;
pub const WEIGHT_CHIRAL: f64 = 1.0;
#[derive(Debug, Clone)]
pub struct Field {
n: usize,
lower: Vec<f64>,
upper: Vec<f64>,
centers: Vec<Center>,
pub weight_chiral: f64,
}
impl Field {
#[must_use]
pub fn new(b: &Bounds, centers: &[Center]) -> Self {
let n = b.len();
for c in centers {
assert!(
(c.atom as usize) < n,
"手性中心的中心原子下标 {} 越界(原子数 {n})",
c.atom
);
for &l in c.real_ligands() {
assert!(
(l as usize) < n,
"手性中心 {} 的配体下标 {l} 越界(原子数 {n})",
c.atom
);
}
}
let mut lower = vec![0.0; n * n];
let mut upper = vec![0.0; n * n];
for i in 0..n {
for j in (i + 1)..n {
lower[i * n + j] = b.lower(i, j);
upper[i * n + j] = b.upper(i, j);
}
}
Self {
n,
lower,
upper,
centers: centers.to_vec(),
weight_chiral: WEIGHT_CHIRAL,
}
}
#[must_use]
pub fn len(&self) -> usize {
self.n
}
#[must_use]
pub fn is_empty(&self) -> bool {
self.n == 0
}
fn distance_term(&self, x: &[f64], g: &mut [f64]) -> f64 {
let n = self.n;
let mut e = 0.0;
for i in 0..n {
for j in (i + 1)..n {
let dx = [
x[3 * i] - x[3 * j],
x[3 * i + 1] - x[3 * j + 1],
x[3 * i + 2] - x[3 * j + 2],
];
let d2 = dx[0] * dx[0] + dx[1] * dx[1] + dx[2] * dx[2];
let (lo, hi) = (self.lower[i * n + j], self.upper[i * n + j]);
let (lo2, hi2) = (lo * lo, hi * hi);
let dedd_over_d = if d2 > hi2 && hi2 > 0.0 {
let val = d2 / hi2 - 1.0;
e += val * val;
4.0 * val / hi2
} else if d2 < lo2 {
let s = lo2 + d2;
let val = 2.0 * lo2 / s - 1.0;
e += val * val;
-8.0 * val * lo2 / (s * s)
} else {
continue;
};
for t in 0..3 {
let f = dedd_over_d * dx[t];
g[3 * i + t] += f;
g[3 * j + t] -= f;
}
}
}
e
}
fn chiral_term(&self, x: &[f64], g: &mut [f64]) -> f64 {
let mut e = 0.0;
for c in &self.centers {
let p: Vec<[f64; 3]> = c
.real_ligands()
.iter()
.map(|&a| {
let a = a as usize;
[x[3 * a], x[3 * a + 1], x[3 * a + 2]]
})
.collect();
let sub = |u: [f64; 3], v: [f64; 3]| [u[0] - v[0], u[1] - v[1], u[2] - v[2]];
let cross = |u: [f64; 3], v: [f64; 3]| {
[
u[1] * v[2] - u[2] * v[1],
u[2] * v[0] - u[0] * v[2],
u[0] * v[1] - u[1] * v[0],
]
};
let ctr = c.atom as usize;
let o = [x[3 * ctr], x[3 * ctr + 1], x[3 * ctr + 2]];
let term = |a: [f64; 3],
b: [f64; 3],
cc: [f64; 3],
idx: [usize; 3],
base: usize,
want: f64,
lo_abs: f64,
hi_abs: f64,
w: f64,
g: &mut [f64]| {
let bxc = cross(b, cc);
let v = a[0] * bxc[0] + a[1] * bxc[1] + a[2] * bxc[2];
let (lo, hi) = if want < 0.0 {
(-hi_abs, -lo_abs)
} else {
(lo_abs, hi_abs)
};
let dev = if v < lo {
v - lo
} else if v > hi {
v - hi
} else {
return 0.0;
};
let k = 2.0 * w * dev;
let cxa = cross(cc, a);
let axb = cross(a, b);
for t in 0..3 {
g[3 * idx[0] + t] += k * bxc[t];
g[3 * idx[1] + t] += k * cxa[t];
g[3 * idx[2] + t] += k * axb[t];
g[3 * base + t] -= k * (bxc[t] + cxa[t] + axb[t]);
}
w * dev * dev
};
if !c.is_three_coordinate() {
let (a, b, cc) = (sub(p[1], p[0]), sub(p[2], p[0]), sub(p[3], p[0]));
let idx = [
c.ligands[1] as usize,
c.ligands[2] as usize,
c.ligands[3] as usize,
];
e += term(
a,
b,
cc,
idx,
c.ligands[0] as usize,
-c.sign,
VOL_LO,
VOL_HI,
self.weight_chiral,
g,
);
}
let (a2, b2, c2) = (sub(p[0], o), sub(p[1], o), sub(p[2], o));
let idx2 = [
c.ligands[0] as usize,
c.ligands[1] as usize,
c.ligands[2] as usize,
];
e += term(
a2,
b2,
c2,
idx2,
ctr,
c.sign,
UMBRELLA_LO,
UMBRELLA_HI,
self.weight_chiral * WEIGHT_UMBRELLA,
g,
);
}
e
}
}
impl Objective for Field {
fn value_and_grad(&self, x: &[f64], grad: &mut [f64]) -> f64 {
if !x.iter().all(|v| v.is_finite()) {
for v in grad.iter_mut() {
*v = f64::NAN;
}
return f64::NAN;
}
for v in grad.iter_mut() {
*v = 0.0;
}
self.distance_term(x, grad) + self.chiral_term(x, grad)
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::optimize::{max_grad_error, minimize, Options};
fn field_from(points: &[[f64; 3]], half: f64, centers: &[Center]) -> Field {
let n = points.len();
let mut b = Bounds::new(n, 0.0, 1000.0);
for i in 0..n {
for j in (i + 1)..n {
let d = ((points[i][0] - points[j][0]).powi(2)
+ (points[i][1] - points[j][1]).powi(2)
+ (points[i][2] - points[j][2]).powi(2))
.sqrt();
b.set_lower(i, j, (d - half).max(0.05));
b.set_upper(i, j, d + half);
}
}
Field::new(&b, centers)
}
#[test]
fn 非有限坐标不许被判成满分() {
let pts = [
[0.0, 0.0, 0.0],
[1.5, 0.0, 0.0],
[0.0, 1.5, 0.0],
[0.0, 0.0, 1.5],
];
let f = field_from(&pts, 0.1, &[]);
let clean: Vec<f64> = pts.iter().flatten().copied().collect();
let mut g = vec![0.0; clean.len()];
assert_eq!(
f.value_and_grad(&clean, &mut g),
0.0,
"界就是照这组点定的,它本身该是零点"
);
assert!(g.iter().all(|v| *v == 0.0), "零点上梯度该是零");
assert!(
max_grad_error(&f, &clean, 1e-5) < 1e-6,
"干净坐标上梯度该对"
);
let r = minimize(&f, &mut clean.clone(), &Options::default());
assert!(r.converged && r.value.is_finite(), "干净坐标上该收敛");
for (tag, bad) in [
("NaN", f64::NAN),
("+inf", f64::INFINITY),
("-inf", f64::NEG_INFINITY),
] {
for poison in [0usize, 5, 11] {
let mut x = clean.clone();
x[poison] = bad;
let mut g = vec![0.0; x.len()];
let e = f.value_and_grad(&x, &mut g);
assert!(e.is_nan(), "{tag}@{poison}:误差该是 NaN,实得 {e}");
assert!(
g.iter().all(|v| v.is_nan()),
"{tag}@{poison}:梯度该整条 NaN,实得 {g:?}"
);
assert!(
max_grad_error(&f, &x, 1e-5).is_nan(),
"{tag}@{poison}:梯度校验不许报出一个有限偏差"
);
let mut xx = x.clone();
let r = minimize(&f, &mut xx, &Options::default());
assert!(!r.converged, "{tag}@{poison}:不许报收敛");
assert!(
!r.grad_norm.is_finite(),
"{tag}@{poison}:梯度范数不许是有限数,实得 {}",
r.grad_norm
);
assert!(
!r.value.is_finite(),
"{tag}@{poison}:目标值不许是有限数,实得 {}",
r.value
);
}
}
}
fn lcg(state: &mut u64) -> f64 {
*state = state
.wrapping_mul(6_364_136_223_846_793_005)
.wrapping_add(1);
((*state >> 11) as f64) / ((1u64 << 53) as f64) * 2.0 - 1.0
}
fn random_points(n: usize, seed: u64, scale: f64) -> Vec<[f64; 3]> {
let mut st = seed;
(0..n)
.map(|_| {
[
lcg(&mut st) * scale,
lcg(&mut st) * scale,
lcg(&mut st) * scale,
]
})
.collect()
}
fn flat(p: &[[f64; 3]]) -> Vec<f64> {
p.iter().flat_map(|q| q.iter().copied()).collect()
}
#[test]
fn 距离项的梯度与能量一致() {
let pts = random_points(10, 5, 3.0);
let f = field_from(&pts, 0.3, &[]);
let mut st = 77u64;
let x: Vec<f64> = flat(&pts).iter().map(|v| v + lcg(&mut st) * 2.0).collect();
let e = max_grad_error(&f, &x, 1e-6);
assert!(e < 1e-6, "距离项梯度对不上:{e:.3e}");
}
fn assert_both_chiral_active(pts: &[[f64; 3]], c: &Center) {
let inside = |v: f64, lo: f64, hi: f64, want: f64| {
if want < 0.0 {
v <= -lo && v >= -hi
} else {
v >= lo && v <= hi
}
};
let l = c.ligands.map(|k| pts[k as usize]);
let vl = crate::chiral::signed_volume(l[0], l[1], l[2], l[3]);
let vc = crate::chiral::center_volume(pts, c);
assert!(
!inside(vl, VOL_LO, VOL_HI, -c.sign),
"sign={} 四配体项没在罚(V={vl}),这一半测了个寂寞",
c.sign
);
assert!(
!inside(vc, UMBRELLA_LO, UMBRELLA_HI, c.sign),
"sign={} 伞形项没在罚(V={vc}),这一半测了个寂寞",
c.sign
);
}
#[test]
fn 手性项的梯度与能量一致() {
let pts: Vec<[f64; 3]> = vec![
[0.03, -0.02, 0.02], [1.10, 0.00, 0.00],
[0.00, 0.95, 0.00],
[-1.05, 0.00, 0.00],
[0.00, -1.00, 0.15],
];
for sign in [-1.0, 1.0] {
let c = Center {
atom: 0,
ligands: [1, 2, 3, 4],
sign,
};
let mut b = Bounds::new(5, 0.01, 1000.0);
for i in 0..5 {
for j in (i + 1)..5 {
b.set_lower(i, j, 0.01);
b.set_upper(i, j, 1000.0);
}
}
let f = Field::new(&b, &[c]);
let x = flat(&pts);
let mut g = vec![0.0; 15];
let e0 = f.value_and_grad(&x, &mut g);
assert!(e0 > 0.0, "sign={sign} 这个构型没违反手性,测不到东西");
assert_both_chiral_active(&pts, &c);
let e = max_grad_error(&f, &x, 1e-6);
assert!(e < 1e-6, "sign={sign} 手性项梯度对不上:{e:.3e}");
}
}
#[test]
fn 两项一起的梯度也一致() {
let pts = random_points(8, 21, 2.5);
let c = Center {
atom: 0,
ligands: [1, 2, 3, 4],
sign: -1.0,
};
let f = field_from(&pts, 0.2, &[c]);
let mut st = 33u64;
let x: Vec<f64> = flat(&pts).iter().map(|v| v + lcg(&mut st) * 1.5).collect();
let px: Vec<[f64; 3]> = (0..pts.len())
.map(|i| [x[3 * i], x[3 * i + 1], x[3 * i + 2]])
.collect();
assert_both_chiral_active(&px, &c);
let e = max_grad_error(&f, &x, 1e-6);
assert!(e < 1e-6, "合起来梯度对不上:{e:.3e}");
}
#[test]
fn 落在界内时罚为零() {
let pts = random_points(12, 3, 4.0);
let f = field_from(&pts, 0.3, &[]);
let x = flat(&pts);
let mut g = vec![0.0; x.len()];
let e = f.value_and_grad(&x, &mut g);
assert_eq!(e, 0.0, "界是照这组点定的,它本身应当零罚");
for (k, v) in g.iter().enumerate() {
assert_eq!(*v, 0.0, "第 {k} 个分量的梯度应当是 0");
}
}
#[test]
fn 精修能把违反压下去() {
let pts = random_points(15, 13, 4.0);
let f = field_from(&pts, 0.25, &[]);
let mut st = 5u64;
let mut x: Vec<f64> = (0..45).map(|_| lcg(&mut st) * 0.3).collect();
let mut g = vec![0.0; 45];
let e0 = f.value_and_grad(&x, &mut g);
let r = minimize(
&f,
&mut x,
&Options {
max_iter: 2000,
grad_tol: 1e-10,
memory: 8,
},
);
assert!(e0 > 1.0, "起点应当很差,实际 {e0:.3e}");
assert!(
r.value < 1e-10,
"没压下去:{:.3e}(起点 {e0:.3e},迭代 {})",
r.value,
r.iterations
);
}
#[test]
fn 宽区间的对也必须进力场() {
let mut b = Bounds::new(2, 0.0, 1000.0);
b.set_lower(0, 1, 1.0);
b.set_upper(0, 1, 20.0);
assert!(
b.upper(0, 1) - b.lower(0, 1) > 5.0,
"构造失效:区间必须宽过 basinThresh 才测得到东西"
);
let f = Field::new(&b, &[]);
let mut g = vec![0.0; 6];
let e = f.value_and_grad(&[0.0, 0.0, 0.0, 30.0, 0.0, 0.0], &mut g);
assert!(
e > 0.0,
"宽区间的对被漏掉了 —— 那是 RDKit basinThresh 的行为,不是我们的"
);
assert!(g[0] != 0.0, "梯度也该有,实际是 {}", g[0]);
}
#[test]
fn 下限那一支是饱和的() {
let mut b = Bounds::new(2, 0.0, 1000.0);
b.set_lower(0, 1, 2.0);
b.set_upper(0, 1, 3.0);
let f = Field::new(&b, &[]);
let mut g = vec![0.0; 6];
let e_zero = f.value_and_grad(&[0.0, 0.0, 0.0, 0.0, 0.0, 0.0], &mut g);
assert!(
(e_zero - 1.0).abs() < 1e-12,
"完全重叠时的罚应当恰好是 1,实得 {e_zero}"
);
let e_far = f.value_and_grad(&[0.0, 0.0, 0.0, 30.0, 0.0, 0.0], &mut g);
assert!(e_far > 50.0, "超上限应当罚得很重,实得 {e_far}");
}
#[test]
fn 手性项能把号翻回来() {
let pts: Vec<[f64; 3]> = vec![
[0.0, 0.0, 0.0],
[0.0, 0.0, 1.2],
[1.13, 0.0, -0.4],
[-0.57, 0.98, -0.4],
[-0.57, -0.98, -0.4],
];
let c = Center {
atom: 0,
ligands: [1, 2, 3, 4],
sign: -1.0,
};
let v0 = crate::chiral::center_volume(&pts, &c);
assert!(v0 > 0.0, "构造有问题:起始体积应当为正,实得 {v0}");
let mut b = Bounds::new(5, 0.5, 1000.0);
for i in 1..5 {
b.set_lower(0, i, 1.0);
b.set_upper(0, i, 1.4);
for j in (i + 1)..5 {
b.set_lower(i, j, 1.6);
b.set_upper(i, j, 4.0);
}
}
let f = Field::new(&b, &[c]);
let mut x = flat(&pts);
let r = minimize(
&f,
&mut x,
&Options {
max_iter: 3000,
grad_tol: 1e-9,
memory: 8,
},
);
let p: Vec<[f64; 3]> = (0..5)
.map(|i| [x[3 * i], x[3 * i + 1], x[3 * i + 2]])
.collect();
let v1 = crate::chiral::center_volume(&p, &c);
assert!(
v1 <= -UMBRELLA_LO + 1e-6,
"手性没翻过来:{v0:.3} → {v1:.3}(目标 ≤ −{UMBRELLA_LO},残值 {:.3e})",
r.value
);
let vl = crate::chiral::signed_volume(p[1], p[2], p[3], p[4]);
assert!(
vl >= VOL_LO - 1e-6,
"四配体那一项也该到位(号与中心基点相反):{vl:.3},目标 ≥ {VOL_LO}"
);
}
#[test]
fn 手性项挡得住伞形翻转() {
let mut pts: Vec<[f64; 3]> = vec![
[0.0, 0.0, 0.0],
[0.0, 0.0, 1.2],
[1.13, 0.0, -0.4],
[-0.57, 0.98, -0.4],
[-0.57, -0.98, -0.4],
];
let c = Center {
atom: 0,
ligands: [1, 2, 3, 4],
sign: 1.0,
};
let good = crate::chiral::center_volume(&pts, &c);
{
let (p, q, r) = (pts[1], pts[2], pts[3]);
let sub = |u: [f64; 3], v: [f64; 3]| [u[0] - v[0], u[1] - v[1], u[2] - v[2]];
let (a, b) = (sub(q, p), sub(r, p));
let n = [
a[1] * b[2] - a[2] * b[1],
a[2] * b[0] - a[0] * b[2],
a[0] * b[1] - a[1] * b[0],
];
let nn = n[0] * n[0] + n[1] * n[1] + n[2] * n[2];
let d = sub(pts[0], p);
let t = (d[0] * n[0] + d[1] * n[1] + d[2] * n[2]) / nn;
for k in 0..3 {
pts[0][k] -= 2.0 * t * n[k];
}
}
let v0 = crate::chiral::center_volume(&pts, &c);
assert!(good > 0.0, "起始那个正常构型该是正的,实得 {good}");
assert!(v0 < 0.0, "构造有问题:翻伞之后体积该为负,实得 {v0}");
let mut b = Bounds::new(5, 0.5, 1000.0);
for i in 1..5 {
b.set_lower(0, i, 1.0);
b.set_upper(0, i, 1.6);
for j in (i + 1)..5 {
b.set_lower(i, j, 1.6);
b.set_upper(i, j, 4.0);
}
}
let f = Field::new(&b, &[c]);
let mut x = flat(&pts);
minimize(
&f,
&mut x,
&Options {
max_iter: 3000,
grad_tol: 1e-9,
memory: 8,
},
);
let p: Vec<[f64; 3]> = (0..5)
.map(|i| [x[3 * i], x[3 * i + 1], x[3 * i + 2]])
.collect();
let v1 = crate::chiral::center_volume(&p, &c);
assert!(v1 > 0.0, "伞形翻转没被拉回来:{v0:.3} → {v1:.3}");
}
}