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

Add Lecture 2 MNIST code and checkpoint; place guidance materials in Lecture 3

Browse files
lecture_2/README.md ADDED
@@ -0,0 +1,115 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # Lecture 2 · MNIST Flow Matching
2
+
3
+ This lecture contains one Python script:
4
+ [flow_matching_unet_lecture.py](flow_matching_unet_lecture.py).
5
+ It covers MNIST downloading and data loading, a simple U-Net class, conditional
6
+ flow-matching training, and Euler sampling. Digit labels are ignored during
7
+ training, so generation is unconditional.
8
+
9
+ ## Train and generate
10
+
11
+ After installing the [repository requirements](../requirements.txt), run this
12
+ command from the repository root:
13
+
14
+ ```bash
15
+ python lecture_2/flow_matching_unet_lecture.py
16
+ ```
17
+
18
+ The script downloads the 60,000-image MNIST training split into `data/`,
19
+ normalizes pixel values to the range [-1, 1], trains for 20 epochs, and saves
20
+ these files in `flow_matching_outputs/`:
21
+
22
+ | File | Contents |
23
+ | --- | --- |
24
+ | `flow_unet_mnist.pt` | Learned model weights, loss history, channel width, and training seed |
25
+ | `samples.png` | A grid of 16 generated images |
26
+ | `trajectory.png` | Four samples progressing from noise to images |
27
+
28
+ Paths are relative to the working directory. Change the constants in section 02
29
+ to adjust epochs, batch size, sampling steps, or output paths. CUDA, Apple MPS,
30
+ and CPU are selected automatically according to availability.
31
+
32
+ ## Saved checkpoint
33
+
34
+ The bundled [flow_unet_mnist.pt](https://huggingface.co/ChatterjeeLab/CIS6270/resolve/main/lecture_2/flow_unet_mnist.pt?download=true)
35
+ was trained using the script's model, training function, and sampler.
36
+
37
+ | Setting | Saved run |
38
+ | --- | --- |
39
+ | Training data | All 60,000 MNIST training images |
40
+ | Completed epochs | 5 |
41
+ | Optimizer updates | 2,345 |
42
+ | Batch size | 128 |
43
+ | Optimizer | AdamW, learning rate 0.0002, weight decay 0 |
44
+ | U-Net widths | 32, 64, 128 channels |
45
+ | Trainable parameters | 471,265 |
46
+ | Training seed | 7 |
47
+ | Sampling | 100 forward Euler steps |
48
+
49
+ The saved run used CPU execution with PyTorch 2.14.0 and TorchVision 0.29.0.
50
+ The checkpoint contains `model`, `losses`, `base_channels`, `seed`, and
51
+ `epochs_completed`. It has five completed epochs; the script's default for a
52
+ new training run is twenty. Optimizer state is not included.
53
+
54
+ ## Generate from the checkpoint
55
+
56
+ In Python or a notebook, with the repository root as the working directory:
57
+
58
+ ```python
59
+ from pathlib import Path
60
+ import torch
61
+ from torchvision.utils import save_image
62
+ from lecture_2.flow_matching_unet_lecture import DEVICE, FlowUNet, sample_images
63
+
64
+ # Read the supplied weights and reconstruct the same U-Net.
65
+ checkpoint = torch.load(
66
+ "lecture_2/flow_unet_mnist.pt", map_location="cpu", weights_only=True,
67
+ )
68
+ model = FlowUNet(base_channels=checkpoint["base_channels"]).to(DEVICE)
69
+ model.load_state_dict(checkpoint["model"])
70
+
71
+ # Generate new images directly from Gaussian noise.
72
+ torch.manual_seed(8)
73
+ samples, trajectory = sample_images(model, num_samples=16, steps=100)
74
+
75
+ # Convert the generated values to display pixels and save the results.
76
+ output = Path("flow_matching_outputs")
77
+ output.mkdir(exist_ok=True)
78
+ save_image(((samples.cpu() + 1) / 2).clamp(0, 1),
79
+ output / "checkpoint_samples.png", nrow=4)
80
+ display_path = ((trajectory[:4] + 1) / 2).clamp(0, 1)
81
+ save_image(display_path.flatten(0, 1), output / "checkpoint_trajectory.png",
82
+ nrow=trajectory.shape[1])
83
+ ```
84
+
85
+ Importing the module does not start training or download MNIST. The images are
86
+ generated entirely from the saved weights and freshly sampled Gaussian noise.
87
+ Clipping occurs only when preparing the PNGs; intermediate ODE states remain
88
+ unclipped.
89
+
90
+ ## Selected generated examples
91
+
92
+ These 16 images were manually selected from 512 outputs of the five-epoch
93
+ checkpoint. They illustrate the clearest generated digits in that batch.
94
+ The sampling example above generates a fresh, unselected batch.
95
+
96
+ ![Selected MNIST generations](mnist_selected_examples.png)
97
+
98
+ The following trajectories show the first four selected images at times
99
+ 0, 0.25, 0.5, 0.75, and 1:
100
+
101
+ ![Selected noise-to-image trajectories](mnist_selected_trajectories.png)
102
+
103
+ ## Script sections
104
+
105
+ | Sections | Lecture material |
106
+ | --- | --- |
107
+ | 01–03 | Imports, configuration, and MNIST data loading |
108
+ | 04 | Noise, interpolation time, intermediate image, and target velocity |
109
+ | 05–07 | Convolution blocks, the U-Net, and its forward pass |
110
+ | 08 | Loss, backpropagation, and optimizer updates |
111
+ | 09 | Sampling by integrating the learned velocity |
112
+ | 10–13 | Complete training run, checkpoint saving, and image export |
113
+
114
+ Continue with [Lecture 3](../lecture_3/README.md) for flow matching, diffusion,
115
+ and guidance using ESM-2 residue embeddings.
lecture_2/flow_matching_unet_lecture.py ADDED
@@ -0,0 +1,175 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """CIS 6270, Lecture 2: unconditional MNIST generation with flow matching.
2
+
3
+ Install: python -m pip install torch torchvision
4
+ Run: python flow_matching_unet_lecture.py
5
+ Notation: x0 = prior noise, x1 = data, xt = state, ut = target, v_pred = prediction.
6
+ """
7
+
8
+ # %% 01. Imports
9
+ from pathlib import Path # Dataset and output paths.
10
+
11
+ import torch # Tensors, gradients, and optimization.
12
+ import torch.nn as nn # Layers and model classes.
13
+ import torch.nn.functional as F # Resize feature maps in the decoder.
14
+ from torch.utils.data import DataLoader # Shuffle and batch real images.
15
+ from torchvision import datasets
16
+ from torchvision.transforms import v2 # Convert and normalize image pixels.
17
+ from torchvision.utils import save_image # Export tensors as PNG grids.
18
+
19
+
20
+ # %% 02. Configuration
21
+ BATCH_SIZE = 128 # Images per training batch.
22
+ EPOCHS = 20 # Passes through the training set.
23
+ LEARNING_RATE = 2e-4 # Optimizer step size.
24
+ BASE_CHANNELS = 32 # U-Net widths: 32, 64, 128.
25
+ SAMPLE_STEPS = 100 # Euler steps from t=0 to t=1.
26
+ NUM_SAMPLES = 16 # Generate a 4-by-4 image grid.
27
+ SEED = 7 # Seed for initialization and training.
28
+ DATA_DIR = Path("data") # MNIST download location.
29
+ OUTPUT_DIR = Path("flow_matching_outputs") # Checkpoint and generated image folder.
30
+
31
+ DEVICE = torch.device(
32
+ "cuda" if torch.cuda.is_available() # NVIDIA GPU.
33
+ else "mps" if torch.backends.mps.is_available() # Apple GPU.
34
+ else "cpu" # CPU fallback.
35
+ )
36
+
37
+
38
+ # %% 03. Load real images
39
+ def make_loader():
40
+ transform = v2.Compose([
41
+ v2.ToImage(), # Channel-first image: [1, 28, 28].
42
+ v2.ToDtype(torch.float32, scale=True), # Integer pixels 0-255 -> floats 0-1.
43
+ v2.Normalize(mean=(0.5,), std=(0.5,)), # Real data pixels: [0, 1] -> [-1, 1].
44
+ ])
45
+ dataset = datasets.MNIST(
46
+ root=DATA_DIR, train=True, download=True, transform=transform, # Use the training split.
47
+ )
48
+ return DataLoader(
49
+ dataset, batch_size=BATCH_SIZE, shuffle=True, # Batch shape: [B, 1, 28, 28].
50
+ num_workers=0, pin_memory=(DEVICE.type == "cuda"), # Load in the main process.
51
+ )
52
+
53
+
54
+ # %% 04. Sample the conditional path
55
+ def sample_conditional_path(x1): # x1 contains normalized real images.
56
+ x0 = torch.randn_like(x1) # Independent Gaussian prior; same shape as x1.
57
+ t = torch.rand(x1.shape[0], device=x1.device) # One random time per image; [B].
58
+ t_image = t[:, None, None, None] # [B, 1, 1, 1]; share time across pixels.
59
+ xt = (1.0 - t_image) * x0 + t_image * x1 # Noise at t=0; real image at t=1.
60
+ ut = x1 - x0 # Derivative of the straight interpolation.
61
+ return xt, t, ut # Network input, time, and target velocity.
62
+
63
+
64
+ # %% 05. Define a simple convolution block
65
+ def conv_block(in_channels, out_channels):
66
+ return nn.Sequential(
67
+ nn.Conv2d(in_channels, out_channels, 3, padding=1), nn.SiLU(), # Keep H and W.
68
+ nn.Conv2d(out_channels, out_channels, 3, padding=1), nn.SiLU(), # Refine features.
69
+ )
70
+
71
+
72
+ # %% 06. Define the small U-Net
73
+ class FlowUNet(nn.Module):
74
+ def __init__(self, base_channels=BASE_CHANNELS):
75
+ super().__init__()
76
+ c = base_channels # Base width; c=32 by default.
77
+ self.encoder1 = conv_block(2, c) # Image + time -> c channels.
78
+ self.encoder2 = conv_block(c, 2 * c) # Encoder: c -> 2c.
79
+ self.middle = conv_block(2 * c, 4 * c) # Bottleneck: 2c -> 4c.
80
+ self.decoder2 = conv_block(6 * c, 2 * c) # Upsampled 4c + skip 2c -> 2c.
81
+ self.decoder1 = conv_block(3 * c, c) # Upsampled 2c + skip c -> c.
82
+ self.output = nn.Conv2d(c, 1, 1) # One signed velocity per pixel.
83
+ self.pool = nn.MaxPool2d(2) # Halve the image height and width.
84
+
85
+ # %% 07. Compute the velocity field
86
+ def forward(self, xt, t):
87
+ t_image = t[:, None, None, None].expand_as(xt) # Constant time channel.
88
+ x = torch.cat([xt, t_image], dim=1) # [B, 2, 28, 28].
89
+ skip1 = self.encoder1(x) # Keep full-resolution features.
90
+ skip2 = self.encoder2(self.pool(skip1)) # 28x28 -> 14x14.
91
+ x = self.middle(self.pool(skip2)) # 14x14 -> 7x7.
92
+ x = F.interpolate(x, size=skip2.shape[-2:], mode="nearest") # 7x7 -> 14x14.
93
+ x = self.decoder2(torch.cat([x, skip2], dim=1)) # Restore encoder detail.
94
+ x = F.interpolate(x, size=skip1.shape[-2:], mode="nearest") # 14x14 -> 28x28.
95
+ x = self.decoder1(torch.cat([x, skip1], dim=1)) # Restore fine detail.
96
+ v_pred = self.output(x) # Predict v_theta(xt, t).
97
+ return v_pred # [B, 1, 28, 28], matching xt.
98
+
99
+
100
+ # %% 08. Train for one epoch
101
+ def train_one_epoch(model, loader, optimizer):
102
+ model.train() # Select training mode.
103
+ total_loss, num_images = 0.0, 0 # Accumulate an image-weighted epoch loss.
104
+ for x1, _ in loader: # Ignore digit labels: unconditional generation.
105
+ x1 = x1.to(DEVICE) # Move real images to the model device.
106
+ xt, t, ut = sample_conditional_path(x1) # Fresh noise and times for this batch.
107
+ v_pred = model(xt, t) # The model receives only xt and t.
108
+ squared_error = (v_pred - ut).square() # Velocity error at every pixel.
109
+ loss = squared_error.flatten(1).sum(1).mean() # Sum pixels, average images.
110
+ optimizer.zero_grad(set_to_none=True) # Clear old gradients.
111
+ loss.backward() # Compute gradients with respect to weights.
112
+ optimizer.step() # Update the U-Net weights.
113
+ total_loss += loss.item() * x1.shape[0] # Weight by the actual batch size.
114
+ num_images += x1.shape[0] # Include the smaller final batch.
115
+ return total_loss / num_images # Average training loss per image.
116
+
117
+ # Save snapshots every quarter of the default 100-step run; keep ODE states unclipped.
118
+ # %% 09. Generate images with Euler integration
119
+ @torch.no_grad() # Sampling needs no gradient tracking.
120
+ def sample_images(model, num_samples=NUM_SAMPLES, steps=SAMPLE_STEPS):
121
+ if num_samples < 1 or steps < 1:
122
+ raise ValueError("num_samples and steps must be positive")
123
+ model.eval() # Select evaluation mode.
124
+ device = next(model.parameters()).device # Use the model device.
125
+ xt = torch.randn(num_samples, 1, 28, 28, device=device) # Initial x0.
126
+ dt = 1.0 / steps # Positive time increment.
127
+ snapshots = [xt.cpu().clone()] # Keep the initial noise for plotting.
128
+ for step in range(steps): # Integrate forward in time.
129
+ t = torch.full((num_samples,), step * dt, device=device) # Shared time grid.
130
+ v_pred = model(xt, t) # Reevaluate velocity at the current state.
131
+ xt = xt + dt * v_pred # Euler: state += time increment * velocity.
132
+ if (step + 1) % max(steps // 4, 1) == 0 or step + 1 == steps:
133
+ snapshots.append(xt.cpu().clone()) # Save intermediate and final states.
134
+ return xt, torch.stack(snapshots, dim=1) # Raw images; [B, snapshots, 1, 28, 28].
135
+
136
+
137
+ # %% 10. Set up the complete run
138
+ def main():
139
+ torch.manual_seed(SEED) # Seed training randomness.
140
+ OUTPUT_DIR.mkdir(parents=True, exist_ok=True) # Create the output folder.
141
+ loader = make_loader() # Download, normalize, and batch MNIST.
142
+ model = FlowUNet().to(DEVICE) # Create the velocity network.
143
+ optimizer = torch.optim.AdamW(
144
+ model.parameters(), lr=LEARNING_RATE, weight_decay=0.0, # Optimize the flow loss.
145
+ )
146
+ num_parameters = sum(p.numel() for p in model.parameters()) # Count weights and biases.
147
+ print(f"device={DEVICE} parameters={num_parameters:,}")
148
+
149
+ # %% 11. Train and save the learned parameters
150
+ losses = [] # Record the training loss after each epoch.
151
+ for epoch in range(1, EPOCHS + 1): # Complete EPOCHS passes through the data.
152
+ loss = train_one_epoch(model, loader, optimizer) # Run the training loop.
153
+ losses.append(loss) # Keep the loss history.
154
+ print(f"epoch={epoch:02d} squared_norm_loss={loss:.3f}")
155
+ torch.save({
156
+ "model": model.state_dict(), "losses": losses, # Weights and training curve.
157
+ "base_channels": BASE_CHANNELS, # Width needed to reconstruct the model.
158
+ "seed": SEED, # Record the training seed.
159
+ }, OUTPUT_DIR / "flow_unet_mnist.pt")
160
+
161
+ # %% 12. Sample and save images
162
+ torch.manual_seed(SEED + 1) # Fix the starting noise for generation.
163
+ samples, trajectory = sample_images(model) # Integrate the learned ODE.
164
+ display_samples = ((samples.cpu() + 1.0) / 2.0).clamp(0.0, 1.0) # Clip for display only.
165
+ save_image(display_samples, OUTPUT_DIR / "samples.png", nrow=4) # Final images.
166
+ display_path = ((trajectory[:4] + 1.0) / 2.0).clamp(0.0, 1.0) # First four sample paths.
167
+ save_image(
168
+ display_path.flatten(0, 1), OUTPUT_DIR / "trajectory.png", # One row per sample.
169
+ nrow=trajectory.shape[1], # Time increases across columns.
170
+ )
171
+
172
+
173
+ # %% 13. Run the script
174
+ if __name__ == "__main__": # Execute when launched directly.
175
+ main() # Load data, train, sample, and save.
lecture_2/flow_unet_mnist.pt ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:20373a7d33489c0b596a4bd2f54750795570eda7dd0bef1f064b51a1367bd2d7
3
+ size 1893377
lecture_2/mnist_selected_examples.png ADDED
lecture_2/mnist_selected_trajectories.png ADDED
lecture_3/GUIDANCE_NOTES.md ADDED
@@ -0,0 +1,156 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
lecture_3/README.md ADDED
@@ -0,0 +1,310 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # Lecture 3 · Continuous Generative Models
2
+
3
+ These examples accompany Lecture 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 lecture_3/esm2_flow_guidance.py --epochs 200 --samples 8
50
+ python lecture_3/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 `lecture_3/.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 lecture_3/esm2_flow_guidance.py --epochs 2 --samples 2 --output outputs/flow_smoke
66
+ python lecture_3/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
+ lecture_3/
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 lecture_3/esm2_flow_guidance.py --data my_sequences.csv --output outputs/my_flow
142
+ python lecture_3/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 lecture_3/esm2_flow_guidance.py --min-polar 16
258
+
259
+ # Remove the minimum-count constraint.
260
+ python lecture_3/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.
lecture_3/esm2_diffusion_guidance.py ADDED
@@ -0,0 +1,264 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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()
lecture_3/esm2_example.csv ADDED
@@ -0,0 +1,65 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
lecture_3/esm2_flow_guidance.py ADDED
@@ -0,0 +1,251 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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()