| 1 | import numpy as np |
| 2 | |
| 3 | class Graph(): |
| 4 | def __init__(self, |
| 5 | layout, |
| 6 | strategy, |
| 7 | pad=0, |
| 8 | max_hop=1, |
| 9 | dilation=1): |
| 10 | |
| 11 | self.max_hop = max_hop |
| 12 | self.dilation = dilation |
| 13 | self.seqlen = pad |
| 14 | self.get_edge(layout) |
| 15 | self.hop_dis = get_hop_distance(self.num_node, self.edge, max_hop=max_hop) |
| 16 | |
| 17 | self.dist_center = self.get_distance_to_center(layout) |
| 18 | self.get_adjacency(strategy) |
| 19 | |
| 20 | def get_distance_to_center(self,layout): |
| 21 | dist_center = np.zeros(self.num_node) |
| 22 | if layout == 'hm36_gt': |
| 23 | for i in range(self.seqlen): |
| 24 | index_start = i*self.num_node_each |
| 25 | dist_center[index_start+0 : index_start+7] = [1, 2, 3, 4, 2, 3, 4] |
| 26 | dist_center[index_start+7 : index_start+11] = [0, 1, 2, 3] |
| 27 | dist_center[index_start+11 : index_start+17] = [2, 3, 4, 2, 3, 4] |
| 28 | return dist_center |
| 29 | |
| 30 | def __str__(self): |
| 31 | return self.A |
| 32 | |
| 33 | def graph_link_between_frames(self,base): |
| 34 | return [((front - 1) + i*self.num_node_each, (back - 1)+ i*self.num_node_each) for i in range(self.seqlen) for (front, back) in base] |
| 35 | |
| 36 | def basic_layout(self,neighbour_base, sym_base): |
| 37 | self.num_node = self.num_node_each * self.seqlen |
| 38 | time_link = [(i * self.num_node_each + j, (i + 1) * self.num_node_each + j) for i in range(self.seqlen - 1) |
| 39 | for j in range(self.num_node_each)] |
| 40 | self.time_link_forward = [(i * self.num_node_each + j, (i + 1) * self.num_node_each + j) for i in |
| 41 | range(self.seqlen - 1) |
| 42 | for j in range(self.num_node_each)] |
| 43 | self.time_link_back = [((i + 1) * self.num_node_each + j, (i) * self.num_node_each + j) for i in |
| 44 | range(self.seqlen - 1) |
| 45 | for j in range(self.num_node_each)] |
| 46 | |
| 47 | self_link = [(i, i) for i in range(self.num_node)] |
| 48 | |
| 49 | self.neighbour_link_all = self.graph_link_between_frames(neighbour_base) |
| 50 | |
| 51 | self.sym_link_all = self.graph_link_between_frames(sym_base) |
| 52 | |
| 53 | return self_link, time_link |
| 54 | |
| 55 | def get_edge(self, layout): |
| 56 | if layout == 'hm36_gt': |
| 57 | self.num_node_each = 17 |
| 58 | |
| 59 | |
| 60 | neighbour_base = [(1, 2), (3, 2), (4, 3), (5, 1), (6, 5), (7, 6), |