Download lecture_4/lecture_core.py from ChatterjeeLab/CIS6270: direct link, hf CLI and curl.
- Browser
- Download file 4.54 kB
-
https://huggingface.co/ChatterjeeLab/CIS6270/resolve/main/lecture_4/lecture_core.py
- Command line
-
hf download hf://ChatterjeeLab/CIS6270/lecture_4/lecture_core.py
-
curl -L -o lecture_core.py https://huggingface.co/ChatterjeeLab/CIS6270/resolve/main/lecture_4/lecture_core.py
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() | |
| 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] | |