"""Unified-constitutive plasticity module (design.md 'UNIFIED CONSTITUTIVE FRAMEWORK'). Box plasticity on principal stretches of the elastic part (the Genesis-family form the data engine uses): trial F_e = F F_p^{-1}; SVD; clamp singular values to [y_lo, y_hi]; excess folds into F_p. Gradient safety: torch.linalg.svd's backward divides by (sigma_i^2 - sigma_j^2) and explodes on the (common!) repeated-stretch states — uniform squash has sigma_1=sigma_2. SafeSVD3 keeps torch's forward but hand-codes the standard SVD adjoint with the denominator clamped away from zero (|den| >= EPS * scale, sign preserved). Verified against finite differences in verify_plas.py, including repeated-sigma cases. Unclamped points NEVER route through the SVD graph (exact identity path), so with the yield window wide open the whole module is a bit-exact no-op vs the elastic-only code — the regression test.""" import torch EPS_DEN = 1e-4 class SafeSVD3(torch.autograd.Function): """batched 3x3 SVD with a repeated-singular-value-safe adjoint.""" @staticmethod def forward(ctx, A): U, S, Vh = torch.linalg.svd(A) # sign-fix to proper rotations is NOT needed for the clamp map (U diag(S) Vh reconstructs A # regardless); keep torch's convention. ctx.save_for_backward(U, S, Vh) return U, S, Vh @staticmethod def backward(ctx, gU, gS, gVh): """first-order SVD perturbation, derived from P = U^T dA V, dsigma_i = P_ii, Omega_U[ij] = (s_j P_ij + s_i P_ji)/(s_j^2 - s_i^2), dU = U Omega_U, Omega_V[ij] = (s_i P_ij + s_j P_ji)/(s_j^2 - s_i^2), dVh = -Omega_V V^T. Collecting + + = sum_ij W_ij P_ij gives gA = U W V^T with W_ij = ( s_j (a_ij - a_ji) + s_i (b_ij - b_ji) ) / (s_j^2 - s_i^2) (i != j) W_ii = gS_i, a = U^T gU, b = -gVh V. (FD-validated in verify_plas.py on sign-invariant observables, incl. repeated sigmas.) Denominator clamped away from zero (repeated-sigma safety), sign preserved.""" U, S, Vh = ctx.saved_tensors V = Vh.transpose(-1, -2) S2 = S * S den = S2.unsqueeze(-2) - S2.unsqueeze(-1) # den[i,j] = s_j^2 - s_i^2 scale = S2.max(dim=-1, keepdim=True).values.unsqueeze(-1) + 1e-30 sgn = torch.where(den >= 0, torch.ones_like(den), -torch.ones_like(den)) den = torch.where(den.abs() < EPS_DEN * scale, EPS_DEN * scale * sgn, den) a = U.transpose(-1, -2) @ gU if gU is not None else torch.zeros_like(U) b = -(gVh @ V) if gVh is not None else torch.zeros_like(U) Si = S.unsqueeze(-1) # s_i down rows Sj = S.unsqueeze(-2) # s_j across cols W = (Sj * (a - a.transpose(-1, -2)) + Si * (b - b.transpose(-1, -2))) / den W = W - torch.diag_embed(torch.diagonal(W, dim1=-2, dim2=-1)) if gS is not None: W = W + torch.diag_embed(gS) return U @ W @ Vh def return_map(Fe_trial, y_lo, y_hi): """box return map on principal stretches. Returns (Fe, need_mask (N,) bool). Points whose stretches sit inside [y_lo, y_hi] take the EXACT identity path (no SVD in their graph, bit-identical passthrough). y_lo/y_hi are 0-dim tensors (backbone params) — gradients flow into them from the clamped points only.""" with torch.no_grad(): # trigger test detached St = torch.linalg.svdvals(Fe_trial) need = ((St < y_lo.detach()) | (St > y_hi.detach())).any(-1) if not bool(need.any()): return Fe_trial, need idx = need.nonzero(as_tuple=True)[0] U, S, Vh = SafeSVD3.apply(Fe_trial[idx]) Sc = torch.clamp(S, min=y_lo, max=y_hi) # grads reach y_lo/y_hi here Fe = Fe_trial.clone() Fe[idx] = U @ torch.diag_embed(Sc) @ Vh return Fe, need