diff --git a/samples/python/mnist/graph_spliter.py b/samples/python/mnist/graph_spliter.py index f857f6d202..2f90bcde10 100644 --- a/samples/python/mnist/graph_spliter.py +++ b/samples/python/mnist/graph_spliter.py @@ -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') \ No newline at end of file +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') +"""