Add Lecture 4 discrete diffusion and Lecture 5 discrete flow matching
Browse filesAdd complete CPU training and generation examples for seven diffusion and guidance methods and eight flow methods. Include synthetic DNA, mathematical checks, slide links, and recorded example outputs. Update the course index and setup instructions.
This view is limited to 50 files because it contains too many changes. See raw diff
- README.md +52 -1
- lecture_4/.gitignore +5 -0
- lecture_4/README.md +85 -0
- lecture_4/SLIDE_CODE_MAP.md +21 -0
- lecture_4/common.py +134 -0
- lecture_4/data/README.md +1 -0
- lecture_4/data/dna_train.tsv +257 -0
- lecture_4/diffusion.py +223 -0
- lecture_4/examples/block.py +7 -0
- lecture_4/examples/cfg.py +7 -0
- lecture_4/examples/classifier_exact.py +7 -0
- lecture_4/examples/classifier_gradient.py +7 -0
- lecture_4/examples/mdlm.py +7 -0
- lecture_4/examples/peptune.py +7 -0
- lecture_4/examples/udlm.py +7 -0
- lecture_4/lecture_core.py +136 -0
- lecture_4/numerical_examples.py +24 -0
- lecture_4/requirements.txt +3 -0
- lecture_4/run.py +127 -0
- lecture_4/run_all.py +15 -0
- lecture_4/tests/test_mathematics.py +50 -0
- lecture_4/verified_examples/README.md +15 -0
- lecture_4/verified_examples/block/config.json +18 -0
- lecture_4/verified_examples/block/losses.json +162 -0
- lecture_4/verified_examples/block/report.json +13 -0
- lecture_4/verified_examples/block/samples.txt +8 -0
- lecture_4/verified_examples/cfg/config.json +18 -0
- lecture_4/verified_examples/cfg/losses.json +162 -0
- lecture_4/verified_examples/cfg/report.json +13 -0
- lecture_4/verified_examples/cfg/samples.txt +8 -0
- lecture_4/verified_examples/classifier-exact/config.json +18 -0
- lecture_4/verified_examples/classifier-exact/losses.json +162 -0
- lecture_4/verified_examples/classifier-exact/report.json +15 -0
- lecture_4/verified_examples/classifier-exact/samples.txt +8 -0
- lecture_4/verified_examples/classifier-gradient/config.json +18 -0
- lecture_4/verified_examples/classifier-gradient/losses.json +162 -0
- lecture_4/verified_examples/classifier-gradient/report.json +15 -0
- lecture_4/verified_examples/classifier-gradient/samples.txt +8 -0
- lecture_4/verified_examples/environment.json +7 -0
- lecture_4/verified_examples/mdlm/config.json +18 -0
- lecture_4/verified_examples/mdlm/losses.json +162 -0
- lecture_4/verified_examples/mdlm/report.json +13 -0
- lecture_4/verified_examples/mdlm/samples.txt +8 -0
- lecture_4/verified_examples/peptune/config.json +18 -0
- lecture_4/verified_examples/peptune/losses.json +162 -0
- lecture_4/verified_examples/peptune/report.json +350 -0
- lecture_4/verified_examples/peptune/samples.txt +3 -0
- lecture_4/verified_examples/udlm/config.json +18 -0
- lecture_4/verified_examples/udlm/losses.json +162 -0
- lecture_4/verified_examples/udlm/report.json +14 -0
README.md
CHANGED
|
@@ -32,13 +32,16 @@ and implementation notes; additional directories will accompany later lectures.
|
|
| 32 |
| --- | --- | --- |
|
| 33 |
| 2 | Unconditional MNIST image generation with flow matching and a simple U-Net | [Guide](lecture_2/README.md) · [Script](lecture_2/flow_matching_unet_lecture.py) · [Saved checkpoint](https://huggingface.co/ChatterjeeLab/CIS6270/resolve/main/lecture_2/flow_unet_mnist.pt?download=true) |
|
| 34 |
| 3 | Flow matching, diffusion, and guidance for ESM-2 residue embeddings | [Guide](lecture_3/README.md) · [Flow matching](lecture_3/esm2_flow_guidance.py) · [Diffusion](lecture_3/esm2_diffusion_guidance.py) |
|
|
|
|
|
|
|
| 35 |
|
| 36 |
## Installation
|
| 37 |
|
| 38 |
Use Python 3.11, or another compatible Python version at least 3.10, in a new
|
| 39 |
virtual environment. The shared `requirements.txt` pins PyTorch 2.9.1,
|
| 40 |
TorchVision 0.24.1, and Transformers 4.57.6. Lecture 2 uses PyTorch and
|
| 41 |
-
TorchVision; Lecture 3 also uses Transformers.
|
|
|
|
| 42 |
|
| 43 |
```bash
|
| 44 |
git clone https://huggingface.co/ChatterjeeLab/CIS6270
|
|
@@ -102,6 +105,45 @@ The [lecture guide](lecture_3/README.md) describes the data format, training
|
|
| 102 |
and sampling settings, property calculations, normalization, and residue-count
|
| 103 |
constraint, with commands for using a custom dataset.
|
| 104 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 105 |
## Repository organization
|
| 106 |
|
| 107 |
| Location | Contents |
|
|
@@ -109,6 +151,8 @@ constraint, with commands for using a custom dataset.
|
|
| 109 |
| Repository root | Course index, installation requirements, and license |
|
| 110 |
| [`lecture_2/`](https://huggingface.co/ChatterjeeLab/CIS6270/tree/main/lecture_2) | One MNIST flow-matching script, guide, trained checkpoint, and selected example images |
|
| 111 |
| [`lecture_3/`](https://huggingface.co/ChatterjeeLab/CIS6270/tree/main/lecture_3) | ESM-2 flow and diffusion guidance scripts, sequence data, guide, and mathematical notes |
|
|
|
|
|
|
|
| 112 |
| [`tests/`](https://huggingface.co/ChatterjeeLab/CIS6270/tree/main/tests) | Offline checks for the Lecture 3 examples |
|
| 113 |
|
| 114 |
Installation instructions and the lecture index are maintained at the
|
|
@@ -119,11 +163,18 @@ references accompany the corresponding code.
|
|
| 119 |
|
| 120 |
```bash
|
| 121 |
python -m unittest discover -s tests -v
|
|
|
|
|
|
|
| 122 |
```
|
| 123 |
|
| 124 |
The Lecture 3 unit tests cover property annotations, scalarization weights, reward
|
| 125 |
gradients, DDPM schedule indexing, and constrained decoding.
|
| 126 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 127 |
## License
|
| 128 |
|
| 129 |
The repository code is distributed under the
|
|
|
|
| 32 |
| --- | --- | --- |
|
| 33 |
| 2 | Unconditional MNIST image generation with flow matching and a simple U-Net | [Guide](lecture_2/README.md) · [Script](lecture_2/flow_matching_unet_lecture.py) · [Saved checkpoint](https://huggingface.co/ChatterjeeLab/CIS6270/resolve/main/lecture_2/flow_unet_mnist.pt?download=true) |
|
| 34 |
| 3 | Flow matching, diffusion, and guidance for ESM-2 residue embeddings | [Guide](lecture_3/README.md) · [Flow matching](lecture_3/esm2_flow_guidance.py) · [Diffusion](lecture_3/esm2_diffusion_guidance.py) |
|
| 35 |
+
| 4 | Discrete diffusion, masked and uniform corruption, block generation, and guidance | [Guide](lecture_4/README.md) · [Training and generation](lecture_4/run.py) · [Slide code map](lecture_4/SLIDE_CODE_MAP.md) |
|
| 36 |
+
| 5 | Discrete flow matching, Dirichlet and Fisher paths, Gumbel-Softmax, rectification, and multi-objective generation | [Guide](lecture_5/README.md) · [Training and generation](lecture_5/run.py) · [Slide code map](lecture_5/SLIDE_CODE_MAP.md) |
|
| 37 |
|
| 38 |
## Installation
|
| 39 |
|
| 40 |
Use Python 3.11, or another compatible Python version at least 3.10, in a new
|
| 41 |
virtual environment. The shared `requirements.txt` pins PyTorch 2.9.1,
|
| 42 |
TorchVision 0.24.1, and Transformers 4.57.6. Lecture 2 uses PyTorch and
|
| 43 |
+
TorchVision; Lecture 3 also uses Transformers. Lectures 4 and 5 use PyTorch,
|
| 44 |
+
NumPy, and SciPy. Their lecture folders also provide minimal requirements.
|
| 45 |
|
| 46 |
```bash
|
| 47 |
git clone https://huggingface.co/ChatterjeeLab/CIS6270
|
|
|
|
| 105 |
and sampling settings, property calculations, normalization, and residue-count
|
| 106 |
constraint, with commands for using a custom dataset.
|
| 107 |
|
| 108 |
+
## Lecture 4 - Discrete diffusion
|
| 109 |
+
|
| 110 |
+
Train small DNA denoisers and generate sequences with MDLM, UDLM, block diffusion,
|
| 111 |
+
classifier-free guidance, exact and gradient-based classifier guidance, and a
|
| 112 |
+
PepTune-style search. The [guide](lecture_4/README.md) includes each method's
|
| 113 |
+
command and mathematical assumptions. The [code map](lecture_4/SLIDE_CODE_MAP.md)
|
| 114 |
+
links the slide walkthroughs to their functions.
|
| 115 |
+
|
| 116 |
+
From the repository root, run the complete MDLM example.
|
| 117 |
+
|
| 118 |
+
```bash
|
| 119 |
+
python lecture_4/run.py --method mdlm --data lecture_4/data/dna_train.tsv --out lecture_4/outputs/mdlm
|
| 120 |
+
python lecture_4/run.py --method mdlm --mode sample --out lecture_4/outputs/mdlm
|
| 121 |
+
```
|
| 122 |
+
|
| 123 |
+
The script trains, saves a checkpoint, and writes generated DNA and loss logs.
|
| 124 |
+
The bundled data are synthetic, and the guidance objectives are explicit toy
|
| 125 |
+
properties. No pretrained model or external dataset is required.
|
| 126 |
+
|
| 127 |
+
## Lecture 5 - Discrete flow matching
|
| 128 |
+
|
| 129 |
+
Start with Gat et al.'s discrete flow matching, then run Dirichlet, Fisher,
|
| 130 |
+
Gumbel-Softmax, rectified flow, ReDi, MOG-DFM, and AReUReDi examples. Each method
|
| 131 |
+
has a complete training and generation command in the [guide](lecture_5/README.md).
|
| 132 |
+
|
| 133 |
+
```bash
|
| 134 |
+
python lecture_5/run.py --method gat --data lecture_5/data/dna_train.tsv --out lecture_5/outputs/gat
|
| 135 |
+
python lecture_5/run.py --method gat --mode sample --out lecture_5/outputs/gat
|
| 136 |
+
```
|
| 137 |
+
|
| 138 |
+
The current [Discrete Generation slide deck](https://docs.google.com/presentation/d/1qcPdF4SYTq4VdTN_5w9GjviLRPr2HFtbuYMpbfLrc_M/edit)
|
| 139 |
+
contains both lectures. Diffusion code through PepTune belongs to Lecture 4.
|
| 140 |
+
Classic DFM and the subsequent flow methods belong to Lecture 5. The code maps
|
| 141 |
+
use stable slide links, so added transition slides do not break the mapping.
|
| 142 |
+
|
| 143 |
+
Both folders include numerical examples, mathematical tests, and saved results
|
| 144 |
+
from seeded CPU runs. The guides explain finite endpoint approximations and
|
| 145 |
+
classroom simplifications for each method.
|
| 146 |
+
|
| 147 |
## Repository organization
|
| 148 |
|
| 149 |
| Location | Contents |
|
|
|
|
| 151 |
| Repository root | Course index, installation requirements, and license |
|
| 152 |
| [`lecture_2/`](https://huggingface.co/ChatterjeeLab/CIS6270/tree/main/lecture_2) | One MNIST flow-matching script, guide, trained checkpoint, and selected example images |
|
| 153 |
| [`lecture_3/`](https://huggingface.co/ChatterjeeLab/CIS6270/tree/main/lecture_3) | ESM-2 flow and diffusion guidance scripts, sequence data, guide, and mathematical notes |
|
| 154 |
+
| [`lecture_4/`](https://huggingface.co/ChatterjeeLab/CIS6270/tree/main/lecture_4) | Seven discrete diffusion and guidance examples, synthetic DNA, slide code map, and verified outputs |
|
| 155 |
+
| [`lecture_5/`](https://huggingface.co/ChatterjeeLab/CIS6270/tree/main/lecture_5) | Eight discrete and simplex flow examples, synthetic DNA, slide code map, and verified outputs |
|
| 156 |
| [`tests/`](https://huggingface.co/ChatterjeeLab/CIS6270/tree/main/tests) | Offline checks for the Lecture 3 examples |
|
| 157 |
|
| 158 |
Installation instructions and the lecture index are maintained at the
|
|
|
|
| 163 |
|
| 164 |
```bash
|
| 165 |
python -m unittest discover -s tests -v
|
| 166 |
+
python -m unittest discover -s lecture_4/tests -v
|
| 167 |
+
python -m unittest discover -s lecture_5/tests -v
|
| 168 |
```
|
| 169 |
|
| 170 |
The Lecture 3 unit tests cover property annotations, scalarization weights, reward
|
| 171 |
gradients, DDPM schedule indexing, and constrained decoding.
|
| 172 |
|
| 173 |
+
Lecture 4 tests check reverse KL losses and guidance calculations. Lecture 5
|
| 174 |
+
tests check the master equation, Fisher geometry, Gumbel path derivatives, and
|
| 175 |
+
MH detailed balance. Run a short end-to-end check of every method from its
|
| 176 |
+
lecture folder with `python run_all.py --quick`.
|
| 177 |
+
|
| 178 |
## License
|
| 179 |
|
| 180 |
The repository code is distributed under the
|
lecture_4/.gitignore
ADDED
|
@@ -0,0 +1,5 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
__pycache__/
|
| 2 |
+
.venv/
|
| 3 |
+
outputs/
|
| 4 |
+
*.pt
|
| 5 |
+
.pytest_cache/
|
lecture_4/README.md
ADDED
|
@@ -0,0 +1,85 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# CIS 6270 - Lecture 4 - Discrete Diffusion
|
| 2 |
+
|
| 3 |
+
Course hub: [ChatterjeeLab/CIS6270 on Hugging Face](https://huggingface.co/ChatterjeeLab/CIS6270).
|
| 4 |
+
|
| 5 |
+
Complete CPU training and generation examples for the code walkthroughs in the [Discrete Generation lecture](https://docs.google.com/presentation/d/1qcPdF4SYTq4VdTN_5w9GjviLRPr2HFtbuYMpbfLrc_M/edit). This folder ends with PepTune. Classic discrete flow matching and the later methods are in [Lecture 5](../lecture_5/README.md).
|
| 6 |
+
|
| 7 |
+
These are small, inspectable implementations of the mathematical mechanisms. The DNA data and two objective functions are synthetic. They are not benchmark reproductions or biological property predictors.
|
| 8 |
+
|
| 9 |
+
## Start with the complete MDLM example
|
| 10 |
+
|
| 11 |
+
Python 3.11 or later, CPU. No downloaded data or pretrained checkpoint is required.
|
| 12 |
+
|
| 13 |
+
From the course repository root, enter this lecture folder.
|
| 14 |
+
|
| 15 |
+
```bash
|
| 16 |
+
cd lecture_4
|
| 17 |
+
python -m venv .venv
|
| 18 |
+
source .venv/bin/activate
|
| 19 |
+
pip install -r requirements.txt
|
| 20 |
+
python run.py --method mdlm --data data/dna_train.tsv --out outputs/mdlm
|
| 21 |
+
```
|
| 22 |
+
|
| 23 |
+
This command loads the DNA examples, reserves the final 20% for validation, trains a transformer denoiser, writes a checkpoint, and generates new sequences. Each output folder contains `config.json`, `losses.json`, `report.json`, `samples.txt`, and `checkpoint.pt`.
|
| 24 |
+
|
| 25 |
+
Inspect `data/dna_train.tsv` before training. `label=1` means GC fraction at least 0.6. It is a deterministic toy label, not measured activity. Omitting `--data` generates synthetic examples at `--length`.
|
| 26 |
+
|
| 27 |
+
```bash
|
| 28 |
+
# Continue with the saved weights for generation only.
|
| 29 |
+
python run.py --method mdlm --mode sample --out outputs/mdlm
|
| 30 |
+
# An entire small experiment using the four-base vocabulary.
|
| 31 |
+
python examples/mdlm.py --length 4 --train-steps 160 --samples 8
|
| 32 |
+
# Run every method, or a short CPU integration check.
|
| 33 |
+
python run_all.py
|
| 34 |
+
python run_all.py --quick
|
| 35 |
+
python -m unittest discover -s tests -v
|
| 36 |
+
python numerical_examples.py
|
| 37 |
+
```
|
| 38 |
+
|
| 39 |
+
For `--mode sample`, use the same `--method`, `--width`, and `--length` as the checkpoint. Checkpoints contain model weights and metadata; they do not contain optimizer state for resuming optimization.
|
| 40 |
+
|
| 41 |
+
## One shared backbone; small method changes
|
| 42 |
+
|
| 43 |
+
`lecture_core.py` contains the functions shown in the slides. `common.py` contains data, optimization, and output utilities. `diffusion.py` adds the complete samplers and guidance/search loops that make those modules runnable. `run.py` connects training to generation. Each file under `examples/` is a complete command-line entry point for one method and accepts the same options.
|
| 44 |
+
|
| 45 |
+
| Example | Training target | Generation procedure |
|
| 46 |
+
| --- | --- | --- |
|
| 47 |
+
| `mdlm` | Weighted cross-entropy at masked positions | Schedule-based reveals; preserve visible bases |
|
| 48 |
+
| `udlm` | Reverse-rate KL including the staying term | Adaptive reverse Euler; visible bases may be revised |
|
| 49 |
+
| `block` | Conditional masked loss on a sampled block, scaled by block count | Finish each block before extending the sequence |
|
| 50 |
+
| `cfg` / `classifier-free` | Label-conditioned MDLM with 15% label dropout | Geometric conditional/unconditional prediction blend |
|
| 51 |
+
| `classifier-exact` | MDLM plus a classifier trained on noisy sequences | Evaluate every replacement's log-value change and multiply rates |
|
| 52 |
+
| `classifier-gradient` | Same denoiser and noisy classifier | Approximate replacement log-value changes with one-hot gradients |
|
| 53 |
+
| `peptune` | MDLM backbone | MCTS selection, expansion, completion, Pareto rewards, and reward backup |
|
| 54 |
+
|
| 55 |
+
```bash
|
| 56 |
+
python examples/udlm.py --out outputs/udlm
|
| 57 |
+
python examples/block.py --block-size 2
|
| 58 |
+
python examples/cfg.py --label 1 --strength 2
|
| 59 |
+
python examples/classifier_exact.py --classifier-steps 300
|
| 60 |
+
python examples/classifier_gradient.py --classifier-steps 300
|
| 61 |
+
python examples/peptune.py --search-steps 100
|
| 62 |
+
```
|
| 63 |
+
|
| 64 |
+
MDLM and block losses sum over positions, then average over sequences. The uniform method uses the slide's linear schedule and integrates its training loss over `t in [0.02, 0.98]`; uniform initialization at 0.98 and stopping at 0.02 approximate the ideal endpoints. The output retains residual corruption. We do not reinterpret the UDLM parameter vector as a clean-token posterior for a final resampling step. The classifier-guided sampler uses adaptive steps to keep outgoing probabilities valid, then closes any residual masks at `t=0.005`. These finite endpoint conventions are reported in the outputs.
|
| 65 |
+
|
| 66 |
+
CFG here is a categorical prediction blend. It is not asserted to equal geometric mixing of every possible rate parameterization. Classifier guidance is approximate because the classifier is learned; the gradient version also approximates finite edit differences. `strength=0` recovers the unguided rate multiplier in the classifier routines.
|
| 67 |
+
|
| 68 |
+
The block implementation omits future blocks from the input entirely. This is a transparent teaching alternative to a full block-causal attention mask and KV-cache implementation. Training prefixes are clean; generation prefixes are sampled.
|
| 69 |
+
|
| 70 |
+
The PepTune example uses the search mechanism on DNA. Every A/C/G/T sequence is valid, so there is no invalid-SMILES penalty. Peptide tokenization, bond-dependent schedules, chemical validity, RoFormer pretraining, and the paper's biological predictors are outside this toy example. Its archive is the non-dominated subset of evaluated DNA candidates; it is not a global Pareto certificate.
|
| 71 |
+
|
| 72 |
+
## Read the math next to the implementation
|
| 73 |
+
|
| 74 |
+
The tests verify the masked reverse KL cancellation, the small-step UDLM KL limit, the numerical CFG example, and the classifier directional derivative. `numerical_examples.py` prints corruption, training-loss, and generation-step calculations from the slides. `SLIDE_CODE_MAP.md` maps every Lecture 4 code unit to its function.
|
| 75 |
+
|
| 76 |
+
The seeded example runs in `verified_examples/` record actual CPU results. Losses from different objectives are not directly comparable; one Monte Carlo validation draw is not a perplexity estimate. Sampling quality should be assessed with more seeds and longer training if extending these examples.
|
| 77 |
+
|
| 78 |
+
## Sources
|
| 79 |
+
|
| 80 |
+
- [MDLM - Simple and Effective Masked Diffusion Language Models](https://arxiv.org/abs/2406.07524)
|
| 81 |
+
- [UDLM - Simple Guidance Mechanisms for Discrete Diffusion Models](https://arxiv.org/abs/2412.10193)
|
| 82 |
+
- [Block Diffusion](https://arxiv.org/abs/2503.09573)
|
| 83 |
+
- [PepTune](https://arxiv.org/abs/2412.17780)
|
| 84 |
+
|
| 85 |
+
The source papers define the full research methods and experiments. Comments identify the classroom simplifications.
|
lecture_4/SLIDE_CODE_MAP.md
ADDED
|
@@ -0,0 +1,21 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# Code walkthrough map
|
| 2 |
+
|
| 3 |
+
Function names are stable even if the combined deck is later split or reordered. Each method has a full training and sampling command in the README.
|
| 4 |
+
|
| 5 |
+
| Slide unit | Code | Location |
|
| 6 |
+
| --- | --- | --- |
|
| 7 |
+
| [Code for the four-letter vocabulary](https://docs.google.com/presentation/d/1qcPdF4SYTq4VdTN_5w9GjviLRPr2HFtbuYMpbfLrc_M/edit#slide=id.cis4r2_u141_s5) | `complete experiment` | `lecture_core.py` |
|
| 8 |
+
| [Code for a small shared DNA network](https://docs.google.com/presentation/d/1qcPdF4SYTq4VdTN_5w9GjviLRPr2HFtbuYMpbfLrc_M/edit#slide=id.cis4r2_u142_s6) | `DNA` | `lecture_core.py` |
|
| 9 |
+
| [Code for the network forward pass](https://docs.google.com/presentation/d/1qcPdF4SYTq4VdTN_5w9GjviLRPr2HFtbuYMpbfLrc_M/edit#slide=id.cis4r2_u143_s5) | `forward` | `lecture_core.py` |
|
| 10 |
+
| [Code for one cross-entropy per DNA position](https://docs.google.com/presentation/d/1qcPdF4SYTq4VdTN_5w9GjviLRPr2HFtbuYMpbfLrc_M/edit#slide=id.cis4r2_u144_s3) | `token_ce` | `lecture_core.py` |
|
| 11 |
+
| [Code for MDLM training loss](https://docs.google.com/presentation/d/1qcPdF4SYTq4VdTN_5w9GjviLRPr2HFtbuYMpbfLrc_M/edit#slide=id.cis4r2_u145_s10) | `mdlm_loss` | `lecture_core.py` |
|
| 12 |
+
| [Code for one reusable optimization step](https://docs.google.com/presentation/d/1qcPdF4SYTq4VdTN_5w9GjviLRPr2HFtbuYMpbfLrc_M/edit#slide=id.cis4r2_u146_s8) | `train_step` | `lecture_core.py` |
|
| 13 |
+
| [Code for sample each categorical DNA vector](https://docs.google.com/presentation/d/1qcPdF4SYTq4VdTN_5w9GjviLRPr2HFtbuYMpbfLrc_M/edit#slide=id.cis4r2_u147_s5) | `draw` | `lecture_core.py` |
|
| 14 |
+
| [Code for MDLM generation](https://docs.google.com/presentation/d/1qcPdF4SYTq4VdTN_5w9GjviLRPr2HFtbuYMpbfLrc_M/edit#slide=id.cis4r2_u148_s13) | `mdlm_sample` | `lecture_core.py` |
|
| 15 |
+
| [Code for UDLM reverse rates](https://docs.google.com/presentation/d/1qcPdF4SYTq4VdTN_5w9GjviLRPr2HFtbuYMpbfLrc_M/edit#slide=id.cis4r2_u149_s7) | `udlm_rates` | `lecture_core.py` |
|
| 16 |
+
| [Code for the continuous UDLM loss](https://docs.google.com/presentation/d/1qcPdF4SYTq4VdTN_5w9GjviLRPr2HFtbuYMpbfLrc_M/edit#slide=id.cis4r2_u150_s10) | `udlm_loss` | `lecture_core.py` |
|
| 17 |
+
| [Code for train one conditional DNA block](https://docs.google.com/presentation/d/1qcPdF4SYTq4VdTN_5w9GjviLRPr2HFtbuYMpbfLrc_M/edit#slide=id.cis4r2_u151_s9) | `block_loss` | `lecture_core.py` |
|
| 18 |
+
| [Code for classifier-free categorical guidance](https://docs.google.com/presentation/d/1qcPdF4SYTq4VdTN_5w9GjviLRPr2HFtbuYMpbfLrc_M/edit#slide=id.cis4r2_u152_s5) | `geometric_cfg` | `lecture_core.py` |
|
| 19 |
+
| [Code for predictor guidance as a rate multiplier](https://docs.google.com/presentation/d/1qcPdF4SYTq4VdTN_5w9GjviLRPr2HFtbuYMpbfLrc_M/edit#slide=id.cis4r2_u153_s4) | `guide_rates` | `lecture_core.py` |
|
| 20 |
+
| [Code for preserve non-dominated DNA candidates](https://docs.google.com/presentation/d/1qcPdF4SYTq4VdTN_5w9GjviLRPr2HFtbuYMpbfLrc_M/edit#slide=id.cis4r2_u154_s6) | `pareto_filter` | `lecture_core.py` |
|
| 21 |
+
| [Code for run a complete MDLM experiment](https://docs.google.com/presentation/d/1qcPdF4SYTq4VdTN_5w9GjviLRPr2HFtbuYMpbfLrc_M/edit#slide=id.cis4r2_u180_s6) | `complete experiment` | `lecture_core.py` and `run.py` |
|
lecture_4/common.py
ADDED
|
@@ -0,0 +1,134 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""CPU teaching utilities shared by the two lecture folders."""
|
| 2 |
+
import csv
|
| 3 |
+
import json
|
| 4 |
+
import random
|
| 5 |
+
from pathlib import Path
|
| 6 |
+
import numpy as np
|
| 7 |
+
import torch
|
| 8 |
+
from torch import nn
|
| 9 |
+
from lecture_core import DNA, K, MASK, encode
|
| 10 |
+
|
| 11 |
+
|
| 12 |
+
def seed_all(seed):
|
| 13 |
+
random.seed(seed)
|
| 14 |
+
np.random.seed(seed)
|
| 15 |
+
torch.manual_seed(seed)
|
| 16 |
+
torch.set_num_threads(1)
|
| 17 |
+
|
| 18 |
+
|
| 19 |
+
def decode(tokens):
|
| 20 |
+
alphabet = 'ACGTm'
|
| 21 |
+
return [''.join(alphabet[int(i)] for i in row) for row in tokens]
|
| 22 |
+
|
| 23 |
+
|
| 24 |
+
def make_data(count=512, length=8, seed=7):
|
| 25 |
+
"""Four synthetic motif families, with independent 8% base mutations."""
|
| 26 |
+
rng = np.random.default_rng(seed)
|
| 27 |
+
motifs = ['ACGT', 'CGTA', 'TATA', 'GCGC']
|
| 28 |
+
strings, labels = [], []
|
| 29 |
+
for _ in range(count):
|
| 30 |
+
motif = motifs[int(rng.integers(4))]
|
| 31 |
+
seq = list((motif * ((length + 3) // 4))[:length])
|
| 32 |
+
for j in range(length):
|
| 33 |
+
if rng.random() < .08:
|
| 34 |
+
seq[j] = 'ACGT'[int(rng.integers(4))]
|
| 35 |
+
strings.append(''.join(seq))
|
| 36 |
+
labels.append(int(sum(x in 'GC' for x in seq) / length >= .6))
|
| 37 |
+
return encode(strings), torch.tensor(labels)
|
| 38 |
+
|
| 39 |
+
|
| 40 |
+
def load_data(path, length):
|
| 41 |
+
if path is None:
|
| 42 |
+
return make_data(length=length)
|
| 43 |
+
with open(path) as f:
|
| 44 |
+
rows = list(csv.DictReader(f, delimiter='\t'))
|
| 45 |
+
strings = [r['sequence'].strip().upper() for r in rows]
|
| 46 |
+
if not strings or any(len(s) != length or set(s) - set('ACGT') for s in strings):
|
| 47 |
+
raise ValueError('All DNA sequences must contain only A/C/G/T and have --length bases.')
|
| 48 |
+
labels = [int(r.get('label', sum(c in 'GC' for c in s) / length >= .6))
|
| 49 |
+
for r, s in zip(rows, strings)]
|
| 50 |
+
if set(labels) - {0, 1}:
|
| 51 |
+
raise ValueError('Labels must be 0 or 1.')
|
| 52 |
+
return encode(strings), torch.tensor(labels)
|
| 53 |
+
|
| 54 |
+
|
| 55 |
+
class ConditionalDNA(DNA):
|
| 56 |
+
"""The slide network plus a label embedding; label 2 means unconditional."""
|
| 57 |
+
def __init__(self, width=32, max_len=64):
|
| 58 |
+
super().__init__(width, max_len)
|
| 59 |
+
self.condition = nn.Embedding(3, width)
|
| 60 |
+
|
| 61 |
+
def forward(self, z, t=None, label=None):
|
| 62 |
+
h = self.token(z) if z.ndim == 2 else self.soft(z)
|
| 63 |
+
pos = torch.arange(z.shape[1], device=z.device)
|
| 64 |
+
h = h + self.position(pos)[None]
|
| 65 |
+
if t is not None:
|
| 66 |
+
h = h + self.time(t[:, None])[:, None]
|
| 67 |
+
if label is None:
|
| 68 |
+
label = torch.full((len(z),), 2, dtype=torch.long, device=z.device)
|
| 69 |
+
h = h + self.condition(label)[:, None]
|
| 70 |
+
return self.output(self.context(h))
|
| 71 |
+
|
| 72 |
+
|
| 73 |
+
def objectives(tokens):
|
| 74 |
+
"""Both toy objectives are maximized: GC fraction and ATAT agreement."""
|
| 75 |
+
onehot = torch.nn.functional.one_hot(tokens.long(), 4).float()
|
| 76 |
+
return soft_objectives(onehot)
|
| 77 |
+
|
| 78 |
+
|
| 79 |
+
def soft_objectives(z):
|
| 80 |
+
gc = (z[..., 1] + z[..., 2]).mean(-1)
|
| 81 |
+
motif = torch.tensor([0, 3], device=z.device).repeat((z.shape[1] + 1) // 2)[:z.shape[1]]
|
| 82 |
+
match = z.gather(-1, motif[None, :, None].expand(z.shape[0], -1, 1)).squeeze(-1).mean(-1)
|
| 83 |
+
return torch.stack((gc, match), -1)
|
| 84 |
+
|
| 85 |
+
|
| 86 |
+
def metrics(tokens):
|
| 87 |
+
strings = decode(tokens)
|
| 88 |
+
score = objectives(tokens).mean(0)
|
| 89 |
+
return dict(valid_dna=all(set(s) <= set('ACGT') for s in strings),
|
| 90 |
+
unique_fraction=len(set(strings)) / len(strings),
|
| 91 |
+
mean_gc=float(score[0]), mean_atat_match=float(score[1]))
|
| 92 |
+
|
| 93 |
+
|
| 94 |
+
def save_run(out, config, losses, samples, extra=None):
|
| 95 |
+
out = Path(out)
|
| 96 |
+
out.mkdir(parents=True, exist_ok=True)
|
| 97 |
+
(out / ('sample_config.json' if config.get('mode') == 'sample' else 'config.json')).write_text(json.dumps(config, indent=2))
|
| 98 |
+
if losses or not (out / 'losses.json').exists():
|
| 99 |
+
(out / 'losses.json').write_text(json.dumps(losses, indent=2))
|
| 100 |
+
(out / 'samples.txt').write_text('\n'.join(decode(samples)) + '\n')
|
| 101 |
+
report = {'metrics': metrics(samples), **(extra or {})}
|
| 102 |
+
(out / 'report.json').write_text(json.dumps(report, indent=2))
|
| 103 |
+
print(json.dumps({'output': str(out), **report}, indent=2))
|
| 104 |
+
return report
|
| 105 |
+
|
| 106 |
+
|
| 107 |
+
def optimize(model, loss_fn, data, steps, batch_size=32, lr=.002):
|
| 108 |
+
optimizer = torch.optim.AdamW(model.parameters(), lr=lr)
|
| 109 |
+
losses = []
|
| 110 |
+
model.train()
|
| 111 |
+
for step in range(steps):
|
| 112 |
+
idx = torch.randint(len(data), (batch_size,))
|
| 113 |
+
optimizer.zero_grad(set_to_none=True)
|
| 114 |
+
loss = loss_fn(model, data[idx], idx)
|
| 115 |
+
if not torch.isfinite(loss):
|
| 116 |
+
raise FloatingPointError(f'Nonfinite loss at step {step}')
|
| 117 |
+
loss.backward()
|
| 118 |
+
torch.nn.utils.clip_grad_norm_(model.parameters(), 1.)
|
| 119 |
+
optimizer.step()
|
| 120 |
+
losses.append(float(loss.detach()))
|
| 121 |
+
model.eval()
|
| 122 |
+
return losses
|
| 123 |
+
|
| 124 |
+
|
| 125 |
+
def project_simplex(z, eps=1e-6):
|
| 126 |
+
"""Euclidean simplex projection, followed by a small interior floor."""
|
| 127 |
+
sorted_z = z.sort(-1, descending=True).values
|
| 128 |
+
cssv = sorted_z.cumsum(-1) - 1.
|
| 129 |
+
k = torch.arange(1, z.shape[-1] + 1, dtype=z.dtype, device=z.device)
|
| 130 |
+
active = sorted_z - cssv / k > 0
|
| 131 |
+
rho = active.sum(-1, keepdim=True).clamp_min(1)
|
| 132 |
+
theta = cssv.gather(-1, rho - 1) / rho
|
| 133 |
+
result = (z - theta).clamp_min(eps)
|
| 134 |
+
return result / result.sum(-1, keepdim=True)
|
lecture_4/data/README.md
ADDED
|
@@ -0,0 +1 @@
|
|
|
|
|
|
|
| 1 |
+
Synthetic teaching DNA only. Each eight-base sequence begins as two copies of ACGT, CGTA, TATA, or GCGC, then each base is replaced independently with probability 0.08. Seed 7. Label 1 means GC fraction >= 0.6. These labels are not measured biological activity. The CLI also generates length-adjusted data when --data is omitted.
|
lecture_4/data/dna_train.tsv
ADDED
|
@@ -0,0 +1,257 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
sequence label
|
| 2 |
+
TAAACTTA 0
|
| 3 |
+
ACGCTCCC 1
|
| 4 |
+
CGTACGTA 0
|
| 5 |
+
ACGTACGT 0
|
| 6 |
+
CGTACGTA 0
|
| 7 |
+
ACGTAAGT 0
|
| 8 |
+
TATAAATA 0
|
| 9 |
+
ACGTACGT 0
|
| 10 |
+
GCGCGCGC 1
|
| 11 |
+
GCGCGCGC 1
|
| 12 |
+
TATATACA 0
|
| 13 |
+
ACGTACGT 0
|
| 14 |
+
ACGTACGT 0
|
| 15 |
+
ACGCGCGC 1
|
| 16 |
+
TCTATATA 0
|
| 17 |
+
GCGCGCGC 1
|
| 18 |
+
CGTACGTA 0
|
| 19 |
+
ACGTACGT 0
|
| 20 |
+
CGTACGTA 0
|
| 21 |
+
GAGCGCGC 1
|
| 22 |
+
TATATATA 0
|
| 23 |
+
ACGTACGT 0
|
| 24 |
+
GCGCGCGC 1
|
| 25 |
+
ACGTACGT 0
|
| 26 |
+
GCGCACGC 1
|
| 27 |
+
CGTACGTA 0
|
| 28 |
+
GGTACGTA 0
|
| 29 |
+
CTTACGTA 0
|
| 30 |
+
ACGTACGT 0
|
| 31 |
+
AATATATA 0
|
| 32 |
+
CGTACGTA 0
|
| 33 |
+
TATATATA 0
|
| 34 |
+
GCGCGCGC 1
|
| 35 |
+
GCGCGCGC 1
|
| 36 |
+
GCGCTCGC 1
|
| 37 |
+
TCTATCTA 0
|
| 38 |
+
GCGCGCGC 1
|
| 39 |
+
CAAATATA 0
|
| 40 |
+
TCGCGCGC 1
|
| 41 |
+
ACGTACGT 0
|
| 42 |
+
AAGTACGT 0
|
| 43 |
+
GCGCGCGC 1
|
| 44 |
+
CGCGCTTA 1
|
| 45 |
+
TACATATA 0
|
| 46 |
+
TATATATA 0
|
| 47 |
+
CGTACGTA 0
|
| 48 |
+
ACGTACGT 0
|
| 49 |
+
GCGTGCGC 1
|
| 50 |
+
ACGTACCT 0
|
| 51 |
+
ACATACAT 0
|
| 52 |
+
ATGGACGT 0
|
| 53 |
+
GCGCGCGC 1
|
| 54 |
+
TATAAATA 0
|
| 55 |
+
CGTACGTA 0
|
| 56 |
+
ATTTACGA 0
|
| 57 |
+
CGTACGTA 0
|
| 58 |
+
GCACGCGC 1
|
| 59 |
+
GCGCGCGC 1
|
| 60 |
+
TTTATATA 0
|
| 61 |
+
ACGTACGT 0
|
| 62 |
+
GCGCGCGT 1
|
| 63 |
+
GCGCGCGC 1
|
| 64 |
+
TATATATA 0
|
| 65 |
+
CTTACGTA 0
|
| 66 |
+
CGTACGTA 0
|
| 67 |
+
GCGCGCGC 1
|
| 68 |
+
CGTACGTA 0
|
| 69 |
+
TGTTCGTA 0
|
| 70 |
+
ACGTACCT 0
|
| 71 |
+
ACGTACGC 1
|
| 72 |
+
GCGCGCGC 1
|
| 73 |
+
CTTATGTA 0
|
| 74 |
+
CGTGCGTG 1
|
| 75 |
+
ACGTACGT 0
|
| 76 |
+
GCGCCCAC 1
|
| 77 |
+
CGTACGTC 1
|
| 78 |
+
TATATATA 0
|
| 79 |
+
GGGCGCGC 1
|
| 80 |
+
TATATATA 0
|
| 81 |
+
TCGCGCGC 1
|
| 82 |
+
ACGTACGT 0
|
| 83 |
+
TATATCTA 0
|
| 84 |
+
GCGCGCGC 1
|
| 85 |
+
TATATATC 0
|
| 86 |
+
GCACGCGC 1
|
| 87 |
+
ACGTACGT 0
|
| 88 |
+
CGTACGTA 0
|
| 89 |
+
TATATATA 0
|
| 90 |
+
CGTACGTA 0
|
| 91 |
+
GGGCGCGC 1
|
| 92 |
+
TAGATATA 0
|
| 93 |
+
CGTACATA 0
|
| 94 |
+
ACGTACGT 0
|
| 95 |
+
TAGAGATA 0
|
| 96 |
+
ACGTACGT 0
|
| 97 |
+
ACGTACGT 0
|
| 98 |
+
TATATATA 0
|
| 99 |
+
ACTTACGT 0
|
| 100 |
+
CGTACCTA 0
|
| 101 |
+
CGTACGTA 0
|
| 102 |
+
GCGCGCGC 1
|
| 103 |
+
GCTCGCGC 1
|
| 104 |
+
GGGCGCCC 1
|
| 105 |
+
TATATATA 0
|
| 106 |
+
ACGTACGT 0
|
| 107 |
+
ACGTACGT 0
|
| 108 |
+
GCGCCCGC 1
|
| 109 |
+
TATATATA 0
|
| 110 |
+
ACGTACGT 0
|
| 111 |
+
TATAGATA 0
|
| 112 |
+
TATATATT 0
|
| 113 |
+
CAGACGTA 0
|
| 114 |
+
CGTACGTA 0
|
| 115 |
+
CGTAGGTA 0
|
| 116 |
+
GCACGACC 1
|
| 117 |
+
CGTACGAA 0
|
| 118 |
+
GTGCGCGC 1
|
| 119 |
+
CGTACGTA 0
|
| 120 |
+
TATATATA 0
|
| 121 |
+
CGTACGTA 0
|
| 122 |
+
GCACGCGC 1
|
| 123 |
+
ACATACGA 0
|
| 124 |
+
CGTAGGTA 0
|
| 125 |
+
ACGTACAT 0
|
| 126 |
+
TATATATA 0
|
| 127 |
+
GATATATA 0
|
| 128 |
+
ACGTACGT 0
|
| 129 |
+
TCTAGATA 0
|
| 130 |
+
CGTACGTA 0
|
| 131 |
+
CGTACGTA 0
|
| 132 |
+
TATATATA 0
|
| 133 |
+
GCGCGCGC 1
|
| 134 |
+
ACGTACGT 0
|
| 135 |
+
TATATATA 0
|
| 136 |
+
CGTACGTA 0
|
| 137 |
+
ACGTACGT 0
|
| 138 |
+
ACGTTCGT 0
|
| 139 |
+
TATATATA 0
|
| 140 |
+
GCGCGCGC 1
|
| 141 |
+
ACGTACGT 0
|
| 142 |
+
TATATATA 0
|
| 143 |
+
ACGCACGT 1
|
| 144 |
+
ACGTGCGT 1
|
| 145 |
+
CGTACGTA 0
|
| 146 |
+
GTGCGCGC 1
|
| 147 |
+
GCGCGCGC 1
|
| 148 |
+
CGTACGTA 0
|
| 149 |
+
GCGCGCGC 1
|
| 150 |
+
TATATATA 0
|
| 151 |
+
GCGCGCGC 1
|
| 152 |
+
CGTACGTA 0
|
| 153 |
+
CGTACGTA 0
|
| 154 |
+
ACGTACGT 0
|
| 155 |
+
GCGCGCGC 1
|
| 156 |
+
CGTACGTA 0
|
| 157 |
+
GCGCGGGA 1
|
| 158 |
+
TATATATA 0
|
| 159 |
+
ACGTACGT 0
|
| 160 |
+
ACGTACGT 0
|
| 161 |
+
GCGCGCGT 1
|
| 162 |
+
ACTTACGT 0
|
| 163 |
+
CGTACGTA 0
|
| 164 |
+
ACGTACGT 0
|
| 165 |
+
TAAATAAA 0
|
| 166 |
+
GCGTGCGC 1
|
| 167 |
+
GCGCGCAC 1
|
| 168 |
+
ACGTACGT 0
|
| 169 |
+
TATATATA 0
|
| 170 |
+
TAAATATA 0
|
| 171 |
+
CGTACGTA 0
|
| 172 |
+
ACGTACGT 0
|
| 173 |
+
TGTACATA 0
|
| 174 |
+
CGTACGGA 1
|
| 175 |
+
CGTCCGTA 1
|
| 176 |
+
GCGCGCGC 1
|
| 177 |
+
ACGTACGT 0
|
| 178 |
+
TATATATA 0
|
| 179 |
+
CGTACGTA 0
|
| 180 |
+
TGTACGTA 0
|
| 181 |
+
ACGTACGT 0
|
| 182 |
+
ACGTACGT 0
|
| 183 |
+
GCGTATGT 0
|
| 184 |
+
ACGTACGT 0
|
| 185 |
+
CGGAAGTC 1
|
| 186 |
+
TATATATA 0
|
| 187 |
+
TATATATA 0
|
| 188 |
+
CGTGCGTA 1
|
| 189 |
+
GCGCGCGC 1
|
| 190 |
+
CGTACGTA 0
|
| 191 |
+
ACGTACGT 0
|
| 192 |
+
ACGTACGT 0
|
| 193 |
+
GCGCGCGC 1
|
| 194 |
+
ACGTACGT 0
|
| 195 |
+
ACGTACGT 0
|
| 196 |
+
ACGTACGT 0
|
| 197 |
+
CGTACGTA 0
|
| 198 |
+
GCGCGCGC 1
|
| 199 |
+
ACGTAGGT 0
|
| 200 |
+
TATATATA 0
|
| 201 |
+
ACATACGT 0
|
| 202 |
+
CCGTACGT 1
|
| 203 |
+
AAGTACGT 0
|
| 204 |
+
CGTACGTA 0
|
| 205 |
+
GCACGCGC 1
|
| 206 |
+
TACATATA 0
|
| 207 |
+
TATATATA 0
|
| 208 |
+
GCGCGCGC 1
|
| 209 |
+
ACGTACGT 0
|
| 210 |
+
GCGTACGT 1
|
| 211 |
+
TATATATA 0
|
| 212 |
+
TATATATA 0
|
| 213 |
+
GCGCGCGC 1
|
| 214 |
+
CGTGCTTA 0
|
| 215 |
+
ACGTACGT 0
|
| 216 |
+
TATATATA 0
|
| 217 |
+
CGTACGTA 0
|
| 218 |
+
TATATATA 0
|
| 219 |
+
GTGCGCGC 1
|
| 220 |
+
CGTACGTA 0
|
| 221 |
+
GCGTACGT 1
|
| 222 |
+
GCGCGCGC 1
|
| 223 |
+
TATATATA 0
|
| 224 |
+
TATATGTA 0
|
| 225 |
+
TATATATA 0
|
| 226 |
+
CATACTAA 0
|
| 227 |
+
ACGAACGT 0
|
| 228 |
+
CGTACGTA 0
|
| 229 |
+
GCGAGCGC 1
|
| 230 |
+
TATATATA 0
|
| 231 |
+
ACGAACCT 0
|
| 232 |
+
GCGCGCGC 1
|
| 233 |
+
CGTCCGTG 1
|
| 234 |
+
GCGCGAGC 1
|
| 235 |
+
CGTACCTA 0
|
| 236 |
+
CGTACGTA 0
|
| 237 |
+
GCGCGCGC 1
|
| 238 |
+
TATATATA 0
|
| 239 |
+
ACGTACGT 0
|
| 240 |
+
ACGTAGGT 0
|
| 241 |
+
GGTACGTA 0
|
| 242 |
+
CGTACGTA 0
|
| 243 |
+
GCGCGCGA 1
|
| 244 |
+
CGTACGTA 0
|
| 245 |
+
CGTACGGA 1
|
| 246 |
+
ACGTACGT 0
|
| 247 |
+
GCGCGCGC 1
|
| 248 |
+
GCGCGCGC 1
|
| 249 |
+
TATATATA 0
|
| 250 |
+
ACGTACGT 0
|
| 251 |
+
GCGCGCGC 1
|
| 252 |
+
ACGTACGT 0
|
| 253 |
+
ATGGACGT 0
|
| 254 |
+
CGTAAGTA 0
|
| 255 |
+
TATATATA 0
|
| 256 |
+
TATATATA 0
|
| 257 |
+
TATATATA 0
|
lecture_4/diffusion.py
ADDED
|
@@ -0,0 +1,223 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""Complete samplers and guidance extensions around the lecture code modules."""
|
| 2 |
+
import math
|
| 3 |
+
import numpy as np
|
| 4 |
+
import torch
|
| 5 |
+
from torch import nn
|
| 6 |
+
import torch.nn.functional as F
|
| 7 |
+
from lecture_core import (K, MASK, draw, token_ce, udlm_rates,
|
| 8 |
+
rate_step, geometric_cfg, pareto_filter)
|
| 9 |
+
from common import objectives
|
| 10 |
+
|
| 11 |
+
|
| 12 |
+
@torch.no_grad()
|
| 13 |
+
def uniform_sample(model, batch, length, steps=100, epsilon=.02):
|
| 14 |
+
"""Adaptive reverse Euler on the same truncated interval as udlm_loss.
|
| 15 |
+
|
| 16 |
+
Uniform initialization at t=1-epsilon and stopping at t=epsilon
|
| 17 |
+
are endpoint approximations. The returned sequence retains residual
|
| 18 |
+
corruption; UDLM output probabilities are not treated as a clean posterior.
|
| 19 |
+
"""
|
| 20 |
+
z = torch.randint(K, (batch, length))
|
| 21 |
+
time = 1. - epsilon
|
| 22 |
+
count = 0
|
| 23 |
+
while time > epsilon + 1e-8:
|
| 24 |
+
t = torch.full((batch,), time)
|
| 25 |
+
rate = udlm_rates(model(z, t).softmax(-1), z, t)
|
| 26 |
+
max_exit = float(rate.sum(-1).max())
|
| 27 |
+
h = min(1. / steps, time - epsilon, .5 / max(max_exit, 1e-8))
|
| 28 |
+
z = rate_step(z, rate, h)
|
| 29 |
+
time -= h
|
| 30 |
+
count += 1
|
| 31 |
+
if count > 100000:
|
| 32 |
+
raise RuntimeError('Adaptive sampler failed to advance.')
|
| 33 |
+
return z
|
| 34 |
+
|
| 35 |
+
|
| 36 |
+
@torch.no_grad()
|
| 37 |
+
def block_sample(model, batch, length, size=2, steps=20):
|
| 38 |
+
"""Generate one block completely before appending the next one."""
|
| 39 |
+
prefix = torch.empty((batch, 0), dtype=torch.long)
|
| 40 |
+
for start in range(0, length, size):
|
| 41 |
+
width = min(size, length - start)
|
| 42 |
+
block = torch.full((batch, width), MASK)
|
| 43 |
+
grid = torch.linspace(1., 0., steps + 1)
|
| 44 |
+
for t, s in zip(grid[:-1], grid[1:]):
|
| 45 |
+
context = torch.cat((prefix, block), 1)
|
| 46 |
+
p = model(context)[:, start:].softmax(-1)
|
| 47 |
+
candidate = draw(p)
|
| 48 |
+
reveal = (block == MASK) & (torch.rand(block.shape) < (t-s)/t)
|
| 49 |
+
block = torch.where(reveal, candidate, block)
|
| 50 |
+
prefix = torch.cat((prefix, block), 1)
|
| 51 |
+
return prefix
|
| 52 |
+
|
| 53 |
+
|
| 54 |
+
def conditional_loss(model, clean, labels, drop=.15):
|
| 55 |
+
t = torch.rand(len(clean)).clamp_min(1e-4)
|
| 56 |
+
mask = torch.rand(clean.shape) < t[:, None]
|
| 57 |
+
context = clean.masked_fill(mask, MASK)
|
| 58 |
+
condition = labels.clone()
|
| 59 |
+
condition[torch.rand(len(clean)) < drop] = 2
|
| 60 |
+
ce = token_ce(model(context, label=condition), clean)
|
| 61 |
+
return (ce * mask / t[:, None]).sum(1).mean()
|
| 62 |
+
|
| 63 |
+
|
| 64 |
+
@torch.no_grad()
|
| 65 |
+
def cfg_sample(model, batch, length, strength=2., label=1, steps=20):
|
| 66 |
+
z = torch.full((batch, length), MASK)
|
| 67 |
+
grid = torch.linspace(1., 0., steps + 1)
|
| 68 |
+
labels = torch.full((batch,), label, dtype=torch.long)
|
| 69 |
+
for t, s in zip(grid[:-1], grid[1:]):
|
| 70 |
+
uncond = model(z).softmax(-1)
|
| 71 |
+
cond = model(z, label=labels).softmax(-1)
|
| 72 |
+
probability = geometric_cfg(uncond, cond, strength)
|
| 73 |
+
candidate = draw(probability)
|
| 74 |
+
reveal = (z == MASK) & (torch.rand(z.shape) < (t-s)/t)
|
| 75 |
+
z = torch.where(reveal, candidate, z)
|
| 76 |
+
return z
|
| 77 |
+
|
| 78 |
+
|
| 79 |
+
class NoisyClassifier(nn.Module):
|
| 80 |
+
"""Predict high-GC class from a masked sequence and its noise level."""
|
| 81 |
+
def __init__(self, length):
|
| 82 |
+
super().__init__()
|
| 83 |
+
self.net = nn.Sequential(nn.Linear(length * 5 + 1, 64),
|
| 84 |
+
nn.SiLU(), nn.Linear(64, 1))
|
| 85 |
+
|
| 86 |
+
def forward(self, z, t):
|
| 87 |
+
features = F.one_hot(z, 5).float() if z.ndim == 2 else z
|
| 88 |
+
return self.net(torch.cat((features.flatten(1), t[:, None]), 1)).squeeze(-1)
|
| 89 |
+
|
| 90 |
+
|
| 91 |
+
def fit_classifier(model, data, labels, steps=200, batch=32):
|
| 92 |
+
opt = torch.optim.Adam(model.parameters(), lr=.003)
|
| 93 |
+
losses = []
|
| 94 |
+
for _ in range(steps):
|
| 95 |
+
idx = torch.randint(len(data), (batch,))
|
| 96 |
+
t = torch.rand(batch)
|
| 97 |
+
noisy = data[idx].masked_fill(torch.rand(batch, data.shape[1]) < t[:, None], MASK)
|
| 98 |
+
loss = F.binary_cross_entropy_with_logits(model(noisy, t), labels[idx].float())
|
| 99 |
+
opt.zero_grad(set_to_none=True)
|
| 100 |
+
loss.backward()
|
| 101 |
+
opt.step()
|
| 102 |
+
losses.append(float(loss.detach()))
|
| 103 |
+
model.eval()
|
| 104 |
+
return losses
|
| 105 |
+
|
| 106 |
+
|
| 107 |
+
def log_success(classifier, z, t, label):
|
| 108 |
+
logit = classifier(z, t)
|
| 109 |
+
return F.logsigmoid(logit if label == 1 else -logit)
|
| 110 |
+
|
| 111 |
+
|
| 112 |
+
def guidance_changes(classifier, z, time, label=1, gradient=False):
|
| 113 |
+
"""Log h(candidate)-log h(current), exact or first-order one-hot."""
|
| 114 |
+
batch, length = z.shape
|
| 115 |
+
t = torch.full((batch,), float(time))
|
| 116 |
+
if gradient:
|
| 117 |
+
with torch.enable_grad():
|
| 118 |
+
soft = F.one_hot(z, 5).float().requires_grad_(True)
|
| 119 |
+
value = log_success(classifier, soft, t, label)
|
| 120 |
+
grad = torch.autograd.grad(value.sum(), soft)[0]
|
| 121 |
+
current = grad.gather(-1, z[..., None])
|
| 122 |
+
return grad[..., :K] - current
|
| 123 |
+
with torch.no_grad():
|
| 124 |
+
current = log_success(classifier, z, t, label)
|
| 125 |
+
delta = torch.empty((batch, length, K))
|
| 126 |
+
for i in range(length):
|
| 127 |
+
for a in range(K):
|
| 128 |
+
edited = z.clone()
|
| 129 |
+
edited[:, i] = a
|
| 130 |
+
delta[:, i, a] = log_success(classifier, edited, t, label) - current
|
| 131 |
+
return delta
|
| 132 |
+
|
| 133 |
+
|
| 134 |
+
@torch.no_grad()
|
| 135 |
+
def classifier_sample(model, classifier, batch, length, strength=1.,
|
| 136 |
+
label=1, steps=40, gradient=False, epsilon=.005):
|
| 137 |
+
"""Rate guidance with adaptive Euler and an explicit endpoint closure."""
|
| 138 |
+
z = torch.full((batch, length), MASK)
|
| 139 |
+
time = 1.
|
| 140 |
+
iterations = 0
|
| 141 |
+
while time > epsilon + 1e-8:
|
| 142 |
+
p = model(z).softmax(-1)
|
| 143 |
+
delta = guidance_changes(classifier, z, time, label, gradient)
|
| 144 |
+
rates = p / time * (strength * delta).clamp(-20, 20).exp()
|
| 145 |
+
rates *= (z == MASK)[..., None]
|
| 146 |
+
exit_rate = rates.sum(-1)
|
| 147 |
+
h = min(1. / steps, time-epsilon, .5 / max(float(exit_rate.max()), 1e-8))
|
| 148 |
+
change = torch.rand(z.shape) < h * exit_rate
|
| 149 |
+
candidate = draw(rates / exit_rate.clamp_min(1e-12)[..., None] + 1e-12)
|
| 150 |
+
z = torch.where(change & (z == MASK), candidate, z)
|
| 151 |
+
time -= h
|
| 152 |
+
iterations += 1
|
| 153 |
+
if iterations > 100000:
|
| 154 |
+
raise RuntimeError('Guided rate integration failed to advance.')
|
| 155 |
+
z = torch.where(z == MASK, draw(model(z).softmax(-1)), z)
|
| 156 |
+
return z
|
| 157 |
+
|
| 158 |
+
|
| 159 |
+
class SearchNode:
|
| 160 |
+
def __init__(self, tokens, parent=None, prior=1.):
|
| 161 |
+
self.tokens, self.parent, self.prior = tokens, parent, prior
|
| 162 |
+
self.children = []
|
| 163 |
+
self.visits = 0
|
| 164 |
+
self.reward = np.zeros(2)
|
| 165 |
+
|
| 166 |
+
|
| 167 |
+
@torch.no_grad()
|
| 168 |
+
def peptune_search(model, length=8, iterations=100, branching=4):
|
| 169 |
+
"""DNA MCTS: selection, expansion, completion, Pareto rewards, backup.
|
| 170 |
+
|
| 171 |
+
All DNA strings are valid. Peptide chemistry, bond-dependent masks,
|
| 172 |
+
RoFormer training, and the PepTune invalid-SMILES penalty are not used.
|
| 173 |
+
"""
|
| 174 |
+
root = SearchNode(torch.full((length,), MASK))
|
| 175 |
+
archive = torch.empty((0, length), dtype=torch.long)
|
| 176 |
+
archive_scores = torch.empty((0, 2))
|
| 177 |
+
trace = []
|
| 178 |
+
for iteration in range(iterations):
|
| 179 |
+
node = root
|
| 180 |
+
while node.children:
|
| 181 |
+
def selection(child):
|
| 182 |
+
mean = child.reward.mean() / max(child.visits, 1)
|
| 183 |
+
bonus = 1.5 * child.prior * math.sqrt(node.visits + 1) / (child.visits + 1)
|
| 184 |
+
return mean + bonus
|
| 185 |
+
node = max(node.children, key=selection)
|
| 186 |
+
if (node.tokens == MASK).any():
|
| 187 |
+
p = model(node.tokens[None]).softmax(-1)[0]
|
| 188 |
+
positions = torch.where(node.tokens == MASK)[0]
|
| 189 |
+
seen = set()
|
| 190 |
+
for _ in range(branching):
|
| 191 |
+
pos = int(positions[torch.randint(len(positions), ())])
|
| 192 |
+
token = int(torch.multinomial(p[pos], 1))
|
| 193 |
+
candidate = node.tokens.clone()
|
| 194 |
+
candidate[pos] = token
|
| 195 |
+
key = tuple(candidate.tolist())
|
| 196 |
+
if key not in seen:
|
| 197 |
+
node.children.append(SearchNode(candidate, node, float(p[pos, token])))
|
| 198 |
+
seen.add(key)
|
| 199 |
+
node = node.children[0]
|
| 200 |
+
rollout = node.tokens.clone()
|
| 201 |
+
while (rollout == MASK).any():
|
| 202 |
+
prob = model(rollout[None]).softmax(-1)
|
| 203 |
+
position = int(torch.where(rollout == MASK)[0][0])
|
| 204 |
+
rollout[position] = draw(prob)[0, position]
|
| 205 |
+
score = objectives(rollout[None])[0]
|
| 206 |
+
reward = ((score >= archive_scores).float().mean(0).numpy()
|
| 207 |
+
if len(archive) else np.ones(2))
|
| 208 |
+
archive = torch.cat((archive, rollout[None]))
|
| 209 |
+
archive_scores = torch.cat((archive_scores, score[None]))
|
| 210 |
+
archive, archive_scores = pareto_filter(archive, archive_scores)
|
| 211 |
+
# Equal-score alternatives remain; remove exact repeated sequences only.
|
| 212 |
+
unique = []; seen = set()
|
| 213 |
+
for i, row in enumerate(archive.tolist()):
|
| 214 |
+
key = tuple(row)
|
| 215 |
+
if key not in seen:
|
| 216 |
+
unique.append(i); seen.add(key)
|
| 217 |
+
archive, archive_scores = archive[unique], archive_scores[unique]
|
| 218 |
+
while node is not None:
|
| 219 |
+
node.visits += 1
|
| 220 |
+
node.reward += reward
|
| 221 |
+
node = node.parent
|
| 222 |
+
trace.append(dict(iteration=iteration, archive_size=len(archive), score=score.tolist()))
|
| 223 |
+
return archive, archive_scores, trace
|
lecture_4/examples/block.py
ADDED
|
@@ -0,0 +1,7 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""Complete block training and generation example."""
|
| 2 |
+
import sys
|
| 3 |
+
from pathlib import Path
|
| 4 |
+
sys.path.insert(0, str(Path(__file__).resolve().parents[1]))
|
| 5 |
+
from run import main
|
| 6 |
+
if __name__ == '__main__':
|
| 7 |
+
main(['--method', 'block', '--out', 'outputs/block'] + sys.argv[1:])
|
lecture_4/examples/cfg.py
ADDED
|
@@ -0,0 +1,7 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""Complete cfg training and generation example."""
|
| 2 |
+
import sys
|
| 3 |
+
from pathlib import Path
|
| 4 |
+
sys.path.insert(0, str(Path(__file__).resolve().parents[1]))
|
| 5 |
+
from run import main
|
| 6 |
+
if __name__ == '__main__':
|
| 7 |
+
main(['--method', 'cfg', '--out', 'outputs/cfg'] + sys.argv[1:])
|
lecture_4/examples/classifier_exact.py
ADDED
|
@@ -0,0 +1,7 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""Complete classifier-exact training and generation example."""
|
| 2 |
+
import sys
|
| 3 |
+
from pathlib import Path
|
| 4 |
+
sys.path.insert(0, str(Path(__file__).resolve().parents[1]))
|
| 5 |
+
from run import main
|
| 6 |
+
if __name__ == '__main__':
|
| 7 |
+
main(['--method', 'classifier-exact', '--out', 'outputs/classifier-exact'] + sys.argv[1:])
|
lecture_4/examples/classifier_gradient.py
ADDED
|
@@ -0,0 +1,7 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""Complete classifier-gradient training and generation example."""
|
| 2 |
+
import sys
|
| 3 |
+
from pathlib import Path
|
| 4 |
+
sys.path.insert(0, str(Path(__file__).resolve().parents[1]))
|
| 5 |
+
from run import main
|
| 6 |
+
if __name__ == '__main__':
|
| 7 |
+
main(['--method', 'classifier-gradient', '--out', 'outputs/classifier-gradient'] + sys.argv[1:])
|
lecture_4/examples/mdlm.py
ADDED
|
@@ -0,0 +1,7 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""Complete mdlm training and generation example."""
|
| 2 |
+
import sys
|
| 3 |
+
from pathlib import Path
|
| 4 |
+
sys.path.insert(0, str(Path(__file__).resolve().parents[1]))
|
| 5 |
+
from run import main
|
| 6 |
+
if __name__ == '__main__':
|
| 7 |
+
main(['--method', 'mdlm', '--out', 'outputs/mdlm'] + sys.argv[1:])
|
lecture_4/examples/peptune.py
ADDED
|
@@ -0,0 +1,7 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""Complete peptune training and generation example."""
|
| 2 |
+
import sys
|
| 3 |
+
from pathlib import Path
|
| 4 |
+
sys.path.insert(0, str(Path(__file__).resolve().parents[1]))
|
| 5 |
+
from run import main
|
| 6 |
+
if __name__ == '__main__':
|
| 7 |
+
main(['--method', 'peptune', '--out', 'outputs/peptune'] + sys.argv[1:])
|
lecture_4/examples/udlm.py
ADDED
|
@@ -0,0 +1,7 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""Complete udlm training and generation example."""
|
| 2 |
+
import sys
|
| 3 |
+
from pathlib import Path
|
| 4 |
+
sys.path.insert(0, str(Path(__file__).resolve().parents[1]))
|
| 5 |
+
from run import main
|
| 6 |
+
if __name__ == '__main__':
|
| 7 |
+
main(['--method', 'udlm', '--out', 'outputs/udlm'] + sys.argv[1:])
|
lecture_4/lecture_core.py
ADDED
|
@@ -0,0 +1,136 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""Small teaching implementations; illustrative data, not paper reproductions."""
|
| 2 |
+
|
| 3 |
+
import math
|
| 4 |
+
|
| 5 |
+
import numpy as np
|
| 6 |
+
|
| 7 |
+
import torch
|
| 8 |
+
|
| 9 |
+
from torch import nn
|
| 10 |
+
|
| 11 |
+
import torch.nn.functional as F
|
| 12 |
+
|
| 13 |
+
from scipy.special import betainc, beta
|
| 14 |
+
|
| 15 |
+
from scipy.stats import rankdata
|
| 16 |
+
|
| 17 |
+
K, MASK = 4, 4
|
| 18 |
+
|
| 19 |
+
ALPHABET = 'ACGT'
|
| 20 |
+
|
| 21 |
+
def encode(strings):
|
| 22 |
+
return torch.tensor([[ALPHABET.index(c) for c in s]
|
| 23 |
+
for s in strings])
|
| 24 |
+
|
| 25 |
+
def draw(prob):
|
| 26 |
+
shape = prob.shape[:-1]
|
| 27 |
+
sample = torch.multinomial(prob.reshape(-1, K), 1)
|
| 28 |
+
return sample.reshape(shape)
|
| 29 |
+
|
| 30 |
+
class DNA(nn.Module):
|
| 31 |
+
def __init__(self, width=32, max_len=64):
|
| 32 |
+
super().__init__()
|
| 33 |
+
self.token = nn.Embedding(K + 1, width)
|
| 34 |
+
self.soft = nn.Linear(K, width)
|
| 35 |
+
self.position = nn.Embedding(max_len, width)
|
| 36 |
+
self.time = nn.Linear(1, width)
|
| 37 |
+
layer = nn.TransformerEncoderLayer(
|
| 38 |
+
width, 4, 2 * width, dropout=0., batch_first=True)
|
| 39 |
+
self.context = nn.TransformerEncoder(layer, 1)
|
| 40 |
+
self.output = nn.Linear(width, K)
|
| 41 |
+
|
| 42 |
+
def forward(self, z, t=None):
|
| 43 |
+
h = self.token(z) if z.ndim == 2 else self.soft(z)
|
| 44 |
+
pos = torch.arange(z.shape[1], device=z.device)
|
| 45 |
+
h = h + self.position(pos)[None]
|
| 46 |
+
if t is not None:
|
| 47 |
+
h = h + self.time(t[:, None])[:, None]
|
| 48 |
+
return self.output(self.context(h))
|
| 49 |
+
|
| 50 |
+
def token_ce(logits, target):
|
| 51 |
+
return F.cross_entropy(logits.transpose(1, 2),
|
| 52 |
+
target, reduction='none')
|
| 53 |
+
|
| 54 |
+
def mdlm_loss(model, clean):
|
| 55 |
+
batch, length = clean.shape
|
| 56 |
+
t = torch.rand(batch).clamp_min(1e-4)
|
| 57 |
+
masked = torch.rand(batch, length) < t[:, None]
|
| 58 |
+
noisy = clean.masked_fill(masked, MASK)
|
| 59 |
+
logits = model(noisy) # optimal predictor needs no t
|
| 60 |
+
ce = token_ce(logits, clean)
|
| 61 |
+
weighted = ce * masked / t[:, None]
|
| 62 |
+
return weighted.sum(1).mean()
|
| 63 |
+
|
| 64 |
+
def train_step(model, optimizer, clean, loss_fn):
|
| 65 |
+
model.train()
|
| 66 |
+
optimizer.zero_grad()
|
| 67 |
+
loss = loss_fn(model, clean)
|
| 68 |
+
loss.backward()
|
| 69 |
+
optimizer.step()
|
| 70 |
+
return loss.item()
|
| 71 |
+
|
| 72 |
+
@torch.no_grad()
|
| 73 |
+
def mdlm_sample(model, batch, length, steps=20):
|
| 74 |
+
model.eval()
|
| 75 |
+
z = torch.full((batch, length), MASK)
|
| 76 |
+
grid = torch.linspace(1., 0., steps + 1)
|
| 77 |
+
for t, s in zip(grid[:-1], grid[1:]):
|
| 78 |
+
prob = model(z).softmax(-1)
|
| 79 |
+
candidate = draw(prob)
|
| 80 |
+
reveal = torch.rand(z.shape) < (t - s) / t
|
| 81 |
+
update = (z == MASK) & reveal
|
| 82 |
+
z = torch.where(update, candidate, z)
|
| 83 |
+
return z
|
| 84 |
+
|
| 85 |
+
def rate_step(z, rates, h):
|
| 86 |
+
exit_rate = rates.sum(-1)
|
| 87 |
+
assert torch.all(h * exit_rate <= 1. + 1e-6)
|
| 88 |
+
prob = h * rates
|
| 89 |
+
prob.scatter_(-1, z[..., None],
|
| 90 |
+
(1. - h * exit_rate)[..., None])
|
| 91 |
+
return draw(prob.clamp_min(0.))
|
| 92 |
+
|
| 93 |
+
def udlm_rates(clean_prob, z, t):
|
| 94 |
+
alpha = 1. - t[:, None, None]
|
| 95 |
+
noisy_prob = alpha * clean_prob + (1. - alpha) / K
|
| 96 |
+
current = noisy_prob.gather(-1, z[..., None])
|
| 97 |
+
rates = noisy_prob / (K * alpha * current)
|
| 98 |
+
return rates.scatter(-1, z[..., None], 0.)
|
| 99 |
+
|
| 100 |
+
def udlm_loss(model, clean):
|
| 101 |
+
t = .02 + .96 * torch.rand(clean.shape[0])
|
| 102 |
+
random = torch.randint(K, clean.shape)
|
| 103 |
+
z = torch.where(torch.rand(clean.shape) < t[:, None],
|
| 104 |
+
random, clean)
|
| 105 |
+
exact = F.one_hot(clean, K).float()
|
| 106 |
+
pred = model(z, t).softmax(-1)
|
| 107 |
+
a, b = udlm_rates(exact, z, t), udlm_rates(pred, z, t)
|
| 108 |
+
term = a * (a.clamp_min(1e-12).log()
|
| 109 |
+
- b.clamp_min(1e-12).log()) + b - a
|
| 110 |
+
return .96 * term.sum((1, 2)).mean()
|
| 111 |
+
|
| 112 |
+
def block_loss(model, clean, start, size):
|
| 113 |
+
prefix, block = clean[:, :start], clean[:, start:start+size]
|
| 114 |
+
t = torch.rand(clean.shape[0]).clamp_min(1e-4)
|
| 115 |
+
mask = torch.rand(block.shape) < t[:, None]
|
| 116 |
+
ctx = torch.cat([prefix, block.masked_fill(mask, MASK)], 1)
|
| 117 |
+
logits = model(ctx)[:, start:]
|
| 118 |
+
loss = token_ce(logits, block) * mask / t[:, None]
|
| 119 |
+
return loss.sum(1).mean()
|
| 120 |
+
|
| 121 |
+
def geometric_cfg(uncond, cond, strength):
|
| 122 |
+
logits = (1. - strength) * uncond.clamp_min(1e-12).log()
|
| 123 |
+
logits += strength * cond.clamp_min(1e-12).log()
|
| 124 |
+
return logits.softmax(-1)
|
| 125 |
+
|
| 126 |
+
def guide_rates(base_rates, log_values, current_log_value,
|
| 127 |
+
strength=1.):
|
| 128 |
+
log_ratio = log_values - current_log_value[..., None]
|
| 129 |
+
return base_rates * (strength * log_ratio).exp()
|
| 130 |
+
|
| 131 |
+
def pareto_filter(sequences, scores):
|
| 132 |
+
# Maximize both objectives; keep equal-score alternatives.
|
| 133 |
+
ge = (scores[:, None] >= scores[None, :]).all(-1)
|
| 134 |
+
gt = (scores[:, None] > scores[None, :]).any(-1)
|
| 135 |
+
dominated = (ge & gt).any(0)
|
| 136 |
+
return sequences[~dominated], scores[~dominated]
|
lecture_4/numerical_examples.py
ADDED
|
@@ -0,0 +1,24 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""Print the four-letter training and sampling calculations from the lecture."""
|
| 2 |
+
import math
|
| 3 |
+
import torch
|
| 4 |
+
from lecture_core import udlm_rates, geometric_cfg
|
| 5 |
+
|
| 6 |
+
dtype = torch.float64
|
| 7 |
+
print('Alphabet order: A C G T; log losses are nats.')
|
| 8 |
+
q = .6*torch.eye(4,dtype=dtype)+.1*torch.ones(4,4,dtype=dtype)
|
| 9 |
+
print('Uniform one-step corruption matrix:\n',q)
|
| 10 |
+
print('Two-step marginal from clean A:',(q@q)[0].tolist())
|
| 11 |
+
reverse=q[0]*q[:,1];reverse/=reverse.sum()
|
| 12 |
+
print('Previous base given clean A and final C:',reverse.tolist())
|
| 13 |
+
loss=2*(-math.log(.6)-math.log(.7))
|
| 14 |
+
print('MDLM ACGT -> AmGm, t=.5, loss:',loss)
|
| 15 |
+
print('MDLM reverse .5 -> .25 at missing C:',[.05,.30,.10,.05,.50])
|
| 16 |
+
z=torch.tensor([[1]]);t=torch.tensor([.5],dtype=dtype)
|
| 17 |
+
a=udlm_rates(torch.tensor([[[1,0,0,0]]],dtype=dtype),z,t)
|
| 18 |
+
b=udlm_rates(torch.tensor([[[.6,.2,.1,.1]]],dtype=dtype),z,t)
|
| 19 |
+
rate_kl=(a*(a.clamp_min(1e-12).log()-b.clamp_min(1e-12).log())+b-a).sum()
|
| 20 |
+
print('UDLM target rates C -> A,C,G,T:',a.flatten().tolist())
|
| 21 |
+
print('UDLM learned rates:',b.flatten().tolist(),'rate loss:',float(rate_kl))
|
| 22 |
+
p=.1*b;p[...,1]=1-.1*b.sum(-1)
|
| 23 |
+
print('UDLM Euler step:',p.flatten().tolist())
|
| 24 |
+
print('CFG strength 2:',geometric_cfg(torch.full((4,),.25),torch.tensor([.1,.2,.6,.1]),2).tolist())
|
lecture_4/requirements.txt
ADDED
|
@@ -0,0 +1,3 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
torch>=2.2
|
| 2 |
+
numpy>=1.24
|
| 3 |
+
scipy>=1.10
|
lecture_4/run.py
ADDED
|
@@ -0,0 +1,127 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
#!/usr/bin/env python3
|
| 2 |
+
"""Train and sample each Lecture 4 method on small, explicit DNA examples."""
|
| 3 |
+
import argparse
|
| 4 |
+
from pathlib import Path
|
| 5 |
+
import torch
|
| 6 |
+
from lecture_core import DNA, mdlm_loss, mdlm_sample, udlm_loss, block_loss
|
| 7 |
+
from common import (seed_all, load_data, optimize, save_run, ConditionalDNA,
|
| 8 |
+
decode, metrics)
|
| 9 |
+
from diffusion import (uniform_sample, block_sample, conditional_loss,
|
| 10 |
+
cfg_sample, NoisyClassifier, fit_classifier,
|
| 11 |
+
classifier_sample, peptune_search)
|
| 12 |
+
|
| 13 |
+
METHODS = ['mdlm', 'udlm', 'block', 'cfg', 'classifier-free',
|
| 14 |
+
'classifier-gradient', 'classifier-exact', 'peptune']
|
| 15 |
+
|
| 16 |
+
|
| 17 |
+
def main(argv=None):
|
| 18 |
+
parser = argparse.ArgumentParser(description=__doc__)
|
| 19 |
+
parser.add_argument('--method', choices=METHODS, default='mdlm')
|
| 20 |
+
parser.add_argument('--mode', choices=['train-sample', 'train', 'sample'], default='train-sample')
|
| 21 |
+
parser.add_argument('--train-steps', type=int, default=300)
|
| 22 |
+
parser.add_argument('--sample-steps', type=int, default=40)
|
| 23 |
+
parser.add_argument('--classifier-steps', type=int, default=300)
|
| 24 |
+
parser.add_argument('--search-steps', type=int, default=100)
|
| 25 |
+
parser.add_argument('--batch-size', type=int, default=32)
|
| 26 |
+
parser.add_argument('--samples', type=int, default=32)
|
| 27 |
+
parser.add_argument('--length', type=int, default=8)
|
| 28 |
+
parser.add_argument('--width', type=int, default=32)
|
| 29 |
+
parser.add_argument('--block-size', type=int, default=2)
|
| 30 |
+
parser.add_argument('--strength', type=float, default=1.5)
|
| 31 |
+
parser.add_argument('--label', type=int, choices=[0, 1], default=1)
|
| 32 |
+
parser.add_argument('--seed', type=int, default=7)
|
| 33 |
+
parser.add_argument('--data', help='TSV with sequence and optional binary label columns')
|
| 34 |
+
parser.add_argument('--out', default='outputs/mdlm')
|
| 35 |
+
args = parser.parse_args(argv)
|
| 36 |
+
if min(args.length, args.sample_steps, args.batch_size, args.samples, args.block_size) < 1:
|
| 37 |
+
parser.error('Lengths, step counts, and batch counts must be positive.')
|
| 38 |
+
if args.width % 4 or args.length > 64:
|
| 39 |
+
parser.error('Width must be divisible by four; length must not exceed 64.')
|
| 40 |
+
if args.mode != 'sample' and args.train_steps < 1:
|
| 41 |
+
parser.error('Training requires at least one step.')
|
| 42 |
+
seed_all(args.seed)
|
| 43 |
+
out = Path(args.out); out.mkdir(parents=True, exist_ok=True)
|
| 44 |
+
method = 'cfg' if args.method == 'classifier-free' else args.method
|
| 45 |
+
model = ConditionalDNA(args.width) if method == 'cfg' else DNA(args.width)
|
| 46 |
+
classifier = None
|
| 47 |
+
losses = []
|
| 48 |
+
report = {'method': method, 'data': 'synthetic DNA; not biological validation'}
|
| 49 |
+
if args.mode == 'sample':
|
| 50 |
+
checkpoint = torch.load(out / 'checkpoint.pt', weights_only=True)
|
| 51 |
+
if checkpoint['method'] != method or checkpoint['length'] != args.length:
|
| 52 |
+
raise ValueError('Checkpoint method/length must match command arguments.')
|
| 53 |
+
model.load_state_dict(checkpoint['model'])
|
| 54 |
+
if 'classifier' in checkpoint:
|
| 55 |
+
classifier = NoisyClassifier(args.length)
|
| 56 |
+
classifier.load_state_dict(checkpoint['classifier'])
|
| 57 |
+
else:
|
| 58 |
+
data, labels = load_data(args.data, args.length)
|
| 59 |
+
split = max(1, int(.8 * len(data)))
|
| 60 |
+
train, train_labels = data[:split], labels[:split]
|
| 61 |
+
if method == 'udlm':
|
| 62 |
+
loss_fn = lambda m, x, idx: udlm_loss(m, x)
|
| 63 |
+
elif method == 'block':
|
| 64 |
+
starts = list(range(0, args.length, args.block_size))
|
| 65 |
+
def loss_fn(m, x, idx):
|
| 66 |
+
start = starts[int(torch.randint(len(starts), ()))]
|
| 67 |
+
return len(starts) * block_loss(m, x, start, args.block_size)
|
| 68 |
+
elif method == 'cfg':
|
| 69 |
+
loss_fn = lambda m, x, idx: conditional_loss(m, x, train_labels[idx])
|
| 70 |
+
else:
|
| 71 |
+
loss_fn = lambda m, x, idx: mdlm_loss(m, x)
|
| 72 |
+
losses = optimize(model, loss_fn, train, args.train_steps, args.batch_size)
|
| 73 |
+
checkpoint = {'model': model.state_dict(), 'method': method,
|
| 74 |
+
'length': args.length, 'width': args.width}
|
| 75 |
+
if method.startswith('classifier-'):
|
| 76 |
+
classifier = NoisyClassifier(args.length)
|
| 77 |
+
classifier_losses = fit_classifier(classifier, train, train_labels,
|
| 78 |
+
args.classifier_steps, args.batch_size)
|
| 79 |
+
checkpoint['classifier'] = classifier.state_dict()
|
| 80 |
+
report['classifier_final_loss'] = classifier_losses[-1]
|
| 81 |
+
torch.save(checkpoint, out / 'checkpoint.pt')
|
| 82 |
+
report['train_loss_first_20_mean'] = sum(losses[:20]) / len(losses[:20])
|
| 83 |
+
report['train_loss_last_20_mean'] = sum(losses[-20:]) / len(losses[-20:])
|
| 84 |
+
# An independent noisy validation estimate, not a perplexity claim.
|
| 85 |
+
val = data[split:]
|
| 86 |
+
if len(val):
|
| 87 |
+
with torch.no_grad():
|
| 88 |
+
if method == 'udlm':
|
| 89 |
+
validation = udlm_loss(model, val)
|
| 90 |
+
elif method == 'cfg':
|
| 91 |
+
validation = conditional_loss(model, val, labels[split:], drop=0.)
|
| 92 |
+
elif method == 'block':
|
| 93 |
+
validation = sum(block_loss(model, val, start, args.block_size)
|
| 94 |
+
for start in starts)
|
| 95 |
+
else:
|
| 96 |
+
validation = mdlm_loss(model, val)
|
| 97 |
+
report['validation_loss_one_mc_draw'] = float(validation)
|
| 98 |
+
if args.mode == 'train':
|
| 99 |
+
save_run(out, vars(args), losses, train[:args.samples],
|
| 100 |
+
{**report, 'sample_file_contains': 'training examples; generation not requested'})
|
| 101 |
+
return report
|
| 102 |
+
model.eval()
|
| 103 |
+
if method == 'udlm':
|
| 104 |
+
samples = uniform_sample(model, args.samples, args.length, args.sample_steps)
|
| 105 |
+
report['endpoint_approximation'] = 't in [0.02, 0.98]; stop at residual noise 0.02, without a posterior interpretation of the UDLM parameter vector'
|
| 106 |
+
elif method == 'block':
|
| 107 |
+
samples = block_sample(model, args.samples, args.length, args.block_size, args.sample_steps)
|
| 108 |
+
elif method == 'cfg':
|
| 109 |
+
samples = cfg_sample(model, args.samples, args.length, args.strength, args.label, args.sample_steps)
|
| 110 |
+
elif method.startswith('classifier-'):
|
| 111 |
+
classifier.eval()
|
| 112 |
+
samples = classifier_sample(model, classifier, args.samples, args.length,
|
| 113 |
+
args.strength, args.label, args.sample_steps,
|
| 114 |
+
gradient=method == 'classifier-gradient')
|
| 115 |
+
report['guidance'] = 'learned noisy classifier; final residual masks use denoiser closure'
|
| 116 |
+
elif method == 'peptune':
|
| 117 |
+
samples, scores, trace = peptune_search(model, args.length, args.search_steps)
|
| 118 |
+
report['archive_scores'] = scores.tolist()
|
| 119 |
+
report['search_trace'] = trace
|
| 120 |
+
report['scope'] = 'DNA MCTS mechanism; not peptide-model training or paper reproduction'
|
| 121 |
+
else:
|
| 122 |
+
samples = mdlm_sample(model, args.samples, args.length, args.sample_steps)
|
| 123 |
+
return save_run(out, vars(args), losses, samples, report)
|
| 124 |
+
|
| 125 |
+
|
| 126 |
+
if __name__ == '__main__':
|
| 127 |
+
main()
|
lecture_4/run_all.py
ADDED
|
@@ -0,0 +1,15 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
#!/usr/bin/env python3
|
| 2 |
+
"""Run every complete training example; pass --quick for a small CPU check."""
|
| 3 |
+
import argparse
|
| 4 |
+
from run import main
|
| 5 |
+
parser = argparse.ArgumentParser()
|
| 6 |
+
parser.add_argument('--quick', action='store_true')
|
| 7 |
+
args = parser.parse_args()
|
| 8 |
+
methods = ['mdlm', 'udlm', 'block', 'cfg', 'classifier-exact', 'classifier-gradient', 'peptune']
|
| 9 |
+
for method in methods:
|
| 10 |
+
command = ['--method', method, '--out', 'outputs/' + method]
|
| 11 |
+
if args.quick:
|
| 12 |
+
command += ['--train-steps', '20', '--samples', '4', '--length', '4',
|
| 13 |
+
'--batch-size', '8', '--sample-steps', '20']
|
| 14 |
+
command += ['--classifier-steps', '20', '--search-steps', '12']
|
| 15 |
+
main(command)
|
lecture_4/tests/test_mathematics.py
ADDED
|
@@ -0,0 +1,50 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
import unittest
|
| 2 |
+
import sys
|
| 3 |
+
from pathlib import Path
|
| 4 |
+
sys.path.insert(0, str(Path(__file__).resolve().parents[1]))
|
| 5 |
+
import torch
|
| 6 |
+
from lecture_core import udlm_rates, geometric_cfg
|
| 7 |
+
from diffusion import guidance_changes, NoisyClassifier
|
| 8 |
+
|
| 9 |
+
|
| 10 |
+
class Mathematics(unittest.TestCase):
|
| 11 |
+
def test_masked_kl_reduces_to_cross_entropy(self):
|
| 12 |
+
r = .5
|
| 13 |
+
d = torch.tensor([.1,.6,.2,.1], dtype=torch.float64)
|
| 14 |
+
q = torch.tensor([0,r,0,0,1-r], dtype=torch.float64)
|
| 15 |
+
p = torch.cat((r*d, torch.tensor([1-r])))
|
| 16 |
+
actual = (q[q>0] * (q[q>0]/p[q>0]).log()).sum()
|
| 17 |
+
self.assertAlmostEqual(float(actual), float(-r*d[1].log()), places=12)
|
| 18 |
+
|
| 19 |
+
def test_udlm_small_step_kl_converges_to_rate_loss(self):
|
| 20 |
+
a = torch.tensor([2.5,.5,.5], dtype=torch.float64)
|
| 21 |
+
b = torch.tensor([17/18,7/18,7/18], dtype=torch.float64)
|
| 22 |
+
target = (a*(a/b).log()+b-a).sum()
|
| 23 |
+
errors = []
|
| 24 |
+
for h in [1e-3,1e-4,1e-5]:
|
| 25 |
+
q = torch.cat((h*a,(1-h*a.sum()).reshape(1)))
|
| 26 |
+
p = torch.cat((h*b,(1-h*b.sum()).reshape(1)))
|
| 27 |
+
local = (q*(q/p).log()).sum()/h
|
| 28 |
+
errors.append(abs(float(local-target)))
|
| 29 |
+
self.assertLess(errors[-1], errors[0]/50)
|
| 30 |
+
self.assertAlmostEqual(float(target), .9071595148, places=8)
|
| 31 |
+
|
| 32 |
+
def test_cfg_dna_example(self):
|
| 33 |
+
p = geometric_cfg(torch.full((4,),.25),torch.tensor([.1,.2,.6,.1]),2.)
|
| 34 |
+
torch.testing.assert_close(p, torch.tensor([1,4,36,1])/42.)
|
| 35 |
+
|
| 36 |
+
def test_classifier_gradient_matches_directional_derivative(self):
|
| 37 |
+
torch.manual_seed(2)
|
| 38 |
+
model = NoisyClassifier(4).double()
|
| 39 |
+
z = torch.tensor([[0,1,2,3]])
|
| 40 |
+
soft = torch.nn.functional.one_hot(z,5).double().requires_grad_(True)
|
| 41 |
+
t = torch.tensor([.5], dtype=torch.float64)
|
| 42 |
+
f = lambda x: torch.nn.functional.logsigmoid(model(x,t)).sum()
|
| 43 |
+
grad = torch.autograd.grad(f(soft),soft)[0]
|
| 44 |
+
direction = torch.zeros_like(soft);direction[0,0,0]=-1;direction[0,0,1]=1
|
| 45 |
+
h=1e-5
|
| 46 |
+
numeric=(f(soft+h*direction)-f(soft-h*direction))/(2*h)
|
| 47 |
+
self.assertAlmostEqual(float(numeric.detach()),float((grad*direction).sum()),places=8)
|
| 48 |
+
|
| 49 |
+
|
| 50 |
+
if __name__ == '__main__': unittest.main()
|
lecture_4/verified_examples/README.md
ADDED
|
@@ -0,0 +1,15 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# Verified CPU examples
|
| 2 |
+
|
| 3 |
+
Actual seed-7 runs: 160 optimizer steps, length 4, batch size 16, 40 sampling steps, and 8 requested samples. ReDi/AReUReDi additionally use 160 teacher updates and 128 recoupled pairs. PepTune returns its archive, whose size can differ from the requested sample count.
|
| 4 |
+
|
| 5 |
+
| Method | Mean first 20 losses | Mean last 20 losses | Valid DNA |
|
| 6 |
+
| --- | ---: | ---: | --- |
|
| 7 |
+
| mdlm | 4.8388 | 2.6139 | True |
|
| 8 |
+
| udlm | 5.1313 | 2.2786 | True |
|
| 9 |
+
| block | 5.3755 | 3.2929 | True |
|
| 10 |
+
| cfg | 4.3721 | 2.3383 | True |
|
| 11 |
+
| classifier-exact | 4.8388 | 2.6139 | True |
|
| 12 |
+
| classifier-gradient | 4.8388 | 2.6139 | True |
|
| 13 |
+
| peptune | 4.8388 | 2.6139 | True |
|
| 14 |
+
|
| 15 |
+
These are execution and learning checks on toy data, not paper benchmark results. Different method losses have different meanings and scales. The sample counts are too small for statistical quality claims.
|
lecture_4/verified_examples/block/config.json
ADDED
|
@@ -0,0 +1,18 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
{
|
| 2 |
+
"method": "block",
|
| 3 |
+
"mode": "train-sample",
|
| 4 |
+
"train_steps": 160,
|
| 5 |
+
"sample_steps": 40,
|
| 6 |
+
"classifier_steps": 160,
|
| 7 |
+
"search_steps": 40,
|
| 8 |
+
"batch_size": 16,
|
| 9 |
+
"samples": 8,
|
| 10 |
+
"length": 4,
|
| 11 |
+
"width": 32,
|
| 12 |
+
"block_size": 2,
|
| 13 |
+
"strength": 1.5,
|
| 14 |
+
"label": 1,
|
| 15 |
+
"seed": 7,
|
| 16 |
+
"data": null,
|
| 17 |
+
"out": "outputs/block"
|
| 18 |
+
}
|
lecture_4/verified_examples/block/losses.json
ADDED
|
@@ -0,0 +1,162 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
[
|
| 2 |
+
4.331284523010254,
|
| 3 |
+
5.4537835121154785,
|
| 4 |
+
4.560967445373535,
|
| 5 |
+
6.563370704650879,
|
| 6 |
+
8.061755180358887,
|
| 7 |
+
3.8685860633850098,
|
| 8 |
+
6.454323768615723,
|
| 9 |
+
7.774837970733643,
|
| 10 |
+
2.178312301635742,
|
| 11 |
+
4.992790222167969,
|
| 12 |
+
12.195539474487305,
|
| 13 |
+
5.040197372436523,
|
| 14 |
+
3.1479640007019043,
|
| 15 |
+
4.1676025390625,
|
| 16 |
+
4.741274356842041,
|
| 17 |
+
6.513577461242676,
|
| 18 |
+
5.152376651763916,
|
| 19 |
+
6.8658881187438965,
|
| 20 |
+
3.170300006866455,
|
| 21 |
+
2.2749924659729004,
|
| 22 |
+
2.5343780517578125,
|
| 23 |
+
4.639561176300049,
|
| 24 |
+
3.381667375564575,
|
| 25 |
+
2.6100592613220215,
|
| 26 |
+
4.106627941131592,
|
| 27 |
+
5.132066249847412,
|
| 28 |
+
4.543891906738281,
|
| 29 |
+
4.116055011749268,
|
| 30 |
+
4.740029811859131,
|
| 31 |
+
11.817144393920898,
|
| 32 |
+
1.597190499305725,
|
| 33 |
+
2.864633321762085,
|
| 34 |
+
3.4319229125976562,
|
| 35 |
+
2.477750778198242,
|
| 36 |
+
2.259556293487549,
|
| 37 |
+
5.4611029624938965,
|
| 38 |
+
1.7751753330230713,
|
| 39 |
+
3.8391871452331543,
|
| 40 |
+
10.745177268981934,
|
| 41 |
+
2.813826084136963,
|
| 42 |
+
3.071927547454834,
|
| 43 |
+
4.419896602630615,
|
| 44 |
+
5.47519588470459,
|
| 45 |
+
15.978719711303711,
|
| 46 |
+
1.765317440032959,
|
| 47 |
+
3.7841482162475586,
|
| 48 |
+
3.9389586448669434,
|
| 49 |
+
4.711722373962402,
|
| 50 |
+
3.904447078704834,
|
| 51 |
+
4.800976753234863,
|
| 52 |
+
2.614999771118164,
|
| 53 |
+
3.3339426517486572,
|
| 54 |
+
7.427157402038574,
|
| 55 |
+
3.459439754486084,
|
| 56 |
+
5.940922260284424,
|
| 57 |
+
6.2956037521362305,
|
| 58 |
+
1.3079692125320435,
|
| 59 |
+
3.5128886699676514,
|
| 60 |
+
4.249375820159912,
|
| 61 |
+
2.7473535537719727,
|
| 62 |
+
2.1603171825408936,
|
| 63 |
+
0.7737497687339783,
|
| 64 |
+
2.2809128761291504,
|
| 65 |
+
4.380037307739258,
|
| 66 |
+
0.9533628225326538,
|
| 67 |
+
2.962036609649658,
|
| 68 |
+
2.988816261291504,
|
| 69 |
+
2.04714298248291,
|
| 70 |
+
7.599990367889404,
|
| 71 |
+
3.762449026107788,
|
| 72 |
+
3.4661433696746826,
|
| 73 |
+
2.0959715843200684,
|
| 74 |
+
1.7144439220428467,
|
| 75 |
+
3.5818352699279785,
|
| 76 |
+
3.0614638328552246,
|
| 77 |
+
0.7883206605911255,
|
| 78 |
+
1.4352097511291504,
|
| 79 |
+
1.7439507246017456,
|
| 80 |
+
0.8709290623664856,
|
| 81 |
+
0.458389014005661,
|
| 82 |
+
3.4738333225250244,
|
| 83 |
+
4.326974391937256,
|
| 84 |
+
6.331646919250488,
|
| 85 |
+
2.4808788299560547,
|
| 86 |
+
3.24576997756958,
|
| 87 |
+
4.848365306854248,
|
| 88 |
+
1.6670418977737427,
|
| 89 |
+
2.3548476696014404,
|
| 90 |
+
1.993442177772522,
|
| 91 |
+
2.5002496242523193,
|
| 92 |
+
0.3059796988964081,
|
| 93 |
+
2.1486732959747314,
|
| 94 |
+
6.24207878112793,
|
| 95 |
+
1.8249398469924927,
|
| 96 |
+
3.396871566772461,
|
| 97 |
+
3.9418718814849854,
|
| 98 |
+
0.29895153641700745,
|
| 99 |
+
16.587575912475586,
|
| 100 |
+
0.5883609652519226,
|
| 101 |
+
3.9876809120178223,
|
| 102 |
+
1.529473066329956,
|
| 103 |
+
6.409687519073486,
|
| 104 |
+
3.979583978652954,
|
| 105 |
+
3.5919885635375977,
|
| 106 |
+
0.8253586292266846,
|
| 107 |
+
2.201484203338623,
|
| 108 |
+
0.2219580113887787,
|
| 109 |
+
0.9641126990318298,
|
| 110 |
+
8.145115852355957,
|
| 111 |
+
4.052484035491943,
|
| 112 |
+
2.1619670391082764,
|
| 113 |
+
3.3841612339019775,
|
| 114 |
+
2.688559055328369,
|
| 115 |
+
0.4211858808994293,
|
| 116 |
+
0.799785852432251,
|
| 117 |
+
3.8083810806274414,
|
| 118 |
+
1.6525324583053589,
|
| 119 |
+
4.437976837158203,
|
| 120 |
+
1.2433699369430542,
|
| 121 |
+
3.8457489013671875,
|
| 122 |
+
4.093887805938721,
|
| 123 |
+
1.0025733709335327,
|
| 124 |
+
1.6318846940994263,
|
| 125 |
+
3.593684196472168,
|
| 126 |
+
2.961604356765747,
|
| 127 |
+
2.3301167488098145,
|
| 128 |
+
0.24246026575565338,
|
| 129 |
+
3.887437582015991,
|
| 130 |
+
2.659883737564087,
|
| 131 |
+
0.6962559223175049,
|
| 132 |
+
1.298153042793274,
|
| 133 |
+
1.4727579355239868,
|
| 134 |
+
2.848424196243286,
|
| 135 |
+
1.1403732299804688,
|
| 136 |
+
4.791389465332031,
|
| 137 |
+
4.210074424743652,
|
| 138 |
+
1.5758970975875854,
|
| 139 |
+
2.269327402114868,
|
| 140 |
+
0.18495622277259827,
|
| 141 |
+
1.9766299724578857,
|
| 142 |
+
1.7147716283798218,
|
| 143 |
+
3.4504735469818115,
|
| 144 |
+
4.908379554748535,
|
| 145 |
+
5.258154392242432,
|
| 146 |
+
0.6683526039123535,
|
| 147 |
+
1.2319836616516113,
|
| 148 |
+
6.596828460693359,
|
| 149 |
+
3.3549842834472656,
|
| 150 |
+
9.503194808959961,
|
| 151 |
+
5.601862907409668,
|
| 152 |
+
3.6974446773529053,
|
| 153 |
+
4.31416654586792,
|
| 154 |
+
1.9783447980880737,
|
| 155 |
+
0.26172930002212524,
|
| 156 |
+
0.2431989312171936,
|
| 157 |
+
2.2771925926208496,
|
| 158 |
+
1.737384557723999,
|
| 159 |
+
3.598238706588745,
|
| 160 |
+
3.4887681007385254,
|
| 161 |
+
1.9730119705200195
|
| 162 |
+
]
|
lecture_4/verified_examples/block/report.json
ADDED
|
@@ -0,0 +1,13 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
{
|
| 2 |
+
"metrics": {
|
| 3 |
+
"valid_dna": true,
|
| 4 |
+
"unique_fraction": 0.75,
|
| 5 |
+
"mean_gc": 0.59375,
|
| 6 |
+
"mean_atat_match": 0.25
|
| 7 |
+
},
|
| 8 |
+
"method": "block",
|
| 9 |
+
"data": "synthetic DNA; not biological validation",
|
| 10 |
+
"train_loss_first_20_mean": 5.375486207008362,
|
| 11 |
+
"train_loss_last_20_mean": 3.2929233014583588,
|
| 12 |
+
"validation_loss_one_mc_draw": 2.7980802059173584
|
| 13 |
+
}
|
lecture_4/verified_examples/block/samples.txt
ADDED
|
@@ -0,0 +1,8 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
ACGC
|
| 2 |
+
ACGT
|
| 3 |
+
ACGT
|
| 4 |
+
CCTA
|
| 5 |
+
ACGT
|
| 6 |
+
CGTA
|
| 7 |
+
ACGA
|
| 8 |
+
GCGC
|
lecture_4/verified_examples/cfg/config.json
ADDED
|
@@ -0,0 +1,18 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
{
|
| 2 |
+
"method": "cfg",
|
| 3 |
+
"mode": "train-sample",
|
| 4 |
+
"train_steps": 160,
|
| 5 |
+
"sample_steps": 40,
|
| 6 |
+
"classifier_steps": 160,
|
| 7 |
+
"search_steps": 40,
|
| 8 |
+
"batch_size": 16,
|
| 9 |
+
"samples": 8,
|
| 10 |
+
"length": 4,
|
| 11 |
+
"width": 32,
|
| 12 |
+
"block_size": 2,
|
| 13 |
+
"strength": 1.5,
|
| 14 |
+
"label": 1,
|
| 15 |
+
"seed": 7,
|
| 16 |
+
"data": null,
|
| 17 |
+
"out": "outputs/cfg"
|
| 18 |
+
}
|
lecture_4/verified_examples/cfg/losses.json
ADDED
|
@@ -0,0 +1,162 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
[
|
| 2 |
+
4.093783855438232,
|
| 3 |
+
6.376791000366211,
|
| 4 |
+
4.311301231384277,
|
| 5 |
+
5.4559221267700195,
|
| 6 |
+
5.326346397399902,
|
| 7 |
+
4.026087760925293,
|
| 8 |
+
2.633061647415161,
|
| 9 |
+
4.690178394317627,
|
| 10 |
+
5.398682117462158,
|
| 11 |
+
7.209366798400879,
|
| 12 |
+
4.6532487869262695,
|
| 13 |
+
5.112819671630859,
|
| 14 |
+
4.506524085998535,
|
| 15 |
+
4.212957382202148,
|
| 16 |
+
4.013753890991211,
|
| 17 |
+
3.131486177444458,
|
| 18 |
+
2.223745346069336,
|
| 19 |
+
2.7967262268066406,
|
| 20 |
+
3.788841724395752,
|
| 21 |
+
3.4797284603118896,
|
| 22 |
+
2.963750123977661,
|
| 23 |
+
4.317558288574219,
|
| 24 |
+
3.753492832183838,
|
| 25 |
+
5.249013900756836,
|
| 26 |
+
4.309139728546143,
|
| 27 |
+
3.7436439990997314,
|
| 28 |
+
10.669134140014648,
|
| 29 |
+
4.060160160064697,
|
| 30 |
+
2.0760574340820312,
|
| 31 |
+
3.706965684890747,
|
| 32 |
+
3.228757381439209,
|
| 33 |
+
2.2654337882995605,
|
| 34 |
+
4.35042667388916,
|
| 35 |
+
3.539440631866455,
|
| 36 |
+
3.0063133239746094,
|
| 37 |
+
2.5612518787384033,
|
| 38 |
+
5.035284519195557,
|
| 39 |
+
3.1380114555358887,
|
| 40 |
+
3.5969271659851074,
|
| 41 |
+
4.42738151550293,
|
| 42 |
+
2.6170473098754883,
|
| 43 |
+
4.362626075744629,
|
| 44 |
+
4.103579998016357,
|
| 45 |
+
2.5357091426849365,
|
| 46 |
+
3.5490846633911133,
|
| 47 |
+
2.5230069160461426,
|
| 48 |
+
2.873248338699341,
|
| 49 |
+
3.158176898956299,
|
| 50 |
+
4.957892894744873,
|
| 51 |
+
2.1993887424468994,
|
| 52 |
+
2.1841094493865967,
|
| 53 |
+
4.587640285491943,
|
| 54 |
+
1.8580347299575806,
|
| 55 |
+
2.267033100128174,
|
| 56 |
+
1.8229238986968994,
|
| 57 |
+
1.7472164630889893,
|
| 58 |
+
3.5524239540100098,
|
| 59 |
+
2.736100435256958,
|
| 60 |
+
3.342960834503174,
|
| 61 |
+
2.218371629714966,
|
| 62 |
+
2.875619411468506,
|
| 63 |
+
2.75288462638855,
|
| 64 |
+
1.6586329936981201,
|
| 65 |
+
3.1111838817596436,
|
| 66 |
+
4.498309135437012,
|
| 67 |
+
2.5214474201202393,
|
| 68 |
+
3.716679811477661,
|
| 69 |
+
2.803637981414795,
|
| 70 |
+
2.1772122383117676,
|
| 71 |
+
2.169358253479004,
|
| 72 |
+
1.9947891235351562,
|
| 73 |
+
2.1023731231689453,
|
| 74 |
+
2.3627333641052246,
|
| 75 |
+
1.7967830896377563,
|
| 76 |
+
1.8754620552062988,
|
| 77 |
+
2.9364476203918457,
|
| 78 |
+
2.1339101791381836,
|
| 79 |
+
4.144628524780273,
|
| 80 |
+
1.9647955894470215,
|
| 81 |
+
2.0155508518218994,
|
| 82 |
+
2.1333699226379395,
|
| 83 |
+
1.434370517730713,
|
| 84 |
+
2.0664165019989014,
|
| 85 |
+
2.257500410079956,
|
| 86 |
+
1.5195107460021973,
|
| 87 |
+
4.65607213973999,
|
| 88 |
+
2.1975350379943848,
|
| 89 |
+
0.781586229801178,
|
| 90 |
+
1.8217566013336182,
|
| 91 |
+
2.1626477241516113,
|
| 92 |
+
1.8518555164337158,
|
| 93 |
+
3.667850971221924,
|
| 94 |
+
2.1525752544403076,
|
| 95 |
+
4.357625961303711,
|
| 96 |
+
1.1739027500152588,
|
| 97 |
+
0.9763473272323608,
|
| 98 |
+
1.2919902801513672,
|
| 99 |
+
2.5940964221954346,
|
| 100 |
+
2.762819290161133,
|
| 101 |
+
1.969483733177185,
|
| 102 |
+
1.4582405090332031,
|
| 103 |
+
1.5817911624908447,
|
| 104 |
+
1.1240861415863037,
|
| 105 |
+
2.102308988571167,
|
| 106 |
+
3.745494842529297,
|
| 107 |
+
2.124908208847046,
|
| 108 |
+
2.0141687393188477,
|
| 109 |
+
1.9382771253585815,
|
| 110 |
+
1.5961400270462036,
|
| 111 |
+
1.9971623420715332,
|
| 112 |
+
2.9211485385894775,
|
| 113 |
+
2.4554712772369385,
|
| 114 |
+
2.6297366619110107,
|
| 115 |
+
1.5343077182769775,
|
| 116 |
+
3.172154664993286,
|
| 117 |
+
2.1179494857788086,
|
| 118 |
+
2.9100096225738525,
|
| 119 |
+
1.7049200534820557,
|
| 120 |
+
4.258328437805176,
|
| 121 |
+
2.4142634868621826,
|
| 122 |
+
2.198406934738159,
|
| 123 |
+
1.5292108058929443,
|
| 124 |
+
5.4681267738342285,
|
| 125 |
+
2.697603225708008,
|
| 126 |
+
3.230813503265381,
|
| 127 |
+
2.773991107940674,
|
| 128 |
+
2.4651236534118652,
|
| 129 |
+
0.9841868877410889,
|
| 130 |
+
0.870418906211853,
|
| 131 |
+
3.471985101699829,
|
| 132 |
+
2.449937582015991,
|
| 133 |
+
3.4545187950134277,
|
| 134 |
+
1.8310009241104126,
|
| 135 |
+
2.834800958633423,
|
| 136 |
+
0.7408347129821777,
|
| 137 |
+
2.966944456100464,
|
| 138 |
+
2.631303548812866,
|
| 139 |
+
2.5108256340026855,
|
| 140 |
+
2.7800302505493164,
|
| 141 |
+
2.462411880493164,
|
| 142 |
+
1.2342865467071533,
|
| 143 |
+
2.3784074783325195,
|
| 144 |
+
1.8845710754394531,
|
| 145 |
+
2.4422521591186523,
|
| 146 |
+
2.492060422897339,
|
| 147 |
+
1.5423519611358643,
|
| 148 |
+
2.8376150131225586,
|
| 149 |
+
3.2689411640167236,
|
| 150 |
+
4.493769645690918,
|
| 151 |
+
1.8553943634033203,
|
| 152 |
+
2.3281853199005127,
|
| 153 |
+
1.4675829410552979,
|
| 154 |
+
2.056398391723633,
|
| 155 |
+
3.948185443878174,
|
| 156 |
+
2.2498135566711426,
|
| 157 |
+
1.085364818572998,
|
| 158 |
+
2.6406710147857666,
|
| 159 |
+
2.0750088691711426,
|
| 160 |
+
1.4774831533432007,
|
| 161 |
+
3.007999897003174
|
| 162 |
+
]
|
lecture_4/verified_examples/cfg/report.json
ADDED
|
@@ -0,0 +1,13 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
{
|
| 2 |
+
"metrics": {
|
| 3 |
+
"valid_dna": true,
|
| 4 |
+
"unique_fraction": 0.625,
|
| 5 |
+
"mean_gc": 0.9375,
|
| 6 |
+
"mean_atat_match": 0.03125
|
| 7 |
+
},
|
| 8 |
+
"method": "cfg",
|
| 9 |
+
"data": "synthetic DNA; not biological validation",
|
| 10 |
+
"train_loss_first_20_mean": 4.372067654132843,
|
| 11 |
+
"train_loss_last_20_mean": 2.3383171617984773,
|
| 12 |
+
"validation_loss_one_mc_draw": 2.0500779151916504
|
| 13 |
+
}
|
lecture_4/verified_examples/cfg/samples.txt
ADDED
|
@@ -0,0 +1,8 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
GCCC
|
| 2 |
+
GCGC
|
| 3 |
+
CGGC
|
| 4 |
+
GCGC
|
| 5 |
+
GCGC
|
| 6 |
+
GCAC
|
| 7 |
+
GCGA
|
| 8 |
+
GCGC
|
lecture_4/verified_examples/classifier-exact/config.json
ADDED
|
@@ -0,0 +1,18 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
{
|
| 2 |
+
"method": "classifier-exact",
|
| 3 |
+
"mode": "train-sample",
|
| 4 |
+
"train_steps": 160,
|
| 5 |
+
"sample_steps": 40,
|
| 6 |
+
"classifier_steps": 160,
|
| 7 |
+
"search_steps": 40,
|
| 8 |
+
"batch_size": 16,
|
| 9 |
+
"samples": 8,
|
| 10 |
+
"length": 4,
|
| 11 |
+
"width": 32,
|
| 12 |
+
"block_size": 2,
|
| 13 |
+
"strength": 1.5,
|
| 14 |
+
"label": 1,
|
| 15 |
+
"seed": 7,
|
| 16 |
+
"data": null,
|
| 17 |
+
"out": "outputs/classifier-exact"
|
| 18 |
+
}
|
lecture_4/verified_examples/classifier-exact/losses.json
ADDED
|
@@ -0,0 +1,162 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
[
|
| 2 |
+
4.01927375793457,
|
| 3 |
+
3.8247132301330566,
|
| 4 |
+
6.5967559814453125,
|
| 5 |
+
3.0379931926727295,
|
| 6 |
+
5.381988048553467,
|
| 7 |
+
4.364372730255127,
|
| 8 |
+
6.268919944763184,
|
| 9 |
+
4.896388530731201,
|
| 10 |
+
2.5558066368103027,
|
| 11 |
+
5.216179847717285,
|
| 12 |
+
4.734250545501709,
|
| 13 |
+
5.792362213134766,
|
| 14 |
+
3.2293550968170166,
|
| 15 |
+
4.882113456726074,
|
| 16 |
+
4.7334160804748535,
|
| 17 |
+
4.435980796813965,
|
| 18 |
+
5.003822326660156,
|
| 19 |
+
5.925942420959473,
|
| 20 |
+
7.791959762573242,
|
| 21 |
+
4.085095405578613,
|
| 22 |
+
3.899994373321533,
|
| 23 |
+
3.2904622554779053,
|
| 24 |
+
3.4589755535125732,
|
| 25 |
+
2.720918655395508,
|
| 26 |
+
3.027308225631714,
|
| 27 |
+
3.0458157062530518,
|
| 28 |
+
3.7035269737243652,
|
| 29 |
+
3.8884239196777344,
|
| 30 |
+
4.001327037811279,
|
| 31 |
+
4.476953983306885,
|
| 32 |
+
2.1432690620422363,
|
| 33 |
+
4.810318946838379,
|
| 34 |
+
2.952962875366211,
|
| 35 |
+
4.420904159545898,
|
| 36 |
+
3.500626564025879,
|
| 37 |
+
4.333303928375244,
|
| 38 |
+
3.9793171882629395,
|
| 39 |
+
1.8058104515075684,
|
| 40 |
+
4.20003080368042,
|
| 41 |
+
3.7151639461517334,
|
| 42 |
+
2.2277348041534424,
|
| 43 |
+
2.7376279830932617,
|
| 44 |
+
3.615410327911377,
|
| 45 |
+
4.500275611877441,
|
| 46 |
+
5.066170692443848,
|
| 47 |
+
3.6637136936187744,
|
| 48 |
+
2.8754518032073975,
|
| 49 |
+
2.339322328567505,
|
| 50 |
+
2.037803888320923,
|
| 51 |
+
2.1924853324890137,
|
| 52 |
+
3.4916043281555176,
|
| 53 |
+
2.519395351409912,
|
| 54 |
+
2.597137451171875,
|
| 55 |
+
2.4708733558654785,
|
| 56 |
+
2.5220284461975098,
|
| 57 |
+
2.5972704887390137,
|
| 58 |
+
3.05387806892395,
|
| 59 |
+
5.856385231018066,
|
| 60 |
+
2.131591320037842,
|
| 61 |
+
3.3278908729553223,
|
| 62 |
+
2.400947093963623,
|
| 63 |
+
5.875596046447754,
|
| 64 |
+
3.1304073333740234,
|
| 65 |
+
3.21421217918396,
|
| 66 |
+
2.4150002002716064,
|
| 67 |
+
2.215695381164551,
|
| 68 |
+
2.139690637588501,
|
| 69 |
+
3.2768962383270264,
|
| 70 |
+
2.847025156021118,
|
| 71 |
+
3.595715045928955,
|
| 72 |
+
1.7762316465377808,
|
| 73 |
+
3.3017048835754395,
|
| 74 |
+
3.1203761100769043,
|
| 75 |
+
2.4890081882476807,
|
| 76 |
+
4.99165153503418,
|
| 77 |
+
3.243589401245117,
|
| 78 |
+
2.338869333267212,
|
| 79 |
+
1.9984252452850342,
|
| 80 |
+
3.838407516479492,
|
| 81 |
+
2.6884045600891113,
|
| 82 |
+
6.401689052581787,
|
| 83 |
+
1.2461563348770142,
|
| 84 |
+
2.748701810836792,
|
| 85 |
+
8.534076690673828,
|
| 86 |
+
3.991495132446289,
|
| 87 |
+
2.489625930786133,
|
| 88 |
+
2.2347571849823,
|
| 89 |
+
2.9534873962402344,
|
| 90 |
+
2.5714945793151855,
|
| 91 |
+
1.147279977798462,
|
| 92 |
+
2.4217231273651123,
|
| 93 |
+
1.3990497589111328,
|
| 94 |
+
2.5761783123016357,
|
| 95 |
+
2.3062100410461426,
|
| 96 |
+
3.08569073677063,
|
| 97 |
+
1.516831874847412,
|
| 98 |
+
3.085536003112793,
|
| 99 |
+
2.8733534812927246,
|
| 100 |
+
1.2969731092453003,
|
| 101 |
+
3.7263309955596924,
|
| 102 |
+
2.435549736022949,
|
| 103 |
+
3.5106093883514404,
|
| 104 |
+
2.7424113750457764,
|
| 105 |
+
3.264981269836426,
|
| 106 |
+
1.683659315109253,
|
| 107 |
+
1.149167776107788,
|
| 108 |
+
2.775585889816284,
|
| 109 |
+
3.208291530609131,
|
| 110 |
+
2.314664602279663,
|
| 111 |
+
3.91782283782959,
|
| 112 |
+
1.7896864414215088,
|
| 113 |
+
2.5307981967926025,
|
| 114 |
+
2.4437763690948486,
|
| 115 |
+
2.521484136581421,
|
| 116 |
+
2.2941794395446777,
|
| 117 |
+
1.6110831499099731,
|
| 118 |
+
3.128873825073242,
|
| 119 |
+
3.0251963138580322,
|
| 120 |
+
2.2769689559936523,
|
| 121 |
+
1.4722540378570557,
|
| 122 |
+
1.7136496305465698,
|
| 123 |
+
5.040107727050781,
|
| 124 |
+
2.9950344562530518,
|
| 125 |
+
4.272597789764404,
|
| 126 |
+
2.686481237411499,
|
| 127 |
+
2.7622010707855225,
|
| 128 |
+
3.0286059379577637,
|
| 129 |
+
2.0905165672302246,
|
| 130 |
+
2.4465155601501465,
|
| 131 |
+
2.035715341567993,
|
| 132 |
+
2.315361499786377,
|
| 133 |
+
1.4022663831710815,
|
| 134 |
+
3.5252881050109863,
|
| 135 |
+
2.6312732696533203,
|
| 136 |
+
3.7926905155181885,
|
| 137 |
+
1.7568085193634033,
|
| 138 |
+
1.798338532447815,
|
| 139 |
+
3.3464834690093994,
|
| 140 |
+
2.249241828918457,
|
| 141 |
+
2.2580409049987793,
|
| 142 |
+
4.0587310791015625,
|
| 143 |
+
2.206573247909546,
|
| 144 |
+
0.6222774982452393,
|
| 145 |
+
1.8168766498565674,
|
| 146 |
+
3.9070353507995605,
|
| 147 |
+
2.2064626216888428,
|
| 148 |
+
1.1733113527297974,
|
| 149 |
+
1.383936882019043,
|
| 150 |
+
2.251495599746704,
|
| 151 |
+
2.962852954864502,
|
| 152 |
+
2.8415677547454834,
|
| 153 |
+
1.3905932903289795,
|
| 154 |
+
3.165285587310791,
|
| 155 |
+
3.477961540222168,
|
| 156 |
+
1.7277320623397827,
|
| 157 |
+
2.4671823978424072,
|
| 158 |
+
5.010605812072754,
|
| 159 |
+
3.3257808685302734,
|
| 160 |
+
2.7084109783172607,
|
| 161 |
+
3.5739824771881104
|
| 162 |
+
]
|
lecture_4/verified_examples/classifier-exact/report.json
ADDED
|
@@ -0,0 +1,15 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
{
|
| 2 |
+
"metrics": {
|
| 3 |
+
"valid_dna": true,
|
| 4 |
+
"unique_fraction": 0.625,
|
| 5 |
+
"mean_gc": 0.84375,
|
| 6 |
+
"mean_atat_match": 0.0625
|
| 7 |
+
},
|
| 8 |
+
"method": "classifier-exact",
|
| 9 |
+
"data": "synthetic DNA; not biological validation",
|
| 10 |
+
"classifier_final_loss": 0.3678973913192749,
|
| 11 |
+
"train_loss_first_20_mean": 4.838834500312805,
|
| 12 |
+
"train_loss_last_20_mean": 2.613932800292969,
|
| 13 |
+
"validation_loss_one_mc_draw": 3.0652260780334473,
|
| 14 |
+
"guidance": "learned noisy classifier; final residual masks use denoiser closure"
|
| 15 |
+
}
|
lecture_4/verified_examples/classifier-exact/samples.txt
ADDED
|
@@ -0,0 +1,8 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
GCGC
|
| 2 |
+
GCGC
|
| 3 |
+
ACGC
|
| 4 |
+
GCGC
|
| 5 |
+
GGTC
|
| 6 |
+
GCAA
|
| 7 |
+
CGTC
|
| 8 |
+
GCGC
|
lecture_4/verified_examples/classifier-gradient/config.json
ADDED
|
@@ -0,0 +1,18 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
{
|
| 2 |
+
"method": "classifier-gradient",
|
| 3 |
+
"mode": "train-sample",
|
| 4 |
+
"train_steps": 160,
|
| 5 |
+
"sample_steps": 40,
|
| 6 |
+
"classifier_steps": 160,
|
| 7 |
+
"search_steps": 40,
|
| 8 |
+
"batch_size": 16,
|
| 9 |
+
"samples": 8,
|
| 10 |
+
"length": 4,
|
| 11 |
+
"width": 32,
|
| 12 |
+
"block_size": 2,
|
| 13 |
+
"strength": 1.5,
|
| 14 |
+
"label": 1,
|
| 15 |
+
"seed": 7,
|
| 16 |
+
"data": null,
|
| 17 |
+
"out": "outputs/classifier-gradient"
|
| 18 |
+
}
|
lecture_4/verified_examples/classifier-gradient/losses.json
ADDED
|
@@ -0,0 +1,162 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
[
|
| 2 |
+
4.01927375793457,
|
| 3 |
+
3.8247132301330566,
|
| 4 |
+
6.5967559814453125,
|
| 5 |
+
3.0379931926727295,
|
| 6 |
+
5.381988048553467,
|
| 7 |
+
4.364372730255127,
|
| 8 |
+
6.268919944763184,
|
| 9 |
+
4.896388530731201,
|
| 10 |
+
2.5558066368103027,
|
| 11 |
+
5.216179847717285,
|
| 12 |
+
4.734250545501709,
|
| 13 |
+
5.792362213134766,
|
| 14 |
+
3.2293550968170166,
|
| 15 |
+
4.882113456726074,
|
| 16 |
+
4.7334160804748535,
|
| 17 |
+
4.435980796813965,
|
| 18 |
+
5.003822326660156,
|
| 19 |
+
5.925942420959473,
|
| 20 |
+
7.791959762573242,
|
| 21 |
+
4.085095405578613,
|
| 22 |
+
3.899994373321533,
|
| 23 |
+
3.2904622554779053,
|
| 24 |
+
3.4589755535125732,
|
| 25 |
+
2.720918655395508,
|
| 26 |
+
3.027308225631714,
|
| 27 |
+
3.0458157062530518,
|
| 28 |
+
3.7035269737243652,
|
| 29 |
+
3.8884239196777344,
|
| 30 |
+
4.001327037811279,
|
| 31 |
+
4.476953983306885,
|
| 32 |
+
2.1432690620422363,
|
| 33 |
+
4.810318946838379,
|
| 34 |
+
2.952962875366211,
|
| 35 |
+
4.420904159545898,
|
| 36 |
+
3.500626564025879,
|
| 37 |
+
4.333303928375244,
|
| 38 |
+
3.9793171882629395,
|
| 39 |
+
1.8058104515075684,
|
| 40 |
+
4.20003080368042,
|
| 41 |
+
3.7151639461517334,
|
| 42 |
+
2.2277348041534424,
|
| 43 |
+
2.7376279830932617,
|
| 44 |
+
3.615410327911377,
|
| 45 |
+
4.500275611877441,
|
| 46 |
+
5.066170692443848,
|
| 47 |
+
3.6637136936187744,
|
| 48 |
+
2.8754518032073975,
|
| 49 |
+
2.339322328567505,
|
| 50 |
+
2.037803888320923,
|
| 51 |
+
2.1924853324890137,
|
| 52 |
+
3.4916043281555176,
|
| 53 |
+
2.519395351409912,
|
| 54 |
+
2.597137451171875,
|
| 55 |
+
2.4708733558654785,
|
| 56 |
+
2.5220284461975098,
|
| 57 |
+
2.5972704887390137,
|
| 58 |
+
3.05387806892395,
|
| 59 |
+
5.856385231018066,
|
| 60 |
+
2.131591320037842,
|
| 61 |
+
3.3278908729553223,
|
| 62 |
+
2.400947093963623,
|
| 63 |
+
5.875596046447754,
|
| 64 |
+
3.1304073333740234,
|
| 65 |
+
3.21421217918396,
|
| 66 |
+
2.4150002002716064,
|
| 67 |
+
2.215695381164551,
|
| 68 |
+
2.139690637588501,
|
| 69 |
+
3.2768962383270264,
|
| 70 |
+
2.847025156021118,
|
| 71 |
+
3.595715045928955,
|
| 72 |
+
1.7762316465377808,
|
| 73 |
+
3.3017048835754395,
|
| 74 |
+
3.1203761100769043,
|
| 75 |
+
2.4890081882476807,
|
| 76 |
+
4.99165153503418,
|
| 77 |
+
3.243589401245117,
|
| 78 |
+
2.338869333267212,
|
| 79 |
+
1.9984252452850342,
|
| 80 |
+
3.838407516479492,
|
| 81 |
+
2.6884045600891113,
|
| 82 |
+
6.401689052581787,
|
| 83 |
+
1.2461563348770142,
|
| 84 |
+
2.748701810836792,
|
| 85 |
+
8.534076690673828,
|
| 86 |
+
3.991495132446289,
|
| 87 |
+
2.489625930786133,
|
| 88 |
+
2.2347571849823,
|
| 89 |
+
2.9534873962402344,
|
| 90 |
+
2.5714945793151855,
|
| 91 |
+
1.147279977798462,
|
| 92 |
+
2.4217231273651123,
|
| 93 |
+
1.3990497589111328,
|
| 94 |
+
2.5761783123016357,
|
| 95 |
+
2.3062100410461426,
|
| 96 |
+
3.08569073677063,
|
| 97 |
+
1.516831874847412,
|
| 98 |
+
3.085536003112793,
|
| 99 |
+
2.8733534812927246,
|
| 100 |
+
1.2969731092453003,
|
| 101 |
+
3.7263309955596924,
|
| 102 |
+
2.435549736022949,
|
| 103 |
+
3.5106093883514404,
|
| 104 |
+
2.7424113750457764,
|
| 105 |
+
3.264981269836426,
|
| 106 |
+
1.683659315109253,
|
| 107 |
+
1.149167776107788,
|
| 108 |
+
2.775585889816284,
|
| 109 |
+
3.208291530609131,
|
| 110 |
+
2.314664602279663,
|
| 111 |
+
3.91782283782959,
|
| 112 |
+
1.7896864414215088,
|
| 113 |
+
2.5307981967926025,
|
| 114 |
+
2.4437763690948486,
|
| 115 |
+
2.521484136581421,
|
| 116 |
+
2.2941794395446777,
|
| 117 |
+
1.6110831499099731,
|
| 118 |
+
3.128873825073242,
|
| 119 |
+
3.0251963138580322,
|
| 120 |
+
2.2769689559936523,
|
| 121 |
+
1.4722540378570557,
|
| 122 |
+
1.7136496305465698,
|
| 123 |
+
5.040107727050781,
|
| 124 |
+
2.9950344562530518,
|
| 125 |
+
4.272597789764404,
|
| 126 |
+
2.686481237411499,
|
| 127 |
+
2.7622010707855225,
|
| 128 |
+
3.0286059379577637,
|
| 129 |
+
2.0905165672302246,
|
| 130 |
+
2.4465155601501465,
|
| 131 |
+
2.035715341567993,
|
| 132 |
+
2.315361499786377,
|
| 133 |
+
1.4022663831710815,
|
| 134 |
+
3.5252881050109863,
|
| 135 |
+
2.6312732696533203,
|
| 136 |
+
3.7926905155181885,
|
| 137 |
+
1.7568085193634033,
|
| 138 |
+
1.798338532447815,
|
| 139 |
+
3.3464834690093994,
|
| 140 |
+
2.249241828918457,
|
| 141 |
+
2.2580409049987793,
|
| 142 |
+
4.0587310791015625,
|
| 143 |
+
2.206573247909546,
|
| 144 |
+
0.6222774982452393,
|
| 145 |
+
1.8168766498565674,
|
| 146 |
+
3.9070353507995605,
|
| 147 |
+
2.2064626216888428,
|
| 148 |
+
1.1733113527297974,
|
| 149 |
+
1.383936882019043,
|
| 150 |
+
2.251495599746704,
|
| 151 |
+
2.962852954864502,
|
| 152 |
+
2.8415677547454834,
|
| 153 |
+
1.3905932903289795,
|
| 154 |
+
3.165285587310791,
|
| 155 |
+
3.477961540222168,
|
| 156 |
+
1.7277320623397827,
|
| 157 |
+
2.4671823978424072,
|
| 158 |
+
5.010605812072754,
|
| 159 |
+
3.3257808685302734,
|
| 160 |
+
2.7084109783172607,
|
| 161 |
+
3.5739824771881104
|
| 162 |
+
]
|
lecture_4/verified_examples/classifier-gradient/report.json
ADDED
|
@@ -0,0 +1,15 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
{
|
| 2 |
+
"metrics": {
|
| 3 |
+
"valid_dna": true,
|
| 4 |
+
"unique_fraction": 0.5,
|
| 5 |
+
"mean_gc": 0.90625,
|
| 6 |
+
"mean_atat_match": 0.03125
|
| 7 |
+
},
|
| 8 |
+
"method": "classifier-gradient",
|
| 9 |
+
"data": "synthetic DNA; not biological validation",
|
| 10 |
+
"classifier_final_loss": 0.3678973913192749,
|
| 11 |
+
"train_loss_first_20_mean": 4.838834500312805,
|
| 12 |
+
"train_loss_last_20_mean": 2.613932800292969,
|
| 13 |
+
"validation_loss_one_mc_draw": 3.0652260780334473,
|
| 14 |
+
"guidance": "learned noisy classifier; final residual masks use denoiser closure"
|
| 15 |
+
}
|
lecture_4/verified_examples/classifier-gradient/samples.txt
ADDED
|
@@ -0,0 +1,8 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
GCGC
|
| 2 |
+
CCGC
|
| 3 |
+
GCTC
|
| 4 |
+
GCGC
|
| 5 |
+
GCGC
|
| 6 |
+
GCAA
|
| 7 |
+
GCGC
|
| 8 |
+
GCGC
|
lecture_4/verified_examples/environment.json
ADDED
|
@@ -0,0 +1,7 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
{
|
| 2 |
+
"python": "3.12.14",
|
| 3 |
+
"torch": "2.14.0+cpu",
|
| 4 |
+
"numpy": "2.3.5",
|
| 5 |
+
"scipy": "1.17.0",
|
| 6 |
+
"device": "cpu"
|
| 7 |
+
}
|
lecture_4/verified_examples/mdlm/config.json
ADDED
|
@@ -0,0 +1,18 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
{
|
| 2 |
+
"method": "mdlm",
|
| 3 |
+
"mode": "train-sample",
|
| 4 |
+
"train_steps": 160,
|
| 5 |
+
"sample_steps": 40,
|
| 6 |
+
"classifier_steps": 160,
|
| 7 |
+
"search_steps": 40,
|
| 8 |
+
"batch_size": 16,
|
| 9 |
+
"samples": 8,
|
| 10 |
+
"length": 4,
|
| 11 |
+
"width": 32,
|
| 12 |
+
"block_size": 2,
|
| 13 |
+
"strength": 1.5,
|
| 14 |
+
"label": 1,
|
| 15 |
+
"seed": 7,
|
| 16 |
+
"data": null,
|
| 17 |
+
"out": "outputs/mdlm"
|
| 18 |
+
}
|
lecture_4/verified_examples/mdlm/losses.json
ADDED
|
@@ -0,0 +1,162 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
[
|
| 2 |
+
4.01927375793457,
|
| 3 |
+
3.8247132301330566,
|
| 4 |
+
6.5967559814453125,
|
| 5 |
+
3.0379931926727295,
|
| 6 |
+
5.381988048553467,
|
| 7 |
+
4.364372730255127,
|
| 8 |
+
6.268919944763184,
|
| 9 |
+
4.896388530731201,
|
| 10 |
+
2.5558066368103027,
|
| 11 |
+
5.216179847717285,
|
| 12 |
+
4.734250545501709,
|
| 13 |
+
5.792362213134766,
|
| 14 |
+
3.2293550968170166,
|
| 15 |
+
4.882113456726074,
|
| 16 |
+
4.7334160804748535,
|
| 17 |
+
4.435980796813965,
|
| 18 |
+
5.003822326660156,
|
| 19 |
+
5.925942420959473,
|
| 20 |
+
7.791959762573242,
|
| 21 |
+
4.085095405578613,
|
| 22 |
+
3.899994373321533,
|
| 23 |
+
3.2904622554779053,
|
| 24 |
+
3.4589755535125732,
|
| 25 |
+
2.720918655395508,
|
| 26 |
+
3.027308225631714,
|
| 27 |
+
3.0458157062530518,
|
| 28 |
+
3.7035269737243652,
|
| 29 |
+
3.8884239196777344,
|
| 30 |
+
4.001327037811279,
|
| 31 |
+
4.476953983306885,
|
| 32 |
+
2.1432690620422363,
|
| 33 |
+
4.810318946838379,
|
| 34 |
+
2.952962875366211,
|
| 35 |
+
4.420904159545898,
|
| 36 |
+
3.500626564025879,
|
| 37 |
+
4.333303928375244,
|
| 38 |
+
3.9793171882629395,
|
| 39 |
+
1.8058104515075684,
|
| 40 |
+
4.20003080368042,
|
| 41 |
+
3.7151639461517334,
|
| 42 |
+
2.2277348041534424,
|
| 43 |
+
2.7376279830932617,
|
| 44 |
+
3.615410327911377,
|
| 45 |
+
4.500275611877441,
|
| 46 |
+
5.066170692443848,
|
| 47 |
+
3.6637136936187744,
|
| 48 |
+
2.8754518032073975,
|
| 49 |
+
2.339322328567505,
|
| 50 |
+
2.037803888320923,
|
| 51 |
+
2.1924853324890137,
|
| 52 |
+
3.4916043281555176,
|
| 53 |
+
2.519395351409912,
|
| 54 |
+
2.597137451171875,
|
| 55 |
+
2.4708733558654785,
|
| 56 |
+
2.5220284461975098,
|
| 57 |
+
2.5972704887390137,
|
| 58 |
+
3.05387806892395,
|
| 59 |
+
5.856385231018066,
|
| 60 |
+
2.131591320037842,
|
| 61 |
+
3.3278908729553223,
|
| 62 |
+
2.400947093963623,
|
| 63 |
+
5.875596046447754,
|
| 64 |
+
3.1304073333740234,
|
| 65 |
+
3.21421217918396,
|
| 66 |
+
2.4150002002716064,
|
| 67 |
+
2.215695381164551,
|
| 68 |
+
2.139690637588501,
|
| 69 |
+
3.2768962383270264,
|
| 70 |
+
2.847025156021118,
|
| 71 |
+
3.595715045928955,
|
| 72 |
+
1.7762316465377808,
|
| 73 |
+
3.3017048835754395,
|
| 74 |
+
3.1203761100769043,
|
| 75 |
+
2.4890081882476807,
|
| 76 |
+
4.99165153503418,
|
| 77 |
+
3.243589401245117,
|
| 78 |
+
2.338869333267212,
|
| 79 |
+
1.9984252452850342,
|
| 80 |
+
3.838407516479492,
|
| 81 |
+
2.6884045600891113,
|
| 82 |
+
6.401689052581787,
|
| 83 |
+
1.2461563348770142,
|
| 84 |
+
2.748701810836792,
|
| 85 |
+
8.534076690673828,
|
| 86 |
+
3.991495132446289,
|
| 87 |
+
2.489625930786133,
|
| 88 |
+
2.2347571849823,
|
| 89 |
+
2.9534873962402344,
|
| 90 |
+
2.5714945793151855,
|
| 91 |
+
1.147279977798462,
|
| 92 |
+
2.4217231273651123,
|
| 93 |
+
1.3990497589111328,
|
| 94 |
+
2.5761783123016357,
|
| 95 |
+
2.3062100410461426,
|
| 96 |
+
3.08569073677063,
|
| 97 |
+
1.516831874847412,
|
| 98 |
+
3.085536003112793,
|
| 99 |
+
2.8733534812927246,
|
| 100 |
+
1.2969731092453003,
|
| 101 |
+
3.7263309955596924,
|
| 102 |
+
2.435549736022949,
|
| 103 |
+
3.5106093883514404,
|
| 104 |
+
2.7424113750457764,
|
| 105 |
+
3.264981269836426,
|
| 106 |
+
1.683659315109253,
|
| 107 |
+
1.149167776107788,
|
| 108 |
+
2.775585889816284,
|
| 109 |
+
3.208291530609131,
|
| 110 |
+
2.314664602279663,
|
| 111 |
+
3.91782283782959,
|
| 112 |
+
1.7896864414215088,
|
| 113 |
+
2.5307981967926025,
|
| 114 |
+
2.4437763690948486,
|
| 115 |
+
2.521484136581421,
|
| 116 |
+
2.2941794395446777,
|
| 117 |
+
1.6110831499099731,
|
| 118 |
+
3.128873825073242,
|
| 119 |
+
3.0251963138580322,
|
| 120 |
+
2.2769689559936523,
|
| 121 |
+
1.4722540378570557,
|
| 122 |
+
1.7136496305465698,
|
| 123 |
+
5.040107727050781,
|
| 124 |
+
2.9950344562530518,
|
| 125 |
+
4.272597789764404,
|
| 126 |
+
2.686481237411499,
|
| 127 |
+
2.7622010707855225,
|
| 128 |
+
3.0286059379577637,
|
| 129 |
+
2.0905165672302246,
|
| 130 |
+
2.4465155601501465,
|
| 131 |
+
2.035715341567993,
|
| 132 |
+
2.315361499786377,
|
| 133 |
+
1.4022663831710815,
|
| 134 |
+
3.5252881050109863,
|
| 135 |
+
2.6312732696533203,
|
| 136 |
+
3.7926905155181885,
|
| 137 |
+
1.7568085193634033,
|
| 138 |
+
1.798338532447815,
|
| 139 |
+
3.3464834690093994,
|
| 140 |
+
2.249241828918457,
|
| 141 |
+
2.2580409049987793,
|
| 142 |
+
4.0587310791015625,
|
| 143 |
+
2.206573247909546,
|
| 144 |
+
0.6222774982452393,
|
| 145 |
+
1.8168766498565674,
|
| 146 |
+
3.9070353507995605,
|
| 147 |
+
2.2064626216888428,
|
| 148 |
+
1.1733113527297974,
|
| 149 |
+
1.383936882019043,
|
| 150 |
+
2.251495599746704,
|
| 151 |
+
2.962852954864502,
|
| 152 |
+
2.8415677547454834,
|
| 153 |
+
1.3905932903289795,
|
| 154 |
+
3.165285587310791,
|
| 155 |
+
3.477961540222168,
|
| 156 |
+
1.7277320623397827,
|
| 157 |
+
2.4671823978424072,
|
| 158 |
+
5.010605812072754,
|
| 159 |
+
3.3257808685302734,
|
| 160 |
+
2.7084109783172607,
|
| 161 |
+
3.5739824771881104
|
| 162 |
+
]
|
lecture_4/verified_examples/mdlm/report.json
ADDED
|
@@ -0,0 +1,13 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
{
|
| 2 |
+
"metrics": {
|
| 3 |
+
"valid_dna": true,
|
| 4 |
+
"unique_fraction": 0.75,
|
| 5 |
+
"mean_gc": 0.625,
|
| 6 |
+
"mean_atat_match": 0.09375
|
| 7 |
+
},
|
| 8 |
+
"method": "mdlm",
|
| 9 |
+
"data": "synthetic DNA; not biological validation",
|
| 10 |
+
"train_loss_first_20_mean": 4.838834500312805,
|
| 11 |
+
"train_loss_last_20_mean": 2.613932800292969,
|
| 12 |
+
"validation_loss_one_mc_draw": 2.5891432762145996
|
| 13 |
+
}
|
lecture_4/verified_examples/mdlm/samples.txt
ADDED
|
@@ -0,0 +1,8 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
CGTA
|
| 2 |
+
ATGT
|
| 3 |
+
GCGC
|
| 4 |
+
TCGC
|
| 5 |
+
GCGC
|
| 6 |
+
GCGC
|
| 7 |
+
TATA
|
| 8 |
+
TCGA
|
lecture_4/verified_examples/peptune/config.json
ADDED
|
@@ -0,0 +1,18 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
{
|
| 2 |
+
"method": "peptune",
|
| 3 |
+
"mode": "train-sample",
|
| 4 |
+
"train_steps": 160,
|
| 5 |
+
"sample_steps": 40,
|
| 6 |
+
"classifier_steps": 160,
|
| 7 |
+
"search_steps": 40,
|
| 8 |
+
"batch_size": 16,
|
| 9 |
+
"samples": 8,
|
| 10 |
+
"length": 4,
|
| 11 |
+
"width": 32,
|
| 12 |
+
"block_size": 2,
|
| 13 |
+
"strength": 1.5,
|
| 14 |
+
"label": 1,
|
| 15 |
+
"seed": 7,
|
| 16 |
+
"data": null,
|
| 17 |
+
"out": "outputs/peptune"
|
| 18 |
+
}
|
lecture_4/verified_examples/peptune/losses.json
ADDED
|
@@ -0,0 +1,162 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
[
|
| 2 |
+
4.01927375793457,
|
| 3 |
+
3.8247132301330566,
|
| 4 |
+
6.5967559814453125,
|
| 5 |
+
3.0379931926727295,
|
| 6 |
+
5.381988048553467,
|
| 7 |
+
4.364372730255127,
|
| 8 |
+
6.268919944763184,
|
| 9 |
+
4.896388530731201,
|
| 10 |
+
2.5558066368103027,
|
| 11 |
+
5.216179847717285,
|
| 12 |
+
4.734250545501709,
|
| 13 |
+
5.792362213134766,
|
| 14 |
+
3.2293550968170166,
|
| 15 |
+
4.882113456726074,
|
| 16 |
+
4.7334160804748535,
|
| 17 |
+
4.435980796813965,
|
| 18 |
+
5.003822326660156,
|
| 19 |
+
5.925942420959473,
|
| 20 |
+
7.791959762573242,
|
| 21 |
+
4.085095405578613,
|
| 22 |
+
3.899994373321533,
|
| 23 |
+
3.2904622554779053,
|
| 24 |
+
3.4589755535125732,
|
| 25 |
+
2.720918655395508,
|
| 26 |
+
3.027308225631714,
|
| 27 |
+
3.0458157062530518,
|
| 28 |
+
3.7035269737243652,
|
| 29 |
+
3.8884239196777344,
|
| 30 |
+
4.001327037811279,
|
| 31 |
+
4.476953983306885,
|
| 32 |
+
2.1432690620422363,
|
| 33 |
+
4.810318946838379,
|
| 34 |
+
2.952962875366211,
|
| 35 |
+
4.420904159545898,
|
| 36 |
+
3.500626564025879,
|
| 37 |
+
4.333303928375244,
|
| 38 |
+
3.9793171882629395,
|
| 39 |
+
1.8058104515075684,
|
| 40 |
+
4.20003080368042,
|
| 41 |
+
3.7151639461517334,
|
| 42 |
+
2.2277348041534424,
|
| 43 |
+
2.7376279830932617,
|
| 44 |
+
3.615410327911377,
|
| 45 |
+
4.500275611877441,
|
| 46 |
+
5.066170692443848,
|
| 47 |
+
3.6637136936187744,
|
| 48 |
+
2.8754518032073975,
|
| 49 |
+
2.339322328567505,
|
| 50 |
+
2.037803888320923,
|
| 51 |
+
2.1924853324890137,
|
| 52 |
+
3.4916043281555176,
|
| 53 |
+
2.519395351409912,
|
| 54 |
+
2.597137451171875,
|
| 55 |
+
2.4708733558654785,
|
| 56 |
+
2.5220284461975098,
|
| 57 |
+
2.5972704887390137,
|
| 58 |
+
3.05387806892395,
|
| 59 |
+
5.856385231018066,
|
| 60 |
+
2.131591320037842,
|
| 61 |
+
3.3278908729553223,
|
| 62 |
+
2.400947093963623,
|
| 63 |
+
5.875596046447754,
|
| 64 |
+
3.1304073333740234,
|
| 65 |
+
3.21421217918396,
|
| 66 |
+
2.4150002002716064,
|
| 67 |
+
2.215695381164551,
|
| 68 |
+
2.139690637588501,
|
| 69 |
+
3.2768962383270264,
|
| 70 |
+
2.847025156021118,
|
| 71 |
+
3.595715045928955,
|
| 72 |
+
1.7762316465377808,
|
| 73 |
+
3.3017048835754395,
|
| 74 |
+
3.1203761100769043,
|
| 75 |
+
2.4890081882476807,
|
| 76 |
+
4.99165153503418,
|
| 77 |
+
3.243589401245117,
|
| 78 |
+
2.338869333267212,
|
| 79 |
+
1.9984252452850342,
|
| 80 |
+
3.838407516479492,
|
| 81 |
+
2.6884045600891113,
|
| 82 |
+
6.401689052581787,
|
| 83 |
+
1.2461563348770142,
|
| 84 |
+
2.748701810836792,
|
| 85 |
+
8.534076690673828,
|
| 86 |
+
3.991495132446289,
|
| 87 |
+
2.489625930786133,
|
| 88 |
+
2.2347571849823,
|
| 89 |
+
2.9534873962402344,
|
| 90 |
+
2.5714945793151855,
|
| 91 |
+
1.147279977798462,
|
| 92 |
+
2.4217231273651123,
|
| 93 |
+
1.3990497589111328,
|
| 94 |
+
2.5761783123016357,
|
| 95 |
+
2.3062100410461426,
|
| 96 |
+
3.08569073677063,
|
| 97 |
+
1.516831874847412,
|
| 98 |
+
3.085536003112793,
|
| 99 |
+
2.8733534812927246,
|
| 100 |
+
1.2969731092453003,
|
| 101 |
+
3.7263309955596924,
|
| 102 |
+
2.435549736022949,
|
| 103 |
+
3.5106093883514404,
|
| 104 |
+
2.7424113750457764,
|
| 105 |
+
3.264981269836426,
|
| 106 |
+
1.683659315109253,
|
| 107 |
+
1.149167776107788,
|
| 108 |
+
2.775585889816284,
|
| 109 |
+
3.208291530609131,
|
| 110 |
+
2.314664602279663,
|
| 111 |
+
3.91782283782959,
|
| 112 |
+
1.7896864414215088,
|
| 113 |
+
2.5307981967926025,
|
| 114 |
+
2.4437763690948486,
|
| 115 |
+
2.521484136581421,
|
| 116 |
+
2.2941794395446777,
|
| 117 |
+
1.6110831499099731,
|
| 118 |
+
3.128873825073242,
|
| 119 |
+
3.0251963138580322,
|
| 120 |
+
2.2769689559936523,
|
| 121 |
+
1.4722540378570557,
|
| 122 |
+
1.7136496305465698,
|
| 123 |
+
5.040107727050781,
|
| 124 |
+
2.9950344562530518,
|
| 125 |
+
4.272597789764404,
|
| 126 |
+
2.686481237411499,
|
| 127 |
+
2.7622010707855225,
|
| 128 |
+
3.0286059379577637,
|
| 129 |
+
2.0905165672302246,
|
| 130 |
+
2.4465155601501465,
|
| 131 |
+
2.035715341567993,
|
| 132 |
+
2.315361499786377,
|
| 133 |
+
1.4022663831710815,
|
| 134 |
+
3.5252881050109863,
|
| 135 |
+
2.6312732696533203,
|
| 136 |
+
3.7926905155181885,
|
| 137 |
+
1.7568085193634033,
|
| 138 |
+
1.798338532447815,
|
| 139 |
+
3.3464834690093994,
|
| 140 |
+
2.249241828918457,
|
| 141 |
+
2.2580409049987793,
|
| 142 |
+
4.0587310791015625,
|
| 143 |
+
2.206573247909546,
|
| 144 |
+
0.6222774982452393,
|
| 145 |
+
1.8168766498565674,
|
| 146 |
+
3.9070353507995605,
|
| 147 |
+
2.2064626216888428,
|
| 148 |
+
1.1733113527297974,
|
| 149 |
+
1.383936882019043,
|
| 150 |
+
2.251495599746704,
|
| 151 |
+
2.962852954864502,
|
| 152 |
+
2.8415677547454834,
|
| 153 |
+
1.3905932903289795,
|
| 154 |
+
3.165285587310791,
|
| 155 |
+
3.477961540222168,
|
| 156 |
+
1.7277320623397827,
|
| 157 |
+
2.4671823978424072,
|
| 158 |
+
5.010605812072754,
|
| 159 |
+
3.3257808685302734,
|
| 160 |
+
2.7084109783172607,
|
| 161 |
+
3.5739824771881104
|
| 162 |
+
]
|
lecture_4/verified_examples/peptune/report.json
ADDED
|
@@ -0,0 +1,350 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
{
|
| 2 |
+
"metrics": {
|
| 3 |
+
"valid_dna": true,
|
| 4 |
+
"unique_fraction": 1.0,
|
| 5 |
+
"mean_gc": 0.75,
|
| 6 |
+
"mean_atat_match": 0.25
|
| 7 |
+
},
|
| 8 |
+
"method": "peptune",
|
| 9 |
+
"data": "synthetic DNA; not biological validation",
|
| 10 |
+
"train_loss_first_20_mean": 4.838834500312805,
|
| 11 |
+
"train_loss_last_20_mean": 2.613932800292969,
|
| 12 |
+
"validation_loss_one_mc_draw": 2.5891432762145996,
|
| 13 |
+
"archive_scores": [
|
| 14 |
+
[
|
| 15 |
+
1.0,
|
| 16 |
+
0.0
|
| 17 |
+
],
|
| 18 |
+
[
|
| 19 |
+
0.5,
|
| 20 |
+
0.5
|
| 21 |
+
],
|
| 22 |
+
[
|
| 23 |
+
0.75,
|
| 24 |
+
0.25
|
| 25 |
+
]
|
| 26 |
+
],
|
| 27 |
+
"search_trace": [
|
| 28 |
+
{
|
| 29 |
+
"iteration": 0,
|
| 30 |
+
"archive_size": 1,
|
| 31 |
+
"score": [
|
| 32 |
+
0.25,
|
| 33 |
+
0.25
|
| 34 |
+
]
|
| 35 |
+
},
|
| 36 |
+
{
|
| 37 |
+
"iteration": 1,
|
| 38 |
+
"archive_size": 2,
|
| 39 |
+
"score": [
|
| 40 |
+
0.5,
|
| 41 |
+
0.0
|
| 42 |
+
]
|
| 43 |
+
},
|
| 44 |
+
{
|
| 45 |
+
"iteration": 2,
|
| 46 |
+
"archive_size": 2,
|
| 47 |
+
"score": [
|
| 48 |
+
0.0,
|
| 49 |
+
0.0
|
| 50 |
+
]
|
| 51 |
+
},
|
| 52 |
+
{
|
| 53 |
+
"iteration": 3,
|
| 54 |
+
"archive_size": 2,
|
| 55 |
+
"score": [
|
| 56 |
+
1.0,
|
| 57 |
+
0.0
|
| 58 |
+
]
|
| 59 |
+
},
|
| 60 |
+
{
|
| 61 |
+
"iteration": 4,
|
| 62 |
+
"archive_size": 2,
|
| 63 |
+
"score": [
|
| 64 |
+
0.5,
|
| 65 |
+
0.5
|
| 66 |
+
]
|
| 67 |
+
},
|
| 68 |
+
{
|
| 69 |
+
"iteration": 5,
|
| 70 |
+
"archive_size": 2,
|
| 71 |
+
"score": [
|
| 72 |
+
1.0,
|
| 73 |
+
0.0
|
| 74 |
+
]
|
| 75 |
+
},
|
| 76 |
+
{
|
| 77 |
+
"iteration": 6,
|
| 78 |
+
"archive_size": 2,
|
| 79 |
+
"score": [
|
| 80 |
+
0.5,
|
| 81 |
+
0.0
|
| 82 |
+
]
|
| 83 |
+
},
|
| 84 |
+
{
|
| 85 |
+
"iteration": 7,
|
| 86 |
+
"archive_size": 2,
|
| 87 |
+
"score": [
|
| 88 |
+
1.0,
|
| 89 |
+
0.0
|
| 90 |
+
]
|
| 91 |
+
},
|
| 92 |
+
{
|
| 93 |
+
"iteration": 8,
|
| 94 |
+
"archive_size": 2,
|
| 95 |
+
"score": [
|
| 96 |
+
0.0,
|
| 97 |
+
0.0
|
| 98 |
+
]
|
| 99 |
+
},
|
| 100 |
+
{
|
| 101 |
+
"iteration": 9,
|
| 102 |
+
"archive_size": 2,
|
| 103 |
+
"score": [
|
| 104 |
+
0.5,
|
| 105 |
+
0.5
|
| 106 |
+
]
|
| 107 |
+
},
|
| 108 |
+
{
|
| 109 |
+
"iteration": 10,
|
| 110 |
+
"archive_size": 2,
|
| 111 |
+
"score": [
|
| 112 |
+
0.5,
|
| 113 |
+
0.5
|
| 114 |
+
]
|
| 115 |
+
},
|
| 116 |
+
{
|
| 117 |
+
"iteration": 11,
|
| 118 |
+
"archive_size": 3,
|
| 119 |
+
"score": [
|
| 120 |
+
0.75,
|
| 121 |
+
0.25
|
| 122 |
+
]
|
| 123 |
+
},
|
| 124 |
+
{
|
| 125 |
+
"iteration": 12,
|
| 126 |
+
"archive_size": 3,
|
| 127 |
+
"score": [
|
| 128 |
+
0.5,
|
| 129 |
+
0.0
|
| 130 |
+
]
|
| 131 |
+
},
|
| 132 |
+
{
|
| 133 |
+
"iteration": 13,
|
| 134 |
+
"archive_size": 3,
|
| 135 |
+
"score": [
|
| 136 |
+
0.5,
|
| 137 |
+
0.5
|
| 138 |
+
]
|
| 139 |
+
},
|
| 140 |
+
{
|
| 141 |
+
"iteration": 14,
|
| 142 |
+
"archive_size": 3,
|
| 143 |
+
"score": [
|
| 144 |
+
0.5,
|
| 145 |
+
0.0
|
| 146 |
+
]
|
| 147 |
+
},
|
| 148 |
+
{
|
| 149 |
+
"iteration": 15,
|
| 150 |
+
"archive_size": 3,
|
| 151 |
+
"score": [
|
| 152 |
+
1.0,
|
| 153 |
+
0.0
|
| 154 |
+
]
|
| 155 |
+
},
|
| 156 |
+
{
|
| 157 |
+
"iteration": 16,
|
| 158 |
+
"archive_size": 3,
|
| 159 |
+
"score": [
|
| 160 |
+
0.5,
|
| 161 |
+
0.5
|
| 162 |
+
]
|
| 163 |
+
},
|
| 164 |
+
{
|
| 165 |
+
"iteration": 17,
|
| 166 |
+
"archive_size": 3,
|
| 167 |
+
"score": [
|
| 168 |
+
0.5,
|
| 169 |
+
0.0
|
| 170 |
+
]
|
| 171 |
+
},
|
| 172 |
+
{
|
| 173 |
+
"iteration": 18,
|
| 174 |
+
"archive_size": 3,
|
| 175 |
+
"score": [
|
| 176 |
+
0.5,
|
| 177 |
+
0.5
|
| 178 |
+
]
|
| 179 |
+
},
|
| 180 |
+
{
|
| 181 |
+
"iteration": 19,
|
| 182 |
+
"archive_size": 3,
|
| 183 |
+
"score": [
|
| 184 |
+
0.5,
|
| 185 |
+
0.5
|
| 186 |
+
]
|
| 187 |
+
},
|
| 188 |
+
{
|
| 189 |
+
"iteration": 20,
|
| 190 |
+
"archive_size": 3,
|
| 191 |
+
"score": [
|
| 192 |
+
0.0,
|
| 193 |
+
0.0
|
| 194 |
+
]
|
| 195 |
+
},
|
| 196 |
+
{
|
| 197 |
+
"iteration": 21,
|
| 198 |
+
"archive_size": 3,
|
| 199 |
+
"score": [
|
| 200 |
+
1.0,
|
| 201 |
+
0.0
|
| 202 |
+
]
|
| 203 |
+
},
|
| 204 |
+
{
|
| 205 |
+
"iteration": 22,
|
| 206 |
+
"archive_size": 3,
|
| 207 |
+
"score": [
|
| 208 |
+
0.5,
|
| 209 |
+
0.5
|
| 210 |
+
]
|
| 211 |
+
},
|
| 212 |
+
{
|
| 213 |
+
"iteration": 23,
|
| 214 |
+
"archive_size": 3,
|
| 215 |
+
"score": [
|
| 216 |
+
0.5,
|
| 217 |
+
0.5
|
| 218 |
+
]
|
| 219 |
+
},
|
| 220 |
+
{
|
| 221 |
+
"iteration": 24,
|
| 222 |
+
"archive_size": 3,
|
| 223 |
+
"score": [
|
| 224 |
+
1.0,
|
| 225 |
+
0.0
|
| 226 |
+
]
|
| 227 |
+
},
|
| 228 |
+
{
|
| 229 |
+
"iteration": 25,
|
| 230 |
+
"archive_size": 3,
|
| 231 |
+
"score": [
|
| 232 |
+
0.5,
|
| 233 |
+
0.0
|
| 234 |
+
]
|
| 235 |
+
},
|
| 236 |
+
{
|
| 237 |
+
"iteration": 26,
|
| 238 |
+
"archive_size": 3,
|
| 239 |
+
"score": [
|
| 240 |
+
0.5,
|
| 241 |
+
0.5
|
| 242 |
+
]
|
| 243 |
+
},
|
| 244 |
+
{
|
| 245 |
+
"iteration": 27,
|
| 246 |
+
"archive_size": 3,
|
| 247 |
+
"score": [
|
| 248 |
+
1.0,
|
| 249 |
+
0.0
|
| 250 |
+
]
|
| 251 |
+
},
|
| 252 |
+
{
|
| 253 |
+
"iteration": 28,
|
| 254 |
+
"archive_size": 3,
|
| 255 |
+
"score": [
|
| 256 |
+
0.5,
|
| 257 |
+
0.0
|
| 258 |
+
]
|
| 259 |
+
},
|
| 260 |
+
{
|
| 261 |
+
"iteration": 29,
|
| 262 |
+
"archive_size": 3,
|
| 263 |
+
"score": [
|
| 264 |
+
0.5,
|
| 265 |
+
0.5
|
| 266 |
+
]
|
| 267 |
+
},
|
| 268 |
+
{
|
| 269 |
+
"iteration": 30,
|
| 270 |
+
"archive_size": 3,
|
| 271 |
+
"score": [
|
| 272 |
+
0.5,
|
| 273 |
+
0.5
|
| 274 |
+
]
|
| 275 |
+
},
|
| 276 |
+
{
|
| 277 |
+
"iteration": 31,
|
| 278 |
+
"archive_size": 3,
|
| 279 |
+
"score": [
|
| 280 |
+
1.0,
|
| 281 |
+
0.0
|
| 282 |
+
]
|
| 283 |
+
},
|
| 284 |
+
{
|
| 285 |
+
"iteration": 32,
|
| 286 |
+
"archive_size": 3,
|
| 287 |
+
"score": [
|
| 288 |
+
0.5,
|
| 289 |
+
0.5
|
| 290 |
+
]
|
| 291 |
+
},
|
| 292 |
+
{
|
| 293 |
+
"iteration": 33,
|
| 294 |
+
"archive_size": 3,
|
| 295 |
+
"score": [
|
| 296 |
+
0.0,
|
| 297 |
+
0.0
|
| 298 |
+
]
|
| 299 |
+
},
|
| 300 |
+
{
|
| 301 |
+
"iteration": 34,
|
| 302 |
+
"archive_size": 3,
|
| 303 |
+
"score": [
|
| 304 |
+
0.5,
|
| 305 |
+
0.5
|
| 306 |
+
]
|
| 307 |
+
},
|
| 308 |
+
{
|
| 309 |
+
"iteration": 35,
|
| 310 |
+
"archive_size": 3,
|
| 311 |
+
"score": [
|
| 312 |
+
0.5,
|
| 313 |
+
0.5
|
| 314 |
+
]
|
| 315 |
+
},
|
| 316 |
+
{
|
| 317 |
+
"iteration": 36,
|
| 318 |
+
"archive_size": 3,
|
| 319 |
+
"score": [
|
| 320 |
+
1.0,
|
| 321 |
+
0.0
|
| 322 |
+
]
|
| 323 |
+
},
|
| 324 |
+
{
|
| 325 |
+
"iteration": 37,
|
| 326 |
+
"archive_size": 3,
|
| 327 |
+
"score": [
|
| 328 |
+
0.5,
|
| 329 |
+
0.5
|
| 330 |
+
]
|
| 331 |
+
},
|
| 332 |
+
{
|
| 333 |
+
"iteration": 38,
|
| 334 |
+
"archive_size": 3,
|
| 335 |
+
"score": [
|
| 336 |
+
0.5,
|
| 337 |
+
0.0
|
| 338 |
+
]
|
| 339 |
+
},
|
| 340 |
+
{
|
| 341 |
+
"iteration": 39,
|
| 342 |
+
"archive_size": 3,
|
| 343 |
+
"score": [
|
| 344 |
+
0.5,
|
| 345 |
+
0.5
|
| 346 |
+
]
|
| 347 |
+
}
|
| 348 |
+
],
|
| 349 |
+
"scope": "DNA MCTS mechanism; not peptide-model training or paper reproduction"
|
| 350 |
+
}
|
lecture_4/verified_examples/peptune/samples.txt
ADDED
|
@@ -0,0 +1,3 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
GCGC
|
| 2 |
+
ACGT
|
| 3 |
+
GCGT
|
lecture_4/verified_examples/udlm/config.json
ADDED
|
@@ -0,0 +1,18 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
{
|
| 2 |
+
"method": "udlm",
|
| 3 |
+
"mode": "train-sample",
|
| 4 |
+
"train_steps": 160,
|
| 5 |
+
"sample_steps": 40,
|
| 6 |
+
"classifier_steps": 300,
|
| 7 |
+
"search_steps": 100,
|
| 8 |
+
"batch_size": 16,
|
| 9 |
+
"samples": 8,
|
| 10 |
+
"length": 4,
|
| 11 |
+
"width": 32,
|
| 12 |
+
"block_size": 2,
|
| 13 |
+
"strength": 1.5,
|
| 14 |
+
"label": 1,
|
| 15 |
+
"seed": 7,
|
| 16 |
+
"data": null,
|
| 17 |
+
"out": "outputs/udlm"
|
| 18 |
+
}
|
lecture_4/verified_examples/udlm/losses.json
ADDED
|
@@ -0,0 +1,162 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
[
|
| 2 |
+
11.047435760498047,
|
| 3 |
+
3.638819932937622,
|
| 4 |
+
4.339240550994873,
|
| 5 |
+
6.615440368652344,
|
| 6 |
+
7.487863540649414,
|
| 7 |
+
3.715148448944092,
|
| 8 |
+
4.2508416175842285,
|
| 9 |
+
5.987113952636719,
|
| 10 |
+
4.574474334716797,
|
| 11 |
+
2.5587449073791504,
|
| 12 |
+
7.678818225860596,
|
| 13 |
+
10.757691383361816,
|
| 14 |
+
3.761280059814453,
|
| 15 |
+
6.221782684326172,
|
| 16 |
+
2.411395788192749,
|
| 17 |
+
3.9185919761657715,
|
| 18 |
+
2.432310104370117,
|
| 19 |
+
5.186282157897949,
|
| 20 |
+
2.995053768157959,
|
| 21 |
+
3.0481975078582764,
|
| 22 |
+
6.479434490203857,
|
| 23 |
+
4.919713497161865,
|
| 24 |
+
2.8192875385284424,
|
| 25 |
+
2.8869221210479736,
|
| 26 |
+
8.790794372558594,
|
| 27 |
+
2.300025463104248,
|
| 28 |
+
3.3890492916107178,
|
| 29 |
+
3.32928466796875,
|
| 30 |
+
3.703833818435669,
|
| 31 |
+
2.642963409423828,
|
| 32 |
+
3.1563267707824707,
|
| 33 |
+
2.6353256702423096,
|
| 34 |
+
3.5707004070281982,
|
| 35 |
+
2.8710131645202637,
|
| 36 |
+
2.8604514598846436,
|
| 37 |
+
3.6414551734924316,
|
| 38 |
+
4.774967193603516,
|
| 39 |
+
2.004885673522949,
|
| 40 |
+
8.153873443603516,
|
| 41 |
+
1.7916756868362427,
|
| 42 |
+
2.389755964279175,
|
| 43 |
+
4.308714389801025,
|
| 44 |
+
3.9405770301818848,
|
| 45 |
+
6.173405170440674,
|
| 46 |
+
2.9836504459381104,
|
| 47 |
+
2.4070658683776855,
|
| 48 |
+
4.4768242835998535,
|
| 49 |
+
4.716032981872559,
|
| 50 |
+
2.6472368240356445,
|
| 51 |
+
2.835066795349121,
|
| 52 |
+
2.540917158126831,
|
| 53 |
+
3.043062925338745,
|
| 54 |
+
5.332056045532227,
|
| 55 |
+
3.2398338317871094,
|
| 56 |
+
4.370035171508789,
|
| 57 |
+
2.7836556434631348,
|
| 58 |
+
4.69325590133667,
|
| 59 |
+
1.9069887399673462,
|
| 60 |
+
3.3850021362304688,
|
| 61 |
+
2.8949105739593506,
|
| 62 |
+
6.395257472991943,
|
| 63 |
+
6.112308979034424,
|
| 64 |
+
2.172653913497925,
|
| 65 |
+
3.731419324874878,
|
| 66 |
+
3.2633914947509766,
|
| 67 |
+
2.7229864597320557,
|
| 68 |
+
3.2528135776519775,
|
| 69 |
+
6.598641395568848,
|
| 70 |
+
2.3790218830108643,
|
| 71 |
+
1.4963032007217407,
|
| 72 |
+
2.819007158279419,
|
| 73 |
+
2.80535888671875,
|
| 74 |
+
2.719237804412842,
|
| 75 |
+
2.0492019653320312,
|
| 76 |
+
3.2515711784362793,
|
| 77 |
+
1.8298454284667969,
|
| 78 |
+
2.3929498195648193,
|
| 79 |
+
2.056596517562866,
|
| 80 |
+
15.28537654876709,
|
| 81 |
+
3.128511667251587,
|
| 82 |
+
2.0100417137145996,
|
| 83 |
+
1.5984283685684204,
|
| 84 |
+
1.9779601097106934,
|
| 85 |
+
3.192647695541382,
|
| 86 |
+
2.118055582046509,
|
| 87 |
+
2.90411114692688,
|
| 88 |
+
3.8938608169555664,
|
| 89 |
+
2.1555304527282715,
|
| 90 |
+
2.7105512619018555,
|
| 91 |
+
3.3428311347961426,
|
| 92 |
+
1.0114843845367432,
|
| 93 |
+
3.9458742141723633,
|
| 94 |
+
5.690505027770996,
|
| 95 |
+
3.6527702808380127,
|
| 96 |
+
1.1432782411575317,
|
| 97 |
+
5.013024806976318,
|
| 98 |
+
2.777940273284912,
|
| 99 |
+
3.2634408473968506,
|
| 100 |
+
4.932656288146973,
|
| 101 |
+
4.210480690002441,
|
| 102 |
+
2.6948583126068115,
|
| 103 |
+
2.285061836242676,
|
| 104 |
+
3.583308696746826,
|
| 105 |
+
2.962766647338867,
|
| 106 |
+
2.7541191577911377,
|
| 107 |
+
2.287432909011841,
|
| 108 |
+
3.1161656379699707,
|
| 109 |
+
2.5217714309692383,
|
| 110 |
+
3.1153383255004883,
|
| 111 |
+
1.9834219217300415,
|
| 112 |
+
2.0762336254119873,
|
| 113 |
+
2.764826536178589,
|
| 114 |
+
2.7770273685455322,
|
| 115 |
+
1.6031702756881714,
|
| 116 |
+
4.0543365478515625,
|
| 117 |
+
2.547313690185547,
|
| 118 |
+
2.8786070346832275,
|
| 119 |
+
2.6822168827056885,
|
| 120 |
+
6.25325345993042,
|
| 121 |
+
2.590442419052124,
|
| 122 |
+
2.122274398803711,
|
| 123 |
+
1.6236686706542969,
|
| 124 |
+
1.8864619731903076,
|
| 125 |
+
2.4540152549743652,
|
| 126 |
+
3.4854989051818848,
|
| 127 |
+
3.156172513961792,
|
| 128 |
+
4.207565784454346,
|
| 129 |
+
3.6062185764312744,
|
| 130 |
+
2.4027581214904785,
|
| 131 |
+
1.1646852493286133,
|
| 132 |
+
4.980432510375977,
|
| 133 |
+
6.687517166137695,
|
| 134 |
+
1.4371603727340698,
|
| 135 |
+
1.6977627277374268,
|
| 136 |
+
3.9442780017852783,
|
| 137 |
+
2.9689416885375977,
|
| 138 |
+
2.100611686706543,
|
| 139 |
+
2.377840995788574,
|
| 140 |
+
3.0913102626800537,
|
| 141 |
+
6.303923606872559,
|
| 142 |
+
2.232814073562622,
|
| 143 |
+
1.8777183294296265,
|
| 144 |
+
1.323475956916809,
|
| 145 |
+
3.0791330337524414,
|
| 146 |
+
1.765347957611084,
|
| 147 |
+
1.629030466079712,
|
| 148 |
+
2.6598291397094727,
|
| 149 |
+
2.1885077953338623,
|
| 150 |
+
2.656402587890625,
|
| 151 |
+
2.2888011932373047,
|
| 152 |
+
2.212771415710449,
|
| 153 |
+
1.9104533195495605,
|
| 154 |
+
4.000855445861816,
|
| 155 |
+
2.3612029552459717,
|
| 156 |
+
1.7967439889907837,
|
| 157 |
+
1.9667776823043823,
|
| 158 |
+
2.8789355754852295,
|
| 159 |
+
1.4103233814239502,
|
| 160 |
+
1.952869176864624,
|
| 161 |
+
3.3791284561157227
|
| 162 |
+
]
|
lecture_4/verified_examples/udlm/report.json
ADDED
|
@@ -0,0 +1,14 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
{
|
| 2 |
+
"metrics": {
|
| 3 |
+
"valid_dna": true,
|
| 4 |
+
"unique_fraction": 0.875,
|
| 5 |
+
"mean_gc": 0.625,
|
| 6 |
+
"mean_atat_match": 0.25
|
| 7 |
+
},
|
| 8 |
+
"method": "udlm",
|
| 9 |
+
"data": "synthetic DNA; not biological validation",
|
| 10 |
+
"train_loss_first_20_mean": 5.131326353549957,
|
| 11 |
+
"train_loss_last_20_mean": 2.2785560965538023,
|
| 12 |
+
"validation_loss_one_mc_draw": 2.245309829711914,
|
| 13 |
+
"endpoint_approximation": "t in [0.02, 0.98]; stop at residual noise 0.02, without a posterior interpretation of the UDLM parameter vector"
|
| 14 |
+
}
|