File size: 20,754 Bytes
59b9222
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
# CLAUDE.md

This file provides guidance to Claude Code (claude.ai/code) when working with code in this repository.

## What this is

TD3B (Transition-Directed Discrete Diffusion for Allosteric Binder Generation) is a sequence-based
generative framework that designs peptide binders with a **specified direction** β€” agonist or
antagonist β€” for a target protein. It extends **TR2-D2** (a masked discrete-diffusion peptide
generator: MDLM backbone + MCTS amortized finetuning) with three additions:

1. a **Direction Oracle** `f_Ο†` that predicts agonist vs. antagonist behavior,
2. a **soft binding-affinity gate** `g_ψ`, and
3. a **gated reward** `R = g_ψ Β· Οƒ(d*Β·(f_Ο† βˆ’ 0.5)/Ο„)` that steers generation toward direction `d* ∈ {+1 agonist, βˆ’1 antagonist}`.

Finetuning distills MCTS-discovered high-reward sequences into the diffusion policy via a total loss
`L = L_WDCE + λ·L_ctr + β·L_KL` (weighted denoising CE + directional contrastive + KL-to-reference).

Paper: arXiv:2605.09810 (LMRL Workshop, ICLR 2026).

## Repo layout: dev vs. OSS, and where the artifacts are

- This repo (`ChatterjeeLab/TD3B-dev`) is the **full/dev code**. The clean code-only public release is
  `ChatterjeeLab/TD3B` on HuggingFace. Keep changes runnable against that OSS release.
- **No checkpoints or data are in git.** Trained checkpoints, train/test CSVs, and generated binders
  ship as a single ~3.4 GB archive `td3b_dev_artifacts.zip` on Google Drive (link in README). Unzip at
  repo root to populate `checkpoints/`, `data/`, `scoring/functions/classifiers/`, `generated_binders/`.
- `.gitattributes` routes all weight/data extensions (`*.ckpt *.pt *.pth *.npy *.csv? *.zip ...`) through
  Git LFS. Only three XGBoost classifier JSONs (`hemolysis`, `nonfouling`, `solubility`) are in-repo;
  `binding-affinity.pt` and `permeability-xgboost.json` come from the archive.

## Environment & commands

```bash
conda env create -f env.yml   # creates env "td3b" (python 3.10, pytorch-cuda 12.1)
conda activate td3b
pip install -e .              # installs the `td3b` package (setup.py)
```

Core deps: PyTorch + Lightning 2.5.5, HuggingFace `transformers` 4.56.2, `fair-esm` 2.0.0 (ESM2),
`rdkit`, `SmilesPE`, `xgboost`, `wandb`, `hydra-core`. There is **no test suite, no linter config, and
no Makefile** β€” do not look for `pytest`/`tox`/`ruff`. Verify changes by actually running inference.

**Inference** (primary OSS entry point β€” generate directional binders):
```bash
python inference.py \
    --ckpt_path checkpoints/td3b.ckpt \
    --val_csv data/test.csv \
    --save_path results/ \
    --seed 42 --num_pool 32 --val_samples_per_target 8 --resample_alpha 0.1
```
For each target row it generates for **both** directions (agonist `d*=+1`, antagonist `d*=-1`), scores
with the oracle + affinity gate, applies softmax(reward/Ξ±) weighted resampling (Algorithm 2), keeps only
RDKit-valid peptides, and writes `results/td3b_results_seed{seed}.csv`.

**Training** (multi-target): edit paths in `launch_multi_target.sh` (`BASE_PATH`, checkpoints, data,
oracle, and `WANDB_ENTITY` β€” blank by default; set your own), then `bash launch_multi_target.sh`. It calls
`finetune_multi_target.py`. Key knobs live in the launch script: `CONTRASTIVE_WEIGHT`(Ξ»), `KL_BETA`(Ξ²),
`SIGMOID_TEMPERATURE`(Ο„), `NUM_ITER`/`NUM_CHILDREN` (MCTS), `TARGETS_PER_MCTS`(K), and the cadence flags
below.

**Baselines** (CG, SMC, TDS, PepTune, Unguided): `cd baselines && ./run.sh <csv> <baseline> <device>`.
Multi-GPU via a 5th/6th arg (`torchrun`). Note this script loads `../pretrained/peptune-pretrained.ckpt`,
**not** `checkpoints/pretrained.ckpt` (see path landmines below).

**Demo**: `notebooks/TD3B_Inference_Demo.ipynb` (Colab T4).

## Architecture β€” the cross-file big picture

The pieces below only make sense together; reading any one file in isolation misses the flow.

### Diffusion backbone β€” `models/`
- `Diffusion(L.LightningModule)` in `models/diffusion.py` is the MDLM core (absorbing-state / masked
  discrete diffusion). The denoiser is `models/roformer.py::Roformer`, a thin wrapper over HuggingFace
  `RoFormerForMaskedLM` (rotary embeddings). Tokens are **SMILES** via `tokenizer/my_tokenizers.py::SMILES_SPE_Tokenizer`.
- Absorbing state = the tokenizer's `[MASK]` id (`self.mask_index`), not a constant. Generation starts
  fully masked (`sample_prior` β†’ all-mask) and reverse-diffuses. SUBS parameterization
  (`subs_parameterization`) forbids predicting MASK and pins already-unmasked positions ("carry-over
  unmasking"). Callers use `single_reverse_step` / `single_noise_removal` (final step guarantees no
  surviving MASK); MCTS uses the `batch_mcts_reverse_step` / `mcts_reverse_step` variants which also
  return per-step policy vs. pretrained log-probs for the importance log-ratio `log_rnd`.
- Config is built by `configs/finetune_config.py::DiffusionConfig`, a **shim** that synthesizes
  duck-typed attribute objects (`type(...)()`) for backward compat. It fixes `parameterization='subs'`,
  `T=0` (continuous-time MDLM loss), `time_conditioning=False`, and `max_position_embeddings=1035`.
  It is a *partial* interface β€” only the finetune/eval/MCTS fields exist; pure-training paths
  (`antithetic_sampling`, `noise.state_dependent`, `vocab`, `model.length`) are absent and would
  `AttributeError` under it.
- `models/noise_schedule.py`: `get_noise(config)` supports geometric/loglinear/cosine/linear, but every
  reverse/sampling step **asserts `loglinear`**. A second hardcoded `LogPolyNoise` (cubic) masks
  peptide-bond tokens more slowly.

### Directional reward β€” `td3b/td3b_scoring.py`
`TD3BRewardFunction.__call__(List[str] of peptide SMILES) -> (rewards, info)`. Internals:
`g_ψ` = `scoring/functions/binding.py::BindingAffinity` (magnitude); `f_Ο†` = the `DirectionalOracle`
(direction prob ∈[0,1] + confidence ΞΊ); reward = `g_ψ Β· Οƒ(d*Β·(f_Ο†βˆ’0.5)/Ο„)`. `create_td3b_reward_function`
is the factory that builds/loads the oracle, caches the encoded protein tokens, maps `'agonist'/'antagonist'`
β†’ `d* = +1/βˆ’1`, and returns the configured reward. `TD3BConfidenceWeighting` provides the
confidence-weighted importance weights used by MCTS (`w = ΞΊΒ·exp(S/Ξ±)`).

### Direction Oracle β€” `td3b/direction_oracle.py`
`DirectionalOracle` wraps `ESM_TR2D2_GPCRClassifier`: **frozen ESM2** (`facebook/esm2_t33_650M_UR50D`,
1280-d β€” downloads from HF unless `esm_cache_dir`/`esm_local_files_only` set) encodes the **protein**;
a **frozen TR2-D2 RoFormer** encodes the **ligand/peptide SMILES**; both project to `d_model=256`, pass
1 self-attention layer each, then **2 stacked bidirectional cross-attention (BMCA) layers**, mean-pool,
concat, MLP β†’ 2 logits. `predict_with_confidence` returns `f_Ο† = p_agonist` and `ΞΊ = max(softmax)`.
Loading needs four assets: the oracle `.pt`, the TR2-D2 ligand checkpoint, and the SMILES tokenizer
vocab+splits. The RoFormer config (768/8/8/1035) is hardcoded and must match the checkpoint.

### Losses β€” `td3b/td3b_losses.py`
`TD3BTotalLoss` = `L_WDCE + λ·L_ctr + β·L_KL` (λ=`contrastive_weight`, β=`kl_beta`, both default 0.1).
`L_WDCE` is computed **externally** by `training/finetune_utils.py::loss_wdce` and passed in β€” it is the
policy-distillation term that reweights MCTS samples by `softmax(log_rnd)`. `L_ctr` is `ContrastiveLoss`
(margin, default) or `InfoNCELoss` over agonist/antagonist embeddings from `extract_embeddings_from_mdlm`
(reaches into `model.backbone.model`, RoFormer last hidden state, **must not** be under `no_grad`).
`L_KL` is per-position categorical KL to a **frozen deepcopy reference model** (the pretrained weights).

### MCTS β€” `mcts/peptide_mcts.py` + `td3b/td3b_mcts.py`
Base `MCTS` does Pareto/multi-objective tree search: root is fully masked; `select` descends via a
non-dominated (not scalar-UCB) set with `rd.choice`; `expand` samples `num_children` one-step
unmaskings, rolls each to a full sequence, filters with `PeptideAnalyzer.is_peptide`, scores valid ones,
and maintains a **Pareto buffer** of finished trajectories (each storing `x_final`, `log_rnd`,
reward, score vector). `TD3B_MCTS` subclasses it: injects the gated `TD3BRewardFunction`, pads the (N,2)
directional score vector to (N,5) so the base Pareto machinery works, folds confidence into `log_rnd`,
and `forward`/`consolidateBuffer` return **7** values (adding `directional_labels`, `confidences`).

### Training loop β€” `finetune_multi_target.py`
Self-contained inline loop (it does **not** call `td3b/td3b_finetune.py::td3b_finetune`, which is a
legacy/unused single-target loop). Per epoch it alternates:
- **MCTS generation phase** (every `resample_every_n_step` epochs, default 10): for each sampled target Γ—
  each direction, build a per-(target,direction) reward, run a fresh `TD3B_MCTS`, push Pareto survivors
  into a replay buffer with `directional_label` **forced to the intended `d*`** (not the oracle guess).
- **Gradient phase** (every epoch): shuffle the buffer, pad variable-length `x` to max-len with
  `mask_index` + build an attention mask, then WDCE + KL every batch, and contrastive only when a batch
  mixes both directions.

Three independent epoch-indexed cadence knobs: `resample_targets_every` (redraw K targets),
`resample_every_n_step` (MCTS phase), `reset_every_n_step` (reset vs. reuse the search tree).
`add_td3b_sampling_to_model(policy)` monkey-patches `sample_finetuned_td3b` onto the Diffusion instance β€”
required before validation/eval works (that method is not native to `Diffusion`).

### Data schema β€” `td3b/data_utils.py`
CSVs use columns `Target_Sequence`, `Ligand_Sequence` (peptide AA string), and `label`
(`agonist`/`antagonist`, mapped to `d*`); `TD3BDataset` also reads `Action`, `Target_UniProt_ID`,
`Ligand_UniProt_ID`. Binders are converted AA→SMILES via RDKit `MolFromSequence`. The multi-target
script uses its own in-file `TargetDataset` (groups by target, stores per-direction **median binder
length** used to set generation length), reading only `Target_Sequence`/`Ligand_Sequence`/`label`.
`inference.py` also reads `Target_UniProt_ID`, `Ligand_SMILES` (for length).

### New entry points (added 2026-07-12) β€” `finetune_on_target.py`, `generate_valid.py`
- **`finetune_on_target.py` (Function A)** β€” user-facing "bring your own target" wrapper; it does **not** reimplement training. It normalizes `--target_seq` (repeatable) / `--targets_csv` into a temp CSV (seeding a placeholder poly-G binder of `--binder_length` residues for any missing direction so `TargetDataset` can set a length prior), then **subprocess-invokes `finetune_multi_target.py`** with `--targets_per_mcts=<#targets>` + `--resample_targets_every 1` (finetune on ONLY those targets), and finally **generates in-process** reusing `inference.load_model`/`sample_sequences`/`score_sequences` + `create_td3b_reward_function` + Algorithm-2 resampling. `--direction` restricts only what is generated (finetuning always searches both). Writes `results/finetune_on_target/binders_<dir>_validity-<on|off>_seed<seed>.csv` + the finetuned ckpt under `results/<run_name>_<ts>/`. Paths default to repo root; missing heavy artifacts fail fast with a Google-Drive hint.
- **`generate_valid.py` + `sampling_strategies.py` (Function B)** β€” sampling-time validity boosters (no retraining) that reuse the model's own `sample_prior`/`single_reverse_step`/`single_noise_removal`/`forward` and change only token SELECTION (temperature β†’ top-k β†’ softmax β†’ top-p), plus a **remask self-correction loop** (remask the lowest-confidence K% of invalid sequences and re-denoise for R rounds) and a **best-of-N** validity-guided rejection wrapper β€” no diffusion math is reimplemented. `sampling_strategies.generate(...)` dispatches `baseline, more_steps, top_p(=nucleus), top_k, low_temp, remask, best_of_n, nucleus_remask` (default `nucleus_remask`). `generate_valid.py` loads the real ckpt via `inference.load_model`, or falls back to `build_random_model` (random-init `Diffusion`, CPU dev) when `--ckpt_path` is absent. Output: valid-only CSV (`idx,sequence,n_chars`) + a printed valid-yield summary.
- **`inference.py --sampler`** (opt-in; default `baseline` = original behavior, byte-for-byte) selects a
  Function-B sampling strategy for the candidate pool, with pass-through knobs
  (`--num_steps --top_p --top_k --temperature --remask_rounds --remask_frac --best_of_n`). Non-baseline
  strategies dispatch to `sampling_strategies.generate`. Function B is also available standalone via
  `generate_valid.py`.

## Landmines

A batch of path/wiring bugs that stopped the OSS release from running was fixed on 2026-07-09.
`python inference.py --help` and `python finetune_multi_target.py --help` now import cleanly (verified
in the `tr2d2-pep` conda env); a full generation run still needs the Google-Drive artifacts + a GPU.

**Fixed β€” do NOT reintroduce:**
- `td3b/td3b_finetune.py`: the `from plotting import ...` (no such module) is now `try/except`-guarded,
  and `loss_wdce` is imported **lazily inside `td3b_finetune()`** to break a `finetune_utils ↔ td3b`
  circular import. Keep both β€” `td3b/__init__.py` eagerly imports `td3b_finetune`, and `finetune_utils`
  imports the `td3b` package at module load, so any module-level `from training.finetune_utils import …`
  in `td3b_finetune.py` re-creates the cycle (it was inference.py's first crash, before plotting).
- Stale `tr2d2-pep/` prefix stripped from the tokenizer loader (`finetune_utils.load_tokenizer`), every
  classifier loader (`scoring/functions/*.py`, `scoring/scoring_functions.py`, `binding.py`), the
  `Diffusion` fallback tokenizer, and the training results dir. Assets now resolve to the README layout
  (`tokenizer/`, `scoring/functions/classifiers/`, `results/`).
- `inference.py` reward wiring rewritten: build `MultiTargetBindingAffinity` + `DirectionalOracle` once,
  wrap each target with `TargetSpecificBindingAffinity`, call `create_td3b_reward_function`. (It used to
  call a non-existent `create_reward_function` signature whose `TypeError` was swallowed β†’ empty CSV.)
- Demo notebook: `from models.diffusion import Diffusion` (was `from diffusion import …`), clone URL now
  points at the OSS HF repo with `GIT_LFS_SKIP_SMUDGE=1`, `total_memory` typo fixed.
- Added `__init__.py` to the 8 dirs `find_packages()` missed; fixed corrupted `configs/peptune_config.yaml`
  key (`batchinohup ng` β†’ `batching`).

**Fixed β€” round 2 (runtime hardening, 2026-07-10; verified dynamically in `tr2d2-pep`):**
- `inference.py` Algorithm-2 resampling: now gates candidates by `finite-reward AND valid-peptide`
  **before** sampling and draws **without replacement** (`k=min(val_samples_per_target, n_eligible)`).
  Previously `replacement=True` + the peaked softmax produced duplicate rows (inflated counts, skewed
  means) and validity was filtered only afterward (could save 0 samples despite valid candidates).
- Checkpoint-load guards (silent-random-weights class): `inference.py::load_model` now raises if the
  backbone loaded **no** weights and warns on partial loads; `direction_oracle.py::TR2D2RoFormerEncoder`
  now unwraps `model_state_dict`/`state_dict` and raises if **zero** RoFormer keys matched (was silently
  leaving the ligand encoder random); `binding.py` tolerates raw/`state_dict`/`model_state_dict` ckpt
  containers; `_load_state_dict_flexible` loudly flags missing **non-ESM** (trained) keys. All guards
  fail only on the impossible-for-a-valid-checkpoint case, so they can't break a good load.

**Must NOT change β€” tied to pretrained weights:** `max_position_embeddings=1035`,
`parameterization='subs'`, `T=0`, `time_conditioning=False`. `Diffusion.forward` hard-fails if `seq_len > 1035`.

**Remaining known quirks (not inference blockers):**
- **Checkpoint generates valid SMILES only at SHORT length (verified end-to-end 2026-07-12 on real weights
  from `/data1/hanqun/TD3B/checkpoints`).** Valid-peptide yield vs generation length (SMILES tokens),
  `is_peptide` over 32–64 samples: L=40 β†’ 28% baseline / 59% `best_of_n`; L=60 β†’ ~5–30%; **Lβ‰₯150 β†’ ~0%
  regardless of step count (128/256/512) or best_of_n.** Real 25–30-residue binders are 150–210 tokens, i.e.
  outside the valid regime β€” that is why a naive `inference.py` run writes an empty CSV. Fix applied:
  `inference.py` now derives length from the binder's *token* count (not char count; reads the `SMILES`
  column, not the nonexistent `Ligand_SMILES`), adds `--seq_length`/`--max_seq_length`, and hints on empty
  output. Working run: `--seq_length 50 --sampler best_of_n --best_of_n 6` β†’ 16 valid oracle-scored binders,
  antagonist direction-accuracy β‰ˆ 1.0. This is a model-capability limit, not a code bug.
- **ESM2 network fetch:** `facebook/esm2_t33_650M_UR50D` is downloaded at inference twice β€” by the oracle
  (`transformers`) and by `BindingAffinity`/`MultiTargetBindingAffinity` (`fair-esm`, no offline flag).
  Offline runs need warm HF + torch-hub caches.
- **Training default paths** (`finetune_multi_target.py:551`, factory fallback `td3b_scoring.py:345`) still
  resolve under `{base_path}/tr2d2-pep/...` with a wrong default oracle filename
  (`best_model_tr2d2_gpcr_fixed.pt` vs. shipped `direction_oracle.pt`). `launch_multi_target.sh` overrides
  all of these with explicit `checkpoints/` args, so training-via-launch-script works; a bare
  `python finetune_multi_target.py` does not.
- `baselines/run.sh` loads `pretrained/peptune-pretrained.ckpt` (not `checkpoints/pretrained.ckpt`) and
  defaults `CSV_PATH="To Be Added"` β€” pass real args.
- `td3b/td3b_finetune.py::td3b_finetune()` is legacy/unused (superseded by `finetune_multi_target.py`);
  still writes to `{base}/TR2-D2/tr2d2-pep/results/...` and needs the optional `plotting` module.
- `TD3BDataset`/`load_td3b_data` (oracle-training path) require CSV columns `Action`, `Ligand_UniProt_ID`
  that the inference/finetune CSVs don't carry.
- Dead/inert flags: `--contrastive_type` (loss always uses `'margin'`), `--num_epoch_for_sampling`,
  `min_affinity_threshold=0.0` (down-weight branch never fires), `use_confidence_weighting` (stored but
  never applied in `compute_gated_reward`).
- **`PeptideAnalyzer.is_peptide` (`utils/app.py`) false-positives** on atom-only / pure-amino-acid-letter
  strings (e.g. `"CCNCCF"`, `"cccccc"` β†’ True) because it checks AA-letter membership before RDKit. Harmless
  for a well-trained model (emits real peptide SMILES) but can inflate `valid_mask` for an untrained/early
  checkpoint. Shared by MCTS/baselines β€” do not change its semantics without checking all callers.
- Oracle confidence κ is actually ∈[0.5,1] (max of a 2-class softmax), not the documented [0,1]; unused at
  inference. `TD3BConfidenceWeighting.compute_importance_weights` uses raw `exp(reward/Ξ±)` which can overflow
  for large affinities β€” off the inference path (MCTS/training only).
- **Device coupling:** `models/roformer.py` and helpers pick their own device independently of `--device`
  (some hardcode `cuda:0`); `resolve_device` in `scoring_functions.py`/`direction_oracle.py` silently falls
  back to `cuda:0`/CPU. Mixed-device errors are easy to create.

**New entry points (2026-07-12) β€” usage caveats:**
- **The validity toggle is a FILTER, not a reward term.** `finetune_on_target.py --validity_reward {on,off}`
  (default `on`) never changes the reward formula (`affinity Γ— direction`); it toggles the
  `PeptideAnalyzer.is_peptide` gate on **both** halves β€” the finetune-side MCTS expansion (forwarded to
  `finetune_multi_target.py --validity_reward` β†’ `args.enforce_validity`) and the generation-side Algorithm-2
  resampling. `off` retains invalid samples (each output row still records `is_valid`).
  `--finetune_validity_hook {on,off}` decouples the finetune-side gate from the generation-side toggle.
- **Function B benchmarks were on a RANDOM-INIT model** β†’ only *relative* trends (which strategy helps most at
  long length) are meaningful; **absolute** valid-yields need the real checkpoint. `generate_valid.py`
  **silently falls back to random init** when `--ckpt_path` is missing/omitted (it prints a NOTE, but the
  numbers are garbage) β€” always pass a real `--ckpt_path` for reportable yields.
- `finetune_on_target.py` runs the finetune half as a **subprocess** (default `WANDB_MODE=disabled`) and locates
  the produced checkpoint by diffing `results/<run_name>_*` dirs before/after (newest `model_final.ckpt`, else
  newest `model_epoch_*.ckpt`); a renamed/failed run dir would break that discovery. Its own oracle/path defaults
  are passed through explicitly so the subprocess never falls back to `finetune_multi_target.py`'s legacy
  `{base}/tr2d2-pep/...` defaults.