From 004632ff8d500e86044c037a3b98f90714aa8fa6 Mon Sep 17 00:00:00 2001 From: Thiago Crepaldi Date: Wed, 2 Dec 2020 15:32:37 -0800 Subject: [PATCH] TEMP: Add support to measure method execution time for perf improvement --- .../orttraining/python/training/_utils.py | 19 +++++++++++++++++++ .../orttraining/python/training/ortmodule.py | 4 ++++ ...training_test_ortmodule_bert_classifier.py | 8 ++++++-- .../orttraining_test_ortmodule_mnist.py | 8 ++++++-- 4 files changed, 35 insertions(+), 4 deletions(-) diff --git a/orttraining/orttraining/python/training/_utils.py b/orttraining/orttraining/python/training/_utils.py index 4444e8327c..361ae87a6b 100644 --- a/orttraining/orttraining/python/training/_utils.py +++ b/orttraining/orttraining/python/training/_utils.py @@ -2,9 +2,28 @@ 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): '''Returns device index from a device''' diff --git a/orttraining/orttraining/python/training/ortmodule.py b/orttraining/orttraining/python/training/ortmodule.py index a1b2a83408..14fddf1f12 100644 --- a/orttraining/orttraining/python/training/ortmodule.py +++ b/orttraining/orttraining/python/training/ortmodule.py @@ -18,6 +18,7 @@ from . import _utils ONNX_OPSET_VERSION = 12 +__TEMP_ENABLE_METHOD_TIMING__ = True # Needed to re-implement PyTorch's cpu,cuda,to methods T = TypeVar('T', bound='Module') @@ -252,6 +253,7 @@ class ORTModule(torch.nn.Module): proc_inputs = [data for data in inputs if data is not None] return _ORTModuleFunction.apply(*self._convert_forward_input_to_list(*proc_inputs, **kwargs)) + @_utils.timeit(enabled=__TEMP_ENABLE_METHOD_TIMING__) def _convert_forward_input_to_list(self, *inputs, **kwargs): '''Creates forward `*inputs` list from user input and PyTorch initializers @@ -274,6 +276,7 @@ class ORTModule(torch.nn.Module): return result + @_utils.timeit(enabled=__TEMP_ENABLE_METHOD_TIMING__) def _convert_forward_input_list_to_dict(self, *inputs): '''Convert forward `*inputs` list to dict @@ -284,6 +287,7 @@ class ORTModule(torch.nn.Module): *self._onnx_graphs_info.initializer_names_to_train] return dict(zip(forward_input_names, inputs)) + @_utils.timeit(enabled=__TEMP_ENABLE_METHOD_TIMING__) def _convert_backward_input_list_to_dict(self, *inputs): '''Convert backward `*inputs` list to dict diff --git a/orttraining/orttraining/test/python/orttraining_test_ortmodule_bert_classifier.py b/orttraining/orttraining/test/python/orttraining_test_ortmodule_bert_classifier.py index 5d6020fcf0..c2ac0473ff 100644 --- a/orttraining/orttraining/test/python/orttraining_test_ortmodule_bert_classifier.py +++ b/orttraining/orttraining/test/python/orttraining_test_ortmodule_bert_classifier.py @@ -433,8 +433,12 @@ def main(): print('\n======== Global stats ========') if not args.pytorch_only: - estimated_export = epoch_0_training - (total_training_time - epoch_0_training)/(args.epochs-1) - print(" Estimated ONNX export took: {:.4f}s".format(estimated_export)) + estimated_export = 0 + if args.epochs > 1: + estimated_export = epoch_0_training - (total_training_time - epoch_0_training)/(args.epochs-1) + print(" Estimated ONNX export took: {:.4f}s".format(estimated_export)) + else: + print(" Estimated ONNX export took: Estimate available when epochs > 1 only") print(" Accumulated training without export took: {:.4f}s".format(total_training_time - estimated_export)) print(" Accumulated training took: {:.4f}s".format(total_training_time)) print(" Accumulated validation took: {:.4f}s".format(total_test_time)) diff --git a/orttraining/orttraining/test/python/orttraining_test_ortmodule_mnist.py b/orttraining/orttraining/test/python/orttraining_test_ortmodule_mnist.py index 8c88218dc3..0a996b08a8 100644 --- a/orttraining/orttraining/test/python/orttraining_test_ortmodule_mnist.py +++ b/orttraining/orttraining/test/python/orttraining_test_ortmodule_mnist.py @@ -193,8 +193,12 @@ def main(): print('\n======== Global stats ========') if not args.pytorch_only: - estimated_export = epoch_0_training - (total_training_time - epoch_0_training)/(args.epochs-1) - print(" Estimated ONNX export took: {:.4f}s".format(estimated_export)) + estimated_export = 0 + if args.epochs > 1: + estimated_export = epoch_0_training - (total_training_time - epoch_0_training)/(args.epochs-1) + print(" Estimated ONNX export took: {:.4f}s".format(estimated_export)) + else: + print(" Estimated ONNX export took: Estimate available when epochs > 1 only") print(" Accumulated training without export took: {:.4f}s".format(total_training_time - estimated_export)) print(" Accumulated training took: {:.4f}s".format(total_training_time)) print(" Accumulated validation took: {:.4f}s".format(total_test_time))