CIS6270 / lecture_4 /README.md
pranamanam's picture
Upload 87 files
0ba9d09 verified
|
Raw History Blame Contribute Delete
10.5 kB

Lecture 4: Discrete Diffusion

Chapter 4 keeps every state inside a finite vocabulary and replaces the continuous displacement of Chapter 3 by a decision about which tokens to erase, reveal or replace. The chapter builds masked diffusion by hand from a survival schedule, then rebuilds it as one member of a family: a rate matrix fixes what a token can do in the next instant, the master equation follows the marginals, the matrix exponential gives the corruption kernel at any time, and time reversal turns the forward rates into the rates that generate. Seven published models come out of that family by changing the rate matrix, the time variable, what the network predicts, or the state space itself, and a further group keeps the rates and changes the order in which positions are filled, the weighting of the loss, or the distribution to sample from. Every file here runs on the shared synthetic DNA of cis6270.data, with t = 0 clean and t = 1 all masks, survival probability alpha_t, and rates R_t(x, y) in the row convention p_dot = p R.

The files

File Command What it trains and generates
ctmc.py python lecture_4/ctmc.py Trains no network, and works every number of the CTMC sections: the two rules, one row for DNA, waiting times, the master equation, the matrix exponential, the R_pi family, eigenvalues, detailed balance and time reversal.
mdlm.py python lecture_4/mdlm.py Trains the denoiser on the weighted cross-entropy at masked positions, generates DNA from all masks. The chapter's thirty-line model.
uniform.py python lecture_4/uniform.py Trains against the rate loss under uniform replacement, generates by tau-leaping from uniformly random tokens.
d3pm.py python lecture_4/d3pm.py Trains against the discrete-time variational bound plus the auxiliary cross-entropy, generates by T ancestral steps.
sedd.py python lecture_4/sedd.py Trains a network to output the concrete score through denoising score entropy, generates from the predicted ratios.
block.py python lecture_4/block.py Trains masked diffusion inside blocks with a clean prefix, generates block by block. --block moves the dial between MDLM and autoregression.
simplex.py python lecture_4/simplex.py Trains a denoiser that reads a Dirichlet belief at every position, generates by thinning the belief, drawing an innovation and mixing the two.
guidance.py python lecture_4/guidance.py Trains a label-dropout denoiser and a classifier on noisy sequences, then samples unguided, with exact classifier guidance, with its gradient approximation, and classifier free.
planner.py python lecture_4/planner.py Trains with the planner weights inside the bound, samples the same weights in a uniform order and in the planner's order with a remasking corrector.
search.py python lecture_4/search.py Trains masked diffusion, then searches the tree of partial sequences under the h-transform and keeps a Pareto archive over two objectives.

Two constructions of the chapter have no file here, and the chapter develops both in full. Per-token schedules, MD4 and GenMD4, give every clean token its own rate beta_a(t) and change nothing but the weight, which becomes w_a(t) = -alpha_dot_t(a)/(1 - alpha_t(a)) inside the sum over tokens. A2D2 lets the length change by adding an insertion jump beside the unmasking jump, which enlarges the state space and keeps the master equation, the reversal and the rate loss exactly as they are.

Every file takes the shared flags of cis6270.runner.common_parser, so --steps, --samples, --sample-steps, --seed and --quiet mean the same thing everywhere. ctmc.py works through rates the chapter writes down, so its --steps is the number of Euler steps in the matrix-exponential comparison and its --samples the number of simulated paths. d3pm.py takes T from --timesteps, since the same number fixes the corruption and the ancestral sampler. search.py reads --sample-steps as the length of one rollout, which is short because we score every leaf of the tree by its own rollouts.

What a default run produces

We took every number below from --seed 0 on a CPU. The worked numbers are exact and repeat on any machine; the trained numbers move a little with the hardware's floating point, and the sample statistics move with the seed. The training data have GC fraction 0.577 and motif fraction 0.777, and a sampler drawing its bases uniformly would land near 0.05 on the motif, since ACGT fits in a length-16 sequence at thirteen places and each one has probability 1/256.

File Worked numbers it prints Trained numbers
ctmc.py exit rate 4 out of A, P(A, A) = 0.6 at h = 0.1, Euler step (0.16, 0.22, 0.28, 0.34), exp(tau R_pi) at 0.625 and 0.125, Euler products 0.480, 0.570, 0.616, reverse rates 2 and 0.5 out of C none
mdlm.py r = 0.5, reverse law (0.3, 0.2, 0, 0, 0.5), w(0.5) = 2, loss term 2.0996, rate error 1.0217 bound 1.25 nats per token, GC 0.56, motif 0.69
uniform.py kernel 0.625 and 0.125, clean factor 5, posterior (0.375, 0.625), conditional rates 5 and 0.2, posterior average 2, model rates 2.76 and 1.16, rate error 0.754 rate loss 19.1 nats per sequence, GC 0.56, motif 0.50
d3pm.py alpha_bar_2 = 0.56, Q_bar_2 at 0.67 and 0.11, the step-1 posterior (0.5795, 0.3523, 0.0341, 0.0341), Gaussian row 0.4088 and 0.0912, exp(0.3 R) first row (0.6747, 0.1964, 0.1122, 0.0167) bound 18.2 nats per sequence, GC 0.57, motif 0.72
sedd.py target ratio 2, score entropy 0.0107, 0.3863, 0.1891 score entropy 18.8 nats per sequence, mean predicted ratio 1.26 toward a conditional target of 5, GC 0.55, motif 0.33
block.py worked bound 2.7567 nats on ACGTAC, integral of w(t)(1 - alpha_t) = 1 bound 1.21 nats per token at L' = 4 and 1.29 at L' = 1, GC 0.57, motif 0.67
simplex.py c_t = 8 and c_s = 20, parameters (5, 1, 1, 1) and (17, 1, 1, 1), standard deviation 0.161 on A, increment (12, 0, 0, 0), W ~ Beta(8, 12) with mean 0.4, simulated standard deviation 0.0779 against the exact 0.0779 cross-entropy 0.42 nats per token, GC 0.57, motif 0.34
guidance.py h = (2, 1, 0.5, 0.5), guided rates (1, 1.5, 0.5), at gamma = 2 rates (0.5, 4.5, 0.5) with exit rate 5.5 and share 0.8182 on G GC 0.58 unguided, 0.68 with the exact classifier, 0.63 with the gradient approximation, 0.74 classifier free
planner.py reveals positions 3 and 6, remasks position 5, both corrector fluxes 1, PAPL weights (0.4, 2.0, 0.6, 1.0), inner sums 2.9 and 2.12 motif 0.61 in a uniform order against 0.86 in the planner's order, GC 0.58 and 0.51
search.py h(z) = 0.35, h(z') = 0.75, factor 2.143, a rate of 1.2 becomes 2.571, the archive keeps P1 and P2 144 rollouts scored, an archive of 3 sequences, best GC 0.88 and best motif match 1.0

On one core, we get about two seconds from a default run of ctmc.py, and between thirty and forty seconds of processor time from every other file. With --steps 20, each file runs in a few seconds.

mdlm.py --steps 3000 lowers the bound from 1.25 to 1.22 nats per token in about three times the time, and leaves the motif fraction inside the spread it already has across seeds, which at seeds 0, 1 and 2 is 0.69, 0.45 and 0.72. The one-position bound of block.py --block 1 is the exact autoregressive log-likelihood with no slack, and at this budget it still measures above the L' = 4 bound, because the estimator draws one block of sixteen and multiplies by sixteen, which is sixteen times noisier at the same number of draws, and because a denoiser restricted to a left prefix has a harder prediction problem than one that reads both sides. sedd.py sits below the other samplers on the motif because its sampler stops at the time floor, as the list below says. simplex.py --sample-steps 256 raises its motif fraction from 0.34 to 0.38, so the reverse grid accounts for only a small part of what holds that number down.

What is simplified

  • The time weight w(t) = 1/t is unbounded at t = 0, so every file draws its times above a small floor. The chapter's integrand is correct as written; the floor keeps one draw from dominating a batch.
  • The samplers that use tau-leaping rescale a step whose total jump probability would exceed one. uniform.py then finishes with the exact posterior step down to t = 0 that the chapter derives for the pi family. sedd.py predicts ratios, with no posterior over clean tokens to step through, so it stops at the time floor, where alpha_t = 0.92 leaves about eight percent of the positions redrawn from pi.
  • uniform.py and sedd.py use the constant schedule beta(t) = 4, which is the schedule of the chapter's worked example and reaches alpha_t = 1/2 at t = (log 2)/4. Its survival probability at t = 1 is 0.018, and that residue is the KL(q_1 || p_1) term of the bound.
  • d3pm.py samples one step index per sequence and multiplies by T, which estimates the sum over steps without evaluating all of it.
  • block.py samples one block per sequence and multiplies by the number of blocks, for the same reason. To train a block of length L' we hide the suffix behind the mask in place of a block-causal attention mask.
  • simplex.py holds the concentration constant at gamma = 1, realizes the sum over clean tokens by one categorical draw as the chapter does, and reads the final token as the corner the belief has reached. Its loss is the chapter's cross-entropy, so its nats per token are the prediction error of the denoiser at the sampled times, and the bounds above are on a likelihood.
  • guidance.py uses a mean-pooling classifier on one-hot inputs, because we need a derivative with respect to the sequence to approximate the gradient, and the label of this data is a thresholded mean over positions.
  • planner.py reveals the number of positions the schedule sets at each step, which is the deterministic version of the chapter's independent reveal draws, and remasks a fixed count.
  • search.py scores --sample-steps rollout steps per completion and expands a fixed number of children, so the tree stays small enough to build in a few seconds.