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 08337b6493 [python][torch] Add batch-first streaming datasets (#9365)
08337b6493 is described below
commit 08337b649391cbdf352d6dd6417efa53c14ec52b
Author: XiaoHongbo <[email protected]>
AuthorDate: Thu Aug 27 13:21:17 2026 +0800
[python][torch] Add batch-first streaming datasets (#9365)
---
docs/docs/pypaimon/pytorch.md | 31 +-
.../pypaimon/read/datasource/torch_dataset.py | 213 +++++++++++++-
paimon-python/pypaimon/read/table_read.py | 71 ++++-
paimon-python/pypaimon/tests/torch_read_test.py | 325 +++++++++++++++++++++
4 files changed, 636 insertions(+), 4 deletions(-)
diff --git a/docs/docs/pypaimon/pytorch.md b/docs/docs/pypaimon/pytorch.md
index 6ab0af5173..9488b20fda 100644
--- a/docs/docs/pypaimon/pytorch.md
+++ b/docs/docs/pypaimon/pytorch.md
@@ -57,7 +57,36 @@ for batch_idx, batch_data in enumerate(dataloader):
When the `streaming` parameter is true, it will iteratively read;
when it is false, it will read the full amount of data into memory.
-**`prefetch_concurrency`** (default: 1): When streaming is true, number of
threads used for parallel prefetch within each DataLoader worker. Set to a
value greater than 1 to partition splits across threads and increase read
throughput. Has no effect when streaming is false.
+**`prefetch_concurrency`** (default: 1): In streaming row mode, controls
+reader threads per DataLoader worker. It has no effect in non-streaming mode.
+
+### Batch Streaming
+
+For batch-oriented training, make the streaming dataset yield batches directly:
+
+```python
+dataset = table_read.to_torch(
+ splits,
+ streaming=True,
+ batch_format="torch",
+ batch_size=1024,
+)
+dataloader = DataLoader(dataset, batch_size=None, num_workers=2)
+
+for batch in dataloader:
+ train(batch["features"], batch["label"])
+```
+
+`batch_format="pyarrow"` yields PyArrow `RecordBatch` objects instead;
+`batch_format="torch"` yields dictionaries of tensors. The default Tensor
+converter supports non-null numeric, boolean, and numeric fixed-size-list
+columns. Use `to_tensor_fn` for other types or custom conversion.
+
+Omit `batch_size` to preserve native reader batches. Otherwise, batches are
+combined or sliced to the requested size. Use `DataLoader(batch_size=None)` to
+disable a second batching step. Batch streaming does not support
`shuffle=True`.
+Numeric tensors may share read-only Arrow buffers; clone them before in-place
+mutation. Batch formats currently require `prefetch_concurrency=1`.
## File Format Metadata Cache
diff --git a/paimon-python/pypaimon/read/datasource/torch_dataset.py
b/paimon-python/pypaimon/read/datasource/torch_dataset.py
index 5eb3485ddd..a8c4f4f7ca 100644
--- a/paimon-python/pypaimon/read/datasource/torch_dataset.py
+++ b/paimon-python/pypaimon/read/datasource/torch_dataset.py
@@ -21,11 +21,15 @@ Module to read a Paimon table into PyTorch Dataset.
import queue
import random
import threading
-from typing import Iterator, List
+import warnings
+from typing import Any, Callable, Iterator, List, Optional
+import pyarrow as pa
import torch
from torch.utils.data import Dataset, IterableDataset
+from pypaimon.read.reader.concat_batch_reader import (
+ _MAX_ARROW_OFFSET, _batch_offset_usage)
from pypaimon.read.split import Split
from pypaimon.read.table_read import TableRead
@@ -100,10 +104,48 @@ class _BaseTorchIterDataset(IterableDataset):
row_dict[field_name] = value
return row_dict
+ def _limit_covers_all_splits(self) -> bool:
+ limit = self.table_read.limit
+ if limit is None:
+ return True
+ total_rows = 0
+ for split in self.splits:
+ physical_row_count = getattr(split, "row_count", None)
+ if (
+ isinstance(physical_row_count, bool)
+ or not isinstance(physical_row_count, int)
+ or physical_row_count < 0
+ ):
+ return False
+ row_count = physical_row_count
+ merged_row_count = getattr(split, "merged_row_count", None)
+ if callable(merged_row_count):
+ try:
+ merged_row_count = merged_row_count()
+ except Exception:
+ merged_row_count = None
+ if (
+ not isinstance(merged_row_count, bool)
+ and isinstance(merged_row_count, int)
+ and 0 <= merged_row_count <= physical_row_count
+ ):
+ row_count = merged_row_count
+ total_rows += row_count
+ if total_rows > limit:
+ return False
+ return True
+
def _worker_splits(self, worker_info) -> List[Split]:
if worker_info is None:
return self.splits
+ # DataLoader workers cannot share a limit budget that may truncate.
+ 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)
@@ -230,6 +272,175 @@ class TorchIterDataset(_BaseTorchIterDataset):
t.join(timeout=self._PREFETCH_JOIN_TIMEOUT_SEC)
+def _concat_record_batches(batches: List[pa.RecordBatch]) -> pa.RecordBatch:
+ if len(batches) == 1:
+ return batches[0]
+ return pa.RecordBatch.from_arrays(
+ [
+ pa.concat_arrays([batch.column(i) for batch in batches])
+ for i in range(batches[0].num_columns)
+ ],
+ schema=batches[0].schema,
+ )
+
+
+def _sized_record_batches(
+ batches: Iterator[pa.RecordBatch],
+ batch_size: Optional[int],
+) -> Iterator[pa.RecordBatch]:
+ if batch_size is None:
+ yield from batches
+ return
+
+ pending: List[pa.RecordBatch] = []
+ pending_rows = 0
+ offset_usage = {}
+ for batch in batches:
+ offset = 0
+ while offset < batch.num_rows:
+ take = min(batch_size - pending_rows, batch.num_rows - offset)
+ piece = batch.slice(offset, take)
+ piece_usage = _batch_offset_usage(piece)
+ if pending and any(
+ offset_usage.get(path, 0) + value > _MAX_ARROW_OFFSET
+ for path, value in piece_usage.items()
+ ):
+ yield _concat_record_batches(pending)
+ pending = []
+ pending_rows = 0
+ offset_usage = {}
+ continue
+
+ pending.append(piece)
+ pending_rows += take
+ offset += take
+ for path, value in piece_usage.items():
+ offset_usage[path] = offset_usage.get(path, 0) + value
+ if pending_rows == batch_size or any(
+ value >= _MAX_ARROW_OFFSET for value in offset_usage.values()
+ ):
+ yield _concat_record_batches(pending)
+ pending = []
+ pending_rows = 0
+ offset_usage = {}
+
+ if pending:
+ yield _concat_record_batches(pending)
+
+
+def _default_to_tensor(batch: pa.RecordBatch) -> dict:
+ tensors = {}
+ for name, array in zip(batch.schema.names, batch.columns):
+ if array.null_count:
+ raise ValueError(
+ "Torch tensor conversion does not support null values in "
+ "column %r; provide to_tensor_fn to handle them." % name
+ )
+
+ if pa.types.is_fixed_size_list(array.type):
+ value_type = array.type.value_type
+ if not (
+ pa.types.is_integer(value_type)
+ or pa.types.is_floating(value_type)
+ or pa.types.is_boolean(value_type)
+ ):
+ raise ValueError(
+ "Torch tensor conversion does not support column %r with "
+ "type %s; provide to_tensor_fn." % (name, array.type)
+ )
+ values = array.values.slice(
+ array.offset * array.type.list_size,
+ len(array) * array.type.list_size,
+ )
+ if values.null_count:
+ raise ValueError(
+ "Torch tensor conversion does not support null list values
"
+ "in column %r; provide to_tensor_fn to handle them." % name
+ )
+ numpy_array = values.to_numpy(zero_copy_only=False).reshape(
+ len(array), array.type.list_size
+ )
+ elif (
+ pa.types.is_integer(array.type)
+ or pa.types.is_floating(array.type)
+ or pa.types.is_boolean(array.type)
+ ):
+ numpy_array = array.to_numpy(zero_copy_only=False)
+ else:
+ raise ValueError(
+ "Torch tensor conversion only supports numeric, boolean, and "
+ "fixed-size-list columns; column %r has type %s. Select "
+ "batch_format='pyarrow' or provide to_tensor_fn."
+ % (name, array.type)
+ )
+ with warnings.catch_warnings():
+ warnings.filterwarnings(
+ "ignore",
+ message="The given NumPy array is not writable",
+ category=UserWarning,
+ )
+ tensors[name] = torch.from_numpy(numpy_array)
+ return tensors
+
+
+class TorchBatchIterDataset(_BaseTorchIterDataset):
+ """Streaming IterableDataset which yields Arrow or Tensor batches."""
+
+ def __init__(
+ self,
+ table_read: TableRead,
+ splits: List[Split],
+ batch_format: str,
+ batch_size: Optional[int],
+ to_tensor_fn: Optional[Callable[[pa.RecordBatch], Any]] = None,
+ ):
+ super().__init__(table_read, splits)
+ self.batch_format = batch_format
+ self.batch_size = batch_size
+ self.to_tensor_fn = to_tensor_fn
+
+ def __iter__(self):
+ worker_info = torch.utils.data.get_worker_info()
+ splits_to_process = self._worker_splits(worker_info)
+ raw_batches = self._arrow_batches_for_splits(splits_to_process)
+
+ batches = _sized_record_batches(
+ self._limit_batches(raw_batches), self.batch_size
+ )
+ for batch in batches:
+ if self.batch_format == "torch":
+ converter = self.to_tensor_fn or _default_to_tensor
+ yield converter(batch)
+ else:
+ yield batch
+
+ def _arrow_batches_for_splits(
+ self, splits: List[Split]
+ ) -> Iterator[pa.RecordBatch]:
+ reader = self.table_read.to_arrow_batch_reader(splits)
+ try:
+ for batch in iter(reader.read_next_batch, None):
+ if batch.num_rows:
+ yield batch
+ finally:
+ close = getattr(reader, "close", None)
+ if close is not None:
+ close()
+
+ def _limit_batches(
+ self, batches: Iterator[pa.RecordBatch]
+ ) -> Iterator[pa.RecordBatch]:
+ remaining = self.table_read.limit
+ for batch in batches:
+ if remaining is not None:
+ if remaining <= 0:
+ return
+ if batch.num_rows > remaining:
+ batch = batch.slice(0, remaining)
+ remaining -= batch.num_rows
+ yield batch
+
+
class TorchShuffledIterDataset(_BaseTorchIterDataset):
"""
PyTorch IterableDataset with Paimon-controlled streaming shuffle.
diff --git a/paimon-python/pypaimon/read/table_read.py
b/paimon-python/pypaimon/read/table_read.py
index c2ba44545a..a8fcf92bb3 100644
--- a/paimon-python/pypaimon/read/table_read.py
+++ b/paimon-python/pypaimon/read/table_read.py
@@ -18,7 +18,7 @@
import os
import threading
from concurrent.futures import ThreadPoolExecutor, as_completed
-from typing import Any, Dict, Iterator, List, Optional
+from typing import Any, Callable, Dict, Iterator, List, Optional
import pandas
import pyarrow
@@ -654,12 +654,79 @@ class TableRead:
streaming: bool = False,
prefetch_concurrency: int = 1,
*,
+ batch_format: str = "row",
+ batch_size: Optional[int] = None,
+ to_tensor_fn: Optional[Callable] = None,
shuffle: bool = False,
seed: int = 0,
buffer_size: int = 1000,
max_buffer_input_splits: int = 10,
) -> "torch.utils.data.Dataset":
- """Wrap Paimon table data to PyTorch Dataset."""
+ """Wrap Paimon table data in a PyTorch Dataset.
+
+ Args:
+ splits: Splits to read.
+ streaming: Whether to stream data.
+ prefetch_concurrency: Reader threads per DataLoader worker in row
+ format.
+ batch_format: ``"row"``, ``"pyarrow"``, or ``"torch"``. Batch
+ formats require streaming.
+ 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.
+ """
+ valid_batch_formats = {"row", "pyarrow", "torch"}
+ if batch_format not in valid_batch_formats:
+ raise ValueError(
+ "batch_format must be one of %s, got %r"
+ % (sorted(valid_batch_formats), batch_format)
+ )
+ if batch_size is not None and (
+ isinstance(batch_size, bool)
+ or not isinstance(batch_size, int)
+ or batch_size <= 0
+ ):
+ raise ValueError("batch_size must be a positive int or None")
+ if batch_format == "row":
+ if batch_size is not None:
+ raise ValueError(
+ "batch_size requires batch_format='pyarrow' or 'torch'"
+ )
+ if to_tensor_fn is not None:
+ raise ValueError("to_tensor_fn requires batch_format='torch'")
+ else:
+ if not streaming:
+ raise ValueError(
+ "batch_format=%r requires streaming=True" % batch_format
+ )
+ if shuffle:
+ raise ValueError(
+ "shuffle=True only supports batch_format='row'"
+ )
+ if batch_format == "pyarrow" and to_tensor_fn is not None:
+ raise ValueError("to_tensor_fn requires batch_format='torch'")
+ if to_tensor_fn is not None and not callable(to_tensor_fn):
+ raise ValueError("to_tensor_fn must be callable")
+ if (
+ isinstance(prefetch_concurrency, bool)
+ or not isinstance(prefetch_concurrency, int)
+ or prefetch_concurrency != 1
+ ):
+ raise ValueError(
+ "batch formats require prefetch_concurrency=1"
+ )
+
+ from pypaimon.read.datasource.torch_dataset import (
+ TorchBatchIterDataset,
+ )
+ return TorchBatchIterDataset(
+ self,
+ splits,
+ batch_format=batch_format,
+ batch_size=batch_size,
+ to_tensor_fn=to_tensor_fn,
+ )
+
if shuffle:
if not streaming:
raise ValueError("shuffle=True only supports streaming=True")
diff --git a/paimon-python/pypaimon/tests/torch_read_test.py
b/paimon-python/pypaimon/tests/torch_read_test.py
index 5f55cb2bc8..2b3b126b6c 100644
--- a/paimon-python/pypaimon/tests/torch_read_test.py
+++ b/paimon-python/pypaimon/tests/torch_read_test.py
@@ -19,9 +19,12 @@ import os
import shutil
import tempfile
import unittest
+from types import SimpleNamespace
+from unittest.mock import patch
import pyarrow as pa
from parameterized import parameterized
+import torch
from torch.utils.data import DataLoader
from pypaimon import CatalogFactory, Schema
@@ -143,6 +146,328 @@ class TorchReadTest(unittest.TestCase):
self.assertEqual(sorted_user_ids, expected_user_ids)
self.assertEqual(sorted_behaviors, expected_behaviors)
+ def test_torch_streaming_pyarrow_batches(self):
+ schema = Schema.from_pyarrow_schema(
+ self.pa_schema, partition_keys=['user_id']
+ )
+ self.catalog.create_table(
+ 'default.test_torch_pyarrow_batches', schema, False
+ )
+ table = self.catalog.get_table(
+ 'default.test_torch_pyarrow_batches'
+ )
+ self._write_test_table(table)
+
+ read_builder = table.new_read_builder().with_projection(
+ ['user_id', 'behavior']
+ )
+ splits = read_builder.new_scan().plan().splits()
+ dataset = read_builder.new_read().to_torch(
+ splits,
+ streaming=True,
+ batch_format='pyarrow',
+ batch_size=3,
+ )
+ dataloader = DataLoader(
+ dataset,
+ batch_size=None,
+ num_workers=2,
+ shuffle=False,
+ )
+
+ batches = list(dataloader)
+ self.assertTrue(batches)
+ self.assertTrue(
+ all(isinstance(batch, pa.RecordBatch) for batch in batches)
+ )
+ self.assertTrue(all(0 < batch.num_rows <= 3 for batch in batches))
+ result = pa.Table.from_batches(batches).sort_by('user_id').to_pydict()
+ self.assertEqual(result['user_id'], list(range(1, 9)))
+ self.assertEqual(result['behavior'], list('abcdefgh'))
+
+ def test_torch_streaming_tensor_batches(self):
+ schema = Schema.from_pyarrow_schema(
+ self.pa_schema, partition_keys=['user_id']
+ )
+ self.catalog.create_table(
+ 'default.test_torch_tensor_batches', schema, False
+ )
+ table = self.catalog.get_table(
+ 'default.test_torch_tensor_batches'
+ )
+ self._write_test_table(table)
+
+ read_builder = table.new_read_builder().with_projection(
+ ['user_id', 'item_id']
+ )
+ splits = read_builder.new_scan().plan().splits()
+ dataset = read_builder.new_read().to_torch(
+ splits,
+ streaming=True,
+ batch_format='torch',
+ batch_size=3,
+ )
+
+ batches = list(dataset)
+ self.assertEqual([len(batch['user_id']) for batch in batches], [3, 3,
2])
+ self.assertTrue(
+ all(batch['user_id'].dtype == torch.int32 for batch in batches)
+ )
+ self.assertTrue(
+ all(batch['item_id'].dtype == torch.int64 for batch in batches)
+ )
+ user_ids = torch.cat(
+ [batch['user_id'] for batch in batches]
+ ).sort().values.tolist()
+ self.assertEqual(user_ids, list(range(1, 9)))
+
+ def test_torch_streaming_batches_respect_limit(self):
+ schema = Schema.from_pyarrow_schema(
+ self.pa_schema, partition_keys=['user_id']
+ )
+ self.catalog.create_table(
+ 'default.test_torch_batch_limit', schema, False
+ )
+ table = self.catalog.get_table('default.test_torch_batch_limit')
+ self._write_test_table(table)
+
+ read_builder = table.new_read_builder().with_projection(
+ ['user_id']
+ ).with_limit(5)
+ splits = read_builder.new_scan().plan().splits()
+ dataset = read_builder.new_read().to_torch(
+ splits,
+ streaming=True,
+ batch_format='pyarrow',
+ batch_size=3,
+ )
+ batches = list(dataset)
+ self.assertEqual([batch.num_rows for batch in batches], [3, 2])
+
+ def test_torch_streaming_batches_respect_limit_with_workers(self):
+ schema = Schema.from_pyarrow_schema(
+ self.pa_schema, partition_keys=['user_id']
+ )
+ self.catalog.create_table(
+ 'default.test_torch_batch_worker_limit', schema, False
+ )
+ table = self.catalog.get_table(
+ 'default.test_torch_batch_worker_limit'
+ )
+ self._write_test_table(table)
+
+ predicate = (
+ table.new_read_builder().new_predicate_builder()
+ .greater_than('item_id', 0)
+ )
+ read_builder = (
+ table.new_read_builder()
+ .with_filter(predicate)
+ .with_projection(['user_id'])
+ .with_limit(5)
+ )
+ splits = read_builder.new_scan().plan().splits()
+ self.assertGreater(len(splits), 1)
+ dataset = read_builder.new_read().to_torch(
+ splits,
+ streaming=True,
+ batch_format='pyarrow',
+ batch_size=3,
+ )
+ self.assertEqual(
+ dataset._worker_splits(SimpleNamespace(id=1, num_workers=2)),
+ [],
+ )
+ batches = list(DataLoader(
+ dataset, batch_size=None, num_workers=2
+ ))
+ user_ids = [
+ value
+ for batch in batches
+ for value in batch.column('user_id').to_pylist()
+ ]
+ self.assertEqual(len(user_ids), 5)
+ self.assertEqual(len(set(user_ids)), 5)
+
+ def test_non_binding_limit_preserves_worker_splits(self):
+ schema = Schema.from_pyarrow_schema(
+ self.pa_schema, partition_keys=['user_id']
+ )
+ self.catalog.create_table(
+ 'default.test_torch_non_binding_limit', schema, False
+ )
+ table = self.catalog.get_table(
+ 'default.test_torch_non_binding_limit'
+ )
+ self._write_test_table(table)
+
+ read_builder = table.new_read_builder().with_limit(1000)
+ splits = read_builder.new_scan().plan().splits()
+ self.assertGreater(len(splits), 1)
+ table_read = read_builder.new_read()
+
+ for batch_format in ['row', 'pyarrow']:
+ dataset = table_read.to_torch(
+ splits,
+ streaming=True,
+ batch_format=batch_format,
+ )
+ assigned = [
+ dataset._worker_splits(
+ SimpleNamespace(id=worker_id, num_workers=2)
+ )
+ for worker_id in range(2)
+ ]
+ self.assertTrue(all(assigned))
+ self.assertCountEqual(
+ [id(split) for group in assigned for split in group],
+ [id(split) for split in splits],
+ )
+
+ def test_non_binding_limit_uses_merged_row_counts(self):
+ from pypaimon.read.datasource.torch_dataset import TorchIterDataset
+
+ table_read = SimpleNamespace(limit=8, read_type=[])
+ splits = [
+ SimpleNamespace(row_count=10, merged_row_count=lambda: 4),
+ SimpleNamespace(row_count=10, merged_row_count=lambda: 4),
+ ]
+ dataset = TorchIterDataset(table_read, splits)
+
+ assigned = [
+ dataset._worker_splits(
+ SimpleNamespace(id=worker_id, num_workers=2)
+ )
+ for worker_id in range(2)
+ ]
+ self.assertTrue(all(assigned))
+ self.assertCountEqual(
+ [id(split) for group in assigned for split in group],
+ [id(split) for split in splits],
+ )
+
+ def test_torch_batch_sizing_respects_arrow_offset_limit(self):
+ from pypaimon.read.datasource.torch_dataset import (
+ _sized_record_batches)
+
+ batches = iter([
+ pa.record_batch([pa.array(['aaaa'])], names=['value']),
+ pa.record_batch([pa.array(['bbbb'])], names=['value']),
+ ])
+ with patch(
+ 'pypaimon.read.datasource.torch_dataset._MAX_ARROW_OFFSET', 4
+ ):
+ actual = list(_sized_record_batches(batches, batch_size=2))
+
+ self.assertEqual(
+ [batch.column('value').to_pylist() for batch in actual],
+ [['aaaa'], ['bbbb']],
+ )
+
+ def test_default_tensor_converter_supports_fixed_size_list(self):
+ from pypaimon.read.datasource.torch_dataset import _default_to_tensor
+
+ values = pa.array([1, 2, 3, 4, 5, 6], type=pa.int32())
+ features = pa.FixedSizeListArray.from_arrays(values, 3)
+ batch = pa.RecordBatch.from_arrays([features], ['features'])
+
+ result = _default_to_tensor(batch)
+
+ self.assertEqual(result['features'].dtype, torch.int32)
+ self.assertEqual(result['features'].tolist(), [[1, 2, 3], [4, 5, 6]])
+
+ def test_torch_streaming_custom_tensor_conversion(self):
+ schema = Schema.from_pyarrow_schema(self.pa_schema)
+ self.catalog.create_table(
+ 'default.test_torch_custom_tensor_batch', schema, False
+ )
+ table = self.catalog.get_table(
+ 'default.test_torch_custom_tensor_batch'
+ )
+ self._write_test_table(table)
+
+ read_builder = table.new_read_builder().with_projection(
+ ['user_id', 'behavior']
+ )
+ splits = read_builder.new_scan().plan().splits()
+
+ def to_tensor(batch):
+ return {
+ 'user_id': torch.from_numpy(
+ batch.column('user_id').to_numpy(zero_copy_only=False)
+ ),
+ 'behavior': batch.column('behavior').to_pylist(),
+ }
+
+ dataset = read_builder.new_read().to_torch(
+ splits,
+ streaming=True,
+ batch_format='torch',
+ batch_size=5,
+ to_tensor_fn=to_tensor,
+ )
+ batches = list(dataset)
+ self.assertEqual([len(batch['user_id']) for batch in batches], [5, 3])
+ self.assertEqual(
+ sorted(value for batch in batches for value in batch['behavior']),
+ list('abcdefgh'),
+ )
+
+ default_dataset = read_builder.new_read().to_torch(
+ splits,
+ streaming=True,
+ batch_format='torch',
+ )
+ with self.assertRaisesRegex(ValueError, "batch_format='pyarrow'"):
+ next(iter(default_dataset))
+
+ def test_torch_batch_options_validation(self):
+ schema = Schema.from_pyarrow_schema(self.pa_schema)
+ self.catalog.create_table(
+ 'default.test_torch_batch_validation', schema, False
+ )
+ table = self.catalog.get_table(
+ 'default.test_torch_batch_validation'
+ )
+ 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 self.assertRaisesRegex(ValueError, 'batch_format must be one of'):
+ table_read.to_torch(
+ splits, streaming=True, batch_format='numpy'
+ )
+ with self.assertRaisesRegex(ValueError, 'requires streaming=True'):
+ table_read.to_torch(splits, batch_format='pyarrow')
+ with self.assertRaisesRegex(ValueError, 'batch_size must be'):
+ table_read.to_torch(
+ splits,
+ streaming=True,
+ batch_format='torch',
+ batch_size=0,
+ )
+ with self.assertRaisesRegex(ValueError, 'batch_size requires'):
+ table_read.to_torch(splits, streaming=True, batch_size=2)
+ with self.assertRaisesRegex(ValueError, 'only supports batch_format'):
+ table_read.to_torch(
+ splits,
+ streaming=True,
+ batch_format='torch',
+ shuffle=True,
+ )
+ for invalid in [0, -1, 1.9, True, 2]:
+ with self.subTest(prefetch_concurrency=invalid):
+ with self.assertRaisesRegex(
+ ValueError, 'prefetch_concurrency'
+ ):
+ table_read.to_torch(
+ splits,
+ streaming=True,
+ batch_format='pyarrow',
+ prefetch_concurrency=invalid,
+ )
+
def test_blob_torch_read(self):
"""Test end-to-end blob functionality using blob descriptors."""
import random