Synthyra/ESMFold2
0498
1from itertools import groupby
2from typing import Any
3
4import numpy as np
5import torch
6
7from .esmfold2_constants import ELEMENT_NUMBER_TO_SYMBOL, MOL_TYPE_NONPOLYMER
8from .esmfold2_molecular_complex import (
9 MolecularComplex,
10 MolecularComplexMetadata,
11)
12
13
14def get_element_symbol(atomic_num: int) -> str:
15 return ELEMENT_NUMBER_TO_SYMBOL.get(atomic_num, "X")
16
17
18def build_molecular_complex_from_features(
19 coords: torch.Tensor,
20 plddt: torch.Tensor,
21 atom_mask: torch.Tensor,
22 ref_element: torch.Tensor,
23 ref_atom_name_chars: torch.Tensor,
24 chain_infos: list,
25 complex_id: str,
26) -> MolecularComplex:
27 """Construct a MolecularComplex from feature-dict tensors and chain metadata.
28
29 Non-polymer chains (ligands) collapse all per-atom tokens into a single
30 residue token whose pLDDT is the per-token average and whose hetero flag
31 is True.
32 """
33 mask_np = atom_mask.bool().cpu().numpy()
34 coords_np = coords.float().cpu().numpy()
35 name_chars_np = ref_atom_name_chars.cpu().numpy()
36 elements_np = ref_element.cpu().numpy()
37 plddt_np = plddt.float().cpu().numpy()
38
39 sequence_tokens: list[str] = []
40 chain_ids_per_token: list[int] = []
41 token_to_atoms: list[list[int]] = []
42 confidence: list[float] = []
43 flat_positions: list[list[float]] = []
44 flat_elements: list[str] = []
45 flat_names: list[str] = []
46 flat_hetero: list[bool] = []
47
48 chain_lookup: dict[int, str] = {}
49 entity_info: dict[int, str] = {}
50 out_atom_cursor = 0
51
52 for ci in chain_infos:
53 chain_lookup[ci.asym_id] = ci.chain_id
54 is_nonpolymer = ci.mol_type == MOL_TYPE_NONPOLYMER
55 entity_info[ci.entity_id] = "non-polymer" if is_nonpolymer else "polymer"
56
57 if is_nonpolymer:
58 residue_name = ci.tokens[0].residue_name if ci.tokens else "LIG"
59 sequence_tokens.append(residue_name)
60 chain_ids_per_token.append(ci.asym_id)
61 avg_plddt = (
62 float(np.mean([plddt_np[ti.token_index] for ti in ci.tokens]))
63 if ci.tokens
64 else 0.0
65 )
66 confidence.append(avg_plddt)
67 token_atom_start = out_atom_cursor
68 for ti in ci.tokens:
69 for atom_idx in range(ti.atom_start, ti.atom_start + ti.atom_count):
70 if not mask_np[atom_idx]:
71 continue
72 flat_positions.append(coords_np[atom_idx].tolist())
73 flat_elements.append(get_element_symbol(int(elements_np[atom_idx])))
74 chars = name_chars_np[atom_idx]
75 name = "".join(
76 chr(int(c) + 32) for c in chars if int(c) != 0
77 ).strip()
78 flat_names.append(name)
79 flat_hetero.append(True)
80 out_atom_cursor += 1
81 token_to_atoms.append([token_atom_start, out_atom_cursor])
82 continue
83
84 # Atom-tokenized modified residues (HYP, MSE, ...) span multiple
85 # tokens per residue; collapse them back to one mmCIF residue.
86 for _residue_index, ti_iter in groupby(
87 ci.tokens, key=lambda t: t.residue_index
88 ):
89 ti_group = list(ti_iter)
90 sequence_tokens.append(ti_group[0].residue_name)
91 chain_ids_per_token.append(ci.asym_id)
92 confidence.append(
93 float(np.mean([plddt_np[ti.token_index] for ti in ti_group]))
94 )
95 token_atom_start = out_atom_cursor
96 for ti in ti_group:
97 for atom_idx in range(ti.atom_start, ti.atom_start + ti.atom_count):
98 if not mask_np[atom_idx]:
99 continue
100 flat_positions.append(coords_np[atom_idx].tolist())
101 flat_elements.append(get_element_symbol(int(elements_np[atom_idx])))
102 chars = name_chars_np[atom_idx]
103 name = "".join(
104 chr(int(c) + 32) for c in chars if int(c) != 0
105 ).strip()
106 flat_names.append(name)
107 flat_hetero.append(False)
108 out_atom_cursor += 1
109 token_to_atoms.append([token_atom_start, out_atom_cursor])
110
111 return MolecularComplex(
112 id=complex_id,
113 sequence=sequence_tokens,
114 atom_positions=np.array(flat_positions, dtype=np.float32).reshape(-1, 3),
115 atom_elements=np.array(flat_elements, dtype=object),
116 token_to_atoms=np.array(token_to_atoms, dtype=np.int32).reshape(-1, 2),
117 chain_id=np.array(chain_ids_per_token, dtype=np.int64),
118 plddt=np.array(confidence, dtype=np.float32),
119 atom_names=np.array(flat_names, dtype=object),
120 atom_hetero=np.array(flat_hetero, dtype=bool),
121 metadata=MolecularComplexMetadata(
122 entity_lookup=entity_info,
123 chain_lookup=chain_lookup,
124 assembly_composition=None,
125 ),
126 )
127
128
129def build_molecular_complex(
130 structure: Any, coords: torch.Tensor, plddt: torch.Tensor, complex_id: str
131) -> MolecularComplex:
132 """Directly constructs a MolecularComplex from model outputs without intermediate files.
133
134 Args:
135 structure: Object with .chains, .residues, .atoms numpy structured arrays.
136 coords: [N_atoms, 3] predicted atom coordinates.
137 plddt: [N_residues] per-residue confidence scores.
138 complex_id: Identifier string for the resulting complex.
139 """
140 flat_positions = []
141 flat_elements = []
142 flat_names = []
143 flat_hetero = []
144
145 sequence_tokens = []
146 token_to_atoms = []
147 chain_ids_per_token = []
148 confidence_scores = []
149
150 chain_lookup = {}
151 entity_info = {}
152
153 global_atom_cursor = 0
154 global_res_cursor = 0
155 atom_array_idx = 0
156
157 for chain in structure.chains:
158 chain_idx_numeric = chain["asym_id"]
159 chain_name_str = str(chain["name"])
160 mol_type = chain["mol_type"]
161
162 chain_lookup[chain_idx_numeric] = chain_name_str
163 entity_info[chain["entity_id"]] = (
164 "polymer" if mol_type != MOL_TYPE_NONPOLYMER else "non-polymer"
165 )
166
167 res_start = chain["res_idx"]
168 res_end = chain["res_idx"] + chain["res_num"]
169 residues = structure.residues[res_start:res_end]
170
171 for residue in residues:
172 res_name = str(residue["name"])
173
174 sequence_tokens.append(res_name)
175 chain_ids_per_token.append(chain_idx_numeric)
176
177 score = plddt[global_res_cursor].item()
178 confidence_scores.append(score)
179 token_start_idx = atom_array_idx
180
181 atom_start = residue["atom_idx"]
182 atom_end = residue["atom_idx"] + residue["atom_num"]
183 atoms = structure.atoms[atom_start:atom_end]
184
185 for atom in atoms:
186 if not atom["is_present"]:
187 continue
188
189 pos = coords[global_atom_cursor].tolist()
190 flat_positions.append(pos)
191
192 elem = get_element_symbol(atom["element"].item())
193 flat_elements.append(elem)
194
195 raw_name = atom["name"]
196 if hasattr(raw_name, "tolist"):
197 raw_name = raw_name.tolist()
198 name_str = "".join([chr(c + 32) for c in raw_name if c != 0])
199 flat_names.append(name_str)
200
201 flat_hetero.append(mol_type == MOL_TYPE_NONPOLYMER)
202
203 global_atom_cursor += 1
204 atom_array_idx += 1
205
206 token_to_atoms.append([token_start_idx, atom_array_idx])
207 global_res_cursor += 1
208
209 return MolecularComplex(
210 id=complex_id,
211 sequence=sequence_tokens,
212 atom_positions=np.array(flat_positions, dtype=np.float32),
213 atom_elements=np.array(flat_elements, dtype=object),
214 token_to_atoms=np.array(token_to_atoms, dtype=np.int32),
215 chain_id=np.array(chain_ids_per_token, dtype=np.int64),
216 plddt=np.array(confidence_scores, dtype=np.float32),
217 atom_names=np.array(flat_names, dtype=object),
218 atom_hetero=np.array(flat_hetero, dtype=bool),
219 metadata=MolecularComplexMetadata(
220 entity_lookup=entity_info,
221 chain_lookup=chain_lookup,
222 assembly_composition=None,
223 ),
224 )
225 