"""Neural tied projector with material tokens (roadmap §1, W4-5). The neural generalization of TiedStrainBasis: z = Enc_θ(e, m) # e: co-rotated 6-dim strain, m: material token ê = Dec_θ(z, m) # rotation stays analytic, handled by the solver One network, weights shared across every tet AND every material — the tied-embedding thesis. A new material is a new token row, not a new network. Latent width k matches the PCA baselines it must beat. Per-material input scaling is part of the model (steel strains are ~30x smaller than foam strains; without it the loss and the network capacity are spent entirely on the softest material). """ import numpy as np import torch import torch.nn as nn class NeuralTiedProjector(nn.Module): def __init__(self, n_materials, k=3, d_token=4, hidden=48): super().__init__() self.k = k self.tokens = nn.Embedding(n_materials, d_token) # per-material log-scale, set from data statistics before training self.log_scale = nn.Parameter(torch.zeros(n_materials), requires_grad=False) self.enc = nn.Sequential( nn.Linear(6 + d_token, hidden), nn.SiLU(), nn.Linear(hidden, hidden), nn.SiLU(), nn.Linear(hidden, k), ) self.dec = nn.Sequential( nn.Linear(k + d_token, hidden), nn.SiLU(), nn.Linear(hidden, hidden), nn.SiLU(), nn.Linear(hidden, 6), ) def set_scales(self, scales): """scales: per-material RMS strain, computed from the training corpus.""" with torch.no_grad(): self.log_scale.copy_(torch.log(torch.as_tensor( scales, dtype=self.log_scale.dtype))) def _run(self, en, t): z = self.enc(torch.cat([en, t], dim=1)) return self.dec(torch.cat([z, t], dim=1)) def forward(self, e, mid): """e (N,6) raw strain, mid (N,) material ids -> reconstruction (N,6). Two architectural guards (not patches): - Rejection form f(e) = e - r(e): the network learns what to REMOVE, so admissible strains pass through at gain ~1. A reconstruction network attenuates within-manifold components (Jacobian gain < 1), which acts as artificial damping inside the solver loop; PCA's within-subspace gain is exactly 1 and this form mirrors it. - Zero-anchoring r(0) = 0: rest strain maps exactly to rest, otherwise reconstruction bias becomes spurious rest deformation. """ s = torch.exp(self.log_scale[mid]).unsqueeze(1) t = self.tokens(mid) en = e / s r = self._run(en, t) - self._run(torch.zeros_like(e), t) return (en - r) * s def n_params(self): return sum(p.numel() for p in self.parameters()) class WarmStartNet(nn.Module): """Learned warm start for the global solve (roadmap §5 step 4). Predicts this step's converged correction (x* - s) per vertex from history vectors (last two corrections, h·v, h²·a). Output is a gated linear combination of those vectors with gains computed from their pairwise dot products — rotation-equivariant by construction (the symmetry-handling lesson: don't make the network learn frame invariance) and tied across every vertex and mesh. The classical warm starts (copy previous correction; linear extrapolation) are exact special cases of this form. In W6 it is deployed as a residual ON TOP of linear extrapolation (x0 = s + 2c1 - c2 + net), so it starts from the best classical predictor and learns only what that predictor misses. Zero correctness risk: it only moves the PD loop's starting iterate; the fixed point is unchanged. """ def __init__(self, n_vec=4, hidden=32): super().__init__() self.n_vec = n_vec n_inv = n_vec * (n_vec + 1) // 2 self.mlp = nn.Sequential( nn.Linear(n_inv, hidden), nn.SiLU(), nn.Linear(hidden, hidden), nn.SiLU(), nn.Linear(hidden, n_vec), ) def forward(self, V): """V (N,n_vec,3) history vectors per vertex -> correction (N,3).""" dots = torch.einsum("nid,njd->nij", V, V) iu = torch.triu_indices(self.n_vec, self.n_vec) inv = dots[:, iu[0], iu[1]] # (N, n_inv) scale2 = dots.diagonal(dim1=1, dim2=2).sum(1).clamp_min(1e-24) gains = self.mlp(inv / scale2.unsqueeze(1)) # nondimensional in return torch.einsum("ni,nid->nd", gains, V) # equivariant out def predict_numpy(self, *vecs): with torch.no_grad(): V = torch.stack([torch.as_tensor(np.ascontiguousarray(v), dtype=torch.float32) for v in vecs], dim=1) return self(V).numpy().astype(np.float64) def n_params(self): return sum(p.numel() for p in self.parameters()) class MaterialProjector: """Adapter binding the network to one material id, exposing the same .project(e)->(N,6) numpy interface as TiedStrainBasis, so it drops straight into PDSolver3D(strain_basis=...).""" def __init__(self, net, material_id, device="cpu"): self.net = net.to(device).eval() self.mid = material_id self.device = device def project(self, e): with torch.no_grad(): et = torch.as_tensor(e, dtype=torch.float32, device=self.device) mid = torch.full((len(e),), self.mid, dtype=torch.long, device=self.device) return self.net(et, mid).cpu().numpy().astype(np.float64)