MCPcopy Create free account
hub / github.com/alibaba/euler / main

Function main

examples/solution/run_solution.py:49–93  ·  view source on GitHub ↗
(_)

Source from the content-addressed store, hash-verified

47
48
49def main(_):
50 flags_obj = tf.flags.FLAGS
51 euler_graph = tf_euler.dataset.get_dataset(flags_obj.dataset)
52 euler_graph.load_graph()
53
54 fanouts = list(map(int, flags_obj.fanouts))
55 assert flags_obj.layers == len(fanouts)
56 dims = [flags_obj.hidden_dim] * (flags_obj.layers + 1)
57 if flags_obj.run_mode == 'train':
58 metapath = [euler_graph.train_edge_type] * flags_obj.layers
59 else:
60 metapath = [euler_graph.all_edge_type] * flags_obj.layers
61 num_steps = int((euler_graph.total_size + 1) // flags_obj.batch_size *
62 flags_obj.num_epochs)
63
64 model = get_solution_model(conv=flags_obj.conv,
65 dataflow=flags_obj.flow,
66 dims=dims,
67 fanouts=fanouts,
68 metapath=metapath,
69 feature_idx=euler_graph.feature_idx,
70 feature_dim=euler_graph.feature_dim,
71 label_idx=euler_graph.label_idx,
72 label_dim=euler_graph.label_dim)
73
74 params = {'train_node_type': euler_graph.train_node_type[0],
75 'batch_size': flags_obj.batch_size,
76 'optimizer': flags_obj.optimizer,
77 'learning_rate': flags_obj.learning_rate,
78 'log_steps': flags_obj.log_steps,
79 'model_dir': flags_obj.model_dir,
80 'id_file': euler_graph.id_file,
81 'infer_dir': flags_obj.model_dir,
82 'total_step': num_steps}
83 config = tf.estimator.RunConfig(log_step_count_steps=None)
84 model_estimator = NodeEstimator(model, params, config)
85
86 if flags_obj.run_mode == 'train':
87 model_estimator.train()
88 elif flags_obj.run_mode == 'evaluate':
89 model_estimator.evaluate()
90 elif flags_obj.run_mode == 'infer':
91 model_estimator.infer()
92 else:
93 raise ValueError('Run mode not exist!')
94
95
96if __name__ == '__main__':

Callers

nothing calls this directly

Calls 6

get_solution_modelFunction · 0.90
NodeEstimatorClass · 0.90
load_graphMethod · 0.80
trainMethod · 0.80
evaluateMethod · 0.80
inferMethod · 0.80

Tested by

no test coverage detected