File size: 5,972 Bytes
ad9fbbf
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
import numpy as np
import torch
import torch.nn as nn
import torch.nn.functional as F

from .self_attention import Normalize, gather_edges, gather_nodes


class PositionalEncodings(nn.Module):
    def __init__(self, num_embeddings):
        super(PositionalEncodings, self).__init__()
        self.num_embeddings = num_embeddings

    def forward(self, E_idx):
        # i-j
        N_batch = E_idx.size(0)
        N_nodes = E_idx.size(1)
        N_neighbors = E_idx.size(2)
        device = E_idx.device
        ii = torch.arange(N_nodes, dtype=torch.float32, device=device).view((1, -1, 1))
        d = (E_idx.float() - ii).unsqueeze(-1)
        # Original Transformer frequencies
        frequency = torch.exp(
            torch.arange(0, self.num_embeddings, 2, dtype=torch.float32)
            * -(np.log(10000.0) / self.num_embeddings)).to(device)

        angles = d * frequency.view((1,1,1,-1))
        E = torch.cat((torch.cos(angles), torch.sin(angles)), -1)
        return E # [N_batch, N_nodes, N_neighbors, num_embeddings]


class EdgeFeatures(nn.Module):
    def __init__(self, edge_features, num_positional_embeddings=16,
        num_rbf=16, top_k=30, augment_eps=0.):
        super(EdgeFeatures, self).__init__()
        self.top_k = top_k
        self.augment_eps = augment_eps 
        self.num_rbf = num_rbf

        # Positional encoding
        self.PE = PositionalEncodings(num_positional_embeddings)
        
        # Embedding and normalization
        self.edge_embedding = nn.Linear(num_positional_embeddings + num_rbf + 7, edge_features, bias=True)
        self.norm_edges = Normalize(edge_features)

    def _dist(self, X, mask, eps=1E-6):
        """ Pairwise euclidean distances """
        mask_2D = torch.unsqueeze(mask,1) * torch.unsqueeze(mask,2) # mask [N, L] => mask_2D [N, L, L]
        dX = torch.unsqueeze(X,1) - torch.unsqueeze(X,2) # X 坐标矩阵 [N, L, 3]   dX 坐标差矩阵 [N, L, L, 3]
        D = mask_2D * torch.sqrt(torch.sum(dX**2, 3) + eps) # 距离矩阵 [N, L, L]

        # Identify k nearest neighbors (including self)
        D_max, _ = torch.max(D, -1, keepdim=True)
        D_adjust = D + (1. - mask_2D) * D_max
        D_neighbors, E_idx = torch.topk(D_adjust, self.top_k, dim=-1, largest=False) # [N, L, k]  D_neighbors为具体距离值(从小到大),E_idx为对应邻居节点的编号

        return D_neighbors, E_idx

    def _rbf(self, D):
        # Distance radial basis function
        D_min, D_max, D_count = 0., 20., self.num_rbf
        D_mu = torch.linspace(D_min, D_max, D_count, device=D.device)
        D_mu = D_mu.view([1,1,1,-1])
        D_sigma = (D_max - D_min) / D_count
        D_expand = torch.unsqueeze(D, -1)
        RBF = torch.exp(-((D_expand - D_mu) / D_sigma)**2)

        return RBF # [B, L, K, self.num_rbf]

    def _quaternions(self, R):
        """ Convert a batch of 3D rotations [R] to quaternions [Q]
            R [...,3,3]
            Q [...,4]
        """
        # Simple Wikipedia version
        # en.wikipedia.org/wiki/Rotation_matrix#Quaternion
        # For other options see math.stackexchange.com/questions/2074316/calculating-rotation-axis-from-rotation-matrix
        diag = torch.diagonal(R, dim1=-2, dim2=-1)
        Rxx, Ryy, Rzz = diag.unbind(-1)
        magnitudes = 0.5 * torch.sqrt(torch.abs(1 + torch.stack([
              Rxx - Ryy - Rzz, 
            - Rxx + Ryy - Rzz, 
            - Rxx - Ryy + Rzz
        ], -1)))
        _R = lambda i,j: R[:,:,:,i,j]
        signs = torch.sign(torch.stack([
            _R(2,1) - _R(1,2),
            _R(0,2) - _R(2,0),
            _R(1,0) - _R(0,1)
        ], -1))
        xyz = signs * magnitudes
        # The relu enforces a non-negative trace
        w = torch.sqrt(F.relu(1 + diag.sum(-1, keepdim=True))) / 2.
        Q = torch.cat((xyz, w), -1)
        Q = F.normalize(Q, dim=-1)

        return Q

    def _orientations(self, X, E_idx, eps=1e-6):
        # Shifted slices of unit vectors
        dX = X[:,1:,:] - X[:,:-1,:]
        U = F.normalize(dX, dim=-1) # 少了第一个(u0)
        u_2 = U[:,:-2,:] # u 1~n-2
        u_1 = U[:,1:-1,:] # u 2~n-1
        # Backbone normals
        n_2 = F.normalize(torch.cross(u_2, u_1), dim=-1) # n 1~n-2

        # Build relative orientations
        o_1 = F.normalize(u_2 - u_1, dim=-1) # b 角平分线向量
        O = torch.stack((o_1, n_2, torch.cross(o_1, n_2)), 2)
        O = O.view(list(O.shape[:2]) + [9])
        O = F.pad(O, (0,0,1,2), 'constant', 0) # [B, L, 9]


        O_neighbors = gather_nodes(O, E_idx) # [B, L, K, 9]
        X_neighbors = gather_nodes(X, E_idx) # [B, L, K, 3]
        
        # Re-view as rotation matrices
        O = O.view(list(O.shape[:2]) + [3,3]) # [B, L, 3, 3]
        O_neighbors = O_neighbors.view(list(O_neighbors.shape[:3]) + [3,3]) # [B, L, K, 3, 3]

        # Rotate into local reference frames
        dX = X_neighbors - X.unsqueeze(-2) # [B, L, K, 3]
        dU = torch.matmul(O.unsqueeze(2), dX.unsqueeze(-1)).squeeze(-1) # [B, L, K, 3]
        dU = F.normalize(dU, dim=-1)
        R = torch.matmul(O.unsqueeze(2).transpose(-1,-2), O_neighbors) # [B, L, K, 3, 3]
        Q = self._quaternions(R) # [B, L, K, 4]

        # Orientation features
        O_features = torch.cat((dU,Q), dim=-1) # [B, L, K, 7]

        return O_features


    def forward(self, X, mask): # X:[B, L, 3]  mask:[B, L]
        # Data augmentation
        if self.training and self.augment_eps > 0:
            X = X + self.augment_eps * torch.randn_like(X)

        # Build k-Nearest Neighbors graph
        D_neighbors, E_idx = self._dist(X, mask)

        # Pairwise features
        RBF = self._rbf(D_neighbors)
        O_features = self._orientations(X, E_idx)

        # Pairwise embeddings
        E_positional = self.PE(E_idx)

        E = torch.cat((E_positional, RBF, O_features), -1)
        E = self.edge_embedding(E)
        E = self.norm_edges(E)

        return E, E_idx # E [B, L, K, d]; E_idx [B, L, K]