import os import polars as pl from datasets import load_dataset import requests from tqdm import tqdm # Dataset loading and conversion to Polars DataFrame ds = load_dataset("Datasets/eve-bio") train_ds = ds["train"] df = pl.from_pandas(train_ds.to_pandas()) # Directory to save PDB files pdb_dir = "uniprot_pdbs" os.makedirs(pdb_dir, exist_ok=True) from transformers import EsmForProteinFolding, AutoTokenizer import torch # 1. 加载模型 (第一次会自动下载权重,约 3GB) print("Loading Local Model...") tokenizer = AutoTokenizer.from_pretrained("huggingface/esmfold_v1") model = EsmForProteinFolding.from_pretrained("huggingface/esmfold_v1") # 如果有 GPU,转到 GPU 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 # Function to download FASTA and fold to PDB 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) # import py3Dmol # view = py3Dmol.view(js='https://3dmol.org/build/3Dmol.js', width=800, height=400) # view.addModel("".join(pdb), 'pdb') # view.setStyle({'model': -1}, {"cartoon": {'color': 'spectrum'}}) with open(pdb_path, "w") as f: f.write("".join(pdb)) # Collect all unique UniProt IDs in the dataframe, skip null/missing 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") # Robust restart: skip already-completed 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