| 76 | |
| 77 | class LstmGcn(nn.Module): |
| 78 | def __init__(self, seq, n_fea, adj_dense): |
| 79 | super(LstmGcn, self).__init__() |
| 80 | self.A = adj_dense |
| 81 | self.nodes = adj_dense.shape[0] |
| 82 | self.encoder = nn.Conv2d(self.nodes, self.nodes, (n_fea, n_fea), device=device) |
| 83 | self.gcn_l1 = nn.Linear(seq - n_fea + 1, seq - n_fea + 1, device=device) |
| 84 | self.gcn_l2 = nn.Linear(seq - n_fea + 1, seq - n_fea + 1, device=device) |
| 85 | self.lstm = nn.LSTM(self.nodes, self.nodes, num_layers=2, batch_first=True) |
| 86 | self.act = nn.ReLU() |
| 87 | self.decoder = nn.Linear(seq - n_fea + 1, 1, device=device) |
| 88 | |
| 89 | # calculate A_delta matrix |
| 90 | deg = torch.sum(adj_dense, dim=0) |
| 91 | deg = torch.diag(deg) |
| 92 | deg_delta = torch.linalg.inv(torch.sqrt(deg)) |
| 93 | a_delta = torch.matmul(torch.matmul(deg_delta, adj_dense), deg_delta) |
| 94 | self.A = a_delta |
| 95 | |
| 96 | def forward(self, occ, prc): # occ.shape = [batch, node, seq] |
| 97 | x = torch.stack([occ, prc], dim=3) |