chq1155 commited on
Commit ·
ee96220
1
Parent(s): b43a758
Add PeptiVerse affinity backend
Browse files- README.md +18 -0
- baselines/run_mcts_tr2d2.py +17 -2
- baselines/run_validation_td3b.py +17 -2
- baselines/sampling_setup.py +9 -4
- env.yml +2 -1
- finetune_multi_target.py +8 -5
- inference.py +6 -3
- launch_multi_target.sh +19 -0
- scoring/affinity_config.py +78 -0
- scoring/functions/binding.py +54 -3
- scoring/functions/peptiverse_binding.py +274 -0
- td3b/__init__.py +12 -1
- tests/test_peptiverse_binding.py +124 -0
- training/finetune_utils.py +7 -2
README.md
CHANGED
|
@@ -179,6 +179,24 @@ Key flags:
|
|
| 179 |
|
| 180 |
Output: prints the valid yield (valid / `num_samples`) with per-round counts, and writes the valid sequences to `--save_path` (default `results/valid_<strategy>_len<L>.csv`) with columns idx, sequence, n_chars.
|
| 181 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 182 |
## Training
|
| 183 |
|
| 184 |
### Multi-target TD3B
|
|
|
|
| 179 |
|
| 180 |
Output: prints the valid yield (valid / `num_samples`) with per-round counts, and writes the valid sequences to `--save_path` (default `results/valid_<strategy>_len<L>.csv`) with columns idx, sequence, n_chars.
|
| 181 |
|
| 182 |
+
### Binding-affinity backends
|
| 183 |
+
|
| 184 |
+
The original TD3B affinity predictor remains the default. To use PeptiVerse's
|
| 185 |
+
pooled target-sequence/binder-SMILES regressor instead, add:
|
| 186 |
+
|
| 187 |
+
```bash
|
| 188 |
+
python inference.py \
|
| 189 |
+
--ckpt_path checkpoints/td3b.ckpt \
|
| 190 |
+
--val_csv data/test.csv \
|
| 191 |
+
--affinity_backend peptiverse
|
| 192 |
+
```
|
| 193 |
+
|
| 194 |
+
The PeptiVerse checkpoint and its ESM-2/ChemBERTa encoders are downloaded from
|
| 195 |
+
Hugging Face on first use. For offline runs, provide
|
| 196 |
+
`--peptiverse_affinity_checkpoint /path/to/best_model.pt` together with
|
| 197 |
+
`--peptiverse_local_files_only`. Training supports the same flags; alternatively,
|
| 198 |
+
set `AFFINITY_BACKEND="peptiverse"` in `launch_multi_target.sh`.
|
| 199 |
+
|
| 200 |
## Training
|
| 201 |
|
| 202 |
### Multi-target TD3B
|
baselines/run_mcts_tr2d2.py
CHANGED
|
@@ -27,7 +27,8 @@ from configs.finetune_config import (
|
|
| 27 |
)
|
| 28 |
from training.finetune_utils import load_tokenizer
|
| 29 |
from training.distributed_utils import setup_distributed, cleanup_distributed, is_main_process
|
| 30 |
-
from scoring.
|
|
|
|
| 31 |
from td3b.direction_oracle import DirectionalOracle
|
| 32 |
from finetune_multi_target_tr2d2_ddp import TR2D2GatedReward, TargetDataset, create_tr2d2_mcts
|
| 33 |
from utils.app import PeptideAnalyzer
|
|
@@ -106,6 +107,18 @@ def _build_args(cfg: Dict[str, Any], cli: argparse.Namespace) -> argparse.Namesp
|
|
| 106 |
merged["exploration"] = cli.exploration
|
| 107 |
if cli.max_sequence_length is not None:
|
| 108 |
merged["max_sequence_length"] = cli.max_sequence_length
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 109 |
|
| 110 |
args = SimpleNamespace(**merged)
|
| 111 |
|
|
@@ -216,6 +229,7 @@ def _nanstd(values: np.ndarray) -> float:
|
|
| 216 |
|
| 217 |
def main() -> None:
|
| 218 |
parser = argparse.ArgumentParser(description="MCTS-based TR2-D2 evaluation.")
|
|
|
|
| 219 |
parser.add_argument("--ckpt_path", required=True, help="Path to finetuned checkpoint (.ckpt)")
|
| 220 |
parser.add_argument("--val_csv", required=True, help="Validation CSV path")
|
| 221 |
parser.add_argument("--device", default="cuda", help="Device string (e.g., cuda:0 or cpu)")
|
|
@@ -258,7 +272,8 @@ def main() -> None:
|
|
| 258 |
|
| 259 |
policy_model = _build_model(args, payload["state_dict"], device)
|
| 260 |
|
| 261 |
-
multi_target_affinity =
|
|
|
|
| 262 |
tokenizer=tokenizer,
|
| 263 |
base_path=args.base_path,
|
| 264 |
device=device,
|
|
|
|
| 27 |
)
|
| 28 |
from training.finetune_utils import load_tokenizer
|
| 29 |
from training.distributed_utils import setup_distributed, cleanup_distributed, is_main_process
|
| 30 |
+
from scoring.affinity_config import add_affinity_arguments, create_affinity_from_args
|
| 31 |
+
from scoring.functions.binding import TargetSpecificBindingAffinity
|
| 32 |
from td3b.direction_oracle import DirectionalOracle
|
| 33 |
from finetune_multi_target_tr2d2_ddp import TR2D2GatedReward, TargetDataset, create_tr2d2_mcts
|
| 34 |
from utils.app import PeptideAnalyzer
|
|
|
|
| 107 |
merged["exploration"] = cli.exploration
|
| 108 |
if cli.max_sequence_length is not None:
|
| 109 |
merged["max_sequence_length"] = cli.max_sequence_length
|
| 110 |
+
for name in (
|
| 111 |
+
"affinity_backend",
|
| 112 |
+
"peptiverse_affinity_checkpoint",
|
| 113 |
+
"peptiverse_repo_id",
|
| 114 |
+
"peptiverse_revision",
|
| 115 |
+
"peptiverse_cache_dir",
|
| 116 |
+
"peptiverse_local_files_only",
|
| 117 |
+
"peptiverse_batch_size",
|
| 118 |
+
):
|
| 119 |
+
value = getattr(cli, name, None)
|
| 120 |
+
if value is not None:
|
| 121 |
+
merged[name] = value
|
| 122 |
|
| 123 |
args = SimpleNamespace(**merged)
|
| 124 |
|
|
|
|
| 229 |
|
| 230 |
def main() -> None:
|
| 231 |
parser = argparse.ArgumentParser(description="MCTS-based TR2-D2 evaluation.")
|
| 232 |
+
add_affinity_arguments(parser, default_backend=None)
|
| 233 |
parser.add_argument("--ckpt_path", required=True, help="Path to finetuned checkpoint (.ckpt)")
|
| 234 |
parser.add_argument("--val_csv", required=True, help="Validation CSV path")
|
| 235 |
parser.add_argument("--device", default="cuda", help="Device string (e.g., cuda:0 or cpu)")
|
|
|
|
| 272 |
|
| 273 |
policy_model = _build_model(args, payload["state_dict"], device)
|
| 274 |
|
| 275 |
+
multi_target_affinity = create_affinity_from_args(
|
| 276 |
+
args=args,
|
| 277 |
tokenizer=tokenizer,
|
| 278 |
base_path=args.base_path,
|
| 279 |
device=device,
|
baselines/run_validation_td3b.py
CHANGED
|
@@ -28,7 +28,8 @@ from configs.finetune_config import (
|
|
| 28 |
from training.finetune_utils import load_tokenizer, create_reward_function
|
| 29 |
from finetune_multi_target import TargetDataset
|
| 30 |
from training.distributed_utils import setup_distributed, cleanup_distributed, is_main_process
|
| 31 |
-
from scoring.
|
|
|
|
| 32 |
from td3b.direction_oracle import DirectionalOracle
|
| 33 |
from utils.app import PeptideAnalyzer
|
| 34 |
|
|
@@ -94,6 +95,18 @@ def _build_args(cfg: Dict[str, Any], cli: argparse.Namespace) -> argparse.Namesp
|
|
| 94 |
merged["sampling_eps"] = cli.sampling_eps
|
| 95 |
if cli.seed is not None:
|
| 96 |
merged["seed"] = cli.seed
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 97 |
|
| 98 |
args = SimpleNamespace(**merged)
|
| 99 |
|
|
@@ -270,6 +283,7 @@ def _nanstd(values: np.ndarray) -> float:
|
|
| 270 |
|
| 271 |
def main() -> None:
|
| 272 |
parser = argparse.ArgumentParser(description="Run TD3B validation from a saved checkpoint.")
|
|
|
|
| 273 |
parser.add_argument("--ckpt_path", required=True, help="Path to saved checkpoint (.ckpt)")
|
| 274 |
parser.add_argument("--val_csv", required=True, help="Validation CSV path")
|
| 275 |
parser.add_argument("--device", default="cuda", help="Device string (e.g., cuda:0 or cpu)")
|
|
@@ -312,7 +326,8 @@ def main() -> None:
|
|
| 312 |
|
| 313 |
policy_model = _build_model(args, payload["state_dict"], device)
|
| 314 |
|
| 315 |
-
multi_target_affinity =
|
|
|
|
| 316 |
tokenizer=tokenizer,
|
| 317 |
base_path=args.base_path,
|
| 318 |
device=device,
|
|
|
|
| 28 |
from training.finetune_utils import load_tokenizer, create_reward_function
|
| 29 |
from finetune_multi_target import TargetDataset
|
| 30 |
from training.distributed_utils import setup_distributed, cleanup_distributed, is_main_process
|
| 31 |
+
from scoring.affinity_config import add_affinity_arguments, create_affinity_from_args
|
| 32 |
+
from scoring.functions.binding import TargetSpecificBindingAffinity
|
| 33 |
from td3b.direction_oracle import DirectionalOracle
|
| 34 |
from utils.app import PeptideAnalyzer
|
| 35 |
|
|
|
|
| 95 |
merged["sampling_eps"] = cli.sampling_eps
|
| 96 |
if cli.seed is not None:
|
| 97 |
merged["seed"] = cli.seed
|
| 98 |
+
for name in (
|
| 99 |
+
"affinity_backend",
|
| 100 |
+
"peptiverse_affinity_checkpoint",
|
| 101 |
+
"peptiverse_repo_id",
|
| 102 |
+
"peptiverse_revision",
|
| 103 |
+
"peptiverse_cache_dir",
|
| 104 |
+
"peptiverse_local_files_only",
|
| 105 |
+
"peptiverse_batch_size",
|
| 106 |
+
):
|
| 107 |
+
value = getattr(cli, name, None)
|
| 108 |
+
if value is not None:
|
| 109 |
+
merged[name] = value
|
| 110 |
|
| 111 |
args = SimpleNamespace(**merged)
|
| 112 |
|
|
|
|
| 283 |
|
| 284 |
def main() -> None:
|
| 285 |
parser = argparse.ArgumentParser(description="Run TD3B validation from a saved checkpoint.")
|
| 286 |
+
add_affinity_arguments(parser, default_backend=None)
|
| 287 |
parser.add_argument("--ckpt_path", required=True, help="Path to saved checkpoint (.ckpt)")
|
| 288 |
parser.add_argument("--val_csv", required=True, help="Validation CSV path")
|
| 289 |
parser.add_argument("--device", default="cuda", help="Device string (e.g., cuda:0 or cpu)")
|
|
|
|
| 326 |
|
| 327 |
policy_model = _build_model(args, payload["state_dict"], device)
|
| 328 |
|
| 329 |
+
multi_target_affinity = create_affinity_from_args(
|
| 330 |
+
args=args,
|
| 331 |
tokenizer=tokenizer,
|
| 332 |
base_path=args.base_path,
|
| 333 |
device=device,
|
baselines/sampling_setup.py
CHANGED
|
@@ -15,8 +15,8 @@ from hydra import compose, initialize_config_dir
|
|
| 15 |
from hydra.core.global_hydra import GlobalHydra
|
| 16 |
|
| 17 |
from models.diffusion import Diffusion
|
|
|
|
| 18 |
from scoring.scoring_functions import ScoringFunctions
|
| 19 |
-
from scoring.functions.binding import MultiTargetBindingAffinity
|
| 20 |
from td3b.direction_oracle import DirectionalOracle, resolve_device
|
| 21 |
from td3b.data_utils import peptide_seq_to_smiles, smiles_token_length
|
| 22 |
|
|
@@ -104,11 +104,14 @@ def load_reward_models(
|
|
| 104 |
base_path: Optional[str] = None,
|
| 105 |
multi_target: bool = False,
|
| 106 |
score_func_names: Optional[List[str]] = None,
|
|
|
|
| 107 |
):
|
| 108 |
-
|
|
|
|
| 109 |
if base_model is None or base_path is None:
|
| 110 |
-
raise ValueError("base_model and base_path are required for
|
| 111 |
-
return
|
|
|
|
| 112 |
tokenizer=base_model.tokenizer,
|
| 113 |
base_path=base_path,
|
| 114 |
device=device,
|
|
@@ -222,6 +225,7 @@ def run_baseline(
|
|
| 222 |
|
| 223 |
def main():
|
| 224 |
parser = argparse.ArgumentParser()
|
|
|
|
| 225 |
parser.add_argument("--ckpt_path", type=str, required=True)
|
| 226 |
parser.add_argument("--device", type=str, default="cuda:0")
|
| 227 |
parser.add_argument("--baseline", type=str, default="cg", choices=["cg", "smc", "tds", "unguided", "peptune"])
|
|
@@ -308,6 +312,7 @@ def main():
|
|
| 308 |
base_model=base_model,
|
| 309 |
base_path=base_path,
|
| 310 |
multi_target=multi_target,
|
|
|
|
| 311 |
)
|
| 312 |
direction_oracle = load_direction_oracle(args, args.device)
|
| 313 |
reward_alpha = args.reward_alpha if args.reward_alpha is not None else args.alpha
|
|
|
|
| 15 |
from hydra.core.global_hydra import GlobalHydra
|
| 16 |
|
| 17 |
from models.diffusion import Diffusion
|
| 18 |
+
from scoring.affinity_config import add_affinity_arguments, create_affinity_from_args
|
| 19 |
from scoring.scoring_functions import ScoringFunctions
|
|
|
|
| 20 |
from td3b.direction_oracle import DirectionalOracle, resolve_device
|
| 21 |
from td3b.data_utils import peptide_seq_to_smiles, smiles_token_length
|
| 22 |
|
|
|
|
| 104 |
base_path: Optional[str] = None,
|
| 105 |
multi_target: bool = False,
|
| 106 |
score_func_names: Optional[List[str]] = None,
|
| 107 |
+
args=None,
|
| 108 |
):
|
| 109 |
+
affinity_backend = getattr(args, "affinity_backend", "original")
|
| 110 |
+
if multi_target or affinity_backend == "peptiverse":
|
| 111 |
if base_model is None or base_path is None:
|
| 112 |
+
raise ValueError("base_model and base_path are required for affinity scoring.")
|
| 113 |
+
return create_affinity_from_args(
|
| 114 |
+
args=args,
|
| 115 |
tokenizer=base_model.tokenizer,
|
| 116 |
base_path=base_path,
|
| 117 |
device=device,
|
|
|
|
| 225 |
|
| 226 |
def main():
|
| 227 |
parser = argparse.ArgumentParser()
|
| 228 |
+
add_affinity_arguments(parser)
|
| 229 |
parser.add_argument("--ckpt_path", type=str, required=True)
|
| 230 |
parser.add_argument("--device", type=str, default="cuda:0")
|
| 231 |
parser.add_argument("--baseline", type=str, default="cg", choices=["cg", "smc", "tds", "unguided", "peptune"])
|
|
|
|
| 312 |
base_model=base_model,
|
| 313 |
base_path=base_path,
|
| 314 |
multi_target=multi_target,
|
| 315 |
+
args=args,
|
| 316 |
)
|
| 317 |
direction_oracle = load_direction_oracle(args, args.device)
|
| 318 |
reward_alpha = args.reward_alpha if args.reward_alpha is not None else args.alpha
|
env.yml
CHANGED
|
@@ -23,6 +23,7 @@ dependencies:
|
|
| 23 |
- lightning==2.5.5
|
| 24 |
- fair-esm==2.0.0
|
| 25 |
- transformers==4.56.2
|
|
|
|
| 26 |
- SmilesPE==0.0.3
|
| 27 |
- scipy==1.13.1
|
| 28 |
- wandb==0.22.0
|
|
@@ -34,4 +35,4 @@ dependencies:
|
|
| 34 |
- seaborn==0.13.2
|
| 35 |
- timm==1.0.20
|
| 36 |
- xgboost==3.0.5
|
| 37 |
-
- loguru==0.7.3
|
|
|
|
| 23 |
- lightning==2.5.5
|
| 24 |
- fair-esm==2.0.0
|
| 25 |
- transformers==4.56.2
|
| 26 |
+
- huggingface-hub
|
| 27 |
- SmilesPE==0.0.3
|
| 28 |
- scipy==1.13.1
|
| 29 |
- wandb==0.22.0
|
|
|
|
| 35 |
- seaborn==0.13.2
|
| 36 |
- timm==1.0.20
|
| 37 |
- xgboost==3.0.5
|
| 38 |
+
- loguru==0.7.3
|
finetune_multi_target.py
CHANGED
|
@@ -36,7 +36,8 @@ sys.path.insert(0, os.path.dirname(os.path.abspath(__file__)))
|
|
| 36 |
from models.diffusion import Diffusion
|
| 37 |
from tokenizer.my_tokenizers import SMILES_SPE_Tokenizer
|
| 38 |
from utils.app import PeptideAnalyzer
|
| 39 |
-
from scoring.
|
|
|
|
| 40 |
from td3b.data_utils import peptide_seq_to_smiles, smiles_token_length
|
| 41 |
|
| 42 |
# TD3B imports
|
|
@@ -193,7 +194,7 @@ class TargetDataset:
|
|
| 193 |
|
| 194 |
def run_validation(
|
| 195 |
policy_model: Diffusion,
|
| 196 |
-
multi_target_affinity
|
| 197 |
directional_oracle: DirectionalOracle,
|
| 198 |
tokenizer: SMILES_SPE_Tokenizer,
|
| 199 |
val_dataset: TargetDataset,
|
|
@@ -394,6 +395,7 @@ def run_validation(
|
|
| 394 |
def parse_args():
|
| 395 |
"""Parse command-line arguments."""
|
| 396 |
parser = argparse.ArgumentParser(description='Multi-Target TD3B Fine-Tuning')
|
|
|
|
| 397 |
|
| 398 |
# Paths
|
| 399 |
path_group = parser.add_argument_group('Paths')
|
|
@@ -681,13 +683,14 @@ def main():
|
|
| 681 |
policy_model = add_td3b_sampling_to_model(policy_model)
|
| 682 |
|
| 683 |
# Multi-target affinity predictor
|
| 684 |
-
multi_target_affinity =
|
|
|
|
| 685 |
tokenizer=tokenizer,
|
| 686 |
base_path=args.base_path,
|
| 687 |
device=device,
|
| 688 |
-
emb_model=policy_model.backbone
|
| 689 |
)
|
| 690 |
-
logger.info("Created
|
| 691 |
|
| 692 |
# Directional oracle (GPCR classifier)
|
| 693 |
for path_label, path in [
|
|
|
|
| 36 |
from models.diffusion import Diffusion
|
| 37 |
from tokenizer.my_tokenizers import SMILES_SPE_Tokenizer
|
| 38 |
from utils.app import PeptideAnalyzer
|
| 39 |
+
from scoring.affinity_config import add_affinity_arguments, create_affinity_from_args
|
| 40 |
+
from scoring.functions.binding import TargetSpecificBindingAffinity
|
| 41 |
from td3b.data_utils import peptide_seq_to_smiles, smiles_token_length
|
| 42 |
|
| 43 |
# TD3B imports
|
|
|
|
| 194 |
|
| 195 |
def run_validation(
|
| 196 |
policy_model: Diffusion,
|
| 197 |
+
multi_target_affinity,
|
| 198 |
directional_oracle: DirectionalOracle,
|
| 199 |
tokenizer: SMILES_SPE_Tokenizer,
|
| 200 |
val_dataset: TargetDataset,
|
|
|
|
| 395 |
def parse_args():
|
| 396 |
"""Parse command-line arguments."""
|
| 397 |
parser = argparse.ArgumentParser(description='Multi-Target TD3B Fine-Tuning')
|
| 398 |
+
add_affinity_arguments(parser)
|
| 399 |
|
| 400 |
# Paths
|
| 401 |
path_group = parser.add_argument_group('Paths')
|
|
|
|
| 683 |
policy_model = add_td3b_sampling_to_model(policy_model)
|
| 684 |
|
| 685 |
# Multi-target affinity predictor
|
| 686 |
+
multi_target_affinity = create_affinity_from_args(
|
| 687 |
+
args=args,
|
| 688 |
tokenizer=tokenizer,
|
| 689 |
base_path=args.base_path,
|
| 690 |
device=device,
|
| 691 |
+
emb_model=policy_model.backbone,
|
| 692 |
)
|
| 693 |
+
logger.info("Created %s binding affinity predictor", args.affinity_backend)
|
| 694 |
|
| 695 |
# Directional oracle (GPCR classifier)
|
| 696 |
for path_label, path in [
|
inference.py
CHANGED
|
@@ -29,10 +29,11 @@ from configs.finetune_config import (
|
|
| 29 |
DiffusionConfig, RoFormerConfig, NoiseConfig,
|
| 30 |
TrainingConfig, SamplingConfig, EvalConfig, OptimConfig, MCTSConfig,
|
| 31 |
)
|
|
|
|
|
|
|
| 32 |
from training.finetune_utils import load_tokenizer
|
| 33 |
from td3b.direction_oracle import DirectionalOracle
|
| 34 |
from td3b.td3b_scoring import create_td3b_reward_function
|
| 35 |
-
from scoring.functions.binding import MultiTargetBindingAffinity, TargetSpecificBindingAffinity
|
| 36 |
from td3b.data_utils import peptide_seq_to_smiles, smiles_token_length
|
| 37 |
from utils.app import PeptideAnalyzer
|
| 38 |
import sampling_strategies
|
|
@@ -138,6 +139,7 @@ def score_sequences(reward_model, sequences: List[str]):
|
|
| 138 |
|
| 139 |
def main():
|
| 140 |
parser = argparse.ArgumentParser(description="TD3B Inference")
|
|
|
|
| 141 |
parser.add_argument("--ckpt_path", type=str, required=True, help="Path to TD3B checkpoint")
|
| 142 |
parser.add_argument("--val_csv", type=str, required=True, help="CSV with Target_Sequence, Ligand_Sequence, label columns")
|
| 143 |
parser.add_argument("--save_path", type=str, default="results", help="Output directory")
|
|
@@ -185,7 +187,6 @@ def main():
|
|
| 185 |
# Load model
|
| 186 |
logger.info(f"Loading model from {args.ckpt_path}")
|
| 187 |
model, tokenizer = load_model(args.ckpt_path, device)
|
| 188 |
-
|
| 189 |
# Load targets
|
| 190 |
logger.info(f"Loading targets from {args.val_csv}")
|
| 191 |
df = pd.read_csv(args.val_csv)
|
|
@@ -219,12 +220,14 @@ def main():
|
|
| 219 |
vocab_path = os.path.join(ROOT_DIR, "tokenizer", "new_vocab.txt")
|
| 220 |
splits_path = os.path.join(ROOT_DIR, "tokenizer", "new_splits.txt")
|
| 221 |
|
| 222 |
-
multi_affinity =
|
|
|
|
| 223 |
tokenizer=tokenizer,
|
| 224 |
base_path=ROOT_DIR,
|
| 225 |
device=device,
|
| 226 |
emb_model=model.backbone,
|
| 227 |
)
|
|
|
|
| 228 |
directional_oracle = DirectionalOracle(
|
| 229 |
model_ckpt=oracle_ckpt,
|
| 230 |
tr2d2_checkpoint=oracle_tr2d2,
|
|
|
|
| 29 |
DiffusionConfig, RoFormerConfig, NoiseConfig,
|
| 30 |
TrainingConfig, SamplingConfig, EvalConfig, OptimConfig, MCTSConfig,
|
| 31 |
)
|
| 32 |
+
from scoring.affinity_config import add_affinity_arguments, create_affinity_from_args
|
| 33 |
+
from scoring.functions.binding import TargetSpecificBindingAffinity
|
| 34 |
from training.finetune_utils import load_tokenizer
|
| 35 |
from td3b.direction_oracle import DirectionalOracle
|
| 36 |
from td3b.td3b_scoring import create_td3b_reward_function
|
|
|
|
| 37 |
from td3b.data_utils import peptide_seq_to_smiles, smiles_token_length
|
| 38 |
from utils.app import PeptideAnalyzer
|
| 39 |
import sampling_strategies
|
|
|
|
| 139 |
|
| 140 |
def main():
|
| 141 |
parser = argparse.ArgumentParser(description="TD3B Inference")
|
| 142 |
+
add_affinity_arguments(parser)
|
| 143 |
parser.add_argument("--ckpt_path", type=str, required=True, help="Path to TD3B checkpoint")
|
| 144 |
parser.add_argument("--val_csv", type=str, required=True, help="CSV with Target_Sequence, Ligand_Sequence, label columns")
|
| 145 |
parser.add_argument("--save_path", type=str, default="results", help="Output directory")
|
|
|
|
| 187 |
# Load model
|
| 188 |
logger.info(f"Loading model from {args.ckpt_path}")
|
| 189 |
model, tokenizer = load_model(args.ckpt_path, device)
|
|
|
|
| 190 |
# Load targets
|
| 191 |
logger.info(f"Loading targets from {args.val_csv}")
|
| 192 |
df = pd.read_csv(args.val_csv)
|
|
|
|
| 220 |
vocab_path = os.path.join(ROOT_DIR, "tokenizer", "new_vocab.txt")
|
| 221 |
splits_path = os.path.join(ROOT_DIR, "tokenizer", "new_splits.txt")
|
| 222 |
|
| 223 |
+
multi_affinity = create_affinity_from_args(
|
| 224 |
+
args=args,
|
| 225 |
tokenizer=tokenizer,
|
| 226 |
base_path=ROOT_DIR,
|
| 227 |
device=device,
|
| 228 |
emb_model=model.backbone,
|
| 229 |
)
|
| 230 |
+
logger.info("Using %s binding affinity predictor", args.affinity_backend)
|
| 231 |
directional_oracle = DirectionalOracle(
|
| 232 |
model_ckpt=oracle_ckpt,
|
| 233 |
tr2d2_checkpoint=oracle_tr2d2,
|
launch_multi_target.sh
CHANGED
|
@@ -17,6 +17,12 @@ VAL_CSV="${BASE_PATH}/data/test.csv" # Optional: create validation split
|
|
| 17 |
# Run configuration
|
| 18 |
RUN_NAME="multi_target_td3b" # Timestamp will be added automatically
|
| 19 |
DEVICE="cuda:0"
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 20 |
# Multi-target sampling
|
| 21 |
TARGETS_PER_MCTS=2 # Number of targets sampled per MCTS round (K)
|
| 22 |
RESAMPLE_TARGETS_EVERY=1 # Resample targets every N epochs
|
|
@@ -70,6 +76,17 @@ EXTRA_ORACLE_ARGS=""
|
|
| 70 |
if [ -n "$ORACLE_ESM_CACHE_DIR" ]; then
|
| 71 |
EXTRA_ORACLE_ARGS="$EXTRA_ORACLE_ARGS --direction_oracle_esm_cache_dir $ORACLE_ESM_CACHE_DIR"
|
| 72 |
fi
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 73 |
if [ "$ORACLE_ESM_LOCAL_FILES_ONLY" -eq 1 ]; then
|
| 74 |
EXTRA_ORACLE_ARGS="$EXTRA_ORACLE_ARGS --direction_oracle_esm_local_files_only"
|
| 75 |
fi
|
|
@@ -107,6 +124,8 @@ CMD="python finetune_multi_target.py \
|
|
| 107 |
--pretrained_checkpoint ${PRETRAINED_CHECKPOINT} \
|
| 108 |
--run_name ${RUN_NAME} \
|
| 109 |
--device ${DEVICE} \
|
|
|
|
|
|
|
| 110 |
\
|
| 111 |
--targets_per_mcts ${TARGETS_PER_MCTS} \
|
| 112 |
--resample_targets_every ${RESAMPLE_TARGETS_EVERY} \
|
|
|
|
| 17 |
# Run configuration
|
| 18 |
RUN_NAME="multi_target_td3b" # Timestamp will be added automatically
|
| 19 |
DEVICE="cuda:0"
|
| 20 |
+
|
| 21 |
+
# Binding-affinity backend ("original" preserves the published TD3B behavior)
|
| 22 |
+
AFFINITY_BACKEND="original" # original or peptiverse
|
| 23 |
+
PEPTIVERSE_AFFINITY_CHECKPOINT="" # Optional local best_model.pt
|
| 24 |
+
PEPTIVERSE_CACHE_DIR="" # Optional Hugging Face cache directory
|
| 25 |
+
PEPTIVERSE_LOCAL_FILES_ONLY=0 # Set to 1 for offline-only loading
|
| 26 |
# Multi-target sampling
|
| 27 |
TARGETS_PER_MCTS=2 # Number of targets sampled per MCTS round (K)
|
| 28 |
RESAMPLE_TARGETS_EVERY=1 # Resample targets every N epochs
|
|
|
|
| 76 |
if [ -n "$ORACLE_ESM_CACHE_DIR" ]; then
|
| 77 |
EXTRA_ORACLE_ARGS="$EXTRA_ORACLE_ARGS --direction_oracle_esm_cache_dir $ORACLE_ESM_CACHE_DIR"
|
| 78 |
fi
|
| 79 |
+
|
| 80 |
+
EXTRA_AFFINITY_ARGS=""
|
| 81 |
+
if [ -n "$PEPTIVERSE_AFFINITY_CHECKPOINT" ]; then
|
| 82 |
+
EXTRA_AFFINITY_ARGS="$EXTRA_AFFINITY_ARGS --peptiverse_affinity_checkpoint $PEPTIVERSE_AFFINITY_CHECKPOINT"
|
| 83 |
+
fi
|
| 84 |
+
if [ -n "$PEPTIVERSE_CACHE_DIR" ]; then
|
| 85 |
+
EXTRA_AFFINITY_ARGS="$EXTRA_AFFINITY_ARGS --peptiverse_cache_dir $PEPTIVERSE_CACHE_DIR"
|
| 86 |
+
fi
|
| 87 |
+
if [ "$PEPTIVERSE_LOCAL_FILES_ONLY" -eq 1 ]; then
|
| 88 |
+
EXTRA_AFFINITY_ARGS="$EXTRA_AFFINITY_ARGS --peptiverse_local_files_only"
|
| 89 |
+
fi
|
| 90 |
if [ "$ORACLE_ESM_LOCAL_FILES_ONLY" -eq 1 ]; then
|
| 91 |
EXTRA_ORACLE_ARGS="$EXTRA_ORACLE_ARGS --direction_oracle_esm_local_files_only"
|
| 92 |
fi
|
|
|
|
| 124 |
--pretrained_checkpoint ${PRETRAINED_CHECKPOINT} \
|
| 125 |
--run_name ${RUN_NAME} \
|
| 126 |
--device ${DEVICE} \
|
| 127 |
+
--affinity_backend ${AFFINITY_BACKEND} \
|
| 128 |
+
${EXTRA_AFFINITY_ARGS} \
|
| 129 |
\
|
| 130 |
--targets_per_mcts ${TARGETS_PER_MCTS} \
|
| 131 |
--resample_targets_every ${RESAMPLE_TARGETS_EVERY} \
|
scoring/affinity_config.py
ADDED
|
@@ -0,0 +1,78 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""Shared CLI and construction helpers for binding-affinity backends."""
|
| 2 |
+
|
| 3 |
+
from scoring.functions.binding import create_multi_target_affinity_predictor
|
| 4 |
+
|
| 5 |
+
|
| 6 |
+
def add_affinity_arguments(parser, default_backend="original"):
|
| 7 |
+
"""Add the common affinity-backend options to an argument parser."""
|
| 8 |
+
group = parser.add_argument_group("Binding Affinity")
|
| 9 |
+
group.add_argument(
|
| 10 |
+
"--affinity_backend",
|
| 11 |
+
choices=("original", "peptiverse"),
|
| 12 |
+
default=default_backend,
|
| 13 |
+
help="Affinity predictor to use; the original TD3B model remains the default.",
|
| 14 |
+
)
|
| 15 |
+
group.add_argument(
|
| 16 |
+
"--peptiverse_affinity_checkpoint",
|
| 17 |
+
default=None,
|
| 18 |
+
help=(
|
| 19 |
+
"Optional local PeptiVerse pooled SMILES-affinity checkpoint. "
|
| 20 |
+
"If omitted, it is downloaded from the Hub."
|
| 21 |
+
),
|
| 22 |
+
)
|
| 23 |
+
group.add_argument(
|
| 24 |
+
"--peptiverse_repo_id",
|
| 25 |
+
default=(
|
| 26 |
+
"ChatterjeeLab/PeptiVerse" if default_backend is not None else None
|
| 27 |
+
),
|
| 28 |
+
help="Hugging Face repository used to download the PeptiVerse checkpoint.",
|
| 29 |
+
)
|
| 30 |
+
group.add_argument(
|
| 31 |
+
"--peptiverse_revision",
|
| 32 |
+
default=None,
|
| 33 |
+
help="Optional Hugging Face revision for reproducible model downloads.",
|
| 34 |
+
)
|
| 35 |
+
group.add_argument(
|
| 36 |
+
"--peptiverse_cache_dir",
|
| 37 |
+
default=None,
|
| 38 |
+
help="Optional cache directory for PeptiVerse and encoder artifacts.",
|
| 39 |
+
)
|
| 40 |
+
group.add_argument(
|
| 41 |
+
"--peptiverse_local_files_only",
|
| 42 |
+
action="store_true",
|
| 43 |
+
default=False if default_backend is not None else None,
|
| 44 |
+
help="Require all PeptiVerse and encoder artifacts to exist locally.",
|
| 45 |
+
)
|
| 46 |
+
group.add_argument(
|
| 47 |
+
"--peptiverse_batch_size",
|
| 48 |
+
type=int,
|
| 49 |
+
default=32 if default_backend is not None else None,
|
| 50 |
+
help="Binder-SMILES embedding batch size for PeptiVerse scoring.",
|
| 51 |
+
)
|
| 52 |
+
return parser
|
| 53 |
+
|
| 54 |
+
|
| 55 |
+
def create_affinity_from_args(args, tokenizer, base_path, device, emb_model=None):
|
| 56 |
+
"""Build the selected multi-target affinity predictor from CLI/config values."""
|
| 57 |
+
return create_multi_target_affinity_predictor(
|
| 58 |
+
backend=getattr(args, "affinity_backend", None) or "original",
|
| 59 |
+
tokenizer=tokenizer,
|
| 60 |
+
base_path=base_path,
|
| 61 |
+
device=device,
|
| 62 |
+
emb_model=emb_model,
|
| 63 |
+
peptiverse_checkpoint=getattr(
|
| 64 |
+
args, "peptiverse_affinity_checkpoint", None
|
| 65 |
+
),
|
| 66 |
+
peptiverse_repo_id=(
|
| 67 |
+
getattr(args, "peptiverse_repo_id", None)
|
| 68 |
+
or "ChatterjeeLab/PeptiVerse"
|
| 69 |
+
),
|
| 70 |
+
peptiverse_revision=getattr(args, "peptiverse_revision", None),
|
| 71 |
+
peptiverse_cache_dir=getattr(args, "peptiverse_cache_dir", None),
|
| 72 |
+
peptiverse_local_files_only=bool(
|
| 73 |
+
getattr(args, "peptiverse_local_files_only", False)
|
| 74 |
+
),
|
| 75 |
+
peptiverse_batch_size=(
|
| 76 |
+
getattr(args, "peptiverse_batch_size", None) or 32
|
| 77 |
+
),
|
| 78 |
+
)
|
scoring/functions/binding.py
CHANGED
|
@@ -1,5 +1,6 @@
|
|
| 1 |
import sys
|
| 2 |
import os, torch
|
|
|
|
| 3 |
import numpy as np
|
| 4 |
import torch
|
| 5 |
import pandas as pd
|
|
@@ -7,6 +8,15 @@ import torch.nn as nn
|
|
| 7 |
import esm
|
| 8 |
from transformers import AutoModelForMaskedLM
|
| 9 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 10 |
|
| 11 |
def _sanitize_token_ids(input_ids: torch.Tensor, vocab_size: int, unk_id: int) -> torch.Tensor:
|
| 12 |
if vocab_size <= 0 or input_ids.numel() == 0:
|
|
@@ -154,7 +164,8 @@ class BindingAffinity:
|
|
| 154 |
self.max_pep_len = getattr(self.pep_model.config, "max_position_embeddings", None)
|
| 155 |
|
| 156 |
self.model = ImprovedBindingPredictor().to(self.device)
|
| 157 |
-
|
|
|
|
| 158 |
map_location=self.device,
|
| 159 |
weights_only=False)
|
| 160 |
_bind_sd = checkpoint.get('model_state_dict', checkpoint.get('state_dict', checkpoint)) \
|
|
@@ -270,7 +281,8 @@ class MultiTargetBindingAffinity:
|
|
| 270 |
|
| 271 |
# Binding affinity prediction model
|
| 272 |
self.model = ImprovedBindingPredictor().to(self.device)
|
| 273 |
-
|
|
|
|
| 274 |
map_location=self.device,
|
| 275 |
weights_only=False)
|
| 276 |
_bind_sd = checkpoint.get('model_state_dict', checkpoint.get('state_dict', checkpoint)) \
|
|
@@ -450,7 +462,7 @@ class TargetSpecificBindingAffinity:
|
|
| 450 |
where only peptide sequences need to be provided.
|
| 451 |
"""
|
| 452 |
|
| 453 |
-
def __init__(self, multi_target_predictor
|
| 454 |
"""
|
| 455 |
Create a target-specific binding affinity predictor.
|
| 456 |
|
|
@@ -484,3 +496,42 @@ class TargetSpecificBindingAffinity:
|
|
| 484 |
List of binding affinity scores
|
| 485 |
"""
|
| 486 |
return self.forward(input_seqs)
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
import sys
|
| 2 |
import os, torch
|
| 3 |
+
from pathlib import Path
|
| 4 |
import numpy as np
|
| 5 |
import torch
|
| 6 |
import pandas as pd
|
|
|
|
| 8 |
import esm
|
| 9 |
from transformers import AutoModelForMaskedLM
|
| 10 |
|
| 11 |
+
from scoring.functions.peptiverse_binding import PeptiVerseBindingAffinity
|
| 12 |
+
|
| 13 |
+
|
| 14 |
+
def _resolve_project_root(base_path):
|
| 15 |
+
base_path = Path(base_path)
|
| 16 |
+
if (base_path / "scoring" / "functions" / "classifiers").is_dir():
|
| 17 |
+
return base_path
|
| 18 |
+
return base_path / "tr2d2-pep"
|
| 19 |
+
|
| 20 |
|
| 21 |
def _sanitize_token_ids(input_ids: torch.Tensor, vocab_size: int, unk_id: int) -> torch.Tensor:
|
| 22 |
if vocab_size <= 0 or input_ids.numel() == 0:
|
|
|
|
| 164 |
self.max_pep_len = getattr(self.pep_model.config, "max_position_embeddings", None)
|
| 165 |
|
| 166 |
self.model = ImprovedBindingPredictor().to(self.device)
|
| 167 |
+
project_root = _resolve_project_root(base_path)
|
| 168 |
+
checkpoint = torch.load(project_root / 'scoring/functions/classifiers/binding-affinity.pt',
|
| 169 |
map_location=self.device,
|
| 170 |
weights_only=False)
|
| 171 |
_bind_sd = checkpoint.get('model_state_dict', checkpoint.get('state_dict', checkpoint)) \
|
|
|
|
| 281 |
|
| 282 |
# Binding affinity prediction model
|
| 283 |
self.model = ImprovedBindingPredictor().to(self.device)
|
| 284 |
+
project_root = _resolve_project_root(base_path)
|
| 285 |
+
checkpoint = torch.load(project_root / 'scoring/functions/classifiers/binding-affinity.pt',
|
| 286 |
map_location=self.device,
|
| 287 |
weights_only=False)
|
| 288 |
_bind_sd = checkpoint.get('model_state_dict', checkpoint.get('state_dict', checkpoint)) \
|
|
|
|
| 462 |
where only peptide sequences need to be provided.
|
| 463 |
"""
|
| 464 |
|
| 465 |
+
def __init__(self, multi_target_predictor, prot_seq: str):
|
| 466 |
"""
|
| 467 |
Create a target-specific binding affinity predictor.
|
| 468 |
|
|
|
|
| 496 |
List of binding affinity scores
|
| 497 |
"""
|
| 498 |
return self.forward(input_seqs)
|
| 499 |
+
|
| 500 |
+
|
| 501 |
+
def create_multi_target_affinity_predictor(
|
| 502 |
+
backend="original",
|
| 503 |
+
tokenizer=None,
|
| 504 |
+
base_path=None,
|
| 505 |
+
device=None,
|
| 506 |
+
emb_model=None,
|
| 507 |
+
peptiverse_checkpoint=None,
|
| 508 |
+
peptiverse_repo_id="ChatterjeeLab/PeptiVerse",
|
| 509 |
+
peptiverse_revision=None,
|
| 510 |
+
peptiverse_cache_dir=None,
|
| 511 |
+
peptiverse_local_files_only=False,
|
| 512 |
+
peptiverse_batch_size=32,
|
| 513 |
+
):
|
| 514 |
+
"""Create the original TD3B or PeptiVerse affinity backend."""
|
| 515 |
+
backend = str(backend).lower()
|
| 516 |
+
if backend == "original":
|
| 517 |
+
if tokenizer is None or base_path is None:
|
| 518 |
+
raise ValueError("The original affinity backend requires tokenizer and base_path.")
|
| 519 |
+
return MultiTargetBindingAffinity(
|
| 520 |
+
tokenizer=tokenizer,
|
| 521 |
+
base_path=base_path,
|
| 522 |
+
device=device,
|
| 523 |
+
emb_model=emb_model,
|
| 524 |
+
)
|
| 525 |
+
if backend == "peptiverse":
|
| 526 |
+
return PeptiVerseBindingAffinity(
|
| 527 |
+
device=device,
|
| 528 |
+
checkpoint_path=peptiverse_checkpoint,
|
| 529 |
+
repo_id=peptiverse_repo_id,
|
| 530 |
+
revision=peptiverse_revision,
|
| 531 |
+
cache_dir=peptiverse_cache_dir,
|
| 532 |
+
local_files_only=peptiverse_local_files_only,
|
| 533 |
+
batch_size=peptiverse_batch_size,
|
| 534 |
+
)
|
| 535 |
+
raise ValueError(
|
| 536 |
+
f"Unknown affinity backend {backend!r}; choose 'original' or 'peptiverse'."
|
| 537 |
+
)
|
scoring/functions/peptiverse_binding.py
ADDED
|
@@ -0,0 +1,274 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""PeptiVerse binding-affinity adapter for TD3B.
|
| 2 |
+
|
| 3 |
+
This module implements the pooled target-sequence/binder-SMILES model published
|
| 4 |
+
in ChatterjeeLab/PeptiVerse without loading PeptiVerse's unrelated predictors.
|
| 5 |
+
"""
|
| 6 |
+
|
| 7 |
+
import logging
|
| 8 |
+
from pathlib import Path
|
| 9 |
+
from typing import Dict, List, Optional
|
| 10 |
+
|
| 11 |
+
import torch
|
| 12 |
+
import torch.nn as nn
|
| 13 |
+
|
| 14 |
+
logger = logging.getLogger(__name__)
|
| 15 |
+
|
| 16 |
+
DEFAULT_REPO_ID = "ChatterjeeLab/PeptiVerse"
|
| 17 |
+
DEFAULT_CHECKPOINT_FILE = (
|
| 18 |
+
"training_classifiers/binding_affinity/"
|
| 19 |
+
"chemberta_smiles_pooled/best_model.pt"
|
| 20 |
+
)
|
| 21 |
+
DEFAULT_ESM_MODEL = "facebook/esm2_t33_650M_UR50D"
|
| 22 |
+
DEFAULT_CHEMBERTA_MODEL = "DeepChem/ChemBERTa-77M-MLM"
|
| 23 |
+
|
| 24 |
+
|
| 25 |
+
class PeptiVersePooledAffinityModel(nn.Module):
|
| 26 |
+
"""PeptiVerse's bidirectional cross-attention affinity head."""
|
| 27 |
+
|
| 28 |
+
def __init__(
|
| 29 |
+
self,
|
| 30 |
+
target_dim: int,
|
| 31 |
+
binder_dim: int,
|
| 32 |
+
hidden_dim: int,
|
| 33 |
+
n_heads: int,
|
| 34 |
+
n_layers: int,
|
| 35 |
+
dropout: float,
|
| 36 |
+
) -> None:
|
| 37 |
+
super().__init__()
|
| 38 |
+
self.t_proj = nn.Sequential(
|
| 39 |
+
nn.Linear(target_dim, hidden_dim), nn.LayerNorm(hidden_dim)
|
| 40 |
+
)
|
| 41 |
+
self.b_proj = nn.Sequential(
|
| 42 |
+
nn.Linear(binder_dim, hidden_dim), nn.LayerNorm(hidden_dim)
|
| 43 |
+
)
|
| 44 |
+
self.layers = nn.ModuleList()
|
| 45 |
+
for _ in range(n_layers):
|
| 46 |
+
self.layers.append(
|
| 47 |
+
nn.ModuleDict(
|
| 48 |
+
{
|
| 49 |
+
"attn_tb": nn.MultiheadAttention(
|
| 50 |
+
hidden_dim, n_heads, dropout=dropout
|
| 51 |
+
),
|
| 52 |
+
"attn_bt": nn.MultiheadAttention(
|
| 53 |
+
hidden_dim, n_heads, dropout=dropout
|
| 54 |
+
),
|
| 55 |
+
"n1t": nn.LayerNorm(hidden_dim),
|
| 56 |
+
"n2t": nn.LayerNorm(hidden_dim),
|
| 57 |
+
"n1b": nn.LayerNorm(hidden_dim),
|
| 58 |
+
"n2b": nn.LayerNorm(hidden_dim),
|
| 59 |
+
"fft": nn.Sequential(
|
| 60 |
+
nn.Linear(hidden_dim, 4 * hidden_dim),
|
| 61 |
+
nn.GELU(),
|
| 62 |
+
nn.Dropout(dropout),
|
| 63 |
+
nn.Linear(4 * hidden_dim, hidden_dim),
|
| 64 |
+
),
|
| 65 |
+
"ffb": nn.Sequential(
|
| 66 |
+
nn.Linear(hidden_dim, 4 * hidden_dim),
|
| 67 |
+
nn.GELU(),
|
| 68 |
+
nn.Dropout(dropout),
|
| 69 |
+
nn.Linear(4 * hidden_dim, hidden_dim),
|
| 70 |
+
),
|
| 71 |
+
}
|
| 72 |
+
)
|
| 73 |
+
)
|
| 74 |
+
|
| 75 |
+
self.shared = nn.Sequential(
|
| 76 |
+
nn.Linear(2 * hidden_dim, hidden_dim),
|
| 77 |
+
nn.GELU(),
|
| 78 |
+
nn.Dropout(dropout),
|
| 79 |
+
)
|
| 80 |
+
self.reg = nn.Linear(hidden_dim, 1)
|
| 81 |
+
self.cls = nn.Linear(hidden_dim, 3)
|
| 82 |
+
|
| 83 |
+
def forward(self, target: torch.Tensor, binder: torch.Tensor):
|
| 84 |
+
target = self.t_proj(target).unsqueeze(0)
|
| 85 |
+
binder = self.b_proj(binder).unsqueeze(0)
|
| 86 |
+
for layer in self.layers:
|
| 87 |
+
target_attn, _ = layer["attn_tb"](target, binder, binder)
|
| 88 |
+
target = layer["n1t"]((target + target_attn).transpose(0, 1)).transpose(0, 1)
|
| 89 |
+
target = layer["n2t"](
|
| 90 |
+
(target + layer["fft"](target)).transpose(0, 1)
|
| 91 |
+
).transpose(0, 1)
|
| 92 |
+
|
| 93 |
+
binder_attn, _ = layer["attn_bt"](binder, target, target)
|
| 94 |
+
binder = layer["n1b"]((binder + binder_attn).transpose(0, 1)).transpose(0, 1)
|
| 95 |
+
binder = layer["n2b"](
|
| 96 |
+
(binder + layer["ffb"](binder)).transpose(0, 1)
|
| 97 |
+
).transpose(0, 1)
|
| 98 |
+
|
| 99 |
+
hidden = self.shared(torch.cat([target[0], binder[0]], dim=-1))
|
| 100 |
+
return self.reg(hidden).squeeze(-1), self.cls(hidden)
|
| 101 |
+
|
| 102 |
+
|
| 103 |
+
class PeptiVerseBindingAffinity:
|
| 104 |
+
"""Score target/peptide-SMILES pairs with PeptiVerse's pK regressor."""
|
| 105 |
+
|
| 106 |
+
backend_name = "peptiverse"
|
| 107 |
+
|
| 108 |
+
def __init__(
|
| 109 |
+
self,
|
| 110 |
+
device=None,
|
| 111 |
+
checkpoint_path: Optional[str] = None,
|
| 112 |
+
repo_id: str = DEFAULT_REPO_ID,
|
| 113 |
+
revision: Optional[str] = None,
|
| 114 |
+
cache_dir: Optional[str] = None,
|
| 115 |
+
local_files_only: bool = False,
|
| 116 |
+
esm_name: str = DEFAULT_ESM_MODEL,
|
| 117 |
+
chemberta_name: str = DEFAULT_CHEMBERTA_MODEL,
|
| 118 |
+
batch_size: int = 32,
|
| 119 |
+
max_protein_length: int = 1022,
|
| 120 |
+
max_smiles_length: int = 512,
|
| 121 |
+
) -> None:
|
| 122 |
+
from transformers import AutoModel, AutoTokenizer, EsmModel, EsmTokenizer
|
| 123 |
+
|
| 124 |
+
self.device = torch.device(
|
| 125 |
+
"cuda" if torch.cuda.is_available() else "cpu"
|
| 126 |
+
) if device is None else torch.device(device)
|
| 127 |
+
self.batch_size = max(1, int(batch_size))
|
| 128 |
+
self.max_protein_length = max_protein_length
|
| 129 |
+
self.max_smiles_length = max_smiles_length
|
| 130 |
+
|
| 131 |
+
resolved_checkpoint = self._resolve_checkpoint(
|
| 132 |
+
checkpoint_path=checkpoint_path,
|
| 133 |
+
repo_id=repo_id,
|
| 134 |
+
revision=revision,
|
| 135 |
+
cache_dir=cache_dir,
|
| 136 |
+
local_files_only=local_files_only,
|
| 137 |
+
)
|
| 138 |
+
checkpoint = torch.load(
|
| 139 |
+
resolved_checkpoint, map_location=self.device, weights_only=False
|
| 140 |
+
)
|
| 141 |
+
if checkpoint.get("mode") != "pooled":
|
| 142 |
+
raise ValueError(
|
| 143 |
+
f"Expected a pooled PeptiVerse checkpoint, got {checkpoint.get('mode')!r}"
|
| 144 |
+
)
|
| 145 |
+
|
| 146 |
+
state_dict = checkpoint["state_dict"]
|
| 147 |
+
params = checkpoint.get("best_params", {})
|
| 148 |
+
model = PeptiVersePooledAffinityModel(
|
| 149 |
+
target_dim=int(state_dict["t_proj.0.weight"].shape[1]),
|
| 150 |
+
binder_dim=int(state_dict["b_proj.0.weight"].shape[1]),
|
| 151 |
+
hidden_dim=int(params.get("hidden_dim", state_dict["t_proj.0.weight"].shape[0])),
|
| 152 |
+
n_heads=int(params.get("n_heads", 4)),
|
| 153 |
+
n_layers=int(params.get("n_layers", self._infer_layers(state_dict))),
|
| 154 |
+
dropout=float(params.get("dropout", 0.0)),
|
| 155 |
+
)
|
| 156 |
+
model.load_state_dict(state_dict, strict=True)
|
| 157 |
+
self.model = model.to(self.device).eval()
|
| 158 |
+
|
| 159 |
+
model_kwargs = {
|
| 160 |
+
"cache_dir": cache_dir,
|
| 161 |
+
"local_files_only": local_files_only,
|
| 162 |
+
}
|
| 163 |
+
self.target_tokenizer = EsmTokenizer.from_pretrained(esm_name, **model_kwargs)
|
| 164 |
+
self.target_encoder = EsmModel.from_pretrained(
|
| 165 |
+
esm_name, add_pooling_layer=False, **model_kwargs
|
| 166 |
+
).to(self.device).eval()
|
| 167 |
+
self.binder_tokenizer = AutoTokenizer.from_pretrained(
|
| 168 |
+
chemberta_name, **model_kwargs
|
| 169 |
+
)
|
| 170 |
+
self.binder_encoder = AutoModel.from_pretrained(
|
| 171 |
+
chemberta_name, **model_kwargs
|
| 172 |
+
).to(self.device).eval()
|
| 173 |
+
self.target_cache: Dict[str, torch.Tensor] = {}
|
| 174 |
+
logger.info("Loaded PeptiVerse affinity checkpoint: %s", resolved_checkpoint)
|
| 175 |
+
|
| 176 |
+
@staticmethod
|
| 177 |
+
def _resolve_checkpoint(
|
| 178 |
+
checkpoint_path: Optional[str],
|
| 179 |
+
repo_id: str,
|
| 180 |
+
revision: Optional[str],
|
| 181 |
+
cache_dir: Optional[str],
|
| 182 |
+
local_files_only: bool,
|
| 183 |
+
) -> str:
|
| 184 |
+
if checkpoint_path is not None:
|
| 185 |
+
path = Path(checkpoint_path).expanduser()
|
| 186 |
+
if not path.is_file():
|
| 187 |
+
raise FileNotFoundError(f"PeptiVerse checkpoint not found: {path}")
|
| 188 |
+
return str(path)
|
| 189 |
+
|
| 190 |
+
from huggingface_hub import hf_hub_download
|
| 191 |
+
|
| 192 |
+
return hf_hub_download(
|
| 193 |
+
repo_id=repo_id,
|
| 194 |
+
filename=DEFAULT_CHECKPOINT_FILE,
|
| 195 |
+
revision=revision,
|
| 196 |
+
cache_dir=cache_dir,
|
| 197 |
+
local_files_only=local_files_only,
|
| 198 |
+
)
|
| 199 |
+
|
| 200 |
+
@staticmethod
|
| 201 |
+
def _infer_layers(state_dict) -> int:
|
| 202 |
+
layer_ids = {
|
| 203 |
+
int(key.split(".")[1])
|
| 204 |
+
for key in state_dict
|
| 205 |
+
if key.startswith("layers.")
|
| 206 |
+
}
|
| 207 |
+
return max(layer_ids) + 1
|
| 208 |
+
|
| 209 |
+
@staticmethod
|
| 210 |
+
def _special_token_ids(tokenizer) -> List[int]:
|
| 211 |
+
values = [
|
| 212 |
+
getattr(tokenizer, f"{name}_token_id", None)
|
| 213 |
+
for name in ("pad", "cls", "sep", "bos", "eos", "mask")
|
| 214 |
+
]
|
| 215 |
+
return sorted({int(value) for value in values if value is not None})
|
| 216 |
+
|
| 217 |
+
@torch.no_grad()
|
| 218 |
+
def _pool(self, texts, tokenizer, encoder, max_length: int) -> torch.Tensor:
|
| 219 |
+
tokens = tokenizer(
|
| 220 |
+
list(texts),
|
| 221 |
+
return_tensors="pt",
|
| 222 |
+
padding=True,
|
| 223 |
+
truncation=True,
|
| 224 |
+
max_length=max_length,
|
| 225 |
+
)
|
| 226 |
+
tokens = {name: value.to(self.device) for name, value in tokens.items()}
|
| 227 |
+
attention_mask = tokens.get(
|
| 228 |
+
"attention_mask", torch.ones_like(tokens["input_ids"])
|
| 229 |
+
).bool()
|
| 230 |
+
valid_mask = attention_mask
|
| 231 |
+
special_ids = self._special_token_ids(tokenizer)
|
| 232 |
+
if special_ids:
|
| 233 |
+
special = torch.tensor(special_ids, device=self.device)
|
| 234 |
+
valid_mask = valid_mask & ~torch.isin(tokens["input_ids"], special)
|
| 235 |
+
hidden = encoder(**tokens).last_hidden_state
|
| 236 |
+
weights = valid_mask.unsqueeze(-1).to(hidden.dtype)
|
| 237 |
+
return (hidden * weights).sum(dim=1) / weights.sum(dim=1).clamp_min(1.0)
|
| 238 |
+
|
| 239 |
+
def get_protein_embedding(self, prot_seq: str) -> torch.Tensor:
|
| 240 |
+
prot_seq = prot_seq.strip()
|
| 241 |
+
if prot_seq not in self.target_cache:
|
| 242 |
+
self.target_cache[prot_seq] = self._pool(
|
| 243 |
+
[prot_seq],
|
| 244 |
+
self.target_tokenizer,
|
| 245 |
+
self.target_encoder,
|
| 246 |
+
self.max_protein_length,
|
| 247 |
+
)
|
| 248 |
+
return self.target_cache[prot_seq]
|
| 249 |
+
|
| 250 |
+
@torch.no_grad()
|
| 251 |
+
def forward(self, input_seqs, prot_seq: str):
|
| 252 |
+
input_seqs = list(input_seqs)
|
| 253 |
+
if not input_seqs:
|
| 254 |
+
return []
|
| 255 |
+
|
| 256 |
+
target = self.get_protein_embedding(prot_seq)
|
| 257 |
+
scores = []
|
| 258 |
+
for start in range(0, len(input_seqs), self.batch_size):
|
| 259 |
+
batch = input_seqs[start:start + self.batch_size]
|
| 260 |
+
binder = self._pool(
|
| 261 |
+
batch,
|
| 262 |
+
self.binder_tokenizer,
|
| 263 |
+
self.binder_encoder,
|
| 264 |
+
self.max_smiles_length,
|
| 265 |
+
)
|
| 266 |
+
affinity, _ = self.model(target.expand(len(batch), -1), binder)
|
| 267 |
+
scores.extend(affinity.detach().cpu().tolist())
|
| 268 |
+
return scores
|
| 269 |
+
|
| 270 |
+
def __call__(self, input_seqs, prot_seq: str):
|
| 271 |
+
return self.forward(input_seqs, prot_seq)
|
| 272 |
+
|
| 273 |
+
def clear_cache(self) -> None:
|
| 274 |
+
self.target_cache.clear()
|
td3b/__init__.py
CHANGED
|
@@ -7,7 +7,6 @@ from .direction_oracle import DirectionalOracle
|
|
| 7 |
from .td3b_scoring import TD3BRewardFunction, TD3BConfidenceWeighting, create_td3b_reward_function
|
| 8 |
from .td3b_losses import ContrastiveLoss, InfoNCELoss, TD3BTotalLoss, extract_embeddings_from_mdlm
|
| 9 |
from .td3b_mcts import TD3B_MCTS, create_td3b_mcts
|
| 10 |
-
from .td3b_finetune import td3b_finetune, add_td3b_sampling_to_model
|
| 11 |
from .data_utils import TD3BDataset, load_td3b_data
|
| 12 |
|
| 13 |
__all__ = [
|
|
@@ -28,3 +27,15 @@ __all__ = [
|
|
| 28 |
]
|
| 29 |
|
| 30 |
__version__ = '0.1.0'
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 7 |
from .td3b_scoring import TD3BRewardFunction, TD3BConfidenceWeighting, create_td3b_reward_function
|
| 8 |
from .td3b_losses import ContrastiveLoss, InfoNCELoss, TD3BTotalLoss, extract_embeddings_from_mdlm
|
| 9 |
from .td3b_mcts import TD3B_MCTS, create_td3b_mcts
|
|
|
|
| 10 |
from .data_utils import TD3BDataset, load_td3b_data
|
| 11 |
|
| 12 |
__all__ = [
|
|
|
|
| 27 |
]
|
| 28 |
|
| 29 |
__version__ = '0.1.0'
|
| 30 |
+
|
| 31 |
+
|
| 32 |
+
def __getattr__(name):
|
| 33 |
+
"""Load fine-tuning helpers lazily to avoid a package import cycle."""
|
| 34 |
+
if name in {'td3b_finetune', 'add_td3b_sampling_to_model'}:
|
| 35 |
+
from .td3b_finetune import add_td3b_sampling_to_model, td3b_finetune
|
| 36 |
+
|
| 37 |
+
return {
|
| 38 |
+
'td3b_finetune': td3b_finetune,
|
| 39 |
+
'add_td3b_sampling_to_model': add_td3b_sampling_to_model,
|
| 40 |
+
}[name]
|
| 41 |
+
raise AttributeError(f"module {__name__!r} has no attribute {name!r}")
|
tests/test_peptiverse_binding.py
ADDED
|
@@ -0,0 +1,124 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
import unittest
|
| 2 |
+
from types import SimpleNamespace
|
| 3 |
+
from unittest.mock import patch
|
| 4 |
+
|
| 5 |
+
import torch
|
| 6 |
+
|
| 7 |
+
from scoring.functions.peptiverse_binding import (
|
| 8 |
+
PeptiVerseBindingAffinity,
|
| 9 |
+
PeptiVersePooledAffinityModel,
|
| 10 |
+
)
|
| 11 |
+
|
| 12 |
+
|
| 13 |
+
class DummyTokenizer:
|
| 14 |
+
pad_token_id = 0
|
| 15 |
+
cls_token_id = 1
|
| 16 |
+
eos_token_id = 2
|
| 17 |
+
sep_token_id = None
|
| 18 |
+
bos_token_id = None
|
| 19 |
+
mask_token_id = None
|
| 20 |
+
|
| 21 |
+
def __call__(self, texts, **kwargs):
|
| 22 |
+
del texts, kwargs
|
| 23 |
+
return {
|
| 24 |
+
"input_ids": torch.tensor([[1, 3, 4, 2, 0]]),
|
| 25 |
+
"attention_mask": torch.tensor([[1, 1, 1, 1, 0]]),
|
| 26 |
+
}
|
| 27 |
+
|
| 28 |
+
|
| 29 |
+
class DummyEncoder(torch.nn.Module):
|
| 30 |
+
def forward(self, input_ids, attention_mask):
|
| 31 |
+
del attention_mask
|
| 32 |
+
hidden = input_ids.float().unsqueeze(-1).repeat(1, 1, 2)
|
| 33 |
+
return SimpleNamespace(last_hidden_state=hidden)
|
| 34 |
+
|
| 35 |
+
|
| 36 |
+
class DummyAffinityHead(torch.nn.Module):
|
| 37 |
+
def forward(self, target, binder):
|
| 38 |
+
del target
|
| 39 |
+
return binder.sum(dim=-1), torch.zeros(len(binder), 3)
|
| 40 |
+
|
| 41 |
+
|
| 42 |
+
class PeptiVerseBindingTests(unittest.TestCase):
|
| 43 |
+
def test_affinity_head_shapes(self):
|
| 44 |
+
model = PeptiVersePooledAffinityModel(
|
| 45 |
+
target_dim=8,
|
| 46 |
+
binder_dim=6,
|
| 47 |
+
hidden_dim=12,
|
| 48 |
+
n_heads=3,
|
| 49 |
+
n_layers=2,
|
| 50 |
+
dropout=0.0,
|
| 51 |
+
).eval()
|
| 52 |
+
affinity, classes = model(torch.randn(4, 8), torch.randn(4, 6))
|
| 53 |
+
self.assertEqual(tuple(affinity.shape), (4,))
|
| 54 |
+
self.assertEqual(tuple(classes.shape), (4, 3))
|
| 55 |
+
|
| 56 |
+
def test_pool_excludes_special_tokens(self):
|
| 57 |
+
predictor = object.__new__(PeptiVerseBindingAffinity)
|
| 58 |
+
predictor.device = torch.device("cpu")
|
| 59 |
+
pooled = predictor._pool(
|
| 60 |
+
["unused"], DummyTokenizer(), DummyEncoder(), max_length=8
|
| 61 |
+
)
|
| 62 |
+
expected = torch.tensor([[3.5, 3.5]])
|
| 63 |
+
self.assertTrue(torch.equal(pooled, expected))
|
| 64 |
+
|
| 65 |
+
def test_factory_keeps_original_as_default(self):
|
| 66 |
+
from scoring.functions import binding
|
| 67 |
+
|
| 68 |
+
original = object()
|
| 69 |
+
with patch.object(
|
| 70 |
+
binding, "MultiTargetBindingAffinity", return_value=original
|
| 71 |
+
) as constructor:
|
| 72 |
+
result = binding.create_multi_target_affinity_predictor(
|
| 73 |
+
tokenizer=object(), base_path="/tmp", device="cpu"
|
| 74 |
+
)
|
| 75 |
+
self.assertIs(result, original)
|
| 76 |
+
constructor.assert_called_once()
|
| 77 |
+
|
| 78 |
+
def test_factory_selects_peptiverse(self):
|
| 79 |
+
from scoring.functions import binding
|
| 80 |
+
|
| 81 |
+
peptiverse = object()
|
| 82 |
+
with patch.object(
|
| 83 |
+
binding, "PeptiVerseBindingAffinity", return_value=peptiverse
|
| 84 |
+
) as constructor:
|
| 85 |
+
result = binding.create_multi_target_affinity_predictor(
|
| 86 |
+
backend="peptiverse",
|
| 87 |
+
device="cpu",
|
| 88 |
+
peptiverse_checkpoint="model.pt",
|
| 89 |
+
)
|
| 90 |
+
self.assertIs(result, peptiverse)
|
| 91 |
+
constructor.assert_called_once_with(
|
| 92 |
+
device="cpu",
|
| 93 |
+
checkpoint_path="model.pt",
|
| 94 |
+
repo_id="ChatterjeeLab/PeptiVerse",
|
| 95 |
+
revision=None,
|
| 96 |
+
cache_dir=None,
|
| 97 |
+
local_files_only=False,
|
| 98 |
+
batch_size=32,
|
| 99 |
+
)
|
| 100 |
+
|
| 101 |
+
def test_forward_batches_binder_smiles(self):
|
| 102 |
+
predictor = object.__new__(PeptiVerseBindingAffinity)
|
| 103 |
+
predictor.batch_size = 2
|
| 104 |
+
predictor.binder_tokenizer = object()
|
| 105 |
+
predictor.binder_encoder = object()
|
| 106 |
+
predictor.max_smiles_length = 16
|
| 107 |
+
predictor.model = DummyAffinityHead()
|
| 108 |
+
predictor.get_protein_embedding = lambda _: torch.zeros(1, 3)
|
| 109 |
+
|
| 110 |
+
batch_sizes = []
|
| 111 |
+
|
| 112 |
+
def fake_pool(texts, tokenizer, encoder, max_length):
|
| 113 |
+
del tokenizer, encoder, max_length
|
| 114 |
+
batch_sizes.append(len(texts))
|
| 115 |
+
return torch.ones(len(texts), 2)
|
| 116 |
+
|
| 117 |
+
predictor._pool = fake_pool
|
| 118 |
+
scores = predictor.forward(["a", "b", "c", "d", "e"], "TARGET")
|
| 119 |
+
self.assertEqual(batch_sizes, [2, 2, 1])
|
| 120 |
+
self.assertEqual(scores, [2.0] * 5)
|
| 121 |
+
|
| 122 |
+
|
| 123 |
+
if __name__ == "__main__":
|
| 124 |
+
unittest.main()
|
training/finetune_utils.py
CHANGED
|
@@ -192,8 +192,13 @@ def load_tokenizer(base_path: str) -> SMILES_SPE_Tokenizer:
|
|
| 192 |
>>> tokenizer = load_tokenizer('To Be Added')
|
| 193 |
"""
|
| 194 |
base_path = Path(base_path)
|
| 195 |
-
|
| 196 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 197 |
|
| 198 |
if not vocab_path.exists():
|
| 199 |
raise FileNotFoundError(f"Vocabulary file not found: {vocab_path}")
|
|
|
|
| 192 |
>>> tokenizer = load_tokenizer('To Be Added')
|
| 193 |
"""
|
| 194 |
base_path = Path(base_path)
|
| 195 |
+
project_root = (
|
| 196 |
+
base_path
|
| 197 |
+
if (base_path / "tokenizer").is_dir()
|
| 198 |
+
else base_path / "tr2d2-pep"
|
| 199 |
+
)
|
| 200 |
+
vocab_path = project_root / "tokenizer" / "new_vocab.txt"
|
| 201 |
+
spe_path = project_root / "tokenizer" / "new_splits.txt"
|
| 202 |
|
| 203 |
if not vocab_path.exists():
|
| 204 |
raise FileNotFoundError(f"Vocabulary file not found: {vocab_path}")
|