sample code change.

This commit is contained in:
Vincent Wang 2020-11-05 05:58:31 +00:00 committed by Thiago Crepaldi
parent 934feb0c99
commit 6d8fde8324

View file

@ -120,7 +120,6 @@ def split_graph(onnx_model):
# MNIST
"""
original_model = onnx.load('mnist_original.onnx')
config = C.ModuleGradientGraphBuilderConfiguration()
weight_names_to_train = set()
@ -132,14 +131,12 @@ for output in original_model.graph.output:
output_names.add(output.name)
config.output_names = output_names
gradient_graph_model = onnx.load_model_from_string(C.ModuleGradientGraphBuilder().build(original_model.SerializeToString(), config))
onnx.save(gradient_graph_model, 'minst_gradient_graph.onnx')
forward_model, backward_model = split_graph(gradient_graph_model)
onnx.save(forward_model, 'mnist_forward.onnx')
onnx.save(backward_model, 'mnist_backward.onnx')
models = [onnx.load_model_from_string(model_as_string) for model_as_string in C.ModuleGradientGraphBuilder().build_and_split(original_model.SerializeToString(), config)]
onnx.save(models[0], 'minst_gradient_graph.onnx')
onnx.save(models[1], 'mnist_forward.onnx')
onnx.save(models[2], 'mnist_backward.onnx')
"""
#BERT
original_model = onnx.load('bert-tiny.onnx')
config = C.ModuleGradientGraphBuilderConfiguration()
@ -152,8 +149,8 @@ for output in original_model.graph.output:
output_names.add(output.name)
config.output_names = output_names
gradient_graph_model = onnx.load_model_from_string(C.ModuleGradientGraphBuilder().build(original_model.SerializeToString(), config))
onnx.save(gradient_graph_model, 'bert_gradient_graph.onnx')
forward_model, backward_model = split_graph(gradient_graph_model)
onnx.save(forward_model, 'bert_forward.onnx')
onnx.save(backward_model, 'bert_backward.onnx')
models = [onnx.load_model_from_string(model_as_string) for model_as_string in C.ModuleGradientGraphBuilder().build_and_split(original_model.SerializeToString(), config)]
onnx.save(models[0], 'bert_gradient_graph.onnx')
onnx.save(models[1], 'bert_forward.onnx')
onnx.save(models[2], 'bert_backward.onnx')
"""