diff --git a/orttraining/orttraining/eager/opgen/opgen/atenops.py b/orttraining/orttraining/eager/opgen/opgen/atenops.py index 147a56911f..606b677ade 100644 --- a/orttraining/orttraining/eager/opgen/opgen/atenops.py +++ b/orttraining/orttraining/eager/opgen/opgen/atenops.py @@ -163,8 +163,10 @@ hand_implemented = { "aten::_local_scalar_dense": MakeTorchFallback(), # This function extracts a scalar value from # a tensor with exactly one value; there's no need to try to do this on an ORT device. # See CPU impl at pytorch/blob/master/aten/src/ATen/native/Scalar.cpp - "aten::gt.Scalar_out": MakeTorchFallback(), - "aten::lt.Scalar_out": MakeTorchFallback(), + "aten::lt.Scalar_out": Cast(Less(A="self", B="other"), to="GetONNXTensorProtoDataType(out.scalar_type())"), + "aten::lt.Tensor_out": Cast(Less(A="self", B="other"), to="GetONNXTensorProtoDataType(out.scalar_type())"), + "aten::gt.Scalar_out": Cast(Greater(A="self", B="other"), to="GetONNXTensorProtoDataType(out.scalar_type())"), + "aten::gt.Tensor_out": Cast(Greater(A="self", B="other"), to="GetONNXTensorProtoDataType(out.scalar_type())"), "aten::equal": SignatureOnly(), "aten::_softmax": Softmax("self", axis="dim"), "aten::argmax.out": SignatureOnly(), @@ -192,4 +194,12 @@ ops = {**ops, **hand_implemented} # Need to enhance the support for onnx type constrains to automatically # resolve whether the op need type promotion. # Will remove this list in the future. -type_promotion_ops = (*type_promotion_ops, "aten::gelu_backward") +type_promotion_ops.append("aten::gelu_backward") +type_promotion_ops.append("aten::gt.Tensor_out") +type_promotion_ops.append("aten::lt.Tensor_out") +type_promotion_ops.append("aten::gt.Scalar_out") +type_promotion_ops.append("aten::lt.Scalar_out") +type_promotion_ops.append("aten::ne.Tensor_out") +type_promotion_ops.append("aten::eq.Tensor_out") +type_promotion_ops.append("aten::ne.Scalar_out") +type_promotion_ops.append("aten::eq.Scalar_out") diff --git a/orttraining/orttraining/eager/ort_aten.cpp b/orttraining/orttraining/eager/ort_aten.cpp index bafabb410b..718c05517d 100644 --- a/orttraining/orttraining/eager/ort_aten.cpp +++ b/orttraining/orttraining/eager/ort_aten.cpp @@ -87,28 +87,39 @@ onnxruntime::MLDataType ort_scalar_type_from_aten( OrtValue create_ort_value( onnxruntime::ORTInvoker& invoker, const at::Scalar& scalar) { - return create_ort_value(invoker, scalar, at::kFloat); + return create_ort_value(invoker, scalar, scalar.type()); } OrtValue create_ort_value( onnxruntime::ORTInvoker& invoker, const at::Scalar& scalar, at::ScalarType type) { - float val = scalar.toFloat(); OrtValue ort_val; onnxruntime::Tensor::InitOrtValue(ort_scalar_type_from_aten(type), onnxruntime::TensorShape({}), invoker.GetCurrentExecutionProvider().GetAllocator(0, OrtMemTypeDefault), ort_val); auto* ort_tensor = ort_val.GetMutable(); switch (type) { - case at::ScalarType::Float: + case at::ScalarType::Float: { + float val = scalar.toFloat(); CopyVectorToTensor(invoker, &val, 1, *ort_tensor); break; + } case at::ScalarType::BFloat16: { at::BFloat16 valBFloat16 = scalar.toBFloat16(); - Ort::BFloat16_t *valOrtBFloat16 = reinterpret_cast(&valBFloat16); + Ort::BFloat16_t* valOrtBFloat16 = reinterpret_cast(&valBFloat16); CopyVectorToTensor(invoker, valOrtBFloat16, 1, *ort_tensor); break; } + case at::ScalarType::Double: { + double val = scalar.toDouble(); + CopyVectorToTensor(invoker, &val, 1, *ort_tensor); + break; + } + case at::ScalarType::Long: { + int64_t val = scalar.toLong(); + CopyVectorToTensor(invoker, &val, 1, *ort_tensor); + break; + } default: // TODO(unknown): support more types // For most at::ScalarType, it should be safe to just call value.to<> diff --git a/orttraining/orttraining/eager/test/ort_ops.py b/orttraining/orttraining/eager/test/ort_ops.py index 983fc9b6f6..f1fd8c5d81 100644 --- a/orttraining/orttraining/eager/test/ort_ops.py +++ b/orttraining/orttraining/eager/test/ort_ops.py @@ -47,17 +47,11 @@ ops = [ # ["isinf", ], # ["det" ]] +math_sign_ops = ["eq", "ne", "lt", "gt"] -def rename_func_to_op(testcase_func, param_num, param): - return f"test_{parameterized.to_safe_name(str(param.args[0]))}" - - -def rename_func_to_inplace(testcase_func, param_num, param): - return f"test_{parameterized.to_safe_name(str(param.args[0]))}_" - - -def rename_func_to_out(testcase_func, param_num, param): - return f"test_{parameterized.to_safe_name(str(param.args[0]))}_out" +# The function renames the test function: ops/math_sign_ops (e.g. abs)+ the test name(e.g. out), results in: test_abs_out +def rename_func(testcase_func, param_num, param): + return f"test_{parameterized.to_safe_name(str(param.args[0]))}{testcase_func.__name__[7:]}" class OrtOpTests(unittest.TestCase): @@ -380,15 +374,15 @@ class OrtOpTests(unittest.TestCase): assert torch.equal(cpu_result, ort_result.cpu()) # @parameterized.expand generate test methods for ops and using name_func we renaming the test to be test_{ops} - @parameterized.expand(ops, name_func=rename_func_to_op) + @parameterized.expand(ops, name_func=rename_func) def test_op(self, test_name, tensor_test=torch.rand(6)): # compile eval- creates a code object that evaluates the operator (for example torch.abs(tensor_test)) and returns its result. cpu_result = eval(compile("torch." + test_name + "(tensor_test)", "", "eval")) ort_result = eval(compile("torch." + test_name + "(tensor_test.to(self.get_device()))", "", "eval")) assert torch.allclose(cpu_result, ort_result.cpu(), equal_nan=True) - @parameterized.expand(ops, name_func=rename_func_to_inplace) - def test_op_inplace(self, test_name, tensor_test=torch.rand(6)): + @parameterized.expand(ops, name_func=rename_func) + def test_op_(self, test_name, tensor_test=torch.rand(6)): device = self.get_device() cpu_tensor = tensor_test @@ -399,7 +393,7 @@ class OrtOpTests(unittest.TestCase): assert torch.allclose(cpu_tensor, ort_tensor.cpu(), equal_nan=True) - @parameterized.expand(ops, name_func=rename_func_to_out) + @parameterized.expand(ops, name_func=rename_func) def test_op_out(self, test_name, tensor_test=torch.rand(6)): ##relu -don't have output if test_name == "relu": @@ -481,95 +475,81 @@ class OrtOpTests(unittest.TestCase): self.assertEqual(cpu_tensor.size(), ort_tensor.size()) self.assertTrue(torch.allclose(cpu_tensor, ort_tensor.cpu())) - def test_abs_out(self): + @parameterized.expand(math_sign_ops, name_func=rename_func) + def test_op_tensor(self, math_sign_ops): device = self.get_device() - cpu_tensor = torch.tensor([-1, -2, 3, -6, -7]) - ort_tensor = cpu_tensor.to(device) - - cpu_out_tensor = torch.tensor([], dtype=torch.long) - ort_out_tensor = cpu_out_tensor.to(device) - - cpu_result = torch.abs(cpu_tensor, out=cpu_out_tensor) - ort_result = torch.abs(ort_tensor, out=ort_out_tensor) - - assert torch.equal(cpu_result, ort_result.cpu()) - assert torch.equal(cpu_out_tensor, ort_out_tensor.cpu()) - assert torch.equal(ort_result.cpu(), ort_out_tensor.cpu()) - - def test_eq_tensor(self): - device = self.get_device() - cpu_a = torch.Tensor([1.0, 1.5, 2.0]) + cpu_a = torch.Tensor([1.0, 1.5, 2.0, 3.5]) ort_a = cpu_a.to(device) - cpu_b = torch.Tensor([1.0, 1.5, 2.1]) + cpu_b = torch.Tensor([1.0, 1.4, 2.1, 2.4]) ort_b = cpu_b.to(device) for tensor_type in {torch.float, torch.bool}: - for func in {"eq", "ne"}: - print(f"Testing {func} with type {tensor_type}") - cpu_out_tensor = torch.tensor([], dtype=tensor_type) - ort_out_tensor = cpu_out_tensor.to(device) - cpu_a_b_eq_result = eval( - compile("torch." + func + "(cpu_a, cpu_b, out=cpu_out_tensor)", "", "eval") - ) - ort_a_b_eq_result = eval( - compile("torch." + func + "(ort_a, ort_b, out=ort_out_tensor)", "", "eval") - ) - assert torch.equal(cpu_a_b_eq_result.to(device), ort_a_b_eq_result) - assert torch.equal(cpu_out_tensor, ort_out_tensor.to("cpu")) - assert ort_out_tensor.dtype == tensor_type + cpu_out_tensor = torch.tensor([], dtype=tensor_type) + ort_out_tensor = cpu_out_tensor.to(device) + cpu_a_b_result = eval( + compile("torch." + math_sign_ops + "(cpu_a, cpu_b, out=cpu_out_tensor)", "", "eval") + ) + ort_a_b_result = eval( + compile("torch." + math_sign_ops + "(ort_a, ort_b, out=ort_out_tensor)", "", "eval") + ) + assert torch.equal(cpu_a_b_result.to(device), ort_a_b_result) + assert torch.equal(cpu_out_tensor, ort_out_tensor.to("cpu")) + assert ort_out_tensor.dtype == tensor_type - def test_eq_scalar(self): + @parameterized.expand(math_sign_ops, name_func=rename_func) + def test_op_scalar(self, math_sign_ops): device = self.get_device() cpu_tensor_int = torch.tensor([1, 1], dtype=torch.int32) - cpu_scalar_int = torch.scalar_tensor(1, dtype=torch.int) - cpu_scalar_int_not = torch.scalar_tensor(2, dtype=torch.int) + cpu_scalar_int_lt = torch.scalar_tensor(2, dtype=torch.int) + cpu_scalar_int_gt = torch.scalar_tensor(0, dtype=torch.int) cpu_tensor_float = torch.tensor([1.1, 1.1], dtype=torch.float32) - cpu_scalar_float = torch.scalar_tensor(1.1, dtype=torch.float32) - cpu_scalar_float_not = torch.scalar_tensor(1.0, dtype=torch.float32) + float_lt = 1.0 + float_gt = 1.2 ort_tensor_int = cpu_tensor_int.to(device) - ort_scalar_int = cpu_scalar_int.to(device) - ort_scalar_int_not = cpu_scalar_int_not.to(device) + ort_scalar_int_lt = cpu_scalar_int_lt.to(device) + ort_scalar_int_gt = cpu_scalar_int_gt.to(device) ort_tensor_float = cpu_tensor_float.to(device) - ort_scalar_float = cpu_scalar_float.to(device) - ort_scalar_float_not = cpu_scalar_float_not.to(device) # compare int to int, float to float - ort only supports same type at the moment cpu_out_tensor = torch.tensor([], dtype=torch.bool) ort_out_tensor = cpu_out_tensor.to(device) - for func in {"eq", "ne"}: - cpu_int_int_result = eval( - compile("torch." + func + "(cpu_tensor_int, cpu_scalar_int, out=cpu_out_tensor)", "", "eval") - ) - cpu_int_int_not_result = eval( - compile("torch." + func + "(cpu_tensor_int, cpu_scalar_int_not)", "", "eval") - ) - cpu_float_float_result = eval( - compile("torch." + func + "(cpu_tensor_float, cpu_scalar_float)", "", "eval") - ) - cpu_float_float_not_result = eval( - compile("torch." + func + "(cpu_tensor_float, cpu_scalar_float_not)", "", "eval") + cpu_int_int_result = eval( + compile( + "torch." + math_sign_ops + "(cpu_tensor_int, cpu_scalar_int_lt, out=cpu_out_tensor)", "", "eval" ) + ) + cpu_int_int_gt_result = eval( + compile("torch." + math_sign_ops + "(cpu_tensor_int, cpu_scalar_int_gt)", "", "eval") + ) + cpu_float_float_lt_result = eval( + compile("torch." + math_sign_ops + "(cpu_tensor_float, float_lt)", "", "eval") + ) + cpu_float_float_gt_result = eval( + compile("torch." + math_sign_ops + "(cpu_tensor_float, float_gt)", "", "eval") + ) - ort_int_int_result = eval( - compile("torch." + func + "(ort_tensor_int, ort_scalar_int, out=ort_out_tensor)", "", "eval") - ) - ort_int_int_not_result = eval( - compile("torch." + func + "(ort_tensor_int, ort_scalar_int_not)", "", "eval") - ) - ort_float_float_result = eval( - compile("torch." + func + "(ort_tensor_float, ort_scalar_float)", "", "eval") - ) - ort_float_float_not_result = eval( - compile("torch." + func + "(ort_tensor_float, ort_scalar_float_not)", "", "eval") + ort_int_int_result = eval( + compile( + "torch." + math_sign_ops + "(ort_tensor_int, ort_scalar_int_lt, out=ort_out_tensor)", "", "eval" ) + ) + ort_int_int_gt_result = eval( + compile("torch." + math_sign_ops + "(ort_tensor_int, ort_scalar_int_gt)", "", "eval") + ) + ort_float_float_lt_result = eval( + compile("torch." + math_sign_ops + "(ort_tensor_float, float_lt)", "", "eval") + ) + ort_float_float_gt_result = eval( + compile("torch." + math_sign_ops + "(ort_tensor_float, float_gt)", "", "eval") + ) - assert torch.equal(cpu_out_tensor, ort_out_tensor.to("cpu")) - assert torch.equal(cpu_int_int_result, ort_int_int_result.to("cpu")) - assert torch.equal(cpu_int_int_not_result, ort_int_int_not_result.to("cpu")) - assert torch.equal(cpu_float_float_result, ort_float_float_result.to("cpu")) - assert torch.equal(cpu_float_float_not_result, ort_float_float_not_result.to("cpu")) + assert torch.equal(cpu_out_tensor, ort_out_tensor.to("cpu")) + assert torch.equal(cpu_int_int_result, ort_int_int_result.to("cpu")) + assert torch.equal(cpu_int_int_gt_result, ort_int_int_gt_result.to("cpu")) + assert torch.equal(cpu_float_float_lt_result, ort_float_float_lt_result.to("cpu")) + assert torch.equal(cpu_float_float_gt_result, ort_float_float_gt_result.to("cpu")) def test_fill(self): device = self.get_device()