Parallel Quadratic Tree Layer / quadratic_tree.py

✓✓ Beats tuned baseline

Raw ⬇ ZIP
  1"""Differentiable quadratic tree elimination and synchronized rake contraction."""
  2from dataclasses import dataclass
  3from typing import Dict, List, Tuple
  4import torch
  5
  6Tensor = torch.Tensor
  7
  8@dataclass
  9class EdgeFactor:
 10    # phi(xp, xc)=1/2 [xp,xc]' H [xp,xc] + [gp,gc]'[xp,xc] + k
 11    Hpp: Tensor
 12    Hpc: Tensor
 13    Hcc: Tensor
 14    gp: Tensor
 15    gc: Tensor
 16    k: Tensor
 17
 18
 19def schur_message(Uc: Tensor, hc: Tensor, kc: Tensor, e: EdgeFactor,
 20                  eps: float = 1e-6):
 21    """Eliminate child xc from Vc(xc)+phi(xp,xc), returning a parent message."""
 22    K = e.Hcc + Uc + eps * torch.eye(Uc.shape[-1], device=Uc.device, dtype=Uc.dtype)
 23    # Child block and parent-child cross block of the combined quadratic.
 24    b = e.gc + hc
 25    S = e.Hpp - e.Hpc @ torch.linalg.solve(K, e.Hpc.transpose(-1, -2))
 26    g = e.gp - e.Hpc @ torch.linalg.solve(K, b.unsqueeze(-1)).squeeze(-1)
 27    kk = e.k + kc - 0.5 * (b * torch.linalg.solve(K, b.unsqueeze(-1)).squeeze(-1)).sum(-1)
 28    # Symmetrization protects tiny floating point asymmetry.
 29    return (S + S.transpose(-1, -2)) * 0.5, g, kk
 30
 31
 32def sequential_postorder(parent: List[int], U: Tensor, h: Tensor, k: Tensor,
 33                          edges: Dict[int, EdgeFactor], eps: float = 1e-6):
 34    """Reference postorder elimination. parent[root] == -1; edge key is child."""
 35    n = len(parent)
 36    children = [[] for _ in range(n)]
 37    root = -1
 38    for c, p in enumerate(parent):
 39        if p < 0: root = c
 40        else: children[p].append(c)
 41    Us, hs, ks = list(U), list(h), list(k)
 42    order = []
 43    def visit(i):
 44        for c in children[i]: visit(c)
 45        order.append(i)
 46    visit(root)
 47    saved = {}
 48    for c in order:
 49        p = parent[c]
 50        if p >= 0:
 51            saved[c] = (Us[c], hs[c], ks[c])
 52            um, hm, km = schur_message(Us[c], hs[c], ks[c], edges[c], eps)
 53            Us[p], hs[p], ks[p] = Us[p] + um, hs[p] + hm, ks[p] + km
 54    return Us[root], hs[root], ks[root], saved, root
 55
 56
 57def rake_levels(parent: List[int], U: Tensor, h: Tensor, k: Tensor,
 58                edges: Dict[int, EdgeFactor], eps: float = 1e-6):
 59    """Synchronous leaf-raking. All ready children in a level are eliminated together."""
 60    n = len(parent)
 61    children = [[] for _ in range(n)]
 62    remaining = [0] * n
 63    root = -1
 64    for c, p in enumerate(parent):
 65        if p < 0: root = c
 66        else: children[p].append(c); remaining[p] += 1
 67    Us, hs, ks = list(U), list(h), list(k)
 68    active = [True] * n
 69    ready = [i for i in range(n) if remaining[i] == 0 and i != root]
 70    saved, levels = {}, []
 71    while ready:
 72        levels.append(list(ready))
 73        nxt = []
 74        # Independent entries in one level could be stacked into one batched solve.
 75        for c in ready:
 76            p = parent[c]
 77            saved[c] = (Us[c], hs[c], ks[c])
 78            um, hm, km = schur_message(Us[c], hs[c], ks[c], edges[c], eps)
 79            Us[p], hs[p], ks[p] = Us[p] + um, hs[p] + hm, ks[p] + km
 80            active[c] = False
 81            remaining[p] -= 1
 82            if remaining[p] == 0 and p != root: nxt.append(p)
 83        ready = nxt
 84    return Us[root], hs[root], ks[root], saved, levels, root
 85
 86
 87def reconstruct(parent, root, U, h, k, edges, saved, eps=1e-6):
 88    """MAP reconstruction after elimination, using stored child conditional blocks."""
 89    children = [[] for _ in parent]
 90    for c,p in enumerate(parent):
 91        if p >= 0: children[p].append(c)
 92    z = [None] * len(parent)
 93    z[root] = -torch.linalg.solve(U[root] + eps*torch.eye(U.shape[-1],device=U.device,dtype=U.dtype), h[root].unsqueeze(-1)).squeeze(-1)
 94    def descend(p):
 95        for c in children[p]:
 96            Uc,hc,kc = saved[c]
 97            e = edges[c]
 98            K=e.Hcc+Uc+eps*torch.eye(Uc.shape[-1],device=Uc.device,dtype=Uc.dtype)
 99            rhs=e.gc+hc+e.Hpc.transpose(-1,-2)@z[p]
100            z[c]=-torch.linalg.solve(K,rhs.unsqueeze(-1)).squeeze(-1)
101            descend(c)
102    descend(root)
103    return torch.stack(z)
104
105
106def make_random_tree(n, d, seed=0, device='cpu', dtype=torch.float64):
107    g=torch.Generator(device=device); g.manual_seed(seed)
108    parent=[-1]+[int(torch.randint(0,i,(1,),generator=g,device=device)) for i in range(1,n)]
109    eye=torch.eye(d,device=device,dtype=dtype)
110    U=[]; h=[]; k=[]; edges={}
111    for i in range(n):
112        L=torch.randn(d,d,generator=g,device=device,dtype=dtype)*.15
113        U.append(L@L.T+.5*eye); h.append(torch.randn(d,generator=g,device=device,dtype=dtype)*.2); k.append(torch.randn((),generator=g,device=device,dtype=dtype)*.1)
114    for c,p in enumerate(parent):
115        if p < 0: continue
116        L=torch.randn(d,d,generator=g,device=device,dtype=dtype)*.12
117        Hcc=L@L.T+.4*eye
118        C=torch.randn(d,d,generator=g,device=device,dtype=dtype)*.08
119        edges[c]=EdgeFactor(torch.zeros_like(eye), C, Hcc, torch.randn(d,generator=g,device=device,dtype=dtype)*.05, torch.randn(d,generator=g,device=device,dtype=dtype)*.05, torch.zeros((),device=device,dtype=dtype))
120    return parent,torch.stack(U),torch.stack(h),torch.stack(k),edges