This is an automated email from the ASF dual-hosted git repository.
JingsongLi pushed a commit to branch master
in repository https://gitbox.apache.org/repos/asf/paimon.git
The following commit(s) were added to refs/heads/master by this push:
new 963528cd45 [python][torch] Support distributed iterable dataset
sharding (#9429)
963528cd45 is described below
commit 963528cd45a183ad166bee14e437b1e8c691c4cc
Author: XiaoHongbo <[email protected]>
AuthorDate: Sun Aug 30 17:25:31 2026 +0800
[python][torch] Support distributed iterable dataset sharding (#9429)
## What changed
- Shard streaming Torch datasets across DDP ranks and DataLoader workers
when `auto_detect_rank=True`.
- Keep rank sharding disabled by default because `to_torch` accepts
caller-planned splits.
- Resolve rank from `torch.distributed`, then torchrun environment
variables, and preserve it in spawned workers.
- Reject binding limits with multiple ranks.
- Keep shuffled reads reproducible per epoch, rank, and worker.
- Document uneven inputs and `DistributedDataParallel.join()`.
## Reference
The API is inspired by
[LanceDataset](https://github.com/lance-format/lance/blob/main/python/python/lance/torch/data.py).
Unlike Lance, this low-level API accepts external splits, so automatic
sharding is opt-in.
## Validation
- 43 Torch tests passed.
- CPU torchrun test executes DDP forward/backward with 2 ranks and 2
spawned DataLoader workers.
- Split assignment is complete and non-overlapping.
- Pre-sharded splits remain unchanged by default.
- flake8 and `git diff --check` passed.
---
docs/docs/pypaimon/pytorch.md | 50 +++
.../pypaimon/read/datasource/torch_dataset.py | 195 +++++++--
paimon-python/pypaimon/read/table_read.py | 27 +-
.../tests/torch_distributed_sharding_worker.py | 102 +++++
paimon-python/pypaimon/tests/torch_read_test.py | 459 +++++++++++++++++++++
5 files changed, 808 insertions(+), 25 deletions(-)
diff --git a/docs/docs/pypaimon/pytorch.md b/docs/docs/pypaimon/pytorch.md
index ac9f6b8d64..8961d76434 100644
--- a/docs/docs/pypaimon/pytorch.md
+++ b/docs/docs/pypaimon/pytorch.md
@@ -60,6 +60,56 @@ 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
+def main():
+ dataset = table_read.to_torch(
+ splits,
+ streaming=True,
+ auto_detect_rank=True,
+ )
+ dataloader = DataLoader(
+ dataset,
+ batch_size=32,
+ num_workers=2,
+ multiprocessing_context="spawn",
+ )
+
+ with model.join():
+ for batch in dataloader:
+ train(batch)
+
+
+if __name__ == "__main__":
+ main()
+```
+
+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.
+Automatic detection uses the default process group. For subgroup DDP, resolve
+the context from the group before creating the DataLoader:
+
+```python
+import torch.distributed as dist
+
+dataset = table_read.to_torch(
+ splits,
+ streaming=True,
+ sharding_rank=dist.get_rank(ddp_group),
+ sharding_world_size=dist.get_world_size(ddp_group),
+)
+```
+
+With multi-worker DDP, use `spawn` (or `forkserver`) and create and iterate the
+DataLoader through an `if __name__ == "__main__":` guarded entry point.
+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 a8c4f4f7ca..daf435887b 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,76 @@ 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,
+ sharding_rank: Optional[int] = None,
+ sharding_world_size: Optional[int] = None,
+):
+ if not isinstance(auto_detect_rank, bool):
+ raise ValueError("auto_detect_rank must be a bool")
+ if sharding_rank is not None or sharding_world_size is not None:
+ if auto_detect_rank:
+ raise ValueError(
+ "explicit sharding context cannot be combined with "
+ "auto_detect_rank=True"
+ )
+ if sharding_rank is None or sharding_world_size is None:
+ raise ValueError(
+ "sharding_rank and sharding_world_size must be set together"
+ )
+ return _validate_distributed_context(
+ sharding_rank, sharding_world_size
+ )
+ 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 +163,55 @@ 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,
+ sharding_rank: Optional[int] = None,
+ sharding_world_size: Optional[int] = None,
+ ):
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.sharding_rank = sharding_rank
+ self.sharding_world_size = sharding_world_size
+ self.rank, self.world_size = _resolve_distributed_context(
+ auto_detect_rank,
+ sharding_rank,
+ sharding_world_size,
+ )
+ self._context_pid = os.getpid()
+
+ def __getstate__(self):
+ state = self.__dict__.copy()
+ rank, world_size = _resolve_distributed_context(
+ self.auto_detect_rank,
+ self.sharding_rank,
+ self.sharding_world_size,
+ )
+ 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,
+ self.sharding_rank,
+ self.sharding_world_size,
+ )
+ 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 +252,25 @@ class _BaseTorchIterDataset(IterableDataset):
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 +290,15 @@ 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,
+ sharding_rank: Optional[int] = None,
+ sharding_world_size: Optional[int] = None,
+ ):
"""
Initialize TorchIterDataset.
@@ -190,7 +309,13 @@ class TorchIterDataset(_BaseTorchIterDataset):
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,
+ sharding_rank,
+ sharding_world_size,
+ )
self.prefetch_concurrency = max(1, int(prefetch_concurrency))
def __iter__(self):
@@ -393,8 +518,17 @@ class TorchBatchIterDataset(_BaseTorchIterDataset):
batch_format: str,
batch_size: Optional[int],
to_tensor_fn: Optional[Callable[[pa.RecordBatch], Any]] = None,
+ auto_detect_rank: bool = False,
+ sharding_rank: Optional[int] = None,
+ sharding_world_size: Optional[int] = None,
):
- super().__init__(table_read, splits)
+ super().__init__(
+ table_read,
+ splits,
+ auto_detect_rank,
+ sharding_rank,
+ sharding_world_size,
+ )
self.batch_format = batch_format
self.batch_size = batch_size
self.to_tensor_fn = to_tensor_fn
@@ -457,8 +591,17 @@ class TorchShuffledIterDataset(_BaseTorchIterDataset):
seed: int = 0,
buffer_size: int = 1000,
max_buffer_input_splits: int = 10,
+ auto_detect_rank: bool = False,
+ sharding_rank: Optional[int] = None,
+ sharding_world_size: Optional[int] = None,
):
- super().__init__(table_read, splits)
+ super().__init__(
+ table_read,
+ splits,
+ auto_detect_rank,
+ sharding_rank,
+ sharding_world_size,
+ )
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 +702,11 @@ class TorchShuffledIterDataset(_BaseTorchIterDataset):
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 a8fcf92bb3..cfe13fd75b 100644
--- a/paimon-python/pypaimon/read/table_read.py
+++ b/paimon-python/pypaimon/read/table_read.py
@@ -661,6 +661,9 @@ class TableRead:
seed: int = 0,
buffer_size: int = 1000,
max_buffer_input_splits: int = 10,
+ auto_detect_rank: bool = False,
+ sharding_rank: Optional[int] = None,
+ sharding_world_size: Optional[int] = None,
) -> "torch.utils.data.Dataset":
"""Wrap Paimon table data in a PyTorch Dataset.
@@ -674,6 +677,9 @@ class TableRead:
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.
+ sharding_rank: Explicit rank in the intended DDP process group.
+ sharding_world_size: Explicit size of that process group.
"""
valid_batch_formats = {"row", "pyarrow", "torch"}
if batch_format not in valid_batch_formats:
@@ -681,6 +687,12 @@ class TableRead:
"batch_format must be one of %s, got %r"
% (sorted(valid_batch_formats), batch_format)
)
+ if (
+ auto_detect_rank
+ or sharding_rank is not None
+ or sharding_world_size is not None
+ ) and not streaming:
+ raise ValueError("distributed sharding requires streaming=True")
if batch_size is not None and (
isinstance(batch_size, bool)
or not isinstance(batch_size, int)
@@ -725,6 +737,9 @@ class TableRead:
batch_format=batch_format,
batch_size=batch_size,
to_tensor_fn=to_tensor_fn,
+ auto_detect_rank=auto_detect_rank,
+ sharding_rank=sharding_rank,
+ sharding_world_size=sharding_world_size,
)
if shuffle:
@@ -739,12 +754,22 @@ class TableRead:
seed=seed,
buffer_size=buffer_size,
max_buffer_input_splits=max_buffer_input_splits,
+ auto_detect_rank=auto_detect_rank,
+ sharding_rank=sharding_rank,
+ sharding_world_size=sharding_world_size,
)
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,
+ sharding_rank=sharding_rank,
+ sharding_world_size=sharding_world_size,
+ )
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 0000000000..f4c21623fe
--- /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 8d0abca713..809fe582ac 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
@@ -30,9 +35,395 @@ from torch.utils.data import DataLoader
from pypaimon import CatalogFactory, Schema
from pypaimon.multimodal.table import MultimodalTable
+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_explicit_context_ignores_global_process_group(self):
+ with 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=3
+ ), patch.object(
+ torch.distributed, "get_world_size", return_value=4
+ ):
+ dataset = TorchIterDataset(
+ self._table_read(),
+ list(range(8)),
+ sharding_rank=1,
+ sharding_world_size=2,
+ )
+ assigned = dataset._worker_splits(None)
+
+ self.assertEqual((dataset.rank, dataset.world_size), (1, 2))
+ self.assertEqual(assigned, [4, 5, 6, 7])
+
+ 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 self.assertRaisesRegex(ValueError, "must be set together"):
+ TorchIterDataset(
+ self._table_read(), [], sharding_rank=0
+ )
+ with self.assertRaisesRegex(ValueError, "cannot be combined"):
+ TorchIterDataset(
+ self._table_read(),
+ [],
+ auto_detect_rank=True,
+ sharding_rank=0,
+ sharding_world_size=1,
+ )
+
+ 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):
@@ -441,6 +832,14 @@ class TorchReadTest(unittest.TestCase):
)
with self.assertRaisesRegex(ValueError, 'requires streaming=True'):
table_read.to_torch(splits, batch_format='pyarrow')
+ with self.assertRaisesRegex(ValueError, 'requires streaming=True'):
+ table_read.to_torch(splits, auto_detect_rank=True)
+ with self.assertRaisesRegex(ValueError, 'requires streaming=True'):
+ table_read.to_torch(
+ splits,
+ sharding_rank=0,
+ sharding_world_size=1,
+ )
with self.assertRaisesRegex(ValueError, 'batch_size must be'):
table_read.to_torch(
splits,
@@ -469,6 +868,66 @@ class TorchReadTest(unittest.TestCase):
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)
+
+ dataset = table_read.to_torch(
+ splits,
+ streaming=True,
+ sharding_rank=1,
+ sharding_world_size=2,
+ )
+ self.assertFalse(dataset.auto_detect_rank)
+ self.assertEqual(
+ dataset._worker_splits(None),
+ splits[(len(splits) + 1) // 2:],
+ )
+ self.assertIsNotNone(table_read.to_torch(splits))
+
def test_blob_torch_read(self):
"""Test end-to-end blob functionality using blob descriptors."""
import random