From 04e13e19d463b1995fe97a5e21742d28aff6c891 Mon Sep 17 00:00:00 2001 From: Ihor Indyk Date: Wed, 22 Jul 2026 19:17:05 -0700 Subject: [PATCH] Fix reference cycle in stats and BatchMapDataset. PiperOrigin-RevId: 952465059 --- CHANGELOG.md | 1 + grain/_src/python/dataset/dataset_test.py | 91 ++++++++++++++++++- grain/_src/python/dataset/stats.py | 6 +- grain/_src/python/dataset/stats_test.py | 1 + .../python/dataset/transformations/batch.py | 10 +- 5 files changed, 98 insertions(+), 11 deletions(-) diff --git a/CHANGELOG.md b/CHANGELOG.md index 7efdcc63e..cdefd1f88 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -17,6 +17,7 @@ changes. Best viewed [here](https://google-grain.readthedocs.io/en/latest/change * Bug fixes: * Fixed bug in DataLoader where sharding remainder was dropped even when ShardOptions.drop_remainder=False. + * Fixes reference cycle in BatchMapDataset. ## Grain 0.2.18 (June 17, 2026) diff --git a/grain/_src/python/dataset/dataset_test.py b/grain/_src/python/dataset/dataset_test.py index c50ebe883..c0078e200 100644 --- a/grain/_src/python/dataset/dataset_test.py +++ b/grain/_src/python/dataset/dataset_test.py @@ -17,7 +17,7 @@ import gc import sys import time -from typing import TypeVar +from typing import Any, Sequence, TypeVar from unittest import mock from absl.testing import absltest @@ -1432,5 +1432,94 @@ def test_get_element_spec_from_iter_dataset(self): self.assertEqual(spec.dtype, np.int64) +def _find_shortest_cycle_recursive( + path: list[Any], + current_obj: Any, + visited_ids: set[int], + object_predicate: Any, +) -> list[Any] | None: + if not object_predicate(current_obj): + return None + + current_id = id(current_obj) + visited_ids.add(current_id) + path.append(current_obj) + + shortest_cycle = None + neighbors = gc.get_referents(current_obj) + for neighbor in neighbors: + neighbor_id = id(neighbor) + if neighbor_id == id(path[0]): + shortest_cycle = list(path) + break + elif neighbor_id not in visited_ids: + cycle = _find_shortest_cycle_recursive( + path, + neighbor, + visited_ids, + object_predicate, + ) + if cycle is not None and ( + shortest_cycle is None or len(cycle) < len(shortest_cycle) + ): + shortest_cycle = cycle + + if shortest_cycle is not None: + visited_ids.remove(current_id) + + path.pop() + return shortest_cycle + + +def _find_cycles(objects: Sequence[Any]) -> list[list[Any]]: + object_ids = {id(o) for o in objects} + object_predicate = lambda obj: id(obj) in object_ids + + path = [] + visited_ids = set() + cycles = [] + for obj in objects: + visited_ids_from_obj = set(visited_ids) + shortest_cycle = _find_shortest_cycle_recursive( + path, + obj, + visited_ids_from_obj, + object_predicate, + ) + if shortest_cycle is not None: + cycles.append(shortest_cycle) + visited_ids = visited_ids.union( + id(cycle_obj) for cycle_obj in shortest_cycle + ) + else: + visited_ids.add(id(obj)) + return cycles + + +class ReferenceCycleTest(parameterized.TestCase): + + def test_batch_map_no_reference_cycles(self): + original_debug_flags = gc.get_debug() + try: + gc.disable() + gc.garbage.clear() + gc.collect() + gc.set_debug(gc.DEBUG_SAVEALL) + + ds = dataset.MapDataset.range(10).batch(batch_size=2) + _ = next(iter(ds)) + + del ds + gc.collect() + + reference_cycles = _find_cycles(gc.garbage) + self.assertEmpty(reference_cycles) + finally: + gc.set_debug(original_debug_flags) + gc.garbage.clear() + gc.collect() + gc.enable() + + if __name__ == "__main__": absltest.main() diff --git a/grain/_src/python/dataset/stats.py b/grain/_src/python/dataset/stats.py index 0d8e03adb..95162399d 100644 --- a/grain/_src/python/dataset/stats.py +++ b/grain/_src/python/dataset/stats.py @@ -764,12 +764,8 @@ def _running_in_colab() -> bool: class _DefaultStats(Stats): """Default implementation for statistics collection that does nothing.""" - def __init__(self, config: StatsConfig, parents: Sequence[Stats]): - super().__init__(config, parents) - - @contextlib.contextmanager def record_self_time(self, *, num_elements: int = 1, offset_ns: int = 0): - yield + return contextlib.nullcontext() def record_output_spec(self, element: T) -> T: return element diff --git a/grain/_src/python/dataset/stats_test.py b/grain/_src/python/dataset/stats_test.py index 7edc9aee0..7c0d36839 100644 --- a/grain/_src/python/dataset/stats_test.py +++ b/grain/_src/python/dataset/stats_test.py @@ -21,6 +21,7 @@ import threading import time from unittest import mock +import weakref from absl import flags from absl.testing import flagsaver diff --git a/grain/_src/python/dataset/transformations/batch.py b/grain/_src/python/dataset/transformations/batch.py index 0f204ca26..4f4cd9827 100644 --- a/grain/_src/python/dataset/transformations/batch.py +++ b/grain/_src/python/dataset/transformations/batch.py @@ -37,6 +37,7 @@ S = TypeVar("S") +@functools.cache def _is_batch_map_pushdown_experiment_enabled() -> bool: return False @@ -435,13 +436,12 @@ def set_slice(self, sl: slice, sequential_slice: bool = False) -> None: dataset.set_slice(self._parent, sl, sequential_slice) self._update_length() - @functools.cached_property - def _get_parent_items_fn(self): + def _get_parent_items(self, items): # Leverage batch pushdown API to retrieve multiple items at once if the # experiment is enabled. if _is_batch_map_pushdown_experiment_enabled(): - return lambda items: self._parent._getitems(list(items)) # pylint: disable=protected-access - return lambda items: [self._parent[i] for i in items] + return self._parent._getitems(list(items)) # pylint: disable=protected-access + return [self._parent[i] for i in items] def __len__(self): return self._length @@ -458,7 +458,7 @@ def __getitem__(self, index): # Add offset for epoch. start += epoch * self._parent_length stop += epoch * self._parent_length - values = self._get_parent_items_fn(range(start, stop)) + values = self._get_parent_items(range(start, stop)) with self._stats.record_self_time(): try: return self._stats.record_output_spec(self._batch_fn(values))