"""Differentiable quadratic tree elimination and synchronized rake contraction.""" from dataclasses import dataclass from typing import Dict, List, Tuple import torch Tensor = torch.Tensor @dataclass class EdgeFactor: # phi(xp, xc)=1/2 [xp,xc]' H [xp,xc] + [gp,gc]'[xp,xc] + k Hpp: Tensor Hpc: Tensor Hcc: Tensor gp: Tensor gc: Tensor k: Tensor def schur_message(Uc: Tensor, hc: Tensor, kc: Tensor, e: EdgeFactor, eps: float = 1e-6): """Eliminate child xc from Vc(xc)+phi(xp,xc), returning a parent message.""" K = e.Hcc + Uc + eps * torch.eye(Uc.shape[-1], device=Uc.device, dtype=Uc.dtype) # Child block and parent-child cross block of the combined quadratic. b = e.gc + hc S = e.Hpp - e.Hpc @ torch.linalg.solve(K, e.Hpc.transpose(-1, -2)) g = e.gp - e.Hpc @ torch.linalg.solve(K, b.unsqueeze(-1)).squeeze(-1) kk = e.k + kc - 0.5 * (b * torch.linalg.solve(K, b.unsqueeze(-1)).squeeze(-1)).sum(-1) # Symmetrization protects tiny floating point asymmetry. return (S + S.transpose(-1, -2)) * 0.5, g, kk def sequential_postorder(parent: List[int], U: Tensor, h: Tensor, k: Tensor, edges: Dict[int, EdgeFactor], eps: float = 1e-6): """Reference postorder elimination. parent[root] == -1; edge key is child.""" n = len(parent) children = [[] for _ in range(n)] root = -1 for c, p in enumerate(parent): if p < 0: root = c else: children[p].append(c) Us, hs, ks = list(U), list(h), list(k) order = [] def visit(i): for c in children[i]: visit(c) order.append(i) visit(root) saved = {} for c in order: p = parent[c] if p >= 0: saved[c] = (Us[c], hs[c], ks[c]) um, hm, km = schur_message(Us[c], hs[c], ks[c], edges[c], eps) Us[p], hs[p], ks[p] = Us[p] + um, hs[p] + hm, ks[p] + km return Us[root], hs[root], ks[root], saved, root def rake_levels(parent: List[int], U: Tensor, h: Tensor, k: Tensor, edges: Dict[int, EdgeFactor], eps: float = 1e-6): """Synchronous leaf-raking. All ready children in a level are eliminated together.""" n = len(parent) children = [[] for _ in range(n)] remaining = [0] * n root = -1 for c, p in enumerate(parent): if p < 0: root = c else: children[p].append(c); remaining[p] += 1 Us, hs, ks = list(U), list(h), list(k) active = [True] * n ready = [i for i in range(n) if remaining[i] == 0 and i != root] saved, levels = {}, [] while ready: levels.append(list(ready)) nxt = [] # Independent entries in one level could be stacked into one batched solve. for c in ready: p = parent[c] saved[c] = (Us[c], hs[c], ks[c]) um, hm, km = schur_message(Us[c], hs[c], ks[c], edges[c], eps) Us[p], hs[p], ks[p] = Us[p] + um, hs[p] + hm, ks[p] + km active[c] = False remaining[p] -= 1 if remaining[p] == 0 and p != root: nxt.append(p) ready = nxt return Us[root], hs[root], ks[root], saved, levels, root def reconstruct(parent, root, U, h, k, edges, saved, eps=1e-6): """MAP reconstruction after elimination, using stored child conditional blocks.""" children = [[] for _ in parent] for c,p in enumerate(parent): if p >= 0: children[p].append(c) z = [None] * len(parent) 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) def descend(p): for c in children[p]: Uc,hc,kc = saved[c] e = edges[c] K=e.Hcc+Uc+eps*torch.eye(Uc.shape[-1],device=Uc.device,dtype=Uc.dtype) rhs=e.gc+hc+e.Hpc.transpose(-1,-2)@z[p] z[c]=-torch.linalg.solve(K,rhs.unsqueeze(-1)).squeeze(-1) descend(c) descend(root) return torch.stack(z) def make_random_tree(n, d, seed=0, device='cpu', dtype=torch.float64): g=torch.Generator(device=device); g.manual_seed(seed) parent=[-1]+[int(torch.randint(0,i,(1,),generator=g,device=device)) for i in range(1,n)] eye=torch.eye(d,device=device,dtype=dtype) U=[]; h=[]; k=[]; edges={} for i in range(n): L=torch.randn(d,d,generator=g,device=device,dtype=dtype)*.15 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) for c,p in enumerate(parent): if p < 0: continue L=torch.randn(d,d,generator=g,device=device,dtype=dtype)*.12 Hcc=L@L.T+.4*eye C=torch.randn(d,d,generator=g,device=device,dtype=dtype)*.08 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)) return parent,torch.stack(U),torch.stack(h),torch.stack(k),edges