import math import torch from torch import nn import torch.nn.functional as F def make_pyramid(x, levels): # x: [B,1,L], average pooling is the dyadic coarse representation. out = [x] for _ in range(levels): out.append(F.avg_pool1d(out[-1], 2, 2)) return out class LocalVelocity(nn.Module): def __init__(self, radius=3, hidden=24): super().__init__() k = 2 * radius + 1 self.net = nn.Sequential( nn.Conv1d(3, hidden, k, padding=radius), nn.SiLU(), nn.Conv1d(hidden, hidden, k, padding=radius), nn.SiLU(), nn.Conv1d(hidden, 1, 1)) def forward(self, z, t, cond=None): b, _, n = z.shape tt = t[:, None, None].expand(b, 1, n) # cond is deliberately ignored by the baseline. return self.net(torch.cat([z, tt, torch.zeros_like(z)], 1)) class PyramidVelocity(nn.Module): def __init__(self, levels=2, radius=2, hidden=20): super().__init__() self.levels = levels self.nets = nn.ModuleList() for s in range(levels + 1): self.nets.append(nn.Sequential( nn.Conv1d(3, hidden, 2 * radius + 1, padding=radius), nn.SiLU(), nn.Conv1d(hidden, hidden, 2 * radius + 1, padding=radius), nn.SiLU(), nn.Conv1d(hidden, 1, 1))) def forward_level(self, z, t, cond, level): b, _, n = z.shape tt = t[:, None, None].expand(b, 1, n) if cond is None: cond = torch.zeros_like(z) return self.nets[level](torch.cat([z, cond, tt], 1)) def correlated_data(n, batch, device, seed=2708): g = torch.Generator(device=device).manual_seed(seed) white = torch.randn(batch, 1, n, generator=g, device=device) # Smooth low-frequency latent plus independent fine structure. low = F.avg_pool1d(F.pad(white, (8, 8), mode='circular'), 17, 1) fine = torch.randn(batch, 1, n, generator=g, device=device) return 1.8 * low + 0.35 * fine def train_and_eval(steps=120, n=32, batch=32, levels=2, radius=2, device='cpu'): torch.manual_seed(2708) data = correlated_data(n, 700, device, 2708) train, val = data[:600], data[600:] base = LocalVelocity(radius=3, hidden=20).to(device) pyr = PyramidVelocity(levels=levels, radius=radius, hidden=20).to(device) ob = torch.optim.Adam(base.parameters(), lr=3e-3) op = torch.optim.Adam(pyr.parameters(), lr=3e-3) for step in range(steps): ix = torch.randint(0, len(train), (batch,), device=device) x1 = train[ix]; x0 = torch.randn_like(x1); t = torch.rand(batch, device=device) z = (1-t[:,None,None])*x0 + t[:,None,None]*x1 target = x1-x0 ob.zero_grad(); lb = (base(z,t)-target).square().mean(); lb.backward(); ob.step() zp = make_pyramid(z, levels); pred_losses=[] op.zero_grad() for s in range(levels+1): cond = None if s == levels else F.interpolate(zp[s+1], size=zp[s].shape[-1], mode='linear', align_corners=False) # Networks are indexed fine=0 through coarse=levels. pred_losses.append((pyr.forward_level(zp[s],t,cond,s)-make_pyramid(target,levels)[s]).square().mean()) lp = sum(pred_losses)/(levels+1); lp.backward(); op.step() with torch.no_grad(): x1=val; x0=torch.randn_like(x1); t=torch.rand(len(val),device=device) z=(1-t[:,None,None])*x0+t[:,None,None]*x1; target=x1-x0 base_mse=(base(z,t)-target).square().mean().item() zp=make_pyramid(z,levels); tp=make_pyramid(target,levels) pm=[] for s in range(levels+1): cond=None if s==levels else F.interpolate(zp[s+1],size=zp[s].shape[-1],mode='linear',align_corners=False) pm.append((pyr.forward_level(zp[s],t,cond,s)-tp[s]).square().mean().item()) # Correlation reconstruction: predict velocity at t=0, xhat=z+v ~= x1. zb=torch.zeros_like(x1); tb=torch.zeros(len(val),device=device) pred_b=base(zb,tb); pred_p=[] zp0=make_pyramid(zb,levels) for s in range(levels+1): cond=None if s==levels else F.interpolate(zp0[s+1],size=zp0[s].shape[-1],mode='linear',align_corners=False) pred_p.append(pyr.forward_level(zp0[s],tb,cond,s)) # Compare covariance correlation at distance 16 to data. def corr(x,d=16): a=x[..., :-d].flatten(); b=x[..., d:].flatten() return F.cosine_similarity(a-a.mean(),b-b.mean(),dim=0).item() target_corr=corr(x1); cb=corr(pred_b); cp=corr(F.interpolate(pred_p[1],size=n,mode='linear',align_corners=False)+pred_p[0]) return {'baseline_velocity_mse':base_mse,'pyramid_velocity_mse_mean':sum(pm)/(levels+1), 'pyramid_level_mse':pm,'target_corr_d16':target_corr,'baseline_corr_d16':cb,'pyramid_corr_d16':cp,'baseline_params':sum(p.numel() for p in base.parameters()),'pyramid_params':sum(p.numel() for p in pyr.parameters())} if __name__ == '__main__': import json try: device = 'cuda' if torch.cuda.is_available() else 'cpu' try: result = train_and_eval(device=device) except Exception as e: if device == 'cuda': device = 'cpu' result = train_and_eval(device=device) else: raise result['device'] = device result['radius_sweep'] = {} for r in [1, 2, 3, 4]: try: q = train_and_eval(steps=60, n=32, batch=32, radius=r, device=device) except Exception: q = train_and_eval(steps=60, n=32, batch=32, radius=r, device='cpu') result['radius_sweep'][str(r)] = q open('flow_results.json','w').write(json.dumps(result, indent=2)) print(json.dumps(result, indent=2)) except Exception as e: print(json.dumps({'error': repr(e)})) raise