Chainsaw / scripts /development /tests /unit /test_get_predictions.py
wuxing0105's picture
Upload folder using huggingface_hub
80a72c3 verified
Raw
History Blame Contribute Delete
2.03 kB
import shutil
import sys
from pathlib import Path
# ruff: noqa: E402
REPO_ROOT = Path(__file__).parents[4].resolve()
sys.path.append(f"{REPO_ROOT}")
from scripts.get_predictions import predict, load_model, PredictionResult
DEFAULT_DISORDERED_DOMAIN_THRESHOLD = 0.35
def get_length_from_pdb_file(pdb_file):
with open(pdb_file) as f:
for line in f:
if line.startswith('SEQRES'):
return int(line.split()[3])
def test_predict_from_pdb_file(tmp_path, capsys):
example_af_id = "AF-A0A1W2PQ64-F1-model_v4"
expected_structure_file = REPO_ROOT / 'scripts' / 'example_files' / f'{example_af_id}.pdb'
expected_nres = get_length_from_pdb_file(expected_structure_file)
expected_result = PredictionResult(
pdb_path=expected_structure_file,
chain_id=example_af_id,
sequence_md5='a126e3d4d1a2dcadaa684287855d19f4',
nres=expected_nres,
ndom=3,
chopping='7-40_95-193,42-91',
confidence=0.8307,
)
tmp_results_dir = tmp_path / "results"
tmp_results_dir.mkdir()
expected_model_dir = REPO_ROOT / 'weight' / 'model_v1'
tmp_model_dir = tmp_path / "models"
tmp_model_dir.mkdir()
shutil.copyfile(str(expected_model_dir / 'weights.pt'), str(tmp_model_dir / 'weights.pt'))
expected_config_dir = REPO_ROOT / 'conf' / 'model_v1'
for config_fname in ['config.json', 'feature_config.json']:
shutil.copyfile(str(expected_config_dir / config_fname), str(tmp_model_dir / config_fname))
tmp_structure_file = tmp_path / f"{example_af_id}.pdb"
shutil.copyfile(str(expected_structure_file), str(tmp_structure_file))
model = load_model(
model_dir=str(tmp_model_dir),
remove_disordered_domain_threshold=DEFAULT_DISORDERED_DOMAIN_THRESHOLD)
result = predict(model, str(tmp_structure_file))
assert normalise_result(result) == normalise_result(expected_result)
def normalise_result(res):
res.confidence = round(res.confidence, 4)
res.pdb_path = '__PDB_PATH__'