RG Pyramid Flow Matching / rg_pyramid_flow.py

Failed on benchmark

Raw ⬇ ZIP
  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