Parallel Quadratic Tree Layer / quadratic_tree.py
Beats tuned baseline
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