RG Pyramid Flow Matching / rg_pyramid_flow.py
Failed on benchmark
1import math
2import torch
3from torch import nn
4import torch.nn.functional as F
5
6
7def make_pyramid(x, levels):
8 # x: [B,1,L], average pooling is the dyadic coarse representation.
9 out = [x]
10 for _ in range(levels):
11 out.append(F.avg_pool1d(out[-1], 2, 2))
12 return out
13
14
15class LocalVelocity(nn.Module):
16 def __init__(self, radius=3, hidden=24):
17 super().__init__()
18 k = 2 * radius + 1
19 self.net = nn.Sequential(
20 nn.Conv1d(3, hidden, k, padding=radius), nn.SiLU(),
21 nn.Conv1d(hidden, hidden, k, padding=radius), nn.SiLU(),
22 nn.Conv1d(hidden, 1, 1))
23
24 def forward(self, z, t, cond=None):
25 b, _, n = z.shape
26 tt = t[:, None, None].expand(b, 1, n)
27 # cond is deliberately ignored by the baseline.
28 return self.net(torch.cat([z, tt, torch.zeros_like(z)], 1))
29
30
31class PyramidVelocity(nn.Module):
32 def __init__(self, levels=2, radius=2, hidden=20):
33 super().__init__()
34 self.levels = levels
35 self.nets = nn.ModuleList()
36 for s in range(levels + 1):
37 self.nets.append(nn.Sequential(
38 nn.Conv1d(3, hidden, 2 * radius + 1, padding=radius), nn.SiLU(),
39 nn.Conv1d(hidden, hidden, 2 * radius + 1, padding=radius), nn.SiLU(),
40 nn.Conv1d(hidden, 1, 1)))
41
42 def forward_level(self, z, t, cond, level):
43 b, _, n = z.shape
44 tt = t[:, None, None].expand(b, 1, n)
45 if cond is None: cond = torch.zeros_like(z)
46 return self.nets[level](torch.cat([z, cond, tt], 1))
47
48
49def correlated_data(n, batch, device, seed=2708):
50 g = torch.Generator(device=device).manual_seed(seed)
51 white = torch.randn(batch, 1, n, generator=g, device=device)
52 # Smooth low-frequency latent plus independent fine structure.
53 low = F.avg_pool1d(F.pad(white, (8, 8), mode='circular'), 17, 1)
54 fine = torch.randn(batch, 1, n, generator=g, device=device)
55 return 1.8 * low + 0.35 * fine
56
57
58def train_and_eval(steps=120, n=32, batch=32, levels=2, radius=2, device='cpu'):
59 torch.manual_seed(2708)
60 data = correlated_data(n, 700, device, 2708)
61 train, val = data[:600], data[600:]
62 base = LocalVelocity(radius=3, hidden=20).to(device)
63 pyr = PyramidVelocity(levels=levels, radius=radius, hidden=20).to(device)
64 ob = torch.optim.Adam(base.parameters(), lr=3e-3)
65 op = torch.optim.Adam(pyr.parameters(), lr=3e-3)
66 for step in range(steps):
67 ix = torch.randint(0, len(train), (batch,), device=device)
68 x1 = train[ix]; x0 = torch.randn_like(x1); t = torch.rand(batch, device=device)
69 z = (1-t[:,None,None])*x0 + t[:,None,None]*x1
70 target = x1-x0
71 ob.zero_grad(); lb = (base(z,t)-target).square().mean(); lb.backward(); ob.step()
72 zp = make_pyramid(z, levels); pred_losses=[]
73 op.zero_grad()
74 for s in range(levels+1):
75 cond = None if s == levels else F.interpolate(zp[s+1], size=zp[s].shape[-1], mode='linear', align_corners=False)
76 # Networks are indexed fine=0 through coarse=levels.
77 pred_losses.append((pyr.forward_level(zp[s],t,cond,s)-make_pyramid(target,levels)[s]).square().mean())
78 lp = sum(pred_losses)/(levels+1); lp.backward(); op.step()
79 with torch.no_grad():
80 x1=val; x0=torch.randn_like(x1); t=torch.rand(len(val),device=device)
81 z=(1-t[:,None,None])*x0+t[:,None,None]*x1; target=x1-x0
82 base_mse=(base(z,t)-target).square().mean().item()
83 zp=make_pyramid(z,levels); tp=make_pyramid(target,levels)
84 pm=[]
85 for s in range(levels+1):
86 cond=None if s==levels else F.interpolate(zp[s+1],size=zp[s].shape[-1],mode='linear',align_corners=False)
87 pm.append((pyr.forward_level(zp[s],t,cond,s)-tp[s]).square().mean().item())
88 # Correlation reconstruction: predict velocity at t=0, xhat=z+v ~= x1.
89 zb=torch.zeros_like(x1); tb=torch.zeros(len(val),device=device)
90 pred_b=base(zb,tb); pred_p=[]
91 zp0=make_pyramid(zb,levels)
92 for s in range(levels+1):
93 cond=None if s==levels else F.interpolate(zp0[s+1],size=zp0[s].shape[-1],mode='linear',align_corners=False)
94 pred_p.append(pyr.forward_level(zp0[s],tb,cond,s))
95 # Compare covariance correlation at distance 16 to data.
96 def corr(x,d=16):
97 a=x[..., :-d].flatten(); b=x[..., d:].flatten()
98 return F.cosine_similarity(a-a.mean(),b-b.mean(),dim=0).item()
99 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])
100 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())}
101
102if __name__ == '__main__':
103 import json
104 try:
105 device = 'cuda' if torch.cuda.is_available() else 'cpu'
106 try:
107 result = train_and_eval(device=device)
108 except Exception as e:
109 if device == 'cuda':
110 device = 'cpu'
111 result = train_and_eval(device=device)
112 else:
113 raise
114 result['device'] = device
115 result['radius_sweep'] = {}
116 for r in [1, 2, 3, 4]:
117 try:
118 q = train_and_eval(steps=60, n=32, batch=32, radius=r, device=device)
119 except Exception:
120 q = train_and_eval(steps=60, n=32, batch=32, radius=r, device='cpu')
121 result['radius_sweep'][str(r)] = q
122 open('flow_results.json','w').write(json.dumps(result, indent=2))
123 print(json.dumps(result, indent=2))
124 except Exception as e:
125 print(json.dumps({'error': repr(e)}))
126 raise