Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
29 changes: 29 additions & 0 deletions py/torch_tensorrt/dynamo/conversion/aten_ops_converters.py
Original file line number Diff line number Diff line change
Expand Up @@ -2821,6 +2821,31 @@ def aten_ops_logical_xor(
)


# `x & False` is False whatever x is, and `x | True` is True whatever x is. For
# those two, TensorRT 11.2 folds the layer down to a constant, and it gets that
# wrong when the constant operand is smaller than the output and has to
# broadcast: the build either aborts with an internal error or the engine
# returns garbage. A Python scalar always reaches the network as a rank-0
# constant, so it always has to broadcast, so these two combinations always hit
# it. Nothing else does, including both XOR cases, since XOR has no value that
# fixes its result.
_ABSORBING_BITWISE_SCALAR = {
torch.ops.aten.bitwise_and.Scalar: False,
torch.ops.aten.bitwise_and.Scalar_Tensor: False,
torch.ops.aten.bitwise_or.Scalar: True,
torch.ops.aten.bitwise_or.Scalar_Tensor: True,
}


def _is_absorbing_bitwise_scalar(target: Target, scalar: Any) -> bool:
"""Would TensorRT mis-evaluate this scalar bitwise op (see above)?"""
if target not in _ABSORBING_BITWISE_SCALAR:
return False
# A non-bool scalar is rejected by the dtype check further down anyway, and
# `1 == True` in Python, so check the type before comparing the value.
return isinstance(scalar, bool) and scalar == _ABSORBING_BITWISE_SCALAR[target]


def bitwise_type_validator(
node: Node, settings: Optional[CompilationSettings] = None
) -> bool:
Expand Down Expand Up @@ -2870,6 +2895,8 @@ def bitwise_type_validator(
lhs_meta = lhs_val.meta.get("tensor_meta")
if lhs_meta is None:
return False
if _is_absorbing_bitwise_scalar(node.target, rhs_val):
return False
return lhs_meta.dtype in supported_type and isinstance(rhs_val, bool)

elif node.target in scalar_tensor_targets:
Expand All @@ -2878,6 +2905,8 @@ def bitwise_type_validator(
rhs_meta = rhs_val.meta.get("tensor_meta")
if rhs_meta is None:
return False
if _is_absorbing_bitwise_scalar(node.target, lhs_val):
return False
return isinstance(lhs_val, bool) and rhs_meta.dtype in supported_type

else:
Expand Down
49 changes: 46 additions & 3 deletions tests/py/dynamo/conversion/test_bitwise_and_aten.py
Original file line number Diff line number Diff line change
@@ -1,9 +1,11 @@
import unittest

import torch
import torch.nn as nn
import torch_tensorrt
from parameterized import parameterized
from torch.export import Dim
from torch.testing._internal.common_utils import run_tests
from torch.testing._internal.common_utils import TestCase, run_tests
from torch_tensorrt import Input
from torch_tensorrt.dynamo.utils import ATOL, RTOL

Expand Down Expand Up @@ -80,10 +82,12 @@ def forward(self, lhs_val, rhs_val):
use_example_tensors=False,
)

# Only True is here; `x & False` is handled by TestBitwiseAndFalseScalar
# below, because the validator keeps it out of TensorRT.
@parameterized.expand(
[
("2d", (5, 3), True),
("3d", (5, 3, 2), False),
("3d", (5, 3, 2), True),
]
)
def test_bitwise_and_scalar(self, _, shape, scalar):
Expand All @@ -104,7 +108,7 @@ def forward(self, tensor):
@parameterized.expand(
[
("2d", (5, 3), True),
("3d", (5, 3, 2), False),
("3d", (5, 3, 2), True),
]
)
def test_bitwise_and_scalar_tensor(self, _, shape, scalar):
Expand Down Expand Up @@ -165,5 +169,44 @@ def forward(self, lhs_val, rhs_val):
)


@unittest.skipIf(not torch.cuda.is_available(), "Skip because CUDA is not available")
class TestBitwiseAndFalseScalar(TestCase):
"""`x & False` must stay in PyTorch and still give the right answer.

DispatchTestCase never consults the capability validator, so these go
through the full compile pipeline and check that the partitioner built no
TensorRT engine at all.
"""

@parameterized.expand(
[
("scalar_2d", (5, 3), False),
("scalar_tensor_3d", (5, 3, 2), True),
]
)
def test_falls_back(self, _, shape, scalar_first):
class bitwise_and(nn.Module):
def forward(self, tensor):
if scalar_first:
return torch.ops.aten.bitwise_and.Scalar_Tensor(False, tensor)
return torch.ops.aten.bitwise_and.Scalar(tensor, False)

mod = bitwise_and().eval().cuda()
inputs = [torch.randint(0, 2, shape, dtype=bool).cuda()]
trt_mod = torch_tensorrt.compile(
mod,
ir="dynamo",
inputs=inputs,
min_block_size=1,
cache_built_engines=False,
reuse_cached_engines=False,
)
acc_count = sum(
1 for name, _ in trt_mod.named_children() if "_run_on_acc" in name
)
self.assertEqual(acc_count, 0)
torch.testing.assert_close(trt_mod(*inputs), mod(*inputs))


if __name__ == "__main__":
run_tests()
50 changes: 47 additions & 3 deletions tests/py/dynamo/conversion/test_bitwise_or_aten.py
Original file line number Diff line number Diff line change
@@ -1,7 +1,10 @@
import unittest

import torch
import torch.nn as nn
import torch_tensorrt
from parameterized import parameterized
from torch.testing._internal.common_utils import run_tests
from torch.testing._internal.common_utils import TestCase, run_tests
from torch_tensorrt import Input

from .harness import DispatchTestCase
Expand Down Expand Up @@ -77,9 +80,11 @@ def forward(self, lhs_val, rhs_val):
use_example_tensors=False,
)

# Only False is here; `x | True` is handled by TestBitwiseOrTrueScalar
# below, because the validator keeps it out of TensorRT.
@parameterized.expand(
[
("2d", (5, 3), True),
("2d", (5, 3), False),
("3d", (5, 3, 2), False),
]
)
Expand All @@ -100,7 +105,7 @@ def forward(self, tensor):

@parameterized.expand(
[
("2d", (5, 3), True),
("2d", (5, 3), False),
("3d", (5, 3, 2), False),
]
)
Expand All @@ -120,5 +125,44 @@ def forward(self, tensor):
)


@unittest.skipIf(not torch.cuda.is_available(), "Skip because CUDA is not available")
class TestBitwiseOrTrueScalar(TestCase):
"""`x | True` must stay in PyTorch and still give the right answer.

DispatchTestCase never consults the capability validator, so these go
through the full compile pipeline and check that the partitioner built no
TensorRT engine at all.
"""

@parameterized.expand(
[
("scalar_2d", (5, 3), False),
("scalar_tensor_3d", (5, 3, 2), True),
]
)
def test_falls_back(self, _, shape, scalar_first):
class bitwise_or(nn.Module):
def forward(self, tensor):
if scalar_first:
return torch.ops.aten.bitwise_or.Scalar_Tensor(True, tensor)
return torch.ops.aten.bitwise_or.Scalar(tensor, True)

mod = bitwise_or().eval().cuda()
inputs = [torch.randint(0, 2, shape, dtype=bool).cuda()]
trt_mod = torch_tensorrt.compile(
mod,
ir="dynamo",
inputs=inputs,
min_block_size=1,
cache_built_engines=False,
reuse_cached_engines=False,
)
acc_count = sum(
1 for name, _ in trt_mod.named_children() if "_run_on_acc" in name
)
self.assertEqual(acc_count, 0)
torch.testing.assert_close(trt_mod(*inputs), mod(*inputs))


if __name__ == "__main__":
run_tests()
63 changes: 63 additions & 0 deletions tests/py/dynamo/conversion/test_bitwise_validator_aten.py
Original file line number Diff line number Diff line change
Expand Up @@ -23,6 +23,18 @@ def make_operand(shape):
return node


def _make_bitwise_scalar_node(target, tensor_shape, scalar, scalar_first=False):
operand = MagicMock()
operand.meta = {
"tensor_meta": SimpleNamespace(dtype=torch.bool, shape=torch.Size(tensor_shape))
}

node = MagicMock()
node.target = target
node.args = (scalar, operand) if scalar_first else (operand, scalar)
return node


class TestBitwiseValidator(unittest.TestCase):
def test_bitwise_and_scalar_tensor_other_falls_back(self):
node = _make_bitwise_node(torch.ops.aten.bitwise_and.Tensor, (2, 3), ())
Expand All @@ -40,6 +52,57 @@ def test_bitwise_xor_scalar_tensor_other_is_supported(self):
node = _make_bitwise_node(torch.ops.aten.bitwise_xor.Tensor, (2, 3), ())
self.assertTrue(bitwise_type_validator(node))

# A scalar operand only breaks when its value fixes the result of the op:
# False for AND, True for OR. The other value, and XOR either way, stay in
# TensorRT.
def test_bitwise_and_false_scalar_falls_back(self):
for target, scalar_first in (
(torch.ops.aten.bitwise_and.Scalar, False),
(torch.ops.aten.bitwise_and.Scalar_Tensor, True),
):
with self.subTest(target=target):
node = _make_bitwise_scalar_node(target, (2, 3), False, scalar_first)
self.assertFalse(bitwise_type_validator(node))

def test_bitwise_and_true_scalar_is_supported(self):
for target, scalar_first in (
(torch.ops.aten.bitwise_and.Scalar, False),
(torch.ops.aten.bitwise_and.Scalar_Tensor, True),
):
with self.subTest(target=target):
node = _make_bitwise_scalar_node(target, (2, 3), True, scalar_first)
self.assertTrue(bitwise_type_validator(node))

def test_bitwise_or_true_scalar_falls_back(self):
for target, scalar_first in (
(torch.ops.aten.bitwise_or.Scalar, False),
(torch.ops.aten.bitwise_or.Scalar_Tensor, True),
):
with self.subTest(target=target):
node = _make_bitwise_scalar_node(target, (2, 3), True, scalar_first)
self.assertFalse(bitwise_type_validator(node))

def test_bitwise_or_false_scalar_is_supported(self):
for target, scalar_first in (
(torch.ops.aten.bitwise_or.Scalar, False),
(torch.ops.aten.bitwise_or.Scalar_Tensor, True),
):
with self.subTest(target=target):
node = _make_bitwise_scalar_node(target, (2, 3), False, scalar_first)
self.assertTrue(bitwise_type_validator(node))

def test_bitwise_xor_scalar_is_supported(self):
for target, scalar_first in (
(torch.ops.aten.bitwise_xor.Scalar, False),
(torch.ops.aten.bitwise_xor.Scalar_Tensor, True),
):
for scalar in (True, False):
with self.subTest(target=target, scalar=scalar):
node = _make_bitwise_scalar_node(
target, (2, 3), scalar, scalar_first
)
self.assertTrue(bitwise_type_validator(node))


if __name__ == "__main__":
run_tests()
Loading