MCPcopy Create free account
hub / github.com/Vegetebird/GraphMLP / Graph

Class Graph

model/block/graph_frames.py:3–128  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

1import numpy as np
2
3class 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),

Callers 1

__init__Method · 0.90

Calls

no outgoing calls

Tested by

no test coverage detected