| |
| from __future__ import annotations |
|
|
| from pathlib import Path |
| from typing import List, Optional, Tuple |
|
|
| from huggingface_hub import snapshot_download |
| from inference import ( |
| PeptiVersePredictor, |
| read_best_manifest_csv, |
| _parse_model_and_emb, |
| EMB_TAG_TO_FOLDER_SUFFIX, |
| ) |
|
|
| |
| |
| |
| root = Path(__file__).resolve().parent |
| MODEL_REPO = "ChatterjeeLab/PeptiVerse" |
| DEFAULT_ASSETS_DIR = Path(root) |
| DEFAULT_MANIFEST = Path("./basic_models.txt") |
|
|
| BANNED_MODELS = {"svm", "enet", "svm_gpu", "enet_gpu"} |
|
|
|
|
| def _norm_prop_disk(prop_key: str) -> str: |
| return "half_life" if prop_key == "halflife" else prop_key |
|
|
|
|
| def _resolve_expected_model_dir( |
| prop_key: str, model_name: str, emb_tag: Optional[str] |
| ) -> str: |
| """ |
| Resolve the subfolder path inside training_classifiers/<property>/. |
| |
| Args: |
| prop_key: manifest property key, e.g. 'hemolysis', 'halflife'. |
| model_name: canonical model string, e.g. 'xgb', 'cnn', 'transformer'. |
| emb_tag: embedding tag from manifest: 'wt', 'peptideclm', or 'chemberta'. |
| None means WT (falls back to 'wt'). |
| """ |
| disk_prop = _norm_prop_disk(prop_key) |
| base = f"training_classifiers/{disk_prop}" |
|
|
| folder_suffix = EMB_TAG_TO_FOLDER_SUFFIX.get(emb_tag or "wt", emb_tag or "wt") |
|
|
| if prop_key == "binding_affinity": |
| return f"{base}/{model_name}" |
|
|
| |
| |
| |
| if prop_key == "halflife" and (emb_tag is None or emb_tag == "wt"): |
| if model_name == "transformer": |
| return f"{base}/transformer_wt_log" |
| if model_name in {"xgb", "xgb_reg"}: |
| return f"{base}/xgb_wt_log" |
|
|
| |
| |
| |
| |
| return f"{base}/{model_name}_{folder_suffix}" |
|
|
|
|
| def build_allow_patterns_from_manifest(manifest_path: Path) -> List[str]: |
| best = read_best_manifest_csv(manifest_path) |
|
|
| allow: List[str] = [] |
|
|
| for prop_key, row in best.items(): |
| for col, parsed in [("wt", row.best_wt), ("smiles", row.best_smiles)]: |
| if parsed is None: |
| continue |
|
|
| model_name, emb_tag = parsed |
|
|
| |
| if model_name in BANNED_MODELS: |
| model_name = "xgb" |
|
|
| model_dir = _resolve_expected_model_dir(prop_key, model_name, emb_tag) |
|
|
| allow += [ |
| f"{model_dir}/best_model.json", |
| f"{model_dir}/best_model.pt", |
| f"{model_dir}/best_model*.joblib", |
| f"{model_dir}/best_model*.json", |
| ] |
|
|
| |
| seen = set() |
| out = [] |
| for p in allow: |
| if p not in seen: |
| out.append(p) |
| seen.add(p) |
| return out |
|
|
|
|
| def download_assets( |
| repo_id: str, |
| manifest_path: Path, |
| out_dir: Path, |
| ) -> Path: |
| out_dir = out_dir.resolve() |
| out_dir.mkdir(parents=True, exist_ok=True) |
|
|
| allow_patterns = build_allow_patterns_from_manifest(manifest_path) |
|
|
| snapshot_download( |
| repo_id=repo_id, |
| local_dir=str(out_dir), |
| local_dir_use_symlinks=False, |
| allow_patterns=allow_patterns, |
| ) |
| return out_dir |
|
|
|
|
| |
| |
| |
| def main(): |
| import argparse |
|
|
| ap = argparse.ArgumentParser( |
| description="Lightweight PeptiVerse inference with on-demand model download." |
| ) |
| ap.add_argument("--repo", default=MODEL_REPO, help="HF repo id containing weights/assets.") |
| ap.add_argument("--manifest", default=str(DEFAULT_MANIFEST), help="Path to best_models.txt") |
| ap.add_argument("--assets", default=str(DEFAULT_ASSETS_DIR), help="Where to store downloaded assets") |
| ap.add_argument("--device", default=None, help="cuda / cpu / cuda:0, etc") |
| ap.add_argument("--dry-run", action="store_true", help="Print allow-patterns without downloading") |
|
|
| ap.add_argument("--property", default="hemolysis") |
| ap.add_argument("--mode", default="wt", choices=["wt", "smiles"]) |
| ap.add_argument("--input", default="GIGAVLKVLTTGLPALISWIKRKRQQ") |
| ap.add_argument("--target_seq", default="GIGAVLKVLTTGLPALISWIKRKRQQ") |
| ap.add_argument("--binder", default="GIGAVLKV") |
| args = ap.parse_args() |
|
|
| manifest_path = Path(args.manifest) |
| if not manifest_path.exists(): |
| raise FileNotFoundError(f"Manifest not found: {manifest_path}") |
|
|
| if args.dry_run: |
| patterns = build_allow_patterns_from_manifest(manifest_path) |
| print(f"Would download {len(patterns)} patterns:") |
| for p in patterns: |
| print(" ", p) |
| return |
|
|
| assets_dir = download_assets( |
| args.repo, manifest_path=manifest_path, out_dir=Path(args.assets) |
| ) |
|
|
| """TEST CODE |
| predictor = PeptiVersePredictor( |
| manifest_path=manifest_path, |
| classifier_weight_root=str(assets_dir), |
| device=args.device, |
| ) |
| if args.property == "binding_affinity": |
| if not args.target_seq or not args.binder: |
| raise ValueError("For binding_affinity, provide --target_seq and --binder.") |
| out = predictor.predict_binding_affinity(args.mode, target_seq=args.target_seq, binder_str=args.binder) |
| else: |
| out = predictor.predict_property(args.property, args.mode, args.input) |
| print(out) |
| """ |
|
|
| if __name__ == "__main__": |
| main() |