mirror of
https://github.com/saymrwulf/onnxruntime.git
synced 2026-07-27 20:02:15 +00:00
Clean ORTModule dev branch (#6944)
This commit is contained in:
parent
48eebed869
commit
5303b33f69
8 changed files with 3 additions and 248 deletions
|
|
@ -2,28 +2,11 @@ import importlib.util
|
|||
import numpy as np
|
||||
import os
|
||||
import sys
|
||||
import time
|
||||
import torch
|
||||
from onnx import TensorProto
|
||||
|
||||
from functools import wraps
|
||||
|
||||
def timeit(enabled=True):
|
||||
def noop_inner(my_func):
|
||||
return my_func
|
||||
|
||||
def inner(my_func):
|
||||
@wraps(my_func)
|
||||
def timed(*args, **kw):
|
||||
tstart = time.time()
|
||||
output = my_func(*args, **kw)
|
||||
tend = time.time()
|
||||
|
||||
print('{}: took {:.3f}ms to execute'.format(my_func.__name__, (tend - tstart) * 1000))
|
||||
return output
|
||||
return timed
|
||||
|
||||
return inner if enabled else noop_inner
|
||||
|
||||
def get_device_index(device):
|
||||
if isinstance(device, str):
|
||||
|
|
@ -42,16 +25,6 @@ def get_device_index_from_input(input):
|
|||
device_index = get_device_index(input.device)
|
||||
return device_index
|
||||
|
||||
def get_device_from_input_args_kwargs(*args, **kwargs):
|
||||
'''Returns device index from first PyTorch Tensor within *args or **kwargs'''
|
||||
|
||||
device = None
|
||||
if args:
|
||||
device = torch.device(args[0].device)
|
||||
if not device and kwargs:
|
||||
device = torch.device(next(iter(kwargs.values())).device)
|
||||
return device
|
||||
|
||||
def get_device_from_module(module):
|
||||
'''Returns the first device found in the `module`'s parameters or None'''
|
||||
device = None
|
||||
|
|
@ -81,15 +54,6 @@ def get_device_str(device):
|
|||
raise RuntimeError('Unsupported device type')
|
||||
return device
|
||||
|
||||
def get_default_device_str(type):
|
||||
if isinstance(type, str):
|
||||
if type == 'cuda':
|
||||
return 'cuda:' + str(torch.cuda.current_device())
|
||||
else:
|
||||
return 'cpu'
|
||||
else:
|
||||
raise RuntimeError('Unsupported device type')
|
||||
|
||||
def get_all_gradients_finite_name_from_session(session):
|
||||
'''Find all_gradients_finite node on Session graph and return its name'''
|
||||
|
||||
|
|
|
|||
|
|
@ -9,8 +9,9 @@ from inspect import signature
|
|||
from torch.utils.dlpack import from_dlpack
|
||||
from torch.utils.cpp_extension import load_inline
|
||||
|
||||
# Needed to re-implement PyTorch's cpu,cuda,to methods
|
||||
from typing import Union, Tuple, Any, Callable, Iterator, Set, Optional, overload, TypeVar, Mapping, Dict
|
||||
# Needed to override PyTorch methods
|
||||
from typing import TypeVar
|
||||
T = TypeVar('T', bound='Module')
|
||||
|
||||
from onnxruntime.capi import _pybind_state as C
|
||||
from onnxruntime.training import register_custom_ops_pytorch_exporter
|
||||
|
|
@ -18,11 +19,6 @@ from . import _utils, _ortmodule_output_transformation
|
|||
|
||||
|
||||
ONNX_OPSET_VERSION = 12
|
||||
__TEMP_ENABLE_METHOD_TIMING__ = False
|
||||
|
||||
# Needed to re-implement PyTorch's cpu,cuda,to methods
|
||||
T = TypeVar('T', bound='Module')
|
||||
|
||||
|
||||
def _create_iobinding(io_binding, inputs, model, device):
|
||||
'''Creates IO binding for a `model` inputs and output'''
|
||||
|
|
@ -346,7 +342,6 @@ class ORTModule(torch.nn.Module):
|
|||
self._is_training = mode
|
||||
self._flattened_output_module.train(mode)
|
||||
|
||||
@_utils.timeit(enabled=__TEMP_ENABLE_METHOD_TIMING__)
|
||||
def _convert_training_graph_input_to_list(self, *inputs, **kwargs):
|
||||
'''Creates forward `*inputs` list from user input and PyTorch initializers
|
||||
|
||||
|
|
|
|||
|
|
@ -1,19 +0,0 @@
|
|||
#!/bin/bash
|
||||
|
||||
cur_dir=$(basename `pwd`)
|
||||
|
||||
if [[ ${cur_dir} != "RelWithDebInfo" ]]
|
||||
then
|
||||
echo "Going to build folder (aka build/Linux/RelWithDebInfo)"
|
||||
cd build/Linux/RelWithDebInfo
|
||||
fi
|
||||
|
||||
echo "Exporting PYTHONPATH to use build dir as onnxruntime package"
|
||||
export PYTHONPATH=$(pwd)
|
||||
|
||||
echo "Copying PyTorch frontend source-code to build folder"
|
||||
cp -Rf ../../../orttraining/orttraining/python/training/* ../../../build/Linux/RelWithDebInfo/onnxruntime/training/
|
||||
|
||||
echo "Running Flexible API (ORTModule)"
|
||||
python ../../../orttraining/orttraining/test/python/orttraining_test_ortmodule_bert_classifier.py --help
|
||||
python ../../../orttraining/orttraining/test/python/orttraining_test_ortmodule_bert_classifier.py $@
|
||||
|
|
@ -1,19 +0,0 @@
|
|||
#!/bin/bash
|
||||
|
||||
cur_dir=$(basename `pwd`)
|
||||
|
||||
if [[ ${cur_dir} != "RelWithDebInfo" ]]
|
||||
then
|
||||
echo "Going to build folder (aka build/Linux/RelWithDebInfo)"
|
||||
cd build/Linux/RelWithDebInfo
|
||||
fi
|
||||
|
||||
echo "Exporting PYTHONPATH to use build dir as onnxruntime package"
|
||||
export PYTHONPATH=$(pwd)
|
||||
|
||||
echo "Copying PyTorch frontend source-code to build folder"
|
||||
cp -Rf ../../../orttraining/orttraining/python/training/* ../../../build/Linux/RelWithDebInfo/onnxruntime/training/
|
||||
|
||||
echo "Running Flexible API (ORTModule)"
|
||||
python ../../../orttraining/orttraining/test/python/orttraining_test_ortmodule_torch_lightning_basic.py --help
|
||||
python ../../../orttraining/orttraining/test/python/orttraining_test_ortmodule_torch_lightning_basic.py $@
|
||||
|
|
@ -1,19 +0,0 @@
|
|||
#!/bin/bash
|
||||
|
||||
cur_dir=$(basename `pwd`)
|
||||
|
||||
if [[ ${cur_dir} != "RelWithDebInfo" ]]
|
||||
then
|
||||
echo "Going to build folder (aka build/Linux/RelWithDebInfo)"
|
||||
cd build/Linux/RelWithDebInfo
|
||||
fi
|
||||
|
||||
echo "Exporting PYTHONPATH to use build dir as onnxruntime package"
|
||||
export PYTHONPATH=$(pwd)
|
||||
|
||||
echo "Copying PyTorch frontend source-code to build folder"
|
||||
cp -Rf ../../../orttraining/orttraining/python/training/* ../../../build/Linux/RelWithDebInfo/onnxruntime/training/
|
||||
|
||||
echo "Running Flexible API (ORTModule)"
|
||||
python ../../../orttraining/orttraining/test/python/orttraining_test_ortmodule_poc.py --help
|
||||
python ../../../orttraining/orttraining/test/python/orttraining_test_ortmodule_poc.py $@
|
||||
|
|
@ -1,19 +0,0 @@
|
|||
#!/bin/bash
|
||||
|
||||
cur_dir=$(basename `pwd`)
|
||||
|
||||
if [[ ${cur_dir} != "RelWithDebInfo" ]]
|
||||
then
|
||||
echo "Going to build folder (aka build/Linux/RelWithDebInfo)"
|
||||
cd build/Linux/RelWithDebInfo
|
||||
fi
|
||||
|
||||
echo "Exporting PYTHONPATH to use build dir as onnxruntime package"
|
||||
export PYTHONPATH=$(pwd)
|
||||
|
||||
echo "Copying PyTorch frontend source-code to build folder"
|
||||
cp -Rf ../../../orttraining/orttraining/python/training/* ../../../build/Linux/RelWithDebInfo/onnxruntime/training/
|
||||
|
||||
echo "Running Flexible API (ORTModule)"
|
||||
python ../../../orttraining/orttraining/test/python/orttraining_test_ortmodule_deepspeed_zero_stage_1.py --help
|
||||
deepspeed ../../../orttraining/orttraining/test/python/orttraining_test_ortmodule_deepspeed_zero_stage_1.py --deepspeed_config ../../../orttraining/orttraining/test/python/orttraining_test_ortmodule_deepspeed_zero_stage_1_config.json $@
|
||||
|
|
@ -1,120 +0,0 @@
|
|||
import onnx
|
||||
import copy
|
||||
from onnx import shape_inference
|
||||
from onnxruntime.capi import _pybind_state as C
|
||||
|
||||
|
||||
def print_list(name, value):
|
||||
print(name + ':', ', '.join(value))
|
||||
|
||||
|
||||
def dim_str(dim):
|
||||
if dim.HasField('dim_value'):
|
||||
return str(dim.dim_value)
|
||||
elif dim.HasField('dim_param'):
|
||||
return dim.dim_param
|
||||
return 'n/a'
|
||||
|
||||
def print_type(name, type):
|
||||
print('[' + name + ']', 'type:', type.tensor_type.elem_type, '| size:', '[' + ','.join([dim_str(d) for d in type.tensor_type.shape.dim]) + ']')
|
||||
|
||||
|
||||
"""
|
||||
# MNIST
|
||||
original_model = onnx.load('mnist_original.onnx')
|
||||
config = C.ModuleGradientGraphBuilderConfiguration()
|
||||
weight_names_to_train = set()
|
||||
for initializer in original_model.graph.initializer:
|
||||
weight_names_to_train.add(initializer.name)
|
||||
config.weight_names_to_train = weight_names_to_train
|
||||
output_names = set()
|
||||
for output in original_model.graph.output:
|
||||
output_names.add(output.name)
|
||||
config.output_names = output_names
|
||||
|
||||
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('BertForSequenceClassification_full_training.onnx')
|
||||
config = C.ModuleGradientGraphBuilderConfiguration()
|
||||
weight_names_to_train = set()
|
||||
for initializer in original_model.graph.initializer:
|
||||
weight_names_to_train.add(initializer.name)
|
||||
config.weight_names_to_train = weight_names_to_train
|
||||
output_names = set()
|
||||
for output in original_model.graph.output:
|
||||
output_names.add(output.name)
|
||||
config.output_names = output_names
|
||||
|
||||
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')
|
||||
"""
|
||||
|
||||
#BERT with loss
|
||||
original_model = onnx.load('bert-tiny-loss.onnx')
|
||||
config = C.ModuleGradientGraphBuilderConfiguration()
|
||||
initializer_names_to_train = []
|
||||
for initializer in original_model.graph.initializer:
|
||||
if initializer.name.startswith('bert.') or initializer.name.startswith('cls.'):
|
||||
initializer_names_to_train.append(initializer.name)
|
||||
config.initializer_names_to_train = initializer_names_to_train
|
||||
input_names_require_grad = []
|
||||
input_names_require_grad.append('input3')
|
||||
config.input_names_require_grad = input_names_require_grad
|
||||
|
||||
module_gradient_graph_builder = C.ModuleGradientGraphBuilder()
|
||||
module_gradient_graph_builder.build_and_split(original_model.SerializeToString(), config)
|
||||
|
||||
forward_model = onnx.load_model_from_string(module_gradient_graph_builder.get_forward_model())
|
||||
backward_model = onnx.load_model_from_string(module_gradient_graph_builder.get_backward_model())
|
||||
onnx.save(onnx.load_model_from_string(module_gradient_graph_builder.get_gradient_model()), 'bert_gradient_graph.onnx')
|
||||
onnx.save(forward_model, 'bert_forward.onnx')
|
||||
onnx.save(backward_model, 'bert_backward.onnx')
|
||||
|
||||
split_graphs_info = module_gradient_graph_builder.get_split_graphs_info()
|
||||
print_list('user_input_names', split_graphs_info.user_input_names)
|
||||
print_list('initializer_names_to_train', split_graphs_info.initializer_names_to_train)
|
||||
print_list('user_output_names', split_graphs_info.user_output_names)
|
||||
print_list('backward_user_input_names', split_graphs_info.backward_user_input_names)
|
||||
print_list('backward_intializer_names_as_input', split_graphs_info.backward_intializer_names_as_input)
|
||||
print_list('intermediate_tensor_names', split_graphs_info.intermediate_tensor_names)
|
||||
print_list('user_output_grad_names', split_graphs_info.user_output_grad_names)
|
||||
print_list('backward_output_grad_names', split_graphs_info.backward_output_grad_names)
|
||||
|
||||
type_map = {}
|
||||
for name in split_graphs_info.user_input_names:
|
||||
type_map[name] = None
|
||||
for name in split_graphs_info.initializer_names_to_train:
|
||||
type_map[name] = None
|
||||
for name in split_graphs_info.user_output_names:
|
||||
type_map[name] = None
|
||||
for name in split_graphs_info.backward_user_input_names:
|
||||
type_map[name] = None
|
||||
for name in split_graphs_info.backward_intializer_names_as_input:
|
||||
type_map[name] = None
|
||||
for name in split_graphs_info.intermediate_tensor_names:
|
||||
type_map[name] = None
|
||||
for name in split_graphs_info.user_output_grad_names:
|
||||
type_map[name] = None
|
||||
for name in split_graphs_info.backward_output_grad_names:
|
||||
type_map[name] = None
|
||||
|
||||
for input in forward_model.graph.input:
|
||||
if input.name in type_map and type_map[input.name] is None:
|
||||
type_map[input.name] = input.type
|
||||
|
||||
for output in forward_model.graph.output:
|
||||
if output.name in type_map and type_map[output.name] is None:
|
||||
type_map[output.name] = output.type
|
||||
output_grad_name = output.name + '_grad'
|
||||
if output_grad_name in type_map and type_map[output_grad_name] is None:
|
||||
type_map[output_grad_name] = output.type
|
||||
|
||||
for key, value in type_map.items():
|
||||
print_type(key, value)
|
||||
|
|
@ -1,14 +1,6 @@
|
|||
# This code is from https://github.com/pytorch/examples/blob/master/mnist/main.py
|
||||
# with modification to do training using onnxruntime as backend on cuda device.
|
||||
|
||||
# To print nodes from ORT backend
|
||||
# Add --cmake_extra_defines onnxruntime_DEBUG_NODE_INPUTS_OUTPUTS=1 to build.sh
|
||||
# export ORT_DEBUG_NODE_IO_NAME_FILTER="SoftmaxCrossEntropyLoss_3_Grad/SoftmaxCrossEntropyLossGrad_0"
|
||||
# export ORT_DEBUG_NODE_IO_NAME_FILTER="SoftmaxCrossEntropyLoss_3"
|
||||
# export ORT_DEBUG_NODE_IO_DUMP_INPUT_DATA=1
|
||||
# export ORT_DEBUG_NODE_IO_DUMP_OUTPUT_DATA=1
|
||||
# See https://github.com/microsoft/onnxruntime/blob/master/onnxruntime/core/framework/debug_node_inputs_outputs_utils.h
|
||||
|
||||
import argparse
|
||||
import os
|
||||
import torch
|
||||
|
|
|
|||
Loading…
Reference in a new issue