CIS6270 / lecture_4 /lecture_core.py
pranamanam's picture
Add Lecture 4 discrete diffusion and Lecture 5 discrete flow matching
c8293a4 verified
Raw History Blame Contribute Delete
4.54 kB
"""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]