-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathtransforms.py
More file actions
57 lines (42 loc) · 1.59 KB
/
Copy pathtransforms.py
File metadata and controls
57 lines (42 loc) · 1.59 KB
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
import numpy as np
import torch
import pdb
def calc_rot_mat(pos):
# Convert to numpy array
if isinstance(pos, torch.Tensor):
pos = pos.numpy()
# Shift pos to be centered
pos_centered = pos - pos.mean(0).reshape((1, -1))
# SVD decomposition
_, _, vt = np.linalg.svd(pos_centered)
# Each row of vt is a principle vector
# vt is orthonormal
# The rotation matrix is simply the transposed vt
# First, get a tentatively transformed result
rot_mat = vt.T
trans = pos_centered @ rot_mat
# Take the farthest node from the origin as the reference node
ref_node = np.argmax(np.linalg.norm(trans, axis=1))
ref_coord = trans[ref_node]
# vt from SVD can have vectors of arbitrary directions
# Invert axes to make ref_node lie in the first quadrant/octant using a mask
mask = np.ones(pos.shape[1])
mask[ref_coord < 0] = -1
mask = mask.reshape((1, pos.shape[1]))
# The final rot_mat and trans
rot_mat = rot_mat * mask
return torch.from_numpy(rot_mat).float()
if __name__ == '__main__':
pos = np.random.rand(10, 3)
new_x = np.random.rand(3)
new_x = new_x / np.linalg.norm(new_x)
new_z = np.cross(new_x, np.random.rand(3))
new_z = new_z / np.linalg.norm(new_z)
new_y = np.cross(new_z, new_x)
new_y = new_y / np.linalg.norm(new_y)
rotation = np.linalg.inv(np.stack([new_x, new_y, new_z]))
_pos = pos @ rotation + np.random.rand(3)
invariant = get_invariant_pos(pos)
_invariant = get_invariant_pos(_pos)
assert np.allclose(invariant['trans'], _invariant['trans'])
print('Passed.')