File size: 3,861 Bytes
24c2ab9
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
"""Rotation extraction from deformation gradients.

Primary path: Mueller et al. 2016, "A Robust Method to Extract the
Rotational Part of Deformations" — iterative quaternion refinement,
warm-started from the previous frame's rotation. Batched over all tets.

Guard path: SVD polar decomposition with reflection fix, used for tets
flagged as inverted/degenerate (det F <= eps). Never feed the solver a
reflection — that was the 2D inversion bug.
"""
import numpy as np


def quat_to_mat(q):
    """Batched quaternion (N,4) [w,x,y,z] -> rotation matrices (N,3,3)."""
    w, x, y, z = q[:, 0], q[:, 1], q[:, 2], q[:, 3]
    R = np.empty((q.shape[0], 3, 3))
    R[:, 0, 0] = 1 - 2 * (y * y + z * z)
    R[:, 0, 1] = 2 * (x * y - w * z)
    R[:, 0, 2] = 2 * (x * z + w * y)
    R[:, 1, 0] = 2 * (x * y + w * z)
    R[:, 1, 1] = 1 - 2 * (x * x + z * z)
    R[:, 1, 2] = 2 * (y * z - w * x)
    R[:, 2, 0] = 2 * (x * z - w * y)
    R[:, 2, 1] = 2 * (y * z + w * x)
    R[:, 2, 2] = 1 - 2 * (x * x + y * y)
    return R


def _quat_mul(a, b):
    w1, x1, y1, z1 = a[:, 0], a[:, 1], a[:, 2], a[:, 3]
    w2, x2, y2, z2 = b[:, 0], b[:, 1], b[:, 2], b[:, 3]
    return np.stack([
        w1 * w2 - x1 * x2 - y1 * y2 - z1 * z2,
        w1 * x2 + x1 * w2 + y1 * z2 - z1 * y2,
        w1 * y2 - x1 * z2 + y1 * w2 + z1 * x2,
        w1 * z2 + x1 * y2 - y1 * x2 + z1 * w2,
    ], axis=1)


def mueller_extract(F, q, iters=8):
    """Extract rotations from F (N,3,3), warm-started at quaternions q (N,4).

    Returns (R (N,3,3), q_new (N,4)). Pure NumPy, batched.
    """
    q = q.copy()
    for _ in range(iters):
        R = quat_to_mat(q)
        # omega = sum_i r_i x f_i / (|sum_i r_i . f_i| + eps), columns i
        cross = np.cross(R.transpose(0, 2, 1)[..., None, :],
                         F.transpose(0, 2, 1)[..., None, :]).sum(axis=(1, 2))
        dot = np.einsum("nij,nij->n", R, F)
        omega = cross / (np.abs(dot) + 1e-9)[:, None]
        angle = np.linalg.norm(omega, axis=1)
        small = angle < 1e-12
        axis = np.where(small[:, None], np.array([1.0, 0, 0]), omega / np.maximum(angle, 1e-12)[:, None])
        dq = np.concatenate([np.cos(angle / 2)[:, None],
                             np.sin(angle / 2)[:, None] * axis], axis=1)
        q = _quat_mul(dq, q)
        q /= np.linalg.norm(q, axis=1, keepdims=True)
        if np.all(angle < 1e-9):
            break
    return quat_to_mat(q), q


def svd_polar(F):
    """Batched SVD polar decomposition with reflection guard. F (N,3,3) -> R."""
    U, s, Vt = np.linalg.svd(F)
    det = np.linalg.det(U @ Vt)
    # flip the smallest singular direction to keep det(R) = +1
    U = U.copy()
    U[det < 0, :, 2] *= -1
    return U @ Vt


def mat_to_quat(R):
    """Batched rotation matrix -> quaternion [w,x,y,z] (numerically safe)."""
    N = R.shape[0]
    q = np.empty((N, 4))
    tr = R[:, 0, 0] + R[:, 1, 1] + R[:, 2, 2]
    # robust branch-free-ish: use the trace branch where valid, else fall back
    w = np.sqrt(np.maximum(1 + tr, 1e-12)) / 2
    q[:, 0] = w
    q[:, 1] = (R[:, 2, 1] - R[:, 1, 2]) / (4 * np.maximum(w, 1e-6))
    q[:, 2] = (R[:, 0, 2] - R[:, 2, 0]) / (4 * np.maximum(w, 1e-6))
    q[:, 3] = (R[:, 1, 0] - R[:, 0, 1]) / (4 * np.maximum(w, 1e-6))
    q /= np.linalg.norm(q, axis=1, keepdims=True)
    return q


def extract_rotations(F, q_cache, det_eps=1e-8, iters=8):
    """Rotation extraction with inversion guard.

    Healthy tets (det F > det_eps): Mueller iterative, warm-started.
    Inverted/degenerate tets: SVD polar with reflection fix, and the
    quaternion cache is reset from the SVD result.
    """
    det = np.linalg.det(F)
    bad = det <= det_eps
    R, q_new = mueller_extract(F, q_cache, iters=iters)
    if bad.any():
        R_bad = svd_polar(F[bad])
        R[bad] = R_bad
        q_new[bad] = mat_to_quat(R_bad)
    return R, q_new, bad