chq1155 commited on
Commit
ee96220
·
1 Parent(s): b43a758

Add PeptiVerse affinity backend

Browse files
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.functions.binding import MultiTargetBindingAffinity, TargetSpecificBindingAffinity
 
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 = MultiTargetBindingAffinity(
 
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.functions.binding import MultiTargetBindingAffinity, TargetSpecificBindingAffinity
 
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 = MultiTargetBindingAffinity(
 
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
- if multi_target:
 
109
  if base_model is None or base_path is None:
110
- raise ValueError("base_model and base_path are required for multi-target affinity.")
111
- return MultiTargetBindingAffinity(
 
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.functions.binding import MultiTargetBindingAffinity, TargetSpecificBindingAffinity
 
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: MultiTargetBindingAffinity,
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 = MultiTargetBindingAffinity(
 
685
  tokenizer=tokenizer,
686
  base_path=args.base_path,
687
  device=device,
688
- emb_model=policy_model.backbone # Use backbone Roformer model (matches v1)
689
  )
690
- logger.info("Created multi-target binding affinity predictor")
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 = MultiTargetBindingAffinity(
 
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
- checkpoint = torch.load(f'{base_path}/scoring/functions/classifiers/binding-affinity.pt',
 
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
- checkpoint = torch.load(f'{base_path}/scoring/functions/classifiers/binding-affinity.pt',
 
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: MultiTargetBindingAffinity, prot_seq: str):
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
- vocab_path = base_path / "tokenizer" / "new_vocab.txt"
196
- spe_path = base_path / "tokenizer" / "new_splits.txt"
 
 
 
 
 
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}")