From 3d009cdde3f60144e877f69408d325e39e931014 Mon Sep 17 00:00:00 2001 From: Wil Brady <25513670+WilBrady@users.noreply.github.com> Date: Thu, 11 Aug 2022 17:00:12 -0400 Subject: [PATCH] Updating binary ops in eager mode to support broadcasting. (#12560) * Updating binary ops in eager mode to support broadcasting. --- .../orttraining/eager/opgen/opgen/atenops.py | 21 +- orttraining/orttraining/eager/ort_aten.cpp | 349 ++++++++++++++++-- orttraining/orttraining/eager/test/ort_ops.py | 19 +- 3 files changed, 348 insertions(+), 41 deletions(-) diff --git a/orttraining/orttraining/eager/opgen/opgen/atenops.py b/orttraining/orttraining/eager/opgen/opgen/atenops.py index 1e77d159b5..5529ad0e0e 100644 --- a/orttraining/orttraining/eager/opgen/opgen/atenops.py +++ b/orttraining/orttraining/eager/opgen/opgen/atenops.py @@ -102,21 +102,6 @@ for unary_op in unary_ops_with_inplace: for unary_op in unary_ops: ops[f"aten::{unary_op}"] = onnx_ops[unary_op]("self") -for binary_op, onnx_op in { - "add": Add("self", Mul("alpha", "other")), - "sub": Sub("self", Mul("alpha", "other")), - "mul": Mul("self", "other"), - "div": Div("self", "other"), -}.items(): - # for Tensor, binary_op.out is used by both binary_op and binary_op_, so we only generate .out - # from testing and call stacks, it also apears scalar ops fall back to the (Tensor) binary_op.out, - # so this is all we need. - name = f"aten::{binary_op}.out" - if name in ops: - raise RuntimeError("Duplicate binary op found in op dictionary.") - ops[f"aten::{binary_op}.out"] = deepcopy(onnx_op) - type_promotion_ops.append(f"aten::{binary_op}.out") - # Notes on Onnx op mapping # # Equal - Onnx spec has the return as a bool tensor, but aten will keep the tensor @@ -182,6 +167,12 @@ hand_implemented = { "aten::squeeze.dim": Squeeze("self", "dim"), "aten::squeeze": SignatureOnly(), "aten::unsqueeze": Unsqueeze(data="self", axes="dim"), + # until the generator is modified to include resizing out based on broadcast result, we hand write + # add, sub, mul, and div. The majority of the code was generated by using the commented onnx ops. + "aten::add.out": SignatureOnly(), # Add("self", Mul("alpha", "other")), + "aten::sub.out": SignatureOnly(), # Sub("self", Mul("alpha", "other")), + "aten::mul.out": SignatureOnly(), # Mul("self", "other"), + "aten::div.out": SignatureOnly(), # Div("self", "other"), } # If the aten op expects a specific output type that differs from self diff --git a/orttraining/orttraining/eager/ort_aten.cpp b/orttraining/orttraining/eager/ort_aten.cpp index 5afae07f4d..c25a574f57 100644 --- a/orttraining/orttraining/eager/ort_aten.cpp +++ b/orttraining/orttraining/eager/ort_aten.cpp @@ -558,6 +558,54 @@ void resize_impl_ort_( return; } +// The following was copied from onnxruntime/core/providers/cuda/math/binary_elementwise_ops.cc +// This method computes the resulting shape from broadcasting. +Status ComputeOutputShape( + const std::string& node_name, + const onnxruntime::TensorShape& lhs_shape, + const onnxruntime::TensorShape& rhs_shape, + onnxruntime::TensorShape& out_shape) { + size_t lhs_rank = lhs_shape.NumDimensions(); + size_t rhs_rank = rhs_shape.NumDimensions(); + size_t out_rank = std::max(lhs_rank, rhs_rank); + + std::vector output_dims(out_rank, 0); + for (size_t i = 0; i < out_rank; ++i) { + int64_t lhs_dim = 1; + if (i < lhs_rank) + lhs_dim = lhs_shape[lhs_rank - 1 - i]; + int64_t rhs_dim = 1; + if (i < rhs_rank) + rhs_dim = rhs_shape[rhs_rank - 1 - i]; + int64_t max = std::max(lhs_dim, rhs_dim); + int64_t min = std::min(lhs_dim, rhs_dim); + int64_t out_dim = (min == 0 ? min : max); // special case a dim value of 0. + if (lhs_dim != out_dim && lhs_dim != 1) + return ORT_MAKE_STATUS(ONNXRUNTIME, FAIL, node_name, ": left operand cannot broadcast on dim ", lhs_rank - 1 - i, + " LeftShape: ", lhs_shape.ToString(), ", RightShape: ", rhs_shape.ToString()); + if (rhs_dim != out_dim && rhs_dim != 1) + return ORT_MAKE_STATUS(ONNXRUNTIME, FAIL, node_name, ": right operand cannot broadcast on dim ", rhs_rank - 1 - i, + " LeftShape: ", lhs_shape.ToString(), ", RightShape: ", rhs_shape.ToString()); + output_dims[out_rank - 1 - i] = out_dim; + } + out_shape = onnxruntime::TensorShape(output_dims); + return Status::OK(); +} + +// Note in the below the user needs to pass in out_shape so they own the memory as IntArrayRef is just a view into it. +at::IntArrayRef BroadcastShape( + const std::string& node_name, + OrtValue lhs, + OrtValue rhs, + onnxruntime::TensorShape& out_shape) { + auto& ort_tensor_lhs = lhs.Get(); + auto& ort_tensor_rhs = rhs.Get(); + auto status = ComputeOutputShape(node_name, ort_tensor_lhs.Shape(), ort_tensor_rhs.Shape(), out_shape); + CHECK_STATUS(status); + auto out_shape_dims = out_shape.GetDims(); + return at::IntArrayRef(&out_shape_dims[0], out_shape_dims.size()); +} + // #pragma region Hand-Implemented ATen Ops namespace aten { @@ -665,13 +713,12 @@ at::Tensor& copy_( auto ort_self = create_ort_value(invoker, self); if (self.scalar_type() != src.scalar_type()) { - if(src.device().type() != at::kORT) - { + if (src.device().type() != at::kORT) { // invoke cast first and then copy for non-ORT device types auto val = at::native::to(src, self.scalar_type()); const auto ort_val = create_ort_value(invoker, val); copy(invoker, ort_val, ort_self); - }else{ + } else { // For ORT device type, the cast operation will perform the copy as well std::vector ort_cast_output(1); ort_cast_output[0] = ort_self; @@ -680,8 +727,8 @@ at::Tensor& copy_( "to", (int64_t)GetONNXTensorProtoDataType(self.scalar_type()), at::kLong); auto status = invoker.Invoke("Cast", - {std::move(ort_src)}, - ort_cast_output, &attrs); + {std::move(ort_src)}, + ort_cast_output, &attrs); CHECK_STATUS(status); } @@ -922,20 +969,20 @@ const at::Tensor& resize_( // aten::cat.out(Tensor[] tensors, int dim=0, *, Tensor(a!) out) -> Tensor(a!) at::Tensor& cat_out( - at::TensorList tensors, - int64_t dim, - // *, - at::Tensor& out) { + at::TensorList tensors, + int64_t dim, + // *, + at::Tensor& out) { ORT_LOG_FN(tensors, dim, out); assert(tensors.size() > 0); if ( std::vector supportedTypes = - {at::kBFloat16, at::kBool, at::kByte, at::kDouble, at::kFloat, at::kHalf, at::kInt, at::kLong, at::kShort}; + {at::kBFloat16, at::kBool, at::kByte, at::kDouble, at::kFloat, at::kHalf, at::kInt, at::kLong, at::kShort}; !IsSupportedType(tensors, supportedTypes)) { return at::native::call_fallback_fn< - &at::native::cpu_fallback, - ATEN_OP(cat_out)>::call(tensors, dim, out); + &at::native::cpu_fallback, + ATEN_OP(cat_out)>::call(tensors, dim, out); } int64_t ndim = tensors[0].dim(); assert(ndim != 0); @@ -960,14 +1007,14 @@ at::Tensor& cat_out( NodeAttributes attrs_0(1); attrs_0["axis"] = create_ort_attribute( - "axis", dim, at::ScalarType::Int); + "axis", dim, at::ScalarType::Int); std::vector ort_outputs_0_Concat(1); ort_outputs_0_Concat[0] = ort_input_out; - auto status = invoker.Invoke("Concat", { - std::move(ort_input_0_tensors), - }, ort_outputs_0_Concat, &attrs_0); + auto status = invoker.Invoke("Concat", + {std::move(ort_input_0_tensors)}, + ort_outputs_0_Concat, &attrs_0); CHECK_STATUS(status); return out; @@ -1176,16 +1223,16 @@ at::Tensor& mm_out( // aten::squeeze(Tensor(a) self) -> Tensor(a) at::Tensor squeeze( - const at::Tensor& self) { + const at::Tensor& self) { ORT_LOG_FN(self); if ( - std::vector supportedTypes = + std::vector supportedTypes = {at::kBFloat16, at::kBool, at::kByte, at::kDouble, at::kFloat, at::kHalf, at::kInt, at::kLong, at::kShort}; !IsSupportedType(self, supportedTypes)) return at::native::call_fallback_fn< - &at::native::cpu_fallback, - ATEN_OP(squeeze)>::call(self); + &at::native::cpu_fallback, + ATEN_OP(squeeze)>::call(self); auto& invoker = GetORTInvoker(self.device()); @@ -1193,15 +1240,267 @@ at::Tensor squeeze( std::vector ort_outputs_0_Squeeze(1); - auto status = invoker.Invoke("Squeeze", { - std::move(ort_input_0_self), - }, ort_outputs_0_Squeeze, nullptr); + auto status = invoker.Invoke("Squeeze", + {std::move(ort_input_0_self)}, + ort_outputs_0_Squeeze, nullptr); CHECK_STATUS(status); at::TensorOptions tensor_options = self.options(); return aten_tensor_from_ort( - std::move(ort_outputs_0_Squeeze[0]), - tensor_options); + std::move(ort_outputs_0_Squeeze[0]), + tensor_options); +} + +// aten::add.out(Tensor self, Tensor other, *, Scalar alpha=1, Tensor(a!) out) -> Tensor(a!) +at::Tensor& add_out( + const at::Tensor& self, + const at::Tensor& other, + // *, + const at::Scalar& alpha, + at::Tensor& out) { + ORT_LOG_FN(self, other, alpha, out); + + auto promoted_type = PromoteScalarTypesWithCategory({other.scalar_type(), self.scalar_type()}, {alpha.type()}); + + if ( + std::vector supportedTypes = + {at::kBFloat16, at::kByte, at::kDouble, at::kFloat, at::kHalf, at::kInt, at::kLong, at::kShort}; + !IsSupportedType(alpha, supportedTypes) || + !IsSupportedType(other, supportedTypes) || + !IsSupportedType(self, supportedTypes) || + !c10::canCast(*promoted_type, out.scalar_type())) { + return at::native::call_fallback_fn< + &at::native::cpu_fallback, + ATEN_OP(add_out)>::call(self, other, alpha, out); + } + auto& invoker = GetORTInvoker(self.device()); + + auto ort_input_0_alpha = create_ort_value(invoker, alpha); + if (alpha.type() != *promoted_type) { + ort_input_0_alpha = CastToType(invoker, ort_input_0_alpha, *promoted_type); + } + auto ort_input_0_other = create_ort_value(invoker, other); + if (other.scalar_type() != *promoted_type) { + ort_input_0_other = CastToType(invoker, ort_input_0_other, *promoted_type); + } + + std::vector ort_outputs_0_Mul(1); + + auto status = invoker.Invoke("Mul", + {std::move(ort_input_0_alpha), std::move(ort_input_0_other)}, + ort_outputs_0_Mul, nullptr); + CHECK_STATUS(status); + + auto ort_input_1_self = create_ort_value(invoker, self); + if (self.scalar_type() != *promoted_type) { + ort_input_1_self = CastToType(invoker, ort_input_1_self, *promoted_type); + } + + // resize the output and then create output ort value to be updated. + onnxruntime::TensorShape out_shape; + resize_output( + invoker, + dynamic_cast(out.unsafeGetTensorImpl()), + BroadcastShape(__func__, ort_input_1_self, ort_outputs_0_Mul[0], out_shape)); + + auto ort_input_out = create_ort_value(invoker, out); + + std::vector ort_outputs_1_Add(1); + if (*promoted_type == out.scalar_type()) { + ort_outputs_1_Add[0] = ort_input_out; + } + + status = invoker.Invoke("Add", + {std::move(ort_input_1_self), std::move(ort_outputs_0_Mul[0])}, + ort_outputs_1_Add, nullptr); + CHECK_STATUS(status); + + if (*promoted_type != out.scalar_type()) { + CastToType_out(invoker, ort_outputs_1_Add[0], ort_input_out, out.scalar_type()); + } + return out; +} + +// aten::sub.out(Tensor self, Tensor other, *, Scalar alpha=1, Tensor(a!) out) -> Tensor(a!) +at::Tensor& sub_out( + const at::Tensor& self, + const at::Tensor& other, + // *, + const at::Scalar& alpha, + at::Tensor& out) { + ORT_LOG_FN(self, other, alpha, out); + + auto promoted_type = PromoteScalarTypesWithCategory({other.scalar_type(), self.scalar_type()}, {alpha.type()}); + + if ( + std::vector supportedTypes = + {at::kBFloat16, at::kByte, at::kDouble, at::kFloat, at::kHalf, at::kInt, at::kLong, at::kShort}; + !IsSupportedType(alpha, supportedTypes) || + !IsSupportedType(other, supportedTypes) || + !IsSupportedType(self, supportedTypes) || + !c10::canCast(*promoted_type, out.scalar_type())) { + return at::native::call_fallback_fn< + &at::native::cpu_fallback, + ATEN_OP(sub_out)>::call(self, other, alpha, out); + } + auto& invoker = GetORTInvoker(self.device()); + + auto ort_input_0_alpha = create_ort_value(invoker, alpha); + if (alpha.type() != *promoted_type) { + ort_input_0_alpha = CastToType(invoker, ort_input_0_alpha, *promoted_type); + } + auto ort_input_0_other = create_ort_value(invoker, other); + if (other.scalar_type() != *promoted_type) { + ort_input_0_other = CastToType(invoker, ort_input_0_other, *promoted_type); + } + + std::vector ort_outputs_0_Mul(1); + + auto status = invoker.Invoke("Mul", + {std::move(ort_input_0_alpha), std::move(ort_input_0_other)}, + ort_outputs_0_Mul, nullptr); + CHECK_STATUS(status); + + auto ort_input_1_self = create_ort_value(invoker, self); + if (self.scalar_type() != *promoted_type) { + ort_input_1_self = CastToType(invoker, ort_input_1_self, *promoted_type); + } + + // resize the output and then create output ort value to be updated. + onnxruntime::TensorShape out_shape; + resize_output( + invoker, + dynamic_cast(out.unsafeGetTensorImpl()), + BroadcastShape(__func__, ort_input_1_self, ort_outputs_0_Mul[0], out_shape)); + + auto ort_input_out = create_ort_value(invoker, out); + + std::vector ort_outputs_1_Sub(1); + if (*promoted_type == out.scalar_type()) { + ort_outputs_1_Sub[0] = ort_input_out; + } + + status = invoker.Invoke("Sub", + {std::move(ort_input_1_self), std::move(ort_outputs_0_Mul[0])}, + ort_outputs_1_Sub, nullptr); + CHECK_STATUS(status); + + if (*promoted_type != out.scalar_type()) { + CastToType_out(invoker, ort_outputs_1_Sub[0], ort_input_out, out.scalar_type()); + } + return out; +} + +// aten::mul.out(Tensor self, Tensor other, *, Tensor(a!) out) -> Tensor(a!) +at::Tensor& mul_out( + const at::Tensor& self, + const at::Tensor& other, + // *, + at::Tensor& out) { + ORT_LOG_FN(self, other, out); + + auto promoted_type = PromoteScalarTypesWithCategory({self.scalar_type(), other.scalar_type()}, {}); + + if ( + std::vector supportedTypes = + {at::kBFloat16, at::kByte, at::kDouble, at::kFloat, at::kHalf, at::kInt, at::kLong, at::kShort}; + !IsSupportedType(other, supportedTypes) || + !IsSupportedType(self, supportedTypes) || + !c10::canCast(*promoted_type, out.scalar_type())) { + return at::native::call_fallback_fn< + &at::native::cpu_fallback, + ATEN_OP(mul_out)>::call(self, other, out); + } + auto& invoker = GetORTInvoker(self.device()); + + auto ort_input_0_self = create_ort_value(invoker, self); + if (self.scalar_type() != *promoted_type) { + ort_input_0_self = CastToType(invoker, ort_input_0_self, *promoted_type); + } + auto ort_input_0_other = create_ort_value(invoker, other); + if (other.scalar_type() != *promoted_type) { + ort_input_0_other = CastToType(invoker, ort_input_0_other, *promoted_type); + } + + // resize the output and then create output ort value to be updated. + onnxruntime::TensorShape out_shape; + resize_output( + invoker, + dynamic_cast(out.unsafeGetTensorImpl()), + BroadcastShape(__func__, ort_input_0_self, ort_input_0_other, out_shape)); + + auto ort_input_out = create_ort_value(invoker, out); + + std::vector ort_outputs_0_Mul(1); + if (*promoted_type == out.scalar_type()) { + ort_outputs_0_Mul[0] = ort_input_out; + } + + auto status = invoker.Invoke("Mul", + {std::move(ort_input_0_self), std::move(ort_input_0_other)}, + ort_outputs_0_Mul, nullptr); + CHECK_STATUS(status); + + if (*promoted_type != out.scalar_type()) { + CastToType_out(invoker, ort_outputs_0_Mul[0], ort_input_out, out.scalar_type()); + } + return out; +} + +// aten::div.out(Tensor self, Tensor other, *, Tensor(a!) out) -> Tensor(a!) +at::Tensor& div_out( + const at::Tensor& self, + const at::Tensor& other, + // *, + at::Tensor& out) { + ORT_LOG_FN(self, other, out); + + auto promoted_type = PromoteScalarTypesWithCategory({self.scalar_type(), other.scalar_type()}, {}); + + if ( + std::vector supportedTypes = + {at::kBFloat16, at::kByte, at::kDouble, at::kFloat, at::kHalf, at::kInt, at::kLong, at::kShort}; + !IsSupportedType(other, supportedTypes) || + !IsSupportedType(self, supportedTypes) || + !c10::canCast(*promoted_type, out.scalar_type())) { + return at::native::call_fallback_fn< + &at::native::cpu_fallback, + ATEN_OP(div_out)>::call(self, other, out); + } + auto& invoker = GetORTInvoker(self.device()); + + auto ort_input_0_self = create_ort_value(invoker, self); + if (self.scalar_type() != *promoted_type) { + ort_input_0_self = CastToType(invoker, ort_input_0_self, *promoted_type); + } + auto ort_input_0_other = create_ort_value(invoker, other); + if (other.scalar_type() != *promoted_type) { + ort_input_0_other = CastToType(invoker, ort_input_0_other, *promoted_type); + } + + // resize the output and then create output ort value to be updated. + onnxruntime::TensorShape out_shape; + resize_output( + invoker, + dynamic_cast(out.unsafeGetTensorImpl()), + BroadcastShape(__func__, ort_input_0_self, ort_input_0_other, out_shape)); + + auto ort_input_out = create_ort_value(invoker, out); + + std::vector ort_outputs_0_Div(1); + if (*promoted_type == out.scalar_type()) { + ort_outputs_0_Div[0] = ort_input_out; + } + + auto status = invoker.Invoke("Div", + {std::move(ort_input_0_self), std::move(ort_input_0_other)}, + ort_outputs_0_Div, nullptr); + CHECK_STATUS(status); + + if (*promoted_type != out.scalar_type()) { + CastToType_out(invoker, ort_outputs_0_Div[0], ort_input_out, out.scalar_type()); + } + return out; } } // namespace aten diff --git a/orttraining/orttraining/eager/test/ort_ops.py b/orttraining/orttraining/eager/test/ort_ops.py index 77e5da6c4c..f4ad21cfb9 100644 --- a/orttraining/orttraining/eager/test/ort_ops.py +++ b/orttraining/orttraining/eager/test/ort_ops.py @@ -489,6 +489,23 @@ class OrtOpTests(unittest.TestCase): assert torch.equal(cpu_result1, ort_result1.cpu()) assert torch.equal(cpu_result2, ort_result2.cpu()) + def test_add_broadcasting(self): + device = self.get_device() + cpu_first = torch.rand(1, 1, 3, 4, 5) + ort_first = cpu_first.to(device) + cpu_last = torch.rand(1, 2, 3, 1, 1) + ort_last = cpu_last.to(device) + cpu_single = torch.rand(5) + ort_single = cpu_single.to(device) + + cpu_result1 = cpu_first + cpu_last # dims = (1,2,3,4,5) is final + ort_result1 = ort_first + ort_last + assert torch.equal(cpu_result1, ort_result1.cpu()) + + cpu_result2 = cpu_result1 + cpu_single + ort_result2 = ort_result1 + ort_single + assert torch.equal(cpu_result2, ort_result2.cpu()) + ################################ parameterized test follow ####################################### # OPS - is a list of [test_operator, tested_tensor=torch.rand (6)]. # The default value for tested_tensor is torch.rand (6)- size of 6 uniform distribution on the interval [0, 1). @@ -670,7 +687,7 @@ class OrtOpTests(unittest.TestCase): @parameterized.expand(binary_ops, name_func=rename_func) def test_op_binary_tensor(self, binary_op, op_sign, alpha_supported): device = self.get_device() - cpu_input = torch.rand(3, 3) + cpu_input = torch.rand(3, 1) # use broadcasting in the second dim. ort_input = cpu_input.to(device) cpu_other = torch.rand(3, 3) ort_other = cpu_other.to(device)