#![allow(unused)]
use std::f32::consts::TAU;
use lin_alg::{
f32::{Mat3 as Mat3F32, Quaternion, Vec3},
f64::Vec3 as Vec3F64,
};
use rand::{
RngExt,
distr::{Distribution, Uniform},
prelude::{SliceRandom, ThreadRng},
};
use rand_distr::Normal;
use crate::{
AtomDynamics, ComputationDevice, MdState, MolDynamics, NATIVE_TO_KCAL,
barostat::SimBox,
solvent::{
WaterMolOpc,
init::{MIN_WATER_O_O_DIST_SQ, n_water_mols},
shrinking_box::ShrinkingBoxCfg,
},
thermostat::{GAS_CONST_R, KB_A2_PS2_PER_K_PER_AMU},
};
pub fn make_water_mols_grid(
cell: &SimBox,
temperature_tgt: f32,
zero_com_drift: bool,
) -> Vec<WaterMolOpc> {
println!("Initializing a solvent grid, as part of template preparation...");
let mut rng = rand::rng();
let distro = Uniform::<f32>::new(0.0, 1.0).unwrap();
let n_mols = n_water_mols(cell, &[]);
let mut result: Vec<WaterMolOpc> = Vec::with_capacity(n_mols);
let lx = cell.bounds_high.x - cell.bounds_low.x;
let ly = cell.bounds_high.y - cell.bounds_low.y;
let lz = cell.bounds_high.z - cell.bounds_low.z;
let base = (n_mols as f32).cbrt().round().max(1.0) as usize;
let n_x = base;
let n_y = base;
let n_z = n_mols.div_ceil(n_x * n_y);
let spacing_x = lx / n_x as f32;
let spacing_y = ly / n_y as f32;
let spacing_z = lz / n_z as f32;
let fault_ratio = 3;
let mut num_added = 0;
let mut loops_used = 0;
'outer: for i in 0..n_mols * fault_ratio {
let a = i % n_x;
let b = (i / n_x) % n_y;
let c = (i / (n_x * n_y)) % n_z;
let posit = Vec3::new(
cell.bounds_low.x + (a as f32 + 0.5) * spacing_x,
cell.bounds_low.y + (b as f32 + 0.5) * spacing_y,
cell.bounds_low.z + (c as f32 + 0.5) * spacing_z,
);
for w in &result {
let dist_sq = (w.o.posit - posit).magnitude_squared();
if dist_sq < MIN_WATER_O_O_DIST_SQ {
loops_used += 1;
continue 'outer;
}
}
result.push(WaterMolOpc::new(
posit,
Vec3::new_zero(),
Quaternion::random(&mut rng, Some(distro)),
));
num_added += 1;
if num_added == n_mols {
break;
}
loops_used += 1;
}
init_velocities(&mut result, temperature_tgt, zero_com_drift, &mut rng);
println!(
"Added {} / {n_mols} solvent mols. Used {loops_used} loops",
result.len()
);
result
}
fn init_velocities(
mols: &mut [WaterMolOpc],
t_target: f32,
zero_com_drift: bool,
rng: &mut ThreadRng,
) {
let kT = KB_A2_PS2_PER_K_PER_AMU * t_target;
for m in mols.iter_mut() {
let (r_com, m_tot) = {
let mut r = Vec3::new_zero();
let mut m_tot = 0.0;
for a in [&m.o, &m.h0, &m.h1] {
r += a.posit * a.mass;
m_tot += a.mass;
}
(r / m_tot, m_tot)
};
let r_0 = m.o.posit - r_com;
let r_h0 = m.h0.posit - r_com;
let r_h1 = m.h1.posit - r_com;
let sigma_v = (kT / m_tot).sqrt();
let n = Normal::new(0.0, sigma_v).unwrap();
let v_com = Vec3::new(n.sample(rng), n.sample(rng), n.sample(rng));
let inertia = |r: Vec3, mass: f32| {
let r2 = r.dot(r);
[
[
mass * (r2 - r.x * r.x),
-mass * r.x * r.y,
-mass * r.x * r.z,
],
[
-mass * r.y * r.x,
mass * (r2 - r.y * r.y),
-mass * r.y * r.z,
],
[
-mass * r.z * r.x,
-mass * r.z * r.y,
mass * (r2 - r.z * r.z),
],
]
};
let mut I_arr = inertia(r_0, m.o.mass);
let add_I = |I: &mut [[f32; 3]; 3], J: [[f32; 3]; 3]| {
for i in 0..3 {
for j in 0..3 {
I[i][j] += J[i][j];
}
}
};
add_I(&mut I_arr, inertia(r_h0, m.h0.mass));
add_I(&mut I_arr, inertia(r_h1, m.h1.mass));
let I = Mat3F32::from_arr(I_arr);
let (eigvecs, eigvals) = I.eigen_vecs_vals();
let L_principal = Vec3::new(
Normal::new(0.0, (kT * eigvals.x.max(0.0)).sqrt())
.unwrap()
.sample(rng),
Normal::new(0.0, (kT * eigvals.y.max(0.0)).sqrt())
.unwrap()
.sample(rng),
Normal::new(0.0, (kT * eigvals.z.max(0.0)).sqrt())
.unwrap()
.sample(rng),
);
let L_world = eigvecs * L_principal; let omega = I.solve_system(L_world);
m.o.vel = v_com + omega.cross(r_0);
m.h0.vel = v_com + omega.cross(r_h0);
m.h1.vel = v_com + omega.cross(r_h1);
}
if zero_com_drift {
remove_com_velocity(mols);
}
let (ke_raw, dof) = kinetic_energy_and_dof(mols, zero_com_drift);
let temperature_meas = (2.0 * ke_raw) / (dof as f32 * GAS_CONST_R as f32);
let lambda = (t_target / temperature_meas).sqrt();
for a in atoms_mut(mols) {
if a.mass > 0.0 {
a.vel *= lambda;
}
}
}
fn kinetic_energy_and_dof(mols: &[WaterMolOpc], zero_com_drift: bool) -> (f32, usize) {
let mut ke = 0.;
for w in mols {
ke += (w.o.mass * w.o.vel.magnitude_squared()) as f64;
ke += (w.h0.mass * w.h0.vel.magnitude_squared()) as f64;
ke += (w.h1.mass * w.h1.vel.magnitude_squared()) as f64;
}
let mut dof = mols.len() * 3;
if zero_com_drift {
dof = dof.saturating_sub(3);
}
(ke as f32 * 0.5 * NATIVE_TO_KCAL, dof)
}
fn temperature_from_water_velocities(mols: &[WaterMolOpc], zero_com_drift: bool) -> Option<f32> {
if mols.is_empty() {
return None;
}
let (ke_raw, dof) = kinetic_energy_and_dof(mols, zero_com_drift);
if ke_raw <= 0.0 || dof == 0 {
return None;
}
Some((2.0 * ke_raw) / (dof as f32 * GAS_CONST_R as f32))
}
fn rigid_body_centroid(atoms: &[AtomDynamics]) -> Vec3F64 {
let mut centroid = Vec3F64::new_zero();
let mut mass_total = 0.;
for atom in atoms {
let mass = atom.mass as f64;
centroid += Vec3F64::from(atom.posit) * mass;
mass_total += mass;
}
if mass_total <= f64::EPSILON {
let sum = atoms.iter().fold(Vec3F64::new_zero(), |acc, atom| {
acc + Vec3F64::from(atom.posit)
});
return sum * (1.0 / atoms.len().max(1) as f64);
}
centroid / mass_total
}
fn water_centroid(water: &WaterMolOpc) -> Vec3F64 {
let atoms = [&water.o, &water.h0, &water.h1];
let mut centroid = Vec3F64::new_zero();
let mut mass_total = 0.;
for atom in atoms {
let mass = atom.mass as f64;
centroid += Vec3F64::from(atom.posit) * mass;
mass_total += mass;
}
centroid / mass_total.max(f64::EPSILON)
}
fn scale_centroid_position(
centroid: Vec3F64,
current_center: Vec3,
new_center: Vec3,
scale_x: f64,
scale_y: f64,
scale_z: f64,
) -> Vec3F64 {
let current_center_f64: Vec3F64 = current_center.into();
let new_center_f64: Vec3F64 = new_center.into();
let relative_pos = centroid - current_center_f64;
let scaled_relative_pos = Vec3F64::new(
relative_pos.x * scale_x,
relative_pos.y * scale_y,
relative_pos.z * scale_z,
);
new_center_f64 + scaled_relative_pos
}
fn shift_rigid_body(atoms: &mut [AtomDynamics], displacement: Vec3) {
for atom in atoms {
atom.posit += displacement;
}
}
fn shift_water_molecule(water: &mut WaterMolOpc, displacement: Vec3) {
water.o.posit += displacement;
water.h0.posit += displacement;
water.h1.posit += displacement;
water.m.posit += displacement;
}
impl MdState {
pub fn shrink_cell_towards(
&mut self,
dev: &ComputationDevice,
target_cell: SimBox,
cfg: ShrinkingBoxCfg,
) -> bool {
let current_cell = self.cell;
let Some(new_cell) = cfg.next_cell(current_cell, target_cell) else {
return false;
};
let scale_x = (new_cell.extent.x / current_cell.extent.x) as f64;
let scale_y = (new_cell.extent.y / current_cell.extent.y) as f64;
let scale_z = (new_cell.extent.z / current_cell.extent.z) as f64;
let current_center = current_cell.center();
let new_center = new_cell.center();
for (mol_i, &start) in self.mol_start_indices.iter().enumerate() {
let end = self
.mol_start_indices
.get(mol_i + 1)
.copied()
.unwrap_or(self.atoms.len());
let mol_atoms = &mut self.atoms[start..end];
let centroid = rigid_body_centroid(mol_atoms);
let new_centroid = scale_centroid_position(
centroid,
current_center,
new_center,
scale_x,
scale_y,
scale_z,
);
shift_rigid_body(mol_atoms, (new_centroid - centroid).into());
}
for water in &mut self.water {
let centroid = water_centroid(water);
let new_centroid = scale_centroid_position(
centroid,
current_center,
new_center,
scale_x,
scale_y,
scale_z,
);
shift_water_molecule(water, (new_centroid - centroid).into());
}
self.cell = new_cell;
self.rebuild_spatial_caches(dev);
true
}
pub fn rebuild_spatial_caches(&mut self, dev: &ComputationDevice) {
self.spme_force_prev = None;
self.build_all_neighbors(dev);
self.regen_pme(dev);
}
pub fn redistribute_interleaved_opc_waters(
&mut self,
dev: &ComputationDevice,
solvent_centers: &[Vec3F64],
solvent_spacing: Vec3F64,
) -> bool {
let water_count = self.water.len();
if water_count == 0 {
return true;
}
let solvent_atom_posits: Vec<_> = self.atoms.iter().map(|atom| atom.posit).collect();
let water_temperature_tgt =
temperature_from_water_velocities(&self.water, false).unwrap_or(300.0);
let placed_water = place_interleaved_opc_waters(
solvent_centers,
solvent_spacing,
&solvent_atom_posits,
water_count,
&self.cell,
water_temperature_tgt,
&mut rand::rng(),
);
if placed_water.len() != water_count {
eprintln!(
"redistribute_interleaved_opc_waters: placed {} / {water_count} waters; \
retaining template-seeded waters.",
placed_water.len()
);
return false;
}
self.water = placed_water;
self.water_pme_sites_forces = vec![[Vec3F64::new_zero(); 3]; self.water.len()];
self.rebuild_spatial_caches(dev);
true
}
}
fn water_conflicts_with_solvent(
water: &WaterMolOpc,
solvent_atom_posits: &[Vec3],
cell: &SimBox,
) -> bool {
const MIN_O_TO_SOLVENT_SQ: f32 = 1.7 * 1.7;
const MIN_H_TO_SOLVENT_SQ: f32 = 1.0 * 1.0;
for solvent_posit in solvent_atom_posits {
let o_diff = water.o.posit - *solvent_posit;
if cell.min_image(o_diff).magnitude_squared() < MIN_O_TO_SOLVENT_SQ {
return true;
}
let h0_diff = water.h0.posit - *solvent_posit;
if cell.min_image(h0_diff).magnitude_squared() < MIN_H_TO_SOLVENT_SQ {
return true;
}
let h1_diff = water.h1.posit - *solvent_posit;
if cell.min_image(h1_diff).magnitude_squared() < MIN_H_TO_SOLVENT_SQ {
return true;
}
}
false
}
fn water_conflicts_with_water(
candidate: &WaterMolOpc,
placed: &[WaterMolOpc],
cell: &SimBox,
) -> bool {
const PBC_MIN_WATER_O_O_DIST_SQ: f32 = 2.8 * 2.8;
for water in placed {
let diff = water.o.posit - candidate.o.posit;
let direct_sq = diff.magnitude_squared();
if direct_sq < MIN_WATER_O_O_DIST_SQ {
return true;
}
let min_image_sq = cell.min_image(diff).magnitude_squared();
if min_image_sq < MIN_WATER_O_O_DIST_SQ {
return true;
}
if min_image_sq < PBC_MIN_WATER_O_O_DIST_SQ && min_image_sq < direct_sq {
return true;
}
}
false
}
fn make_interleaved_water_offsets() -> Vec<Vec3F64> {
let mut offsets = Vec::new();
for &scale in &[0.22_f64, 0.38_f64] {
for dx in -1..=1 {
for dy in -1..=1 {
for dz in -1..=1 {
if dx == 0 && dy == 0 && dz == 0 {
continue;
}
offsets.push(Vec3F64::new(
dx as f64 * scale,
dy as f64 * scale,
dz as f64 * scale,
));
}
}
}
}
offsets.sort_by(|a, b| {
let a_nonzero = usize::from(a.x != 0.0) + usize::from(a.y != 0.0) + usize::from(a.z != 0.0);
let b_nonzero = usize::from(b.x != 0.0) + usize::from(b.y != 0.0) + usize::from(b.z != 0.0);
a_nonzero
.cmp(&b_nonzero)
.then_with(|| a.magnitude_squared().total_cmp(&b.magnitude_squared()))
});
offsets
}
fn place_interleaved_opc_waters(
solvent_centers: &[Vec3F64],
solvent_spacing: Vec3F64,
solvent_atom_posits: &[Vec3],
water_count: usize,
cell: &SimBox,
temperature_tgt: f32,
rng: &mut ThreadRng,
) -> Vec<WaterMolOpc> {
const MAX_WATER_ROT_ATTEMPTS: usize = 12;
const MIN_CENTER_OFFSET_FRAC: f64 = 0.12;
const MAX_FALLBACK_ATTEMPTS_PER_WATER: usize = 40;
if water_count == 0 || solvent_centers.is_empty() {
return Vec::new();
}
let offsets_unit = make_interleaved_water_offsets();
let cell_offsets: Vec<Vec<Vec3F64>> = solvent_centers
.iter()
.map(|_| {
let mut offsets = offsets_unit.clone();
offsets.shuffle(rng);
offsets
})
.collect();
let mut placed = Vec::with_capacity(water_count);
let distro = Uniform::<f32>::new(0.0, 1.0).unwrap();
let try_place_candidate =
|o_posit: Vec3, placed: &mut Vec<WaterMolOpc>, rng: &mut ThreadRng| -> bool {
if !cell.contains(o_posit) {
return false;
}
for _ in 0..MAX_WATER_ROT_ATTEMPTS {
let candidate = WaterMolOpc::new(
o_posit,
Vec3::new_zero(),
Quaternion::random(rng, Some(distro)),
);
if water_conflicts_with_solvent(&candidate, solvent_atom_posits, cell) {
continue;
}
if water_conflicts_with_water(&candidate, placed, cell) {
continue;
}
placed.push(candidate);
return true;
}
false
};
'round_robin: for candidate_idx in 0..offsets_unit.len() {
for (cell_idx, center) in solvent_centers.iter().enumerate() {
if placed.len() == water_count {
break 'round_robin;
}
let offset = cell_offsets[cell_idx][candidate_idx];
let o_posit = Vec3::new(
(center.x + offset.x * solvent_spacing.x) as f32,
(center.y + offset.y * solvent_spacing.y) as f32,
(center.z + offset.z * solvent_spacing.z) as f32,
);
let _ = try_place_candidate(o_posit, &mut placed, rng);
}
}
let fallback_attempts = water_count * MAX_FALLBACK_ATTEMPTS_PER_WATER;
for _ in 0..fallback_attempts {
if placed.len() == water_count {
break;
}
let center = solvent_centers[rng.random_range(0..solvent_centers.len())];
let dx = rng.random_range(-0.42_f64..0.42_f64);
let dy = rng.random_range(-0.42_f64..0.42_f64);
let dz = rng.random_range(-0.42_f64..0.42_f64);
if dx.abs() < MIN_CENTER_OFFSET_FRAC
&& dy.abs() < MIN_CENTER_OFFSET_FRAC
&& dz.abs() < MIN_CENTER_OFFSET_FRAC
{
continue;
}
let o_posit = Vec3::new(
(center.x + dx * solvent_spacing.x) as f32,
(center.y + dy * solvent_spacing.y) as f32,
(center.z + dz * solvent_spacing.z) as f32,
);
let _ = try_place_candidate(o_posit, &mut placed, rng);
}
init_velocities(&mut placed, temperature_tgt, true, rng);
placed
}
pub fn atoms_mut(mols: &mut [WaterMolOpc]) -> impl Iterator<Item = &mut AtomDynamics> {
mols.iter_mut()
.flat_map(|m| [&mut m.o, &mut m.h0, &mut m.h1].into_iter())
}
fn remove_com_velocity(mols: &mut [WaterMolOpc]) {
let mut p = Vec3::new_zero();
let mut m_tot = 0.0;
for a in atoms_mut(mols) {
p += a.vel * a.mass;
m_tot += a.mass;
}
let v_com = p / m_tot;
for a in atoms_mut(mols) {
a.vel -= v_com;
}
}
#[derive(Clone, Copy, Debug, PartialEq)]
pub enum CustomSolventCount {
Specified(usize),
Auto(f64),
}