diff --git a/docs/docs/pypaimon/pytorch.md b/docs/docs/pypaimon/pytorch.md index 9488b20fda1e..806df0f38b00 100644 --- a/docs/docs/pypaimon/pytorch.md +++ b/docs/docs/pypaimon/pytorch.md @@ -60,6 +60,30 @@ when it is false, it will read the full amount of data into memory. **`prefetch_concurrency`** (default: 1): In streaming row mode, controls reader threads per DataLoader worker. It has no effect in non-streaming mode. +### Distributed Sharding + +Streaming reads shard splits across DDP ranks and DataLoader workers: + +```python +dataset = table_read.to_torch( + splits, + streaming=True, + auto_detect_rank=True, +) +dataloader = DataLoader(dataset, batch_size=32, num_workers=2) + +with model.join(): + for batch in dataloader: + train(batch) +``` + +Automatic rank sharding is opt-in. Enable it only when every rank receives the +same ordered, complete splits from one snapshot; leave it disabled for splits +already sharded by the application. +A rank may receive fewer rows because splits have different sizes; `join()` +keeps DDP collectives aligned while preserving every row without duplication. +A limit that may truncate the input is rejected when multiple ranks are active. + ### Batch Streaming For batch-oriented training, make the streaming dataset yield batches directly: diff --git a/paimon-python/pypaimon/read/datasource/torch_dataset.py b/paimon-python/pypaimon/read/datasource/torch_dataset.py index a8c4f4f7cab4..6fd9da97f0c1 100644 --- a/paimon-python/pypaimon/read/datasource/torch_dataset.py +++ b/paimon-python/pypaimon/read/datasource/torch_dataset.py @@ -18,6 +18,7 @@ """ Module to read a Paimon table into PyTorch Dataset. """ +import os import queue import random import threading @@ -40,6 +41,59 @@ def _share_epoch_with_torch_workers(value): return torch.tensor(value, dtype=torch.long).share_memory_() +def _validate_distributed_context(rank: int, world_size: int): + if isinstance(rank, bool) or not isinstance(rank, int): + raise ValueError("rank must be an int") + if isinstance(world_size, bool) or not isinstance(world_size, int): + raise ValueError("world_size must be an int") + if world_size <= 0: + raise ValueError("world_size must be greater than 0") + if rank < 0 or rank >= world_size: + raise ValueError("rank must satisfy 0 <= rank < world_size") + return rank, world_size + + +def _resolve_distributed_context(auto_detect_rank: bool): + if not isinstance(auto_detect_rank, bool): + raise ValueError("auto_detect_rank must be a bool") + if not auto_detect_rank: + return 0, 1 + + distributed = getattr(torch, "distributed", None) + if ( + distributed is not None + and distributed.is_available() + and distributed.is_initialized() + ): + rank = distributed.get_rank() + world_size = distributed.get_world_size() + return _validate_distributed_context(rank, world_size) + + env_rank = os.environ.get("RANK") + env_world_size = os.environ.get("WORLD_SIZE") + if env_rank is not None or env_world_size is not None: + if env_rank is None or env_world_size is None: + raise ValueError( + "RANK and WORLD_SIZE environment variables must be set together" + ) + try: + rank, world_size = int(env_rank), int(env_world_size) + except ValueError: + raise ValueError( + "RANK and WORLD_SIZE environment variables must be integers" + ) + return _validate_distributed_context(rank, world_size) + + return 0, 1 + + +def _balanced_slice(values: List[Any], shard_id: int, shard_count: int): + base_size, remainder = divmod(len(values), shard_count) + start = shard_id * base_size + min(shard_id, remainder) + size = base_size + (1 if shard_id < remainder else 0) + return values[start:start + size] + + class TorchDataset(Dataset): """ PyTorch Dataset implementation for reading Paimon table data. @@ -92,10 +146,41 @@ class _BaseTorchIterDataset(IterableDataset): Shared helpers for streaming PyTorch datasets backed by Paimon splits. """ - def __init__(self, table_read: TableRead, splits: List[Split]): + def __init__( + self, + table_read: TableRead, + splits: List[Split], + auto_detect_rank: bool = False, + ): self.table_read = table_read self.splits = splits self.field_names = [field.name for field in table_read.read_type] + self.auto_detect_rank = auto_detect_rank + self.rank, self.world_size = _resolve_distributed_context(auto_detect_rank) + self._context_pid = os.getpid() + + def __getstate__(self): + state = self.__dict__.copy() + rank, world_size = _resolve_distributed_context(self.auto_detect_rank) + state["rank"] = rank + state["world_size"] = world_size + state["_context_pid"] = os.getpid() + return state + + def _distributed_context(self): + rank, world_size = _resolve_distributed_context( + self.auto_detect_rank + ) + current_pid = os.getpid() + if ( + current_pid != self._context_pid + and world_size == 1 + and self.world_size > 1 + ): + return self.rank, self.world_size + self.rank, self.world_size = rank, world_size + self._context_pid = current_pid + return rank, world_size def _row_to_dict(self, offset_row) -> dict: row_dict = {} @@ -136,30 +221,25 @@ def _limit_covers_all_splits(self) -> bool: return True def _worker_splits(self, worker_info) -> List[Split]: - if worker_info is None: - return self.splits + rank, world_size = self._distributed_context() + worker_id = worker_info.id if worker_info is not None else 0 + num_workers = worker_info.num_workers if worker_info is not None else 1 - # DataLoader workers cannot share a limit budget that may truncate. + if self.table_read.limit == 0: + return [] if ( self.table_read.limit is not None and not self._limit_covers_all_splits() ): - return self.splits if worker_info.id == 0 else [] - - worker_id = worker_info.id - num_workers = worker_info.num_workers - total_splits = len(self.splits) - splits_per_worker = total_splits // num_workers - remainder = total_splits % num_workers - - if worker_id < remainder: - start_idx = worker_id * (splits_per_worker + 1) - end_idx = start_idx + splits_per_worker + 1 - else: - start_idx = worker_id * splits_per_worker + remainder - end_idx = start_idx + splits_per_worker + if world_size > 1: + raise ValueError( + "limit is not supported with distributed Torch sharding" + ) + # A binding limit cannot be shared safely. + return self.splits if worker_id == 0 else [] - return self.splits[start_idx:end_idx] + rank_splits = _balanced_slice(self.splits, rank, world_size) + return _balanced_slice(rank_splits, worker_id, num_workers) class TorchIterDataset(_BaseTorchIterDataset): @@ -179,7 +259,13 @@ class TorchIterDataset(_BaseTorchIterDataset): _PREFETCH_GET_TIMEOUT_SEC = 300.0 _PREFETCH_JOIN_TIMEOUT_SEC = 5.0 - def __init__(self, table_read: TableRead, splits: List[Split], prefetch_concurrency: int = 1): + def __init__( + self, + table_read: TableRead, + splits: List[Split], + prefetch_concurrency: int = 1, + auto_detect_rank: bool = False, + ): """ Initialize TorchIterDataset. @@ -190,7 +276,7 @@ def __init__(self, table_read: TableRead, splits: List[Split], prefetch_concurre this worker (default 1). When > 1, splits are partitioned across threads to increase read throughput. """ - super().__init__(table_read, splits) + super().__init__(table_read, splits, auto_detect_rank) self.prefetch_concurrency = max(1, int(prefetch_concurrency)) def __iter__(self): @@ -393,8 +479,9 @@ def __init__( batch_format: str, batch_size: Optional[int], to_tensor_fn: Optional[Callable[[pa.RecordBatch], Any]] = None, + auto_detect_rank: bool = False, ): - super().__init__(table_read, splits) + super().__init__(table_read, splits, auto_detect_rank) self.batch_format = batch_format self.batch_size = batch_size self.to_tensor_fn = to_tensor_fn @@ -457,8 +544,9 @@ def __init__( seed: int = 0, buffer_size: int = 1000, max_buffer_input_splits: int = 10, + auto_detect_rank: bool = False, ): - super().__init__(table_read, splits) + super().__init__(table_read, splits, auto_detect_rank) self.seed = self._require_int(seed, "seed") self.buffer_size = self._require_positive_int(buffer_size, "buffer_size") self.max_buffer_input_splits = self._require_positive_int( @@ -559,7 +647,11 @@ def _iter_buffer_shuffled_rows( rows: Iterator[dict], worker_id: int, ) -> Iterator[dict]: - rng = random.Random(self.seed + self.epoch * 1000003 + worker_id) + rank, world_size = self._distributed_context() + rng_seed = self.seed + self.epoch * 1000003 + worker_id + if world_size > 1: + rng_seed = "%d:%d" % (rng_seed, rank) + rng = random.Random(rng_seed) buffer = [] for row in rows: if len(buffer) < self.buffer_size: diff --git a/paimon-python/pypaimon/read/table_read.py b/paimon-python/pypaimon/read/table_read.py index a8fcf92bb333..4491b5dd2997 100644 --- a/paimon-python/pypaimon/read/table_read.py +++ b/paimon-python/pypaimon/read/table_read.py @@ -661,6 +661,7 @@ def to_torch( seed: int = 0, buffer_size: int = 1000, max_buffer_input_splits: int = 10, + auto_detect_rank: bool = False, ) -> "torch.utils.data.Dataset": """Wrap Paimon table data in a PyTorch Dataset. @@ -674,6 +675,7 @@ def to_torch( batch_size: Rows per batch; ``None`` preserves reader batches. to_tensor_fn: Optional RecordBatch converter for Torch batches. shuffle: Whether to shuffle rows; supported only in row format. + auto_detect_rank: Whether streaming reads shard by DDP rank. """ valid_batch_formats = {"row", "pyarrow", "torch"} if batch_format not in valid_batch_formats: @@ -725,6 +727,7 @@ def to_torch( batch_format=batch_format, batch_size=batch_size, to_tensor_fn=to_tensor_fn, + auto_detect_rank=auto_detect_rank, ) if shuffle: @@ -739,12 +742,18 @@ def to_torch( seed=seed, buffer_size=buffer_size, max_buffer_input_splits=max_buffer_input_splits, + auto_detect_rank=auto_detect_rank, ) return dataset if streaming: from pypaimon.read.datasource.torch_dataset import TorchIterDataset - dataset = TorchIterDataset(self, splits, prefetch_concurrency) + dataset = TorchIterDataset( + self, + splits, + prefetch_concurrency, + auto_detect_rank=auto_detect_rank, + ) return dataset else: from pypaimon.read.datasource.torch_dataset import TorchDataset diff --git a/paimon-python/pypaimon/tests/torch_distributed_sharding_worker.py b/paimon-python/pypaimon/tests/torch_distributed_sharding_worker.py new file mode 100644 index 000000000000..f4c21623fe91 --- /dev/null +++ b/paimon-python/pypaimon/tests/torch_distributed_sharding_worker.py @@ -0,0 +1,102 @@ +# Licensed to the Apache Software Foundation (ASF) under one +# or more contributor license agreements. See the NOTICE file +# distributed with this work for additional information +# regarding copyright ownership. The ASF licenses this file +# to you under the Apache License, Version 2.0 (the +# "License"); you may not use this file except in compliance +# with the License. You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, +# software distributed under the License is distributed on an +# "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY +# KIND, either express or implied. See the License for the +# specific language governing permissions and limitations +# under the License. + +import json +import os +import sys +from types import SimpleNamespace + +import torch +from torch.utils.data import DataLoader +from torch.nn.parallel import DistributedDataParallel + +from pypaimon.read.datasource.torch_dataset import TorchIterDataset + + +class _OffsetRow: + def __init__(self, values): + self._values = values + + def get_field(self, index): + return self._values[index] + + +class _TableRead: + limit = None + read_type = [ + SimpleNamespace(name="split_id"), + SimpleNamespace(name="rank"), + SimpleNamespace(name="worker"), + ] + + def __init__(self): + self.rank = None + + def to_iterator(self, splits): + worker_info = torch.utils.data.get_worker_info() + worker_id = worker_info.id if worker_info is not None else 0 + for split_id in splits: + yield _OffsetRow([split_id, self.rank, worker_id]) + + +def main(): + output_dir = sys.argv[1] + rank_env = os.environ.pop("RANK") + world_size_env = os.environ.pop("WORLD_SIZE") + table_read = _TableRead() + dataset = TorchIterDataset( + table_read, + list(range(11)), + auto_detect_rank=True, + ) + os.environ["RANK"] = rank_env + os.environ["WORLD_SIZE"] = world_size_env + torch.distributed.init_process_group("gloo") + rank = torch.distributed.get_rank() + try: + table_read.rank = rank + os.environ.pop("RANK") + os.environ.pop("WORLD_SIZE") + loader = DataLoader( + dataset, + batch_size=None, + num_workers=2, + multiprocessing_context="spawn", + ) + model = DistributedDataParallel(torch.nn.Linear(1, 1)) + optimizer = torch.optim.SGD(model.parameters(), lr=0.01) + rows = [] + with model.join(): + for row in loader: + rows.append(row) + value = torch.tensor([[float(row["split_id"])]]) + model(value).sum().backward() + optimizer.step() + optimizer.zero_grad() + with open( + os.path.join(output_dir, "rank-%d.json" % rank), + "w", + encoding="utf-8", + ) as result_file: + json.dump(rows, result_file) + torch.distributed.barrier() + finally: + torch.distributed.destroy_process_group() + + +if __name__ == "__main__": + main() diff --git a/paimon-python/pypaimon/tests/torch_read_test.py b/paimon-python/pypaimon/tests/torch_read_test.py index 2b3b126b6c2f..6e62c975a4f7 100644 --- a/paimon-python/pypaimon/tests/torch_read_test.py +++ b/paimon-python/pypaimon/tests/torch_read_test.py @@ -15,8 +15,13 @@ # specific language governing permissions and limitations # under the License. +import json +import multiprocessing import os +import pickle import shutil +import subprocess +import sys import tempfile import unittest from types import SimpleNamespace @@ -29,9 +34,362 @@ from pypaimon import CatalogFactory, Schema +from pypaimon.read.datasource.torch_dataset import ( + TorchIterDataset, + TorchShuffledIterDataset, + _resolve_distributed_context, +) from pypaimon.table.file_store_table import FileStoreTable +def _collect_spawned_worker_splits(dataset, output): + os.environ.pop("RANK", None) + os.environ.pop("WORLD_SIZE", None) + output.put(dataset._worker_splits(None)) + + +class TorchDistributedShardingTest(unittest.TestCase): + @staticmethod + def _table_read(limit=None): + return SimpleNamespace(limit=limit, read_type=[]) + + @staticmethod + def _worker(worker_id, num_workers): + return SimpleNamespace(id=worker_id, num_workers=num_workers) + + def _dataset( + self, + splits, + limit=None, + dataset_type=TorchIterDataset, + **kwargs + ): + return dataset_type( + self._table_read(limit), + splits, + auto_detect_rank=True, + **kwargs + ) + + def _assignments(self, split_count, world_size, num_workers): + splits = list(range(split_count)) + assignments = {} + for rank in range(world_size): + dataset = self._dataset(splits) + for worker_id in range(num_workers): + with patch( + "pypaimon.read.datasource.torch_dataset." + "_resolve_distributed_context", + return_value=(rank, world_size), + ): + assignments[(rank, worker_id)] = dataset._worker_splits( + self._worker(worker_id, num_workers) + ) + return assignments + + def assertCompleteNonOverlapping(self, assignments, expected): + assigned = [ + split for splits in assignments.values() for split in splits + ] + self.assertCountEqual(assigned, expected) + self.assertEqual(len(assigned), len(set(assigned))) + + @parameterized.expand([ + ("single", 7, 1, 1, [7]), + ("workers", 10, 1, 3, [4, 3, 3]), + ("ranks", 10, 3, 1, [4, 3, 3]), + ("rank_workers", 17, 3, 2, [3, 3, 3, 3, 3, 2]), + ("uneven", 11, 2, 2, [3, 3, 3, 2]), + ("sparse", 3, 2, 3, [1, 1, 0, 1, 0, 0]), + ]) + def test_balanced_assignments( + self, _, split_count, world_size, num_workers, expected_sizes + ): + assignments = self._assignments( + split_count, world_size, num_workers + ) + self.assertCompleteNonOverlapping( + assignments, list(range(split_count)) + ) + self.assertEqual( + [len(splits) for splits in assignments.values()], expected_sizes + ) + + def test_binding_limit_rejects_distributed_sharding(self): + splits = [SimpleNamespace(row_count=10) for _ in range(4)] + dataset = self._dataset(splits, limit=5) + with patch( + "pypaimon.read.datasource.torch_dataset." + "_resolve_distributed_context", + return_value=(0, 2), + ), self.assertRaisesRegex(ValueError, "limit is not supported"): + dataset._worker_splits(None) + + def test_zero_limit_returns_no_splits(self): + dataset = self._dataset([SimpleNamespace(row_count=10)], limit=0) + with patch( + "pypaimon.read.datasource.torch_dataset." + "_resolve_distributed_context", + return_value=(1, 2), + ): + self.assertEqual(dataset._worker_splits(None), []) + + def test_initialized_distributed_context_precedes_environment(self): + with patch.dict( + os.environ, {"RANK": "4", "WORLD_SIZE": "5"}, clear=True + ), patch.object( + torch.distributed, "is_available", return_value=True + ), patch.object( + torch.distributed, "is_initialized", return_value=True + ), patch.object( + torch.distributed, "get_rank", return_value=1 + ), patch.object( + torch.distributed, "get_world_size", return_value=3 + ): + context = _resolve_distributed_context(True) + + self.assertEqual(context, (1, 3)) + + def test_context_is_resolved_after_dataset_construction(self): + dataset = TorchIterDataset( + self._table_read(), list(range(8)), auto_detect_rank=True + ) + with patch.dict( + os.environ, {"RANK": "2", "WORLD_SIZE": "4"}, clear=True + ), patch.object( + torch.distributed, "is_available", return_value=True + ), patch.object( + torch.distributed, "is_initialized", return_value=False + ): + context = _resolve_distributed_context(True) + assigned = dataset._worker_splits(None) + + self.assertEqual(context, (2, 4)) + self.assertEqual(assigned, [4, 5]) + + def test_constructor_context_is_preserved_in_spawned_worker(self): + with patch( + "pypaimon.read.datasource.torch_dataset." + "_resolve_distributed_context", + return_value=(0, 1), + ): + dataset = TorchIterDataset( + self._table_read(), list(range(8)), auto_detect_rank=True + ) + + context = multiprocessing.get_context("spawn") + output = context.Queue() + process = context.Process( + target=_collect_spawned_worker_splits, + args=(dataset, output), + ) + with patch( + "pypaimon.read.datasource.torch_dataset." + "_resolve_distributed_context", + return_value=(1, 2), + ): + process.start() + process.join(30) + if process.is_alive(): + process.terminate() + process.join() + self.assertEqual(process.exitcode, 0) + self.assertEqual(output.get(timeout=5), [4, 5, 6, 7]) + output.close() + + def test_same_process_uses_latest_context(self): + with patch( + "pypaimon.read.datasource.torch_dataset." + "_resolve_distributed_context", + return_value=(1, 2), + ): + dataset = TorchIterDataset( + self._table_read(), list(range(8)), auto_detect_rank=True + ) + + with patch( + "pypaimon.read.datasource.torch_dataset." + "_resolve_distributed_context", + return_value=(0, 1), + ): + self.assertEqual(dataset._worker_splits(None), list(range(8))) + + def test_auto_falls_back_to_single_process(self): + with patch.dict(os.environ, {}, clear=True), patch.object( + torch.distributed, "is_available", return_value=False + ): + context = _resolve_distributed_context(True) + + self.assertEqual(context, (0, 1)) + + def test_disabled_preserves_worker_sharding(self): + splits = list(range(8)) + with patch.dict( + os.environ, {"RANK": "1", "WORLD_SIZE": "2"}, clear=True + ), patch.object( + torch.distributed, "is_available", return_value=True + ), patch.object( + torch.distributed, "is_initialized", return_value=True + ): + dataset = TorchIterDataset( + self._table_read(), splits, auto_detect_rank=False + ) + assigned = dataset._worker_splits(self._worker(1, 2)) + + self.assertEqual(assigned, list(range(4, 8))) + + def test_shuffled_dataset_is_reproducible_and_rank_local(self): + splits = list(range(20)) + datasets = [ + self._dataset( + splits, + dataset_type=TorchShuffledIterDataset, + seed=17, + buffer_size=20, + ) + for _ in range(2) + ] + local_splits = [] + for rank, dataset in enumerate(datasets): + with patch( + "pypaimon.read.datasource.torch_dataset." + "_resolve_distributed_context", + return_value=(rank, 2), + ): + local_splits.append(dataset._worker_splits(None)) + self.assertTrue(set(local_splits[0]).isdisjoint(local_splits[1])) + self.assertCountEqual(local_splits[0] + local_splits[1], splits) + restored = pickle.loads(pickle.dumps(datasets[1])) + self.assertTrue(restored.auto_detect_rank) + + rows = [{"id": value} for value in range(20)] + with patch( + "pypaimon.read.datasource.torch_dataset." + "_resolve_distributed_context", + return_value=(0, 2), + ): + first = list(datasets[0]._iter_buffer_shuffled_rows(iter(rows), 0)) + repeat = list(datasets[0]._iter_buffer_shuffled_rows(iter(rows), 0)) + other_worker = list( + datasets[0]._iter_buffer_shuffled_rows(iter(rows), 1) + ) + with patch( + "pypaimon.read.datasource.torch_dataset." + "_resolve_distributed_context", + return_value=(1, 2), + ): + other_rank = list( + datasets[1]._iter_buffer_shuffled_rows(iter(rows), 0) + ) + self.assertEqual(first, repeat) + self.assertNotEqual(first, other_rank) + self.assertNotEqual(first, other_worker) + + datasets[0].set_epoch(1) + with patch( + "pypaimon.read.datasource.torch_dataset." + "_resolve_distributed_context", + return_value=(0, 2), + ): + next_epoch = list( + datasets[0]._iter_buffer_shuffled_rows(iter(rows), 0) + ) + self.assertNotEqual(first, next_epoch) + + def test_invalid_distributed_context(self): + with self.assertRaisesRegex(ValueError, "auto_detect_rank"): + TorchIterDataset( + self._table_read(), [], auto_detect_rank="auto" + ) + + with patch.dict(os.environ, {"RANK": "one"}, clear=True), patch.object( + torch.distributed, "is_available", return_value=False + ), self.assertRaisesRegex(ValueError, "must be set together"): + _resolve_distributed_context(True) + + with patch.dict( + os.environ, {"RANK": "one", "WORLD_SIZE": "2"}, clear=True + ), patch.object( + torch.distributed, "is_available", return_value=False + ), self.assertRaisesRegex(ValueError, "must be integers"): + _resolve_distributed_context(True) + + for environment, message in [ + ({"RANK": "0", "WORLD_SIZE": "0"}, "greater than 0"), + ({"RANK": "2", "WORLD_SIZE": "2"}, "0 <= rank"), + ]: + with self.subTest(environment=environment), patch.dict( + os.environ, environment, clear=True + ), patch.object( + torch.distributed, "is_available", return_value=False + ), self.assertRaisesRegex(ValueError, message): + _resolve_distributed_context(True) + + @unittest.skipUnless( + torch.distributed.is_available(), "torch.distributed is unavailable" + ) + def test_torchrun_rank_and_worker_sharding(self): + script = os.path.join( + os.path.dirname(__file__), "torch_distributed_sharding_worker.py" + ) + python_root = os.path.abspath( + os.path.join(os.path.dirname(__file__), "..", "..") + ) + with tempfile.TemporaryDirectory() as output_dir: + env = os.environ.copy() + env["PYTHONPATH"] = os.pathsep.join( + filter(None, [python_root, env.get("PYTHONPATH")]) + ) + process = subprocess.run( + [ + sys.executable, + "-m", + "torch.distributed.run", + "--standalone", + "--nproc-per-node=2", + script, + output_dir, + ], + env=env, + stdout=subprocess.PIPE, + stderr=subprocess.PIPE, + text=True, + timeout=180, + ) + self.assertEqual( + process.returncode, + 0, + "torchrun failed:\n%s\n%s" % ( + process.stdout, process.stderr + ), + ) + rows = [] + for rank in range(2): + with open( + os.path.join(output_dir, "rank-%d.json" % rank), + encoding="utf-8", + ) as result_file: + rows.extend(json.load(result_file)) + + split_ids = [row["split_id"] for row in rows] + self.assertCountEqual(split_ids, list(range(11))) + self.assertEqual(len(split_ids), len(set(split_ids))) + assignments = {} + for row in rows: + assignments.setdefault( + (row["rank"], row["worker"]), [] + ).append(row["split_id"]) + self.assertEqual( + {key: sorted(values) for key, values in assignments.items()}, + { + (0, 0): [0, 1, 2], + (0, 1): [3, 4, 5], + (1, 0): [6, 7, 8], + (1, 1): [9, 10], + }, + ) + + class TorchReadTest(unittest.TestCase): @classmethod def setUpClass(cls): @@ -468,6 +826,54 @@ def test_torch_batch_options_validation(self): prefetch_concurrency=invalid, ) + def test_torch_distributed_sharding_public_api(self): + schema = Schema.from_pyarrow_schema( + self.pa_schema, partition_keys=['user_id'] + ) + self.catalog.create_table( + 'default.test_torch_distributed_api', schema, False + ) + table = self.catalog.get_table( + 'default.test_torch_distributed_api' + ) + self._write_test_table(table) + read_builder = table.new_read_builder().with_projection(['user_id']) + splits = read_builder.new_scan().plan().splits() + table_read = read_builder.new_read() + + with patch( + "pypaimon.read.datasource.torch_dataset." + "_resolve_distributed_context", + return_value=(1, 2), + ): + datasets = [ + table_read.to_torch( + splits, + streaming=True, + batch_format=batch_format, + shuffle=batch_format == 'row' and shuffle, + auto_detect_rank=True, + ) + for batch_format, shuffle in [ + ('row', False), + ('row', True), + ('pyarrow', False), + ] + ] + expected = splits[(len(splits) + 1) // 2:] + for dataset in datasets: + self.assertTrue(dataset.auto_detect_rank) + self.assertEqual(dataset._worker_splits(None), expected) + + pre_sharded = splits[::2] + with patch.dict( + os.environ, {"RANK": "1", "WORLD_SIZE": "2"}, clear=True + ): + dataset = table_read.to_torch(pre_sharded, streaming=True) + self.assertFalse(dataset.auto_detect_rank) + self.assertEqual(dataset._worker_splits(None), pre_sharded) + self.assertIsNotNone(table_read.to_torch(splits)) + def test_blob_torch_read(self): """Test end-to-end blob functionality using blob descriptors.""" import random