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
3 changes: 3 additions & 0 deletions .github/workflows/tests.yml
Original file line number Diff line number Diff line change
Expand Up @@ -5,6 +5,9 @@ on: [push]
jobs:
build-linux:
runs-on: ubuntu-latest
defaults:
run:
shell: bash -l {0}

steps:
- name: Checkout repository
Expand Down
83 changes: 37 additions & 46 deletions src/molearn/analysis/analyser.py
Original file line number Diff line number Diff line change
Expand Up @@ -142,6 +142,25 @@ def _set_metadata(self, data: PDBData, bundle: DatasetBundle) -> None:
system_check = (self.n_atoms == bundle.dataset.shape[1] and self.atoms == data.atoms)
assert system_check, "Datasets have different number of atoms or atom types. Have you selected the same atoms?"

atom_order = [tuple(a) for a in data.get_atominfo()]
if not hasattr(self, "atom_order"):
self.atom_order = atom_order
elif self.atom_order != atom_order:
i = next((k for k in range(min(len(self.atom_order), len(atom_order)))
if self.atom_order[k] != atom_order[k]), 0)
raise ValueError(
f"dataset atoms are in a different order from the datasets already "
f"loaded. First difference at index {i}: expected "
f"{self.atom_order[i]}, got {atom_order[i]} (each entry is "
f"[name, resname, resid]).\n"
f"Index-based analyses (get_inversions, get_bondlengths) would be "
f"silently wrong. Reorder the coordinates to match, e.g.\n"
f" ref = reference_data.get_atominfo()\n"
f" pos = {{tuple(a): i for i, a in enumerate(this_data.get_atominfo())}}\n"
f" perm = [pos[tuple(a)] for a in ref]\n"
f" coords = coords[:, perm]"
)

def _prepare_bundle(self, data: PDBData) -> DatasetBundle:
dataset = data.dataset
if dataset.ndim != 3:
Expand Down Expand Up @@ -404,26 +423,9 @@ def get_inversions(self, key) -> dict[str, np.ndarray]:
% missing
)

# Get atom indices
mol_df = self.mol.data
indices = dict()
for resid in mol_df.resid.unique():
resname = mol_df[mol_df["resid"] == resid].resname.unique()[0]
if not resname == "GLY":
N_id = mol_df[
(mol_df["resid"] == resid) & (mol_df["name"] == "N")
].index[0]
C_id = mol_df[
(mol_df["resid"] == resid) & (mol_df["name"] == "C")
].index[0]
CA_id = mol_df[
(mol_df["resid"] == resid) & (mol_df["name"] == "CA")
].index[0]
CB_id = mol_df[
(mol_df["resid"] == resid) & (mol_df["name"] == "CB")
].index[0]
indices[resname + str(resid)] = (N_id, CA_id, C_id, CB_id)
idx = np.asarray(list(indices.values()))
has_cb = self.indices["CB"] >= 0 # False for glycine
idx = np.stack([self.indices["N"][has_cb], self.indices["CA"][has_cb],
self.indices["C"][has_cb], self.indices["CB"][has_cb]], axis=1)

if key in self._datasets.keys():
dataset = self.get_dataset(key, scale=True)
Expand Down Expand Up @@ -481,26 +483,15 @@ def get_bondlengths(self, key) -> dict[str, dict[str, np.ndarray]]:
else:
raise ValueError("Selected atoms should contain at least N, CA, and C.")

mol_df = self.mol.data
for resid in mol_df.resid.unique():
resname = mol_df[mol_df["resid"] == resid].resname.unique()[0]

N_id = mol_df[(mol_df["resid"] == resid) & (mol_df["name"] == "N")].index[0]
CA_id = mol_df[(mol_df["resid"] == resid) & (mol_df["name"] == "CA")].index[0]
C_id = mol_df[(mol_df["resid"] == resid) & (mol_df["name"] == "C")].index[0]
indices["N-CA"].append((N_id, CA_id))
indices["CA-C"].append((CA_id, C_id))
if resname != "GLY" and "CB" in self.atoms:
CB_id = mol_df[
(mol_df["resid"] == resid) & (mol_df["name"] == "CB")
].index[0]
indices["CA-CB"].append((CA_id, CB_id))

if resid != max(mol_df.resid.unique()):
next_N_id = mol_df[
(mol_df["resid"] == (resid + 1)) & (mol_df["name"] == "N")
].index[0]
indices["C-N"].append((C_id, next_N_id))
N, CA, C = self.indices["N"], self.indices["CA"], self.indices["C"]
indices["N-CA"] = list(zip(N, CA))
indices["CA-C"] = list(zip(CA, C))
# NOTE: consecutive residues are bonded regardless of chain, so a C-N
# "bond" is reported across chain breaks.
indices["C-N"] = list(zip(C[:-1], N[1:]))
if "CA-CB" in indices:
has_cb = self.indices["CB"] >= 0
indices["CA-CB"] = list(zip(CA[has_cb], self.indices["CB"][has_cb]))

# Look for the key in self._datasets and self._encoded
if key in self._datasets.keys():
Expand Down Expand Up @@ -547,19 +538,19 @@ def _get_dihedrals(self, data):
CA = data[:, self.indices['CA'], :].numpy()
C = data[:, self.indices['C'], :].numpy()
C_prev = np.roll(C, shift=1, axis=1)
C_next = np.roll(C, shift=-1, axis=1)
N_next = np.roll(N, shift=-1, axis=1)
CA_next = np.roll(CA, shift=-1, axis=1)

# φ: C_{i-1}, N_i, CA_i, C_i
phi = self._dihedrals(C_prev[:, 1:], N[:, 1:], CA[:, 1:], C[:, 1:])
# ψ: N_i, CA_i, C_i, N_{i+1}
psi = self._dihedrals(N[:, :-1], CA[:, :-1], C[:, :-1], N_next[:, :-1])
# ω: C_i, N_{i+1}, CA_{i+1}, C_{i+1}
omega = self._dihedrals(C[:, :-1], N_next[:, :-1], CA[:, :-1], C_next[:, :-1])
# ω: CA_i, C_i, N_{i+1}, CA_{i+1}
omega = self._dihedrals(CA[:, :-1], C[:, :-1], N_next[:, :-1], CA_next[:, :-1])
dihedrals = {"Phi": phi, "Psi": psi, "Omega": omega}

if 'CB' in self.atoms:
valid = (self.indices['CB'] > 0)
valid = (self.indices['CB'] >= 0) # 0 is a valid atom index
CB_atoms = self.indices['CB'][valid]
CB = data[:, CB_atoms, :].numpy()
N_v = N[:, valid, :]
Expand Down Expand Up @@ -607,7 +598,7 @@ def _get_angles(self, data):
"O-C-N": O_C_N_next,
}
if 'CB' in self.atoms:
valid = (self.indices['CB'] > 0)
valid = (self.indices['CB'] >= 0) # 0 is a valid atom index
CB_atoms = self.indices['CB'][valid]
CB = data[:, CB_atoms, :].numpy()
N_v = N[:, valid, :]
Expand Down Expand Up @@ -878,7 +869,7 @@ def _dihedrals(p0, p1, p2, p3):
v /= np.linalg.norm(v, axis=-1, keepdims=True)
w /= np.linalg.norm(w, axis=-1, keepdims=True)
x = np.sum(v * w, axis=-1)
y = np.sum(np.cross(b1, v), axis=-1) * np.sum(w, axis=-1)
y = np.sum(np.cross(b1, v) * w, axis=-1)
return np.arctan2(y, x)

@staticmethod
Expand Down
60 changes: 31 additions & 29 deletions src/molearn/data/pdb_data.py
Original file line number Diff line number Diff line change
Expand Up @@ -81,35 +81,34 @@ def _prepare_coordinates(self) -> np.ndarray:
)

def _compute_backbone_indices(self):
n_indices, ca_indices, cb_indices, c_indices, o_indices = [], [], [], [], []
for i, atom in enumerate(self._mol.atoms):
if atom.name == 'N':
if not len(n_indices) == len(ca_indices) == len(c_indices) == len(o_indices):
raise ValueError("Inconsistent number of N, CA, C, and O atoms in the trajectory.")
if len(cb_indices) < len(n_indices):
cb_indices.append(-1)
n_indices.append(i)
elif atom.name == 'CA':
ca_indices.append(i)
elif atom.name == 'C':
c_indices.append(i)
elif atom.name == 'O':
o_indices.append(i)
elif atom.name == 'CB':
cb_indices.append(i)
else:
raise ValueError(f"Unknown atom name: {atom.name}. Check atom selection.")
if not len(n_indices) == len(ca_indices) == len(c_indices) == len(o_indices):
raise ValueError("Inconsistent number of N, CA, C, and O atoms in the trajectory.")
if len(cb_indices) < len(ca_indices):
cb_indices.append(-1)
self.indices = {
"N": torch.as_tensor(n_indices, dtype=torch.long),
"CA": torch.as_tensor(ca_indices, dtype=torch.long),
"C": torch.as_tensor(c_indices, dtype=torch.long),
"O": torch.as_tensor(o_indices, dtype=torch.long),
"CB": torch.as_tensor(cb_indices, dtype=torch.long),
}
"""Per-residue index of each backbone atom within the selected atoms.

Returns one row per residue for each of N, CA, C, O and CB, holding the
position of that atom in the coordinate array. CB is -1 where the residue
has none (glycine, or a selection that excluded it).
"""
atoms = self._mol.atoms
names = np.asarray(atoms.names)
slot = np.unique(np.asarray(atoms.resindices), return_inverse=True)[1]
n_residues = int(slot.max()) + 1 if len(slot) else 0

indices = {}
for name in ("N", "CA", "C", "O", "CB"):
column = np.full(n_residues, -1, dtype=np.int64)
found = np.flatnonzero(names == name)
column[slot[found]] = found
indices[name] = column

incomplete = [n for n in ("N", "CA", "C", "O") if (indices[n] < 0).any()]
if incomplete:
n_missing = {n: int((indices[n] < 0).sum()) for n in incomplete}
raise ValueError(
f"{n_residues} residues selected but some lack backbone atoms "
f"{n_missing} (atom: number of residues missing it). Check the atom "
f"selection, and that the structure has no incomplete residues."
)

self.indices = {k: torch.as_tensor(v) for k, v in indices.items()}
self.cb_valid_idx = self.indices["CB"][self.indices["CB"] >= 0]

def _standardise_coordinates(self, coords: np.ndarray) -> np.ndarray:
Expand Down Expand Up @@ -256,6 +255,9 @@ def atomselect(self, atoms: str | list[str]):
else:
raise ValueError("Unsupported atom selection")
self._mol.atoms = self._mol.select_atoms(selection_string)
# get_atominfo() caches; the selection has just changed the atom set/order
if hasattr(self, "atominfo"):
del self.atominfo

def prepare_dataset(self, std=None, mean=None) -> torch.Tensor:
"""
Expand Down
Loading