diff --git a/.github/workflows/tests.yml b/.github/workflows/tests.yml index f164097..616f746 100644 --- a/.github/workflows/tests.yml +++ b/.github/workflows/tests.yml @@ -5,6 +5,9 @@ on: [push] jobs: build-linux: runs-on: ubuntu-latest + defaults: + run: + shell: bash -l {0} steps: - name: Checkout repository diff --git a/src/molearn/analysis/analyser.py b/src/molearn/analysis/analyser.py index c9058c8..12e2b4c 100644 --- a/src/molearn/analysis/analyser.py +++ b/src/molearn/analysis/analyser.py @@ -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: @@ -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) @@ -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(): @@ -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, :] @@ -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, :] @@ -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 diff --git a/src/molearn/data/pdb_data.py b/src/molearn/data/pdb_data.py index f632266..b5f3e29 100644 --- a/src/molearn/data/pdb_data.py +++ b/src/molearn/data/pdb_data.py @@ -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: @@ -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: """