Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
148 changes: 86 additions & 62 deletions CodeEntropy/levels/axes.py
Original file line number Diff line number Diff line change
Expand Up @@ -144,12 +144,14 @@ def get_residue_axes(
masses=ua_masses,
dimensions=data_container.dimensions[:3],
)
rot_axes, moment_of_inertia = self.get_custom_principal_axes(moi_tensor)
rot_axes, moment_of_inertia = self.get_principal_axes_from_tensor(
moi_tensor
)
trans_axes = rot_axes # per original convention
rot_center = np.array(residue.center_of_mass())
else:
make_whole(data_container.atoms)
trans_axes = data_container.atoms.principal_axes()
trans_axes = self.get_principal_axes_from_group(data_container.atoms)
if len(edge_atom_set) == 1:
edge_atom = edge_atom_set[0]
rot_center, rot_axes = self.get_terminal_axes(
Expand Down Expand Up @@ -217,14 +219,14 @@ def get_residue_axes_from_topology(
masses=topology.residue_ua_masses,
dimensions=dimensions,
)
rot_axes, moment_of_inertia = self.get_custom_principal_axes(
rot_axes, moment_of_inertia = self.get_principal_axes_from_tensor(
moment_of_inertia_tensor
)
trans_axes = rot_axes
else:
make_whole(mol.atoms)
trans_axes = mol.atoms.principal_axes()
rot_axes, moment_of_inertia = self.get_vanilla_axes(residue_atoms)
trans_axes = self.get_principal_axes_from_group(mol.atoms)
rot_axes, moment_of_inertia = self.get_molecule_axes(residue_atoms)
center = residue_atoms.center_of_mass(unwrap=True)

return trans_axes, rot_axes, center, moment_of_inertia
Expand Down Expand Up @@ -281,7 +283,7 @@ def get_UA_axes(self, data_container, index: int, res_position):
# only the one residue => use principal axes
residue = data_container
trans_center = data_container.atoms.center_of_mass(unwrap=True)
trans_axes = data_container.atoms.principal_axes()
trans_axes = self.get_principal_axes_from_group(data_container.atoms)
else:
# residue of interest has at least one neighbour
if res_position == -1 or res_position == 1:
Expand Down Expand Up @@ -416,12 +418,12 @@ def get_UA_axes_from_topology(
masses=topology.residue_ua_masses,
dimensions=dimensions,
)
trans_axes, _moment_of_inertia = self.get_custom_principal_axes(
trans_axes, _moment_of_inertia = self.get_principal_axes_from_tensor(
moment_of_inertia_tensor
)
else:
make_whole(residue_atoms)
trans_axes = residue_atoms.principal_axes()
trans_axes = self.get_principal_axes_from_group(residue_atoms)

center = heavy_atom.position
rot_axes, moment_of_inertia = self.get_bonded_axes_from_topology(
Expand Down Expand Up @@ -478,26 +480,26 @@ def get_bonded_axes_from_topology(
ua_all = u.atoms[topology.ua_all_atom_indices]

if len(heavy_bonded) == 0:
custom_axes, custom_moment_of_inertia = self.get_vanilla_axes(ua_all)
custom_axes, custom_moment_of_inertia = self.get_molecule_axes(ua_all)

if len(heavy_bonded) == 1 and len(light_bonded) == 0:
custom_axes = self.get_custom_axes(
custom_axes = self.get_bonded_vector_axes(
a=heavy_atom.position,
b_list=[heavy_bonded[0].position],
c=np.zeros(3),
dimensions=dimensions,
)

if len(heavy_bonded) == 1 and len(light_bonded) >= 1:
custom_axes = self.get_custom_axes(
custom_axes = self.get_bonded_vector_axes(
a=heavy_atom.position,
b_list=[heavy_bonded[0].position],
c=light_bonded[0].position,
dimensions=dimensions,
)

if len(heavy_bonded) >= 2:
custom_axes = self.get_custom_axes(
custom_axes = self.get_bonded_vector_axes(
a=heavy_atom.position,
b_list=heavy_bonded.positions,
c=heavy_bonded[1].position,
Expand Down Expand Up @@ -611,7 +613,7 @@ def get_terminal_axes(self, residue, edge, dimensions):
if len(bonded_atoms) == 0:
# there is only one heavy atom in the residue
rot_center = edge.position
rot_axes = residue.atoms.principal_axes()
rot_axes = self.get_principal_axes_from_group(residue.atoms)
else:
average_bonded = np.zeros(3)
for bonded_atom in bonded_atoms:
Expand All @@ -632,7 +634,7 @@ def get_terminal_axes(self, residue, edge, dimensions):
)
else:
rot_center = edge.position
rot_axes = self.get_custom_axes(
rot_axes = self.get_bonded_vector_axes(
a=edge.position,
b_list=[average_bonded],
c=np.zeros(3),
Expand Down Expand Up @@ -672,7 +674,7 @@ def get_non_terminal_axes(self, residue, edges, dimensions):
)
else:
rot_center = (edges[0].position + edges[1].position) / 2
rot_axes = self.get_custom_axes(
rot_axes = self.get_bonded_vector_axes(
a=rot_center,
b_list=[edges[0].position],
c=np.zeros(3),
Expand Down Expand Up @@ -749,11 +751,11 @@ def get_bonded_axes(self, system, atom, dimensions: np.ndarray):

# case1
if len(heavy_bonded) == 0:
custom_axes, custom_moment_of_inertia = self.get_vanilla_axes(ua_all)
custom_axes, custom_moment_of_inertia = self.get_molecule_axes(ua_all)

# case2
if len(heavy_bonded) == 1 and len(light_bonded) == 0:
custom_axes = self.get_custom_axes(
custom_axes = self.get_bonded_vector_axes(
a=atom.position,
b_list=[heavy_bonded[0].position],
c=np.zeros(3),
Expand All @@ -762,7 +764,7 @@ def get_bonded_axes(self, system, atom, dimensions: np.ndarray):

# case3
if len(heavy_bonded) == 1 and len(light_bonded) >= 1:
custom_axes = self.get_custom_axes(
custom_axes = self.get_bonded_vector_axes(
a=atom.position,
b_list=[heavy_bonded[0].position],
c=light_bonded[0].position,
Expand All @@ -772,7 +774,7 @@ def get_bonded_axes(self, system, atom, dimensions: np.ndarray):
# case4 (not used in original 2019 code; case5 used instead)
# case5
if len(heavy_bonded) >= 2:
custom_axes = self.get_custom_axes(
custom_axes = self.get_bonded_vector_axes(
a=atom.position,
b_list=heavy_bonded.positions,
c=heavy_bonded[1].position,
Expand Down Expand Up @@ -812,14 +814,15 @@ def find_bonded_atoms(self, atom_idx: int, system):
bonded_H_atoms = bonded_atoms.select_atoms("mass 1 to 1.1")
return bonded_heavy_atoms, bonded_H_atoms

def get_vanilla_axes(self, molecule):
"""Get principal axes and sorted principal moments (vanilla method).
def get_molecule_axes(self, molecule):
"""Get principal axes and sorted principal moments.

Compute the principal axes and moments of inertia for a molecule using
MDAnalysis built-in functionality.
Compute the principal axes and moments of inertia for a molecule from its
moment of inertia tensor.

The original description is preserved:
- The molecule is made whole to ensure correct handling of PBC.
- The molecule is made whole to ensure correct handling of PBC. The tensor
is taken after that, so it does not depend on where the molecule sits
relative to the periodic boundary.
- The moments are obtained by diagonalising the moment of inertia tensor.
- Eigenvalues are returned sorted from largest to smallest magnitude.

Expand All @@ -832,17 +835,10 @@ def get_vanilla_axes(self, molecule):
- principal_axes: (3, 3) axes.
- moment_of_inertia: (3,) moments sorted descending by absolute value.
"""
moment_of_inertia_tensor = molecule.moment_of_inertia(unwrap=True)
make_whole(molecule.atoms)
principal_axes = molecule.principal_axes()

eigenvalues, _ = np.linalg.eig(moment_of_inertia_tensor)
order = np.argsort(np.abs(eigenvalues))[::-1]
moment_of_inertia = eigenvalues[order]
return self.get_principal_axes_from_tensor(molecule.moment_of_inertia())

return principal_axes, moment_of_inertia

def get_custom_axes(
def get_bonded_vector_axes(
self,
a: np.ndarray,
b_list: Sequence[np.ndarray],
Expand Down Expand Up @@ -1074,44 +1070,72 @@ def get_moment_of_inertia_tensor(

return moment_of_inertia_tensor

def get_custom_principal_axes(
self, moment_of_inertia_tensor: np.ndarray
def get_principal_axes_from_tensor(
self,
moment_of_inertia_tensor: np.ndarray,
degeneracy_rtol: float = 100 * np.finfo(np.float32).eps,
reference_vector: Sequence[float] = (1.0, 2.0, 3.0),
) -> tuple[np.ndarray, np.ndarray]:
"""Compute principal axes and moments from a custom MOI tensor.
"""Compute the principal axes and moments of a moment of inertia tensor.

Principal axes and centre of axes from the ordered eigenvalues and
eigenvectors of a moment of inertia tensor. This function allows for a
custom moment of inertia tensor to be used, which isn't possible with the
built-in MDAnalysis principal_axes() function.
The eigenvectors of a symmetric tensor are only defined up to a sign, and,
for equal moments, up to a rotation within the degenerate eigenspace. The
axes are returned in a canonical form, so that a given tensor always
produces the same axes. The convention is:

Original behaviour preserved:

- Eigenvalues are sorted by descending absolute magnitude.
- Eigenvectors are transposed so axes are returned as rows.
- Z axis is flipped to enforce the same handedness convention as the
original implementation.
- Moments are sorted by descending absolute value.
- Within a degenerate eigenspace (moments equal to within
``degeneracy_rtol`` of the largest), the basis is the eigenvectors of a
reference tensor built from ``reference_vector``, projected onto that
eigenspace.
- The first two axes are signed to have a positive projection onto
``reference_vector``, and the third is their cross product.

Args:
moment_of_inertia_tensor: (3, 3) custom inertia tensor.
moment_of_inertia_tensor: (3, 3) symmetric inertia tensor.
degeneracy_rtol: Relative tolerance, with respect to the largest
moment, for treating moments as degenerate. The default is a
multiple of the float32 precision of the coordinates.
reference_vector: Fixed direction used to choose axis signs and the
basis within a degenerate eigenspace. Its components must be
distinct, so that no symmetric or axis-aligned direction is
perpendicular to it or ties on its components.

Returns:
Tuple[np.ndarray, np.ndarray]:
- principal_axes: (3, 3) principal axes (rows).
- moment_of_inertia: (3,) principal moments.
- principal_moments: (3,) moments sorted by descending magnitude.
"""
eigenvalues, eigenvectors = np.linalg.eig(moment_of_inertia_tensor)
order = np.abs(eigenvalues).argsort()[::-1] # descending order
transposed = np.transpose(eigenvectors) # columns -> rows
moment_of_inertia = eigenvalues[order]
principal_axes = transposed[order]

# point z axis in correct direction, as per original code
cross_xy = np.cross(principal_axes[0], principal_axes[1])
dot_z = float(np.dot(cross_xy, principal_axes[2]))
if dot_z < 0:
principal_axes[2] *= -1

return principal_axes, moment_of_inertia
principal_moments, eigenvectors = np.linalg.eigh(moment_of_inertia_tensor)
order = np.argsort(np.abs(principal_moments))[::-1]
principal_moments = principal_moments[order]
eigenvectors = eigenvectors[:, order]

reference = np.asarray(reference_vector, dtype=float)
gaps = np.abs(np.diff(np.abs(principal_moments)))
boundaries = (
np.flatnonzero(gaps > degeneracy_rtol * np.abs(principal_moments[0])) + 1
)
for subspace in np.split(np.arange(3), boundaries):
if len(subspace) > 1:
basis = eigenvectors[:, subspace]
_, rotation = np.linalg.eigh(basis.T @ np.diag(reference) @ basis)
eigenvectors[:, subspace] = basis @ rotation

for axis in (0, 1):
if eigenvectors[:, axis] @ reference < 0:
eigenvectors[:, axis] *= -1

z_axis = np.cross(eigenvectors[:, 0], eigenvectors[:, 1])
principal_axes = np.array([eigenvectors[:, 0], eigenvectors[:, 1], z_axis])
return principal_axes, principal_moments

def get_principal_axes_from_group(self, group) -> np.ndarray:
"""Return the canonical principal axes (rows) of an atom group."""
principal_axes, _ = self.get_principal_axes_from_tensor(
group.atoms.moment_of_inertia()
)
return principal_axes

def get_UA_masses(self, molecule) -> list[float]:
"""Return united-atom (UA) masses for a molecule.
Expand Down
12 changes: 6 additions & 6 deletions CodeEntropy/levels/nodes/covariance.py
Original file line number Diff line number Diff line change
Expand Up @@ -506,8 +506,8 @@ def _build_ua_vectors(
# principal axes
make_whole(residue.atoms)
make_whole(bead)
trans_axes = residue.atoms.principal_axes()
rot_axes, moi = axes_manager.get_vanilla_axes(bead)
trans_axes = axes_manager.get_principal_axes_from_group(residue.atoms)
rot_axes, moi = axes_manager.get_molecule_axes(bead)
center = bead.center_of_mass(unwrap=True)

force_vecs.append(
Expand Down Expand Up @@ -658,8 +658,8 @@ def _get_residue_axes(
make_whole(mol.atoms)
make_whole(bead)

trans_axes = mol.atoms.principal_axes()
rot_axes, moi = axes_manager.get_vanilla_axes(bead)
trans_axes = axes_manager.get_principal_axes_from_group(mol.atoms)
rot_axes, moi = axes_manager.get_molecule_axes(bead)
center = bead.center_of_mass(unwrap=True)
return (
np.asarray(trans_axes),
Expand Down Expand Up @@ -688,8 +688,8 @@ def _get_polymer_axes(
make_whole(mol.atoms)
make_whole(bead)

trans_axes = mol.atoms.principal_axes()
rot_axes, moi = axes_manager.get_vanilla_axes(bead)
trans_axes = axes_manager.get_principal_axes_from_group(mol.atoms)
rot_axes, moi = axes_manager.get_molecule_axes(bead)
center = bead.center_of_mass(unwrap=True)

return (
Expand Down
10 changes: 5 additions & 5 deletions tests/regression/baselines/benzaldehyde/axes_off.json
Original file line number Diff line number Diff line change
Expand Up @@ -2,15 +2,15 @@
"groups": {
"0": {
"components": {
"united_atom:Transvibrational": 0.08982962903796131,
"united_atom:Rovibrational": 32.16018134884085,
"residue:FTmat-Transvibrational": 88.7671666695003,
"residue:FTmat-Rovibrational": 61.61036267672132,
"united_atom:Transvibrational": 0.06950663196232491,
"united_atom:Rovibrational": 35.19655714759131,
"residue:FTmat-Transvibrational": 86.6413638447789,
"residue:FTmat-Rovibrational": 62.24303760058523,
"united_atom:Conformational": 0.0,
"residue:Conformational": 0.0,
"residue:Orientational": 20.481571492615355
},
"total": 203.1091118167158
"total": 204.6320367175331
}
}
}
Original file line number Diff line number Diff line change
Expand Up @@ -2,15 +2,15 @@
"groups": {
"0": {
"components": {
"united_atom:Transvibrational": 0.07119323721997475,
"united_atom:Rovibrational": 49.68669738152346,
"residue:Transvibrational": 69.48692941204929,
"residue:Rovibrational": 68.46147102540942,
"united_atom:Transvibrational": 0.06817494136874551,
"united_atom:Rovibrational": 49.686697381523466,
"residue:Transvibrational": 69.25273989458984,
"residue:Rovibrational": 68.64810133422232,
"united_atom:Conformational": 0.0,
"residue:Conformational": 0.0,
"residue:Orientational": 20.481571492615355
},
"total": 208.1878625488175
"total": 208.13728504431973
}
}
}
10 changes: 5 additions & 5 deletions tests/regression/baselines/benzaldehyde/default.json
Original file line number Diff line number Diff line change
Expand Up @@ -2,15 +2,15 @@
"groups": {
"0": {
"components": {
"united_atom:Transvibrational": 24.88518233240474,
"united_atom:Rovibrational": 27.950376507672583,
"residue:FTmat-Transvibrational": 71.03412922724692,
"residue:FTmat-Rovibrational": 59.44169664956799,
"united_atom:Transvibrational": 24.921822103314035,
"united_atom:Rovibrational": 27.950376507672594,
"residue:FTmat-Transvibrational": 70.97090266225699,
"residue:FTmat-Rovibrational": 59.588669973122094,
"united_atom:Conformational": 7.522643899702263,
"residue:Conformational": 0.0,
"residue:Orientational": 59.43920971558428
},
"total": 250.27323833217878
"total": 250.39362486165226
}
}
}
Loading
Loading