mirror of
https://github.com/saymrwulf/onnxruntime.git
synced 2026-07-30 20:18:08 +00:00
Implementing aten::gt.Scalar_out and aten::lt.Scalar_out (#12181)
* Implementing aten::gt.Scalar_out and aten::lt.Scalar_out * modified the code according to code review
This commit is contained in:
parent
007ef42749
commit
7dc45bc311
3 changed files with 90 additions and 89 deletions
|
|
@ -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")
|
||||
|
|
|
|||
|
|
@ -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<onnxruntime::Tensor>();
|
||||
switch (type) {
|
||||
case at::ScalarType::Float:
|
||||
case at::ScalarType::Float: {
|
||||
float val = scalar.toFloat();
|
||||
CopyVectorToTensor<float>(invoker, &val, 1, *ort_tensor);
|
||||
break;
|
||||
}
|
||||
case at::ScalarType::BFloat16: {
|
||||
at::BFloat16 valBFloat16 = scalar.toBFloat16();
|
||||
Ort::BFloat16_t *valOrtBFloat16 = reinterpret_cast<Ort::BFloat16_t *>(&valBFloat16);
|
||||
Ort::BFloat16_t* valOrtBFloat16 = reinterpret_cast<Ort::BFloat16_t*>(&valBFloat16);
|
||||
CopyVectorToTensor<Ort::BFloat16_t>(invoker, valOrtBFloat16, 1, *ort_tensor);
|
||||
break;
|
||||
}
|
||||
case at::ScalarType::Double: {
|
||||
double val = scalar.toDouble();
|
||||
CopyVectorToTensor<double>(invoker, &val, 1, *ort_tensor);
|
||||
break;
|
||||
}
|
||||
case at::ScalarType::Long: {
|
||||
int64_t val = scalar.toLong();
|
||||
CopyVectorToTensor<int64_t>(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<>
|
||||
|
|
|
|||
|
|
@ -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)", "<string>", "eval"))
|
||||
ort_result = eval(compile("torch." + test_name + "(tensor_test.to(self.get_device()))", "<string>", "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)", "<string>", "eval")
|
||||
)
|
||||
ort_a_b_eq_result = eval(
|
||||
compile("torch." + func + "(ort_a, ort_b, out=ort_out_tensor)", "<string>", "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)", "<string>", "eval")
|
||||
)
|
||||
ort_a_b_result = eval(
|
||||
compile("torch." + math_sign_ops + "(ort_a, ort_b, out=ort_out_tensor)", "<string>", "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)", "<string>", "eval")
|
||||
)
|
||||
cpu_int_int_not_result = eval(
|
||||
compile("torch." + func + "(cpu_tensor_int, cpu_scalar_int_not)", "<string>", "eval")
|
||||
)
|
||||
cpu_float_float_result = eval(
|
||||
compile("torch." + func + "(cpu_tensor_float, cpu_scalar_float)", "<string>", "eval")
|
||||
)
|
||||
cpu_float_float_not_result = eval(
|
||||
compile("torch." + func + "(cpu_tensor_float, cpu_scalar_float_not)", "<string>", "eval")
|
||||
cpu_int_int_result = eval(
|
||||
compile(
|
||||
"torch." + math_sign_ops + "(cpu_tensor_int, cpu_scalar_int_lt, out=cpu_out_tensor)", "<string>", "eval"
|
||||
)
|
||||
)
|
||||
cpu_int_int_gt_result = eval(
|
||||
compile("torch." + math_sign_ops + "(cpu_tensor_int, cpu_scalar_int_gt)", "<string>", "eval")
|
||||
)
|
||||
cpu_float_float_lt_result = eval(
|
||||
compile("torch." + math_sign_ops + "(cpu_tensor_float, float_lt)", "<string>", "eval")
|
||||
)
|
||||
cpu_float_float_gt_result = eval(
|
||||
compile("torch." + math_sign_ops + "(cpu_tensor_float, float_gt)", "<string>", "eval")
|
||||
)
|
||||
|
||||
ort_int_int_result = eval(
|
||||
compile("torch." + func + "(ort_tensor_int, ort_scalar_int, out=ort_out_tensor)", "<string>", "eval")
|
||||
)
|
||||
ort_int_int_not_result = eval(
|
||||
compile("torch." + func + "(ort_tensor_int, ort_scalar_int_not)", "<string>", "eval")
|
||||
)
|
||||
ort_float_float_result = eval(
|
||||
compile("torch." + func + "(ort_tensor_float, ort_scalar_float)", "<string>", "eval")
|
||||
)
|
||||
ort_float_float_not_result = eval(
|
||||
compile("torch." + func + "(ort_tensor_float, ort_scalar_float_not)", "<string>", "eval")
|
||||
ort_int_int_result = eval(
|
||||
compile(
|
||||
"torch." + math_sign_ops + "(ort_tensor_int, ort_scalar_int_lt, out=ort_out_tensor)", "<string>", "eval"
|
||||
)
|
||||
)
|
||||
ort_int_int_gt_result = eval(
|
||||
compile("torch." + math_sign_ops + "(ort_tensor_int, ort_scalar_int_gt)", "<string>", "eval")
|
||||
)
|
||||
ort_float_float_lt_result = eval(
|
||||
compile("torch." + math_sign_ops + "(ort_tensor_float, float_lt)", "<string>", "eval")
|
||||
)
|
||||
ort_float_float_gt_result = eval(
|
||||
compile("torch." + math_sign_ops + "(ort_tensor_float, float_gt)", "<string>", "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()
|
||||
|
|
|
|||
Loading…
Reference in a new issue