torch-structure-manipulation 0.6.0__py3-none-any.whl

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
@@ -0,0 +1,54 @@
1
+ """Atomic structure data, bonding annotations, and structure transforms."""
2
+
3
+ from importlib.metadata import PackageNotFoundError, version
4
+
5
+ from .atomic_structure import AtomicStructure
6
+ from .bonding import (
7
+ annotate_bonding_environments,
8
+ classify_structure_composition,
9
+ get_scattering_provider_keys,
10
+ )
11
+ from .structure_transforms import (
12
+ apply_rotation,
13
+ apply_rotation_to_coords,
14
+ apply_translation,
15
+ apply_translation_to_coords,
16
+ ball_query_atoms,
17
+ calculate_center_from_tensors,
18
+ center_structure,
19
+ center_structure_from_coords,
20
+ df_to_atomxyz,
21
+ df_to_atomzyx,
22
+ find_atoms_in_ball,
23
+ get_nucleic_acid_residues,
24
+ get_protein_residues,
25
+ remove_sidechains,
26
+ separate_protein_rna,
27
+ )
28
+
29
+ try:
30
+ __version__ = version("torch-structure-manipulation")
31
+ except PackageNotFoundError:
32
+ __version__ = "uninstalled"
33
+
34
+ __all__ = [
35
+ "AtomicStructure",
36
+ "annotate_bonding_environments",
37
+ "apply_rotation",
38
+ "apply_rotation_to_coords",
39
+ "apply_translation",
40
+ "apply_translation_to_coords",
41
+ "ball_query_atoms",
42
+ "calculate_center_from_tensors",
43
+ "center_structure",
44
+ "center_structure_from_coords",
45
+ "classify_structure_composition",
46
+ "df_to_atomxyz",
47
+ "df_to_atomzyx",
48
+ "find_atoms_in_ball",
49
+ "get_nucleic_acid_residues",
50
+ "get_protein_residues",
51
+ "get_scattering_provider_keys",
52
+ "remove_sidechains",
53
+ "separate_protein_rna",
54
+ ]
@@ -0,0 +1,228 @@
1
+ """Tensor-backed atomic structure data."""
2
+
3
+ from __future__ import annotations
4
+
5
+ from dataclasses import dataclass, replace
6
+ from typing import TYPE_CHECKING
7
+
8
+ import gemmi
9
+ import torch
10
+
11
+ if TYPE_CHECKING:
12
+ import pandas as pd
13
+
14
+
15
+ @dataclass(frozen=True, slots=True)
16
+ class AtomicStructure:
17
+ """Lightweight numerical representation of an atomic structure.
18
+
19
+ Positions use array-friendly ``(z, y, x)`` order. Numerical fields may have
20
+ arbitrary broadcast-compatible batch dimensions before their atom dimension.
21
+ Text metadata is shared across batches and held in immutable tuples.
22
+
23
+ Bonding metadata (``bonded_environments``, ``molecule_types``) is not batched:
24
+ one tuple is shared by every batch member. That matches ensemble-of-poses use
25
+ cases (same chemistry, different coordinates) but not batched structures with
26
+ different chemistry. See :meth:`from_annotated_dataframe` and the bonded-factor
27
+ notes on
28
+ :func:`torch_calculate_electrostatic_potential.potential_from_structure_3d`.
29
+ """
30
+
31
+ positions_zyx: torch.Tensor
32
+ atomic_numbers: torch.Tensor
33
+ elements: tuple[str, ...]
34
+ atom_names: tuple[str, ...]
35
+ b_factors: torch.Tensor
36
+ occupancies: torch.Tensor
37
+ bonded_environments: tuple[str, ...] | None = None
38
+ molecule_types: tuple[str, ...] | None = None
39
+
40
+ def __post_init__(self) -> None:
41
+ """Validate atom dimensions and batch broadcasting."""
42
+ if self.positions_zyx.ndim < 2 or self.positions_zyx.shape[-1] != 3:
43
+ raise ValueError("positions_zyx must have shape (..., n_atoms, 3)")
44
+ n_atoms = self.positions_zyx.shape[-2]
45
+ if self.atomic_numbers.ndim < 1 or self.atomic_numbers.shape[-1] != n_atoms:
46
+ raise ValueError(
47
+ "atomic_numbers must have shape (..., n_atoms) with the same "
48
+ "number of atoms as positions_zyx"
49
+ )
50
+ if len(self.elements) != n_atoms or len(self.atom_names) != n_atoms:
51
+ raise ValueError("text metadata must have one value per atom")
52
+ numerical_fields = {
53
+ "b_factors": self.b_factors,
54
+ "occupancies": self.occupancies,
55
+ }
56
+ for name, value in numerical_fields.items():
57
+ if value.ndim > 0 and value.shape[-1] != n_atoms:
58
+ raise ValueError(f"{name} must be scalar or have shape (..., n_atoms)")
59
+ batch_shapes = [
60
+ self.positions_zyx.shape[:-2],
61
+ self.atomic_numbers.shape[:-1],
62
+ *[
63
+ value.shape[:-1] if value.ndim > 0 else ()
64
+ for value in numerical_fields.values()
65
+ ],
66
+ ]
67
+ try:
68
+ torch.broadcast_shapes(*batch_shapes) # type: ignore[no-untyped-call]
69
+ except RuntimeError as error:
70
+ raise ValueError(
71
+ "numerical AtomicStructure fields have incompatible batch shapes"
72
+ ) from error
73
+ if (
74
+ self.bonded_environments is not None
75
+ and len(self.bonded_environments) != n_atoms
76
+ ):
77
+ raise ValueError("bonded_environments must have one value per atom")
78
+ if self.molecule_types is not None and len(self.molecule_types) != n_atoms:
79
+ raise ValueError("molecule_types must have one value per atom")
80
+
81
+ @classmethod
82
+ def from_dataframe(
83
+ cls,
84
+ df: pd.DataFrame,
85
+ *,
86
+ device: torch.device | str | None = None,
87
+ dtype: torch.dtype = torch.float32,
88
+ ) -> AtomicStructure:
89
+ """Construct from an mmdf-compatible DataFrame.
90
+
91
+ Required columns are ``x``, ``y``, ``z``, and ``element``. Atom names
92
+ come from ``atom`` when present. ``b_isotropic`` and ``occupancy``
93
+ default to zero and one, respectively. Optional ``bonded_environments``
94
+ and ``molecule_type`` columns are preserved when present.
95
+ """
96
+ required = {"x", "y", "z", "element"}
97
+ missing = sorted(required.difference(df.columns))
98
+ if missing:
99
+ raise ValueError(f"missing required structure columns: {missing}")
100
+
101
+ elements = tuple(str(value).strip().upper() for value in df["element"])
102
+ if "atomic_number" in df:
103
+ atomic_number_values = [int(value) for value in df["atomic_number"]]
104
+ else:
105
+ atomic_number_values = [
106
+ gemmi.Element(element).atomic_number for element in elements
107
+ ]
108
+ unknown = sorted(
109
+ element
110
+ for element, atomic_number in zip(
111
+ elements, atomic_number_values, strict=True
112
+ )
113
+ if atomic_number == 0
114
+ )
115
+ if unknown:
116
+ raise ValueError(f"unknown element symbols: {unknown}")
117
+
118
+ positions = torch.as_tensor(
119
+ df.loc[:, ["z", "y", "x"]].to_numpy(copy=True),
120
+ dtype=dtype,
121
+ device=device,
122
+ )
123
+ atomic_numbers = torch.tensor(
124
+ atomic_number_values,
125
+ dtype=torch.int64,
126
+ device=device,
127
+ )
128
+ b_values = df["b_isotropic"] if "b_isotropic" in df else [0.0] * len(df)
129
+ occupancy_values = df["occupancy"] if "occupancy" in df else [1.0] * len(df)
130
+ atom_names = (
131
+ tuple(str(value).strip() for value in df["atom"])
132
+ if "atom" in df
133
+ else ("",) * len(df)
134
+ )
135
+ bonded = (
136
+ tuple(str(value) for value in df["bonded_environments"])
137
+ if "bonded_environments" in df
138
+ else None
139
+ )
140
+ molecule_types = (
141
+ tuple(str(value) for value in df["molecule_type"])
142
+ if "molecule_type" in df
143
+ else None
144
+ )
145
+ return cls(
146
+ positions_zyx=positions,
147
+ atomic_numbers=atomic_numbers,
148
+ elements=elements,
149
+ atom_names=atom_names,
150
+ b_factors=torch.as_tensor(b_values, dtype=dtype, device=device),
151
+ occupancies=torch.as_tensor(occupancy_values, dtype=dtype, device=device),
152
+ bonded_environments=bonded,
153
+ molecule_types=molecule_types,
154
+ )
155
+
156
+ @classmethod
157
+ def from_annotated_dataframe(
158
+ cls,
159
+ df: pd.DataFrame,
160
+ *,
161
+ include_hydrogens: bool = True,
162
+ device: torch.device | str | None = None,
163
+ dtype: torch.dtype = torch.float32,
164
+ ) -> AtomicStructure:
165
+ """Annotate bonding metadata, then construct from the result.
166
+
167
+ This is the usual entry point for Peng bonded scattering factors: it
168
+ calls :func:`~torch_structure_manipulation.annotate_bonding_environments`
169
+ to add ``bonded_environments`` and ``molecule_type`` columns, then
170
+ delegates to :meth:`from_dataframe`.
171
+
172
+ The input must include ``chain``, ``residue_id``, ``residue``, ``atom``,
173
+ and ``element`` in addition to the coordinate columns required by
174
+ :meth:`from_dataframe`.
175
+ """
176
+ from .bonding import annotate_bonding_environments
177
+
178
+ annotated = annotate_bonding_environments(
179
+ df, include_hydrogens=include_hydrogens
180
+ )
181
+ return cls.from_dataframe(annotated, device=device, dtype=dtype)
182
+
183
+ @property
184
+ def num_atoms(self) -> int:
185
+ """Number of atoms in each structure."""
186
+ return self.positions_zyx.shape[-2]
187
+
188
+ @property
189
+ def batch_shape(self) -> torch.Size:
190
+ """Broadcasted batch shape of all numerical fields."""
191
+ batch_shapes = [
192
+ self.positions_zyx.shape[:-2],
193
+ self.atomic_numbers.shape[:-1],
194
+ self.b_factors.shape[:-1] if self.b_factors.ndim > 0 else (),
195
+ self.occupancies.shape[:-1] if self.occupancies.ndim > 0 else (),
196
+ ]
197
+ return torch.Size(
198
+ torch.broadcast_shapes(*batch_shapes) # type: ignore[no-untyped-call]
199
+ )
200
+
201
+ @property
202
+ def device(self) -> torch.device:
203
+ """Device containing the atomic positions."""
204
+ return self.positions_zyx.device
205
+
206
+ def with_positions(self, positions_zyx: torch.Tensor) -> AtomicStructure:
207
+ """Return a copy with replacement, broadcast-compatible positions."""
208
+ if positions_zyx.ndim < 2 or positions_zyx.shape[-2:] != (self.num_atoms, 3):
209
+ raise ValueError(
210
+ "replacement positions must have shape (..., n_atoms, 3) with the "
211
+ "same number of atoms"
212
+ )
213
+ return replace(self, positions_zyx=positions_zyx)
214
+
215
+ def to(
216
+ self,
217
+ device: torch.device | str | None = None,
218
+ dtype: torch.dtype | None = None,
219
+ ) -> AtomicStructure:
220
+ """Return a copy with numerical tensors moved to a device and dtype."""
221
+ floating_dtype = self.positions_zyx.dtype if dtype is None else dtype
222
+ return replace(
223
+ self,
224
+ positions_zyx=self.positions_zyx.to(device=device, dtype=floating_dtype),
225
+ atomic_numbers=self.atomic_numbers.to(device=device),
226
+ b_factors=self.b_factors.to(device=device, dtype=floating_dtype),
227
+ occupancies=self.occupancies.to(device=device, dtype=floating_dtype),
228
+ )
@@ -0,0 +1,260 @@
1
+ """Annotate atomic bonding environments from packaged residue templates."""
2
+
3
+ from __future__ import annotations
4
+
5
+ import json
6
+ from importlib.resources import files
7
+ from itertools import pairwise
8
+ from typing import Any
9
+
10
+ import pandas as pd
11
+
12
+ from .structure_transforms import (
13
+ get_nucleic_acid_residues,
14
+ get_protein_residues,
15
+ )
16
+
17
+ _REQUIRED_COLUMNS = {"atom", "chain", "element", "residue", "residue_id"}
18
+ NUCLEIC_ACID_RESIDUES = get_nucleic_acid_residues()
19
+ PROTEIN_RESIDUES = get_protein_residues()
20
+
21
+
22
+ def _load_templates() -> tuple[
23
+ dict[str, dict[str, list[str]]], dict[str, dict[str, list[str]]]
24
+ ]:
25
+ resource = files(__package__).joinpath("bonding_data.json")
26
+ data: dict[str, Any] = json.loads(resource.read_text(encoding="utf-8"))
27
+ return data["protein"], data["rna"]
28
+
29
+
30
+ _PROTEIN_BONDING, _RNA_BONDING = _load_templates()
31
+
32
+
33
+ def annotate_bonding_environments(
34
+ df: pd.DataFrame, include_hydrogens: bool = True
35
+ ) -> pd.DataFrame:
36
+ """Return a copy with ``bonded_environments`` and per-atom ``molecule_type``.
37
+
38
+ The input follows mmdf naming conventions but is not tied to mmdf itself.
39
+ Bonding templates are keyed by residue and atom names; adjacent numeric
40
+ residue IDs are used to add peptide and phosphodiester bonds.
41
+ """
42
+ missing = sorted(_REQUIRED_COLUMNS.difference(df.columns))
43
+ if missing:
44
+ raise ValueError(f"missing required structure columns: {missing}")
45
+
46
+ result = df.copy()
47
+ if result.empty:
48
+ result["bonded_environments"] = pd.Series(dtype="object", index=result.index)
49
+ result["molecule_type"] = pd.Series(dtype="object", index=result.index)
50
+ return result
51
+
52
+ residues = [str(value).strip().upper() for value in result["residue"]]
53
+ atoms = [str(value).strip().upper() for value in result["atom"]]
54
+ elements = [str(value).strip().upper() for value in result["element"]]
55
+ residue_ids = [str(value) for value in result["residue_id"]]
56
+ chains = [str(value) for value in result["chain"]]
57
+
58
+ residue_lookup: dict[tuple[str, str], dict[str, str]] = {}
59
+ residue_names: dict[tuple[str, str], str] = {}
60
+ for residue, atom, element, residue_id, chain in zip(
61
+ residues, atoms, elements, residue_ids, chains, strict=True
62
+ ):
63
+ key = (chain, residue_id)
64
+ residue_lookup.setdefault(key, {})[atom] = element
65
+ residue_names.setdefault(key, residue)
66
+
67
+ next_residue, previous_residue = _build_residue_order(residue_names)
68
+ molecule_types = [_molecule_type(residue) for residue in residues]
69
+ environments = [
70
+ _environment_for_atom(
71
+ residue=residue,
72
+ atom=atom,
73
+ element=element,
74
+ key=(chain, residue_id),
75
+ residue_lookup=residue_lookup,
76
+ next_residue=next_residue,
77
+ previous_residue=previous_residue,
78
+ include_hydrogens=include_hydrogens,
79
+ )
80
+ for residue, atom, element, residue_id, chain in zip(
81
+ residues, atoms, elements, residue_ids, chains, strict=True
82
+ )
83
+ ]
84
+ result["bonded_environments"] = environments
85
+ result["molecule_type"] = molecule_types
86
+ return result
87
+
88
+
89
+ def classify_structure_composition(df: pd.DataFrame) -> str:
90
+ """Summarize residue composition for the whole table.
91
+
92
+ Returns an aggregate label such as ``"protein"``, ``"rna"``,
93
+ ``"rna+protein"``, or ``"other"``. This is descriptive metadata only and
94
+ is not a valid per-atom scattering-provider key.
95
+ """
96
+ residues = {str(value).strip().upper() for value in df["residue"]}
97
+ has_protein = bool(residues & PROTEIN_RESIDUES)
98
+ has_rna = bool(residues & NUCLEIC_ACID_RESIDUES)
99
+ if has_protein and has_rna:
100
+ return "rna+protein"
101
+ if has_protein:
102
+ return "protein"
103
+ if has_rna:
104
+ return "rna"
105
+ return "other"
106
+
107
+
108
+ def get_scattering_provider_keys(df: pd.DataFrame) -> list[str]:
109
+ """Return the Peng scattering-provider key for each atom.
110
+
111
+ Each value is ``"protein"``, ``"rna"``, or ``"other"`` and matches the
112
+ ``molecule_type`` column written by :func:`annotate_bonding_environments`.
113
+ """
114
+ return [_molecule_type(str(residue).strip().upper()) for residue in df["residue"]]
115
+
116
+
117
+ def _environment_for_atom(
118
+ *,
119
+ residue: str,
120
+ atom: str,
121
+ element: str,
122
+ key: tuple[str, str],
123
+ residue_lookup: dict[tuple[str, str], dict[str, str]],
124
+ next_residue: dict[tuple[str, str], tuple[str, str]],
125
+ previous_residue: dict[tuple[str, str], tuple[str, str]],
126
+ include_hydrogens: bool,
127
+ ) -> str:
128
+ template = _template_for(residue)
129
+ bonded_names = list(template.get(atom, ()))
130
+ if residue in _RNA_BONDING and atom == "O3'" and key in next_residue:
131
+ bonded_names = [name for name in bonded_names if name != "HO3'"]
132
+
133
+ bonded_elements: list[str] = []
134
+ residue_atoms = residue_lookup.get(key, {})
135
+ for bonded_name in bonded_names:
136
+ bonded_element = _find_element(bonded_name, residue_atoms)
137
+ if bonded_element:
138
+ if include_hydrogens or bonded_element != "H":
139
+ bonded_elements.append(bonded_element)
140
+ elif include_hydrogens and _is_hydrogen_name(bonded_name):
141
+ bonded_elements.append("H")
142
+
143
+ if residue in _PROTEIN_BONDING and atom == "C" and key in next_residue:
144
+ _append_atom_element(
145
+ bonded_elements, "N", residue_lookup[next_residue[key]], include_hydrogens
146
+ )
147
+ elif residue in _PROTEIN_BONDING and atom == "N" and key in previous_residue:
148
+ _append_atom_element(
149
+ bonded_elements,
150
+ "C",
151
+ residue_lookup[previous_residue[key]],
152
+ include_hydrogens,
153
+ )
154
+ elif residue in _RNA_BONDING and atom == "O3'" and key in next_residue:
155
+ _append_atom_element(
156
+ bonded_elements, "P", residue_lookup[next_residue[key]], include_hydrogens
157
+ )
158
+ elif residue in _RNA_BONDING and atom == "P" and key in previous_residue:
159
+ _append_atom_element(
160
+ bonded_elements,
161
+ "O3'",
162
+ residue_lookup[previous_residue[key]],
163
+ include_hydrogens,
164
+ )
165
+
166
+ bonded_key = "".join(sorted(bonded_elements))
167
+ category = _oxygen_carbon_category(
168
+ residue, atom, element, bonded_key, key, residue_lookup, next_residue
169
+ )
170
+ suffix = f", {category}" if category is not None else ""
171
+ return f"{element}({bonded_key}{suffix})"
172
+
173
+
174
+ def _oxygen_carbon_category(
175
+ residue: str,
176
+ atom: str,
177
+ element: str,
178
+ bonded_key: str,
179
+ key: tuple[str, str],
180
+ residue_lookup: dict[tuple[str, str], dict[str, str]],
181
+ next_residue: dict[tuple[str, str], tuple[str, str]],
182
+ ) -> str | None:
183
+ if element != "O" or bonded_key != "C":
184
+ return None
185
+ if atom == "OXT" or (residue == "ASP" and atom in {"OD1", "OD2"}):
186
+ return "carboxyl"
187
+ if residue == "GLU" and atom in {"OE1", "OE2"}:
188
+ return "carboxyl"
189
+ if (
190
+ residue in _PROTEIN_BONDING
191
+ and atom == "O"
192
+ and key in next_residue
193
+ and _find_element("N", residue_lookup[next_residue[key]])
194
+ ):
195
+ return "amide"
196
+ return None
197
+
198
+
199
+ def _build_residue_order(
200
+ residue_names: dict[tuple[str, str], str],
201
+ ) -> tuple[
202
+ dict[tuple[str, str], tuple[str, str]],
203
+ dict[tuple[str, str], tuple[str, str]],
204
+ ]:
205
+ by_chain: dict[str, list[tuple[float, tuple[str, str]]]] = {}
206
+ for key in residue_names:
207
+ try:
208
+ numeric_id = float(key[1])
209
+ except ValueError:
210
+ continue
211
+ by_chain.setdefault(key[0], []).append((numeric_id, key))
212
+
213
+ next_residue: dict[tuple[str, str], tuple[str, str]] = {}
214
+ previous_residue: dict[tuple[str, str], tuple[str, str]] = {}
215
+ for residues in by_chain.values():
216
+ residues.sort(key=lambda item: item[0])
217
+ for (left_id, left_key), (right_id, right_key) in pairwise(residues):
218
+ if abs(right_id - left_id - 1.0) < 1e-6:
219
+ next_residue[left_key] = right_key
220
+ previous_residue[right_key] = left_key
221
+ return next_residue, previous_residue
222
+
223
+
224
+ def _template_for(residue: str) -> dict[str, list[str]]:
225
+ if residue in _PROTEIN_BONDING:
226
+ return _PROTEIN_BONDING[residue]
227
+ return _RNA_BONDING.get(residue, {})
228
+
229
+
230
+ def _molecule_type(residue: str) -> str:
231
+ if residue in PROTEIN_RESIDUES:
232
+ return "protein"
233
+ if residue in NUCLEIC_ACID_RESIDUES:
234
+ return "rna"
235
+ return "other"
236
+
237
+
238
+ def _append_atom_element(
239
+ elements: list[str],
240
+ atom_name: str,
241
+ residue_atoms: dict[str, str],
242
+ include_hydrogens: bool,
243
+ ) -> None:
244
+ element = _find_element(atom_name, residue_atoms)
245
+ if element and (include_hydrogens or element != "H"):
246
+ elements.append(element)
247
+
248
+
249
+ def _find_element(atom_name: str, residue_atoms: dict[str, str]) -> str:
250
+ if atom_name in residue_atoms:
251
+ return residue_atoms[atom_name]
252
+ normalized = atom_name.replace("'", "").replace("*", "")
253
+ for candidate, element in residue_atoms.items():
254
+ if candidate.replace("'", "").replace("*", "") == normalized:
255
+ return element
256
+ return ""
257
+
258
+
259
+ def _is_hydrogen_name(atom_name: str) -> bool:
260
+ return atom_name.upper().startswith("H")