diff --git a/README.md b/README.md index 03b05b2d504b8b67e4b101ec2cca9d53c8b572b1..aca70a2e5e0aa9096e8ce379564d959b8a6adaa4 100644 --- a/README.md +++ b/README.md @@ -32,13 +32,16 @@ and implementation notes; additional directories will accompany later lectures. | --- | --- | --- | | 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) | | 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) | +| 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) | +| 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) | ## Installation Use Python 3.11, or another compatible Python version at least 3.10, in a new virtual environment. The shared `requirements.txt` pins PyTorch 2.9.1, TorchVision 0.24.1, and Transformers 4.57.6. Lecture 2 uses PyTorch and -TorchVision; Lecture 3 also uses Transformers. +TorchVision; Lecture 3 also uses Transformers. Lectures 4 and 5 use PyTorch, +NumPy, and SciPy. Their lecture folders also provide minimal requirements. ```bash git clone https://huggingface.co/ChatterjeeLab/CIS6270 @@ -102,6 +105,45 @@ The [lecture guide](lecture_3/README.md) describes the data format, training and sampling settings, property calculations, normalization, and residue-count constraint, with commands for using a custom dataset. +## Lecture 4 - Discrete diffusion + +Train small DNA denoisers and generate sequences with MDLM, UDLM, block diffusion, +classifier-free guidance, exact and gradient-based classifier guidance, and a +PepTune-style search. The [guide](lecture_4/README.md) includes each method's +command and mathematical assumptions. The [code map](lecture_4/SLIDE_CODE_MAP.md) +links the slide walkthroughs to their functions. + +From the repository root, run the complete MDLM example. + +```bash +python lecture_4/run.py --method mdlm --data lecture_4/data/dna_train.tsv --out lecture_4/outputs/mdlm +python lecture_4/run.py --method mdlm --mode sample --out lecture_4/outputs/mdlm +``` + +The script trains, saves a checkpoint, and writes generated DNA and loss logs. +The bundled data are synthetic, and the guidance objectives are explicit toy +properties. No pretrained model or external dataset is required. + +## Lecture 5 - Discrete flow matching + +Start with Gat et al.'s discrete flow matching, then run Dirichlet, Fisher, +Gumbel-Softmax, rectified flow, ReDi, MOG-DFM, and AReUReDi examples. Each method +has a complete training and generation command in the [guide](lecture_5/README.md). + +```bash +python lecture_5/run.py --method gat --data lecture_5/data/dna_train.tsv --out lecture_5/outputs/gat +python lecture_5/run.py --method gat --mode sample --out lecture_5/outputs/gat +``` + +The current [Discrete Generation slide deck](https://docs.google.com/presentation/d/1qcPdF4SYTq4VdTN_5w9GjviLRPr2HFtbuYMpbfLrc_M/edit) +contains both lectures. Diffusion code through PepTune belongs to Lecture 4. +Classic DFM and the subsequent flow methods belong to Lecture 5. The code maps +use stable slide links, so added transition slides do not break the mapping. + +Both folders include numerical examples, mathematical tests, and saved results +from seeded CPU runs. The guides explain finite endpoint approximations and +classroom simplifications for each method. + ## Repository organization | Location | Contents | @@ -109,6 +151,8 @@ constraint, with commands for using a custom dataset. | Repository root | Course index, installation requirements, and license | | [`lecture_2/`](https://huggingface.co/ChatterjeeLab/CIS6270/tree/main/lecture_2) | One MNIST flow-matching script, guide, trained checkpoint, and selected example images | | [`lecture_3/`](https://huggingface.co/ChatterjeeLab/CIS6270/tree/main/lecture_3) | ESM-2 flow and diffusion guidance scripts, sequence data, guide, and mathematical notes | +| [`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 | +| [`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 | | [`tests/`](https://huggingface.co/ChatterjeeLab/CIS6270/tree/main/tests) | Offline checks for the Lecture 3 examples | Installation instructions and the lecture index are maintained at the @@ -119,11 +163,18 @@ references accompany the corresponding code. ```bash python -m unittest discover -s tests -v +python -m unittest discover -s lecture_4/tests -v +python -m unittest discover -s lecture_5/tests -v ``` The Lecture 3 unit tests cover property annotations, scalarization weights, reward gradients, DDPM schedule indexing, and constrained decoding. +Lecture 4 tests check reverse KL losses and guidance calculations. Lecture 5 +tests check the master equation, Fisher geometry, Gumbel path derivatives, and +MH detailed balance. Run a short end-to-end check of every method from its +lecture folder with `python run_all.py --quick`. + ## License The repository code is distributed under the diff --git a/lecture_4/.gitignore b/lecture_4/.gitignore new file mode 100644 index 0000000000000000000000000000000000000000..876670e6a75ee4b613025a193d0791312ac8679b --- /dev/null +++ b/lecture_4/.gitignore @@ -0,0 +1,5 @@ +__pycache__/ +.venv/ +outputs/ +*.pt +.pytest_cache/ diff --git a/lecture_4/README.md b/lecture_4/README.md new file mode 100644 index 0000000000000000000000000000000000000000..a678db55b9163c172a434d5f4184994246af4b0e --- /dev/null +++ b/lecture_4/README.md @@ -0,0 +1,85 @@ +# CIS 6270 - Lecture 4 - Discrete Diffusion + +Course hub: [ChatterjeeLab/CIS6270 on Hugging Face](https://huggingface.co/ChatterjeeLab/CIS6270). + +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). + +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. + +## Start with the complete MDLM example + +Python 3.11 or later, CPU. No downloaded data or pretrained checkpoint is required. + +From the course repository root, enter this lecture folder. + +```bash +cd lecture_4 +python -m venv .venv +source .venv/bin/activate +pip install -r requirements.txt +python run.py --method mdlm --data data/dna_train.tsv --out outputs/mdlm +``` + +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`. + +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`. + +```bash +# Continue with the saved weights for generation only. +python run.py --method mdlm --mode sample --out outputs/mdlm +# An entire small experiment using the four-base vocabulary. +python examples/mdlm.py --length 4 --train-steps 160 --samples 8 +# Run every method, or a short CPU integration check. +python run_all.py +python run_all.py --quick +python -m unittest discover -s tests -v +python numerical_examples.py +``` + +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. + +## One shared backbone; small method changes + +`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. + +| Example | Training target | Generation procedure | +| --- | --- | --- | +| `mdlm` | Weighted cross-entropy at masked positions | Schedule-based reveals; preserve visible bases | +| `udlm` | Reverse-rate KL including the staying term | Adaptive reverse Euler; visible bases may be revised | +| `block` | Conditional masked loss on a sampled block, scaled by block count | Finish each block before extending the sequence | +| `cfg` / `classifier-free` | Label-conditioned MDLM with 15% label dropout | Geometric conditional/unconditional prediction blend | +| `classifier-exact` | MDLM plus a classifier trained on noisy sequences | Evaluate every replacement's log-value change and multiply rates | +| `classifier-gradient` | Same denoiser and noisy classifier | Approximate replacement log-value changes with one-hot gradients | +| `peptune` | MDLM backbone | MCTS selection, expansion, completion, Pareto rewards, and reward backup | + +```bash +python examples/udlm.py --out outputs/udlm +python examples/block.py --block-size 2 +python examples/cfg.py --label 1 --strength 2 +python examples/classifier_exact.py --classifier-steps 300 +python examples/classifier_gradient.py --classifier-steps 300 +python examples/peptune.py --search-steps 100 +``` + +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. + +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. + +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. + +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. + +## Read the math next to the implementation + +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. + +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. + +## Sources + +- [MDLM - Simple and Effective Masked Diffusion Language Models](https://arxiv.org/abs/2406.07524) +- [UDLM - Simple Guidance Mechanisms for Discrete Diffusion Models](https://arxiv.org/abs/2412.10193) +- [Block Diffusion](https://arxiv.org/abs/2503.09573) +- [PepTune](https://arxiv.org/abs/2412.17780) + +The source papers define the full research methods and experiments. Comments identify the classroom simplifications. diff --git a/lecture_4/SLIDE_CODE_MAP.md b/lecture_4/SLIDE_CODE_MAP.md new file mode 100644 index 0000000000000000000000000000000000000000..d6390182cc34b4fe2368747cffdd18f148a200b0 --- /dev/null +++ b/lecture_4/SLIDE_CODE_MAP.md @@ -0,0 +1,21 @@ +# Code walkthrough map + +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. + +| Slide unit | Code | Location | +| --- | --- | --- | +| [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` | +| [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` | +| [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` | +| [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` | +| [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` | +| [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` | +| [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` | +| [Code for MDLM generation](https://docs.google.com/presentation/d/1qcPdF4SYTq4VdTN_5w9GjviLRPr2HFtbuYMpbfLrc_M/edit#slide=id.cis4r2_u148_s13) | `mdlm_sample` | `lecture_core.py` | +| [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` | +| [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` | +| [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` | +| [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` | +| [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` | +| [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` | +| [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` | diff --git a/lecture_4/common.py b/lecture_4/common.py new file mode 100644 index 0000000000000000000000000000000000000000..277eca7d7ad53417be3e3ad2bdba319714c9386c --- /dev/null +++ b/lecture_4/common.py @@ -0,0 +1,134 @@ +"""CPU teaching utilities shared by the two lecture folders.""" +import csv +import json +import random +from pathlib import Path +import numpy as np +import torch +from torch import nn +from lecture_core import DNA, K, MASK, encode + + +def seed_all(seed): + random.seed(seed) + np.random.seed(seed) + torch.manual_seed(seed) + torch.set_num_threads(1) + + +def decode(tokens): + alphabet = 'ACGTm' + return [''.join(alphabet[int(i)] for i in row) for row in tokens] + + +def make_data(count=512, length=8, seed=7): + """Four synthetic motif families, with independent 8% base mutations.""" + rng = np.random.default_rng(seed) + motifs = ['ACGT', 'CGTA', 'TATA', 'GCGC'] + strings, labels = [], [] + for _ in range(count): + motif = motifs[int(rng.integers(4))] + seq = list((motif * ((length + 3) // 4))[:length]) + for j in range(length): + if rng.random() < .08: + seq[j] = 'ACGT'[int(rng.integers(4))] + strings.append(''.join(seq)) + labels.append(int(sum(x in 'GC' for x in seq) / length >= .6)) + return encode(strings), torch.tensor(labels) + + +def load_data(path, length): + if path is None: + return make_data(length=length) + with open(path) as f: + rows = list(csv.DictReader(f, delimiter='\t')) + strings = [r['sequence'].strip().upper() for r in rows] + if not strings or any(len(s) != length or set(s) - set('ACGT') for s in strings): + raise ValueError('All DNA sequences must contain only A/C/G/T and have --length bases.') + labels = [int(r.get('label', sum(c in 'GC' for c in s) / length >= .6)) + for r, s in zip(rows, strings)] + if set(labels) - {0, 1}: + raise ValueError('Labels must be 0 or 1.') + return encode(strings), torch.tensor(labels) + + +class ConditionalDNA(DNA): + """The slide network plus a label embedding; label 2 means unconditional.""" + def __init__(self, width=32, max_len=64): + super().__init__(width, max_len) + self.condition = nn.Embedding(3, width) + + def forward(self, z, t=None, label=None): + h = self.token(z) if z.ndim == 2 else self.soft(z) + pos = torch.arange(z.shape[1], device=z.device) + h = h + self.position(pos)[None] + if t is not None: + h = h + self.time(t[:, None])[:, None] + if label is None: + label = torch.full((len(z),), 2, dtype=torch.long, device=z.device) + h = h + self.condition(label)[:, None] + return self.output(self.context(h)) + + +def objectives(tokens): + """Both toy objectives are maximized: GC fraction and ATAT agreement.""" + onehot = torch.nn.functional.one_hot(tokens.long(), 4).float() + return soft_objectives(onehot) + + +def soft_objectives(z): + gc = (z[..., 1] + z[..., 2]).mean(-1) + motif = torch.tensor([0, 3], device=z.device).repeat((z.shape[1] + 1) // 2)[:z.shape[1]] + match = z.gather(-1, motif[None, :, None].expand(z.shape[0], -1, 1)).squeeze(-1).mean(-1) + return torch.stack((gc, match), -1) + + +def metrics(tokens): + strings = decode(tokens) + score = objectives(tokens).mean(0) + return dict(valid_dna=all(set(s) <= set('ACGT') for s in strings), + unique_fraction=len(set(strings)) / len(strings), + mean_gc=float(score[0]), mean_atat_match=float(score[1])) + + +def save_run(out, config, losses, samples, extra=None): + out = Path(out) + out.mkdir(parents=True, exist_ok=True) + (out / ('sample_config.json' if config.get('mode') == 'sample' else 'config.json')).write_text(json.dumps(config, indent=2)) + if losses or not (out / 'losses.json').exists(): + (out / 'losses.json').write_text(json.dumps(losses, indent=2)) + (out / 'samples.txt').write_text('\n'.join(decode(samples)) + '\n') + report = {'metrics': metrics(samples), **(extra or {})} + (out / 'report.json').write_text(json.dumps(report, indent=2)) + print(json.dumps({'output': str(out), **report}, indent=2)) + return report + + +def optimize(model, loss_fn, data, steps, batch_size=32, lr=.002): + optimizer = torch.optim.AdamW(model.parameters(), lr=lr) + losses = [] + model.train() + for step in range(steps): + idx = torch.randint(len(data), (batch_size,)) + optimizer.zero_grad(set_to_none=True) + loss = loss_fn(model, data[idx], idx) + if not torch.isfinite(loss): + raise FloatingPointError(f'Nonfinite loss at step {step}') + loss.backward() + torch.nn.utils.clip_grad_norm_(model.parameters(), 1.) + optimizer.step() + losses.append(float(loss.detach())) + model.eval() + return losses + + +def project_simplex(z, eps=1e-6): + """Euclidean simplex projection, followed by a small interior floor.""" + sorted_z = z.sort(-1, descending=True).values + cssv = sorted_z.cumsum(-1) - 1. + k = torch.arange(1, z.shape[-1] + 1, dtype=z.dtype, device=z.device) + active = sorted_z - cssv / k > 0 + rho = active.sum(-1, keepdim=True).clamp_min(1) + theta = cssv.gather(-1, rho - 1) / rho + result = (z - theta).clamp_min(eps) + return result / result.sum(-1, keepdim=True) diff --git a/lecture_4/data/README.md b/lecture_4/data/README.md new file mode 100644 index 0000000000000000000000000000000000000000..34c37c778eaa416a7cf9568f6e382f5771601b1c --- /dev/null +++ b/lecture_4/data/README.md @@ -0,0 +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. diff --git a/lecture_4/data/dna_train.tsv b/lecture_4/data/dna_train.tsv new file mode 100644 index 0000000000000000000000000000000000000000..e90e1227dd16f4e0f12e30bf40d10c531030dcc2 --- /dev/null +++ b/lecture_4/data/dna_train.tsv @@ -0,0 +1,257 @@ +sequence label +TAAACTTA 0 +ACGCTCCC 1 +CGTACGTA 0 +ACGTACGT 0 +CGTACGTA 0 +ACGTAAGT 0 +TATAAATA 0 +ACGTACGT 0 +GCGCGCGC 1 +GCGCGCGC 1 +TATATACA 0 +ACGTACGT 0 +ACGTACGT 0 +ACGCGCGC 1 +TCTATATA 0 +GCGCGCGC 1 +CGTACGTA 0 +ACGTACGT 0 +CGTACGTA 0 +GAGCGCGC 1 +TATATATA 0 +ACGTACGT 0 +GCGCGCGC 1 +ACGTACGT 0 +GCGCACGC 1 +CGTACGTA 0 +GGTACGTA 0 +CTTACGTA 0 +ACGTACGT 0 +AATATATA 0 +CGTACGTA 0 +TATATATA 0 +GCGCGCGC 1 +GCGCGCGC 1 +GCGCTCGC 1 +TCTATCTA 0 +GCGCGCGC 1 +CAAATATA 0 +TCGCGCGC 1 +ACGTACGT 0 +AAGTACGT 0 +GCGCGCGC 1 +CGCGCTTA 1 +TACATATA 0 +TATATATA 0 +CGTACGTA 0 +ACGTACGT 0 +GCGTGCGC 1 +ACGTACCT 0 +ACATACAT 0 +ATGGACGT 0 +GCGCGCGC 1 +TATAAATA 0 +CGTACGTA 0 +ATTTACGA 0 +CGTACGTA 0 +GCACGCGC 1 +GCGCGCGC 1 +TTTATATA 0 +ACGTACGT 0 +GCGCGCGT 1 +GCGCGCGC 1 +TATATATA 0 +CTTACGTA 0 +CGTACGTA 0 +GCGCGCGC 1 +CGTACGTA 0 +TGTTCGTA 0 +ACGTACCT 0 +ACGTACGC 1 +GCGCGCGC 1 +CTTATGTA 0 +CGTGCGTG 1 +ACGTACGT 0 +GCGCCCAC 1 +CGTACGTC 1 +TATATATA 0 +GGGCGCGC 1 +TATATATA 0 +TCGCGCGC 1 +ACGTACGT 0 +TATATCTA 0 +GCGCGCGC 1 +TATATATC 0 +GCACGCGC 1 +ACGTACGT 0 +CGTACGTA 0 +TATATATA 0 +CGTACGTA 0 +GGGCGCGC 1 +TAGATATA 0 +CGTACATA 0 +ACGTACGT 0 +TAGAGATA 0 +ACGTACGT 0 +ACGTACGT 0 +TATATATA 0 +ACTTACGT 0 +CGTACCTA 0 +CGTACGTA 0 +GCGCGCGC 1 +GCTCGCGC 1 +GGGCGCCC 1 +TATATATA 0 +ACGTACGT 0 +ACGTACGT 0 +GCGCCCGC 1 +TATATATA 0 +ACGTACGT 0 +TATAGATA 0 +TATATATT 0 +CAGACGTA 0 +CGTACGTA 0 +CGTAGGTA 0 +GCACGACC 1 +CGTACGAA 0 +GTGCGCGC 1 +CGTACGTA 0 +TATATATA 0 +CGTACGTA 0 +GCACGCGC 1 +ACATACGA 0 +CGTAGGTA 0 +ACGTACAT 0 +TATATATA 0 +GATATATA 0 +ACGTACGT 0 +TCTAGATA 0 +CGTACGTA 0 +CGTACGTA 0 +TATATATA 0 +GCGCGCGC 1 +ACGTACGT 0 +TATATATA 0 +CGTACGTA 0 +ACGTACGT 0 +ACGTTCGT 0 +TATATATA 0 +GCGCGCGC 1 +ACGTACGT 0 +TATATATA 0 +ACGCACGT 1 +ACGTGCGT 1 +CGTACGTA 0 +GTGCGCGC 1 +GCGCGCGC 1 +CGTACGTA 0 +GCGCGCGC 1 +TATATATA 0 +GCGCGCGC 1 +CGTACGTA 0 +CGTACGTA 0 +ACGTACGT 0 +GCGCGCGC 1 +CGTACGTA 0 +GCGCGGGA 1 +TATATATA 0 +ACGTACGT 0 +ACGTACGT 0 +GCGCGCGT 1 +ACTTACGT 0 +CGTACGTA 0 +ACGTACGT 0 +TAAATAAA 0 +GCGTGCGC 1 +GCGCGCAC 1 +ACGTACGT 0 +TATATATA 0 +TAAATATA 0 +CGTACGTA 0 +ACGTACGT 0 +TGTACATA 0 +CGTACGGA 1 +CGTCCGTA 1 +GCGCGCGC 1 +ACGTACGT 0 +TATATATA 0 +CGTACGTA 0 +TGTACGTA 0 +ACGTACGT 0 +ACGTACGT 0 +GCGTATGT 0 +ACGTACGT 0 +CGGAAGTC 1 +TATATATA 0 +TATATATA 0 +CGTGCGTA 1 +GCGCGCGC 1 +CGTACGTA 0 +ACGTACGT 0 +ACGTACGT 0 +GCGCGCGC 1 +ACGTACGT 0 +ACGTACGT 0 +ACGTACGT 0 +CGTACGTA 0 +GCGCGCGC 1 +ACGTAGGT 0 +TATATATA 0 +ACATACGT 0 +CCGTACGT 1 +AAGTACGT 0 +CGTACGTA 0 +GCACGCGC 1 +TACATATA 0 +TATATATA 0 +GCGCGCGC 1 +ACGTACGT 0 +GCGTACGT 1 +TATATATA 0 +TATATATA 0 +GCGCGCGC 1 +CGTGCTTA 0 +ACGTACGT 0 +TATATATA 0 +CGTACGTA 0 +TATATATA 0 +GTGCGCGC 1 +CGTACGTA 0 +GCGTACGT 1 +GCGCGCGC 1 +TATATATA 0 +TATATGTA 0 +TATATATA 0 +CATACTAA 0 +ACGAACGT 0 +CGTACGTA 0 +GCGAGCGC 1 +TATATATA 0 +ACGAACCT 0 +GCGCGCGC 1 +CGTCCGTG 1 +GCGCGAGC 1 +CGTACCTA 0 +CGTACGTA 0 +GCGCGCGC 1 +TATATATA 0 +ACGTACGT 0 +ACGTAGGT 0 +GGTACGTA 0 +CGTACGTA 0 +GCGCGCGA 1 +CGTACGTA 0 +CGTACGGA 1 +ACGTACGT 0 +GCGCGCGC 1 +GCGCGCGC 1 +TATATATA 0 +ACGTACGT 0 +GCGCGCGC 1 +ACGTACGT 0 +ATGGACGT 0 +CGTAAGTA 0 +TATATATA 0 +TATATATA 0 +TATATATA 0 diff --git a/lecture_4/diffusion.py b/lecture_4/diffusion.py new file mode 100644 index 0000000000000000000000000000000000000000..8e5fce05a739f0aa78fe1cf388e180de61011f93 --- /dev/null +++ b/lecture_4/diffusion.py @@ -0,0 +1,223 @@ +"""Complete samplers and guidance extensions around the lecture code modules.""" +import math +import numpy as np +import torch +from torch import nn +import torch.nn.functional as F +from lecture_core import (K, MASK, draw, token_ce, udlm_rates, + rate_step, geometric_cfg, pareto_filter) +from common import objectives + + +@torch.no_grad() +def uniform_sample(model, batch, length, steps=100, epsilon=.02): + """Adaptive reverse Euler on the same truncated interval as udlm_loss. + + Uniform initialization at t=1-epsilon and stopping at t=epsilon + are endpoint approximations. The returned sequence retains residual + corruption; UDLM output probabilities are not treated as a clean posterior. + """ + z = torch.randint(K, (batch, length)) + time = 1. - epsilon + count = 0 + while time > epsilon + 1e-8: + t = torch.full((batch,), time) + rate = udlm_rates(model(z, t).softmax(-1), z, t) + max_exit = float(rate.sum(-1).max()) + h = min(1. / steps, time - epsilon, .5 / max(max_exit, 1e-8)) + z = rate_step(z, rate, h) + time -= h + count += 1 + if count > 100000: + raise RuntimeError('Adaptive sampler failed to advance.') + return z + + +@torch.no_grad() +def block_sample(model, batch, length, size=2, steps=20): + """Generate one block completely before appending the next one.""" + prefix = torch.empty((batch, 0), dtype=torch.long) + for start in range(0, length, size): + width = min(size, length - start) + block = torch.full((batch, width), MASK) + grid = torch.linspace(1., 0., steps + 1) + for t, s in zip(grid[:-1], grid[1:]): + context = torch.cat((prefix, block), 1) + p = model(context)[:, start:].softmax(-1) + candidate = draw(p) + reveal = (block == MASK) & (torch.rand(block.shape) < (t-s)/t) + block = torch.where(reveal, candidate, block) + prefix = torch.cat((prefix, block), 1) + return prefix + + +def conditional_loss(model, clean, labels, drop=.15): + t = torch.rand(len(clean)).clamp_min(1e-4) + mask = torch.rand(clean.shape) < t[:, None] + context = clean.masked_fill(mask, MASK) + condition = labels.clone() + condition[torch.rand(len(clean)) < drop] = 2 + ce = token_ce(model(context, label=condition), clean) + return (ce * mask / t[:, None]).sum(1).mean() + + +@torch.no_grad() +def cfg_sample(model, batch, length, strength=2., label=1, steps=20): + z = torch.full((batch, length), MASK) + grid = torch.linspace(1., 0., steps + 1) + labels = torch.full((batch,), label, dtype=torch.long) + for t, s in zip(grid[:-1], grid[1:]): + uncond = model(z).softmax(-1) + cond = model(z, label=labels).softmax(-1) + probability = geometric_cfg(uncond, cond, strength) + candidate = draw(probability) + reveal = (z == MASK) & (torch.rand(z.shape) < (t-s)/t) + z = torch.where(reveal, candidate, z) + return z + + +class NoisyClassifier(nn.Module): + """Predict high-GC class from a masked sequence and its noise level.""" + def __init__(self, length): + super().__init__() + self.net = nn.Sequential(nn.Linear(length * 5 + 1, 64), + nn.SiLU(), nn.Linear(64, 1)) + + def forward(self, z, t): + features = F.one_hot(z, 5).float() if z.ndim == 2 else z + return self.net(torch.cat((features.flatten(1), t[:, None]), 1)).squeeze(-1) + + +def fit_classifier(model, data, labels, steps=200, batch=32): + opt = torch.optim.Adam(model.parameters(), lr=.003) + losses = [] + for _ in range(steps): + idx = torch.randint(len(data), (batch,)) + t = torch.rand(batch) + noisy = data[idx].masked_fill(torch.rand(batch, data.shape[1]) < t[:, None], MASK) + loss = F.binary_cross_entropy_with_logits(model(noisy, t), labels[idx].float()) + opt.zero_grad(set_to_none=True) + loss.backward() + opt.step() + losses.append(float(loss.detach())) + model.eval() + return losses + + +def log_success(classifier, z, t, label): + logit = classifier(z, t) + return F.logsigmoid(logit if label == 1 else -logit) + + +def guidance_changes(classifier, z, time, label=1, gradient=False): + """Log h(candidate)-log h(current), exact or first-order one-hot.""" + batch, length = z.shape + t = torch.full((batch,), float(time)) + if gradient: + with torch.enable_grad(): + soft = F.one_hot(z, 5).float().requires_grad_(True) + value = log_success(classifier, soft, t, label) + grad = torch.autograd.grad(value.sum(), soft)[0] + current = grad.gather(-1, z[..., None]) + return grad[..., :K] - current + with torch.no_grad(): + current = log_success(classifier, z, t, label) + delta = torch.empty((batch, length, K)) + for i in range(length): + for a in range(K): + edited = z.clone() + edited[:, i] = a + delta[:, i, a] = log_success(classifier, edited, t, label) - current + return delta + + +@torch.no_grad() +def classifier_sample(model, classifier, batch, length, strength=1., + label=1, steps=40, gradient=False, epsilon=.005): + """Rate guidance with adaptive Euler and an explicit endpoint closure.""" + z = torch.full((batch, length), MASK) + time = 1. + iterations = 0 + while time > epsilon + 1e-8: + p = model(z).softmax(-1) + delta = guidance_changes(classifier, z, time, label, gradient) + rates = p / time * (strength * delta).clamp(-20, 20).exp() + rates *= (z == MASK)[..., None] + exit_rate = rates.sum(-1) + h = min(1. / steps, time-epsilon, .5 / max(float(exit_rate.max()), 1e-8)) + change = torch.rand(z.shape) < h * exit_rate + candidate = draw(rates / exit_rate.clamp_min(1e-12)[..., None] + 1e-12) + z = torch.where(change & (z == MASK), candidate, z) + time -= h + iterations += 1 + if iterations > 100000: + raise RuntimeError('Guided rate integration failed to advance.') + z = torch.where(z == MASK, draw(model(z).softmax(-1)), z) + return z + + +class SearchNode: + def __init__(self, tokens, parent=None, prior=1.): + self.tokens, self.parent, self.prior = tokens, parent, prior + self.children = [] + self.visits = 0 + self.reward = np.zeros(2) + + +@torch.no_grad() +def peptune_search(model, length=8, iterations=100, branching=4): + """DNA MCTS: selection, expansion, completion, Pareto rewards, backup. + + All DNA strings are valid. Peptide chemistry, bond-dependent masks, + RoFormer training, and the PepTune invalid-SMILES penalty are not used. + """ + root = SearchNode(torch.full((length,), MASK)) + archive = torch.empty((0, length), dtype=torch.long) + archive_scores = torch.empty((0, 2)) + trace = [] + for iteration in range(iterations): + node = root + while node.children: + def selection(child): + mean = child.reward.mean() / max(child.visits, 1) + bonus = 1.5 * child.prior * math.sqrt(node.visits + 1) / (child.visits + 1) + return mean + bonus + node = max(node.children, key=selection) + if (node.tokens == MASK).any(): + p = model(node.tokens[None]).softmax(-1)[0] + positions = torch.where(node.tokens == MASK)[0] + seen = set() + for _ in range(branching): + pos = int(positions[torch.randint(len(positions), ())]) + token = int(torch.multinomial(p[pos], 1)) + candidate = node.tokens.clone() + candidate[pos] = token + key = tuple(candidate.tolist()) + if key not in seen: + node.children.append(SearchNode(candidate, node, float(p[pos, token]))) + seen.add(key) + node = node.children[0] + rollout = node.tokens.clone() + while (rollout == MASK).any(): + prob = model(rollout[None]).softmax(-1) + position = int(torch.where(rollout == MASK)[0][0]) + rollout[position] = draw(prob)[0, position] + score = objectives(rollout[None])[0] + reward = ((score >= archive_scores).float().mean(0).numpy() + if len(archive) else np.ones(2)) + archive = torch.cat((archive, rollout[None])) + archive_scores = torch.cat((archive_scores, score[None])) + archive, archive_scores = pareto_filter(archive, archive_scores) + # Equal-score alternatives remain; remove exact repeated sequences only. + unique = []; seen = set() + for i, row in enumerate(archive.tolist()): + key = tuple(row) + if key not in seen: + unique.append(i); seen.add(key) + archive, archive_scores = archive[unique], archive_scores[unique] + while node is not None: + node.visits += 1 + node.reward += reward + node = node.parent + trace.append(dict(iteration=iteration, archive_size=len(archive), score=score.tolist())) + return archive, archive_scores, trace diff --git a/lecture_4/examples/block.py b/lecture_4/examples/block.py new file mode 100644 index 0000000000000000000000000000000000000000..34194d252b0ae68aee918ca1376c1477ca024bf2 --- /dev/null +++ b/lecture_4/examples/block.py @@ -0,0 +1,7 @@ +"""Complete block training and generation example.""" +import sys +from pathlib import Path +sys.path.insert(0, str(Path(__file__).resolve().parents[1])) +from run import main +if __name__ == '__main__': + main(['--method', 'block', '--out', 'outputs/block'] + sys.argv[1:]) diff --git a/lecture_4/examples/cfg.py b/lecture_4/examples/cfg.py new file mode 100644 index 0000000000000000000000000000000000000000..503f9de483e3784615dc12a916b2900d33dd705c --- /dev/null +++ b/lecture_4/examples/cfg.py @@ -0,0 +1,7 @@ +"""Complete cfg training and generation example.""" +import sys +from pathlib import Path +sys.path.insert(0, str(Path(__file__).resolve().parents[1])) +from run import main +if __name__ == '__main__': + main(['--method', 'cfg', '--out', 'outputs/cfg'] + sys.argv[1:]) diff --git a/lecture_4/examples/classifier_exact.py b/lecture_4/examples/classifier_exact.py new file mode 100644 index 0000000000000000000000000000000000000000..ff9256aa4cfb5316e502c1ea7c0dfc574332ba78 --- /dev/null +++ b/lecture_4/examples/classifier_exact.py @@ -0,0 +1,7 @@ +"""Complete classifier-exact training and generation example.""" +import sys +from pathlib import Path +sys.path.insert(0, str(Path(__file__).resolve().parents[1])) +from run import main +if __name__ == '__main__': + main(['--method', 'classifier-exact', '--out', 'outputs/classifier-exact'] + sys.argv[1:]) diff --git a/lecture_4/examples/classifier_gradient.py b/lecture_4/examples/classifier_gradient.py new file mode 100644 index 0000000000000000000000000000000000000000..6624c3a6fb107ce9e7210ad417031a1fa57fc0e4 --- /dev/null +++ b/lecture_4/examples/classifier_gradient.py @@ -0,0 +1,7 @@ +"""Complete classifier-gradient training and generation example.""" +import sys +from pathlib import Path +sys.path.insert(0, str(Path(__file__).resolve().parents[1])) +from run import main +if __name__ == '__main__': + main(['--method', 'classifier-gradient', '--out', 'outputs/classifier-gradient'] + sys.argv[1:]) diff --git a/lecture_4/examples/mdlm.py b/lecture_4/examples/mdlm.py new file mode 100644 index 0000000000000000000000000000000000000000..df1ffab5cf868d55e33761a2fb473604c4c271af --- /dev/null +++ b/lecture_4/examples/mdlm.py @@ -0,0 +1,7 @@ +"""Complete mdlm training and generation example.""" +import sys +from pathlib import Path +sys.path.insert(0, str(Path(__file__).resolve().parents[1])) +from run import main +if __name__ == '__main__': + main(['--method', 'mdlm', '--out', 'outputs/mdlm'] + sys.argv[1:]) diff --git a/lecture_4/examples/peptune.py b/lecture_4/examples/peptune.py new file mode 100644 index 0000000000000000000000000000000000000000..e1866f95df5c7f37657a57344d744f0f022cfe0a --- /dev/null +++ b/lecture_4/examples/peptune.py @@ -0,0 +1,7 @@ +"""Complete peptune training and generation example.""" +import sys +from pathlib import Path +sys.path.insert(0, str(Path(__file__).resolve().parents[1])) +from run import main +if __name__ == '__main__': + main(['--method', 'peptune', '--out', 'outputs/peptune'] + sys.argv[1:]) diff --git a/lecture_4/examples/udlm.py b/lecture_4/examples/udlm.py new file mode 100644 index 0000000000000000000000000000000000000000..f18f9974dbef7c65f1a5a6078531922db473cb8f --- /dev/null +++ b/lecture_4/examples/udlm.py @@ -0,0 +1,7 @@ +"""Complete udlm training and generation example.""" +import sys +from pathlib import Path +sys.path.insert(0, str(Path(__file__).resolve().parents[1])) +from run import main +if __name__ == '__main__': + main(['--method', 'udlm', '--out', 'outputs/udlm'] + sys.argv[1:]) diff --git a/lecture_4/lecture_core.py b/lecture_4/lecture_core.py new file mode 100644 index 0000000000000000000000000000000000000000..262490b4bbca24e99e92d5ca62fc5408cb480413 --- /dev/null +++ b/lecture_4/lecture_core.py @@ -0,0 +1,136 @@ +"""Small teaching implementations; illustrative data, not paper reproductions.""" + +import math + +import numpy as np + +import torch + +from torch import nn + +import torch.nn.functional as F + +from scipy.special import betainc, beta + +from scipy.stats import rankdata + +K, MASK = 4, 4 + +ALPHABET = 'ACGT' + +def encode(strings): + return torch.tensor([[ALPHABET.index(c) for c in s] + for s in strings]) + +def draw(prob): + shape = prob.shape[:-1] + sample = torch.multinomial(prob.reshape(-1, K), 1) + return sample.reshape(shape) + +class DNA(nn.Module): + def __init__(self, width=32, max_len=64): + super().__init__() + self.token = nn.Embedding(K + 1, width) + self.soft = nn.Linear(K, width) + self.position = nn.Embedding(max_len, width) + self.time = nn.Linear(1, width) + layer = nn.TransformerEncoderLayer( + width, 4, 2 * width, dropout=0., batch_first=True) + self.context = nn.TransformerEncoder(layer, 1) + self.output = nn.Linear(width, K) + + def forward(self, z, t=None): + h = self.token(z) if z.ndim == 2 else self.soft(z) + pos = torch.arange(z.shape[1], device=z.device) + h = h + self.position(pos)[None] + if t is not None: + h = h + self.time(t[:, None])[:, None] + return self.output(self.context(h)) + +def token_ce(logits, target): + return F.cross_entropy(logits.transpose(1, 2), + target, reduction='none') + +def mdlm_loss(model, clean): + batch, length = clean.shape + t = torch.rand(batch).clamp_min(1e-4) + masked = torch.rand(batch, length) < t[:, None] + noisy = clean.masked_fill(masked, MASK) + logits = model(noisy) # optimal predictor needs no t + ce = token_ce(logits, clean) + weighted = ce * masked / t[:, None] + return weighted.sum(1).mean() + +def train_step(model, optimizer, clean, loss_fn): + model.train() + optimizer.zero_grad() + loss = loss_fn(model, clean) + loss.backward() + optimizer.step() + return loss.item() + +@torch.no_grad() +def mdlm_sample(model, batch, length, steps=20): + model.eval() + z = torch.full((batch, length), MASK) + grid = torch.linspace(1., 0., steps + 1) + for t, s in zip(grid[:-1], grid[1:]): + prob = model(z).softmax(-1) + candidate = draw(prob) + reveal = torch.rand(z.shape) < (t - s) / t + update = (z == MASK) & reveal + z = torch.where(update, candidate, z) + return z + +def rate_step(z, rates, h): + exit_rate = rates.sum(-1) + assert torch.all(h * exit_rate <= 1. + 1e-6) + prob = h * rates + prob.scatter_(-1, z[..., None], + (1. - h * exit_rate)[..., None]) + return draw(prob.clamp_min(0.)) + +def udlm_rates(clean_prob, z, t): + alpha = 1. - t[:, None, None] + noisy_prob = alpha * clean_prob + (1. - alpha) / K + current = noisy_prob.gather(-1, z[..., None]) + rates = noisy_prob / (K * alpha * current) + return rates.scatter(-1, z[..., None], 0.) + +def udlm_loss(model, clean): + t = .02 + .96 * torch.rand(clean.shape[0]) + random = torch.randint(K, clean.shape) + z = torch.where(torch.rand(clean.shape) < t[:, None], + random, clean) + exact = F.one_hot(clean, K).float() + pred = model(z, t).softmax(-1) + a, b = udlm_rates(exact, z, t), udlm_rates(pred, z, t) + term = a * (a.clamp_min(1e-12).log() + - b.clamp_min(1e-12).log()) + b - a + return .96 * term.sum((1, 2)).mean() + +def block_loss(model, clean, start, size): + prefix, block = clean[:, :start], clean[:, start:start+size] + t = torch.rand(clean.shape[0]).clamp_min(1e-4) + mask = torch.rand(block.shape) < t[:, None] + ctx = torch.cat([prefix, block.masked_fill(mask, MASK)], 1) + logits = model(ctx)[:, start:] + loss = token_ce(logits, block) * mask / t[:, None] + return loss.sum(1).mean() + +def geometric_cfg(uncond, cond, strength): + logits = (1. - strength) * uncond.clamp_min(1e-12).log() + logits += strength * cond.clamp_min(1e-12).log() + return logits.softmax(-1) + +def guide_rates(base_rates, log_values, current_log_value, + strength=1.): + log_ratio = log_values - current_log_value[..., None] + return base_rates * (strength * log_ratio).exp() + +def pareto_filter(sequences, scores): + # Maximize both objectives; keep equal-score alternatives. + ge = (scores[:, None] >= scores[None, :]).all(-1) + gt = (scores[:, None] > scores[None, :]).any(-1) + dominated = (ge & gt).any(0) + return sequences[~dominated], scores[~dominated] diff --git a/lecture_4/numerical_examples.py b/lecture_4/numerical_examples.py new file mode 100644 index 0000000000000000000000000000000000000000..ecb4e0f78fbb6b5cc3ed7288bcb38622e3c28fb1 --- /dev/null +++ b/lecture_4/numerical_examples.py @@ -0,0 +1,24 @@ +"""Print the four-letter training and sampling calculations from the lecture.""" +import math +import torch +from lecture_core import udlm_rates, geometric_cfg + +dtype = torch.float64 +print('Alphabet order: A C G T; log losses are nats.') +q = .6*torch.eye(4,dtype=dtype)+.1*torch.ones(4,4,dtype=dtype) +print('Uniform one-step corruption matrix:\n',q) +print('Two-step marginal from clean A:',(q@q)[0].tolist()) +reverse=q[0]*q[:,1];reverse/=reverse.sum() +print('Previous base given clean A and final C:',reverse.tolist()) +loss=2*(-math.log(.6)-math.log(.7)) +print('MDLM ACGT -> AmGm, t=.5, loss:',loss) +print('MDLM reverse .5 -> .25 at missing C:',[.05,.30,.10,.05,.50]) +z=torch.tensor([[1]]);t=torch.tensor([.5],dtype=dtype) +a=udlm_rates(torch.tensor([[[1,0,0,0]]],dtype=dtype),z,t) +b=udlm_rates(torch.tensor([[[.6,.2,.1,.1]]],dtype=dtype),z,t) +rate_kl=(a*(a.clamp_min(1e-12).log()-b.clamp_min(1e-12).log())+b-a).sum() +print('UDLM target rates C -> A,C,G,T:',a.flatten().tolist()) +print('UDLM learned rates:',b.flatten().tolist(),'rate loss:',float(rate_kl)) +p=.1*b;p[...,1]=1-.1*b.sum(-1) +print('UDLM Euler step:',p.flatten().tolist()) +print('CFG strength 2:',geometric_cfg(torch.full((4,),.25),torch.tensor([.1,.2,.6,.1]),2).tolist()) diff --git a/lecture_4/requirements.txt b/lecture_4/requirements.txt new file mode 100644 index 0000000000000000000000000000000000000000..629699a0a61cd18d3bae56bef8d774c89ce2cfb4 --- /dev/null +++ b/lecture_4/requirements.txt @@ -0,0 +1,3 @@ +torch>=2.2 +numpy>=1.24 +scipy>=1.10 diff --git a/lecture_4/run.py b/lecture_4/run.py new file mode 100644 index 0000000000000000000000000000000000000000..e8210f1800e178fd8e7fa3c9bdc401e93c9b45ca --- /dev/null +++ b/lecture_4/run.py @@ -0,0 +1,127 @@ +#!/usr/bin/env python3 +"""Train and sample each Lecture 4 method on small, explicit DNA examples.""" +import argparse +from pathlib import Path +import torch +from lecture_core import DNA, mdlm_loss, mdlm_sample, udlm_loss, block_loss +from common import (seed_all, load_data, optimize, save_run, ConditionalDNA, + decode, metrics) +from diffusion import (uniform_sample, block_sample, conditional_loss, + cfg_sample, NoisyClassifier, fit_classifier, + classifier_sample, peptune_search) + +METHODS = ['mdlm', 'udlm', 'block', 'cfg', 'classifier-free', + 'classifier-gradient', 'classifier-exact', 'peptune'] + + +def main(argv=None): + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument('--method', choices=METHODS, default='mdlm') + parser.add_argument('--mode', choices=['train-sample', 'train', 'sample'], default='train-sample') + parser.add_argument('--train-steps', type=int, default=300) + parser.add_argument('--sample-steps', type=int, default=40) + parser.add_argument('--classifier-steps', type=int, default=300) + parser.add_argument('--search-steps', type=int, default=100) + parser.add_argument('--batch-size', type=int, default=32) + parser.add_argument('--samples', type=int, default=32) + parser.add_argument('--length', type=int, default=8) + parser.add_argument('--width', type=int, default=32) + parser.add_argument('--block-size', type=int, default=2) + parser.add_argument('--strength', type=float, default=1.5) + parser.add_argument('--label', type=int, choices=[0, 1], default=1) + parser.add_argument('--seed', type=int, default=7) + parser.add_argument('--data', help='TSV with sequence and optional binary label columns') + parser.add_argument('--out', default='outputs/mdlm') + args = parser.parse_args(argv) + if min(args.length, args.sample_steps, args.batch_size, args.samples, args.block_size) < 1: + parser.error('Lengths, step counts, and batch counts must be positive.') + if args.width % 4 or args.length > 64: + parser.error('Width must be divisible by four; length must not exceed 64.') + if args.mode != 'sample' and args.train_steps < 1: + parser.error('Training requires at least one step.') + seed_all(args.seed) + out = Path(args.out); out.mkdir(parents=True, exist_ok=True) + method = 'cfg' if args.method == 'classifier-free' else args.method + model = ConditionalDNA(args.width) if method == 'cfg' else DNA(args.width) + classifier = None + losses = [] + report = {'method': method, 'data': 'synthetic DNA; not biological validation'} + if args.mode == 'sample': + checkpoint = torch.load(out / 'checkpoint.pt', weights_only=True) + if checkpoint['method'] != method or checkpoint['length'] != args.length: + raise ValueError('Checkpoint method/length must match command arguments.') + model.load_state_dict(checkpoint['model']) + if 'classifier' in checkpoint: + classifier = NoisyClassifier(args.length) + classifier.load_state_dict(checkpoint['classifier']) + else: + data, labels = load_data(args.data, args.length) + split = max(1, int(.8 * len(data))) + train, train_labels = data[:split], labels[:split] + if method == 'udlm': + loss_fn = lambda m, x, idx: udlm_loss(m, x) + elif method == 'block': + starts = list(range(0, args.length, args.block_size)) + def loss_fn(m, x, idx): + start = starts[int(torch.randint(len(starts), ()))] + return len(starts) * block_loss(m, x, start, args.block_size) + elif method == 'cfg': + loss_fn = lambda m, x, idx: conditional_loss(m, x, train_labels[idx]) + else: + loss_fn = lambda m, x, idx: mdlm_loss(m, x) + losses = optimize(model, loss_fn, train, args.train_steps, args.batch_size) + checkpoint = {'model': model.state_dict(), 'method': method, + 'length': args.length, 'width': args.width} + if method.startswith('classifier-'): + classifier = NoisyClassifier(args.length) + classifier_losses = fit_classifier(classifier, train, train_labels, + args.classifier_steps, args.batch_size) + checkpoint['classifier'] = classifier.state_dict() + report['classifier_final_loss'] = classifier_losses[-1] + torch.save(checkpoint, out / 'checkpoint.pt') + report['train_loss_first_20_mean'] = sum(losses[:20]) / len(losses[:20]) + report['train_loss_last_20_mean'] = sum(losses[-20:]) / len(losses[-20:]) + # An independent noisy validation estimate, not a perplexity claim. + val = data[split:] + if len(val): + with torch.no_grad(): + if method == 'udlm': + validation = udlm_loss(model, val) + elif method == 'cfg': + validation = conditional_loss(model, val, labels[split:], drop=0.) + elif method == 'block': + validation = sum(block_loss(model, val, start, args.block_size) + for start in starts) + else: + validation = mdlm_loss(model, val) + report['validation_loss_one_mc_draw'] = float(validation) + if args.mode == 'train': + save_run(out, vars(args), losses, train[:args.samples], + {**report, 'sample_file_contains': 'training examples; generation not requested'}) + return report + model.eval() + if method == 'udlm': + samples = uniform_sample(model, args.samples, args.length, args.sample_steps) + report['endpoint_approximation'] = 't in [0.02, 0.98]; stop at residual noise 0.02, without a posterior interpretation of the UDLM parameter vector' + elif method == 'block': + samples = block_sample(model, args.samples, args.length, args.block_size, args.sample_steps) + elif method == 'cfg': + samples = cfg_sample(model, args.samples, args.length, args.strength, args.label, args.sample_steps) + elif method.startswith('classifier-'): + classifier.eval() + samples = classifier_sample(model, classifier, args.samples, args.length, + args.strength, args.label, args.sample_steps, + gradient=method == 'classifier-gradient') + report['guidance'] = 'learned noisy classifier; final residual masks use denoiser closure' + elif method == 'peptune': + samples, scores, trace = peptune_search(model, args.length, args.search_steps) + report['archive_scores'] = scores.tolist() + report['search_trace'] = trace + report['scope'] = 'DNA MCTS mechanism; not peptide-model training or paper reproduction' + else: + samples = mdlm_sample(model, args.samples, args.length, args.sample_steps) + return save_run(out, vars(args), losses, samples, report) + + +if __name__ == '__main__': + main() diff --git a/lecture_4/run_all.py b/lecture_4/run_all.py new file mode 100644 index 0000000000000000000000000000000000000000..86290aa9e3e5e646931aba2615210ce228954991 --- /dev/null +++ b/lecture_4/run_all.py @@ -0,0 +1,15 @@ +#!/usr/bin/env python3 +"""Run every complete training example; pass --quick for a small CPU check.""" +import argparse +from run import main +parser = argparse.ArgumentParser() +parser.add_argument('--quick', action='store_true') +args = parser.parse_args() +methods = ['mdlm', 'udlm', 'block', 'cfg', 'classifier-exact', 'classifier-gradient', 'peptune'] +for method in methods: + command = ['--method', method, '--out', 'outputs/' + method] + if args.quick: + command += ['--train-steps', '20', '--samples', '4', '--length', '4', + '--batch-size', '8', '--sample-steps', '20'] + command += ['--classifier-steps', '20', '--search-steps', '12'] + main(command) diff --git a/lecture_4/tests/test_mathematics.py b/lecture_4/tests/test_mathematics.py new file mode 100644 index 0000000000000000000000000000000000000000..1501091f7cbeab65c5e85eaee71a588e38362d2a --- /dev/null +++ b/lecture_4/tests/test_mathematics.py @@ -0,0 +1,50 @@ +import unittest +import sys +from pathlib import Path +sys.path.insert(0, str(Path(__file__).resolve().parents[1])) +import torch +from lecture_core import udlm_rates, geometric_cfg +from diffusion import guidance_changes, NoisyClassifier + + +class Mathematics(unittest.TestCase): + def test_masked_kl_reduces_to_cross_entropy(self): + r = .5 + d = torch.tensor([.1,.6,.2,.1], dtype=torch.float64) + q = torch.tensor([0,r,0,0,1-r], dtype=torch.float64) + p = torch.cat((r*d, torch.tensor([1-r]))) + actual = (q[q>0] * (q[q>0]/p[q>0]).log()).sum() + self.assertAlmostEqual(float(actual), float(-r*d[1].log()), places=12) + + def test_udlm_small_step_kl_converges_to_rate_loss(self): + a = torch.tensor([2.5,.5,.5], dtype=torch.float64) + b = torch.tensor([17/18,7/18,7/18], dtype=torch.float64) + target = (a*(a/b).log()+b-a).sum() + errors = [] + for h in [1e-3,1e-4,1e-5]: + q = torch.cat((h*a,(1-h*a.sum()).reshape(1))) + p = torch.cat((h*b,(1-h*b.sum()).reshape(1))) + local = (q*(q/p).log()).sum()/h + errors.append(abs(float(local-target))) + self.assertLess(errors[-1], errors[0]/50) + self.assertAlmostEqual(float(target), .9071595148, places=8) + + def test_cfg_dna_example(self): + p = geometric_cfg(torch.full((4,),.25),torch.tensor([.1,.2,.6,.1]),2.) + torch.testing.assert_close(p, torch.tensor([1,4,36,1])/42.) + + def test_classifier_gradient_matches_directional_derivative(self): + torch.manual_seed(2) + model = NoisyClassifier(4).double() + z = torch.tensor([[0,1,2,3]]) + soft = torch.nn.functional.one_hot(z,5).double().requires_grad_(True) + t = torch.tensor([.5], dtype=torch.float64) + f = lambda x: torch.nn.functional.logsigmoid(model(x,t)).sum() + grad = torch.autograd.grad(f(soft),soft)[0] + direction = torch.zeros_like(soft);direction[0,0,0]=-1;direction[0,0,1]=1 + h=1e-5 + numeric=(f(soft+h*direction)-f(soft-h*direction))/(2*h) + self.assertAlmostEqual(float(numeric.detach()),float((grad*direction).sum()),places=8) + + +if __name__ == '__main__': unittest.main() diff --git a/lecture_4/verified_examples/README.md b/lecture_4/verified_examples/README.md new file mode 100644 index 0000000000000000000000000000000000000000..8f1defccc1d307b78a0c780ea14a30930cc1de17 --- /dev/null +++ b/lecture_4/verified_examples/README.md @@ -0,0 +1,15 @@ +# Verified CPU examples + +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. + +| Method | Mean first 20 losses | Mean last 20 losses | Valid DNA | +| --- | ---: | ---: | --- | +| mdlm | 4.8388 | 2.6139 | True | +| udlm | 5.1313 | 2.2786 | True | +| block | 5.3755 | 3.2929 | True | +| cfg | 4.3721 | 2.3383 | True | +| classifier-exact | 4.8388 | 2.6139 | True | +| classifier-gradient | 4.8388 | 2.6139 | True | +| peptune | 4.8388 | 2.6139 | True | + +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. diff --git a/lecture_4/verified_examples/block/config.json b/lecture_4/verified_examples/block/config.json new file mode 100644 index 0000000000000000000000000000000000000000..a700a16f7636fd12f4232349c3e7cd9f4c9f9753 --- /dev/null +++ b/lecture_4/verified_examples/block/config.json @@ -0,0 +1,18 @@ +{ + "method": "block", + "mode": "train-sample", + "train_steps": 160, + "sample_steps": 40, + "classifier_steps": 160, + "search_steps": 40, + "batch_size": 16, + "samples": 8, + "length": 4, + "width": 32, + "block_size": 2, + "strength": 1.5, + "label": 1, + "seed": 7, + "data": null, + "out": "outputs/block" +} \ No newline at end of file diff --git a/lecture_4/verified_examples/block/losses.json b/lecture_4/verified_examples/block/losses.json new file mode 100644 index 0000000000000000000000000000000000000000..c471aad4ecb2ef58b48a6ba1d59a302fec3bec3a --- /dev/null +++ b/lecture_4/verified_examples/block/losses.json @@ -0,0 +1,162 @@ +[ + 4.331284523010254, + 5.4537835121154785, + 4.560967445373535, + 6.563370704650879, + 8.061755180358887, + 3.8685860633850098, + 6.454323768615723, + 7.774837970733643, + 2.178312301635742, + 4.992790222167969, + 12.195539474487305, + 5.040197372436523, + 3.1479640007019043, + 4.1676025390625, + 4.741274356842041, + 6.513577461242676, + 5.152376651763916, + 6.8658881187438965, + 3.170300006866455, + 2.2749924659729004, + 2.5343780517578125, + 4.639561176300049, + 3.381667375564575, + 2.6100592613220215, + 4.106627941131592, + 5.132066249847412, + 4.543891906738281, + 4.116055011749268, + 4.740029811859131, + 11.817144393920898, + 1.597190499305725, + 2.864633321762085, + 3.4319229125976562, + 2.477750778198242, + 2.259556293487549, + 5.4611029624938965, + 1.7751753330230713, + 3.8391871452331543, + 10.745177268981934, + 2.813826084136963, + 3.071927547454834, + 4.419896602630615, + 5.47519588470459, + 15.978719711303711, + 1.765317440032959, + 3.7841482162475586, + 3.9389586448669434, + 4.711722373962402, + 3.904447078704834, + 4.800976753234863, + 2.614999771118164, + 3.3339426517486572, + 7.427157402038574, + 3.459439754486084, + 5.940922260284424, + 6.2956037521362305, + 1.3079692125320435, + 3.5128886699676514, + 4.249375820159912, + 2.7473535537719727, + 2.1603171825408936, + 0.7737497687339783, + 2.2809128761291504, + 4.380037307739258, + 0.9533628225326538, + 2.962036609649658, + 2.988816261291504, + 2.04714298248291, + 7.599990367889404, + 3.762449026107788, + 3.4661433696746826, + 2.0959715843200684, + 1.7144439220428467, + 3.5818352699279785, + 3.0614638328552246, + 0.7883206605911255, + 1.4352097511291504, + 1.7439507246017456, + 0.8709290623664856, + 0.458389014005661, + 3.4738333225250244, + 4.326974391937256, + 6.331646919250488, + 2.4808788299560547, + 3.24576997756958, + 4.848365306854248, + 1.6670418977737427, + 2.3548476696014404, + 1.993442177772522, + 2.5002496242523193, + 0.3059796988964081, + 2.1486732959747314, + 6.24207878112793, + 1.8249398469924927, + 3.396871566772461, + 3.9418718814849854, + 0.29895153641700745, + 16.587575912475586, + 0.5883609652519226, + 3.9876809120178223, + 1.529473066329956, + 6.409687519073486, + 3.979583978652954, + 3.5919885635375977, + 0.8253586292266846, + 2.201484203338623, + 0.2219580113887787, + 0.9641126990318298, + 8.145115852355957, + 4.052484035491943, + 2.1619670391082764, + 3.3841612339019775, + 2.688559055328369, + 0.4211858808994293, + 0.799785852432251, + 3.8083810806274414, + 1.6525324583053589, + 4.437976837158203, + 1.2433699369430542, + 3.8457489013671875, + 4.093887805938721, + 1.0025733709335327, + 1.6318846940994263, + 3.593684196472168, + 2.961604356765747, + 2.3301167488098145, + 0.24246026575565338, + 3.887437582015991, + 2.659883737564087, + 0.6962559223175049, + 1.298153042793274, + 1.4727579355239868, + 2.848424196243286, + 1.1403732299804688, + 4.791389465332031, + 4.210074424743652, + 1.5758970975875854, + 2.269327402114868, + 0.18495622277259827, + 1.9766299724578857, + 1.7147716283798218, + 3.4504735469818115, + 4.908379554748535, + 5.258154392242432, + 0.6683526039123535, + 1.2319836616516113, + 6.596828460693359, + 3.3549842834472656, + 9.503194808959961, + 5.601862907409668, + 3.6974446773529053, + 4.31416654586792, + 1.9783447980880737, + 0.26172930002212524, + 0.2431989312171936, + 2.2771925926208496, + 1.737384557723999, + 3.598238706588745, + 3.4887681007385254, + 1.9730119705200195 +] \ No newline at end of file diff --git a/lecture_4/verified_examples/block/report.json b/lecture_4/verified_examples/block/report.json new file mode 100644 index 0000000000000000000000000000000000000000..d8df3aa5e0e6ffce9cb0b0d1d9221e172a01fad9 --- /dev/null +++ b/lecture_4/verified_examples/block/report.json @@ -0,0 +1,13 @@ +{ + "metrics": { + "valid_dna": true, + "unique_fraction": 0.75, + "mean_gc": 0.59375, + "mean_atat_match": 0.25 + }, + "method": "block", + "data": "synthetic DNA; not biological validation", + "train_loss_first_20_mean": 5.375486207008362, + "train_loss_last_20_mean": 3.2929233014583588, + "validation_loss_one_mc_draw": 2.7980802059173584 +} \ No newline at end of file diff --git a/lecture_4/verified_examples/block/samples.txt b/lecture_4/verified_examples/block/samples.txt new file mode 100644 index 0000000000000000000000000000000000000000..1efae56abb042f70ca55f888a31558c15cea1640 --- /dev/null +++ b/lecture_4/verified_examples/block/samples.txt @@ -0,0 +1,8 @@ +ACGC +ACGT +ACGT +CCTA +ACGT +CGTA +ACGA +GCGC diff --git a/lecture_4/verified_examples/cfg/config.json b/lecture_4/verified_examples/cfg/config.json new file mode 100644 index 0000000000000000000000000000000000000000..431cb32edf7d5ae1eef7792d5820332a0da0da7f --- /dev/null +++ b/lecture_4/verified_examples/cfg/config.json @@ -0,0 +1,18 @@ +{ + "method": "cfg", + "mode": "train-sample", + "train_steps": 160, + "sample_steps": 40, + "classifier_steps": 160, + "search_steps": 40, + "batch_size": 16, + "samples": 8, + "length": 4, + "width": 32, + "block_size": 2, + "strength": 1.5, + "label": 1, + "seed": 7, + "data": null, + "out": "outputs/cfg" +} \ No newline at end of file diff --git a/lecture_4/verified_examples/cfg/losses.json b/lecture_4/verified_examples/cfg/losses.json new file mode 100644 index 0000000000000000000000000000000000000000..744590809a90d61bf0915cb4b029d3159f4178fa --- /dev/null +++ b/lecture_4/verified_examples/cfg/losses.json @@ -0,0 +1,162 @@ +[ + 4.093783855438232, + 6.376791000366211, + 4.311301231384277, + 5.4559221267700195, + 5.326346397399902, + 4.026087760925293, + 2.633061647415161, + 4.690178394317627, + 5.398682117462158, + 7.209366798400879, + 4.6532487869262695, + 5.112819671630859, + 4.506524085998535, + 4.212957382202148, + 4.013753890991211, + 3.131486177444458, + 2.223745346069336, + 2.7967262268066406, + 3.788841724395752, + 3.4797284603118896, + 2.963750123977661, + 4.317558288574219, + 3.753492832183838, + 5.249013900756836, + 4.309139728546143, + 3.7436439990997314, + 10.669134140014648, + 4.060160160064697, + 2.0760574340820312, + 3.706965684890747, + 3.228757381439209, + 2.2654337882995605, + 4.35042667388916, + 3.539440631866455, + 3.0063133239746094, + 2.5612518787384033, + 5.035284519195557, + 3.1380114555358887, + 3.5969271659851074, + 4.42738151550293, + 2.6170473098754883, + 4.362626075744629, + 4.103579998016357, + 2.5357091426849365, + 3.5490846633911133, + 2.5230069160461426, + 2.873248338699341, + 3.158176898956299, + 4.957892894744873, + 2.1993887424468994, + 2.1841094493865967, + 4.587640285491943, + 1.8580347299575806, + 2.267033100128174, + 1.8229238986968994, + 1.7472164630889893, + 3.5524239540100098, + 2.736100435256958, + 3.342960834503174, + 2.218371629714966, + 2.875619411468506, + 2.75288462638855, + 1.6586329936981201, + 3.1111838817596436, + 4.498309135437012, + 2.5214474201202393, + 3.716679811477661, + 2.803637981414795, + 2.1772122383117676, + 2.169358253479004, + 1.9947891235351562, + 2.1023731231689453, + 2.3627333641052246, + 1.7967830896377563, + 1.8754620552062988, + 2.9364476203918457, + 2.1339101791381836, + 4.144628524780273, + 1.9647955894470215, + 2.0155508518218994, + 2.1333699226379395, + 1.434370517730713, + 2.0664165019989014, + 2.257500410079956, + 1.5195107460021973, + 4.65607213973999, + 2.1975350379943848, + 0.781586229801178, + 1.8217566013336182, + 2.1626477241516113, + 1.8518555164337158, + 3.667850971221924, + 2.1525752544403076, + 4.357625961303711, + 1.1739027500152588, + 0.9763473272323608, + 1.2919902801513672, + 2.5940964221954346, + 2.762819290161133, + 1.969483733177185, + 1.4582405090332031, + 1.5817911624908447, + 1.1240861415863037, + 2.102308988571167, + 3.745494842529297, + 2.124908208847046, + 2.0141687393188477, + 1.9382771253585815, + 1.5961400270462036, + 1.9971623420715332, + 2.9211485385894775, + 2.4554712772369385, + 2.6297366619110107, + 1.5343077182769775, + 3.172154664993286, + 2.1179494857788086, + 2.9100096225738525, + 1.7049200534820557, + 4.258328437805176, + 2.4142634868621826, + 2.198406934738159, + 1.5292108058929443, + 5.4681267738342285, + 2.697603225708008, + 3.230813503265381, + 2.773991107940674, + 2.4651236534118652, + 0.9841868877410889, + 0.870418906211853, + 3.471985101699829, + 2.449937582015991, + 3.4545187950134277, + 1.8310009241104126, + 2.834800958633423, + 0.7408347129821777, + 2.966944456100464, + 2.631303548812866, + 2.5108256340026855, + 2.7800302505493164, + 2.462411880493164, + 1.2342865467071533, + 2.3784074783325195, + 1.8845710754394531, + 2.4422521591186523, + 2.492060422897339, + 1.5423519611358643, + 2.8376150131225586, + 3.2689411640167236, + 4.493769645690918, + 1.8553943634033203, + 2.3281853199005127, + 1.4675829410552979, + 2.056398391723633, + 3.948185443878174, + 2.2498135566711426, + 1.085364818572998, + 2.6406710147857666, + 2.0750088691711426, + 1.4774831533432007, + 3.007999897003174 +] \ No newline at end of file diff --git a/lecture_4/verified_examples/cfg/report.json b/lecture_4/verified_examples/cfg/report.json new file mode 100644 index 0000000000000000000000000000000000000000..6e076300e58d58d5f8f80ccf8cbc73f0997a33e4 --- /dev/null +++ b/lecture_4/verified_examples/cfg/report.json @@ -0,0 +1,13 @@ +{ + "metrics": { + "valid_dna": true, + "unique_fraction": 0.625, + "mean_gc": 0.9375, + "mean_atat_match": 0.03125 + }, + "method": "cfg", + "data": "synthetic DNA; not biological validation", + "train_loss_first_20_mean": 4.372067654132843, + "train_loss_last_20_mean": 2.3383171617984773, + "validation_loss_one_mc_draw": 2.0500779151916504 +} \ No newline at end of file diff --git a/lecture_4/verified_examples/cfg/samples.txt b/lecture_4/verified_examples/cfg/samples.txt new file mode 100644 index 0000000000000000000000000000000000000000..0b23e6b08069dfc930cee41d03d95c186fa6aee8 --- /dev/null +++ b/lecture_4/verified_examples/cfg/samples.txt @@ -0,0 +1,8 @@ +GCCC +GCGC +CGGC +GCGC +GCGC +GCAC +GCGA +GCGC diff --git a/lecture_4/verified_examples/classifier-exact/config.json b/lecture_4/verified_examples/classifier-exact/config.json new file mode 100644 index 0000000000000000000000000000000000000000..75b08015930c9c46f4378ec4b195ef85c2d1cff2 --- /dev/null +++ b/lecture_4/verified_examples/classifier-exact/config.json @@ -0,0 +1,18 @@ +{ + "method": "classifier-exact", + "mode": "train-sample", + "train_steps": 160, + "sample_steps": 40, + "classifier_steps": 160, + "search_steps": 40, + "batch_size": 16, + "samples": 8, + "length": 4, + "width": 32, + "block_size": 2, + "strength": 1.5, + "label": 1, + "seed": 7, + "data": null, + "out": "outputs/classifier-exact" +} \ No newline at end of file diff --git a/lecture_4/verified_examples/classifier-exact/losses.json b/lecture_4/verified_examples/classifier-exact/losses.json new file mode 100644 index 0000000000000000000000000000000000000000..b635568a3b9a5a1e1389b46e3643b1aad6d81020 --- /dev/null +++ b/lecture_4/verified_examples/classifier-exact/losses.json @@ -0,0 +1,162 @@ +[ + 4.01927375793457, + 3.8247132301330566, + 6.5967559814453125, + 3.0379931926727295, + 5.381988048553467, + 4.364372730255127, + 6.268919944763184, + 4.896388530731201, + 2.5558066368103027, + 5.216179847717285, + 4.734250545501709, + 5.792362213134766, + 3.2293550968170166, + 4.882113456726074, + 4.7334160804748535, + 4.435980796813965, + 5.003822326660156, + 5.925942420959473, + 7.791959762573242, + 4.085095405578613, + 3.899994373321533, + 3.2904622554779053, + 3.4589755535125732, + 2.720918655395508, + 3.027308225631714, + 3.0458157062530518, + 3.7035269737243652, + 3.8884239196777344, + 4.001327037811279, + 4.476953983306885, + 2.1432690620422363, + 4.810318946838379, + 2.952962875366211, + 4.420904159545898, + 3.500626564025879, + 4.333303928375244, + 3.9793171882629395, + 1.8058104515075684, + 4.20003080368042, + 3.7151639461517334, + 2.2277348041534424, + 2.7376279830932617, + 3.615410327911377, + 4.500275611877441, + 5.066170692443848, + 3.6637136936187744, + 2.8754518032073975, + 2.339322328567505, + 2.037803888320923, + 2.1924853324890137, + 3.4916043281555176, + 2.519395351409912, + 2.597137451171875, + 2.4708733558654785, + 2.5220284461975098, + 2.5972704887390137, + 3.05387806892395, + 5.856385231018066, + 2.131591320037842, + 3.3278908729553223, + 2.400947093963623, + 5.875596046447754, + 3.1304073333740234, + 3.21421217918396, + 2.4150002002716064, + 2.215695381164551, + 2.139690637588501, + 3.2768962383270264, + 2.847025156021118, + 3.595715045928955, + 1.7762316465377808, + 3.3017048835754395, + 3.1203761100769043, + 2.4890081882476807, + 4.99165153503418, + 3.243589401245117, + 2.338869333267212, + 1.9984252452850342, + 3.838407516479492, + 2.6884045600891113, + 6.401689052581787, + 1.2461563348770142, + 2.748701810836792, + 8.534076690673828, + 3.991495132446289, + 2.489625930786133, + 2.2347571849823, + 2.9534873962402344, + 2.5714945793151855, + 1.147279977798462, + 2.4217231273651123, + 1.3990497589111328, + 2.5761783123016357, + 2.3062100410461426, + 3.08569073677063, + 1.516831874847412, + 3.085536003112793, + 2.8733534812927246, + 1.2969731092453003, + 3.7263309955596924, + 2.435549736022949, + 3.5106093883514404, + 2.7424113750457764, + 3.264981269836426, + 1.683659315109253, + 1.149167776107788, + 2.775585889816284, + 3.208291530609131, + 2.314664602279663, + 3.91782283782959, + 1.7896864414215088, + 2.5307981967926025, + 2.4437763690948486, + 2.521484136581421, + 2.2941794395446777, + 1.6110831499099731, + 3.128873825073242, + 3.0251963138580322, + 2.2769689559936523, + 1.4722540378570557, + 1.7136496305465698, + 5.040107727050781, + 2.9950344562530518, + 4.272597789764404, + 2.686481237411499, + 2.7622010707855225, + 3.0286059379577637, + 2.0905165672302246, + 2.4465155601501465, + 2.035715341567993, + 2.315361499786377, + 1.4022663831710815, + 3.5252881050109863, + 2.6312732696533203, + 3.7926905155181885, + 1.7568085193634033, + 1.798338532447815, + 3.3464834690093994, + 2.249241828918457, + 2.2580409049987793, + 4.0587310791015625, + 2.206573247909546, + 0.6222774982452393, + 1.8168766498565674, + 3.9070353507995605, + 2.2064626216888428, + 1.1733113527297974, + 1.383936882019043, + 2.251495599746704, + 2.962852954864502, + 2.8415677547454834, + 1.3905932903289795, + 3.165285587310791, + 3.477961540222168, + 1.7277320623397827, + 2.4671823978424072, + 5.010605812072754, + 3.3257808685302734, + 2.7084109783172607, + 3.5739824771881104 +] \ No newline at end of file diff --git a/lecture_4/verified_examples/classifier-exact/report.json b/lecture_4/verified_examples/classifier-exact/report.json new file mode 100644 index 0000000000000000000000000000000000000000..2ed766c3e844c772e65bfcbd4fea39844dd551f1 --- /dev/null +++ b/lecture_4/verified_examples/classifier-exact/report.json @@ -0,0 +1,15 @@ +{ + "metrics": { + "valid_dna": true, + "unique_fraction": 0.625, + "mean_gc": 0.84375, + "mean_atat_match": 0.0625 + }, + "method": "classifier-exact", + "data": "synthetic DNA; not biological validation", + "classifier_final_loss": 0.3678973913192749, + "train_loss_first_20_mean": 4.838834500312805, + "train_loss_last_20_mean": 2.613932800292969, + "validation_loss_one_mc_draw": 3.0652260780334473, + "guidance": "learned noisy classifier; final residual masks use denoiser closure" +} \ No newline at end of file diff --git a/lecture_4/verified_examples/classifier-exact/samples.txt b/lecture_4/verified_examples/classifier-exact/samples.txt new file mode 100644 index 0000000000000000000000000000000000000000..51f456975a73f6b90289f8c5510e3f7a71a08779 --- /dev/null +++ b/lecture_4/verified_examples/classifier-exact/samples.txt @@ -0,0 +1,8 @@ +GCGC +GCGC +ACGC +GCGC +GGTC +GCAA +CGTC +GCGC diff --git a/lecture_4/verified_examples/classifier-gradient/config.json b/lecture_4/verified_examples/classifier-gradient/config.json new file mode 100644 index 0000000000000000000000000000000000000000..0ec6ceaced02451fa49d57fe15cd66e43897f9e0 --- /dev/null +++ b/lecture_4/verified_examples/classifier-gradient/config.json @@ -0,0 +1,18 @@ +{ + "method": "classifier-gradient", + "mode": "train-sample", + "train_steps": 160, + "sample_steps": 40, + "classifier_steps": 160, + "search_steps": 40, + "batch_size": 16, + "samples": 8, + "length": 4, + "width": 32, + "block_size": 2, + "strength": 1.5, + "label": 1, + "seed": 7, + "data": null, + "out": "outputs/classifier-gradient" +} \ No newline at end of file diff --git a/lecture_4/verified_examples/classifier-gradient/losses.json b/lecture_4/verified_examples/classifier-gradient/losses.json new file mode 100644 index 0000000000000000000000000000000000000000..b635568a3b9a5a1e1389b46e3643b1aad6d81020 --- /dev/null +++ b/lecture_4/verified_examples/classifier-gradient/losses.json @@ -0,0 +1,162 @@ +[ + 4.01927375793457, + 3.8247132301330566, + 6.5967559814453125, + 3.0379931926727295, + 5.381988048553467, + 4.364372730255127, + 6.268919944763184, + 4.896388530731201, + 2.5558066368103027, + 5.216179847717285, + 4.734250545501709, + 5.792362213134766, + 3.2293550968170166, + 4.882113456726074, + 4.7334160804748535, + 4.435980796813965, + 5.003822326660156, + 5.925942420959473, + 7.791959762573242, + 4.085095405578613, + 3.899994373321533, + 3.2904622554779053, + 3.4589755535125732, + 2.720918655395508, + 3.027308225631714, + 3.0458157062530518, + 3.7035269737243652, + 3.8884239196777344, + 4.001327037811279, + 4.476953983306885, + 2.1432690620422363, + 4.810318946838379, + 2.952962875366211, + 4.420904159545898, + 3.500626564025879, + 4.333303928375244, + 3.9793171882629395, + 1.8058104515075684, + 4.20003080368042, + 3.7151639461517334, + 2.2277348041534424, + 2.7376279830932617, + 3.615410327911377, + 4.500275611877441, + 5.066170692443848, + 3.6637136936187744, + 2.8754518032073975, + 2.339322328567505, + 2.037803888320923, + 2.1924853324890137, + 3.4916043281555176, + 2.519395351409912, + 2.597137451171875, + 2.4708733558654785, + 2.5220284461975098, + 2.5972704887390137, + 3.05387806892395, + 5.856385231018066, + 2.131591320037842, + 3.3278908729553223, + 2.400947093963623, + 5.875596046447754, + 3.1304073333740234, + 3.21421217918396, + 2.4150002002716064, + 2.215695381164551, + 2.139690637588501, + 3.2768962383270264, + 2.847025156021118, + 3.595715045928955, + 1.7762316465377808, + 3.3017048835754395, + 3.1203761100769043, + 2.4890081882476807, + 4.99165153503418, + 3.243589401245117, + 2.338869333267212, + 1.9984252452850342, + 3.838407516479492, + 2.6884045600891113, + 6.401689052581787, + 1.2461563348770142, + 2.748701810836792, + 8.534076690673828, + 3.991495132446289, + 2.489625930786133, + 2.2347571849823, + 2.9534873962402344, + 2.5714945793151855, + 1.147279977798462, + 2.4217231273651123, + 1.3990497589111328, + 2.5761783123016357, + 2.3062100410461426, + 3.08569073677063, + 1.516831874847412, + 3.085536003112793, + 2.8733534812927246, + 1.2969731092453003, + 3.7263309955596924, + 2.435549736022949, + 3.5106093883514404, + 2.7424113750457764, + 3.264981269836426, + 1.683659315109253, + 1.149167776107788, + 2.775585889816284, + 3.208291530609131, + 2.314664602279663, + 3.91782283782959, + 1.7896864414215088, + 2.5307981967926025, + 2.4437763690948486, + 2.521484136581421, + 2.2941794395446777, + 1.6110831499099731, + 3.128873825073242, + 3.0251963138580322, + 2.2769689559936523, + 1.4722540378570557, + 1.7136496305465698, + 5.040107727050781, + 2.9950344562530518, + 4.272597789764404, + 2.686481237411499, + 2.7622010707855225, + 3.0286059379577637, + 2.0905165672302246, + 2.4465155601501465, + 2.035715341567993, + 2.315361499786377, + 1.4022663831710815, + 3.5252881050109863, + 2.6312732696533203, + 3.7926905155181885, + 1.7568085193634033, + 1.798338532447815, + 3.3464834690093994, + 2.249241828918457, + 2.2580409049987793, + 4.0587310791015625, + 2.206573247909546, + 0.6222774982452393, + 1.8168766498565674, + 3.9070353507995605, + 2.2064626216888428, + 1.1733113527297974, + 1.383936882019043, + 2.251495599746704, + 2.962852954864502, + 2.8415677547454834, + 1.3905932903289795, + 3.165285587310791, + 3.477961540222168, + 1.7277320623397827, + 2.4671823978424072, + 5.010605812072754, + 3.3257808685302734, + 2.7084109783172607, + 3.5739824771881104 +] \ No newline at end of file diff --git a/lecture_4/verified_examples/classifier-gradient/report.json b/lecture_4/verified_examples/classifier-gradient/report.json new file mode 100644 index 0000000000000000000000000000000000000000..27bdc27c718afe92cf03f0097ec885274d29cae2 --- /dev/null +++ b/lecture_4/verified_examples/classifier-gradient/report.json @@ -0,0 +1,15 @@ +{ + "metrics": { + "valid_dna": true, + "unique_fraction": 0.5, + "mean_gc": 0.90625, + "mean_atat_match": 0.03125 + }, + "method": "classifier-gradient", + "data": "synthetic DNA; not biological validation", + "classifier_final_loss": 0.3678973913192749, + "train_loss_first_20_mean": 4.838834500312805, + "train_loss_last_20_mean": 2.613932800292969, + "validation_loss_one_mc_draw": 3.0652260780334473, + "guidance": "learned noisy classifier; final residual masks use denoiser closure" +} \ No newline at end of file diff --git a/lecture_4/verified_examples/classifier-gradient/samples.txt b/lecture_4/verified_examples/classifier-gradient/samples.txt new file mode 100644 index 0000000000000000000000000000000000000000..1d666dcfabb974a45143bbab40e6b4b52133f27b --- /dev/null +++ b/lecture_4/verified_examples/classifier-gradient/samples.txt @@ -0,0 +1,8 @@ +GCGC +CCGC +GCTC +GCGC +GCGC +GCAA +GCGC +GCGC diff --git a/lecture_4/verified_examples/environment.json b/lecture_4/verified_examples/environment.json new file mode 100644 index 0000000000000000000000000000000000000000..8cffcc8fa542f4980ddd9dbc5fe3929cee0a59bd --- /dev/null +++ b/lecture_4/verified_examples/environment.json @@ -0,0 +1,7 @@ +{ + "python": "3.12.14", + "torch": "2.14.0+cpu", + "numpy": "2.3.5", + "scipy": "1.17.0", + "device": "cpu" +} \ No newline at end of file diff --git a/lecture_4/verified_examples/mdlm/config.json b/lecture_4/verified_examples/mdlm/config.json new file mode 100644 index 0000000000000000000000000000000000000000..e21c2af1ef1b53f2b2174645d450ca5641919afb --- /dev/null +++ b/lecture_4/verified_examples/mdlm/config.json @@ -0,0 +1,18 @@ +{ + "method": "mdlm", + "mode": "train-sample", + "train_steps": 160, + "sample_steps": 40, + "classifier_steps": 160, + "search_steps": 40, + "batch_size": 16, + "samples": 8, + "length": 4, + "width": 32, + "block_size": 2, + "strength": 1.5, + "label": 1, + "seed": 7, + "data": null, + "out": "outputs/mdlm" +} \ No newline at end of file diff --git a/lecture_4/verified_examples/mdlm/losses.json b/lecture_4/verified_examples/mdlm/losses.json new file mode 100644 index 0000000000000000000000000000000000000000..b635568a3b9a5a1e1389b46e3643b1aad6d81020 --- /dev/null +++ b/lecture_4/verified_examples/mdlm/losses.json @@ -0,0 +1,162 @@ +[ + 4.01927375793457, + 3.8247132301330566, + 6.5967559814453125, + 3.0379931926727295, + 5.381988048553467, + 4.364372730255127, + 6.268919944763184, + 4.896388530731201, + 2.5558066368103027, + 5.216179847717285, + 4.734250545501709, + 5.792362213134766, + 3.2293550968170166, + 4.882113456726074, + 4.7334160804748535, + 4.435980796813965, + 5.003822326660156, + 5.925942420959473, + 7.791959762573242, + 4.085095405578613, + 3.899994373321533, + 3.2904622554779053, + 3.4589755535125732, + 2.720918655395508, + 3.027308225631714, + 3.0458157062530518, + 3.7035269737243652, + 3.8884239196777344, + 4.001327037811279, + 4.476953983306885, + 2.1432690620422363, + 4.810318946838379, + 2.952962875366211, + 4.420904159545898, + 3.500626564025879, + 4.333303928375244, + 3.9793171882629395, + 1.8058104515075684, + 4.20003080368042, + 3.7151639461517334, + 2.2277348041534424, + 2.7376279830932617, + 3.615410327911377, + 4.500275611877441, + 5.066170692443848, + 3.6637136936187744, + 2.8754518032073975, + 2.339322328567505, + 2.037803888320923, + 2.1924853324890137, + 3.4916043281555176, + 2.519395351409912, + 2.597137451171875, + 2.4708733558654785, + 2.5220284461975098, + 2.5972704887390137, + 3.05387806892395, + 5.856385231018066, + 2.131591320037842, + 3.3278908729553223, + 2.400947093963623, + 5.875596046447754, + 3.1304073333740234, + 3.21421217918396, + 2.4150002002716064, + 2.215695381164551, + 2.139690637588501, + 3.2768962383270264, + 2.847025156021118, + 3.595715045928955, + 1.7762316465377808, + 3.3017048835754395, + 3.1203761100769043, + 2.4890081882476807, + 4.99165153503418, + 3.243589401245117, + 2.338869333267212, + 1.9984252452850342, + 3.838407516479492, + 2.6884045600891113, + 6.401689052581787, + 1.2461563348770142, + 2.748701810836792, + 8.534076690673828, + 3.991495132446289, + 2.489625930786133, + 2.2347571849823, + 2.9534873962402344, + 2.5714945793151855, + 1.147279977798462, + 2.4217231273651123, + 1.3990497589111328, + 2.5761783123016357, + 2.3062100410461426, + 3.08569073677063, + 1.516831874847412, + 3.085536003112793, + 2.8733534812927246, + 1.2969731092453003, + 3.7263309955596924, + 2.435549736022949, + 3.5106093883514404, + 2.7424113750457764, + 3.264981269836426, + 1.683659315109253, + 1.149167776107788, + 2.775585889816284, + 3.208291530609131, + 2.314664602279663, + 3.91782283782959, + 1.7896864414215088, + 2.5307981967926025, + 2.4437763690948486, + 2.521484136581421, + 2.2941794395446777, + 1.6110831499099731, + 3.128873825073242, + 3.0251963138580322, + 2.2769689559936523, + 1.4722540378570557, + 1.7136496305465698, + 5.040107727050781, + 2.9950344562530518, + 4.272597789764404, + 2.686481237411499, + 2.7622010707855225, + 3.0286059379577637, + 2.0905165672302246, + 2.4465155601501465, + 2.035715341567993, + 2.315361499786377, + 1.4022663831710815, + 3.5252881050109863, + 2.6312732696533203, + 3.7926905155181885, + 1.7568085193634033, + 1.798338532447815, + 3.3464834690093994, + 2.249241828918457, + 2.2580409049987793, + 4.0587310791015625, + 2.206573247909546, + 0.6222774982452393, + 1.8168766498565674, + 3.9070353507995605, + 2.2064626216888428, + 1.1733113527297974, + 1.383936882019043, + 2.251495599746704, + 2.962852954864502, + 2.8415677547454834, + 1.3905932903289795, + 3.165285587310791, + 3.477961540222168, + 1.7277320623397827, + 2.4671823978424072, + 5.010605812072754, + 3.3257808685302734, + 2.7084109783172607, + 3.5739824771881104 +] \ No newline at end of file diff --git a/lecture_4/verified_examples/mdlm/report.json b/lecture_4/verified_examples/mdlm/report.json new file mode 100644 index 0000000000000000000000000000000000000000..a1747e88a0494c0acdc763d5743db489d11f1381 --- /dev/null +++ b/lecture_4/verified_examples/mdlm/report.json @@ -0,0 +1,13 @@ +{ + "metrics": { + "valid_dna": true, + "unique_fraction": 0.75, + "mean_gc": 0.625, + "mean_atat_match": 0.09375 + }, + "method": "mdlm", + "data": "synthetic DNA; not biological validation", + "train_loss_first_20_mean": 4.838834500312805, + "train_loss_last_20_mean": 2.613932800292969, + "validation_loss_one_mc_draw": 2.5891432762145996 +} \ No newline at end of file diff --git a/lecture_4/verified_examples/mdlm/samples.txt b/lecture_4/verified_examples/mdlm/samples.txt new file mode 100644 index 0000000000000000000000000000000000000000..c27c562b3957ef54e11ced2008e9d130f048a379 --- /dev/null +++ b/lecture_4/verified_examples/mdlm/samples.txt @@ -0,0 +1,8 @@ +CGTA +ATGT +GCGC +TCGC +GCGC +GCGC +TATA +TCGA diff --git a/lecture_4/verified_examples/peptune/config.json b/lecture_4/verified_examples/peptune/config.json new file mode 100644 index 0000000000000000000000000000000000000000..81d8daf42e9d08dddd8ca7d5d34b29e2eeb4d8b8 --- /dev/null +++ b/lecture_4/verified_examples/peptune/config.json @@ -0,0 +1,18 @@ +{ + "method": "peptune", + "mode": "train-sample", + "train_steps": 160, + "sample_steps": 40, + "classifier_steps": 160, + "search_steps": 40, + "batch_size": 16, + "samples": 8, + "length": 4, + "width": 32, + "block_size": 2, + "strength": 1.5, + "label": 1, + "seed": 7, + "data": null, + "out": "outputs/peptune" +} \ No newline at end of file diff --git a/lecture_4/verified_examples/peptune/losses.json b/lecture_4/verified_examples/peptune/losses.json new file mode 100644 index 0000000000000000000000000000000000000000..b635568a3b9a5a1e1389b46e3643b1aad6d81020 --- /dev/null +++ b/lecture_4/verified_examples/peptune/losses.json @@ -0,0 +1,162 @@ +[ + 4.01927375793457, + 3.8247132301330566, + 6.5967559814453125, + 3.0379931926727295, + 5.381988048553467, + 4.364372730255127, + 6.268919944763184, + 4.896388530731201, + 2.5558066368103027, + 5.216179847717285, + 4.734250545501709, + 5.792362213134766, + 3.2293550968170166, + 4.882113456726074, + 4.7334160804748535, + 4.435980796813965, + 5.003822326660156, + 5.925942420959473, + 7.791959762573242, + 4.085095405578613, + 3.899994373321533, + 3.2904622554779053, + 3.4589755535125732, + 2.720918655395508, + 3.027308225631714, + 3.0458157062530518, + 3.7035269737243652, + 3.8884239196777344, + 4.001327037811279, + 4.476953983306885, + 2.1432690620422363, + 4.810318946838379, + 2.952962875366211, + 4.420904159545898, + 3.500626564025879, + 4.333303928375244, + 3.9793171882629395, + 1.8058104515075684, + 4.20003080368042, + 3.7151639461517334, + 2.2277348041534424, + 2.7376279830932617, + 3.615410327911377, + 4.500275611877441, + 5.066170692443848, + 3.6637136936187744, + 2.8754518032073975, + 2.339322328567505, + 2.037803888320923, + 2.1924853324890137, + 3.4916043281555176, + 2.519395351409912, + 2.597137451171875, + 2.4708733558654785, + 2.5220284461975098, + 2.5972704887390137, + 3.05387806892395, + 5.856385231018066, + 2.131591320037842, + 3.3278908729553223, + 2.400947093963623, + 5.875596046447754, + 3.1304073333740234, + 3.21421217918396, + 2.4150002002716064, + 2.215695381164551, + 2.139690637588501, + 3.2768962383270264, + 2.847025156021118, + 3.595715045928955, + 1.7762316465377808, + 3.3017048835754395, + 3.1203761100769043, + 2.4890081882476807, + 4.99165153503418, + 3.243589401245117, + 2.338869333267212, + 1.9984252452850342, + 3.838407516479492, + 2.6884045600891113, + 6.401689052581787, + 1.2461563348770142, + 2.748701810836792, + 8.534076690673828, + 3.991495132446289, + 2.489625930786133, + 2.2347571849823, + 2.9534873962402344, + 2.5714945793151855, + 1.147279977798462, + 2.4217231273651123, + 1.3990497589111328, + 2.5761783123016357, + 2.3062100410461426, + 3.08569073677063, + 1.516831874847412, + 3.085536003112793, + 2.8733534812927246, + 1.2969731092453003, + 3.7263309955596924, + 2.435549736022949, + 3.5106093883514404, + 2.7424113750457764, + 3.264981269836426, + 1.683659315109253, + 1.149167776107788, + 2.775585889816284, + 3.208291530609131, + 2.314664602279663, + 3.91782283782959, + 1.7896864414215088, + 2.5307981967926025, + 2.4437763690948486, + 2.521484136581421, + 2.2941794395446777, + 1.6110831499099731, + 3.128873825073242, + 3.0251963138580322, + 2.2769689559936523, + 1.4722540378570557, + 1.7136496305465698, + 5.040107727050781, + 2.9950344562530518, + 4.272597789764404, + 2.686481237411499, + 2.7622010707855225, + 3.0286059379577637, + 2.0905165672302246, + 2.4465155601501465, + 2.035715341567993, + 2.315361499786377, + 1.4022663831710815, + 3.5252881050109863, + 2.6312732696533203, + 3.7926905155181885, + 1.7568085193634033, + 1.798338532447815, + 3.3464834690093994, + 2.249241828918457, + 2.2580409049987793, + 4.0587310791015625, + 2.206573247909546, + 0.6222774982452393, + 1.8168766498565674, + 3.9070353507995605, + 2.2064626216888428, + 1.1733113527297974, + 1.383936882019043, + 2.251495599746704, + 2.962852954864502, + 2.8415677547454834, + 1.3905932903289795, + 3.165285587310791, + 3.477961540222168, + 1.7277320623397827, + 2.4671823978424072, + 5.010605812072754, + 3.3257808685302734, + 2.7084109783172607, + 3.5739824771881104 +] \ No newline at end of file diff --git a/lecture_4/verified_examples/peptune/report.json b/lecture_4/verified_examples/peptune/report.json new file mode 100644 index 0000000000000000000000000000000000000000..643ecceb2482e98f373228794000f29f97843d20 --- /dev/null +++ b/lecture_4/verified_examples/peptune/report.json @@ -0,0 +1,350 @@ +{ + "metrics": { + "valid_dna": true, + "unique_fraction": 1.0, + "mean_gc": 0.75, + "mean_atat_match": 0.25 + }, + "method": "peptune", + "data": "synthetic DNA; not biological validation", + "train_loss_first_20_mean": 4.838834500312805, + "train_loss_last_20_mean": 2.613932800292969, + "validation_loss_one_mc_draw": 2.5891432762145996, + "archive_scores": [ + [ + 1.0, + 0.0 + ], + [ + 0.5, + 0.5 + ], + [ + 0.75, + 0.25 + ] + ], + "search_trace": [ + { + "iteration": 0, + "archive_size": 1, + "score": [ + 0.25, + 0.25 + ] + }, + { + "iteration": 1, + "archive_size": 2, + "score": [ + 0.5, + 0.0 + ] + }, + { + "iteration": 2, + "archive_size": 2, + "score": [ + 0.0, + 0.0 + ] + }, + { + "iteration": 3, + "archive_size": 2, + "score": [ + 1.0, + 0.0 + ] + }, + { + "iteration": 4, + "archive_size": 2, + "score": [ + 0.5, + 0.5 + ] + }, + { + "iteration": 5, + "archive_size": 2, + "score": [ + 1.0, + 0.0 + ] + }, + { + "iteration": 6, + "archive_size": 2, + "score": [ + 0.5, + 0.0 + ] + }, + { + "iteration": 7, + "archive_size": 2, + "score": [ + 1.0, + 0.0 + ] + }, + { + "iteration": 8, + "archive_size": 2, + "score": [ + 0.0, + 0.0 + ] + }, + { + "iteration": 9, + "archive_size": 2, + "score": [ + 0.5, + 0.5 + ] + }, + { + "iteration": 10, + "archive_size": 2, + "score": [ + 0.5, + 0.5 + ] + }, + { + "iteration": 11, + "archive_size": 3, + "score": [ + 0.75, + 0.25 + ] + }, + { + "iteration": 12, + "archive_size": 3, + "score": [ + 0.5, + 0.0 + ] + }, + { + "iteration": 13, + "archive_size": 3, + "score": [ + 0.5, + 0.5 + ] + }, + { + "iteration": 14, + "archive_size": 3, + "score": [ + 0.5, + 0.0 + ] + }, + { + "iteration": 15, + "archive_size": 3, + "score": [ + 1.0, + 0.0 + ] + }, + { + "iteration": 16, + "archive_size": 3, + "score": [ + 0.5, + 0.5 + ] + }, + { + "iteration": 17, + "archive_size": 3, + "score": [ + 0.5, + 0.0 + ] + }, + { + "iteration": 18, + "archive_size": 3, + "score": [ + 0.5, + 0.5 + ] + }, + { + "iteration": 19, + "archive_size": 3, + "score": [ + 0.5, + 0.5 + ] + }, + { + "iteration": 20, + "archive_size": 3, + "score": [ + 0.0, + 0.0 + ] + }, + { + "iteration": 21, + "archive_size": 3, + "score": [ + 1.0, + 0.0 + ] + }, + { + "iteration": 22, + "archive_size": 3, + "score": [ + 0.5, + 0.5 + ] + }, + { + "iteration": 23, + "archive_size": 3, + "score": [ + 0.5, + 0.5 + ] + }, + { + "iteration": 24, + "archive_size": 3, + "score": [ + 1.0, + 0.0 + ] + }, + { + "iteration": 25, + "archive_size": 3, + "score": [ + 0.5, + 0.0 + ] + }, + { + "iteration": 26, + "archive_size": 3, + "score": [ + 0.5, + 0.5 + ] + }, + { + "iteration": 27, + "archive_size": 3, + "score": [ + 1.0, + 0.0 + ] + }, + { + "iteration": 28, + "archive_size": 3, + "score": [ + 0.5, + 0.0 + ] + }, + { + "iteration": 29, + "archive_size": 3, + "score": [ + 0.5, + 0.5 + ] + }, + { + "iteration": 30, + "archive_size": 3, + "score": [ + 0.5, + 0.5 + ] + }, + { + "iteration": 31, + "archive_size": 3, + "score": [ + 1.0, + 0.0 + ] + }, + { + "iteration": 32, + "archive_size": 3, + "score": [ + 0.5, + 0.5 + ] + }, + { + "iteration": 33, + "archive_size": 3, + "score": [ + 0.0, + 0.0 + ] + }, + { + "iteration": 34, + "archive_size": 3, + "score": [ + 0.5, + 0.5 + ] + }, + { + "iteration": 35, + "archive_size": 3, + "score": [ + 0.5, + 0.5 + ] + }, + { + "iteration": 36, + "archive_size": 3, + "score": [ + 1.0, + 0.0 + ] + }, + { + "iteration": 37, + "archive_size": 3, + "score": [ + 0.5, + 0.5 + ] + }, + { + "iteration": 38, + "archive_size": 3, + "score": [ + 0.5, + 0.0 + ] + }, + { + "iteration": 39, + "archive_size": 3, + "score": [ + 0.5, + 0.5 + ] + } + ], + "scope": "DNA MCTS mechanism; not peptide-model training or paper reproduction" +} \ No newline at end of file diff --git a/lecture_4/verified_examples/peptune/samples.txt b/lecture_4/verified_examples/peptune/samples.txt new file mode 100644 index 0000000000000000000000000000000000000000..49f3a6ceea9dd1ec151c4aeb5cefc8ca0e2d03b5 --- /dev/null +++ b/lecture_4/verified_examples/peptune/samples.txt @@ -0,0 +1,3 @@ +GCGC +ACGT +GCGT diff --git a/lecture_4/verified_examples/udlm/config.json b/lecture_4/verified_examples/udlm/config.json new file mode 100644 index 0000000000000000000000000000000000000000..aca4ef436c54fccf9e624d18206526acc8d387ea --- /dev/null +++ b/lecture_4/verified_examples/udlm/config.json @@ -0,0 +1,18 @@ +{ + "method": "udlm", + "mode": "train-sample", + "train_steps": 160, + "sample_steps": 40, + "classifier_steps": 300, + "search_steps": 100, + "batch_size": 16, + "samples": 8, + "length": 4, + "width": 32, + "block_size": 2, + "strength": 1.5, + "label": 1, + "seed": 7, + "data": null, + "out": "outputs/udlm" +} \ No newline at end of file diff --git a/lecture_4/verified_examples/udlm/losses.json b/lecture_4/verified_examples/udlm/losses.json new file mode 100644 index 0000000000000000000000000000000000000000..1de39c4c49329232af75d1ab979c4c17a71c1c44 --- /dev/null +++ b/lecture_4/verified_examples/udlm/losses.json @@ -0,0 +1,162 @@ +[ + 11.047435760498047, + 3.638819932937622, + 4.339240550994873, + 6.615440368652344, + 7.487863540649414, + 3.715148448944092, + 4.2508416175842285, + 5.987113952636719, + 4.574474334716797, + 2.5587449073791504, + 7.678818225860596, + 10.757691383361816, + 3.761280059814453, + 6.221782684326172, + 2.411395788192749, + 3.9185919761657715, + 2.432310104370117, + 5.186282157897949, + 2.995053768157959, + 3.0481975078582764, + 6.479434490203857, + 4.919713497161865, + 2.8192875385284424, + 2.8869221210479736, + 8.790794372558594, + 2.300025463104248, + 3.3890492916107178, + 3.32928466796875, + 3.703833818435669, + 2.642963409423828, + 3.1563267707824707, + 2.6353256702423096, + 3.5707004070281982, + 2.8710131645202637, + 2.8604514598846436, + 3.6414551734924316, + 4.774967193603516, + 2.004885673522949, + 8.153873443603516, + 1.7916756868362427, + 2.389755964279175, + 4.308714389801025, + 3.9405770301818848, + 6.173405170440674, + 2.9836504459381104, + 2.4070658683776855, + 4.4768242835998535, + 4.716032981872559, + 2.6472368240356445, + 2.835066795349121, + 2.540917158126831, + 3.043062925338745, + 5.332056045532227, + 3.2398338317871094, + 4.370035171508789, + 2.7836556434631348, + 4.69325590133667, + 1.9069887399673462, + 3.3850021362304688, + 2.8949105739593506, + 6.395257472991943, + 6.112308979034424, + 2.172653913497925, + 3.731419324874878, + 3.2633914947509766, + 2.7229864597320557, + 3.2528135776519775, + 6.598641395568848, + 2.3790218830108643, + 1.4963032007217407, + 2.819007158279419, + 2.80535888671875, + 2.719237804412842, + 2.0492019653320312, + 3.2515711784362793, + 1.8298454284667969, + 2.3929498195648193, + 2.056596517562866, + 15.28537654876709, + 3.128511667251587, + 2.0100417137145996, + 1.5984283685684204, + 1.9779601097106934, + 3.192647695541382, + 2.118055582046509, + 2.90411114692688, + 3.8938608169555664, + 2.1555304527282715, + 2.7105512619018555, + 3.3428311347961426, + 1.0114843845367432, + 3.9458742141723633, + 5.690505027770996, + 3.6527702808380127, + 1.1432782411575317, + 5.013024806976318, + 2.777940273284912, + 3.2634408473968506, + 4.932656288146973, + 4.210480690002441, + 2.6948583126068115, + 2.285061836242676, + 3.583308696746826, + 2.962766647338867, + 2.7541191577911377, + 2.287432909011841, + 3.1161656379699707, + 2.5217714309692383, + 3.1153383255004883, + 1.9834219217300415, + 2.0762336254119873, + 2.764826536178589, + 2.7770273685455322, + 1.6031702756881714, + 4.0543365478515625, + 2.547313690185547, + 2.8786070346832275, + 2.6822168827056885, + 6.25325345993042, + 2.590442419052124, + 2.122274398803711, + 1.6236686706542969, + 1.8864619731903076, + 2.4540152549743652, + 3.4854989051818848, + 3.156172513961792, + 4.207565784454346, + 3.6062185764312744, + 2.4027581214904785, + 1.1646852493286133, + 4.980432510375977, + 6.687517166137695, + 1.4371603727340698, + 1.6977627277374268, + 3.9442780017852783, + 2.9689416885375977, + 2.100611686706543, + 2.377840995788574, + 3.0913102626800537, + 6.303923606872559, + 2.232814073562622, + 1.8777183294296265, + 1.323475956916809, + 3.0791330337524414, + 1.765347957611084, + 1.629030466079712, + 2.6598291397094727, + 2.1885077953338623, + 2.656402587890625, + 2.2888011932373047, + 2.212771415710449, + 1.9104533195495605, + 4.000855445861816, + 2.3612029552459717, + 1.7967439889907837, + 1.9667776823043823, + 2.8789355754852295, + 1.4103233814239502, + 1.952869176864624, + 3.3791284561157227 +] \ No newline at end of file diff --git a/lecture_4/verified_examples/udlm/report.json b/lecture_4/verified_examples/udlm/report.json new file mode 100644 index 0000000000000000000000000000000000000000..f727cf5f4b47f9ef9458d107de3f3a608b01076d --- /dev/null +++ b/lecture_4/verified_examples/udlm/report.json @@ -0,0 +1,14 @@ +{ + "metrics": { + "valid_dna": true, + "unique_fraction": 0.875, + "mean_gc": 0.625, + "mean_atat_match": 0.25 + }, + "method": "udlm", + "data": "synthetic DNA; not biological validation", + "train_loss_first_20_mean": 5.131326353549957, + "train_loss_last_20_mean": 2.2785560965538023, + "validation_loss_one_mc_draw": 2.245309829711914, + "endpoint_approximation": "t in [0.02, 0.98]; stop at residual noise 0.02, without a posterior interpretation of the UDLM parameter vector" +} \ No newline at end of file diff --git a/lecture_4/verified_examples/udlm/samples.txt b/lecture_4/verified_examples/udlm/samples.txt new file mode 100644 index 0000000000000000000000000000000000000000..eb173dab7c8a8e7bacda63d3a5d838e63fa47b20 --- /dev/null +++ b/lecture_4/verified_examples/udlm/samples.txt @@ -0,0 +1,8 @@ +ACGT +ACGT +CGTA +ACTT +GCGT +ACGC +GCGC +CCTG diff --git a/lecture_5/.gitignore b/lecture_5/.gitignore new file mode 100644 index 0000000000000000000000000000000000000000..876670e6a75ee4b613025a193d0791312ac8679b --- /dev/null +++ b/lecture_5/.gitignore @@ -0,0 +1,5 @@ +__pycache__/ +.venv/ +outputs/ +*.pt +.pytest_cache/ diff --git a/lecture_5/README.md b/lecture_5/README.md new file mode 100644 index 0000000000000000000000000000000000000000..e9187e6f8fadfa67008a1f4da650045d2040f810 --- /dev/null +++ b/lecture_5/README.md @@ -0,0 +1,87 @@ +# CIS 6270 - Lecture 5 - Discrete Flow Matching + +Course hub: [ChatterjeeLab/CIS6270 on Hugging Face](https://huggingface.co/ChatterjeeLab/CIS6270). + +Complete CPU examples for the flow-matching code in the [Discrete Generation lecture](https://docs.google.com/presentation/d/1qcPdF4SYTq4VdTN_5w9GjviLRPr2HFtbuYMpbfLrc_M/edit). Begin with **Gat et al.'s classic discrete flow matching**, then change the path, loss, or guidance module. Lecture 4's diffusion examples are in [lecture_4/](../lecture_4/README.md). The slide deck remains combined. + +The examples use synthetic DNA and small transformers to make training and generation inspectable. They are teaching implementations, not replications of large-scale paper experiments. + +## Run classic DFM first + +Python 3.11 or later, CPU. Data generation and training require no external downloads beyond the dependencies. + +From the course repository root, enter this lecture folder. + +```bash +cd lecture_5 +python -m venv .venv +source .venv/bin/activate +pip install -r requirements.txt +python run.py --method gat --data data/dna_train.tsv --out outputs/gat +python run.py --method gat --mode sample --out outputs/gat +``` + +The first command trains on independent source/target pairs, validates on held-out synthetic DNA, saves a checkpoint, and generates sequences using the learned jump rates. Outputs include `config.json`, `losses.json`, `report.json`, `samples.txt`, and `checkpoint.pt`. Generation-only commands must use the checkpoint's method, width, and sequence length. Optimizer state is not saved for resumed training. + +```bash +python examples/gat.py --length 4 --train-steps 160 --samples 8 +python run_all.py +python run_all.py --quick +python -m unittest discover -s tests -v +python numerical_examples.py +``` + +## Change only the relevant module + +`lecture_core.py` contains the code modules from the slides. `common.py` contains the backbone-independent data and training utilities. `flows.py` implements complete integration, teacher recoupling, guidance, and refinement loops. `run.py` selects the loss and sampler. Each `examples/*.py` file is a runnable method-specific entry point. + +| Example | Training | Generation | +| --- | --- | --- | +| `gat` | Endpoint-token cross-entropy on a categorical mixture path | Euler updates from posterior-derived jump rates | +| `dirichlet` | Endpoint prediction from Dirichlet samples | Posterior average of Beta-CDF-derived conditional fields | +| `fisher` | Tangent velocity regression on square-root geodesics | Tangent integration, positive-orthant correction, spherical normalization | +| `gumbel` | Regression to the full noisy path derivative | Integrate from an approximate finite-temperature base and decode | +| `rectified` | Straight source-to-one-hot probability-path regression | Integrate the learned probability-vector field | +| `redi` | Train a Gat teacher, retain source/endpoints, train a new student | Generate from the recoupled student's rates | +| `mog-dfm` | Train a Gat backbone | Rank/direction rate reweighting and an adaptive acute hypercone | +| `areuredi` | Gat teacher followed by ReDi student training | Student generation, then annealed locally balanced MH refinement | + +```bash +python examples/dirichlet.py +python examples/fisher.py --guidance 0.1 +python examples/gumbel.py +python examples/rectified.py +python examples/redi.py --teacher-steps 300 --pairs 256 +python examples/mog_dfm.py --preference 0.7 0.3 --strength 1 +python examples/areuredi.py --preference 0.5 0.5 --refine-steps 100 +``` + +All objectives in `common.objectives` are maximized: GC fraction and agreement with a repeating ATAT motif. They deliberately conflict. `--guidance` adds the differentiable toy GC objective to a simplex field; it is not a learned biological classifier. + +## Numerical and theorem boundaries + +- **Gat:** the linear schedule uses `t=n/steps`, so the singular rate at `t=1` is never evaluated. Its Euler update is valid at each position because `h <= 1-t`. Parallel token updates have finite-step joint approximation error. +- **Dirichlet:** concentration increment ranges from 0 to 8. The Beta-CDF concentration derivative is evaluated with a centered finite difference. A finite concentration endpoint is not a vertex distribution; the final output uses argmax. +- **Fisher:** the metric has the factor four under the square-root map. Spherical tangency does not ensure positive coordinates under a finite neural Euler step. The sampler records negative-coordinate corrections before projection and normalization. +- **Gumbel:** the slide construction fixes the full perturbed logits and cools temperature. Defaults are `beta=1`, `tau_max=4`, decay 4. A finite starting temperature does not remove endpoint dependence exactly; the sampler's data-independent noise base is an approximation. The finite noisy endpoint need not select the training letter. The code does not claim exact endpoint matching. `numerical_examples.py` shows why noise scale and temperature play different roles. +- **Continuous rectification:** this example trains one rectified field. The convex-cost theorem concerns the exact conditional-mean field and preserved marginals; a finite learned projected solver is approximate. +- **ReDi:** the code performs one teacher-recoupling round and stores `teacher_pairs.pt`. It does not assume that recoupling automatically decreases conditional total correlation. The lecture states the extra projection-compatibility condition needed for the data-processing argument. +- **MOG-DFM:** candidate ranks are normalized within each position. The code evaluates all positions before a parallel Euler step, rather than using the paper's random-position loop. The cone adapts within 10 to 89 degrees and does not use an outside-cone fallback, so the acute-cone local progress condition is retained. The reported guarantee is local weighted progress, not improvement in every property or global Pareto optimality. +- **AReUReDi:** the MH demonstration uses an explicitly evaluable, smoothed position-wise reference fit to the DNA training data. A denoiser alone is not treated as a joint density. The full forward/reverse proposal and reference ratios remain in acceptance. Finite annealing is not a proof of equilibrium concentration or global optimization. + +The simplex solver reports how often a proposed update needed a nonnegativity correction. Increase `--sample-steps` when studying discretization; projection itself modifies the numerical dynamics. + +## Numerical checks and slide mapping + +The tests verify the exact master equation on all 16 two-base states, Fisher midpoint/tangency, Gumbel's path derivative with a fixed noise realization, and MH detailed balance on the complete two-base state space. `SLIDE_CODE_MAP.md` links all Lecture 5 code units to the matching functions. `verified_examples/` contains actual seeded CPU training and generation results for every method. + +## Sources + +- [Gat et al. - Discrete Flow Matching](https://arxiv.org/abs/2407.15595) +- [Dirichlet Flow Matching](https://arxiv.org/abs/2402.05841) +- [Fisher Flow Matching](https://arxiv.org/abs/2405.14664) +- [Gumbel-Softmax Flow Matching](https://arxiv.org/abs/2503.17361) +- [Rectified Flow](https://arxiv.org/abs/2209.03003) +- [ReDi](https://arxiv.org/abs/2507.15897) +- [MOG-DFM](https://arxiv.org/abs/2505.07086) +- [AReUReDi](https://arxiv.org/abs/2510.00352) diff --git a/lecture_5/SLIDE_CODE_MAP.md b/lecture_5/SLIDE_CODE_MAP.md new file mode 100644 index 0000000000000000000000000000000000000000..8d7d989e51ff6a96abc5fc0300ede2ef6a66f185 --- /dev/null +++ b/lecture_5/SLIDE_CODE_MAP.md @@ -0,0 +1,24 @@ +# Code walkthrough map + +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. + +| Slide unit | Code | Location | +| --- | --- | --- | +| [Code for classic DFM training pairs](https://docs.google.com/presentation/d/1qcPdF4SYTq4VdTN_5w9GjviLRPr2HFtbuYMpbfLrc_M/edit#slide=id.cis4r2_u155_s9) | `gat_loss` | `lecture_core.py` | +| [Code for convert a posterior into Gat jump rates](https://docs.google.com/presentation/d/1qcPdF4SYTq4VdTN_5w9GjviLRPr2HFtbuYMpbfLrc_M/edit#slide=id.cis4r2_u156_s4) | `gat_rates` | `lecture_core.py` | +| [Code for one valid categorical Euler step](https://docs.google.com/presentation/d/1qcPdF4SYTq4VdTN_5w9GjviLRPr2HFtbuYMpbfLrc_M/edit#slide=id.cis4r2_u157_s7) | `rate_step` | `lecture_core.py` | +| [Code for classic DFM generation](https://docs.google.com/presentation/d/1qcPdF4SYTq4VdTN_5w9GjviLRPr2HFtbuYMpbfLrc_M/edit#slide=id.cis4r2_u158_s12) | `gat_sample` | `lecture_core.py` | +| [Code for Dirichlet posterior training](https://docs.google.com/presentation/d/1qcPdF4SYTq4VdTN_5w9GjviLRPr2HFtbuYMpbfLrc_M/edit#slide=id.cis4r2_u159_s7) | `dirichlet_loss` | `lecture_core.py` | +| [Code for the Dirichlet velocity module](https://docs.google.com/presentation/d/1qcPdF4SYTq4VdTN_5w9GjviLRPr2HFtbuYMpbfLrc_M/edit#slide=id.cis4r2_u160_s9) | `dirichlet_velocity` | `lecture_core.py` | +| [Code for a Fisher conditional path](https://docs.google.com/presentation/d/1qcPdF4SYTq4VdTN_5w9GjviLRPr2HFtbuYMpbfLrc_M/edit#slide=id.cis4r2_u161_s9) | `fisher_path` | `lecture_core.py` | +| [Code for Fisher velocity regression](https://docs.google.com/presentation/d/1qcPdF4SYTq4VdTN_5w9GjviLRPr2HFtbuYMpbfLrc_M/edit#slide=id.cis4r2_u162_s9) | `fisher_loss` | `lecture_core.py` | +| [Code for the Gumbel conditional path](https://docs.google.com/presentation/d/1qcPdF4SYTq4VdTN_5w9GjviLRPr2HFtbuYMpbfLrc_M/edit#slide=id.cis4r2_u163_s9) | `gumbel_path` | `lecture_core.py` | +| [Code for Gumbel velocity regression](https://docs.google.com/presentation/d/1qcPdF4SYTq4VdTN_5w9GjviLRPr2HFtbuYMpbfLrc_M/edit#slide=id.cis4r2_u164_s7) | `gumbel_loss` | `lecture_core.py` | +| [Code for objective scores modify the base rates](https://docs.google.com/presentation/d/1qcPdF4SYTq4VdTN_5w9GjviLRPr2HFtbuYMpbfLrc_M/edit#slide=id.cis4r2_u165_s12) | `mog_rates` | `lecture_core.py` | +| [Code for continuous rectified-flow training](https://docs.google.com/presentation/d/1qcPdF4SYTq4VdTN_5w9GjviLRPr2HFtbuYMpbfLrc_M/edit#slide=id.cis4r2_u166_s8) | `rectified_loss` | `lecture_core.py` | +| [Code for ReDi replaces the endpoint pairing](https://docs.google.com/presentation/d/1qcPdF4SYTq4VdTN_5w9GjviLRPr2HFtbuYMpbfLrc_M/edit#slide=id.cis4r2_u167_s6) | `redi_pairs` | `lecture_core.py` | +| [Code for the preference scalarization](https://docs.google.com/presentation/d/1qcPdF4SYTq4VdTN_5w9GjviLRPr2HFtbuYMpbfLrc_M/edit#slide=id.cis4r2_u168_s3) | `maxmin` | `lecture_core.py` | +| [Code for enumerate every single-base DNA edit](https://docs.google.com/presentation/d/1qcPdF4SYTq4VdTN_5w9GjviLRPr2HFtbuYMpbfLrc_M/edit#slide=id.cis4r2_u169_s3) | `dna_neighbors` | `lecture_core.py` | +| [Code for a normalized locally balanced proposal](https://docs.google.com/presentation/d/1qcPdF4SYTq4VdTN_5w9GjviLRPr2HFtbuYMpbfLrc_M/edit#slide=id.cis4r2_u170_s6) | `proposal` | `lecture_core.py` | +| [Code for correct the refinement proposal](https://docs.google.com/presentation/d/1qcPdF4SYTq4VdTN_5w9GjviLRPr2HFtbuYMpbfLrc_M/edit#slide=id.cis4r2_u171_s12) | `mh_refine` | `lecture_core.py` | +| [Code for run a complete Gat DFM experiment](https://docs.google.com/presentation/d/1qcPdF4SYTq4VdTN_5w9GjviLRPr2HFtbuYMpbfLrc_M/edit#slide=id.cis4r2_u181_s5) | `complete experiment` | `lecture_core.py` and `run.py` | diff --git a/lecture_5/common.py b/lecture_5/common.py new file mode 100644 index 0000000000000000000000000000000000000000..277eca7d7ad53417be3e3ad2bdba319714c9386c --- /dev/null +++ b/lecture_5/common.py @@ -0,0 +1,134 @@ +"""CPU teaching utilities shared by the two lecture folders.""" +import csv +import json +import random +from pathlib import Path +import numpy as np +import torch +from torch import nn +from lecture_core import DNA, K, MASK, encode + + +def seed_all(seed): + random.seed(seed) + np.random.seed(seed) + torch.manual_seed(seed) + torch.set_num_threads(1) + + +def decode(tokens): + alphabet = 'ACGTm' + return [''.join(alphabet[int(i)] for i in row) for row in tokens] + + +def make_data(count=512, length=8, seed=7): + """Four synthetic motif families, with independent 8% base mutations.""" + rng = np.random.default_rng(seed) + motifs = ['ACGT', 'CGTA', 'TATA', 'GCGC'] + strings, labels = [], [] + for _ in range(count): + motif = motifs[int(rng.integers(4))] + seq = list((motif * ((length + 3) // 4))[:length]) + for j in range(length): + if rng.random() < .08: + seq[j] = 'ACGT'[int(rng.integers(4))] + strings.append(''.join(seq)) + labels.append(int(sum(x in 'GC' for x in seq) / length >= .6)) + return encode(strings), torch.tensor(labels) + + +def load_data(path, length): + if path is None: + return make_data(length=length) + with open(path) as f: + rows = list(csv.DictReader(f, delimiter='\t')) + strings = [r['sequence'].strip().upper() for r in rows] + if not strings or any(len(s) != length or set(s) - set('ACGT') for s in strings): + raise ValueError('All DNA sequences must contain only A/C/G/T and have --length bases.') + labels = [int(r.get('label', sum(c in 'GC' for c in s) / length >= .6)) + for r, s in zip(rows, strings)] + if set(labels) - {0, 1}: + raise ValueError('Labels must be 0 or 1.') + return encode(strings), torch.tensor(labels) + + +class ConditionalDNA(DNA): + """The slide network plus a label embedding; label 2 means unconditional.""" + def __init__(self, width=32, max_len=64): + super().__init__(width, max_len) + self.condition = nn.Embedding(3, width) + + def forward(self, z, t=None, label=None): + h = self.token(z) if z.ndim == 2 else self.soft(z) + pos = torch.arange(z.shape[1], device=z.device) + h = h + self.position(pos)[None] + if t is not None: + h = h + self.time(t[:, None])[:, None] + if label is None: + label = torch.full((len(z),), 2, dtype=torch.long, device=z.device) + h = h + self.condition(label)[:, None] + return self.output(self.context(h)) + + +def objectives(tokens): + """Both toy objectives are maximized: GC fraction and ATAT agreement.""" + onehot = torch.nn.functional.one_hot(tokens.long(), 4).float() + return soft_objectives(onehot) + + +def soft_objectives(z): + gc = (z[..., 1] + z[..., 2]).mean(-1) + motif = torch.tensor([0, 3], device=z.device).repeat((z.shape[1] + 1) // 2)[:z.shape[1]] + match = z.gather(-1, motif[None, :, None].expand(z.shape[0], -1, 1)).squeeze(-1).mean(-1) + return torch.stack((gc, match), -1) + + +def metrics(tokens): + strings = decode(tokens) + score = objectives(tokens).mean(0) + return dict(valid_dna=all(set(s) <= set('ACGT') for s in strings), + unique_fraction=len(set(strings)) / len(strings), + mean_gc=float(score[0]), mean_atat_match=float(score[1])) + + +def save_run(out, config, losses, samples, extra=None): + out = Path(out) + out.mkdir(parents=True, exist_ok=True) + (out / ('sample_config.json' if config.get('mode') == 'sample' else 'config.json')).write_text(json.dumps(config, indent=2)) + if losses or not (out / 'losses.json').exists(): + (out / 'losses.json').write_text(json.dumps(losses, indent=2)) + (out / 'samples.txt').write_text('\n'.join(decode(samples)) + '\n') + report = {'metrics': metrics(samples), **(extra or {})} + (out / 'report.json').write_text(json.dumps(report, indent=2)) + print(json.dumps({'output': str(out), **report}, indent=2)) + return report + + +def optimize(model, loss_fn, data, steps, batch_size=32, lr=.002): + optimizer = torch.optim.AdamW(model.parameters(), lr=lr) + losses = [] + model.train() + for step in range(steps): + idx = torch.randint(len(data), (batch_size,)) + optimizer.zero_grad(set_to_none=True) + loss = loss_fn(model, data[idx], idx) + if not torch.isfinite(loss): + raise FloatingPointError(f'Nonfinite loss at step {step}') + loss.backward() + torch.nn.utils.clip_grad_norm_(model.parameters(), 1.) + optimizer.step() + losses.append(float(loss.detach())) + model.eval() + return losses + + +def project_simplex(z, eps=1e-6): + """Euclidean simplex projection, followed by a small interior floor.""" + sorted_z = z.sort(-1, descending=True).values + cssv = sorted_z.cumsum(-1) - 1. + k = torch.arange(1, z.shape[-1] + 1, dtype=z.dtype, device=z.device) + active = sorted_z - cssv / k > 0 + rho = active.sum(-1, keepdim=True).clamp_min(1) + theta = cssv.gather(-1, rho - 1) / rho + result = (z - theta).clamp_min(eps) + return result / result.sum(-1, keepdim=True) diff --git a/lecture_5/data/README.md b/lecture_5/data/README.md new file mode 100644 index 0000000000000000000000000000000000000000..34c37c778eaa416a7cf9568f6e382f5771601b1c --- /dev/null +++ b/lecture_5/data/README.md @@ -0,0 +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. diff --git a/lecture_5/data/dna_train.tsv b/lecture_5/data/dna_train.tsv new file mode 100644 index 0000000000000000000000000000000000000000..e90e1227dd16f4e0f12e30bf40d10c531030dcc2 --- /dev/null +++ b/lecture_5/data/dna_train.tsv @@ -0,0 +1,257 @@ +sequence label +TAAACTTA 0 +ACGCTCCC 1 +CGTACGTA 0 +ACGTACGT 0 +CGTACGTA 0 +ACGTAAGT 0 +TATAAATA 0 +ACGTACGT 0 +GCGCGCGC 1 +GCGCGCGC 1 +TATATACA 0 +ACGTACGT 0 +ACGTACGT 0 +ACGCGCGC 1 +TCTATATA 0 +GCGCGCGC 1 +CGTACGTA 0 +ACGTACGT 0 +CGTACGTA 0 +GAGCGCGC 1 +TATATATA 0 +ACGTACGT 0 +GCGCGCGC 1 +ACGTACGT 0 +GCGCACGC 1 +CGTACGTA 0 +GGTACGTA 0 +CTTACGTA 0 +ACGTACGT 0 +AATATATA 0 +CGTACGTA 0 +TATATATA 0 +GCGCGCGC 1 +GCGCGCGC 1 +GCGCTCGC 1 +TCTATCTA 0 +GCGCGCGC 1 +CAAATATA 0 +TCGCGCGC 1 +ACGTACGT 0 +AAGTACGT 0 +GCGCGCGC 1 +CGCGCTTA 1 +TACATATA 0 +TATATATA 0 +CGTACGTA 0 +ACGTACGT 0 +GCGTGCGC 1 +ACGTACCT 0 +ACATACAT 0 +ATGGACGT 0 +GCGCGCGC 1 +TATAAATA 0 +CGTACGTA 0 +ATTTACGA 0 +CGTACGTA 0 +GCACGCGC 1 +GCGCGCGC 1 +TTTATATA 0 +ACGTACGT 0 +GCGCGCGT 1 +GCGCGCGC 1 +TATATATA 0 +CTTACGTA 0 +CGTACGTA 0 +GCGCGCGC 1 +CGTACGTA 0 +TGTTCGTA 0 +ACGTACCT 0 +ACGTACGC 1 +GCGCGCGC 1 +CTTATGTA 0 +CGTGCGTG 1 +ACGTACGT 0 +GCGCCCAC 1 +CGTACGTC 1 +TATATATA 0 +GGGCGCGC 1 +TATATATA 0 +TCGCGCGC 1 +ACGTACGT 0 +TATATCTA 0 +GCGCGCGC 1 +TATATATC 0 +GCACGCGC 1 +ACGTACGT 0 +CGTACGTA 0 +TATATATA 0 +CGTACGTA 0 +GGGCGCGC 1 +TAGATATA 0 +CGTACATA 0 +ACGTACGT 0 +TAGAGATA 0 +ACGTACGT 0 +ACGTACGT 0 +TATATATA 0 +ACTTACGT 0 +CGTACCTA 0 +CGTACGTA 0 +GCGCGCGC 1 +GCTCGCGC 1 +GGGCGCCC 1 +TATATATA 0 +ACGTACGT 0 +ACGTACGT 0 +GCGCCCGC 1 +TATATATA 0 +ACGTACGT 0 +TATAGATA 0 +TATATATT 0 +CAGACGTA 0 +CGTACGTA 0 +CGTAGGTA 0 +GCACGACC 1 +CGTACGAA 0 +GTGCGCGC 1 +CGTACGTA 0 +TATATATA 0 +CGTACGTA 0 +GCACGCGC 1 +ACATACGA 0 +CGTAGGTA 0 +ACGTACAT 0 +TATATATA 0 +GATATATA 0 +ACGTACGT 0 +TCTAGATA 0 +CGTACGTA 0 +CGTACGTA 0 +TATATATA 0 +GCGCGCGC 1 +ACGTACGT 0 +TATATATA 0 +CGTACGTA 0 +ACGTACGT 0 +ACGTTCGT 0 +TATATATA 0 +GCGCGCGC 1 +ACGTACGT 0 +TATATATA 0 +ACGCACGT 1 +ACGTGCGT 1 +CGTACGTA 0 +GTGCGCGC 1 +GCGCGCGC 1 +CGTACGTA 0 +GCGCGCGC 1 +TATATATA 0 +GCGCGCGC 1 +CGTACGTA 0 +CGTACGTA 0 +ACGTACGT 0 +GCGCGCGC 1 +CGTACGTA 0 +GCGCGGGA 1 +TATATATA 0 +ACGTACGT 0 +ACGTACGT 0 +GCGCGCGT 1 +ACTTACGT 0 +CGTACGTA 0 +ACGTACGT 0 +TAAATAAA 0 +GCGTGCGC 1 +GCGCGCAC 1 +ACGTACGT 0 +TATATATA 0 +TAAATATA 0 +CGTACGTA 0 +ACGTACGT 0 +TGTACATA 0 +CGTACGGA 1 +CGTCCGTA 1 +GCGCGCGC 1 +ACGTACGT 0 +TATATATA 0 +CGTACGTA 0 +TGTACGTA 0 +ACGTACGT 0 +ACGTACGT 0 +GCGTATGT 0 +ACGTACGT 0 +CGGAAGTC 1 +TATATATA 0 +TATATATA 0 +CGTGCGTA 1 +GCGCGCGC 1 +CGTACGTA 0 +ACGTACGT 0 +ACGTACGT 0 +GCGCGCGC 1 +ACGTACGT 0 +ACGTACGT 0 +ACGTACGT 0 +CGTACGTA 0 +GCGCGCGC 1 +ACGTAGGT 0 +TATATATA 0 +ACATACGT 0 +CCGTACGT 1 +AAGTACGT 0 +CGTACGTA 0 +GCACGCGC 1 +TACATATA 0 +TATATATA 0 +GCGCGCGC 1 +ACGTACGT 0 +GCGTACGT 1 +TATATATA 0 +TATATATA 0 +GCGCGCGC 1 +CGTGCTTA 0 +ACGTACGT 0 +TATATATA 0 +CGTACGTA 0 +TATATATA 0 +GTGCGCGC 1 +CGTACGTA 0 +GCGTACGT 1 +GCGCGCGC 1 +TATATATA 0 +TATATGTA 0 +TATATATA 0 +CATACTAA 0 +ACGAACGT 0 +CGTACGTA 0 +GCGAGCGC 1 +TATATATA 0 +ACGAACCT 0 +GCGCGCGC 1 +CGTCCGTG 1 +GCGCGAGC 1 +CGTACCTA 0 +CGTACGTA 0 +GCGCGCGC 1 +TATATATA 0 +ACGTACGT 0 +ACGTAGGT 0 +GGTACGTA 0 +CGTACGTA 0 +GCGCGCGA 1 +CGTACGTA 0 +CGTACGGA 1 +ACGTACGT 0 +GCGCGCGC 1 +GCGCGCGC 1 +TATATATA 0 +ACGTACGT 0 +GCGCGCGC 1 +ACGTACGT 0 +ATGGACGT 0 +CGTAAGTA 0 +TATATATA 0 +TATATATA 0 +TATATATA 0 diff --git a/lecture_5/examples/areuredi.py b/lecture_5/examples/areuredi.py new file mode 100644 index 0000000000000000000000000000000000000000..4ee08f90348710561f6a631094d41f3b6a6057c4 --- /dev/null +++ b/lecture_5/examples/areuredi.py @@ -0,0 +1,7 @@ +"""Complete areuredi training and generation example.""" +import sys +from pathlib import Path +sys.path.insert(0, str(Path(__file__).resolve().parents[1])) +from run import main +if __name__ == '__main__': + main(['--method', 'areuredi', '--out', 'outputs/areuredi'] + sys.argv[1:]) diff --git a/lecture_5/examples/dirichlet.py b/lecture_5/examples/dirichlet.py new file mode 100644 index 0000000000000000000000000000000000000000..9a0df582b782b24de2526ae8715c4419d0d884e6 --- /dev/null +++ b/lecture_5/examples/dirichlet.py @@ -0,0 +1,7 @@ +"""Complete dirichlet training and generation example.""" +import sys +from pathlib import Path +sys.path.insert(0, str(Path(__file__).resolve().parents[1])) +from run import main +if __name__ == '__main__': + main(['--method', 'dirichlet', '--out', 'outputs/dirichlet'] + sys.argv[1:]) diff --git a/lecture_5/examples/fisher.py b/lecture_5/examples/fisher.py new file mode 100644 index 0000000000000000000000000000000000000000..4f6095d22e59e0d2c4f9d9aff4c08c99307df55b --- /dev/null +++ b/lecture_5/examples/fisher.py @@ -0,0 +1,7 @@ +"""Complete fisher training and generation example.""" +import sys +from pathlib import Path +sys.path.insert(0, str(Path(__file__).resolve().parents[1])) +from run import main +if __name__ == '__main__': + main(['--method', 'fisher', '--out', 'outputs/fisher'] + sys.argv[1:]) diff --git a/lecture_5/examples/gat.py b/lecture_5/examples/gat.py new file mode 100644 index 0000000000000000000000000000000000000000..336955c0f59624c6f5e14e46236dbb61288242ab --- /dev/null +++ b/lecture_5/examples/gat.py @@ -0,0 +1,7 @@ +"""Complete gat training and generation example.""" +import sys +from pathlib import Path +sys.path.insert(0, str(Path(__file__).resolve().parents[1])) +from run import main +if __name__ == '__main__': + main(['--method', 'gat', '--out', 'outputs/gat'] + sys.argv[1:]) diff --git a/lecture_5/examples/gumbel.py b/lecture_5/examples/gumbel.py new file mode 100644 index 0000000000000000000000000000000000000000..d92d79e5efe07fb514cfcae571d651da62cea55c --- /dev/null +++ b/lecture_5/examples/gumbel.py @@ -0,0 +1,7 @@ +"""Complete gumbel training and generation example.""" +import sys +from pathlib import Path +sys.path.insert(0, str(Path(__file__).resolve().parents[1])) +from run import main +if __name__ == '__main__': + main(['--method', 'gumbel', '--out', 'outputs/gumbel'] + sys.argv[1:]) diff --git a/lecture_5/examples/mog_dfm.py b/lecture_5/examples/mog_dfm.py new file mode 100644 index 0000000000000000000000000000000000000000..e5b8b58ab0641775367839df9ae600f131b4bbea --- /dev/null +++ b/lecture_5/examples/mog_dfm.py @@ -0,0 +1,7 @@ +"""Complete mog-dfm training and generation example.""" +import sys +from pathlib import Path +sys.path.insert(0, str(Path(__file__).resolve().parents[1])) +from run import main +if __name__ == '__main__': + main(['--method', 'mog-dfm', '--out', 'outputs/mog-dfm'] + sys.argv[1:]) diff --git a/lecture_5/examples/rectified.py b/lecture_5/examples/rectified.py new file mode 100644 index 0000000000000000000000000000000000000000..521d784ed4985992391d73cdfab4cc04ff001c4f --- /dev/null +++ b/lecture_5/examples/rectified.py @@ -0,0 +1,7 @@ +"""Complete rectified training and generation example.""" +import sys +from pathlib import Path +sys.path.insert(0, str(Path(__file__).resolve().parents[1])) +from run import main +if __name__ == '__main__': + main(['--method', 'rectified', '--out', 'outputs/rectified'] + sys.argv[1:]) diff --git a/lecture_5/examples/redi.py b/lecture_5/examples/redi.py new file mode 100644 index 0000000000000000000000000000000000000000..2fd370ba706d9e45c6115e4ccec746343e8d76be --- /dev/null +++ b/lecture_5/examples/redi.py @@ -0,0 +1,7 @@ +"""Complete redi training and generation example.""" +import sys +from pathlib import Path +sys.path.insert(0, str(Path(__file__).resolve().parents[1])) +from run import main +if __name__ == '__main__': + main(['--method', 'redi', '--out', 'outputs/redi'] + sys.argv[1:]) diff --git a/lecture_5/flows.py b/lecture_5/flows.py new file mode 100644 index 0000000000000000000000000000000000000000..e251d41195543820ccdb7a22258dd30d76c6e6d1 --- /dev/null +++ b/lecture_5/flows.py @@ -0,0 +1,188 @@ +"""Full teaching samplers for simplex, rectification, and guided jump flows.""" +import math +import numpy as np +import torch +import torch.nn.functional as F +from lecture_core import (K, draw, gat_rates, rate_step, gat_sample, + dirichlet_velocity, dna_neighbors, mh_refine) +from common import project_simplex, objectives, soft_objectives, decode + + +@torch.no_grad() +def simplex_sample(model, method, batch, length, steps=100, a_max=8., + beta=1., tau_max=4., guidance=0.): + """Finite Euler discretization with explicit simplex/orthant projection. + + Projection is a numerical safeguard, not the exact continuous dynamics. + Gumbel has an approximate finite-temperature base and endpoint. + """ + base = torch.distributions.Dirichlet(torch.ones(K)) + z = base.sample((batch, length)) + corrections = 0 + if method == 'gumbel': + uniform = torch.rand(batch, length, K).clamp(1e-6, 1-1e-6) + g = -(-uniform.log()).log() + z = (g / (beta * tau_max)).softmax(-1) + y = z.sqrt() if method == 'fisher' else z + end = a_max if method == 'dirichlet' else 1. + dt = end / steps + for i in range(steps): + time = i * dt + t = torch.full((batch,), time) + if method == 'dirichlet': + posterior = model(y, t).softmax(-1) + velocity = dirichlet_velocity(y, posterior, time) + elif method == 'fisher': + raw = model(y, t) + velocity = raw - y * (raw * y).sum(-1, keepdim=True) + else: + raw = model(y, t) + velocity = raw - raw.mean(-1, keepdim=True) + if guidance: + # A differentiable toy objective on probabilities, not a trained + # biochemical predictor. Choose a direction toward higher GC. + with torch.enable_grad(): + state = y.detach().requires_grad_(True) + probability = state.square() if method == 'fisher' else state + score = soft_objectives(probability)[:, 0].sum() + gradient = torch.autograd.grad(score, state)[0] + if method == 'fisher': + gradient -= y * (gradient * y).sum(-1, keepdim=True) + else: + gradient -= gradient.mean(-1, keepdim=True) + velocity += guidance * gradient + candidate = y + dt * velocity + if not torch.isfinite(candidate).all(): + raise FloatingPointError('Nonfinite simplex integration state.') + corrections += int((candidate < 0).any(-1).sum()) + if method == 'fisher': + candidate = candidate.clamp_min(1e-6) + y = candidate / candidate.norm(dim=-1, keepdim=True) + else: + y = project_simplex(candidate) + prob = y.square() if method == 'fisher' else y + result = prob.argmax(-1) + return result, {'negative_coordinate_corrections': corrections, + 'max_probability_sum_error': float((prob.sum(-1)-1).abs().max()), + 'finite_endpoint': True, 'decode': 'argmax'} + + +def rectified_training_loss(model, clean): + source = torch.distributions.Dirichlet(torch.ones(K)).sample(clean.shape) + target = F.one_hot(clean, K).float() + t = torch.rand(len(clean)) + z = (1-t[:, None, None])*source + t[:, None, None]*target + raw = model(z, t) + prediction = raw - raw.mean(-1, keepdim=True) + return (prediction-(target-source)).square().sum(-1).mean() + + +@torch.no_grad() +def make_teacher_pairs(teacher, count, length, steps=60, chunk=64): + sources, targets = [], [] + for start in range(0, count, chunk): + batch = min(chunk, count-start) + source = torch.randint(K, (batch, length)) + target = gat_sample(teacher, batch, length, steps, source=source) + sources.append(source); targets.append(target) + return torch.cat(sources), torch.cat(targets) + + +def mog_multiplier(changes, preference, strength=1., angle=math.pi/3): + """Rank/direction score and acute-cone test; all objectives maximize.""" + from scipy.stats import rankdata + ranks = rankdata(changes, axis=0, method='average') / len(changes) + rank_score = ranks.mean(1) + norm = np.linalg.norm(changes, axis=1) + cosine = (changes @ preference) / np.maximum(norm*np.linalg.norm(preference), 1e-12) + direction = changes @ preference + # Formula matches the lecture: z-standardized rank and direction terms. + def standardize(x): + return (x-x.mean()) / max(float(x.std()), 1e-12) + score = standardize(rank_score) + standardize(direction) + keep = (norm > 0) & (cosine >= np.cos(angle)) + return strength * np.exp(np.clip(score, -20, 20)) * keep + + +@torch.no_grad() +def mog_sample(model, batch, length, steps=50, preference=(.7, .3), strength=1.): + """Evaluate all single-base edits, reweight rates, then use adaptive Euler. + + All-position scoring differs from the paper's random-position loop. + Rank normalization is per position. The adaptive cone stays acute; + no outside-cone fallback is used, so the local theorem remains valid. + """ + z = torch.randint(K, (batch, length)) + t = 0. + rejected = 0 + angle = math.pi/3 + ema = .5 + count = 0 + while t < 1-1e-6: + tb = torch.full((batch,), t) + rates = gat_rates(model(z, tb).softmax(-1), z, tb) + rejected_step = 0 + for b in range(batch): + for i in range(length): + candidates, destinations = [], [] + for a in range(K): + if a == int(z[b, i]): + continue + y = z[b].clone(); y[i] = a + candidates.append(y); destinations.append(a) + changes = (objectives(torch.stack(candidates)) - objectives(z[b:b+1])).numpy() + multiplier = mog_multiplier(changes, np.array(preference), strength, angle) + rejected_step += int(np.sum(multiplier == 0)) + for a, value in zip(destinations, multiplier): + rates[b, i, a] *= float(value) + rejected += rejected_step + rejection = rejected_step / (batch * length * (K-1)) + ema = .9 * ema + .1 * rejection + angle = float(np.clip(angle * np.exp(.02 * (ema-.5)), + np.deg2rad(10), np.deg2rad(89))) + max_exit = float(rates.sum(-1).max()) + h = min(1./steps, 1-t, .5/max(max_exit, 1e-8)) + z = rate_step(z, rates, h) + t += h + count += 1 + if count > 50000: + raise RuntimeError('MOG integration failed to advance.') + return z, {'filtered_candidates': rejected, 'adaptive_steps': count, + 'preference': list(preference), 'final_cone_angle_degrees': float(np.rad2deg(angle)), + 'guarantee': 'local weighted progress under the acute-cone assumptions'} + + +def reference_model(data, pseudocount=.5): + """A tractable positive product reference for exact MH ratios.""" + counts = F.one_hot(data, K).float().sum(0) + pseudocount + probability = counts / counts.sum(-1, keepdim=True) + def log_probability(sequence): + idx = torch.tensor(['ACGT'.index(c) for c in sequence]) + return float(probability.log()[torch.arange(len(idx)), idx].sum()) + return probability, log_probability + + +def refine_areuredi(tokens, data, steps=100, preference=(.5,.5), eta_max=8., seed=7): + """ReDi endpoint refinement with normalized proposals and exact MH ratios. + + The reference is explicitly fitted and evaluable; it is not a denoiser + joint density. Annealing is finite, so no global-optimum claim is made. + """ + from lecture_core import encode + _, log_ref = reference_model(data) + weight = np.array(preference) + def score(sequence): + values = objectives(encode([sequence]))[0].numpy() + return float(np.min(weight * values)) + rng = np.random.default_rng(seed) + strings = decode(tokens) + changes = 0 + for step in range(steps): + eta = eta_max * (step+1) / steps + for i, sequence in enumerate(strings): + new = mh_refine(sequence, score, log_ref, eta, rng) + changes += int(new != sequence) + strings[i] = new + return encode(strings), {'accepted_edits': changes, 'eta_final': eta_max, + 'reference': 'smoothed position-wise empirical DNA frequencies', + 'scope': 'finite annealing demonstration, not an equilibrium certificate'} diff --git a/lecture_5/lecture_core.py b/lecture_5/lecture_core.py new file mode 100644 index 0000000000000000000000000000000000000000..b6749a1ae69426823078585397f4b4f40d7fb5a3 --- /dev/null +++ b/lecture_5/lecture_core.py @@ -0,0 +1,203 @@ +"""Small teaching implementations; illustrative data, not paper reproductions.""" + +import math + +import numpy as np + +import torch + +from torch import nn + +import torch.nn.functional as F + +from scipy.special import betainc, beta + +from scipy.stats import rankdata + +K, MASK = 4, 4 + +ALPHABET = 'ACGT' + +def encode(strings): + return torch.tensor([[ALPHABET.index(c) for c in s] + for s in strings]) + +def draw(prob): + shape = prob.shape[:-1] + sample = torch.multinomial(prob.reshape(-1, K), 1) + return sample.reshape(shape) + +class DNA(nn.Module): + def __init__(self, width=32, max_len=64): + super().__init__() + self.token = nn.Embedding(K + 1, width) + self.soft = nn.Linear(K, width) + self.position = nn.Embedding(max_len, width) + self.time = nn.Linear(1, width) + layer = nn.TransformerEncoderLayer( + width, 4, 2 * width, dropout=0., batch_first=True) + self.context = nn.TransformerEncoder(layer, 1) + self.output = nn.Linear(width, K) + + def forward(self, z, t=None): + h = self.token(z) if z.ndim == 2 else self.soft(z) + pos = torch.arange(z.shape[1], device=z.device) + h = h + self.position(pos)[None] + if t is not None: + h = h + self.time(t[:, None])[:, None] + return self.output(self.context(h)) + +def token_ce(logits, target): + return F.cross_entropy(logits.transpose(1, 2), + target, reduction='none') + +def train_step(model, optimizer, clean, loss_fn): + model.train() + optimizer.zero_grad() + loss = loss_fn(model, clean) + loss.backward() + optimizer.step() + return loss.item() + +def gat_loss(model, target, source=None): + if source is None: + source = torch.randint(K, target.shape) + t = torch.rand(target.shape[0]) + use_target = torch.rand(target.shape) < t[:, None] + z = torch.where(use_target, target, source) + logits = model(z, t) + return token_ce(logits, target).sum(1).mean() + +def gat_rates(prob, z, t): + rates = prob / (1. - t[:, None, None]) + return rates.scatter(-1, z[..., None], 0.) + +def rate_step(z, rates, h): + exit_rate = rates.sum(-1) + assert torch.all(h * exit_rate <= 1. + 1e-6) + prob = h * rates + prob.scatter_(-1, z[..., None], + (1. - h * exit_rate)[..., None]) + return draw(prob.clamp_min(0.)) + +@torch.no_grad() +def gat_sample(model, batch, length, steps=20, source=None): + model.eval() + z = (torch.randint(K, (batch, length)) + if source is None else source.clone()) + h = 1. / steps + for n in range(steps): + t = torch.full((batch,), n / steps) + prob = model(z, t).softmax(-1) + rates = gat_rates(prob, z, t) + z = rate_step(z, rates, h) + return z + +def dirichlet_loss(model, clean): + a = 8. * torch.rand(clean.shape[0]) + target = F.one_hot(clean, K).float() + concentration = 1. + a[:, None, None] * target + z = torch.distributions.Dirichlet(concentration).sample() + return token_ce(model(z, a), clean).sum(1).mean() + +def dirichlet_velocity(z, posterior, a): + r = z.detach().double().numpy().clip(1e-7, 1. - 1e-7) + da = 1e-4 + derivative = (betainc(1+a+da, K-1, r) + - betainc(1+a-da, K-1, r)) / (2*da) + density = r**a * (1-r)**(K-2) / beta(1+a, K-1) + c = torch.as_tensor(-derivative / ((1-r)*density), + dtype=z.dtype) + weight = posterior * c + return weight - z * weight.sum(-1, keepdim=True) + +def fisher_path(z0, target, t): + y0, y1 = z0.sqrt(), F.one_hot(target, K).float() + cosine = (y0 * y1).sum(-1, keepdim=True) + w = cosine.clamp(-1+1e-6, 1-1e-6).acos() + s = t[:, None, None] + y = (((1-s)*w).sin()*y0 + (s*w).sin()*y1) / w.sin() + velocity = w * (-((1-s)*w).cos()*y0 + + (s*w).cos()*y1) / w.sin() + return y, velocity + +def fisher_loss(model, clean): + base = torch.distributions.Dirichlet(torch.ones(K)) + z0 = base.sample(clean.shape) + t = torch.rand(clean.shape[0]) + y, target = fisher_path(z0, clean, t) + raw = model(y, t) + tangent = raw - y * (raw * y).sum(-1, keepdim=True) + return 4. * (tangent - target).square().sum(-1).mean() + +def gumbel_path(target, t, beta_noise=1., tau_max=4., decay=4.): + uniform = torch.rand(*target.shape, K).clamp(1e-6, 1-1e-6) + gumbel = -(-uniform.log()).log() + a = F.one_hot(target, K).float() + gumbel / beta_noise + tau = tau_max * (-decay*t[:, None, None]).exp() + z = (a / tau).softmax(-1) + velocity = (decay/tau) * z * (a - (z*a).sum(-1, keepdim=True)) + return z, velocity + +def gumbel_loss(model, clean): + t = torch.rand(clean.shape[0]) + z, target = gumbel_path(clean, t) + raw = model(z, t) + tangent = raw - raw.mean(-1, keepdim=True) + return (tangent - target).square().sum(-1).mean() + +def mog_rates(base, changes, preference, importance, + strength, angle): + ranks = rankdata(changes, axis=0, method='average') + ranks = ranks / len(changes) + norm = np.linalg.norm(changes, axis=1) + denom = np.maximum(norm*np.linalg.norm(preference), 1e-12) + cosine = (changes @ preference) / denom + direction = changes @ preference + standardize = lambda x: (x-x.mean()) / max(x.std(), 1e-12) + score = (standardize((ranks*importance).mean(1)) + + standardize(direction)) + keep = (norm > 0) & (cosine >= np.cos(angle)) + return base * np.exp(strength*score) * keep + +def rectified_loss(model, source, target): + t = torch.rand(source.shape[0]) + z = (1-t[:, None, None])*source + t[:, None, None]*target + conditional_velocity = target - source + raw = model(z, t) + tangent = raw - raw.mean(-1, keepdim=True) + return (tangent-conditional_velocity).square().sum(-1).mean() + +@torch.no_grad() +def redi_pairs(teacher, batch, length, teacher_steps=100): + source = torch.randint(K, (batch, length)) + target = gat_sample(teacher, batch, length, + teacher_steps, source=source) + return source, target + +def maxmin(scores, preference): + return np.min(scores * preference, axis=-1) + +def dna_neighbors(sequence): + return [sequence[:i] + b + sequence[i+1:] + for i, old in enumerate(sequence) + for b in ALPHABET if b != old] + +def proposal(sequence, score, eta): + candidates = dna_neighbors(sequence) + log_weight = .5 * eta * np.array([score(y)-score(sequence) + for y in candidates]) + weight = np.exp(log_weight - log_weight.max()) + return candidates, weight / weight.sum() + +def mh_refine(sequence, score, log_ref, eta, rng): + candidates, forward = proposal(sequence, score, eta) + index = rng.choice(len(candidates), p=forward) + candidate = candidates[index] + reverse_candidates, reverse = proposal(candidate, score, eta) + q_reverse = reverse[reverse_candidates.index(sequence)] + log_ratio = log_ref(candidate) - log_ref(sequence) + log_ratio += eta * (score(candidate)-score(sequence)) + log_ratio += np.log(q_reverse) - np.log(forward[index]) + accept = np.log(rng.random()) < min(0., log_ratio) + return candidate if accept else sequence diff --git a/lecture_5/numerical_examples.py b/lecture_5/numerical_examples.py new file mode 100644 index 0000000000000000000000000000000000000000..0364b5e8001d7c260dc9a9b1366481237e7f93bb --- /dev/null +++ b/lecture_5/numerical_examples.py @@ -0,0 +1,26 @@ +"""Transparent DNA path, rate, geometry, and guidance calculations.""" +import math +import torch +from lecture_core import encode, gat_rates, fisher_path, dirichlet_velocity + +print('Alphabet order: A C G T.') +source=encode(['TGCA']);target=encode(['ACGT']);u=torch.tensor([[.1,.7,.2,.9]]) +print('Gat sampled t=.25 intermediate:',torch.where(u<.25,target,source).tolist(),'= AGGA') +print('Gat target probability loss:',-math.log(.7*.6*.8*.5)) +rates=gat_rates(torch.tensor([[[.1,.6,.2,.1]]]),torch.tensor([[2]]),torch.tensor([.5])) +p=.1*rates;p[...,2]=1-.1*rates.sum(-1) +print('Gat Euler probabilities from G:',p.flatten().tolist()) +z=torch.tensor([[[.1,.2,.4,.3]]]);posterior=torch.tensor([[[0.,0.,1.,0.]]]) +v=dirichlet_velocity(z,posterior,3.) +print('Dirichlet conditional G velocity:',v.flatten().tolist()) +y,v=fisher_path(torch.full((1,1,4),.25),torch.tensor([[2]]),torch.tensor([.5])) +print('Fisher midpoint probabilities:',y.square().flatten().tolist()) +print('Fisher spherical velocity:',v.flatten().tolist()) +a=torch.tensor([0.,math.log(2),math.log(4),0.]);z=a.softmax(-1) +print('Gumbel realized softmax at tau=1:',z.tolist()) +print('Same noisy logits at tau=.5:',(2*a).softmax(-1).tolist()) +print('Temperature derivative -1 velocity:',(z*(a-(z*a).sum())).tolist()) +for beta in [1.,4.,8.]: + print('DNA endpoint winner probability, beta=',beta,':',math.exp(beta)/(math.exp(beta)+3)) +print('ReDi independent matching-pair TC:',math.log(4),'nats; deterministic informative coupling TC: 0') +print('Preference (.8,.2), change (.3,-.1): weighted gain .22 despite second-objective loss') diff --git a/lecture_5/requirements.txt b/lecture_5/requirements.txt new file mode 100644 index 0000000000000000000000000000000000000000..629699a0a61cd18d3bae56bef8d774c89ce2cfb4 --- /dev/null +++ b/lecture_5/requirements.txt @@ -0,0 +1,3 @@ +torch>=2.2 +numpy>=1.24 +scipy>=1.10 diff --git a/lecture_5/run.py b/lecture_5/run.py new file mode 100644 index 0000000000000000000000000000000000000000..9bcf07a6ffe18fd44a0aa37d4e20f67f55a7ae0a --- /dev/null +++ b/lecture_5/run.py @@ -0,0 +1,115 @@ +#!/usr/bin/env python3 +"""Train and generate DNA with classic DFM before the simplex extensions.""" +import argparse +from pathlib import Path +import torch +from lecture_core import (DNA, gat_loss, gat_sample, dirichlet_loss, + fisher_loss, gumbel_loss) +from common import seed_all, load_data, optimize, save_run +from flows import (simplex_sample, rectified_training_loss, make_teacher_pairs, + mog_sample, refine_areuredi) + +METHODS = ['gat', 'dirichlet', 'fisher', 'gumbel', 'rectified', + 'redi', 'mog-dfm', 'areuredi'] + + +def main(argv=None): + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument('--method', choices=METHODS, default='gat') + parser.add_argument('--mode', choices=['train-sample', 'train', 'sample'], default='train-sample') + parser.add_argument('--train-steps', type=int, default=300) + parser.add_argument('--teacher-steps', type=int, default=300) + parser.add_argument('--sample-steps', type=int, default=100) + parser.add_argument('--refine-steps', type=int, default=100) + parser.add_argument('--pairs', type=int, default=256) + parser.add_argument('--batch-size', type=int, default=32) + parser.add_argument('--samples', type=int, default=32) + parser.add_argument('--length', type=int, default=8) + parser.add_argument('--width', type=int, default=32) + parser.add_argument('--guidance', type=float, default=0., help='Toy GC objective for simplex samplers') + parser.add_argument('--strength', type=float, default=1., help='MOG rate multiplier beta') + parser.add_argument('--preference', type=float, nargs=2, default=[.7, .3]) + parser.add_argument('--seed', type=int, default=7) + parser.add_argument('--data', help='TSV with DNA sequence column') + parser.add_argument('--out', default='outputs/gat') + args = parser.parse_args(argv) + if min(args.length, args.sample_steps, args.samples, args.batch_size, args.pairs) < 1: + parser.error('Lengths and sample/batch counts must be positive.') + if args.width % 4 or args.length > 64: + parser.error('Width must be divisible by four; length must be at most 64.') + if min(args.preference) <= 0 or args.strength <= 0: + parser.error('Preference weights and MOG strength must be positive.') + if args.mode != 'sample' and min(args.train_steps, args.teacher_steps) < 1: + parser.error('Training step counts must be positive.') + preference = [p / sum(args.preference) for p in args.preference] + seed_all(args.seed) + out = Path(args.out); out.mkdir(parents=True, exist_ok=True) + data, labels = load_data(args.data, args.length) + split = max(1, int(.8 * len(data))) + train = data[:split] + model = DNA(args.width) + losses = [] + report = {'method': args.method, 'data': 'synthetic DNA; not biological validation'} + loss_fns = {'gat': gat_loss, 'dirichlet': dirichlet_loss, + 'fisher': fisher_loss, 'gumbel': gumbel_loss, + 'rectified': rectified_training_loss, 'mog-dfm': gat_loss} + if args.mode == 'sample': + checkpoint = torch.load(out / 'checkpoint.pt', weights_only=True) + if checkpoint['method'] != args.method or checkpoint['length'] != args.length: + raise ValueError('Checkpoint method/length must match command arguments.') + model.load_state_dict(checkpoint['model']) + if 'reference_data' in checkpoint: + train = checkpoint['reference_data'] + else: + if args.method in ['redi', 'areuredi']: + teacher = DNA(args.width) + teacher_losses = optimize(teacher, lambda m,x,i: gat_loss(m,x), + train, args.teacher_steps, args.batch_size) + source, target = make_teacher_pairs(teacher, args.pairs, args.length, + min(args.sample_steps, 100)) + torch.save({'source': source, 'target': target}, out / 'teacher_pairs.pt') + losses = optimize(model, lambda m,x,i: gat_loss(m,x,source[i]), + target, args.train_steps, args.batch_size) + report['teacher_loss_last_20_mean'] = sum(teacher_losses[-20:]) / len(teacher_losses[-20:]) + report['paired_examples'] = args.pairs + report['redi_scope'] = 'one teacher-recoupling round; no guaranteed TC reduction' + else: + loss_fn = loss_fns[args.method] + losses = optimize(model, lambda m,x,i: loss_fn(m,x), train, + args.train_steps, args.batch_size) + checkpoint = {'model': model.state_dict(), 'method': args.method, + 'length': args.length, 'width': args.width} + if args.method == 'areuredi': + checkpoint['reference_data'] = train + torch.save(checkpoint, out / 'checkpoint.pt') + report['train_loss_first_20_mean'] = sum(losses[:20]) / len(losses[:20]) + report['train_loss_last_20_mean'] = sum(losses[-20:]) / len(losses[-20:]) + if args.method in loss_fns and len(data[split:]): + with torch.no_grad(): + report['validation_loss_one_mc_draw'] = float(loss_fns[args.method](model, data[split:])) + if args.mode == 'train': + return save_run(out, vars(args), losses, train[:args.samples], + {**report, 'sample_file_contains': 'training examples; generation not requested'}) + model.eval() + if args.method in ['dirichlet', 'fisher', 'gumbel', 'rectified']: + samples, extra = simplex_sample(model, args.method, args.samples, + args.length, args.sample_steps, guidance=args.guidance) + report.update(extra) + if args.method == 'gumbel': + report['gumbel_scope'] = ('beta=1, tau_max=4, decay=4; finite noisy endpoint ' + 'and approximate data-independent base; no exact data endpoint claim') + elif args.method == 'mog-dfm': + samples, extra = mog_sample(model, args.samples, args.length, + args.sample_steps, preference, args.strength) + report.update(extra) + else: + samples = gat_sample(model, args.samples, args.length, args.sample_steps) + if args.method == 'areuredi': + samples, extra = refine_areuredi(samples, train, args.refine_steps, + preference, seed=args.seed) + report.update(extra) + return save_run(out, vars(args), losses, samples, report) + + +if __name__ == '__main__': + main() diff --git a/lecture_5/run_all.py b/lecture_5/run_all.py new file mode 100644 index 0000000000000000000000000000000000000000..f16d76e0a291f12897818e38a417cc232ec2f1be --- /dev/null +++ b/lecture_5/run_all.py @@ -0,0 +1,15 @@ +#!/usr/bin/env python3 +"""Run every complete training example; pass --quick for a small CPU check.""" +import argparse +from run import main +parser = argparse.ArgumentParser() +parser.add_argument('--quick', action='store_true') +args = parser.parse_args() +methods = ['gat', 'dirichlet', 'fisher', 'gumbel', 'rectified', 'redi', 'mog-dfm', 'areuredi'] +for method in methods: + command = ['--method', method, '--out', 'outputs/' + method] + if args.quick: + command += ['--train-steps', '20', '--samples', '4', '--length', '4', + '--batch-size', '8', '--sample-steps', '20'] + command += ['--teacher-steps', '20', '--pairs', '32', '--refine-steps', '12'] + main(command) diff --git a/lecture_5/tests/test_mathematics.py b/lecture_5/tests/test_mathematics.py new file mode 100644 index 0000000000000000000000000000000000000000..580fab5a0dd5b4e4d26fdc7b822a497470cd0938 --- /dev/null +++ b/lecture_5/tests/test_mathematics.py @@ -0,0 +1,66 @@ +import itertools +import unittest +import sys +from pathlib import Path +sys.path.insert(0, str(Path(__file__).resolve().parents[1])) +import numpy as np +import torch +from lecture_core import fisher_path, gumbel_path, proposal +from flows import reference_model + + +class Mathematics(unittest.TestCase): + def test_gat_posterior_average_matches_full_master_equation(self): + states = list(itertools.product(range(4), repeat=2)) + targets = [(a,a) for a in range(4)] + def marginal(t): + p=np.zeros(16); numerator=np.zeros((16,16)) + for a in states: + for b in targets: + for xi,x in enumerate(states): + probability=np.prod([(1-t)*(x[j]==a[j])+t*(x[j]==b[j]) for j in range(2)])/64 + p[xi]+=probability + for j in range(2): + if x[j]!=b[j]: + y=list(x);y[j]=b[j];yi=states.index(tuple(y)) + numerator[xi,yi]+=probability/(1-t) + rates=numerator/p[:,None] + np.fill_diagonal(rates,-rates.sum(1)) + return p,rates + t=.4;h=1e-5 + p,r=marginal(t) + derivative=(marginal(t+h)[0]-marginal(t-h)[0])/(2*h) + np.testing.assert_allclose(p@r,derivative,atol=1e-9) + + def test_fisher_midpoint_and_tangent(self): + y,v=fisher_path(torch.full((1,1,4),.25),torch.tensor([[2]]),torch.tensor([.5])) + torch.testing.assert_close(y.square(),torch.tensor([[[1/12,1/12,3/4,1/12]]])) + self.assertLess(float((y*v).sum(-1).abs().max()),1e-6) + + def test_gumbel_velocity_keeps_the_realized_noise(self): + target=torch.tensor([[0,1,2,3]]) + def path(t): + torch.manual_seed(12) + return gumbel_path(target,torch.tensor([t])) + z,v=path(.3);h=1e-3 + numeric=(path(.3+h)[0]-path(.3-h)[0])/(2*h) + torch.testing.assert_close(v,numeric,atol=2e-4,rtol=2e-3) + self.assertLess(float(v.sum(-1).abs().max()),1e-6) + + def test_mh_detailed_balance_on_all_two_base_sequences(self): + states=[''.join(x) for x in itertools.product('ACGT',repeat=2)] + score=lambda s:(s.count('G')+.5*s.count('C'))/2 + weight=np.array([np.exp(3*score(s)) for s in states]);pi=weight/weight.sum() + matrix=np.zeros((16,16)) + for i,x in enumerate(states): + neighbors,q=proposal(x,score,3.) + for y,forward in zip(neighbors,q): + j=states.index(y);back,reverse=proposal(y,score,3.) + accept=min(1,pi[j]*reverse[back.index(x)]/(pi[i]*forward)) + matrix[i,j]=forward*accept + matrix[i,i]=1-matrix[i].sum() + np.testing.assert_allclose(pi[:,None]*matrix,pi[None,:]*matrix.T,atol=1e-14) + np.testing.assert_allclose(pi@matrix,pi,atol=1e-14) + + +if __name__ == '__main__': unittest.main() diff --git a/lecture_5/verified_examples/README.md b/lecture_5/verified_examples/README.md new file mode 100644 index 0000000000000000000000000000000000000000..62afe93b49c8dfd27bec16d3733cbc4e354b77b0 --- /dev/null +++ b/lecture_5/verified_examples/README.md @@ -0,0 +1,16 @@ +# Verified CPU examples + +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. + +| Method | Mean first 20 losses | Mean last 20 losses | Valid DNA | +| --- | ---: | ---: | --- | +| gat | 4.2671 | 2.6596 | True | +| dirichlet | 5.1298 | 1.7297 | True | +| fisher | 3.8173 | 2.2441 | True | +| gumbel | 0.8744 | 0.3556 | True | +| rectified | 0.8557 | 0.4507 | True | +| redi | 3.9590 | 2.1978 | True | +| mog-dfm | 4.2671 | 2.6596 | True | +| areuredi | 3.9590 | 2.1978 | True | + +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. diff --git a/lecture_5/verified_examples/areuredi/config.json b/lecture_5/verified_examples/areuredi/config.json new file mode 100644 index 0000000000000000000000000000000000000000..36c21e131495cbf76dad95d8eb217186b54a1b8f --- /dev/null +++ b/lecture_5/verified_examples/areuredi/config.json @@ -0,0 +1,22 @@ +{ + "method": "areuredi", + "mode": "train-sample", + "train_steps": 160, + "teacher_steps": 160, + "sample_steps": 40, + "refine_steps": 40, + "pairs": 128, + "batch_size": 16, + "samples": 8, + "length": 4, + "width": 32, + "guidance": 0.0, + "strength": 1.0, + "preference": [ + 0.7, + 0.3 + ], + "seed": 7, + "data": null, + "out": "outputs/areuredi" +} \ No newline at end of file diff --git a/lecture_5/verified_examples/areuredi/losses.json b/lecture_5/verified_examples/areuredi/losses.json new file mode 100644 index 0000000000000000000000000000000000000000..befbff9e0a2f8f8bf295ae6322b677ae798d7f65 --- /dev/null +++ b/lecture_5/verified_examples/areuredi/losses.json @@ -0,0 +1,162 @@ +[ + 5.578631401062012, + 5.215188026428223, + 4.979227542877197, + 4.527420997619629, + 4.682675838470459, + 4.541810989379883, + 4.201014041900635, + 4.809026718139648, + 3.5231478214263916, + 3.2159132957458496, + 3.903719902038574, + 3.392885684967041, + 3.4545810222625732, + 3.26749324798584, + 3.8055408000946045, + 3.170132637023926, + 3.293631076812744, + 2.7804863452911377, + 2.633892059326172, + 4.202793598175049, + 2.587411642074585, + 2.8439173698425293, + 2.620241165161133, + 2.5747973918914795, + 2.8458690643310547, + 2.6336047649383545, + 2.693552255630493, + 2.5792150497436523, + 3.147822618484497, + 2.2284865379333496, + 3.2001423835754395, + 2.6833009719848633, + 2.679579734802246, + 2.5627329349517822, + 2.251354932785034, + 3.9288926124572754, + 2.149650812149048, + 3.063689708709717, + 2.279712677001953, + 4.056086540222168, + 2.963582754135132, + 2.6804754734039307, + 1.967834234237671, + 2.4215853214263916, + 2.9922292232513428, + 2.785958766937256, + 3.5931544303894043, + 2.694758892059326, + 2.2316219806671143, + 2.0388429164886475, + 2.147338390350342, + 2.3730430603027344, + 1.6726641654968262, + 2.4648900032043457, + 1.438779592514038, + 2.2980430126190186, + 2.843613862991333, + 1.9195117950439453, + 2.389061450958252, + 3.2739145755767822, + 2.0545103549957275, + 1.9256659746170044, + 3.3023366928100586, + 2.801600694656372, + 2.824629545211792, + 2.492016315460205, + 2.8223798274993896, + 1.6702975034713745, + 2.814354419708252, + 2.0587551593780518, + 2.781656265258789, + 1.737116813659668, + 1.610327124595642, + 2.4434523582458496, + 2.191859006881714, + 3.4145593643188477, + 3.0424160957336426, + 2.804013252258301, + 2.3747494220733643, + 2.493913173675537, + 1.9928159713745117, + 2.0292046070098877, + 2.24965500831604, + 3.6124138832092285, + 2.4276280403137207, + 1.9320605993270874, + 2.759394645690918, + 1.8870341777801514, + 3.647275447845459, + 1.537461519241333, + 2.9912829399108887, + 1.6703108549118042, + 2.2091546058654785, + 1.9724771976470947, + 1.7340176105499268, + 1.7092127799987793, + 2.3674867153167725, + 2.40274715423584, + 2.2620091438293457, + 1.8866041898727417, + 3.7603864669799805, + 2.1125717163085938, + 2.674081325531006, + 1.9604092836380005, + 1.7161355018615723, + 2.0131330490112305, + 2.309553384780884, + 2.6360435485839844, + 2.5695250034332275, + 2.009622812271118, + 1.8004002571105957, + 1.3521816730499268, + 1.6246684789657593, + 1.3006975650787354, + 2.041378974914551, + 2.0425312519073486, + 2.045912027359009, + 3.1553938388824463, + 2.1102423667907715, + 1.0823614597320557, + 3.032719612121582, + 1.9022471904754639, + 2.118574857711792, + 2.1357498168945312, + 3.15908145904541, + 2.7499589920043945, + 3.0981345176696777, + 2.4977638721466064, + 2.40651273727417, + 2.120854616165161, + 1.4740090370178223, + 1.7393325567245483, + 2.018094539642334, + 1.570328712463379, + 2.010258436203003, + 1.5725524425506592, + 2.523509979248047, + 2.2055015563964844, + 1.5145868062973022, + 1.9369760751724243, + 2.6806631088256836, + 2.1069183349609375, + 1.9884653091430664, + 3.1166341304779053, + 2.874910831451416, + 1.8844469785690308, + 1.6240601539611816, + 2.5054893493652344, + 1.6134421825408936, + 2.26651930809021, + 2.565666913986206, + 2.329096555709839, + 1.6169782876968384, + 2.5725347995758057, + 2.094921588897705, + 1.9747111797332764, + 1.7718333005905151, + 2.446991443634033, + 1.5896382331848145, + 2.3323352336883545 +] \ No newline at end of file diff --git a/lecture_5/verified_examples/areuredi/report.json b/lecture_5/verified_examples/areuredi/report.json new file mode 100644 index 0000000000000000000000000000000000000000..87a0b34b38fa4dfe07134a90441b142042b7792d --- /dev/null +++ b/lecture_5/verified_examples/areuredi/report.json @@ -0,0 +1,19 @@ +{ + "metrics": { + "valid_dna": true, + "unique_fraction": 0.875, + "mean_gc": 0.4375, + "mean_atat_match": 0.21875 + }, + "method": "areuredi", + "data": "synthetic DNA; not biological validation", + "teacher_loss_last_20_mean": 2.72806202173233, + "paired_examples": 128, + "redi_scope": "one teacher-recoupling round; no guaranteed TC reduction", + "train_loss_first_20_mean": 3.9589606523513794, + "train_loss_last_20_mean": 2.1978128612041474, + "accepted_edits": 181, + "eta_final": 8.0, + "reference": "smoothed position-wise empirical DNA frequencies", + "scope": "finite annealing demonstration, not an equilibrium certificate" +} \ No newline at end of file diff --git a/lecture_5/verified_examples/areuredi/samples.txt b/lecture_5/verified_examples/areuredi/samples.txt new file mode 100644 index 0000000000000000000000000000000000000000..d782ccf9a67637fc6a7629165a060a0c88a11d7d --- /dev/null +++ b/lecture_5/verified_examples/areuredi/samples.txt @@ -0,0 +1,8 @@ +CTTT +GCGC +ACAA +ACGA +GAGC +TGTA +TCTT +TCTT diff --git a/lecture_5/verified_examples/dirichlet/config.json b/lecture_5/verified_examples/dirichlet/config.json new file mode 100644 index 0000000000000000000000000000000000000000..ef93e07fff8259eec1949c98c8f0a6915195fc6d --- /dev/null +++ b/lecture_5/verified_examples/dirichlet/config.json @@ -0,0 +1,22 @@ +{ + "method": "dirichlet", + "mode": "train-sample", + "train_steps": 160, + "teacher_steps": 160, + "sample_steps": 40, + "refine_steps": 40, + "pairs": 128, + "batch_size": 16, + "samples": 8, + "length": 4, + "width": 32, + "guidance": 0.0, + "strength": 1.0, + "preference": [ + 0.7, + 0.3 + ], + "seed": 7, + "data": null, + "out": "outputs/dirichlet" +} \ No newline at end of file diff --git a/lecture_5/verified_examples/dirichlet/losses.json b/lecture_5/verified_examples/dirichlet/losses.json new file mode 100644 index 0000000000000000000000000000000000000000..b56c7c45f0b584b6796e70a0412ce9d277ebb2e1 --- /dev/null +++ b/lecture_5/verified_examples/dirichlet/losses.json @@ -0,0 +1,162 @@ +[ + 5.987035274505615, + 5.7057271003723145, + 5.4776153564453125, + 5.435891628265381, + 5.582164287567139, + 5.3017377853393555, + 5.315533638000488, + 5.035096645355225, + 5.051508903503418, + 5.2037739753723145, + 5.109261989593506, + 5.108462810516357, + 4.778862953186035, + 4.777516841888428, + 4.974148750305176, + 4.8705244064331055, + 4.619065761566162, + 4.653285980224609, + 4.839997291564941, + 4.7690043449401855, + 4.481639862060547, + 4.536505699157715, + 4.780343055725098, + 4.4304728507995605, + 4.663294315338135, + 4.257389068603516, + 4.1362128257751465, + 4.409623622894287, + 4.034981727600098, + 4.11863899230957, + 4.146239280700684, + 4.428400993347168, + 3.9758212566375732, + 3.926872491836548, + 3.5168545246124268, + 3.9071707725524902, + 3.735076665878296, + 3.8633854389190674, + 3.630112886428833, + 3.444991111755371, + 3.599454879760742, + 3.0867366790771484, + 3.204212188720703, + 4.141026973724365, + 3.969395637512207, + 3.0596437454223633, + 3.5696041584014893, + 3.088737964630127, + 3.611302137374878, + 3.785858154296875, + 3.3500888347625732, + 2.989191770553589, + 3.21290922164917, + 2.6051502227783203, + 2.8940303325653076, + 2.719219446182251, + 2.355130910873413, + 2.667929172515869, + 2.6372263431549072, + 2.487605571746826, + 2.2540249824523926, + 3.2833821773529053, + 2.2793779373168945, + 2.491062879562378, + 1.6186134815216064, + 1.9055742025375366, + 2.983051300048828, + 2.754556894302368, + 2.1230244636535645, + 2.6531665325164795, + 2.091691732406616, + 2.476653575897217, + 2.002650260925293, + 1.9472030401229858, + 2.250502109527588, + 2.285085439682007, + 1.957679271697998, + 2.7570881843566895, + 1.6728074550628662, + 2.135037422180176, + 1.627407431602478, + 2.2475335597991943, + 3.223327875137329, + 2.079972982406616, + 2.17677903175354, + 1.4659333229064941, + 1.854305624961853, + 2.135460138320923, + 2.5259783267974854, + 2.1924796104431152, + 2.6706223487854004, + 1.7147475481033325, + 1.5733774900436401, + 1.5839917659759521, + 1.8101816177368164, + 1.6969491243362427, + 2.1939094066619873, + 2.2701213359832764, + 1.3590651750564575, + 1.8197208642959595, + 2.5202126502990723, + 1.112671136856079, + 1.6572060585021973, + 1.9186416864395142, + 2.151210308074951, + 2.294743061065674, + 1.6319775581359863, + 2.6904213428497314, + 1.4546334743499756, + 1.5501033067703247, + 1.0871210098266602, + 2.4776058197021484, + 1.789743185043335, + 1.3015778064727783, + 0.7730430364608765, + 1.2221710681915283, + 1.3628941774368286, + 1.5448253154754639, + 1.8459429740905762, + 1.5920908451080322, + 1.8741068840026855, + 1.606163740158081, + 1.0870177745819092, + 1.9514107704162598, + 1.4525673389434814, + 2.047329902648926, + 2.161097288131714, + 1.0058730840682983, + 1.4372437000274658, + 1.685389518737793, + 1.0957973003387451, + 1.4873944520950317, + 0.9507176280021667, + 0.9506893157958984, + 1.4974210262298584, + 1.191109538078308, + 2.463002920150757, + 2.1995906829833984, + 2.0290000438690186, + 2.622788906097412, + 2.003711700439453, + 2.577554941177368, + 2.561497211456299, + 1.0081901550292969, + 1.6362948417663574, + 1.4071522951126099, + 1.8568885326385498, + 1.614437460899353, + 2.286831855773926, + 1.1085795164108276, + 0.9553455114364624, + 1.2254221439361572, + 2.3875327110290527, + 2.3000073432922363, + 1.2774486541748047, + 1.0425993204116821, + 2.0665574073791504, + 1.6575270891189575, + 1.2717294692993164, + 2.348963975906372 +] \ No newline at end of file diff --git a/lecture_5/verified_examples/dirichlet/report.json b/lecture_5/verified_examples/dirichlet/report.json new file mode 100644 index 0000000000000000000000000000000000000000..a5f2cc44189a6d9c3d84f64e3b069c342bd68be6 --- /dev/null +++ b/lecture_5/verified_examples/dirichlet/report.json @@ -0,0 +1,17 @@ +{ + "metrics": { + "valid_dna": true, + "unique_fraction": 0.875, + "mean_gc": 0.3125, + "mean_atat_match": 0.125 + }, + "method": "dirichlet", + "data": "synthetic DNA; not biological validation", + "train_loss_first_20_mean": 5.129810786247253, + "train_loss_last_20_mean": 1.7297136068344117, + "validation_loss_one_mc_draw": 1.699185848236084, + "negative_coordinate_corrections": 0, + "max_probability_sum_error": 1.1920928955078125e-07, + "finite_endpoint": true, + "decode": "argmax" +} \ No newline at end of file diff --git a/lecture_5/verified_examples/dirichlet/samples.txt b/lecture_5/verified_examples/dirichlet/samples.txt new file mode 100644 index 0000000000000000000000000000000000000000..11e4b8b86f958ad2206d3cf11c47a43d2499a432 --- /dev/null +++ b/lecture_5/verified_examples/dirichlet/samples.txt @@ -0,0 +1,8 @@ +ACTT +TATA +ACTC +AGGA +CGTA +CATA +TATA +CCTA diff --git a/lecture_5/verified_examples/environment.json b/lecture_5/verified_examples/environment.json new file mode 100644 index 0000000000000000000000000000000000000000..8cffcc8fa542f4980ddd9dbc5fe3929cee0a59bd --- /dev/null +++ b/lecture_5/verified_examples/environment.json @@ -0,0 +1,7 @@ +{ + "python": "3.12.14", + "torch": "2.14.0+cpu", + "numpy": "2.3.5", + "scipy": "1.17.0", + "device": "cpu" +} \ No newline at end of file diff --git a/lecture_5/verified_examples/fisher/config.json b/lecture_5/verified_examples/fisher/config.json new file mode 100644 index 0000000000000000000000000000000000000000..c9817bfb1fdbb46be1cefbb248723c893188f90c --- /dev/null +++ b/lecture_5/verified_examples/fisher/config.json @@ -0,0 +1,22 @@ +{ + "method": "fisher", + "mode": "train-sample", + "train_steps": 160, + "teacher_steps": 160, + "sample_steps": 40, + "refine_steps": 40, + "pairs": 128, + "batch_size": 16, + "samples": 8, + "length": 4, + "width": 32, + "guidance": 0.0, + "strength": 1.0, + "preference": [ + 0.7, + 0.3 + ], + "seed": 7, + "data": null, + "out": "outputs/fisher" +} \ No newline at end of file diff --git a/lecture_5/verified_examples/fisher/losses.json b/lecture_5/verified_examples/fisher/losses.json new file mode 100644 index 0000000000000000000000000000000000000000..bc5f1e32ada3e4b818828aa8c4a41c6cea912c9f --- /dev/null +++ b/lecture_5/verified_examples/fisher/losses.json @@ -0,0 +1,162 @@ +[ + 8.571112632751465, + 7.303001880645752, + 5.1533427238464355, + 4.744162559509277, + 4.073383331298828, + 3.9271888732910156, + 4.31433629989624, + 3.0154306888580322, + 2.3777942657470703, + 3.170802116394043, + 2.6646761894226074, + 2.915005683898926, + 2.8402624130249023, + 3.812424421310425, + 2.3469979763031006, + 2.5703535079956055, + 3.515812635421753, + 2.982490062713623, + 3.0468573570251465, + 3.000063419342041, + 3.406219482421875, + 3.3296115398406982, + 3.011011838912964, + 2.480292797088623, + 2.9582724571228027, + 3.1977603435516357, + 3.128955602645874, + 2.2359578609466553, + 2.5828840732574463, + 3.4687106609344482, + 2.096220016479492, + 2.69463849067688, + 3.3743786811828613, + 2.262096881866455, + 2.928584098815918, + 2.6014180183410645, + 2.602999448776245, + 2.8589024543762207, + 3.174992084503174, + 2.7553186416625977, + 2.3450777530670166, + 2.7961537837982178, + 3.3228542804718018, + 2.9487850666046143, + 2.7978155612945557, + 3.455988883972168, + 2.2441558837890625, + 2.8949103355407715, + 2.880878448486328, + 2.186849594116211, + 2.235743522644043, + 1.8804774284362793, + 3.455813407897949, + 2.0829925537109375, + 2.082059144973755, + 2.530484914779663, + 2.7628636360168457, + 2.6144649982452393, + 2.476773262023926, + 2.4828391075134277, + 2.4270236492156982, + 2.7848620414733887, + 3.1417713165283203, + 2.309669256210327, + 2.755452871322632, + 3.49771785736084, + 2.934539318084717, + 2.7419092655181885, + 2.101997137069702, + 2.705397844314575, + 2.4323384761810303, + 2.167492389678955, + 2.7036519050598145, + 2.2807931900024414, + 2.9349067211151123, + 2.5758087635040283, + 2.4565770626068115, + 3.107597827911377, + 2.366063356399536, + 2.8201122283935547, + 2.8875114917755127, + 2.0656650066375732, + 3.2820560932159424, + 2.921267509460449, + 3.381413698196411, + 2.3997411727905273, + 2.3265788555145264, + 2.066448450088501, + 3.3638381958007812, + 1.8956438302993774, + 2.981762409210205, + 3.0908000469207764, + 2.0562901496887207, + 3.068679094314575, + 2.1612884998321533, + 2.5689961910247803, + 2.6655631065368652, + 2.81134295463562, + 3.134366512298584, + 2.2510604858398438, + 3.04536509513855, + 2.682065963745117, + 2.325636386871338, + 2.155775308609009, + 2.218646764755249, + 2.392900228500366, + 3.41904878616333, + 2.6010544300079346, + 2.3730106353759766, + 2.9705750942230225, + 2.698889970779419, + 2.8097193241119385, + 2.1913466453552246, + 2.81672739982605, + 2.9725401401519775, + 2.805398464202881, + 2.403073787689209, + 2.6580448150634766, + 2.9282894134521484, + 2.842071771621704, + 2.3050999641418457, + 2.810702323913574, + 2.1849422454833984, + 2.2562265396118164, + 3.07672119140625, + 2.0688652992248535, + 2.294011116027832, + 3.3593554496765137, + 2.9014484882354736, + 2.057595729827881, + 2.7754175662994385, + 3.2037341594696045, + 2.950990676879883, + 2.6158878803253174, + 2.0368189811706543, + 2.10182785987854, + 2.1909384727478027, + 2.7789814472198486, + 2.8794353008270264, + 2.876467227935791, + 1.9936639070510864, + 3.069582939147949, + 2.346599817276001, + 1.8762702941894531, + 1.5599921941757202, + 1.7584424018859863, + 2.0794010162353516, + 2.459878444671631, + 3.011709213256836, + 2.4862372875213623, + 2.082012176513672, + 2.2711989879608154, + 2.4660708904266357, + 2.0117135047912598, + 2.49086332321167, + 2.6844394207000732, + 1.8233208656311035, + 2.498573064804077, + 2.1096994876861572, + 1.8029230833053589 +] \ No newline at end of file diff --git a/lecture_5/verified_examples/fisher/report.json b/lecture_5/verified_examples/fisher/report.json new file mode 100644 index 0000000000000000000000000000000000000000..8fe6e4e0099e1bca8a787923fe19017cd14b332e --- /dev/null +++ b/lecture_5/verified_examples/fisher/report.json @@ -0,0 +1,17 @@ +{ + "metrics": { + "valid_dna": true, + "unique_fraction": 0.875, + "mean_gc": 0.78125, + "mean_atat_match": 0.15625 + }, + "method": "fisher", + "data": "synthetic DNA; not biological validation", + "train_loss_first_20_mean": 3.8172749519348144, + "train_loss_last_20_mean": 2.24412961602211, + "validation_loss_one_mc_draw": 2.5719263553619385, + "negative_coordinate_corrections": 283, + "max_probability_sum_error": 2.384185791015625e-07, + "finite_endpoint": true, + "decode": "argmax" +} \ No newline at end of file diff --git a/lecture_5/verified_examples/fisher/samples.txt b/lecture_5/verified_examples/fisher/samples.txt new file mode 100644 index 0000000000000000000000000000000000000000..097181dad8d6803f8fd11f6ddccf4e2b0c416aac --- /dev/null +++ b/lecture_5/verified_examples/fisher/samples.txt @@ -0,0 +1,8 @@ +ACGA +CCGC +CCGC +GCGA +ACGC +CCGT +GCGC +ACGT diff --git a/lecture_5/verified_examples/gat/config.json b/lecture_5/verified_examples/gat/config.json new file mode 100644 index 0000000000000000000000000000000000000000..3d7abfe1ea4c5028525b3a9e034c5db09f161746 --- /dev/null +++ b/lecture_5/verified_examples/gat/config.json @@ -0,0 +1,22 @@ +{ + "method": "gat", + "mode": "train-sample", + "train_steps": 160, + "teacher_steps": 160, + "sample_steps": 40, + "refine_steps": 40, + "pairs": 128, + "batch_size": 16, + "samples": 8, + "length": 4, + "width": 32, + "guidance": 0.0, + "strength": 1.0, + "preference": [ + 0.7, + 0.3 + ], + "seed": 7, + "data": null, + "out": "outputs/gat" +} \ No newline at end of file diff --git a/lecture_5/verified_examples/gat/losses.json b/lecture_5/verified_examples/gat/losses.json new file mode 100644 index 0000000000000000000000000000000000000000..eea81b1433d184a0005eb8e06c290dbfa4ae80ca --- /dev/null +++ b/lecture_5/verified_examples/gat/losses.json @@ -0,0 +1,162 @@ +[ + 5.555057525634766, + 5.890437602996826, + 4.9087815284729, + 4.1881232261657715, + 4.83707857131958, + 4.51339054107666, + 4.665221214294434, + 4.274393558502197, + 4.344203948974609, + 4.0454254150390625, + 4.4479660987854, + 4.175173759460449, + 3.908479690551758, + 4.424745559692383, + 3.710341215133667, + 3.0476884841918945, + 3.789707899093628, + 3.167095184326172, + 4.319497108459473, + 3.128598213195801, + 3.988326072692871, + 3.800382614135742, + 2.937274932861328, + 3.5654163360595703, + 3.1350417137145996, + 3.3353466987609863, + 3.536527156829834, + 3.8737754821777344, + 3.6239640712738037, + 3.4370455741882324, + 3.745927095413208, + 3.221271276473999, + 2.7261011600494385, + 3.674633026123047, + 3.3985674381256104, + 3.094015121459961, + 3.678131103515625, + 2.2317159175872803, + 2.205465793609619, + 2.597670078277588, + 2.588956832885742, + 2.7397661209106445, + 2.59556245803833, + 2.917344331741333, + 2.5570318698883057, + 2.386134147644043, + 3.1508893966674805, + 2.4510059356689453, + 4.685295104980469, + 2.9762532711029053, + 2.8647146224975586, + 1.7768700122833252, + 2.915055274963379, + 3.9934372901916504, + 2.89289927482605, + 3.744565010070801, + 2.373737335205078, + 3.3278493881225586, + 2.8167989253997803, + 2.7180516719818115, + 3.4837515354156494, + 3.799549102783203, + 3.4675912857055664, + 2.3154969215393066, + 4.505186080932617, + 1.955189824104309, + 2.515925407409668, + 2.2730512619018555, + 2.814613103866577, + 1.8026926517486572, + 3.208705186843872, + 3.4525582790374756, + 3.1925339698791504, + 2.355441093444824, + 4.377381324768066, + 2.403754711151123, + 3.8744754791259766, + 2.23490047454834, + 2.883716106414795, + 3.6779189109802246, + 2.032773494720459, + 2.265133857727051, + 3.280043601989746, + 2.532355308532715, + 2.523271083831787, + 2.462573766708374, + 3.8266568183898926, + 2.1587255001068115, + 2.038233757019043, + 3.216428279876709, + 3.1186740398406982, + 2.5141170024871826, + 3.333552837371826, + 3.366112232208252, + 3.6347179412841797, + 2.3267645835876465, + 4.024384498596191, + 3.129969358444214, + 3.784872055053711, + 2.7616593837738037, + 3.108980417251587, + 3.7265422344207764, + 2.9501824378967285, + 2.2087435722351074, + 4.133711338043213, + 3.426473617553711, + 3.052323579788208, + 3.112914562225342, + 3.7782933712005615, + 4.0156474113464355, + 2.726518392562866, + 3.0648770332336426, + 3.1597461700439453, + 2.0522677898406982, + 4.124749660491943, + 2.3789560794830322, + 2.617194890975952, + 3.1719017028808594, + 3.822108745574951, + 3.3801960945129395, + 3.1706669330596924, + 2.4255475997924805, + 2.65407395362854, + 3.5186941623687744, + 3.1759512424468994, + 3.3725812435150146, + 3.144639492034912, + 2.2421658039093018, + 4.015059947967529, + 2.498589277267456, + 2.3411831855773926, + 2.8789939880371094, + 2.6798653602600098, + 2.2761776447296143, + 2.7247633934020996, + 2.2538399696350098, + 2.367486000061035, + 3.8682806491851807, + 3.0465571880340576, + 2.8480899333953857, + 2.8016717433929443, + 2.4772744178771973, + 2.2610678672790527, + 3.1808559894561768, + 2.6285760402679443, + 2.3866126537323, + 3.704378604888916, + 1.7231093645095825, + 2.2585153579711914, + 2.0817971229553223, + 2.890470266342163, + 2.369269371032715, + 3.0134172439575195, + 2.607996702194214, + 3.0741591453552246, + 2.1917362213134766, + 3.8198983669281006, + 2.6407318115234375, + 2.5884804725646973, + 2.4910426139831543 +] \ No newline at end of file diff --git a/lecture_5/verified_examples/gat/report.json b/lecture_5/verified_examples/gat/report.json new file mode 100644 index 0000000000000000000000000000000000000000..825ce1b8ea33d23d4e01b4ed9f1afa125c767502 --- /dev/null +++ b/lecture_5/verified_examples/gat/report.json @@ -0,0 +1,13 @@ +{ + "metrics": { + "valid_dna": true, + "unique_fraction": 0.75, + "mean_gc": 0.375, + "mean_atat_match": 0.125 + }, + "method": "gat", + "data": "synthetic DNA; not biological validation", + "train_loss_first_20_mean": 4.267070317268372, + "train_loss_last_20_mean": 2.6595530688762663, + "validation_loss_one_mc_draw": 3.0547287464141846 +} \ No newline at end of file diff --git a/lecture_5/verified_examples/gat/samples.txt b/lecture_5/verified_examples/gat/samples.txt new file mode 100644 index 0000000000000000000000000000000000000000..76290f4ba84095a1f8bad2ec0e993899bcb2ba7b --- /dev/null +++ b/lecture_5/verified_examples/gat/samples.txt @@ -0,0 +1,8 @@ +ACGT +TCTA +CGTA +CGTA +CATA +TATA +ACGT +CCTA diff --git a/lecture_5/verified_examples/gumbel/config.json b/lecture_5/verified_examples/gumbel/config.json new file mode 100644 index 0000000000000000000000000000000000000000..542b9ec6a04b7900b9e7fafba09044af8783fd95 --- /dev/null +++ b/lecture_5/verified_examples/gumbel/config.json @@ -0,0 +1,22 @@ +{ + "method": "gumbel", + "mode": "train-sample", + "train_steps": 160, + "teacher_steps": 160, + "sample_steps": 40, + "refine_steps": 40, + "pairs": 128, + "batch_size": 16, + "samples": 8, + "length": 4, + "width": 32, + "guidance": 0.0, + "strength": 1.0, + "preference": [ + 0.7, + 0.3 + ], + "seed": 7, + "data": null, + "out": "outputs/gumbel" +} \ No newline at end of file diff --git a/lecture_5/verified_examples/gumbel/losses.json b/lecture_5/verified_examples/gumbel/losses.json new file mode 100644 index 0000000000000000000000000000000000000000..33958412705a5207271db2d0dc269d7409702832 --- /dev/null +++ b/lecture_5/verified_examples/gumbel/losses.json @@ -0,0 +1,162 @@ +[ + 1.2242249250411987, + 1.3620853424072266, + 1.0061757564544678, + 1.0545729398727417, + 0.9724909067153931, + 0.5293634533882141, + 0.6660601496696472, + 1.0677317380905151, + 0.878416895866394, + 0.9330573678016663, + 1.0132296085357666, + 0.8705434203147888, + 0.6601288318634033, + 0.6662362813949585, + 0.7980155944824219, + 0.8221703767776489, + 0.6677682995796204, + 0.5565563440322876, + 0.8239631652832031, + 0.915525496006012, + 0.8262391686439514, + 0.9338005185127258, + 0.7041582465171814, + 0.9623083472251892, + 0.7411048412322998, + 0.9122228622436523, + 0.6384052038192749, + 0.7176005244255066, + 0.6149242520332336, + 0.7930399775505066, + 0.9014930725097656, + 0.6472141742706299, + 0.5843918323516846, + 0.5959654450416565, + 0.8754173517227173, + 0.7718366980552673, + 0.5473393797874451, + 0.6955640912055969, + 0.754619836807251, + 0.5733585953712463, + 0.8471091389656067, + 0.7447681427001953, + 0.7545447945594788, + 0.6515666842460632, + 0.7182407379150391, + 0.5480296611785889, + 0.614464521408081, + 0.782268762588501, + 0.6048372387886047, + 0.6472717523574829, + 0.5385202765464783, + 0.6479076147079468, + 0.6813259124755859, + 0.6917023062705994, + 0.4879384934902191, + 0.5533413887023926, + 0.707999050617218, + 0.6003414392471313, + 0.7006964683532715, + 0.5734584331512451, + 0.7012950778007507, + 0.5581093430519104, + 0.4966667890548706, + 0.6264007687568665, + 0.5599461793899536, + 0.5360195636749268, + 0.665524959564209, + 0.6365910172462463, + 0.578106164932251, + 0.5579401254653931, + 0.5634976029396057, + 0.6342311501502991, + 0.5109550952911377, + 0.599521815776825, + 0.6203869581222534, + 0.611301839351654, + 0.5039989352226257, + 0.5459249019622803, + 0.5482620000839233, + 0.5299906730651855, + 0.5089408159255981, + 0.5460392236709595, + 0.41642990708351135, + 0.5682117938995361, + 0.46522098779678345, + 0.6535881757736206, + 0.5205237865447998, + 0.5230173468589783, + 0.5246025919914246, + 0.5621505975723267, + 0.44565266370773315, + 0.584469199180603, + 0.5001013278961182, + 0.553354024887085, + 0.5152371525764465, + 0.44525519013404846, + 0.5020580291748047, + 0.446787029504776, + 0.44008758664131165, + 0.386381596326828, + 0.5054974555969238, + 0.48175373673439026, + 0.3925288915634155, + 0.47026005387306213, + 0.40446269512176514, + 0.5167693495750427, + 0.393154501914978, + 0.4570925235748291, + 0.42173269391059875, + 0.35426291823387146, + 0.4976067841053009, + 0.4924037754535675, + 0.47807833552360535, + 0.4965772330760956, + 0.33257365226745605, + 0.29500123858451843, + 0.40001818537712097, + 0.36628803610801697, + 0.3501744568347931, + 0.24306155741214752, + 0.4280138909816742, + 0.33857619762420654, + 0.4133177101612091, + 0.31040486693382263, + 0.3238377273082733, + 0.37552645802497864, + 0.40614137053489685, + 0.3577098846435547, + 0.3461991846561432, + 0.3121543526649475, + 0.38914811611175537, + 0.42639264464378357, + 0.3369901478290558, + 0.4701955020427704, + 0.36834174394607544, + 0.2963407039642334, + 0.23096537590026855, + 0.3539124131202698, + 0.38701286911964417, + 0.4892623722553253, + 0.3251381516456604, + 0.35162293910980225, + 0.3250449001789093, + 0.34927764534950256, + 0.3196697235107422, + 0.28305667638778687, + 0.4241291880607605, + 0.41060760617256165, + 0.4208475947380066, + 0.4529426693916321, + 0.24722036719322205, + 0.43186840415000916, + 0.34614986181259155, + 0.35659974813461304, + 0.20478856563568115, + 0.3144701421260834, + 0.4556642472743988, + 0.38099414110183716, + 0.31331413984298706, + 0.3990325331687927 +] \ No newline at end of file diff --git a/lecture_5/verified_examples/gumbel/report.json b/lecture_5/verified_examples/gumbel/report.json new file mode 100644 index 0000000000000000000000000000000000000000..1d2625fd0833a71eb7483ae5c05ad37588d069f5 --- /dev/null +++ b/lecture_5/verified_examples/gumbel/report.json @@ -0,0 +1,18 @@ +{ + "metrics": { + "valid_dna": true, + "unique_fraction": 1.0, + "mean_gc": 0.4375, + "mean_atat_match": 0.3125 + }, + "method": "gumbel", + "data": "synthetic DNA; not biological validation", + "train_loss_first_20_mean": 0.8744158446788788, + "train_loss_last_20_mean": 0.355621962249279, + "validation_loss_one_mc_draw": 0.3075082004070282, + "negative_coordinate_corrections": 468, + "max_probability_sum_error": 1.1920928955078125e-07, + "finite_endpoint": true, + "decode": "argmax", + "gumbel_scope": "beta=1, tau_max=4, decay=4; finite noisy endpoint and approximate data-independent base; no exact data endpoint claim" +} \ No newline at end of file diff --git a/lecture_5/verified_examples/gumbel/samples.txt b/lecture_5/verified_examples/gumbel/samples.txt new file mode 100644 index 0000000000000000000000000000000000000000..26ba5d5f9c100c570db6cf3061f22007b0938b8a --- /dev/null +++ b/lecture_5/verified_examples/gumbel/samples.txt @@ -0,0 +1,8 @@ +TTAC +CCCC +TTGC +AGCG +CATT +TTAT +GTTG +AAGA diff --git a/lecture_5/verified_examples/mog-dfm/config.json b/lecture_5/verified_examples/mog-dfm/config.json new file mode 100644 index 0000000000000000000000000000000000000000..c71a1d043df6499fcac0f0029b4698f98940c390 --- /dev/null +++ b/lecture_5/verified_examples/mog-dfm/config.json @@ -0,0 +1,22 @@ +{ + "method": "mog-dfm", + "mode": "train-sample", + "train_steps": 160, + "teacher_steps": 160, + "sample_steps": 40, + "refine_steps": 40, + "pairs": 128, + "batch_size": 16, + "samples": 8, + "length": 4, + "width": 32, + "guidance": 0.0, + "strength": 1.0, + "preference": [ + 0.7, + 0.3 + ], + "seed": 7, + "data": null, + "out": "outputs/mog-dfm" +} \ No newline at end of file diff --git a/lecture_5/verified_examples/mog-dfm/losses.json b/lecture_5/verified_examples/mog-dfm/losses.json new file mode 100644 index 0000000000000000000000000000000000000000..eea81b1433d184a0005eb8e06c290dbfa4ae80ca --- /dev/null +++ b/lecture_5/verified_examples/mog-dfm/losses.json @@ -0,0 +1,162 @@ +[ + 5.555057525634766, + 5.890437602996826, + 4.9087815284729, + 4.1881232261657715, + 4.83707857131958, + 4.51339054107666, + 4.665221214294434, + 4.274393558502197, + 4.344203948974609, + 4.0454254150390625, + 4.4479660987854, + 4.175173759460449, + 3.908479690551758, + 4.424745559692383, + 3.710341215133667, + 3.0476884841918945, + 3.789707899093628, + 3.167095184326172, + 4.319497108459473, + 3.128598213195801, + 3.988326072692871, + 3.800382614135742, + 2.937274932861328, + 3.5654163360595703, + 3.1350417137145996, + 3.3353466987609863, + 3.536527156829834, + 3.8737754821777344, + 3.6239640712738037, + 3.4370455741882324, + 3.745927095413208, + 3.221271276473999, + 2.7261011600494385, + 3.674633026123047, + 3.3985674381256104, + 3.094015121459961, + 3.678131103515625, + 2.2317159175872803, + 2.205465793609619, + 2.597670078277588, + 2.588956832885742, + 2.7397661209106445, + 2.59556245803833, + 2.917344331741333, + 2.5570318698883057, + 2.386134147644043, + 3.1508893966674805, + 2.4510059356689453, + 4.685295104980469, + 2.9762532711029053, + 2.8647146224975586, + 1.7768700122833252, + 2.915055274963379, + 3.9934372901916504, + 2.89289927482605, + 3.744565010070801, + 2.373737335205078, + 3.3278493881225586, + 2.8167989253997803, + 2.7180516719818115, + 3.4837515354156494, + 3.799549102783203, + 3.4675912857055664, + 2.3154969215393066, + 4.505186080932617, + 1.955189824104309, + 2.515925407409668, + 2.2730512619018555, + 2.814613103866577, + 1.8026926517486572, + 3.208705186843872, + 3.4525582790374756, + 3.1925339698791504, + 2.355441093444824, + 4.377381324768066, + 2.403754711151123, + 3.8744754791259766, + 2.23490047454834, + 2.883716106414795, + 3.6779189109802246, + 2.032773494720459, + 2.265133857727051, + 3.280043601989746, + 2.532355308532715, + 2.523271083831787, + 2.462573766708374, + 3.8266568183898926, + 2.1587255001068115, + 2.038233757019043, + 3.216428279876709, + 3.1186740398406982, + 2.5141170024871826, + 3.333552837371826, + 3.366112232208252, + 3.6347179412841797, + 2.3267645835876465, + 4.024384498596191, + 3.129969358444214, + 3.784872055053711, + 2.7616593837738037, + 3.108980417251587, + 3.7265422344207764, + 2.9501824378967285, + 2.2087435722351074, + 4.133711338043213, + 3.426473617553711, + 3.052323579788208, + 3.112914562225342, + 3.7782933712005615, + 4.0156474113464355, + 2.726518392562866, + 3.0648770332336426, + 3.1597461700439453, + 2.0522677898406982, + 4.124749660491943, + 2.3789560794830322, + 2.617194890975952, + 3.1719017028808594, + 3.822108745574951, + 3.3801960945129395, + 3.1706669330596924, + 2.4255475997924805, + 2.65407395362854, + 3.5186941623687744, + 3.1759512424468994, + 3.3725812435150146, + 3.144639492034912, + 2.2421658039093018, + 4.015059947967529, + 2.498589277267456, + 2.3411831855773926, + 2.8789939880371094, + 2.6798653602600098, + 2.2761776447296143, + 2.7247633934020996, + 2.2538399696350098, + 2.367486000061035, + 3.8682806491851807, + 3.0465571880340576, + 2.8480899333953857, + 2.8016717433929443, + 2.4772744178771973, + 2.2610678672790527, + 3.1808559894561768, + 2.6285760402679443, + 2.3866126537323, + 3.704378604888916, + 1.7231093645095825, + 2.2585153579711914, + 2.0817971229553223, + 2.890470266342163, + 2.369269371032715, + 3.0134172439575195, + 2.607996702194214, + 3.0741591453552246, + 2.1917362213134766, + 3.8198983669281006, + 2.6407318115234375, + 2.5884804725646973, + 2.4910426139831543 +] \ No newline at end of file diff --git a/lecture_5/verified_examples/mog-dfm/report.json b/lecture_5/verified_examples/mog-dfm/report.json new file mode 100644 index 0000000000000000000000000000000000000000..adf18b9d9e7d93497effb0245966754d1f5f9096 --- /dev/null +++ b/lecture_5/verified_examples/mog-dfm/report.json @@ -0,0 +1,21 @@ +{ + "metrics": { + "valid_dna": true, + "unique_fraction": 0.875, + "mean_gc": 0.8125, + "mean_atat_match": 0.0625 + }, + "method": "mog-dfm", + "data": "synthetic DNA; not biological validation", + "train_loss_first_20_mean": 4.267070317268372, + "train_loss_last_20_mean": 2.6595530688762663, + "validation_loss_one_mc_draw": 3.0547287464141846, + "filtered_candidates": 3210, + "adaptive_steps": 42, + "preference": [ + 0.7, + 0.3 + ], + "final_cone_angle_degrees": 73.28139809162991, + "guarantee": "local weighted progress under the acute-cone assumptions" +} \ No newline at end of file diff --git a/lecture_5/verified_examples/mog-dfm/samples.txt b/lecture_5/verified_examples/mog-dfm/samples.txt new file mode 100644 index 0000000000000000000000000000000000000000..e8fe7eb74c7f5212fcb03a77cfdbc2528c1bd04d --- /dev/null +++ b/lecture_5/verified_examples/mog-dfm/samples.txt @@ -0,0 +1,8 @@ +ACGG +GCGC +CGCT +CCTC +GGGG +CATA +GCGC +CGGG diff --git a/lecture_5/verified_examples/rectified/config.json b/lecture_5/verified_examples/rectified/config.json new file mode 100644 index 0000000000000000000000000000000000000000..b682980bdea4da7f38285caa7374d7a37951c85b --- /dev/null +++ b/lecture_5/verified_examples/rectified/config.json @@ -0,0 +1,22 @@ +{ + "method": "rectified", + "mode": "train-sample", + "train_steps": 160, + "teacher_steps": 160, + "sample_steps": 40, + "refine_steps": 40, + "pairs": 128, + "batch_size": 16, + "samples": 8, + "length": 4, + "width": 32, + "guidance": 0.0, + "strength": 1.0, + "preference": [ + 0.7, + 0.3 + ], + "seed": 7, + "data": null, + "out": "outputs/rectified" +} \ No newline at end of file diff --git a/lecture_5/verified_examples/rectified/losses.json b/lecture_5/verified_examples/rectified/losses.json new file mode 100644 index 0000000000000000000000000000000000000000..65090334fcbed5b7ea0c280fbe8d15d0674a6ff5 --- /dev/null +++ b/lecture_5/verified_examples/rectified/losses.json @@ -0,0 +1,162 @@ +[ + 1.5004026889801025, + 1.3279709815979004, + 1.020445704460144, + 0.9433400630950928, + 0.8730323910713196, + 0.8370198011398315, + 0.828363299369812, + 0.7445297241210938, + 0.7042967677116394, + 0.7865955829620361, + 0.7114711403846741, + 0.7578641772270203, + 0.7066575884819031, + 0.8381931185722351, + 0.7460746169090271, + 0.7295529842376709, + 0.8562743663787842, + 0.7675602436065674, + 0.6833122968673706, + 0.7502613067626953, + 0.7138526439666748, + 0.7552366256713867, + 0.7716754674911499, + 0.701257586479187, + 0.8012855648994446, + 0.7691758275032043, + 0.7751699090003967, + 0.5887867212295532, + 0.6236982941627502, + 0.7256749868392944, + 0.5828669667243958, + 0.6786181926727295, + 0.7846888303756714, + 0.5736732482910156, + 0.6665273308753967, + 0.644498348236084, + 0.5881228446960449, + 0.6510074138641357, + 0.690257728099823, + 0.6035992503166199, + 0.5631512403488159, + 0.5767524838447571, + 0.6770502328872681, + 0.7123513221740723, + 0.631034791469574, + 0.6635925769805908, + 0.5014064311981201, + 0.6941777467727661, + 0.6142820119857788, + 0.48644065856933594, + 0.4690500497817993, + 0.46637943387031555, + 0.7518098950386047, + 0.49232983589172363, + 0.3956705331802368, + 0.5027762651443481, + 0.587986946105957, + 0.619992733001709, + 0.5430957078933716, + 0.5696991086006165, + 0.501618504524231, + 0.5913504958152771, + 0.5609534382820129, + 0.45804500579833984, + 0.5298258662223816, + 0.6507817506790161, + 0.574570894241333, + 0.5613725781440735, + 0.3827498257160187, + 0.6147360801696777, + 0.5094749927520752, + 0.4917214512825012, + 0.5207288265228271, + 0.4132007956504822, + 0.5481823682785034, + 0.4948666989803314, + 0.4673762321472168, + 0.6916881203651428, + 0.43736565113067627, + 0.4895283281803131, + 0.6169437170028687, + 0.41678938269615173, + 0.6408990621566772, + 0.6173636317253113, + 0.591915488243103, + 0.3778156638145447, + 0.4527539014816284, + 0.3762066066265106, + 0.6339781284332275, + 0.3565376102924347, + 0.5479102730751038, + 0.5159654021263123, + 0.4224727749824524, + 0.5880125164985657, + 0.3577631711959839, + 0.41638559103012085, + 0.5254880785942078, + 0.4846354126930237, + 0.6233035922050476, + 0.34998732805252075, + 0.6874790191650391, + 0.48790356516838074, + 0.4745773673057556, + 0.38737213611602783, + 0.43984824419021606, + 0.4269067347049713, + 0.5833062529563904, + 0.4814833998680115, + 0.42202267050743103, + 0.6122362613677979, + 0.5279317498207092, + 0.47669070959091187, + 0.4828556478023529, + 0.6017410159111023, + 0.5857915878295898, + 0.5456337332725525, + 0.45868223905563354, + 0.4599453806877136, + 0.574993371963501, + 0.4864501357078552, + 0.3898826241493225, + 0.6481937170028687, + 0.39759117364883423, + 0.4148366451263428, + 0.5928701758384705, + 0.37306398153305054, + 0.446571409702301, + 0.6392943263053894, + 0.562612771987915, + 0.4197731912136078, + 0.5282881259918213, + 0.6784601211547852, + 0.5052616000175476, + 0.49076318740844727, + 0.3642507791519165, + 0.42254331707954407, + 0.41095465421676636, + 0.5811480283737183, + 0.5557343363761902, + 0.5879122614860535, + 0.40378907322883606, + 0.649993360042572, + 0.5146540999412537, + 0.39203402400016785, + 0.2682815194129944, + 0.3100827634334564, + 0.3964778184890747, + 0.4988245964050293, + 0.6458528637886047, + 0.46725594997406006, + 0.3653281033039093, + 0.4387623965740204, + 0.44214093685150146, + 0.45757776498794556, + 0.48578232526779175, + 0.5535528659820557, + 0.4033154547214508, + 0.528259813785553, + 0.4303647577762604, + 0.36250683665275574 +] \ No newline at end of file diff --git a/lecture_5/verified_examples/rectified/report.json b/lecture_5/verified_examples/rectified/report.json new file mode 100644 index 0000000000000000000000000000000000000000..2104c568502c03acdf158331a7824d85eeced589 --- /dev/null +++ b/lecture_5/verified_examples/rectified/report.json @@ -0,0 +1,17 @@ +{ + "metrics": { + "valid_dna": true, + "unique_fraction": 0.875, + "mean_gc": 0.75, + "mean_atat_match": 0.125 + }, + "method": "rectified", + "data": "synthetic DNA; not biological validation", + "train_loss_first_20_mean": 0.855660942196846, + "train_loss_last_20_mean": 0.4507418662309647, + "validation_loss_one_mc_draw": 0.4995970129966736, + "negative_coordinate_corrections": 733, + "max_probability_sum_error": 1.1920928955078125e-07, + "finite_endpoint": true, + "decode": "argmax" +} \ No newline at end of file diff --git a/lecture_5/verified_examples/rectified/samples.txt b/lecture_5/verified_examples/rectified/samples.txt new file mode 100644 index 0000000000000000000000000000000000000000..c12968cd694ecdb0fd784ad934e015900c2e83dc --- /dev/null +++ b/lecture_5/verified_examples/rectified/samples.txt @@ -0,0 +1,8 @@ +ATGT +CCGC +CGGC +GCGA +GCGC +CATA +GCGC +ACGC diff --git a/lecture_5/verified_examples/redi/config.json b/lecture_5/verified_examples/redi/config.json new file mode 100644 index 0000000000000000000000000000000000000000..5f0ece6929b7fff2a41616628ebc5d3f93845a42 --- /dev/null +++ b/lecture_5/verified_examples/redi/config.json @@ -0,0 +1,22 @@ +{ + "method": "redi", + "mode": "train-sample", + "train_steps": 160, + "teacher_steps": 160, + "sample_steps": 40, + "refine_steps": 40, + "pairs": 128, + "batch_size": 16, + "samples": 8, + "length": 4, + "width": 32, + "guidance": 0.0, + "strength": 1.0, + "preference": [ + 0.7, + 0.3 + ], + "seed": 7, + "data": null, + "out": "outputs/redi" +} \ No newline at end of file diff --git a/lecture_5/verified_examples/redi/losses.json b/lecture_5/verified_examples/redi/losses.json new file mode 100644 index 0000000000000000000000000000000000000000..befbff9e0a2f8f8bf295ae6322b677ae798d7f65 --- /dev/null +++ b/lecture_5/verified_examples/redi/losses.json @@ -0,0 +1,162 @@ +[ + 5.578631401062012, + 5.215188026428223, + 4.979227542877197, + 4.527420997619629, + 4.682675838470459, + 4.541810989379883, + 4.201014041900635, + 4.809026718139648, + 3.5231478214263916, + 3.2159132957458496, + 3.903719902038574, + 3.392885684967041, + 3.4545810222625732, + 3.26749324798584, + 3.8055408000946045, + 3.170132637023926, + 3.293631076812744, + 2.7804863452911377, + 2.633892059326172, + 4.202793598175049, + 2.587411642074585, + 2.8439173698425293, + 2.620241165161133, + 2.5747973918914795, + 2.8458690643310547, + 2.6336047649383545, + 2.693552255630493, + 2.5792150497436523, + 3.147822618484497, + 2.2284865379333496, + 3.2001423835754395, + 2.6833009719848633, + 2.679579734802246, + 2.5627329349517822, + 2.251354932785034, + 3.9288926124572754, + 2.149650812149048, + 3.063689708709717, + 2.279712677001953, + 4.056086540222168, + 2.963582754135132, + 2.6804754734039307, + 1.967834234237671, + 2.4215853214263916, + 2.9922292232513428, + 2.785958766937256, + 3.5931544303894043, + 2.694758892059326, + 2.2316219806671143, + 2.0388429164886475, + 2.147338390350342, + 2.3730430603027344, + 1.6726641654968262, + 2.4648900032043457, + 1.438779592514038, + 2.2980430126190186, + 2.843613862991333, + 1.9195117950439453, + 2.389061450958252, + 3.2739145755767822, + 2.0545103549957275, + 1.9256659746170044, + 3.3023366928100586, + 2.801600694656372, + 2.824629545211792, + 2.492016315460205, + 2.8223798274993896, + 1.6702975034713745, + 2.814354419708252, + 2.0587551593780518, + 2.781656265258789, + 1.737116813659668, + 1.610327124595642, + 2.4434523582458496, + 2.191859006881714, + 3.4145593643188477, + 3.0424160957336426, + 2.804013252258301, + 2.3747494220733643, + 2.493913173675537, + 1.9928159713745117, + 2.0292046070098877, + 2.24965500831604, + 3.6124138832092285, + 2.4276280403137207, + 1.9320605993270874, + 2.759394645690918, + 1.8870341777801514, + 3.647275447845459, + 1.537461519241333, + 2.9912829399108887, + 1.6703108549118042, + 2.2091546058654785, + 1.9724771976470947, + 1.7340176105499268, + 1.7092127799987793, + 2.3674867153167725, + 2.40274715423584, + 2.2620091438293457, + 1.8866041898727417, + 3.7603864669799805, + 2.1125717163085938, + 2.674081325531006, + 1.9604092836380005, + 1.7161355018615723, + 2.0131330490112305, + 2.309553384780884, + 2.6360435485839844, + 2.5695250034332275, + 2.009622812271118, + 1.8004002571105957, + 1.3521816730499268, + 1.6246684789657593, + 1.3006975650787354, + 2.041378974914551, + 2.0425312519073486, + 2.045912027359009, + 3.1553938388824463, + 2.1102423667907715, + 1.0823614597320557, + 3.032719612121582, + 1.9022471904754639, + 2.118574857711792, + 2.1357498168945312, + 3.15908145904541, + 2.7499589920043945, + 3.0981345176696777, + 2.4977638721466064, + 2.40651273727417, + 2.120854616165161, + 1.4740090370178223, + 1.7393325567245483, + 2.018094539642334, + 1.570328712463379, + 2.010258436203003, + 1.5725524425506592, + 2.523509979248047, + 2.2055015563964844, + 1.5145868062973022, + 1.9369760751724243, + 2.6806631088256836, + 2.1069183349609375, + 1.9884653091430664, + 3.1166341304779053, + 2.874910831451416, + 1.8844469785690308, + 1.6240601539611816, + 2.5054893493652344, + 1.6134421825408936, + 2.26651930809021, + 2.565666913986206, + 2.329096555709839, + 1.6169782876968384, + 2.5725347995758057, + 2.094921588897705, + 1.9747111797332764, + 1.7718333005905151, + 2.446991443634033, + 1.5896382331848145, + 2.3323352336883545 +] \ No newline at end of file diff --git a/lecture_5/verified_examples/redi/report.json b/lecture_5/verified_examples/redi/report.json new file mode 100644 index 0000000000000000000000000000000000000000..9b286ea1847e93d704113b5420e509b8fa262ffa --- /dev/null +++ b/lecture_5/verified_examples/redi/report.json @@ -0,0 +1,15 @@ +{ + "metrics": { + "valid_dna": true, + "unique_fraction": 0.875, + "mean_gc": 0.40625, + "mean_atat_match": 0.03125 + }, + "method": "redi", + "data": "synthetic DNA; not biological validation", + "teacher_loss_last_20_mean": 2.72806202173233, + "paired_examples": 128, + "redi_scope": "one teacher-recoupling round; no guaranteed TC reduction", + "train_loss_first_20_mean": 3.9589606523513794, + "train_loss_last_20_mean": 2.1978128612041474 +} \ No newline at end of file diff --git a/lecture_5/verified_examples/redi/samples.txt b/lecture_5/verified_examples/redi/samples.txt new file mode 100644 index 0000000000000000000000000000000000000000..b64bc0c42f06a70bb0a5e5f601b3e1ba58d84e21 --- /dev/null +++ b/lecture_5/verified_examples/redi/samples.txt @@ -0,0 +1,8 @@ +GCGA +TCCA +TATT +CGTA +GCTA +GCGC +TATA +TATA diff --git a/requirements.txt b/requirements.txt index dfe1a53848b6b012226b48d498b4f7961186d8c9..527a04073982472c0dd6d88fc7483c86e0ac9d63 100644 --- a/requirements.txt +++ b/requirements.txt @@ -2,3 +2,5 @@ torch==2.9.1 torchvision==0.24.1 transformers==4.57.6 +numpy>=1.24 +scipy>=1.10