CoolFace
Modelpublic

Synthyra/ESMFold2

sourceHugging Facemitupdated 3d agoView on Hugging Face
0likes498downloads
esmfold2_output.py225 linesDownload Raw Back to root
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