| import os |
| import polars as pl |
| from datasets import load_dataset |
| import requests |
| from tqdm import tqdm |
|
|
| |
| ds = load_dataset("Datasets/eve-bio") |
| train_ds = ds["train"] |
| df = pl.from_pandas(train_ds.to_pandas()) |
|
|
| |
| pdb_dir = "uniprot_pdbs" |
| os.makedirs(pdb_dir, exist_ok=True) |
|
|
| from transformers import EsmForProteinFolding, AutoTokenizer |
| import torch |
|
|
| |
| print("Loading Local Model...") |
| tokenizer = AutoTokenizer.from_pretrained("huggingface/esmfold_v1") |
| model = EsmForProteinFolding.from_pretrained("huggingface/esmfold_v1") |
|
|
| |
| device = "cuda" if torch.cuda.is_available() else "cpu" |
| model = model.to(device) |
| |
| model.esm = model.esm.half() |
|
|
| def fold_local(sequence): |
| inputs = tokenizer([sequence], return_tensors="pt", add_special_tokens=False)['input_ids'] |
| tokenized_input = inputs.to(device) |
|
|
| with torch.no_grad(): |
| output = model(tokenized_input) |
| print(output) |
| return convert_outputs_to_pdb(output) |
|
|
| from transformers.models.esm.openfold_utils.protein import to_pdb, Protein as OFProtein |
| from transformers.models.esm.openfold_utils.feats import atom14_to_atom37 |
|
|
| def convert_outputs_to_pdb(outputs): |
| final_atom_positions = atom14_to_atom37(outputs["positions"][-1], outputs) |
| outputs = {k: v.to("cpu").numpy() for k, v in outputs.items()} |
| final_atom_positions = final_atom_positions.cpu().numpy() |
| final_atom_mask = outputs["atom37_atom_exists"] |
| pdbs = [] |
| for i in range(outputs["aatype"].shape[0]): |
| aa = outputs["aatype"][i] |
| pred_pos = final_atom_positions[i] |
| mask = final_atom_mask[i] |
| resid = outputs["residue_index"][i] + 1 |
| pred = OFProtein( |
| aatype=aa, |
| atom_positions=pred_pos, |
| atom_mask=mask, |
| residue_index=resid, |
| b_factors=outputs["plddt"][i], |
| chain_index=outputs["chain_index"][i] if "chain_index" in outputs else None, |
| ) |
| pdbs.append(to_pdb(pred)) |
| return pdbs |
|
|
| |
| def download_and_fold_uniprot(uniprot_id, pdb_path): |
| url = f"https://rest.uniprot.org/uniprotkb/{uniprot_id}.fasta" |
| resp = requests.get(url) |
| print("analysing", uniprot_id) |
| if resp.status_code == 200: |
| print(resp.text) |
| lines = resp.text.splitlines() |
| sequence = "".join(lines[1:]) |
| print("\n提取出的纯序列:") |
| print(sequence[:50] + "...") |
| else: |
| print("下载失败") |
| test_seq = sequence |
| output_file = pdb_path |
| |
| pdb = fold_local(test_seq) |
|
|
| |
|
|
| |
| |
| |
| with open(pdb_path, "w") as f: |
| f.write("".join(pdb)) |
|
|
| |
| uniprot_ids = df["target__uniprot_id"].unique() |
| uniprot_ids = [x for x in uniprot_ids if isinstance(x, str) and x.strip()] |
|
|
| for uniprot_id in tqdm(uniprot_ids, desc="Processing UniProt IDs"): |
| try: |
| pdb_path = os.path.join(pdb_dir, f"{uniprot_id}.pdb") |
| |
| if os.path.exists(pdb_path) and os.path.getsize(pdb_path) > 500: |
| continue |
| print(uniprot_id) |
| download_and_fold_uniprot(uniprot_id, pdb_path) |
| except: |
| pass |
|
|
|
|