CIS6270 / lecture_6 /expanding.py
pranamanam's picture
Upload 87 files
0ba9d09 verified
Raw History Blame Contribute Delete
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
@torch.no_grad()
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()