TW-DTI / DEV /data_analyses.py
graenys's picture
Upload 14 files
ebcab53 verified
Raw
History Blame Contribute Delete
3.59 kB
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