mirror of
https://github.com/saymrwulf/onnxruntime.git
synced 2026-07-29 20:14:01 +00:00
sample code change.
This commit is contained in:
parent
934feb0c99
commit
6d8fde8324
1 changed files with 10 additions and 13 deletions
|
|
@ -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')
|
||||
"""
|
||||
|
|
|
|||
Loading…
Reference in a new issue