Download lecture_6/expanding.py from ChatterjeeLab/CIS6270: direct link, hf CLI and curl.
- Browser
- Download file 11.4 kB
-
https://huggingface.co/ChatterjeeLab/CIS6270/resolve/main/lecture_6/expanding.py
- Command line
-
hf download hf://ChatterjeeLab/CIS6270/lecture_6/expanding.py
-
curl -L -o expanding.py https://huggingface.co/ChatterjeeLab/CIS6270/resolve/main/lecture_6/expanding.py
11.4 kB
| """Expanding flow maps, where the state gains coordinates as time runs. | |
| The Jacobian of a map between states of different length is rectangular, so it has no determinant | |
| and no inverse, and Liouville's formula, the exact likelihood and the inverse map all stop | |
| applying. If we split the move into an expansion that fills new coordinates with source noise and a | |
| transport that acts at the longer length, the Jacobian is square again. Each coordinate then carries | |
| its own clock, rescaled to the interval it occupies, and the composition identity is still | |
| available, because a finite difference between two maps needs no derivative in time. | |
| Chapter 6, Section 6.32 of the notes. Trained on the synthetic DNA of | |
| ``cis6270.data.synthetic_dna``, grown from half its length to all of it. Prints the local clocks | |
| of the chapter's figure, 0.800, 0.692 and 0.333 at a global time of 0.8, the composition weight | |
| 0.125 at (0.3, 0.6, 0.9), and samples from the trained map. | |
| python lecture_6/expanding.py # 1500 steps, about half a minute | |
| python lecture_6/expanding.py --steps 20 # a few seconds, for a first look | |
| """ | |
| # %% 01. Imports and configuration | |
| import math | |
| import torch | |
| from cis6270.data import DNA_SIZE, decode_dna, synthetic_dna | |
| from cis6270.metrics import gc_fraction, motif_fraction, token_entropy, unique_fraction | |
| from cis6270.nets import TwoTimeMLP, count_parameters | |
| from cis6270.runner import common_parser, pick_device, print_report, save_run, set_seed, train | |
| torch.set_num_threads(1) # These networks are small, and splitting one matrix product costs more | |
| # than it saves. ``tests/conftest.py`` sets the same thing for the test session, and this line | |
| # acts when a reader runs the script directly. | |
| BOOK_BIRTHS = (0.0, 0.0, 0.35, 0.7) # The insertion times of the chapter's four-tile figure. | |
| # %% 02. Local clocks | |
| # | |
| # A coordinate inserted at t_ins has 1 - t_ins of global time left, so its own clock runs at | |
| # 1/(1 - t_ins) and reaches one exactly at the data end. Before its insertion the clock reads | |
| # zero and the coordinate holds its source value. | |
| def local_time(t, birth): | |
| """t_i = (t - t_ins)/(1 - t_ins), clamped to [0, 1] so an unborn coordinate reads zero.""" | |
| return ((t - birth) / (1.0 - birth).clamp(min=1e-6)).clamp(0.0, 1.0) | |
| def worked_clocks(t: float = 0.8): | |
| """The clocks, the speeds and the interpolant shares of the chapter's figure at one time.""" | |
| births = torch.tensor(BOOK_BIRTHS) | |
| clocks = local_time(torch.tensor(t), births) # [4]. | |
| speeds = 1.0 / (1.0 - births) # [4] dt_i/dt. | |
| return { | |
| "insertion times": [round(v, 2) for v in births.tolist()], | |
| "local clocks at t=0.8": [round(v, 3) for v in clocks.tolist()], | |
| "clock speeds": [round(v, 3) for v in speeds.tolist()], | |
| "source share": [round(v, 3) for v in (1.0 - clocks).tolist()], | |
| "last clock from its speed": float(speeds[-1] * (1.0 - t)), | |
| "live coordinates at 0.8": int((births <= t).sum()), | |
| "live coordinates at 0.6": int((births <= 0.6).sum()), | |
| } | |
| # %% 03. The composition weight under growth | |
| # | |
| # The convexity weight is a ratio of remaining times and makes no reference to the dimension, so | |
| # the derivation at fixed length goes through once every leg has been lifted to the final length. | |
| # A coordinate born at zero carries the chapter's omega; a later birth carries its own. | |
| def omega(s, u, t): | |
| """omega_{s,u,t} = (u-s)(1-t)/((t-s)(1-u)), the share of the earlier leg.""" | |
| return (u - s) * (1.0 - t) / (((t - s) * (1.0 - u)).clamp(min=1e-9)) | |
| def worked_omega(s: float = 0.3, u: float = 0.6, t: float = 0.9): | |
| """The weight at the chapter's triple, for coordinates born at three different times.""" | |
| out = {"omega at (0.3, 0.6, 0.9)": float(omega(*(torch.tensor([v]) for v in (s, u, t))))} | |
| for birth in (0.0, 0.35, 0.7): | |
| b = torch.tensor([birth]) | |
| local = [local_time(torch.tensor([v]), b) for v in (s, u, t)] | |
| out[f"omega for a birth at {birth}"] = float(omega(*local)) | |
| return out | |
| # %% 04. A state whose length grows, and where the fresh noise enters | |
| # | |
| # Every coordinate that will ever appear gets its source value from one draw at the final length, | |
| # and a leg that stops before an insertion leaves that coordinate at its source value. One draw | |
| # then serves all three maps of the composition identity. | |
| def lifted_state(source, data_onehot, t, births): | |
| """The state at global time t, each coordinate interpolated on its own clock. [B, L, K].""" | |
| clock = local_time(t[:, None], births[None, :])[:, :, None] # [B, L, 1]. | |
| return (1.0 - clock) * source + clock * data_onehot # [B, L, K]. | |
| def expansion_log_density(source, births, s: float, t: float): | |
| """The source log density of the coordinates inserted inside (s, t], the expansion term.""" | |
| inserted = (births > s) & (births <= t) # [L] which coordinates appear on this interval. | |
| per_coordinate = (-0.5 * source ** 2 - 0.5 * math.log(2.0 * math.pi)).sum(-1) # [B, L]. | |
| return float(per_coordinate[:, inserted].sum(-1).mean()), int(inserted.sum()) | |
| # %% 05. The model, a two-time denoiser that also reads each coordinate's clock | |
| def predict(model, x, s, t, births, length, vocab): | |
| """psi_{s,t}(x) at every position, with each coordinate's own clock as an extra channel. | |
| The network has the same width in and out, so every position carries K + 1 channels, the K | |
| coordinates of the state and the local clock. We drop the clock channel of the output. | |
| """ | |
| clock = local_time(s[:, None], births[None, :])[:, :, None] # [B, L, 1] where each stands. | |
| features = torch.cat([x, clock], dim=-1).flatten(1) # [B, L*(K+1)]. | |
| logits = model(features, s, t).view(-1, length, vocab + 1)[..., :vocab] # [B, L, K]. | |
| return torch.softmax(logits, dim=-1) # [B, L, K] on the simplex at every position. | |
| def step_map(x, prediction, s, t, births): | |
| """The convex form, with the weights read on each coordinate's own clock.""" | |
| s_local = local_time(s[:, None], births[None, :])[:, :, None] # [B, L, 1]. | |
| t_local = local_time(t[:, None], births[None, :])[:, :, None] # [B, L, 1]. | |
| keep = (1.0 - t_local) / (1.0 - s_local).clamp(min=1e-6) | |
| return keep * x + (1.0 - keep) * prediction # [B, L, K]. | |
| def make_loss(model, data, births, length, vocab, batch_size, device): | |
| eye = torch.eye(vocab, device=device) | |
| def loss_fn(step: int) -> torch.Tensor: | |
| index = torch.randint(0, len(data), (batch_size,), device=device) | |
| ids = data[index] # [B, L]. | |
| x1 = eye[ids] # [B, L, K] the data, lifted. | |
| source = torch.randn(batch_size, length, vocab, device=device) # Drawn once, at full length. | |
| t_diag = torch.rand(batch_size, device=device) | |
| xt = lifted_state(source, x1, t_diag, births) | |
| diagonal = -(predict(model, xt, t_diag, t_diag, births, length, vocab).clamp_min(1e-9) | |
| .log().gather(-1, ids[..., None])).mean() # Cross-entropy on the diagonal. | |
| s = torch.rand(batch_size, device=device) | |
| t = s + (1.0 - s) * torch.rand(batch_size, device=device) | |
| u = s + (t - s) * torch.rand(batch_size, device=device) | |
| xs = lifted_state(source, x1, s, births) # The same source draw serves all three legs. | |
| with torch.no_grad(): | |
| first = predict(model, xs, s, u, births, length, vocab) | |
| second = predict(model, step_map(xs, first, s, u, births), u, t, births, length, vocab) | |
| s_l, u_l, t_l = (local_time(v[:, None], births[None, :]) for v in (s, u, t)) | |
| weight = omega(s_l, u_l, t_l)[:, :, None] # [B, L, 1] one weight per coordinate. | |
| target = weight * first + (1.0 - weight) * second # Still on the simplex. | |
| student = predict(model, xs, s, t, births, length, vocab).clamp_min(1e-9) | |
| return diagonal + (target * (target.clamp_min(1e-9) / student).log()).sum(-1).mean() | |
| return loss_fn | |
| # %% 06. Sampling | |
| def sample(model, count, births, length, vocab, steps, device): | |
| """Walk a uniform grid, with every coordinate stepping on its own clock.""" | |
| x = torch.randn(count, length, vocab, device=device) # The source values, drawn once. | |
| grid = torch.linspace(0.0, 1.0, steps + 1) | |
| for a, b in zip(grid[:-1], grid[1:]): | |
| s, t = a.expand(count).to(device), b.expand(count).to(device) | |
| x = step_map(x, predict(model, x, s, t, births, length, vocab), s, t, births) | |
| return x.argmax(-1) # [N, L]. | |
| # %% 07. The report | |
| def main() -> None: | |
| parser = common_parser(__doc__.split("\n")[0]) | |
| parser.set_defaults(steps=1500, batch_size=128, samples=64, length=16, sample_steps=4) | |
| args = parser.parse_args() | |
| set_seed(args.seed) | |
| device = pick_device(args.device) | |
| tokens, _ = synthetic_dna(num_sequences=1024, length=args.length, seed=args.seed) | |
| data = tokens.to(device) # [1024, L] the full-length sequences. | |
| # Half the coordinates exist at the source; the rest are inserted at evenly spaced times. | |
| half = args.length // 2 | |
| births = torch.cat([torch.zeros(half, device=device), | |
| torch.linspace(0.1, 0.8, args.length - half, device=device)]) # [L]. | |
| report = worked_clocks() | |
| report.update(worked_omega()) | |
| report["coordinates at the source"] = int((births == 0).sum()) | |
| report["coordinates at the data"] = args.length | |
| source = torch.randn(256, args.length, DNA_SIZE) | |
| added, count = expansion_log_density(source, births.cpu(), 0.0, 0.5) | |
| report["coordinates inserted in (0, 0.5]"] = count | |
| report["source log density of those"] = added | |
| if not args.quiet: | |
| print(f"local clocks at t = 0.8: {report['local clocks at t=0.8']} with speeds " | |
| f"{report['clock speeds']}") | |
| print(f"omega at (0.3, 0.6, 0.9) = {report['omega at (0.3, 0.6, 0.9)']:.3f}") | |
| model = TwoTimeMLP(dim=args.length * (DNA_SIZE + 1), width=args.width).to(device) | |
| if not args.quiet: | |
| print(f"\ntraining the expanding map parameters={count_parameters(model):,}") | |
| losses = train(model, make_loss(model, data, births, args.length, DNA_SIZE, args.batch_size, | |
| device), steps=args.steps, lr=args.lr, quiet=args.quiet) | |
| samples: list[str] = [] | |
| early, late = births == 0, births > 0 # Coordinates born at the source and born later. | |
| for steps in (1, args.sample_steps): | |
| torch.manual_seed(args.seed + 1) | |
| drawn = sample(model, args.samples, births, args.length, DNA_SIZE, steps, device) | |
| report[f"{steps}-step GC fraction"] = gc_fraction(drawn) | |
| report[f"{steps}-step GC, born at the source"] = gc_fraction(drawn[:, early]) | |
| report[f"{steps}-step GC, born later"] = gc_fraction(drawn[:, late]) | |
| report[f"{steps}-step motif fraction"] = motif_fraction(drawn) | |
| report[f"{steps}-step token entropy"] = token_entropy(drawn, DNA_SIZE) | |
| report[f"{steps}-step unique fraction"] = unique_fraction(drawn) | |
| samples.append(f"{steps} step(s): {decode_dna(drawn[:1])[0]} GC {gc_fraction(drawn):.3f}") | |
| report["training GC fraction"] = gc_fraction(data) | |
| report["final loss"] = sum(losses[-20:]) / len(losses[-20:]) | |
| print_report("Expanding flow maps, Chapter 6", report, samples) | |
| save_run(args.out, model=model, config=vars(args), report=report, losses=losses, samples=samples) | |
| if __name__ == "__main__": | |
| main() | |