Skip to content
Draft
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
2 changes: 2 additions & 0 deletions src/metaxy/utils/__init__.py
Original file line number Diff line number Diff line change
@@ -1,7 +1,9 @@
"""Utility modules for Metaxy."""

from metaxy.utils.batched_writer import BatchedMetadataWriter
from metaxy.utils.collect_batches import collect_batches

__all__ = [
"BatchedMetadataWriter",
"collect_batches",
]
78 changes: 78 additions & 0 deletions src/metaxy/utils/collect_batches.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,78 @@
from collections.abc import Iterator
from typing import TYPE_CHECKING, Any, cast

import narwhals as nw
import polars as pl
import pyarrow as pa
from narwhals.typing import Frame

if TYPE_CHECKING:
import ibis.expr.types as ir


def collect_batches(
df: Frame,
chunk_size: int | None = None,
**kwargs: Any,
) -> Iterator[nw.DataFrame[Any]]:
"""
Collect batches of data from a DataFrame or LazyFrame.

Uses native batch iteration when available (Polars, Ibis) to avoid
recomputation overhead.

Parameters:
df: The frame to collect batches from.
chunk_size: The size of each batch. If None, collect everything at once.
**kwargs: Additional keyword arguments to pass to the backend-specific collect method.

Yields:
An iterator over the collected batches.
"""
if isinstance(df, nw.DataFrame):
df = df.lazy()

assert isinstance(df, nw.LazyFrame)

if df.implementation == nw.Implementation.POLARS:
yield from _collect_batches_polars(df, chunk_size, **kwargs)
elif df.implementation == nw.Implementation.IBIS:
yield from _collect_batches_ibis(df, chunk_size, **kwargs)
else:
raise NotImplementedError(
f"collect_batches is not supported for {df.implementation}. Supported backends: Polars, Ibis."
)


def _collect_batches_polars(
df: nw.LazyFrame[Any],
chunk_size: int | None,
**kwargs: Any,
) -> Iterator[nw.DataFrame[Any]]:
"""Collect batches using Polars native collect_batches."""
df_polars = cast(pl.LazyFrame, df.to_native())

if chunk_size is None:
yield nw.from_native(df_polars.collect(**kwargs))
else:
for batch in df_polars.collect_batches(chunk_size=chunk_size, **kwargs):
yield nw.from_native(batch)


def _collect_batches_ibis(
df: nw.LazyFrame[Any],
chunk_size: int | None,
**kwargs: Any,
) -> Iterator[nw.DataFrame[Any]]:
"""Collect batches using Ibis native to_pyarrow_batches."""
ibis_table = cast("ir.Table", df.to_native())

if chunk_size is None:
yield df.collect()
else:
batch_reader = ibis_table.to_pyarrow_batches(chunk_size=chunk_size, **kwargs)
for batch in batch_reader:
yield nw.from_native(pa.Table.from_batches([batch]))


__all__ = ["collect_batches"]
108 changes: 108 additions & 0 deletions tests/utils/test_collect_batches.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,108 @@
"""Tests for the collect_batches utility function."""

import ibis
import narwhals as nw
import polars as pl
import pytest

from metaxy.utils import collect_batches


class TestCollectBatchesPolars:
"""Tests for collect_batches with Polars backend."""

def test_eager_no_chunk_size(self):
"""Without chunk_size, yields the whole frame."""
df = pl.DataFrame({"a": [1, 2, 3], "b": [4, 5, 6]})
nw_df = nw.from_native(df)

batches = list(collect_batches(nw_df))

assert len(batches) == 1
assert batches[0].to_polars().equals(df)

def test_lazy_with_chunk_size(self):
"""LazyFrame is split into batches of the requested size."""
df = pl.DataFrame({"a": list(range(10)), "b": list(range(10, 20))})
nw_lf = nw.from_native(df.lazy())

batches = list(collect_batches(nw_lf, chunk_size=3))

assert len(batches) == 4 # 3+3+3+1
# Verify all data is present
combined = pl.concat([b.to_polars() for b in batches])
assert combined.sort("a").equals(df)

def test_lazy_no_chunk_size(self):
"""LazyFrame without chunk_size yields the whole frame."""
df = pl.DataFrame({"a": [1, 2, 3], "b": [4, 5, 6]})
nw_lf = nw.from_native(df.lazy())

batches = list(collect_batches(nw_lf))

assert len(batches) == 1
assert batches[0].to_polars().equals(df)


class TestCollectBatchesIbis:
"""Tests for collect_batches with Ibis (DuckDB) backend."""

@pytest.fixture
def ibis_connection(self, tmp_path):
"""Create an in-memory DuckDB connection via Ibis."""
return ibis.duckdb.connect(tmp_path / "test.duckdb")

def test_no_chunk_size(self, ibis_connection):
"""Without chunk_size, yields the whole frame."""
df = pl.DataFrame({"a": [1, 2, 3], "b": [4, 5, 6]})
ibis_connection.create_table("test_data", df.to_pandas(), overwrite=True)
ibis_table = ibis_connection.table("test_data")
nw_frame = nw.from_native(ibis_table)

batches = list(collect_batches(nw_frame))

assert len(batches) == 1
result = batches[0].to_polars().sort("a")
assert result["a"].to_list() == [1, 2, 3]

def test_with_chunk_size(self, ibis_connection):
"""Ibis backend uses native to_pyarrow_batches for chunking."""
df = pl.DataFrame({"a": list(range(10)), "b": list(range(10, 20))})
ibis_connection.create_table("test_data", df.to_pandas(), overwrite=True)
ibis_table = ibis_connection.table("test_data")
nw_frame = nw.from_native(ibis_table)

batches = list(collect_batches(nw_frame, chunk_size=3))

# Should have all data across batches
combined = pl.concat([b.to_polars() for b in batches])
assert len(combined) == 10
# Verify all values are present
assert set(combined["a"].to_list()) == set(range(10))

def test_preserves_sorted_order(self, ibis_connection):
"""When data is sorted, batches maintain that order."""
df = pl.DataFrame({"id": list(range(10)), "val": list(range(10, 20))})
ibis_connection.create_table("test_data", df.to_pandas(), overwrite=True)
ibis_table = ibis_connection.table("test_data").order_by("id")
nw_frame = nw.from_native(ibis_table)

batches = list(collect_batches(nw_frame, chunk_size=3))

# Combine batches in order and verify sequence is preserved
combined = pl.concat([b.to_polars() for b in batches])
assert combined["id"].to_list() == list(range(10))


class TestCollectBatchesUnsupported:
"""Tests for unsupported backends."""

def test_unsupported_backend_raises(self):
"""Unsupported backends raise NotImplementedError."""
import pandas as pd

df = pd.DataFrame({"a": [1, 2, 3]})
nw_df = nw.from_native(df)

with pytest.raises(NotImplementedError, match="not supported"):
list(collect_batches(nw_df))
Loading