diff --git a/orttraining/orttraining/python/training/torchdynamo/register_backend.py b/orttraining/orttraining/python/training/torchdynamo/register_backend.py index db6f895573..a9450e119c 100644 --- a/orttraining/orttraining/python/training/torchdynamo/register_backend.py +++ b/orttraining/orttraining/python/training/torchdynamo/register_backend.py @@ -8,16 +8,37 @@ from torch._dynamo.optimizations.training import aot_autograd from .ort_backend import OrtBackend -# This is the underlying compiler for ALL training graphs if the user uses ORT -# to optimize training graphs. -# A global compiler is used here, so that cached compilation results can be reused -# across different graphs (each graph calls this compiler once). +# This should be the underlying compiler for ALL graphs if +# the user uses ORT to accelerate PyTorch via Dynamo. +# By using a global compiler for all graphs, cached compilation +# results can be reused when encountering the identical graphs. DEFAULT_BACKEND = OrtBackend() -# Wrap ORT as a compiler in Dynamo. +# Wrap ORT as a compiler in Dynamo for training (i.e., when .backward is called). # -# Under the hood, AotOrt.compile is called inside functorch. See aot_function +# Under the hood, OrtBackend.compile is called inside functorch. See aot_function # and aot_module in aot_autograd.py in PyTorch repo for more details. Basically, -# AotOrt.compile is mapped to forward graph compiler, fw_compile, and backward +# OrtBackend.compile is mapped to forward graph compiler, fw_compile, and backward # graph compiler, bw_compile, in aot_autograd.py. +# +# Example usage: +# import torch +# from onnxruntime.training.torchdynamo.register_backend import aot_ort +# model = torch.nn.Linear(2, 2) +# compiled_model = torch._dynamo.optimize(aot_ort)(model) +# result = compiled_model(torch.rand(2, 2, dtype=torch.float) +# result.sum().backward() aot_ort = aot_autograd(fw_compiler=DEFAULT_BACKEND, partition_fn=min_cut_rematerialization_partition) + +# Declare ORT as a compiler in Dynamo for inference (i.e., when .backward is NOT called). +# +# ort is usually faster than aot_ort for inference because the graphs generated by aot_autograd +# mechanism are very different than the original graphs. Therefore, some ORT's graph transformers +# are not applicable. +# +# Example usage: +# import torch +# from onnxruntime.training.torchdynamo.register_backend import ort +# model = torch.nn.Linear(2, 2) +# compiled_model = torch._dynamo.optimize(ort)(model) +ort = DEFAULT_BACKEND diff --git a/orttraining/orttraining/test/python/orttraining_test_dort.py b/orttraining/orttraining/test/python/orttraining_test_dort.py index 8ec8261087..3c2166a735 100644 --- a/orttraining/orttraining/test/python/orttraining_test_dort.py +++ b/orttraining/orttraining/test/python/orttraining_test_dort.py @@ -7,7 +7,7 @@ import torch from torch import nn from torch.nn import functional as F -from onnxruntime.training.torchdynamo.register_backend import aot_ort +from onnxruntime.training.torchdynamo.register_backend import ort, aot_ort class TestTorchDynamoOrt(unittest.TestCase): @@ -18,10 +18,11 @@ class TestTorchDynamoOrt(unittest.TestCase): torch.manual_seed(42) def test_elementwise_model(self): + torch._dynamo.reset() """Test DORT with a pure function.""" def run_elementwise_model(): - # A function to test. + # A function to test DORT. def elementwise_model(tensor_x: torch.Tensor): tensor_w = tensor_x.relu() tensor_y = tensor_w * tensor_w + 1.5 @@ -54,7 +55,37 @@ class TestTorchDynamoOrt(unittest.TestCase): for _ in range(5): run_elementwise_model() + def test_elementwise_model_for_inference(self): + torch._dynamo.reset() + + # A function to test DORT for inference (i.e., the compiled function + # doesn't have backward pass). + def elementwise_model(tensor_x: torch.Tensor): + tensor_w = tensor_x.relu() + tensor_y = tensor_w * tensor_w + 1.5 + tensor_z = tensor_y + tensor_x + tensor_p = tensor_z * tensor_x + tensor_q = tensor_p.relu() + return tensor_q + + @torch._dynamo.optimize(ort) + def optimized_elementwise_model(tensor_x: torch.Tensor): + return elementwise_model(tensor_x) + + def run(fun, list_x): + tensor_x = torch.tensor(list_x, dtype=torch.float32).requires_grad_() + tensor_y = fun(tensor_x) + return tensor_y + + # Baseline. + tensor_y = run(elementwise_model, [-1.0, 2.0]) + # ORT result. + tensor_y_new = run(optimized_elementwise_model, [-1.0, 2.0]) + + torch.testing.assert_close(tensor_y, tensor_y_new) + def test_to_copy(self): + torch._dynamo.reset() """Test DORT with aten::_to_copy.""" def run_to_copy(): @@ -88,6 +119,7 @@ class TestTorchDynamoOrt(unittest.TestCase): run_to_copy() def test_mnist_model(self): + torch._dynamo.reset() """Test DORT with a simple nn.Module.""" def run_mnist_model():