From db6a821869abdd4149f5f0635e7d183d913d143c Mon Sep 17 00:00:00 2001 From: Bowen Bao Date: Mon, 24 Aug 2020 14:31:08 -0700 Subject: [PATCH] Enable example transformer test with dynamic size inputs (#4888) Co-authored-by: Thiago Crepaldi --- cmake/onnxruntime_python.cmake | 6 ++++++ .../python/experimental/orttrainer.py | 15 +++++++++++++-- orttraining/orttraining/python/ort_trainer.py | 17 +++++++++++++++-- .../orttraining_test_orttrainer_frontend.py | 4 ++++ samples/python/pytorch_transformer/pt_model.py | 4 ++-- 5 files changed, 40 insertions(+), 6 deletions(-) diff --git a/cmake/onnxruntime_python.cmake b/cmake/onnxruntime_python.cmake index dffc396938..8c9726cb8b 100644 --- a/cmake/onnxruntime_python.cmake +++ b/cmake/onnxruntime_python.cmake @@ -182,6 +182,9 @@ if (onnxruntime_ENABLE_TRAINING) file(GLOB onnxruntime_python_optim_srcs CONFIGURE_DEPENDS "${ORTTRAINING_SOURCE_DIR}/python/experimental/optim/*.py" ) + file(GLOB onnxruntime_python_train_tools_srcs CONFIGURE_DEPENDS + "${REPO_ROOT}/tools/python/register_custom_ops_pytorch_exporter.py" + ) else() file(GLOB onnxruntime_python_capi_training_srcs CONFIGURE_DEPENDS "${ONNXRUNTIME_ROOT}/python/training/*.py" @@ -284,6 +287,9 @@ if (onnxruntime_ENABLE_TRAINING) COMMAND ${CMAKE_COMMAND} -E copy ${onnxruntime_python_optim_srcs} $/onnxruntime/experimental/optim/ + COMMAND ${CMAKE_COMMAND} -E copy + ${onnxruntime_python_train_tools_srcs} + $/onnxruntime/experimental/ ) endif() diff --git a/orttraining/orttraining/python/experimental/orttrainer.py b/orttraining/orttraining/python/experimental/orttrainer.py index fb2c5b1d41..a64eb371a3 100644 --- a/orttraining/orttraining/python/experimental/orttrainer.py +++ b/orttraining/orttraining/python/experimental/orttrainer.py @@ -453,7 +453,11 @@ class ORTTrainer(object): # Do an inference to grab output types model.eval() with torch.no_grad(): - sample_outputs = model(*sample_inputs) + # Deepcopy inputs, since input values may change after model run. + sample_inputs_copy = copy.deepcopy(sample_inputs) + # Deepcopy model, in case model is stateful and changes after model run. + model_copy = copy.deepcopy(model) + sample_outputs = model_copy(*sample_inputs_copy) model.train() if isinstance(sample_outputs, torch.Tensor): sample_outputs = [sample_outputs] @@ -470,7 +474,14 @@ class ORTTrainer(object): # Export the model to ONNX f = io.BytesIO() - torch.onnx._export(model, tuple(sample_inputs), f, + # Deepcopy inputs, since input values may change after model run. + sample_inputs_copy = copy.deepcopy(sample_inputs) + + # Enable contrib ops export from PyTorch + from onnxruntime.experimental import register_custom_ops_pytorch_exporter + register_custom_ops_pytorch_exporter.register_custom_op() + + torch.onnx._export(model, tuple(sample_inputs_copy), f, input_names=[input.name for input in self.model_desc.inputs], output_names=[output.name for output in self.model_desc.outputs], opset_version=self.options._internal_use.onnx_opset_version, diff --git a/orttraining/orttraining/python/ort_trainer.py b/orttraining/orttraining/python/ort_trainer.py index 99d9d81bf2..729d1762ab 100644 --- a/orttraining/orttraining/python/ort_trainer.py +++ b/orttraining/orttraining/python/ort_trainer.py @@ -316,7 +316,12 @@ def convert_model_loss_fn_to_onnx(model, loss_fn, model_desc, device, inputs, op model.eval() with torch.no_grad(): - sample_outputs = model(*sample_inputs) + import copy + # Deepcopy inputs, since input values may change after model run. + sample_inputs_copy = copy.deepcopy(sample_inputs) + # Deepcopy model, in case model is stateful and changes after model run. + model_copy = copy.deepcopy(model) + sample_outputs = model_copy(*sample_inputs_copy) if isinstance(sample_outputs, torch.Tensor): sample_outputs = [sample_outputs] for sample_output, output_desc in zip(sample_outputs, model_desc.outputs_): @@ -336,7 +341,15 @@ def convert_model_loss_fn_to_onnx(model, loss_fn, model_desc, device, inputs, op if LooseVersion(torch.__version__) >= LooseVersion('1.6.0'): other_export_options['training'] = torch.onnx.TrainingMode.TRAINING - torch.onnx._export(model, tuple(sample_inputs), f, + # Deepcopy inputs, since input values may change after model run. + import copy + sample_inputs_copy = copy.deepcopy(sample_inputs) + + # Enable contrib ops export from PyTorch + from onnxruntime.experimental import register_custom_ops_pytorch_exporter + register_custom_ops_pytorch_exporter.register_custom_op() + + torch.onnx._export(model, tuple(sample_inputs_copy), f, input_names=input_names, output_names=output_names, opset_version=opset_version, diff --git a/orttraining/orttraining/test/python/orttraining_test_orttrainer_frontend.py b/orttraining/orttraining/test/python/orttraining_test_orttrainer_frontend.py index faa781b962..26424be89d 100644 --- a/orttraining/orttraining/test/python/orttraining_test_orttrainer_frontend.py +++ b/orttraining/orttraining/test/python/orttraining_test_orttrainer_frontend.py @@ -774,6 +774,10 @@ def testORTTrainerDynamicShape(dynamic_axes): total_steps = 10 for i in range(total_steps): data, targets = batcher_fn(train_data, i) + if dynamic_axes: + # Forcing batches with different sizes to exercise dynamic shapes + data = data[:-(i+1)] + targets = targets[:-(i+1)*data.size(1)] _, _ = trainer.train_step(data, targets) assert trainer._onnx_model is not None diff --git a/samples/python/pytorch_transformer/pt_model.py b/samples/python/pytorch_transformer/pt_model.py index 98898ab49e..63a6c3fbd4 100644 --- a/samples/python/pytorch_transformer/pt_model.py +++ b/samples/python/pytorch_transformer/pt_model.py @@ -32,9 +32,9 @@ class TransformerModel(nn.Module): self.decoder.weight.data.uniform_(-initrange, initrange) def forward(self, input1): - if self.input1_mask is None or self.input1_mask.size(0) != len(input1): + if self.input1_mask is None or self.input1_mask.size(0) != input1.size(0): device = input1.device - mask = self._generate_square_subsequent_mask(len(input1)).to(device) + mask = self._generate_square_subsequent_mask(input1.size(0)).to(device) self.input1_mask = mask input1 = self.encoder(input1) * math.sqrt(self.ninp)