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

Add Lecture 4 discrete diffusion and Lecture 5 discrete flow matching

Browse files

Add 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
Files changed (50) hide show
  1. README.md +52 -1
  2. lecture_4/.gitignore +5 -0
  3. lecture_4/README.md +85 -0
  4. lecture_4/SLIDE_CODE_MAP.md +21 -0
  5. lecture_4/common.py +134 -0
  6. lecture_4/data/README.md +1 -0
  7. lecture_4/data/dna_train.tsv +257 -0
  8. lecture_4/diffusion.py +223 -0
  9. lecture_4/examples/block.py +7 -0
  10. lecture_4/examples/cfg.py +7 -0
  11. lecture_4/examples/classifier_exact.py +7 -0
  12. lecture_4/examples/classifier_gradient.py +7 -0
  13. lecture_4/examples/mdlm.py +7 -0
  14. lecture_4/examples/peptune.py +7 -0
  15. lecture_4/examples/udlm.py +7 -0
  16. lecture_4/lecture_core.py +136 -0
  17. lecture_4/numerical_examples.py +24 -0
  18. lecture_4/requirements.txt +3 -0
  19. lecture_4/run.py +127 -0
  20. lecture_4/run_all.py +15 -0
  21. lecture_4/tests/test_mathematics.py +50 -0
  22. lecture_4/verified_examples/README.md +15 -0
  23. lecture_4/verified_examples/block/config.json +18 -0
  24. lecture_4/verified_examples/block/losses.json +162 -0
  25. lecture_4/verified_examples/block/report.json +13 -0
  26. lecture_4/verified_examples/block/samples.txt +8 -0
  27. lecture_4/verified_examples/cfg/config.json +18 -0
  28. lecture_4/verified_examples/cfg/losses.json +162 -0
  29. lecture_4/verified_examples/cfg/report.json +13 -0
  30. lecture_4/verified_examples/cfg/samples.txt +8 -0
  31. lecture_4/verified_examples/classifier-exact/config.json +18 -0
  32. lecture_4/verified_examples/classifier-exact/losses.json +162 -0
  33. lecture_4/verified_examples/classifier-exact/report.json +15 -0
  34. lecture_4/verified_examples/classifier-exact/samples.txt +8 -0
  35. lecture_4/verified_examples/classifier-gradient/config.json +18 -0
  36. lecture_4/verified_examples/classifier-gradient/losses.json +162 -0
  37. lecture_4/verified_examples/classifier-gradient/report.json +15 -0
  38. lecture_4/verified_examples/classifier-gradient/samples.txt +8 -0
  39. lecture_4/verified_examples/environment.json +7 -0
  40. lecture_4/verified_examples/mdlm/config.json +18 -0
  41. lecture_4/verified_examples/mdlm/losses.json +162 -0
  42. lecture_4/verified_examples/mdlm/report.json +13 -0
  43. lecture_4/verified_examples/mdlm/samples.txt +8 -0
  44. lecture_4/verified_examples/peptune/config.json +18 -0
  45. lecture_4/verified_examples/peptune/losses.json +162 -0
  46. lecture_4/verified_examples/peptune/report.json +350 -0
  47. lecture_4/verified_examples/peptune/samples.txt +3 -0
  48. lecture_4/verified_examples/udlm/config.json +18 -0
  49. lecture_4/verified_examples/udlm/losses.json +162 -0
  50. 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
+ }