Download model/utils/folding.py from OneScience-Group/IgFold: direct link, hf CLI and curl.
- Browser
- Download file 6.38 kB
-
https://huggingface.co/OneScience-Group/IgFold/resolve/main/model/utils/folding.py
- Command line
-
hf download hf://OneScience-Group/IgFold/model/utils/folding.py
-
curl -L -o folding.py https://huggingface.co/OneScience-Group/IgFold/resolve/main/model/utils/folding.py
6.38 kB
| import os | |
| from typing import List | |
| from einops import rearrange | |
| import torch | |
| import numpy as np | |
| from igfold.model.interface import IgFoldInput | |
| from igfold.utils.fasta import get_fasta_chain_dict | |
| from igfold.utils.general import exists | |
| from igfold.utils.pdb import get_atom_coords, save_PDB, write_pdb_bfactor, cdr_indices | |
| def get_sequence_dict( | |
| sequences, | |
| fasta_file, | |
| ): | |
| if exists(sequences) and exists(fasta_file): | |
| print("Both sequences and fasta file provided. Using fasta file.") | |
| seq_dict = get_fasta_chain_dict(fasta_file) | |
| elif not exists(sequences) and exists(fasta_file): | |
| seq_dict = get_fasta_chain_dict(fasta_file) | |
| elif exists(sequences): | |
| seq_dict = sequences | |
| else: | |
| exit("Must provide sequences or fasta file.") | |
| return seq_dict | |
| def process_template( | |
| pdb_file, | |
| fasta_file, | |
| ignore_cdrs=None, | |
| ignore_chain=None, | |
| ): | |
| temp_coords, temp_mask = None, None | |
| if exists(pdb_file): | |
| temp_coords = get_atom_coords( | |
| pdb_file, | |
| fasta_file=fasta_file, | |
| ) | |
| temp_coords = torch.stack( | |
| [ | |
| temp_coords['N'], temp_coords['CA'], temp_coords['C'], | |
| temp_coords['CB'] | |
| ], | |
| dim=1, | |
| ).view(-1, 3).unsqueeze(0) | |
| temp_mask = torch.ones(temp_coords.shape[:2]).bool() | |
| temp_mask[temp_coords.isnan().any(-1)] = False | |
| temp_mask[temp_coords.sum(-1) == 0] = False | |
| if exists(ignore_cdrs): | |
| cdr_names = ["h1", "h2", "h3", "l1", "l2", "l3"] | |
| if ignore_cdrs == False: | |
| cdr_names = [] | |
| elif isinstance(ignore_cdrs, list): | |
| cdr_names = ignore_cdrs | |
| elif isinstance(ignore_cdrs, str): | |
| cdr_names = [ignore_cdrs] | |
| for cdr in cdr_names: | |
| cdr_range = cdr_indices(pdb_file, cdr) | |
| temp_mask[:, (cdr_range[0] - 1) * 4:(cdr_range[1] + 2) * | |
| 4] = False | |
| if exists(ignore_chain) and ignore_chain in ["H", "L"]: | |
| seq_dict = get_fasta_chain_dict(fasta_file) | |
| hlen = len(seq_dict["H"]) | |
| if ignore_chain == "H": | |
| temp_mask[:, :hlen * 4] = False | |
| elif ignore_chain == "L": | |
| temp_mask[:, hlen * 4:] = False | |
| return temp_coords, temp_mask | |
| def process_prediction( | |
| model_out, | |
| pdb_file, | |
| fasta_file, | |
| skip_pdb=False, | |
| do_refine=True, | |
| use_openmm=False, | |
| do_renum=False, | |
| ): | |
| prmsd = rearrange( | |
| model_out.prmsd, | |
| "b (l a) -> b l a", | |
| a=4, | |
| ) | |
| model_out.prmsd = prmsd | |
| if skip_pdb: | |
| return model_out | |
| coords = model_out.coords.squeeze(0).detach() | |
| res_rmsd = prmsd.square().mean(dim=-1).sqrt().squeeze(0) | |
| seq_dict = get_fasta_chain_dict(fasta_file) | |
| full_seq = "".join(list(seq_dict.values())) | |
| chains = list(seq_dict.keys()) | |
| delims = np.cumsum([len(s) for s in seq_dict.values()]).tolist() | |
| write_pdb = not do_refine or use_openmm | |
| pdb_string = save_PDB( | |
| pdb_file, | |
| coords, | |
| full_seq, | |
| chains=chains, | |
| atoms=['N', 'CA', 'C', 'CB', 'O'], | |
| error=res_rmsd, | |
| delim=delims, | |
| write_pdb=write_pdb, | |
| ) | |
| if do_refine: | |
| if use_openmm: | |
| try: | |
| from igfold.refine.openmm_ref import refine | |
| refine_input = [pdb_file] | |
| except: | |
| exit("OpenMM not installed. Please install OpenMM to use refinement.") | |
| else: | |
| try: | |
| from igfold.refine.pyrosetta_ref import refine | |
| refine_input = [pdb_file, pdb_string] | |
| except: | |
| exit("PyRosetta not installed. Please install PyRosetta to use refinement.") | |
| refine(*refine_input) | |
| if do_renum: | |
| try: | |
| from igfold.utils.abnumber_ import renumber_pdb | |
| except: | |
| exit("AbNumber not installed. Please install AbNumber to use renumbering.") | |
| renumber_pdb( | |
| pdb_file, | |
| pdb_file, | |
| ) | |
| write_pdb_bfactor( | |
| pdb_file, | |
| pdb_file, | |
| bfactor=res_rmsd, | |
| ) | |
| return model_out | |
| def fold( | |
| antiberty, | |
| models, | |
| pdb_file, | |
| fasta_file=None, | |
| sequences=None, | |
| template_pdb=None, | |
| ignore_cdrs=None, | |
| ignore_chain=None, | |
| skip_pdb=False, | |
| do_refine=True, | |
| use_openmm=False, | |
| do_renum=True, | |
| truncate_sequences=False, | |
| ): | |
| seq_dict = get_sequence_dict( | |
| sequences, | |
| fasta_file, | |
| ) | |
| if truncate_sequences: | |
| try: | |
| from igfold.utils.abnumber_ import truncate_seq | |
| except: | |
| exit("AbNumber not installed. Please install AbNumber to use truncation.") | |
| seq_dict = {k: truncate_seq(v) for k, v in seq_dict.items()} | |
| if not exists(fasta_file): | |
| fasta_file = pdb_file.replace(".pdb", ".fasta") | |
| with open(fasta_file, "w") as f: | |
| for chain, seq in seq_dict.items(): | |
| f.write(">{}\n{}\n".format( | |
| chain, | |
| seq, | |
| )) | |
| embeddings, attentions = antiberty.embed( | |
| seq_dict.values(), | |
| return_attention=True, | |
| ) | |
| embeddings = [e[1:-1].unsqueeze(0) for e in embeddings] | |
| attentions = [a[:, :, 1:-1, 1:-1].unsqueeze(0) for a in attentions] | |
| temp_coords, temp_mask = process_template( | |
| template_pdb, | |
| fasta_file, | |
| ignore_cdrs=ignore_cdrs, | |
| ignore_chain=ignore_chain, | |
| ) | |
| model_in = IgFoldInput( | |
| embeddings=embeddings, | |
| attentions=attentions, | |
| template_coords=temp_coords, | |
| template_mask=temp_mask, | |
| return_embeddings=True, | |
| ) | |
| model_outs, scores = [], [] | |
| with torch.no_grad(): | |
| for i, model in enumerate(models): | |
| model_out = model(model_in) | |
| model_out = model.gradient_refine(model_in, model_out) | |
| scores.append(model_out.prmsd.quantile(0.9)) | |
| model_outs.append(model_out) | |
| best_model_i = scores.index(min(scores)) | |
| model_out = model_outs[best_model_i] | |
| process_prediction( | |
| model_out, | |
| pdb_file, | |
| fasta_file, | |
| skip_pdb=skip_pdb, | |
| do_refine=do_refine, | |
| use_openmm=use_openmm, | |
| do_renum=do_renum, | |
| ) | |
| return model_out | |