pranamanam commited on
Commit
7efd209
·
verified ·
1 Parent(s): e5ce544

Complete rename of lectures_02_03 to lecture_3

Browse files
lectures_02_03/GUIDANCE_NOTES.md DELETED
@@ -1,156 +0,0 @@
1
- # ESM-2 flow matching and diffusion guidance
2
-
3
- Two independent teaching scripts, each with encoding, normalization, a small
4
- conditional generator, a time-conditioned reward predictor, training, three
5
- guidance examples, and constrained sequence decoding. Neither imports the other.
6
-
7
- ## Run
8
-
9
- Install Python 3.10+ and the dependencies in your own environment:
10
-
11
- ```sh
12
- pip install torch transformers==4.57.6
13
- python esm2_flow_guidance.py --epochs 200 --samples 8
14
- python esm2_diffusion_guidance.py --epochs 200 --samples 8
15
- ```
16
-
17
- Keep both scripts and `esm2_example.csv` together. Their default paths are
18
- relative to the scripts, so the commands also work with absolute script paths.
19
- The first run downloads the public `facebook/esm2_t6_8M_UR50D` checkpoint into
20
- `.esm2_cache`. CUDA is used when available; otherwise the scripts use CPU.
21
- These examples were executed with PyTorch 2.9.1 and Transformers 4.57.6.
22
-
23
- ## Motivation
24
-
25
- We will generate protein-sequence representations with a flow and a diffusion
26
- model, then guide generation toward specified properties. Using frozen ESM-2
27
- residue embeddings, we can compare classifier-free conditioning with
28
- single-objective and weighted multi-objective reward steering in the same latent
29
- space. After generation, we decode the latents into amino acids and enforce an
30
- explicit residue-count constraint. This separates learning the sequence
31
- distribution, expressing preferences, and guaranteeing a discrete output rule.
32
-
33
- ## Where the rewards come from
34
-
35
- The bundled dataset contains 64 synthetic sequences of length 24. These are
36
- teaching examples, not natural proteins or experimentally validated peptides.
37
-
38
- For sequence s of length L, `composition_proxies()` computes:
39
-
40
- - r1 = (number of K and R minus number of D and E) / L.
41
- This is a side-chain charge-count proxy, not a complete pH-dependent charge model.
42
- - r2 = number of residues in `DEHKNQRST` / L.
43
- This is a defined polar/charged composition fraction, not measured solubility.
44
- - The bundled class label c is 1 when r1 > 0, and 0 otherwise.
45
-
46
- CSV columns are `sequence,c,r1,r2`. Both classes 0 and 1 must be present. If r1
47
- and r2 are omitted, the scripts compute the composition proxies directly.
48
- Supplied r1/r2 columns override the proxy calculation, allowing measured
49
- objectives or external predictor outputs. Orient each objective so larger is
50
- better before supplying it. For an undesirable quantity, negate it first.
51
-
52
- All sequences in one run must have the same length, between 1 and 128, using
53
- the 20 canonical amino acids. The default hard minimum is 12, so shorter
54
- sequences require a smaller `--min-polar` value.
55
-
56
- ## Normalization and reward prediction
57
-
58
- - Preserve one ESM-2 embedding per residue, giving [N, L, 320]. Remove BOS/EOS.
59
- - Standardize each latent feature using its mean and standard deviation over
60
- training sequences and residue positions.
61
- - Standardize each property separately: r_tilde = (r - mean) / std.
62
- - Save these statistics. Do not refit them on generated, validation, or test data.
63
- - Train a small predictor on intermediate latents and their clean-sequence
64
- standardized property labels. Its two outputs estimate expected endpoint
65
- rewards from the current state and time.
66
-
67
- The examples use all 64 sequences for training and do not claim held-out quality.
68
- For research, split and cluster data before fitting statistics or networks, then
69
- assess prediction and sequence reconstruction on held-out sequences.
70
-
71
- ## Three sampling modes in each script
72
-
73
- 1. CFG: `w=2`, `eta=0`. Combine unconditional and class-conditioned predictions
74
- as F_uncond + w * (F_cond - F_uncond). Here w=0 is unconditional and w=1 is
75
- ordinary conditional sampling. Drop the class to null index 2 in 20% of
76
- training examples. CFG does not use the external reward predictor.
77
- 2. Single objective: `w=0`, `eta=1`, `lambdas=(1,0)`.
78
- 3. Multiple objectives: `w=0`, `eta=1`, `lambdas=(0.7,0.3)`.
79
-
80
- The weights are normalized to sum to one. Lambda controls the relative
81
- tradeoff in standardized property units; eta controls overall steering strength.
82
- Set both w and eta positive to combine CFG and reward steering. Edit the three
83
- calls in `main()` to change the demonstration settings.
84
-
85
- Flow training uses Z_0=noise, Z_1=data, a straight conditional path, and target
86
- velocity Z_1-Z_0. Sampling integrates from t=0 to t=1 with Euler steps. Reward
87
- steering adds kappa(t) times the reward gradient to the velocity, with the chosen
88
- schedule kappa(t)=4*eta*t*(1-t). This is heuristic velocity steering, not a claim
89
- of exact sampling from a reward-tilted density.
90
-
91
- DDPM training uses Z_0=data and predicts the Gaussian noise used to construct
92
- Z_k. Sampling runs k=1000 down to 1. The small noise predictor includes the
93
- Gaussian-reference skip sqrt(1-alpha_bar_k)*Z_k and learns an additive correction
94
- scaled by sqrt(alpha_bar_k). This is a parameterization of the noise predictor;
95
- the target and DDPM equations remain noise-prediction equations. It lets the
96
- small network pass through high-dimensional noise without reconstructing every
97
- coordinate through its narrow hidden layer.
98
-
99
- For DDPM steering, epsilon_guided = epsilon_CFG - eta *
100
- sqrt(1-alpha_bar_k) * gradient(R_lambda). This follows the score-to-noise
101
- conversion s = -epsilon / sqrt(1-alpha_bar_k). The posterior standard deviation
102
- used for the reverse random increment is a different quantity.
103
-
104
- In both scripts, gradient ascent uses a learned expected-reward predictor.
105
- It is not the exact log conditional likelihood or log exponential-reward
106
- expectation needed for exact conditional or reward-tilted sampling. The
107
- generator and reward weights are frozen during sampling; gradients are enabled
108
- only for the current latent. No gradients through a discrete sequence scorer
109
- are required. A research implementation must validate surrogate quality and
110
- recheck the true objectives after decoding.
111
-
112
- ## One decoder at the end
113
-
114
- The identical `decode()` function appears in both files to keep each script
115
- standalone. Both samplers use this same final decoding operation.
116
-
117
- 1. Undo latent standardization.
118
- 2. Apply ESM-2's frozen language-model head and retain the 20 amino-acid logits.
119
- 3. Take the highest-logit residue at each position.
120
- 4. If fewer than M positions contain a residue in `DEHKNQRST`, replace exactly
121
- the missing number using the lowest logit-cost changes to that set.
122
-
123
- For the default M=12, every decoded 24-residue output contains at least 12
124
- members of the selected set. `--min-polar 0` removes the constraint. The procedure
125
- maximizes the sum of the fixed per-position logits subject to this minimum
126
- count. It does not enforce the constraint along the latent trajectory, preserve
127
- an exact conditioned generative distribution, or guarantee physical solubility.
128
-
129
- The language-model head is a simple available decoder, not a mathematically
130
- exact inverse of ESM-2. Generated latents can be off the encoding manifold.
131
- For serious protein generation, validate decoding and consider training a
132
- dedicated sequence decoder. Property improvements in the surrogate need not
133
- survive discretization or the constrained substitutions.
134
-
135
- ## Outputs and checks
136
-
137
- Each script writes `cfg.fasta`, `single.fasta`, `multi.fasta`, and `results.pt`
138
- to its own output directory. The tensor file contains standardized generated
139
- latents, both networks' state dictionaries, normalization statistics, ESM model
140
- name, latent shape, and constraint settings. The terminal prints example
141
- sequences and re-evaluated mean composition proxies.
142
-
143
- Both examples have been executed through actual ESM-2 encoding, 200 training
144
- epochs, all three guidance modes, and final decoding. Checks included finite
145
- outputs, reward input gradients, normalization, schedule indexing, and the
146
- hard constraint. For small test cases, constrained decoding was compared with
147
- exhaustive search over polar/nonpolar assignments. These checks establish code
148
- mechanics, not biological validity or reliable property optimization.
149
-
150
- ## Primary references
151
-
152
- - ESM model: https://huggingface.co/facebook/esm2_t6_8M_UR50D
153
- - ESM implementation: https://huggingface.co/docs/transformers/model_doc/esm
154
- - Flow Matching: https://arxiv.org/abs/2210.02747
155
- - DDPM: https://arxiv.org/abs/2006.11239
156
- - Classifier-Free Diffusion Guidance: https://arxiv.org/abs/2207.12598
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
lectures_02_03/README.md DELETED
@@ -1,310 +0,0 @@
1
- # Lectures 2 and 3 · Continuous Generative Models
2
-
3
- These examples accompany Lectures 2 and 3 of CIS 6270. We implement flow matching and DDPM diffusion using frozen ESM-2
4
- residue embeddings, with shared property definitions and normalization to
5
- compare guidance across the two sampling processes.
6
-
7
- Each example includes data loading, latent and property normalization, a small
8
- conditional model, training, classifier-free guidance, single-objective reward
9
- steering, weighted multi-objective steering, and final sequence decoding with a
10
- hard residue-count constraint.
11
-
12
- ## Examples
13
-
14
- | Example | What the model learns | Sampling |
15
- | --- | --- | --- |
16
- | [Flow matching](esm2_flow_guidance.py) | A velocity field between noise and clean ESM-2 latents | Euler integration from noise to data |
17
- | [Diffusion](esm2_diffusion_guidance.py) | The noise added to clean ESM-2 latents | A 1,000-step DDPM reverse chain |
18
- | [Mathematical and implementation notes](GUIDANCE_NOTES.md) | Normalization, guidance conventions, decoding, and limitations | Companion reading for both scripts |
19
-
20
- Each script contains the complete training and sampling implementation,
21
- including the final sequence decoder, and can run independently.
22
-
23
- ## Quick start
24
-
25
- Use **Python 3.11**, or another compatible Python version at least 3.10. A fresh
26
- virtual environment is recommended. The dependencies are pinned to the versions
27
- used to test these examples: PyTorch 2.9.1 and Transformers 4.57.6.
28
-
29
- ### 1. Download and install
30
-
31
- ```bash
32
- git clone https://huggingface.co/ChatterjeeLab/CIS6270
33
- cd CIS6270
34
-
35
- python3 -m venv .venv
36
- source .venv/bin/activate
37
- python -m pip install --upgrade pip
38
- python -m pip install -r requirements.txt
39
- ```
40
-
41
- On Windows PowerShell, use `python -m venv .venv` and activate with
42
- `.venv\Scripts\Activate.ps1`.
43
-
44
- ### 2. Train and sample
45
-
46
- Run either example, or both:
47
-
48
- ```bash
49
- python lectures_02_03/esm2_flow_guidance.py --epochs 200 --samples 8
50
- python lectures_02_03/esm2_diffusion_guidance.py --epochs 200 --samples 8
51
- ```
52
-
53
- Each command trains a generator and a property predictor from scratch, then
54
- generates sequences using three guidance settings. The first run downloads the
55
- public [ESM-2 8M checkpoint](https://huggingface.co/facebook/esm2_t6_8M_UR50D)
56
- into `lectures_02_03/.esm2_cache/`, which subsequent runs reuse. Both the repository
57
- and checkpoint are publicly accessible.
58
-
59
- The scripts use CUDA when available and CPU otherwise. Apple Silicon currently
60
- uses the CPU path. Runtime depends on hardware and the first-run download.
61
-
62
- For a short end-to-end installation check:
63
-
64
- ```bash
65
- python lectures_02_03/esm2_flow_guidance.py --epochs 2 --samples 2 --output outputs/flow_smoke
66
- python lectures_02_03/esm2_diffusion_guidance.py --epochs 2 --samples 2 --output outputs/diffusion_smoke
67
- ```
68
-
69
- The two-epoch runs exercise encoding, training, sampling, and decoding.
70
-
71
- ### 3. Find the generated sequences
72
-
73
- By default, outputs are written beside the scripts:
74
-
75
- ```text
76
- lectures_02_03/
77
- ├── esm2_flow_outputs/
78
- │ ├── cfg.fasta
79
- │ ├── single.fasta
80
- │ ├── multi.fasta
81
- │ └── results.pt
82
- └── esm2_diffusion_outputs/
83
- ├── cfg.fasta
84
- ├── single.fasta
85
- ├── multi.fasta
86
- └── results.pt
87
- ```
88
-
89
- The FASTA files contain decoded sequences. `results.pt` contains generated
90
- standardized latents, generator and reward-predictor state dictionaries,
91
- normalization statistics, the ESM checkpoint name, latent dimensions, and
92
- constraint settings. The terminal also prints re-evaluated composition
93
- properties of the decoded sequences.
94
-
95
- Each run initializes new generator and reward-predictor parameters. Repeating
96
- a command with the same output directory replaces its saved files; choose a
97
- new `--output` directory to retain results from separate experiments.
98
-
99
- ## Implementation
100
-
101
- 1. **Encode sequences.** Use frozen ESM-2 to obtain one 320-dimensional vector
102
- per residue, retaining position information and removing BOS/EOS tokens.
103
- 2. **Normalize.** Standardize latent coordinates and each property using
104
- training-set statistics.
105
- 3. **Train.** Learn a conditional generative field and a time-conditioned
106
- predictor of the clean sequence's standardized properties.
107
- 4. **Guide generation.** Compare CFG, one reward, and a weighted reward sum.
108
- 5. **Decode once at the end.** Undo latent normalization, apply the frozen ESM-2
109
- language-model head, and enforce the selected residue-count constraint.
110
- 6. **Re-evaluate.** Calculate the composition properties on the actual decoded
111
- sequences so comparisons between guidance settings reflect the final
112
- amino-acid outputs.
113
-
114
- ## Sequence data and property labels
115
-
116
- For the teaching examples, we use [64 synthetic sequences of length
117
- 24](esm2_example.csv). We calculate both property labels directly
118
- from residue counts and assign the conditioning class according to the first
119
- property:
120
-
121
- | Column | Definition | Interpretation |
122
- | --- | --- | --- |
123
- | `sequence` | A canonical amino-acid sequence | Input to frozen ESM-2 |
124
- | `r1` | `(count(K) + count(R) - count(D) - count(E)) / length` | A simple charge-count proxy |
125
- | `r2` | `count(residues in DEHKNQRST) / length` | A defined polar/charged fraction |
126
- | `c` | `1` if `r1 > 0`, otherwise `0` | The binary conditioning class |
127
-
128
- ### Use your own data
129
-
130
- Supply a CSV with `sequence,c,r1,r2` columns:
131
-
132
- ```csv
133
- sequence,c,r1,r2
134
- QYWSDSWWESQMMSPWYPMSPLSV,0,-0.08333333,0.41666667
135
- CKSEFQPPHLMGHDFFACEMRNFK,0,0.00000000,0.45833333
136
- WPYGEHMLADNNVVKKRLQQWCFI,1,0.04166667,0.41666667
137
- KVVGALPIESFYTAKMESIAVEVI,0,-0.04166667,0.33333333
138
- ```
139
-
140
- ```bash
141
- python lectures_02_03/esm2_flow_guidance.py --data my_sequences.csv --output outputs/my_flow
142
- python lectures_02_03/esm2_diffusion_guidance.py --data my_sequences.csv --output outputs/my_diffusion
143
- ```
144
-
145
- - Include at least four sequences and both class labels, 0 and 1.
146
- - Use one fixed sequence length, at most 128 residues, and the 20 canonical
147
- amino acids to match the fixed-length latent tensors in both models.
148
- - Supply finite property values and orient each objective so **larger is better**.
149
- For a quantity to minimize, negate it before writing the CSV.
150
- - If both `r1` and `r2` are omitted, `composition_proxies()` calculates the
151
- two example properties directly. Supplied reward columns override them.
152
- - Experimental measurements or scores from a separate predictor can fill the
153
- reward columns. We fit a differentiable surrogate to these labels and
154
- evaluate its gradients with respect to the intermediate latent during sampling.
155
- - The printed composition proxies always retain their defined count-based
156
- meaning, even when custom reward labels are supplied. Re-evaluate custom
157
- objectives with the corresponding assay or scorer after decoding.
158
-
159
- The bundled demonstration uses all 64 sequences for training. For studies with
160
- held-out evaluation, partition sequences by cluster before fitting normalization
161
- statistics or model parameters.
162
-
163
- ## Property normalization and scalarization
164
-
165
- Each objective is standardized independently:
166
-
167
- ```text
168
- r̃ₘ(s) = [rₘ(s) − μₘ] / max(σₘ, 10⁻⁶)
169
- ```
170
-
171
- Here μₘ and σₘ are the training-set mean and standard deviation of property m.
172
- Standardization expresses each property in units of its training-set variation,
173
- so the scalarization weights specify relative preferences on a common scale.
174
- Reuse the saved training statistics for new examples to maintain that scale
175
- throughout evaluation and generation.
176
-
177
- The reward predictor estimates standardized clean-sequence properties from an
178
- intermediate latent and its time. During sampling, we combine its predictions:
179
-
180
- ```text
181
- Rλ(z,t) = λ₁ r̂₁(z,t) + λ₂ r̂₂(z,t)
182
- λ₁ ≥ 0, λ₂ ≥ 0, λ₁ + λ₂ = 1
183
- ```
184
-
185
- The nonnegative tradeoff weights λ are normalized to sum to one. Overall
186
- steering strength η is a separate parameter.
187
-
188
- ## Guidance configurations
189
-
190
- Both scripts run these settings in `main()`:
191
-
192
- | Output | CFG strength `w` | Reward strength `eta` | Property weights `lambdas` |
193
- | --- | ---: | ---: | --- |
194
- | `cfg.fasta` | 2.0, class 1 | 0.0 | Unused |
195
- | `single.fasta` | 0.0 | 1.0 | `(1.0, 0.0)` |
196
- | `multi.fasta` | 0.0 | 1.0 | `(0.7, 0.3)` |
197
-
198
- **Classifier-free guidance.** During training, we replace 20% of class labels
199
- with null class index 2. At sampling time, we combine the conditional and
200
- unconditional predictions as `F_uncond + w * (F_cond - F_uncond)`.
201
- Thus `w=0` is unconditional, `w=1` is ordinary conditional sampling, and `w>1` amplifies the conditional
202
- difference. F is a velocity for flow matching and predicted noise for DDPM.
203
-
204
- **Reward steering.** Freeze both networks, enable gradients only for the
205
- current latent, calculate the scalarized reward, and differentiate it with
206
- respect to that latent. Set `(1, 0)` for the first objective alone or change
207
- the weights to express a tradeoff. This correction uses the property predictor,
208
- while CFG uses the conditional and unconditional generative predictions.
209
-
210
- To change these settings, edit the three `sample(...)` calls in `main()`.
211
- To combine CFG with reward steering, set both strengths:
212
-
213
- ```python
214
- latent = sample(model, reward_model, n=8, c=1,
215
- w=2.0, eta=1.0, lambdas=(0.7, 0.3))
216
- ```
217
-
218
- Each call resets the sampling seed to compare methods using the same random
219
- draws for the same batch size.
220
-
221
- ### Guidance in the sampling dynamics
222
-
223
- | | Flow matching | DDPM diffusion |
224
- | --- | --- | --- |
225
- | Clean-data endpoint | Z₁ | Z₀ |
226
- | Training target | Velocity Z₁ − Z₀ | Injected Gaussian noise ε |
227
- | Generation direction | t = 0 → 1 | k = 1000 → 1 |
228
- | Reward correction | Add κ(t)∇Rλ to velocity | Subtract η√(1−ᾱₖ)∇Rλ from predicted noise |
229
- | Numerical step | Euler ODE update | Reverse mean plus posterior Gaussian noise |
230
-
231
- For the flow, we choose `κ(t)=4ηt(1−t)` to taper reward steering near the noise
232
- and data endpoints. For DDPM, we convert the reward-gradient score correction
233
- into a noise-prediction correction using `s=−ε/√(1−ᾱₖ)`. Both implementations
234
- use heuristic gradient steering based on learned estimates of endpoint properties.
235
-
236
- The DDPM predictor includes a Gaussian-reference skip plus a learned correction
237
- so the small MLP can carry high-dimensional noise. The training target remains
238
- noise, and the sampler uses the standard DDPM reverse equations.
239
-
240
- ## Constrained sequence decoding
241
-
242
- The default decoder requires **at least 12 of the 24 output residues** to
243
- belong to:
244
-
245
- ```text
246
- 𝒫 = {D, E, H, K, N, Q, R, S, T}
247
- ```
248
-
249
- We begin with the highest-logit residue at each position and count the members
250
- of 𝒫. When the count falls below the specified minimum, we replace the required
251
- number of residues using the lowest-cost substitutions into 𝒫. This discrete
252
- decoding procedure maximizes the sum of the fixed per-position logits subject
253
- to the minimum-count constraint.
254
-
255
- ```bash
256
- # Require at least 16 selected residues.
257
- python lectures_02_03/esm2_flow_guidance.py --min-polar 16
258
-
259
- # Remove the minimum-count constraint.
260
- python lectures_02_03/esm2_diffusion_guidance.py --min-polar 0
261
- ```
262
-
263
- Set `--min-polar` between zero and the sequence length to specify the minimum
264
- number of selected polar or charged residues in each decoded sequence.
265
-
266
- ## Verification
267
-
268
- Run the offline unit checks after installing dependencies:
269
-
270
- ```bash
271
- python -m unittest discover -s tests -v
272
- ```
273
-
274
- The five unit tests cover dataset annotations, scalarization weights, reward
275
- input gradients, DDPM schedule indexing, and constrained decoding. For the
276
- decoder test, we compare the selected sequences with exhaustive solutions on
277
- small examples. All unit tests run locally with the installed dependencies.
278
-
279
- We also checked both scripts end to end using the ESM-2 checkpoint, 200 training
280
- epochs, and all three guidance modes. All 48 decoded sequences in that run
281
- satisfied the specified minimum residue count.
282
-
283
- ## Troubleshooting
284
-
285
- - **Download failure:** the first run needs internet access to the ESM-2
286
- checkpoint. Cached subsequent runs can use `HF_HUB_OFFLINE=1` once all files
287
- have been downloaded.
288
- - **Dependency conflicts:** use a clean virtual environment and the pinned
289
- requirements. For a specific CUDA build, follow the official
290
- [PyTorch installation instructions](https://pytorch.org/get-started/locally/).
291
- - **Input validation error:** check equal lengths, canonical residues, both
292
- class labels, finite rewards, and a minimum count no greater than the length.
293
- - **Weak or repetitive samples:** inspect training convergence, evaluate
294
- reconstruction through the ESM-2 head, and assess property-predictor accuracy
295
- on held-out sequences before adjusting guidance strength. Compare guidance
296
- settings using properties recalculated after sequence decoding.
297
-
298
- ## References
299
-
300
- - [ESM-2 checkpoint](https://huggingface.co/facebook/esm2_t6_8M_UR50D)
301
- and [Transformers ESM documentation](https://huggingface.co/docs/transformers/model_doc/esm).
302
- - [Flow Matching for Generative Modeling](https://arxiv.org/abs/2210.02747).
303
- - [Denoising Diffusion Probabilistic Models](https://arxiv.org/abs/2006.11239).
304
- - [Classifier-Free Diffusion Guidance](https://arxiv.org/abs/2207.12598).
305
-
306
- ## License
307
-
308
- Repository code is distributed under the [MIT License](https://huggingface.co/ChatterjeeLab/CIS6270/blob/main/LICENSE), matching the
309
- repository's license setting. ESM-2 weights are downloaded separately and remain
310
- subject to their original distribution terms.
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
lectures_02_03/esm2_diffusion_guidance.py DELETED
@@ -1,264 +0,0 @@
1
- """Standalone ESM-2 diffusion guidance example for CIS 6270.
2
- Input CSV: sequence,c,r1,r2. Sequences have one fixed length; c is 0 or 1.
3
- Both objectives are oriented so larger is better. The bundled CSV is synthetic.
4
- Run: python esm2_diffusion_guidance.py --data esm2_example.csv --epochs 200
5
- Install: pip install torch transformers==4.57.6
6
- Outputs: guided residue latents, model weights, and decoded amino-acid sequences.
7
- """
8
- import argparse
9
- import csv
10
- from pathlib import Path
11
-
12
- import torch
13
- from torch import nn
14
- import torch.nn.functional as F
15
- from torch.utils.data import DataLoader, TensorDataset
16
- from transformers import AutoTokenizer, EsmForMaskedLM
17
-
18
- ROOT = Path(__file__).resolve().parent
19
- DEVICE = torch.device("cuda" if torch.cuda.is_available() else "cpu")
20
- ESM_NAME = "facebook/esm2_t6_8M_UR50D"
21
- AMINO_ACIDS = "ACDEFGHIKLMNPQRSTVWY"
22
- POLAR_RESIDUES = "DEHKNQRST" # Operational polar/charged set; not a solubility assay.
23
- BATCH_SIZE, HIDDEN, LEARNING_RATE = 16, 128, 1e-3
24
- CONDITION_DROP = 0.2 # Drop c during training so the same model learns the null condition.
25
-
26
-
27
- def composition_proxies(sequences):
28
- # Transparent teaching rewards; these are not measured activity or solubility.
29
- return torch.tensor([
30
- [(sum(a in "KR" for a in s) - sum(a in "DE" for a in s)) / len(s),
31
- sum(a in POLAR_RESIDUES for a in s) / len(s)]
32
- for s in sequences
33
- ], dtype=torch.float32)
34
-
35
-
36
- # 1. Load annotated sequences and encode frozen ESM-2 residue vectors.
37
- @torch.no_grad()
38
- def load_data(path):
39
- with Path(path).open(newline="") as handle:
40
- rows = list(csv.DictReader(handle))
41
- sequences = [row["sequence"].strip().upper() for row in rows]
42
- if len(rows) < 4 or any(not s or set(s) - set(AMINO_ACIDS) for s in sequences):
43
- raise ValueError("Supply at least four sequences using the 20 standard amino acids")
44
- lengths = {len(s) for s in sequences}
45
- if len(lengths) != 1 or max(lengths) > 128:
46
- raise ValueError("This compact example requires one fixed sequence length, at most 128")
47
- c = torch.tensor([int(row["c"]) for row in rows], dtype=torch.long)
48
- # With sequence,c only, compute example rewards directly from residue counts.
49
- # Optional r1/r2 columns override them with supplied labels, such as measurements.
50
- if {"r1", "r2"}.issubset(rows[0]):
51
- r = torch.tensor([[float(row["r1"]), float(row["r2"])] for row in rows])
52
- else:
53
- r = composition_proxies(sequences)
54
- if set(c.tolist()) != {0, 1} or not torch.isfinite(r).all():
55
- raise ValueError("Include both c=0 and c=1, with finite r1 and r2")
56
-
57
- tokenizer = AutoTokenizer.from_pretrained(ESM_NAME, cache_dir=ROOT / ".esm2_cache")
58
- esm = EsmForMaskedLM.from_pretrained(
59
- ESM_NAME, cache_dir=ROOT / ".esm2_cache", use_safetensors=True
60
- ).to(DEVICE).eval().requires_grad_(False)
61
- encoded = []
62
- for start in range(0, len(sequences), BATCH_SIZE):
63
- tokens = tokenizer(sequences[start:start + BATCH_SIZE], return_tensors="pt")
64
- tokens = {key: value.to(DEVICE) for key, value in tokens.items()}
65
- hidden = esm.esm(**tokens).last_hidden_state
66
- encoded.append(hidden[:, 1:-1].cpu()) # Remove BOS/EOS; retain all L residue positions.
67
- z = torch.cat(encoded) # [N, L, D], with D=320 for this checkpoint.
68
- z_mean = z.mean((0, 1), keepdim=True)
69
- z_std = z.std((0, 1), correction=0, keepdim=True).clamp_min(1e-4)
70
- r_mean, r_std = r.mean(0), r.std(0, correction=0).clamp_min(1e-6)
71
- dataset = TensorDataset((z - z_mean) / z_std, c, (r - r_mean) / r_std)
72
- stats = {"z_mean": z_mean, "z_std": z_std, "r_mean": r_mean, "r_std": r_std}
73
- return dataset, esm, tokenizer, stats
74
-
75
-
76
- # 2. Small model classes: flattening lets each output depend on the entire sequence.
77
- class DiffusionModel(nn.Module):
78
- def __init__(self, length, dim):
79
- super().__init__()
80
- self.length, self.dim = length, dim
81
- self.time = nn.Sequential(nn.Linear(1, 32), nn.SiLU(), nn.Linear(32, 32))
82
- self.condition = nn.Embedding(3, 16) # Indices 0,1 are classes; 2 is the null class.
83
- self.net = nn.Sequential(
84
- nn.Linear(length * dim + 48, HIDDEN), nn.SiLU(),
85
- nn.Linear(HIDDEN, HIDDEN), nn.SiLU(),
86
- nn.Linear(HIDDEN, length * dim),
87
- )
88
-
89
- def forward(self, z, t, c):
90
- time = self.time(t[:, None])
91
- inputs = torch.cat([z.flatten(1), time, self.condition(c)], dim=1)
92
- k = (t * K).round().long().clamp(0, K)
93
- a = alpha_bars[k, None, None]
94
- return (1 - a).sqrt() * z + a.sqrt() * self.net(inputs).reshape_as(z)
95
- # Gaussian-reference noise prediction plus a learned correction.
96
- # The skip carries all coordinates; the correction vanishes near pure noise.
97
-
98
-
99
- class RewardModel(nn.Module):
100
- def __init__(self, length, dim):
101
- super().__init__()
102
- self.net = nn.Sequential(
103
- nn.Linear(length * dim + 1, HIDDEN), nn.SiLU(),
104
- nn.Linear(HIDDEN, HIDDEN), nn.SiLU(), nn.Linear(HIDDEN, 2),
105
- )
106
-
107
- def forward(self, z, t):
108
- return self.net(torch.cat([z.flatten(1), t[:, None]], dim=1))
109
- # Two predicted standardized endpoint objectives, given the current latent and time.
110
-
111
-
112
- # The same DDPM schedule as Section 7; k=0 denotes the clean latent.
113
- K = 1000
114
- betas = torch.cat([torch.zeros(1), torch.linspace(1e-4, 0.02, K)]).to(DEVICE)
115
- alphas = 1.0 - betas
116
- alpha_bars = alphas.cumprod(0)
117
- previous = torch.cat([torch.ones(1, device=DEVICE), alpha_bars[:-1]])
118
- posterior_variances = betas * (1 - previous) / (1 - alpha_bars).clamp_min(1e-20)
119
-
120
-
121
- # 3. DDPM: corrupt Z_0 at step k and regress the exact noise epsilon.
122
- def train(dataset, epochs):
123
- _, length, dim = dataset.tensors[0].shape
124
- loader = DataLoader(dataset, batch_size=BATCH_SIZE, shuffle=True)
125
- model = DiffusionModel(length, dim).to(DEVICE)
126
- reward_model = RewardModel(length, dim).to(DEVICE)
127
- optimizer = torch.optim.Adam(list(model.parameters()) + list(reward_model.parameters()), lr=LEARNING_RATE)
128
- for epoch in range(epochs):
129
- total = 0.0
130
- for z0, c, r_tilde in loader:
131
- z0, c, r_tilde = z0.to(DEVICE), c.to(DEVICE), r_tilde.to(DEVICE)
132
- k = torch.randint(1, K + 1, (len(z0),), device=DEVICE)
133
- t = k.float() / K # Normalized k is only the network's time input.
134
- a = alpha_bars[k, None, None]
135
- epsilon = torch.randn_like(z0)
136
- zk = a.sqrt() * z0 + (1 - a).sqrt() * epsilon
137
- dropped = c.masked_fill(torch.rand(len(c), device=DEVICE) < CONDITION_DROP, 2)
138
- loss_diffusion = F.mse_loss(model(zk, t, dropped), epsilon)
139
- loss_reward = F.mse_loss(reward_model(zk, t), r_tilde)
140
- loss = loss_diffusion + loss_reward
141
- optimizer.zero_grad(set_to_none=True)
142
- loss.backward()
143
- optimizer.step()
144
- total += loss.item()
145
- if (epoch + 1) % 50 == 0 or epoch + 1 == epochs:
146
- print(f"epoch {epoch+1}: combined training loss {total / len(loader):.4f}")
147
- for network in (model, reward_model):
148
- network.eval().requires_grad_(False)
149
- for parameter in network.parameters():
150
- parameter.grad = None
151
- return model, reward_model
152
-
153
-
154
- # 4. R_lambda = sum_m lambda_m * r_tilde_m; differentiate the current latent only.
155
- def reward_gradient(reward_model, z, t, lambdas):
156
- with torch.enable_grad(): # Re-enable input gradients inside sampling.
157
- state = z.detach().requires_grad_(True)
158
- R_lambda = (reward_model(state, t) * lambdas).sum(dim=1)
159
- grad = torch.autograd.grad(R_lambda.sum(), state)[0]
160
- return grad.detach() # Do not retain graphs across sampling steps.
161
-
162
-
163
- def normalize_weights(lambdas):
164
- values = torch.as_tensor(lambdas, dtype=torch.float32, device=DEVICE)
165
- if values.shape != (2,) or not torch.isfinite(values).all() or (values < 0).any() or values.sum() <= 0:
166
- raise ValueError("Use two finite, nonnegative weights with a positive sum")
167
- return values / values.sum() # lambda controls tradeoffs, eta controls strength.
168
-
169
-
170
- # 5. Sampling: CFG first; optional reward gradient then changes predicted noise.
171
- @torch.no_grad()
172
- def sample(model, reward_model, n=8, c=1, w=0.0, eta=0.0, lambdas=(1.0, 0.0)):
173
- if c not in (0, 1) or w < 0 or eta < 0:
174
- raise ValueError("Use c=0/1 and nonnegative w and eta")
175
- lambdas = normalize_weights(lambdas)
176
- torch.manual_seed(123) # Same starting and reverse noise across comparisons.
177
- z = torch.randn(n, model.length, model.dim, device=DEVICE)
178
- null = torch.full((n,), 2, dtype=torch.long, device=DEVICE)
179
- condition = torch.full((n,), c, dtype=torch.long, device=DEVICE)
180
- for k in range(K, 0, -1):
181
- t = torch.full((n,), k / K, device=DEVICE)
182
- eps_uncond = model(z, t, null)
183
- eps = eps_uncond
184
- if w != 0:
185
- eps = eps_uncond + w * (model(z, t, condition) - eps_uncond)
186
- sigma = (1 - alpha_bars[k]).sqrt() # FORWARD corruption standard deviation.
187
- if eta != 0:
188
- grad_R = reward_gradient(reward_model, z, t, lambdas)
189
- eps = eps - eta * sigma * grad_R # s_guided=s_theta+eta*grad R; s=-epsilon/sigma.
190
- mean = (z - betas[k] * eps / sigma) / alphas[k].sqrt()
191
- z = mean + posterior_variances[k].sqrt() * torch.randn_like(z) if k > 1 else mean
192
- return z # No fresh noise at k=1.
193
-
194
-
195
- # 6. One decoder, used only AFTER the flow or diffusion trajectory is complete.
196
- @torch.no_grad()
197
- def decode(z, esm, tokenizer, stats, min_polar=12):
198
- latent = z * stats["z_std"].to(z.device) + stats["z_mean"].to(z.device)
199
- logits = esm.lm_head(latent) # Frozen head produces one vocabulary distribution per residue.
200
- aa_ids = torch.tensor(tokenizer.convert_tokens_to_ids(list(AMINO_ACIDS)), device=z.device)
201
- logits = logits.index_select(-1, aa_ids) # Restrict the vocabulary to the 20 amino acids.
202
- if not 0 <= min_polar <= z.shape[1]:
203
- raise ValueError("min_polar must lie between zero and the sequence length")
204
- polar = torch.tensor([a in POLAR_RESIDUES for a in AMINO_ACIDS], device=z.device)
205
- polar_ids = polar.nonzero().flatten()
206
- best_scores, choices = logits.max(dim=-1) # Start from unrestricted amino-acid argmax.
207
- polar_scores, local = logits[..., polar_ids].max(dim=-1)
208
- polar_choices = polar_ids[local] # Best polar/charged amino acid at each position.
209
- for i in range(len(z)):
210
- already_polar = polar[choices[i]]
211
- missing = max(0, min_polar - int(already_polar.sum()))
212
- if missing:
213
- cost = (best_scores[i] - polar_scores[i]).masked_fill(already_polar, float("inf"))
214
- positions = cost.topk(missing, largest=False).indices
215
- choices[i, positions] = polar_choices[i, positions]
216
- return ["".join(AMINO_ACIDS[i] for i in row) for row in choices.cpu().tolist()]
217
- # Exact maximum-logit decode subject to at least min_polar selected residues.
218
- # This enforces composition, not measured solubility; no gradient through argmax.
219
-
220
-
221
- # 7. Run CFG, single-objective steering, and scalarized multi-objective steering in order.
222
- def main():
223
- parser = argparse.ArgumentParser(description=__doc__)
224
- parser.add_argument("--data", type=Path, default=ROOT / "esm2_example.csv")
225
- parser.add_argument("--epochs", type=int, default=200)
226
- parser.add_argument("--samples", type=int, default=8)
227
- parser.add_argument("--min-polar", type=int, default=12)
228
- parser.add_argument("--output", type=Path, default=ROOT / "esm2_diffusion_outputs")
229
- args = parser.parse_args()
230
- if args.epochs < 1 or args.samples < 1:
231
- parser.error("epochs and samples must be positive")
232
- torch.manual_seed(7)
233
- if DEVICE.type == "cpu":
234
- torch.set_num_threads(2)
235
- dataset, esm, tokenizer, stats = load_data(args.data)
236
- if not 0 <= args.min_polar <= dataset.tensors[0].shape[1]:
237
- parser.error("min-polar must lie between zero and the sequence length")
238
- model, reward_model = train(dataset, args.epochs)
239
- outputs = {
240
- "cfg": sample(model, reward_model, args.samples, c=1, w=2.0),
241
- "single": sample(model, reward_model, args.samples, eta=1.0, lambdas=(1.0, 0.0)),
242
- "multi": sample(model, reward_model, args.samples, eta=1.0, lambdas=(0.7, 0.3)),
243
- }
244
- args.output.mkdir(parents=True, exist_ok=True)
245
- for name, latent in outputs.items():
246
- if not torch.isfinite(latent).all():
247
- raise RuntimeError(f"Nonfinite {name} output; reduce guidance or check training")
248
- sequences = decode(latent, esm, tokenizer, stats, args.min_polar)
249
- assert all(sum(a in POLAR_RESIDUES for a in s) >= args.min_polar for s in sequences)
250
- fasta = "".join(f">{name}_{i+1}\n{seq}\n" for i, seq in enumerate(sequences))
251
- (args.output / f"{name}.fasta").write_text(fasta)
252
- count = sum(a in POLAR_RESIDUES for a in sequences[0])
253
- print(f"{name}: {sequences[0]} polar/charged residues={count}")
254
- print(" decoded mean composition proxies:", composition_proxies(sequences).mean(0).tolist())
255
- torch.save({"standardized_latents": {k: v.cpu() for k, v in outputs.items()},
256
- "model": model.state_dict(), "reward_model": reward_model.state_dict(),
257
- "stats": stats, "esm_name": ESM_NAME,
258
- "length": model.length, "dim": model.dim, "min_polar": args.min_polar,
259
- "polar_residues": POLAR_RESIDUES}, args.output / "results.pt")
260
- print(f"Saved latent tensors and FASTA sequences to {args.output}")
261
-
262
-
263
- if __name__ == "__main__":
264
- main()
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
lectures_02_03/esm2_example.csv DELETED
@@ -1,65 +0,0 @@
1
- sequence,c,r1,r2
2
- QYWSDSWWESQMMSPWYPMSPLSV,0,-0.08333333,0.41666667
3
- CKSEFQPPHLMGHDFFACEMRNFK,0,0.00000000,0.45833333
4
- WPYGEHMLADNNVVKKRLQQWCFI,1,0.04166667,0.41666667
5
- KVVGALPIESFYTAKMESIAVEVI,0,-0.04166667,0.33333333
6
- MEFHAYSLGPASRWSSRHGYFTNL,1,0.04166667,0.45833333
7
- TCYLFCLAQIVMFGRAYVGGYWAK,1,0.08333333,0.16666667
8
- VFSHIVACCEIETRDTCVWNHMYM,0,-0.08333333,0.41666667
9
- QRNQTEIRTLIYMPYGNKTYLRGS,1,0.12500000,0.54166667
10
- WDNQRIKTALAGHTGVDSGWMHNF,0,0.00000000,0.50000000
11
- PPVCRGWQMPRAALWVMRCRFPRQ,1,0.20833333,0.29166667
12
- HPNAKNYGETIFFGLWGNDRKNLV,1,0.04166667,0.45833333
13
- CMKEIYESFLDETTIWVIHMIAAR,0,-0.08333333,0.41666667
14
- RWDTGMWGKDCIHRPIQKKDPTWA,1,0.08333333,0.50000000
15
- WCCQLMAYPPAKMGHPTFDHPLPK,1,0.04166667,0.29166667
16
- QPVRDFVLEVGQSILESHSLKVPS,0,-0.04166667,0.50000000
17
- GREMRGFQKTPNATKKEGTYKPAP,1,0.16666667,0.54166667
18
- ENEWNGSRHYWCVRWLGTLKILTF,1,0.04166667,0.45833333
19
- AVSQIAPHARQQASSQENKHGETT,0,0.00000000,0.66666667
20
- NADTREKMPSNDLSWWLYNMFMPA,0,-0.04166667,0.45833333
21
- SLWFAIMPYVAKYVIYGHVEPTAT,0,0.00000000,0.25000000
22
- YGWHYTLLMWYDERCGFLGRGRAQ,1,0.04166667,0.33333333
23
- EEQEKNLDDMNSLDFSQNLSLFEY,0,-0.25000000,0.66666667
24
- ETWDHYPQYKCVKCKNNHCPHVAY,1,0.04166667,0.50000000
25
- YHIMRKVPSTNEFRPLAHNNCTLQ,1,0.08333333,0.54166667
26
- DFRMNQYGSKDYPKTTGAREHLFM,1,0.04166667,0.54166667
27
- IQWWSDVQTMQCMADEGMHHQYFV,0,-0.12500000,0.45833333
28
- WAQMMMVRQIDCRNFIRHTVYSNF,1,0.08333333,0.45833333
29
- DMTFGHSSDAILAKVYDSKNCPQL,0,-0.04166667,0.50000000
30
- DRSICNKIPPSQTIPFMGPFISQN,1,0.04166667,0.45833333
31
- QNHLSDAEIKYCIKWNVLNGKMLP,1,0.04166667,0.45833333
32
- RQYCPEPSWSQHTNKPACMKVEEC,0,0.00000000,0.54166667
33
- IWWWHVNCQCRDDHEHCPLCYAPI,0,-0.08333333,0.37500000
34
- SDKSVTHIMERDIWPIIAWLRDRV,0,0.00000000,0.50000000
35
- WSDAWDVYIWILAMNCEVGRMYYC,0,-0.08333333,0.25000000
36
- VAKELIRNYGDEFCWLVNSSMKED,0,-0.08333333,0.50000000
37
- VTWANPHINITFCDPMHIIRIQCQ,0,0.00000000,0.41666667
38
- GYYWVTHLVHMVDIWFSDMPIHYE,0,-0.12500000,0.33333333
39
- TRALRHDDVWVDYYFIWTALTYQF,0,-0.04166667,0.41666667
40
- CGNGSDFDTFIQMWICACGMCIDK,0,-0.08333333,0.33333333
41
- NGQMACTQENWAWSVFHFEWWCPK,0,-0.04166667,0.41666667
42
- GSRETKHEVTADCPFQWNGRALKK,1,0.08333333,0.58333333
43
- CHGKYTFTFGMEIRDMPAWYGHFV,0,0.00000000,0.33333333
44
- GSYDNPLAIKWQMVCNSSRMDWTM,0,0.00000000,0.45833333
45
- FRIAQIRSKIGYLQNIFVMSLLLR,1,0.16666667,0.37500000
46
- SCCDWRQMQDRVGPAMMKEWACGE,0,-0.04166667,0.41666667
47
- FFSELMKCHYHYCYAWRRPIAKDW,1,0.08333333,0.37500000
48
- AWGPVLDFNIQSDIFYPCRTDMGL,0,-0.08333333,0.33333333
49
- PHTCRRIMMMEQDLTNFPLAYCTQ,0,0.00000000,0.45833333
50
- DQINDGSAHENQYCQNPSDPFVWK,0,-0.12500000,0.58333333
51
- QMRSHFQYGCGWMITPMRFPKNMA,1,0.12500000,0.37500000
52
- EECYALAGDYDFKVRIQHAQCTWQ,0,-0.08333333,0.45833333
53
- ITKQMWVNDVWDTEVRHHSFMTLF,0,-0.04166667,0.54166667
54
- EHAKYVKNSNCYNKENNVDLFYCN,0,0.00000000,0.58333333
55
- SYYVPKTKQRLTTSRCGLKGVVKA,1,0.25000000,0.50000000
56
- TDAKGWHYNWHWPLHQFCHLPVHA,0,0.00000000,0.41666667
57
- YIVRMSDSFHSGHDHDKKNLIMYH,0,0.00000000,0.58333333
58
- CTHVQNPNSIETNCHEHYDFGTSM,0,-0.12500000,0.62500000
59
- RWWYQNHLMCYNSKDACYHPWINM,1,0.04166667,0.41666667
60
- RKSNGGKSLIQNAGITCVKWCYGV,1,0.16666667,0.41666667
61
- WWNQITTPHWCKWWASHWMQYSRW,1,0.08333333,0.45833333
62
- IDTICWQVCILVQESDAKDVFEGN,0,-0.16666667,0.45833333
63
- NFHEWVLWELTGQWRVWSRQHSKE,0,0.00000000,0.58333333
64
- VTYCMMAKFAGMDYEQKICVHVAC,0,0.00000000,0.29166667
65
- KVKTTLKVEREAPPYTRMSLSAWT,1,0.12500000,0.54166667
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
lectures_02_03/esm2_flow_guidance.py DELETED
@@ -1,251 +0,0 @@
1
- """Standalone ESM-2 flow matching guidance example for CIS 6270.
2
- Input CSV: sequence,c,r1,r2. Sequences have one fixed length; c is 0 or 1.
3
- Both objectives are oriented so larger is better. The bundled CSV is synthetic.
4
- Run: python esm2_flow_guidance.py --data esm2_example.csv --epochs 200
5
- Install: pip install torch transformers==4.57.6
6
- Outputs: guided residue latents, model weights, and decoded amino-acid sequences.
7
- """
8
- import argparse
9
- import csv
10
- from pathlib import Path
11
-
12
- import torch
13
- from torch import nn
14
- import torch.nn.functional as F
15
- from torch.utils.data import DataLoader, TensorDataset
16
- from transformers import AutoTokenizer, EsmForMaskedLM
17
-
18
- ROOT = Path(__file__).resolve().parent
19
- DEVICE = torch.device("cuda" if torch.cuda.is_available() else "cpu")
20
- ESM_NAME = "facebook/esm2_t6_8M_UR50D"
21
- AMINO_ACIDS = "ACDEFGHIKLMNPQRSTVWY"
22
- POLAR_RESIDUES = "DEHKNQRST" # Operational polar/charged set; not a solubility assay.
23
- BATCH_SIZE, HIDDEN, LEARNING_RATE = 16, 128, 1e-3
24
- CONDITION_DROP = 0.2 # Drop c during training so the same model learns the null condition.
25
-
26
-
27
- def composition_proxies(sequences):
28
- # Transparent teaching rewards; these are not measured activity or solubility.
29
- return torch.tensor([
30
- [(sum(a in "KR" for a in s) - sum(a in "DE" for a in s)) / len(s),
31
- sum(a in POLAR_RESIDUES for a in s) / len(s)]
32
- for s in sequences
33
- ], dtype=torch.float32)
34
-
35
-
36
- # 1. Load annotated sequences and encode frozen ESM-2 residue vectors.
37
- @torch.no_grad()
38
- def load_data(path):
39
- with Path(path).open(newline="") as handle:
40
- rows = list(csv.DictReader(handle))
41
- sequences = [row["sequence"].strip().upper() for row in rows]
42
- if len(rows) < 4 or any(not s or set(s) - set(AMINO_ACIDS) for s in sequences):
43
- raise ValueError("Supply at least four sequences using the 20 standard amino acids")
44
- lengths = {len(s) for s in sequences}
45
- if len(lengths) != 1 or max(lengths) > 128:
46
- raise ValueError("This compact example requires one fixed sequence length, at most 128")
47
- c = torch.tensor([int(row["c"]) for row in rows], dtype=torch.long)
48
- # With sequence,c only, compute example rewards directly from residue counts.
49
- # Optional r1/r2 columns override them with supplied labels, such as measurements.
50
- if {"r1", "r2"}.issubset(rows[0]):
51
- r = torch.tensor([[float(row["r1"]), float(row["r2"])] for row in rows])
52
- else:
53
- r = composition_proxies(sequences)
54
- if set(c.tolist()) != {0, 1} or not torch.isfinite(r).all():
55
- raise ValueError("Include both c=0 and c=1, with finite r1 and r2")
56
-
57
- tokenizer = AutoTokenizer.from_pretrained(ESM_NAME, cache_dir=ROOT / ".esm2_cache")
58
- esm = EsmForMaskedLM.from_pretrained(
59
- ESM_NAME, cache_dir=ROOT / ".esm2_cache", use_safetensors=True
60
- ).to(DEVICE).eval().requires_grad_(False)
61
- encoded = []
62
- for start in range(0, len(sequences), BATCH_SIZE):
63
- tokens = tokenizer(sequences[start:start + BATCH_SIZE], return_tensors="pt")
64
- tokens = {key: value.to(DEVICE) for key, value in tokens.items()}
65
- hidden = esm.esm(**tokens).last_hidden_state
66
- encoded.append(hidden[:, 1:-1].cpu()) # Remove BOS/EOS; retain all L residue positions.
67
- z = torch.cat(encoded) # [N, L, D], with D=320 for this checkpoint.
68
- z_mean = z.mean((0, 1), keepdim=True)
69
- z_std = z.std((0, 1), correction=0, keepdim=True).clamp_min(1e-4)
70
- r_mean, r_std = r.mean(0), r.std(0, correction=0).clamp_min(1e-6)
71
- dataset = TensorDataset((z - z_mean) / z_std, c, (r - r_mean) / r_std)
72
- stats = {"z_mean": z_mean, "z_std": z_std, "r_mean": r_mean, "r_std": r_std}
73
- return dataset, esm, tokenizer, stats
74
-
75
-
76
- # 2. Small model classes: flattening lets each output depend on the entire sequence.
77
- class FlowModel(nn.Module):
78
- def __init__(self, length, dim):
79
- super().__init__()
80
- self.length, self.dim = length, dim
81
- self.time = nn.Sequential(nn.Linear(1, 32), nn.SiLU(), nn.Linear(32, 32))
82
- self.skip = nn.Linear(32, 1) # Preserve full-dimensional state/noise through a time gate.
83
- self.condition = nn.Embedding(3, 16) # Indices 0,1 are classes; 2 is the null class.
84
- self.net = nn.Sequential(
85
- nn.Linear(length * dim + 48, HIDDEN), nn.SiLU(),
86
- nn.Linear(HIDDEN, HIDDEN), nn.SiLU(),
87
- nn.Linear(HIDDEN, length * dim),
88
- )
89
-
90
- def forward(self, z, t, c):
91
- time = self.time(t[:, None])
92
- inputs = torch.cat([z.flatten(1), time, self.condition(c)], dim=1)
93
- return self.skip(time)[:, :, None] * z + self.net(inputs).reshape_as(z)
94
- # A time-dependent linear part plus a learned nonlinear velocity correction.
95
-
96
-
97
- class RewardModel(nn.Module):
98
- def __init__(self, length, dim):
99
- super().__init__()
100
- self.net = nn.Sequential(
101
- nn.Linear(length * dim + 1, HIDDEN), nn.SiLU(),
102
- nn.Linear(HIDDEN, HIDDEN), nn.SiLU(), nn.Linear(HIDDEN, 2),
103
- )
104
-
105
- def forward(self, z, t):
106
- return self.net(torch.cat([z.flatten(1), t[:, None]], dim=1))
107
- # Two predicted standardized endpoint objectives, given the current latent and time.
108
-
109
-
110
- # 3. Flow matching: Z_t=(1-t)Z_0+tZ_1; target velocity is Z_1-Z_0.
111
- def train(dataset, epochs):
112
- _, length, dim = dataset.tensors[0].shape
113
- loader = DataLoader(dataset, batch_size=BATCH_SIZE, shuffle=True)
114
- model = FlowModel(length, dim).to(DEVICE)
115
- reward_model = RewardModel(length, dim).to(DEVICE)
116
- optimizer = torch.optim.Adam(list(model.parameters()) + list(reward_model.parameters()), lr=LEARNING_RATE)
117
- for epoch in range(epochs):
118
- total = 0.0
119
- for z1, c, r_tilde in loader:
120
- z1, c, r_tilde = z1.to(DEVICE), c.to(DEVICE), r_tilde.to(DEVICE)
121
- z0 = torch.randn_like(z1) # Flow convention: Z_0=noise, Z_1=clean latent.
122
- t = torch.rand(len(z1), device=DEVICE)
123
- zt = (1 - t[:, None, None]) * z0 + t[:, None, None] * z1
124
- dropped = c.masked_fill(torch.rand(len(c), device=DEVICE) < CONDITION_DROP, 2)
125
- loss_flow = F.mse_loss(model(zt, t, dropped), z1 - z0)
126
- loss_reward = F.mse_loss(reward_model(zt, t), r_tilde)
127
- loss = loss_flow + loss_reward # Independent parameter sets; one optimizer suffices.
128
- optimizer.zero_grad(set_to_none=True)
129
- loss.backward()
130
- optimizer.step()
131
- total += loss.item()
132
- if (epoch + 1) % 50 == 0 or epoch + 1 == epochs:
133
- print(f"epoch {epoch+1}: combined training loss {total / len(loader):.4f}")
134
- for network in (model, reward_model):
135
- network.eval().requires_grad_(False) # Freeze weights; still allow gradients of input z.
136
- for parameter in network.parameters():
137
- parameter.grad = None
138
- return model, reward_model
139
-
140
-
141
- # 4. R_lambda = sum_m lambda_m * r_tilde_m; differentiate the current latent only.
142
- def reward_gradient(reward_model, z, t, lambdas):
143
- with torch.enable_grad(): # Re-enable input gradients inside sampling.
144
- state = z.detach().requires_grad_(True)
145
- R_lambda = (reward_model(state, t) * lambdas).sum(dim=1)
146
- grad = torch.autograd.grad(R_lambda.sum(), state)[0]
147
- return grad.detach() # Do not retain graphs across sampling steps.
148
-
149
-
150
- def normalize_weights(lambdas):
151
- values = torch.as_tensor(lambdas, dtype=torch.float32, device=DEVICE)
152
- if values.shape != (2,) or not torch.isfinite(values).all() or (values < 0).any() or values.sum() <= 0:
153
- raise ValueError("Use two finite, nonnegative weights with a positive sum")
154
- return values / values.sum() # lambda controls tradeoffs, eta controls strength.
155
-
156
-
157
- # 5. Sampling: CFG first; optional reward gradient then changes the velocity.
158
- @torch.no_grad()
159
- def sample(model, reward_model, n=8, c=1, w=0.0, eta=0.0, lambdas=(1.0, 0.0), steps=200):
160
- if c not in (0, 1) or steps < 1 or w < 0 or eta < 0:
161
- raise ValueError("Use c=0/1, positive steps, and nonnegative w and eta")
162
- lambdas = normalize_weights(lambdas)
163
- torch.manual_seed(123) # Compare methods from the same initial noise.
164
- z = torch.randn(n, model.length, model.dim, device=DEVICE)
165
- null = torch.full((n,), 2, dtype=torch.long, device=DEVICE)
166
- condition = torch.full((n,), c, dtype=torch.long, device=DEVICE)
167
- dt = 1.0 / steps
168
- for step in range(steps):
169
- t = torch.full((n,), step * dt, device=DEVICE)
170
- v_uncond = model(z, t, null)
171
- v = v_uncond
172
- if w != 0:
173
- v = v_uncond + w * (model(z, t, condition) - v_uncond)
174
- if eta != 0:
175
- grad_R = reward_gradient(reward_model, z, t, lambdas)
176
- kappa = eta * 4 * t[:, None, None] * (1 - t[:, None, None])
177
- v = v + kappa * grad_R # Chosen steering rule; not exact conditional transport.
178
- z = z + dt * v # Integrate forward from t=0 to t=1.
179
- return z
180
-
181
-
182
- # 6. One decoder, used only AFTER the flow or diffusion trajectory is complete.
183
- @torch.no_grad()
184
- def decode(z, esm, tokenizer, stats, min_polar=12):
185
- latent = z * stats["z_std"].to(z.device) + stats["z_mean"].to(z.device)
186
- logits = esm.lm_head(latent) # Frozen head produces one vocabulary distribution per residue.
187
- aa_ids = torch.tensor(tokenizer.convert_tokens_to_ids(list(AMINO_ACIDS)), device=z.device)
188
- logits = logits.index_select(-1, aa_ids) # Restrict the vocabulary to the 20 amino acids.
189
- if not 0 <= min_polar <= z.shape[1]:
190
- raise ValueError("min_polar must lie between zero and the sequence length")
191
- polar = torch.tensor([a in POLAR_RESIDUES for a in AMINO_ACIDS], device=z.device)
192
- polar_ids = polar.nonzero().flatten()
193
- best_scores, choices = logits.max(dim=-1) # Start from unrestricted amino-acid argmax.
194
- polar_scores, local = logits[..., polar_ids].max(dim=-1)
195
- polar_choices = polar_ids[local] # Best polar/charged amino acid at each position.
196
- for i in range(len(z)):
197
- already_polar = polar[choices[i]]
198
- missing = max(0, min_polar - int(already_polar.sum()))
199
- if missing:
200
- cost = (best_scores[i] - polar_scores[i]).masked_fill(already_polar, float("inf"))
201
- positions = cost.topk(missing, largest=False).indices
202
- choices[i, positions] = polar_choices[i, positions]
203
- return ["".join(AMINO_ACIDS[i] for i in row) for row in choices.cpu().tolist()]
204
- # Exact maximum-logit decode subject to at least min_polar selected residues.
205
- # This enforces composition, not measured solubility; no gradient through argmax.
206
-
207
-
208
- # 7. Run CFG, single-objective steering, and scalarized multi-objective steering in order.
209
- def main():
210
- parser = argparse.ArgumentParser(description=__doc__)
211
- parser.add_argument("--data", type=Path, default=ROOT / "esm2_example.csv")
212
- parser.add_argument("--epochs", type=int, default=200)
213
- parser.add_argument("--samples", type=int, default=8)
214
- parser.add_argument("--min-polar", type=int, default=12)
215
- parser.add_argument("--output", type=Path, default=ROOT / "esm2_flow_outputs")
216
- args = parser.parse_args()
217
- if args.epochs < 1 or args.samples < 1:
218
- parser.error("epochs and samples must be positive")
219
- torch.manual_seed(7)
220
- if DEVICE.type == "cpu":
221
- torch.set_num_threads(2)
222
- dataset, esm, tokenizer, stats = load_data(args.data)
223
- if not 0 <= args.min_polar <= dataset.tensors[0].shape[1]:
224
- parser.error("min-polar must lie between zero and the sequence length")
225
- model, reward_model = train(dataset, args.epochs)
226
- outputs = {
227
- "cfg": sample(model, reward_model, args.samples, c=1, w=2.0),
228
- "single": sample(model, reward_model, args.samples, eta=1.0, lambdas=(1.0, 0.0)),
229
- "multi": sample(model, reward_model, args.samples, eta=1.0, lambdas=(0.7, 0.3)),
230
- }
231
- args.output.mkdir(parents=True, exist_ok=True)
232
- for name, latent in outputs.items():
233
- if not torch.isfinite(latent).all():
234
- raise RuntimeError(f"Nonfinite {name} output; reduce guidance or check training")
235
- sequences = decode(latent, esm, tokenizer, stats, args.min_polar)
236
- assert all(sum(a in POLAR_RESIDUES for a in s) >= args.min_polar for s in sequences)
237
- fasta = "".join(f">{name}_{i+1}\n{seq}\n" for i, seq in enumerate(sequences))
238
- (args.output / f"{name}.fasta").write_text(fasta)
239
- count = sum(a in POLAR_RESIDUES for a in sequences[0])
240
- print(f"{name}: {sequences[0]} polar/charged residues={count}")
241
- print(" decoded mean composition proxies:", composition_proxies(sequences).mean(0).tolist())
242
- torch.save({"standardized_latents": {k: v.cpu() for k, v in outputs.items()},
243
- "model": model.state_dict(), "reward_model": reward_model.state_dict(),
244
- "stats": stats, "esm_name": ESM_NAME,
245
- "length": model.length, "dim": model.dim, "min_polar": args.min_polar,
246
- "polar_residues": POLAR_RESIDUES}, args.output / "results.pt")
247
- print(f"Saved latent tensors and FASTA sequences to {args.output}")
248
-
249
-
250
- if __name__ == "__main__":
251
- main()