Skip to content

Commit 9de41a5

Browse files
ArrayRecord Teamcopybara-github
authored andcommitted
Fix determinism bug in ArrayRecordDataSource
PiperOrigin-RevId: 896769930
1 parent c909611 commit 9de41a5

2 files changed

Lines changed: 54 additions & 32 deletions

File tree

python/array_record_data_source.py

Lines changed: 49 additions & 27 deletions
Original file line numberDiff line numberDiff line change
@@ -38,6 +38,7 @@ def __getitem__(self, record_keys: Sequence[int]) -> Sequence[T]:
3838
import itertools
3939
import os
4040
import pathlib
41+
import queue
4142
import re
4243
import typing
4344
from typing import Any, Callable, Iterator, List, Mapping, Protocol, Sequence, SupportsIndex, Tuple, TypeVar, Union
@@ -271,7 +272,9 @@ def __init__(
271272
self._read_instructions = _get_read_instructions(paths)
272273
self._paths = [ri.filename for ri in self._read_instructions]
273274
# We open readers lazily when we need to read from them.
274-
self._readers = [None] * len(self._read_instructions)
275+
# To ensure thread safety while allowing concurrent reads, we maintain a
276+
# pool of readers for each shard.
277+
self._reader_pools = [queue.LifoQueue() for _ in self._read_instructions]
275278
self._num_records = sum(
276279
map(lambda x: x.num_records, self._read_instructions)
277280
)
@@ -286,10 +289,11 @@ def __enter__(self):
286289

287290
def __exit__(self, exc_type, exc_value, traceback):
288291
logging.debug("__exit__ for ArrayRecordDataSource is called.")
289-
for reader in self._readers:
290-
if reader:
291-
reader.close()
292-
self._readers = [None] * len(self._read_instructions)
292+
for pool in self._reader_pools:
293+
while not pool.empty():
294+
reader = pool.get()
295+
if reader:
296+
reader.close()
293297

294298
def __len__(self) -> int:
295299
return self._num_records
@@ -329,21 +333,29 @@ def _split_keys_per_reader(
329333
positions_and_indices[reader_idx] = [(position, idx)]
330334
return positions_and_indices
331335

332-
def _ensure_reader_exists(self, reader_idx: int) -> None:
333-
"""Threadsafe method to create corresponding reader if it doesn't exist."""
334-
if self._readers[reader_idx] is not None:
335-
return
336-
filename = self._read_instructions[reader_idx].filename
337-
reader = _create_reader(filename, self._reader_options_string)
338-
_check_group_size(filename, reader)
339-
self._readers[reader_idx] = reader
336+
def _get_reader(self, reader_idx: int) -> Any:
337+
"""Gets a reader from the pool or creates a new one."""
338+
try:
339+
return self._reader_pools[reader_idx].get_nowait()
340+
except queue.Empty:
341+
filename = self._read_instructions[reader_idx].filename
342+
reader = _create_reader(filename, self._reader_options_string)
343+
_check_group_size(filename, reader)
344+
return reader
345+
346+
def _release_reader(self, reader_idx: int, reader: Any) -> None:
347+
"""Returns a reader to the pool."""
348+
self._reader_pools[reader_idx].put(reader)
340349

341350
def __getitem__(self, record_key: SupportsIndex) -> bytes:
342351
reader_idx, position = self._reader_idx_and_position(record_key)
343-
self._ensure_reader_exists(reader_idx)
344-
if hasattr(self._readers[reader_idx], "read"):
345-
return self._readers[reader_idx].read([position])[0]
346-
return self._readers[reader_idx][position]
352+
reader = self._get_reader(reader_idx)
353+
try:
354+
if hasattr(reader, "read"):
355+
return reader.read([position])[0]
356+
return reader[position]
357+
finally:
358+
self._release_reader(reader_idx, reader)
347359

348360
def __getitems__(
349361
self, record_keys: Sequence[SupportsIndex]
@@ -352,14 +364,16 @@ def read_records(
352364
reader_idx: int, reader_positions_and_indices: Sequence[Tuple[int, int]]
353365
) -> Sequence[Tuple[Any, int]]:
354366
"""Reads records using the given reader keeping track of the indices."""
355-
# Initialize readers lazily when we need to read from them.
356-
self._ensure_reader_exists(reader_idx)
357-
positions, indices = list(zip(*reader_positions_and_indices))
358-
if hasattr(self._readers[reader_idx], "read"):
359-
records = self._readers[reader_idx].read(positions) # pytype: disable=attribute-error
360-
else:
361-
records = [self._readers[reader_idx][p] for p in positions]
362-
return list(zip(records, indices))
367+
reader = self._get_reader(reader_idx)
368+
try:
369+
positions, indices = list(zip(*reader_positions_and_indices))
370+
if hasattr(reader, "read"):
371+
records = reader.read(positions) # pytype: disable=attribute-error
372+
else:
373+
records = [reader[p] for p in positions]
374+
return list(zip(records, indices))
375+
finally:
376+
self._release_reader(reader_idx, reader)
363377

364378
positions_and_indices = self._split_keys_per_reader(record_keys)
365379
num_threads = _get_flag_value(_GRAIN_NUM_THREADS_FETCHING_RECORDS)
@@ -390,15 +404,23 @@ def read_records(
390404
def __getstate__(self):
391405
logging.debug("__getstate__ for ArrayRecordDataSource is called.")
392406
state = self.__dict__.copy()
393-
del state["_readers"]
407+
del state["_reader_pools"]
394408
return state
395409

396410
def __setstate__(self, state):
397411
logging.debug("__setstate__ for ArrayRecordDataSource is called.")
398412
self.__dict__.update(state)
399413
# We open readers lazily when we need to read from them. Thus, we don't
400414
# need to re-open the same files as before pickling.
401-
self._readers = [None] * len(self._read_instructions)
415+
self._reader_pools = [queue.LifoQueue() for _ in self._read_instructions]
416+
417+
def _peek_readers(self) -> List[Any]:
418+
"""Returns a list of readers (one per shard) or None (for testing)."""
419+
readers = []
420+
for pool in self._reader_pools:
421+
with pool.mutex:
422+
readers.append(pool.queue[-1] if pool.queue else None)
423+
return readers
402424

403425
def __repr__(self) -> str:
404426
"""Storing a hash of paths since paths can be a very long list."""

python/array_record_data_source_test.py

Lines changed: 5 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -109,7 +109,7 @@ def test_array_record_data_source_single_path(self):
109109
) as ar:
110110
actual_data = [ar[x] for x in indices_to_read]
111111
self.assertEqual(expected_data, actual_data)
112-
self.assertTrue(all(reader is None for reader in ar._readers))
112+
self.assertTrue(all(reader is None for reader in ar._peek_readers()))
113113

114114
def test_array_record_data_source_string_read_instructions(self):
115115
indices_to_read = [0, 1, 2, 3, 4]
@@ -132,7 +132,7 @@ def test_array_record_data_source_reverse_order(self):
132132
]) as ar:
133133
actual_data = [ar[x] for x in indices_to_read]
134134
self.assertEqual(expected_data, actual_data)
135-
self.assertTrue(all(reader is None for reader in ar._readers))
135+
self.assertTrue(all(reader is None for reader in ar._peek_readers()))
136136

137137
def test_array_record_data_source_random_order(self):
138138
# some random permutation
@@ -144,7 +144,7 @@ def test_array_record_data_source_random_order(self):
144144
]) as ar:
145145
actual_data = [ar[x] for x in indices_to_read]
146146
self.assertEqual(expected_data, actual_data)
147-
self.assertTrue(all(reader is None for reader in ar._readers))
147+
self.assertTrue(all(reader is None for reader in ar._peek_readers()))
148148

149149
def test_array_record_data_source_random_order_batched(self):
150150
# some random permutation
@@ -156,7 +156,7 @@ def test_array_record_data_source_random_order_batched(self):
156156
]) as ar:
157157
actual_data = ar.__getitems__(indices_to_read)
158158
self.assertEqual(expected_data, actual_data)
159-
self.assertTrue(all(reader is None for reader in ar._readers))
159+
self.assertTrue(all(reader is None for reader in ar._peek_readers()))
160160

161161
def test_array_record_data_source_file_instructions(self):
162162
file_instruction_one = DummyFileInstruction(
@@ -187,7 +187,7 @@ def test_array_record_data_source_file_instructions(self):
187187
actual_data = [ar[x] for x in indices_to_read]
188188

189189
self.assertEqual(expected_data, actual_data)
190-
self.assertTrue(all(reader is None for reader in ar._readers))
190+
self.assertTrue(all(reader is None for reader in ar._peek_readers()))
191191

192192
def test_array_record_source_reader_idx_and_position(self):
193193
file_instructions = [

0 commit comments

Comments
 (0)