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