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

Reply via email to