From 15c67ddbf01f4398f28962d4334b00b773a77e7d Mon Sep 17 00:00:00 2001 From: ashbhandare Date: Thu, 1 Apr 2021 16:05:17 -0700 Subject: [PATCH] Make output 1 of ConcatTraining Optional and place on CPU (#7199) * Optional input 1 on CPU ConcatTraining * Rename output_1 --- .../orttraining/core/graph/training_op_defs.cc | 11 +++++++---- .../training_ops/cpu/tensor/concat_op_test.cc | 17 +++++++++++++++++ .../training_ops/cpu/tensor/concat.cc | 15 +++++++-------- .../training_ops/cuda/tensor/concat.cc | 8 ++++++-- 4 files changed, 37 insertions(+), 14 deletions(-) diff --git a/orttraining/orttraining/core/graph/training_op_defs.cc b/orttraining/orttraining/core/graph/training_op_defs.cc index c2bcce3eef..b5dcca390c 100644 --- a/orttraining/orttraining/core/graph/training_op_defs.cc +++ b/orttraining/orttraining/core/graph/training_op_defs.cc @@ -1332,7 +1332,8 @@ Example 4: .Output(1, "per_input_length", "Vector of length of each concatenated " "input along the 'axis' dimension", - "Tint") + "Tint", + OpSchema::Optional) .TypeConstraint( "T", OpSchema::all_tensor_types(), @@ -1368,9 +1369,11 @@ Example 4: output_shape->add_dim(); } - ONNX_NAMESPACE::TensorShapeProto per_input_len_shape; - per_input_len_shape.add_dim()->set_dim_value(numInputs); - updateOutputShape(ctx, 1, per_input_len_shape); + if (ctx.getNumOutputs() > 1) { + ONNX_NAMESPACE::TensorShapeProto per_input_len_shape; + per_input_len_shape.add_dim()->set_dim_value(numInputs); + updateOutputShape(ctx, 1, per_input_len_shape); + } for (size_t i = 0; i < numInputs; i++) { const auto& shape = ctx.getInputType(i)->tensor_type().shape(); diff --git a/orttraining/orttraining/test/training_ops/cpu/tensor/concat_op_test.cc b/orttraining/orttraining/test/training_ops/cpu/tensor/concat_op_test.cc index 26bf993014..44a27548ae 100644 --- a/orttraining/orttraining/test/training_ops/cpu/tensor/concat_op_test.cc +++ b/orttraining/orttraining/test/training_ops/cpu/tensor/concat_op_test.cc @@ -67,5 +67,22 @@ TEST(ConcatTrainingOpTest, Concat3D_same_len) { test.Run(); } +TEST(ConcatTrainingOpTest, Concat2D_optional_output1) { + OpTester test("ConcatTraining", 1, kMSDomain); + test.AddAttribute("axis", int64_t{1}); + + std::vector dims{4, 1}; + test.AddInput("input1", dims, {11.0f, 21.0f, 31.0f, 41.0f}); + test.AddInput("input2", {4, 2}, {12.0f, 13.0f, 22.0f, 23.0f, 32.0f, 33.0f, 42.0f, 43.0f}); + test.AddInput("input3", dims, {14.0f, 24.0f, 34.0f, 44.0f}); + test.AddOutput("concat_result", {4, 4}, + {11.0f, 12.0f, 13.0f, 14.0f, + 21.0f, 22.0f, 23.0f, 24.0f, + 31.0f, 32.0f, 33.0f, 34.0f, + 41.0f, 42.0f, 43.0f, 44.0f}); + test.AddMissingOptionalOutput(); + test.Run(); +} + } // namespace test } // namespace onnxruntime diff --git a/orttraining/orttraining/training_ops/cpu/tensor/concat.cc b/orttraining/orttraining/training_ops/cpu/tensor/concat.cc index d3bf9502ba..80b296ae85 100644 --- a/orttraining/orttraining/training_ops/cpu/tensor/concat.cc +++ b/orttraining/orttraining/training_ops/cpu/tensor/concat.cc @@ -39,15 +39,14 @@ Status ConcatTraining::Compute(OpKernelContext* ctx) const { if (p.output_num_elements == 0) return Status::OK(); - // Create output tensor for 'per_input_length' - std::vector per_input_length(input_count); - for (int i = 0; i < input_count; ++i) { - per_input_length[i] = input_tensors[i]->Shape()[p.axis]; + // Create optional output tensor for 'per_input_length' + Tensor* per_input_length_tensor = ctx->Output(1, {input_count}); + if (per_input_length_tensor) { + int64_t* per_input_length = per_input_length_tensor->template MutableData(); + for (int i = 0; i < input_count; ++i) { + per_input_length[i] = input_tensors[i]->Shape()[p.axis]; + } } - Tensor* output_1_tensor = ctx->Output(1, {input_count}); - int64_t* output_1_tensor_data = output_1_tensor->template MutableData(); - std::copy(per_input_length.begin(), per_input_length.end(), output_1_tensor_data); - // Compute values to be placed in the output tensor return ComputeImpl(p); } diff --git a/orttraining/orttraining/training_ops/cuda/tensor/concat.cc b/orttraining/orttraining/training_ops/cuda/tensor/concat.cc index 0404c1fa4e..736b8d9c52 100644 --- a/orttraining/orttraining/training_ops/cuda/tensor/concat.cc +++ b/orttraining/orttraining/training_ops/cuda/tensor/concat.cc @@ -11,6 +11,7 @@ ONNX_OPERATOR_KERNEL_EX(ConcatTraining, 1, kCudaExecutionProvider, KernelDefBuilder() + .OutputMemoryType(1) .TypeConstraint("T", DataTypeImpl::AllFixedSizeTensorTypes()), ConcatTraining); @@ -71,8 +72,11 @@ Status ConcatTraining::ComputeInternal(OpKernelContext* ctx) const { input_ptr.GpuPtr(), p.output_num_elements)); - Tensor* output_1_tensor = ctx->Output(1, {input_count}); - CUDA_RETURN_IF_ERROR(cudaMemcpyAsync(output_1_tensor->template MutableData(), concat_sizes_gpu.GpuPtr(), input_count * sizeof(int64_t), cudaMemcpyDeviceToDevice, Stream())); + // Create optional output tensor for 'per_input_length' + Tensor* per_input_length_tensor = ctx->Output(1, {input_count}); + if (per_input_length_tensor) { + std::copy(concat_sizes.begin(), concat_sizes.end(), per_input_length_tensor->template MutableData()); + } return Status::OK(); }