MCPcopy Create free account
hub / github.com/PetarV-/GAT / load_data

Function load_data

utils/process.py:45–96  ·  view source on GitHub ↗

Load data.

(dataset_str)

Source from the content-addressed store, hash-verified

43 return np.array(mask, dtype=np.bool)
44
45def load_data(dataset_str): # {'pubmed', 'citeseer', 'cora'}
46 """Load data."""
47 names = ['x', 'y', 'tx', 'ty', 'allx', 'ally', 'graph']
48 objects = []
49 for i in range(len(names)):
50 with open("data/ind.{}.{}".format(dataset_str, names[i]), 'rb') as f:
51 if sys.version_info > (3, 0):
52 objects.append(pkl.load(f, encoding='latin1'))
53 else:
54 objects.append(pkl.load(f))
55
56 x, y, tx, ty, allx, ally, graph = tuple(objects)
57 test_idx_reorder = parse_index_file("data/ind.{}.test.index".format(dataset_str))
58 test_idx_range = np.sort(test_idx_reorder)
59
60 if dataset_str == 'citeseer':
61 # Fix citeseer dataset (there are some isolated nodes in the graph)
62 # Find isolated nodes, add them as zero-vecs into the right position
63 test_idx_range_full = range(min(test_idx_reorder), max(test_idx_reorder)+1)
64 tx_extended = sp.lil_matrix((len(test_idx_range_full), x.shape[1]))
65 tx_extended[test_idx_range-min(test_idx_range), :] = tx
66 tx = tx_extended
67 ty_extended = np.zeros((len(test_idx_range_full), y.shape[1]))
68 ty_extended[test_idx_range-min(test_idx_range), :] = ty
69 ty = ty_extended
70
71 features = sp.vstack((allx, tx)).tolil()
72 features[test_idx_reorder, :] = features[test_idx_range, :]
73 adj = nx.adjacency_matrix(nx.from_dict_of_lists(graph))
74
75 labels = np.vstack((ally, ty))
76 labels[test_idx_reorder, :] = labels[test_idx_range, :]
77
78 idx_test = test_idx_range.tolist()
79 idx_train = range(len(y))
80 idx_val = range(len(y), len(y)+500)
81
82 train_mask = sample_mask(idx_train, labels.shape[0])
83 val_mask = sample_mask(idx_val, labels.shape[0])
84 test_mask = sample_mask(idx_test, labels.shape[0])
85
86 y_train = np.zeros(labels.shape)
87 y_val = np.zeros(labels.shape)
88 y_test = np.zeros(labels.shape)
89 y_train[train_mask, :] = labels[train_mask, :]
90 y_val[val_mask, :] = labels[val_mask, :]
91 y_test[test_mask, :] = labels[test_mask, :]
92
93 print(adj.shape)
94 print(features.shape)
95
96 return adj, features, y_train, y_val, y_test, train_mask, val_mask, test_mask
97
98def load_random_data(size):
99

Callers

nothing calls this directly

Calls 2

parse_index_fileFunction · 0.85
sample_maskFunction · 0.85

Tested by

no test coverage detected