Download lecture_2/flow_matching_2d.py from ChatterjeeLab/CIS6270: direct link, hf CLI and curl.
- Browser
- Download file 14.4 kB
-
https://huggingface.co/ChatterjeeLab/CIS6270/resolve/main/lecture_2/flow_matching_2d.py
- Command line
-
hf download hf://ChatterjeeLab/CIS6270/lecture_2/flow_matching_2d.py
-
curl -L -o flow_matching_2d.py https://huggingface.co/ChatterjeeLab/CIS6270/resolve/main/lecture_2/flow_matching_2d.py
14.4 kB
| """Flow matching on four Gaussians, the whole method end to end. | |
| A coupling pairs a Gaussian prior sample with a data sample, a straight line joins the pair, and the | |
| derivative of that line is the regression target. We recover the marginal velocity field by least | |
| squares against those per-pair velocities, and then take Euler steps along the learned field to turn | |
| fresh noise into data. The four-Gaussian target is small enough that the marginal velocity has a | |
| closed form, so we can measure the conditional flow-matching equivalence here. | |
| Chapter 2, Sections 2.11 to 2.16 and 2.19 of the notes. Trained on | |
| ``cis6270.data.gaussian_mixture_2d``. Prints the worked endpoint pair, the average at the crossing, | |
| the two losses with their constant gap, the agreement of their gradients, and the mode coverage and | |
| mean exact log density of the generated points. | |
| python lecture_2/flow_matching_2d.py # train and generate, about ten seconds | |
| python lecture_2/flow_matching_2d.py --steps 20 # a first look, with no time to learn | |
| """ | |
| # %% 01. Imports and configuration | |
| import argparse | |
| import torch | |
| from cis6270.data import GAUSSIAN_CENTERS, GAUSSIAN_STD, gaussian_mixture_2d | |
| from cis6270.metrics import mean_log_density, mode_report | |
| from cis6270.nets import VelocityMLP, count_parameters | |
| from cis6270.runner import print_report, save_run, set_seed, train | |
| DIM = 2 # d, the width of the state. | |
| TRAIN_STEPS = 1000 # Optimizer steps, each on one fresh batch of endpoint pairs. | |
| BATCH_SIZE = 256 # Endpoint pairs per step. | |
| LEARNING_RATE = 3e-3 # Optimizer step size. | |
| WIDTH = 128 # Hidden width of the velocity network. | |
| SAMPLE_STEPS = 64 # Euler steps from t = 0 to t = 1, so Delta t = 1/64. | |
| NUM_SAMPLES = 1024 # Points generated for the mode coverage and the log density. | |
| DATA_POINTS = 4096 # Size of the training set drawn from the four-Gaussian target. | |
| EQUIVALENCE_POINTS = 16384 # Pairs used for the Monte Carlo check of Theorem 2.1. | |
| GRID = torch.tensor([0.0, 0.25, 0.50, 0.75, 1.0]) # [5] the chapter's grid of times. | |
| # %% 02. The coupling | |
| # | |
| # Section 2.11. A coupling is any joint law over the endpoint pair whose marginals are the prior and | |
| # the data. We draw the two ends with no reference to each other. That gives the independent | |
| # coupling. | |
| def sample_coupling(data, batch_size): | |
| """Draw (X_0, X_1) ~ pi with pi(x_0, x_1) = p_0(x_0) p_1(x_1). Returns two [B, d] tensors.""" | |
| index = torch.randint(0, len(data), (batch_size,), device=data.device) # [B] which data point. | |
| x_1 = data[index] # [B, d] the data endpoint, drawn from p_1. | |
| x_0 = torch.randn_like(x_1) # [B, d] the prior endpoint, drawn from p_0 = N(0, I). | |
| return x_0, x_1 | |
| # %% 03. The conditional path and its velocity | |
| # | |
| # Section 2.12. The straight segment at constant speed, and its derivative with the endpoints held | |
| # fixed. The velocity is called conditional because only someone who knows both endpoints has it. | |
| def conditional_path(x_0, x_1, t): | |
| """X_t = (1 - t) X_0 + t X_1 and U_t = X_1 - X_0. Returns two [B, d] tensors.""" | |
| x_t = (1.0 - t)[:, None] * x_0 + t[:, None] * x_1 # [B, d] convex combination of the endpoints. | |
| u_t = x_1 - x_0 # [B, d] the displacement, the same vector at every time. | |
| return x_t, u_t | |
| def worked_pair(): | |
| """The chapter's one-dimensional pair x_0 = -2, x_1 = 3, evaluated on the grid of times.""" | |
| x_0 = torch.full((len(GRID), 1), -2.0) # [5, 1] the same prior endpoint at every grid time. | |
| x_1 = torch.full((len(GRID), 1), 3.0) # [5, 1] the same data endpoint at every grid time. | |
| x_t, u_t = conditional_path(x_0, x_1, GRID) # The line x_t = -2 + 5t and its slope. | |
| return { | |
| "worked_pair_path": [round(float(v), 4) for v in x_t[:, 0]], | |
| "worked_pair_velocity": float(u_t[0, 0]), # The slope 5, so 1.25 units in every quarter. | |
| } | |
| # %% 04. The marginal velocity | |
| # | |
| # Section 2.13. Several conditional paths pass through one state at one time, each with its own | |
| # U_t, and a deterministic field assigns one vector there, the average v_t(x) = E[U_t | X_t = x]. | |
| # We can take that average exactly for this target, and Section 2.14 handles the case where we | |
| # cannot. Once we condition on component k, (X_1, X_t) is jointly Gaussian, so the posterior over | |
| # components is a softmax of the component log densities at X_t and v_t follows in closed form. | |
| def marginal_velocity(x_t, t): | |
| """v_t(x) = E[U_t | X_t = x] for a standard Gaussian prior and the four-Gaussian target.""" | |
| centers = GAUSSIAN_CENTERS.to(x_t) # [M, d] the four means of the mixture. | |
| variance = (1.0 - t) ** 2 + (t * GAUSSIAN_STD) ** 2 # [B] variance of X_t given component k. | |
| offset = x_t[:, None, :] - t[:, None, None] * centers # [B, M, d] distance to each shifted mean. | |
| log_weight = -0.5 * (offset ** 2).sum(-1) / variance[:, None] # [B, M] up to a shared constant. | |
| weight = torch.softmax(log_weight, dim=1) # [B, M] posterior over components given X_t. | |
| mean_center = (weight[:, :, None] * centers).sum(1) # [B, d] posterior mean of the center. | |
| # We write E[X_1 | X_t] in terms of that posterior mean and cancel the factor (1 - t) that | |
| # U_t = (X_1 - X_t)/(1 - t) would otherwise divide by. The expression that remains stays | |
| # finite at t = 1. | |
| return ((1.0 - t)[:, None] * mean_center | |
| - (1.0 - t - t * GAUSSIAN_STD ** 2)[:, None] * x_t) / variance[:, None] # [B, d]. | |
| def crossing_average(): | |
| """Two paths of Section 2.13 that meet at x = 2 at t = 0.5, with velocities 4 and -2.""" | |
| x_0 = torch.tensor([[0.0], [3.0]]) # [2, 1] the two prior endpoints. | |
| x_1 = torch.tensor([[4.0], [1.0]]) # [2, 1] the two data endpoints. | |
| x_t, u_t = conditional_path(x_0, x_1, torch.full((2,), 0.5)) # Both states land on x = 2. | |
| return { | |
| "crossing_state": float(x_t.mean()), # Both paths occupy this state at t = 0.5. | |
| "crossing_velocities": [round(float(v), 4) for v in u_t[:, 0]], | |
| "crossing_average_velocity": float(u_t.mean()), # The equally weighted average, 1. | |
| } | |
| # %% 05. The two objectives | |
| # | |
| # Sections 2.14 and 2.15. The flow-matching loss regresses on the marginal velocity, and we do | |
| # not have it for real data. The conditional loss regresses on the sampled U_t, and we do have it. | |
| def objectives(model, x_0, x_1, t): | |
| """L_CFM, L_FM, the gap E||U_t - v_t||^2 and the cross term, on one batch of endpoint pairs.""" | |
| x_t, u_t = conditional_path(x_0, x_1, t) # [B, d] the state and its conditional velocity. | |
| v_pred = model(x_t, t) # [B, d] the network reads the state and the time, never the endpoints. | |
| v_t = marginal_velocity(x_t, t) # [B, d] the target of the loss we cannot evaluate in general. | |
| loss_cfm = (v_pred - u_t).square().sum(1).mean() # E||v_theta(X_t, t) - U_t||^2. | |
| loss_fm = (v_pred - v_t).square().sum(1).mean() # E||v_theta(X_t, t) - v_t(X_t)||^2. | |
| gap = (u_t - v_t).square().sum(1).mean() # The conditional variance of the target. | |
| # We expand the square to get L_CFM = L_FM + cross + gap on any batch. The proof of Theorem 2.1 | |
| # rests on the cross term having conditional mean zero, so on a large batch it is near zero. | |
| cross = 2.0 * ((v_pred - v_t) * (v_t - u_t)).sum(1).mean() | |
| return loss_cfm, loss_fm, gap, cross | |
| # %% 06. The equivalence, measured | |
| # | |
| # Section 2.15. At the crossing the two losses are parabolas in the prediction a, nine units apart | |
| # at every a, with one derivative between them. The same identity holds in expectation on the | |
| # two-dimensional problem, up to the Monte Carlo error of the batch that estimates it. | |
| def crossing_gap(predictions=(0.0, 1.0, 2.0)): | |
| """The conditional and marginal losses at x = 2, t = 0.5, where the targets are 4 and -2.""" | |
| targets = torch.tensor([4.0, -2.0]) # [2] the two conditional velocities meeting there. | |
| marginal = targets.mean() # The conditional mean of the target, equal to 1. | |
| conditional_values, marginal_values = [], [] | |
| for a in predictions: # One candidate prediction at a time. | |
| conditional_values.append(round(float((a - targets).square().mean()), 4)) | |
| marginal_values.append(round(float((a - marginal) ** 2), 4)) | |
| return { | |
| "crossing_conditional_loss": conditional_values, | |
| "crossing_marginal_loss": marginal_values, | |
| "crossing_constant_gap": [round(c - m, 4) for c, m in | |
| zip(conditional_values, marginal_values)], | |
| "crossing_target_variance": float(targets.var(unbiased=False)), # The same constant, 9. | |
| } | |
| def equivalence_check(model, data, num_points): | |
| """Theorem 2.1 on the four-Gaussian problem: the gap in value and the agreement in gradient.""" | |
| x_0, x_1 = sample_coupling(data, num_points) # [N, d] pairs from the independent coupling. | |
| t = torch.rand(num_points, device=x_1.device) # [N] one uniform time per pair. | |
| loss_cfm, loss_fm, gap, cross = objectives(model, x_0, x_1, t) | |
| parameters = list(model.parameters()) | |
| grad_cfm = torch.autograd.grad(loss_cfm, parameters, retain_graph=True) # Gradient of L_CFM. | |
| grad_fm = torch.autograd.grad(loss_fm, parameters) # Gradient of L_FM, on the same batch. | |
| flat_cfm = torch.cat([g.flatten() for g in grad_cfm]) # [P] one vector per objective. | |
| flat_fm = torch.cat([g.flatten() for g in grad_fm]) # [P]. | |
| return { | |
| "loss_cfm": float(loss_cfm.detach()), | |
| "loss_fm": float(loss_fm.detach()), | |
| "loss_gap": float(gap.detach()), # The constant that separates the two objectives. | |
| "loss_cross_term": float(cross.detach()), # Zero in expectation, so small on this batch. | |
| "loss_fm_plus_gap": float((loss_fm + gap).detach()), # L_CFM once the cross term is added. | |
| "gradient_cosine": float(torch.nn.functional.cosine_similarity(flat_cfm, flat_fm, dim=0)), | |
| # The two gradients differ by the gradient of that cross term. That difference falls | |
| # like one over the square root of the batch size. | |
| "gradient_relative_difference": float((flat_cfm - flat_fm).norm() / flat_fm.norm()), | |
| } | |
| # %% 07. Training | |
| # | |
| # Section 2.19. One iteration draws a pair and a time, builds the state in closed form, subtracts | |
| # the endpoints, and takes a gradient step. We never solve an ODE to build a training example. | |
| def make_loss(model, data, batch_size): | |
| """The closure runner.train calls, returning L_CFM on one fresh batch.""" | |
| def loss_fn(step): | |
| x_0, x_1 = sample_coupling(data, batch_size) # [B, d] a fresh pair per step. | |
| t = torch.rand(batch_size, device=x_1.device) # [B] t ~ Uniform(0, 1), one per pair. | |
| x_t, u_t = conditional_path(x_0, x_1, t) # [B, d] the state and the regression target. | |
| return (model(x_t, t) - u_t).square().sum(1).mean() # ||v_theta(X_t, t) - U_t||^2. | |
| return loss_fn | |
| # %% 08. Euler sampling | |
| # | |
| # Section 2.19. We start from a fresh prior sample and integrate the learned field, so one sample | |
| # costs as many forward passes as there are steps. | |
| def sample(model, num_samples, steps): | |
| """Integrate dX_t/dt = v_theta(X_t, t) from t = 0 to t = 1 on a uniform grid. Returns [N, d].""" | |
| device = next(model.parameters()).device # Use the model's own device. | |
| x_t = torch.randn(num_samples, DIM, device=device) # [N, d] the draw X_0 ~ p_0. | |
| dt = 1.0 / steps # Delta t, the width of one interval. | |
| for step in range(steps): # One network evaluation per step. | |
| t = torch.full((num_samples,), step * dt, device=device) # [N] the shared grid time. | |
| x_t = x_t + dt * model(x_t, t) # Euler: state += time increment times velocity. | |
| return x_t | |
| # %% 09. The report | |
| def main(): | |
| parser = argparse.ArgumentParser(description=__doc__.split("\n")[0]) | |
| parser.add_argument("--steps", type=int, default=TRAIN_STEPS, help="optimizer steps") | |
| parser.add_argument("--batch-size", type=int, default=BATCH_SIZE, help="pairs per step") | |
| parser.add_argument("--lr", type=float, default=LEARNING_RATE, help="learning rate") | |
| parser.add_argument("--width", type=int, default=WIDTH, help="hidden width of the network") | |
| parser.add_argument("--samples", type=int, default=NUM_SAMPLES, help="points to generate") | |
| parser.add_argument("--sample-steps", type=int, default=SAMPLE_STEPS, help="Euler steps") | |
| parser.add_argument("--seed", type=int, default=0, help="random seed") | |
| parser.add_argument("--out", type=str, default="", help="directory for the run's output") | |
| parser.add_argument("--quiet", action="store_true", help="print only the final report") | |
| args = parser.parse_args() | |
| set_seed(args.seed) # Seed first, so the report repeats under the same seed. | |
| data = gaussian_mixture_2d(DATA_POINTS, seed=args.seed) # [N, d] the target p_1 = p_data. | |
| model = VelocityMLP(dim=DIM, width=args.width) # v_theta(x, t), the learned velocity field. | |
| report = {"parameters": count_parameters(model)} | |
| report.update(worked_pair()) # Section 2.12. | |
| report.update(crossing_average()) # Section 2.13. | |
| report.update(crossing_gap()) # Section 2.15. | |
| # The gradient comparison runs before training, where the two gradients are far from zero. | |
| report.update(equivalence_check(model, data, EQUIVALENCE_POINTS)) # Section 2.15. | |
| losses = train(model, make_loss(model, data, args.batch_size), | |
| steps=args.steps, lr=args.lr, quiet=args.quiet) | |
| generated = sample(model, args.samples, args.sample_steps) # [N, d] fresh points. | |
| coverage = mode_report(generated) # Nearest mode, how many modes, how evenly they are visited. | |
| report.update({ | |
| "final_loss": sum(losses[-100:]) / len(losses[-100:]), # Averaged over the last steps. | |
| "modes_found": coverage["modes_found"], | |
| "mode_entropy": coverage["mode_entropy"], # log 4 = 1.3863 when all four are equally used. | |
| "nearest_mode_distance": coverage["nearest_distance"], | |
| "mode_fractions": coverage["mode_fractions"], | |
| "mean_log_density": mean_log_density(generated), # Exact log density of the target. | |
| "data_mean_log_density": mean_log_density(data), # The same measurement on real points. | |
| }) | |
| print_report("flow matching on four Gaussians", report) | |
| save_run(args.out, model=model, config=vars(args), report=report, losses=losses) | |
| # %% 10. Run the script | |
| if __name__ == "__main__": | |
| main() | |