Skip to content

Commit 1283670

Browse files
guilhermeleobaspull[bot]
authored andcommitted
Implement mp_subscript / sq_concat / bool_impl for ConstantVariable (#192825)
Pull Request resolved: #192825 Approved by: https://github.com/hameerabbasi
1 parent 5154e00 commit 1283670

9 files changed

Lines changed: 146 additions & 9 deletions

File tree

‎test/dynamo/test_getitem.py‎

Lines changed: 14 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -840,6 +840,20 @@ def fn(x):
840840
x = torch.randn(4)
841841
self.assertEqual(fn(x), self._compile(fn, x))
842842

843+
def test_str_subscript_symbolic_index(self):
844+
# A non-constant key must fall through to the generic "unsupported
845+
# subscript" graph break, not leak AsPythonConstantNotImplementedError.
846+
def fn(t):
847+
i = t.item()
848+
torch._check(i >= 0)
849+
torch._check(i < 3)
850+
return "abc"[i]
851+
852+
with self.assertRaisesRegex(
853+
torch._dynamo.exc.Unsupported, "does not yet support subscripting 'str'"
854+
):
855+
self._compile(fn, torch.tensor(1))
856+
843857
# ===================================================================
844858
# Explicit __getitem__ dunder call path tests
845859
# Exercises: obj.__getitem__(key) → LOAD_ATTR + CALL, which may

‎test/dynamo/test_nb_bool.py‎

Lines changed: 40 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -3,6 +3,8 @@
33

44
import collections
55
import enum
6+
import sys
7+
import unittest
68

79
import torch
810

@@ -123,6 +125,20 @@ def test_empty_range(self):
123125
def test_nonempty_range(self):
124126
self.assertEqual(bool(range(5)), True)
125127

128+
@make_dynamo_test
129+
def test_empty_deque(self):
130+
self.assertEqual(bool(collections.deque()), False)
131+
132+
@make_dynamo_test
133+
def test_nonempty_deque(self):
134+
self.assertEqual(bool(collections.deque([1, 2])), True)
135+
136+
@unittest.skipIf(sys.version_info >= (3, 11), "deque lost nb_bool in 3.11")
137+
@make_dynamo_test
138+
def test_deque_dunder_bool(self):
139+
self.assertEqual(collections.deque().__bool__(), False)
140+
self.assertEqual(collections.deque([1, 2]).__bool__(), True)
141+
126142
# --- dict subclasses ---
127143

128144
@make_dynamo_test
@@ -312,6 +328,30 @@ def fn(x):
312328
compiled = torch.compile(fn, backend="eager", fullgraph=True)
313329
self.assertEqual(fn(x), compiled(x))
314330

331+
# --- int-typed VTs with no runtime object ---
332+
333+
def test_fake_id_bool(self):
334+
# id() of an object minted during tracing yields a FakeIdVariable,
335+
# whose python_type is int and so fills nb_bool.
336+
def fn(x):
337+
tmp = [1, 2, 3]
338+
return x + 1, bool(id(tmp))
339+
340+
x = torch.randn(4)
341+
compiled = torch.compile(fn, backend="eager", fullgraph=True)
342+
self.assertEqual(compiled(x), fn(x))
343+
344+
def test_data_ptr_bool_graph_breaks(self):
345+
# A data pointer is only known at runtime, so its truth value cannot be
346+
# decided at trace time (torch.empty(0).data_ptr() is 0).
347+
def fn(x):
348+
return bool(x.data_ptr())
349+
350+
with self.assertRaisesRegex(
351+
torch._dynamo.exc.Unsupported, "Data pointer truth value"
352+
):
353+
torch.compile(fn, backend="eager", fullgraph=True)(torch.randn(4))
354+
315355
# --- Tensor (TensorVariable path) ---
316356

317357
def test_tensor_nonzero(self):

‎torch/_dynamo/graph_break_registry.json‎

Lines changed: 18 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -3732,6 +3732,14 @@
37323732
]
37333733
}
37343734
],
3735+
"GB9118": [
3736+
{
3737+
"Gb_type": "Data pointer truth value",
3738+
"Context": "nb_bool_impl {self}",
3739+
"Explanation": "Dynamo cannot decide the truth value of a data pointer because the address is only known at runtime.",
3740+
"Hints": []
3741+
}
3742+
],
37353743
"GB0233": [
37363744
{
37373745
"Gb_type": "Attempted to use strided NestedTensor",
@@ -5933,6 +5941,16 @@
59335941
"Hints": []
59345942
}
59355943
],
5944+
"GB2582": [
5945+
{
5946+
"Gb_type": "Missing nb_bool_impl override",
5947+
"Context": "nb_bool_impl {self}",
5948+
"Explanation": "{type(self).__name__} does not implement nb_bool_impl. Add a nb_bool_impl override to {type(self).__name__}.",
5949+
"Hints": [
5950+
"This is likely to be a Dynamo bug. Please report an issue to PyTorch."
5951+
]
5952+
}
5953+
],
59365954
"GB0353": [
59375955
{
59385956
"Gb_type": "rewrite_signature: cannot trace optional function input",

‎torch/_dynamo/variables/base.py‎

Lines changed: 7 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -1813,13 +1813,15 @@ def is_python_constant(self) -> bool:
18131813
except NotImplementedError:
18141814
return False
18151815

1816-
def nb_bool_impl(self, tx: InstructionTranslatorBase) -> VariableTracker | None:
1816+
def nb_bool_impl(self, tx: InstructionTranslatorBase) -> VariableTracker:
18171817
# Mirrors CPython's tp_as_number->nb_bool slot.
18181818
# https://github.com/python/cpython/blob/c09ccd9c429/Objects/object.c#L2135-L2158
1819-
#
1820-
# Returns None when the type has no nb_bool, causing generic_is_true to
1821-
# fall through to length check, then truthy default.
1822-
return None
1819+
unimplemented(
1820+
gb_type="Missing nb_bool_impl override",
1821+
context=f"nb_bool_impl {self}",
1822+
explanation=f"{type(self).__name__} does not implement nb_bool_impl. Add a nb_bool_impl override to {type(self).__name__}.",
1823+
hints=[*graph_break_hints.DYNAMO_BUG],
1824+
)
18231825

18241826
def is_hashable(self) -> bool:
18251827
"""Whether the underlying Python object is hashable.

‎torch/_dynamo/variables/constant.py‎

Lines changed: 43 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -257,6 +257,20 @@ def mp_length_impl(self, tx: InstructionTranslatorBase) -> VariableTracker:
257257
"""Mapping length - delegates to len_impl for constants."""
258258
return self.len_impl(tx)
259259

260+
def mp_subscript_impl(
261+
self,
262+
tx: InstructionTranslatorBase,
263+
key: VariableTracker,
264+
) -> VariableTracker:
265+
from .object_protocol import type_implements_mp_subscript
266+
267+
if type_implements_mp_subscript(type(self.value)) and key.is_python_constant():
268+
try:
269+
return ConstantVariable.create(self.value[key.as_python_constant()])
270+
except Exception as e:
271+
raise_observed_exception(type(e), tx, args=list(e.args))
272+
return super().mp_subscript_impl(tx, key)
273+
260274
def const_getattr(
261275
self, tx: InstructionTranslatorBase, name: str
262276
) -> VariableTracker:
@@ -267,6 +281,18 @@ def const_getattr(
267281
raise NotImplementedError
268282
return member
269283

284+
def sq_concat_impl(
285+
self, tx: InstructionTranslatorBase, other: VariableTracker
286+
) -> VariableTracker:
287+
from .object_protocol import type_implements_sq_concat
288+
289+
if type_implements_sq_concat(type(self.value)) and other.is_python_constant():
290+
try:
291+
return ConstantVariable.create(self.value + other.as_python_constant())
292+
except Exception as e:
293+
raise_observed_exception(type(e), tx, args=list(e.args))
294+
return super().sq_concat_impl(tx, other)
295+
270296
def sq_contains_impl(self, tx: InstructionTranslatorBase, item: VariableTracker):
271297
"""Sequence contains for constants."""
272298
if item.is_python_constant():
@@ -468,6 +494,18 @@ def get_id(self, tx: InstructionTranslatorBase) -> int | None:
468494
def get_real_python_backed_value(self) -> object:
469495
return self.value
470496

497+
def nb_bool_impl(
498+
self,
499+
tx: InstructionTranslatorBase,
500+
) -> VariableTracker:
501+
# CPython: int, float, and bool define nb_bool (returns self for int,
502+
# bool(self) for float). All other constant types do not.
503+
from .object_protocol import type_implements_nb_bool
504+
505+
if type_implements_nb_bool(type(self.value)):
506+
return ConstantVariable.create(bool(self.value))
507+
return super().nb_bool_impl(tx)
508+
471509
def nb_index_impl(
472510
self,
473511
tx: InstructionTranslatorBase,
@@ -913,6 +951,11 @@ def tp_repr_impl(self, tx: InstructionTranslatorBase) -> VariableTracker:
913951
# FakeIdVariable already resolves same-kind id()/hash() comparisons.
914952
return ConstantVariable.create(repr(self.value))
915953

954+
def nb_bool_impl(self, tx: InstructionTranslatorBase) -> VariableTracker:
955+
# Mirrors long_bool. The fake value is only meaningful at compile time,
956+
# but its truthiness is, like tp_repr_impl, a plain function of it.
957+
return ConstantVariable.create(bool(self.value))
958+
916959
def tp_richcompare_impl(
917960
self, tx: InstructionTranslatorBase, other: VariableTracker, op: str
918961
) -> VariableTracker:

‎torch/_dynamo/variables/lists.py‎

Lines changed: 12 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -806,6 +806,10 @@ def stop(self) -> int:
806806
def step(self) -> int:
807807
return guard_if_dyn(self.items[2])
808808

809+
def nb_bool_impl(self, tx: "InstructionTranslatorBase") -> VariableTracker:
810+
# ref: range_bool in https://github.com/python/cpython/blob/v3.13.0/Objects/rangeobject.c#L740-L744
811+
return ConstantVariable.create(self.range_length() != 0)
812+
809813
def range_length(self) -> int:
810814
lo = self.start()
811815
hi = self.stop()
@@ -1416,6 +1420,14 @@ def tp_richcompare_impl(
14161420
) -> VariableTracker:
14171421
return self._seq_richcompare(tx, other, op, collections.deque)
14181422

1423+
if sys.version_info < (3, 11):
1424+
1425+
def nb_bool_impl(self, tx: "InstructionTranslatorBase") -> VariableTracker:
1426+
# deque fills nb_bool (deque_bool: Py_SIZE(deque) != 0) up to Python
1427+
# 3.10; CPython GH-32397 dropped the slot in 3.11, so newer versions
1428+
# fall through to sq_length in generic_is_true and never reach here.
1429+
return ConstantVariable.create(len(self.items) > 0)
1430+
14191431
def is_hashable(self) -> bool:
14201432
return False
14211433

‎torch/_dynamo/variables/object_protocol.py‎

Lines changed: 1 addition & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -401,9 +401,7 @@ def generic_is_true(
401401
raise_observed_exception(type(e), tx, args=[str(e)])
402402

403403
if obj.tp_as_number.nb_bool:
404-
result = obj.nb_bool_impl(tx)
405-
if result is not None:
406-
return result
404+
return obj.nb_bool_impl(tx)
407405

408406
try:
409407
length = generic_size(tx, obj)

‎torch/_dynamo/variables/tensor.py‎

Lines changed: 10 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -3671,6 +3671,16 @@ def tp_richcompare_impl(
36713671
hints=[],
36723672
)
36733673

3674+
def nb_bool_impl(self, tx: "InstructionTranslatorBase") -> VariableTracker:
3675+
"""DataPtr nb_bool: mirrors long_bool, but the address is runtime-only."""
3676+
unimplemented(
3677+
gb_type="Data pointer truth value",
3678+
context=f"nb_bool_impl {self}",
3679+
explanation="Dynamo cannot decide the truth value of a data pointer "
3680+
"because the address is only known at runtime.",
3681+
hints=[],
3682+
)
3683+
36743684
def reconstruct(self, codegen: "PyCodegen") -> None:
36753685
codegen(self.from_tensor)
36763686
codegen.load_method(self.method_name)

‎torch/_dynamo/variables/user_defined.py‎

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -1896,7 +1896,7 @@ def guard_as_python_constant(self) -> object:
18961896
def nb_bool_impl(
18971897
self,
18981898
tx: "InstructionTranslatorBase",
1899-
) -> "VariableTracker | None":
1899+
) -> "VariableTracker":
19001900
# Mirrors slot_nb_bool:
19011901
# https://github.com/python/cpython/blob/c09ccd9c429/Objects/typeobject.c#L9408-L9458
19021902
res = self._maybe_call_special(tx, "__bool__", [])

0 commit comments

Comments
 (0)