Add Lecture 2 MNIST code and checkpoint; place guidance materials in Lecture 3
Browse files- lecture_2/README.md +115 -0
- lecture_2/flow_matching_unet_lecture.py +175 -0
- lecture_2/flow_unet_mnist.pt +3 -0
- lecture_2/mnist_selected_examples.png +0 -0
- lecture_2/mnist_selected_trajectories.png +0 -0
- lecture_3/GUIDANCE_NOTES.md +156 -0
- lecture_3/README.md +310 -0
- lecture_3/esm2_diffusion_guidance.py +264 -0
- lecture_3/esm2_example.csv +65 -0
- lecture_3/esm2_flow_guidance.py +251 -0
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 |
+

|
| 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 |
+

|
| 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()
|