Skip to content
Merged
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
81 changes: 57 additions & 24 deletions src/narwhals/_dask/expr_list.py
Comment thread
FBruzzesi marked this conversation as resolved.
Original file line number Diff line number Diff line change
Expand Up @@ -2,47 +2,80 @@

from typing import TYPE_CHECKING

import pandas as pd

from narwhals._compliant import LazyExprNamespace
from narwhals._compliant.any_namespace import ListNamespace
from narwhals._pandas_like.series import PandasLikeSeries
from narwhals._pandas_like.utils import is_dtype_pyarrow
from narwhals._utils import not_implemented
from narwhals._utils import Implementation, not_implemented

if TYPE_CHECKING:
from collections.abc import Callable

import dask.dataframe.dask_expr as dx
import pandas as pd

from narwhals._dask.dataframe import Incomplete
from narwhals._dask.expr import DaskExpr
from narwhals._utils import Version
from narwhals.typing import NonNestedLiteral


def _list_contains_partition(partition: pd.Series, item: NonNestedLiteral) -> pd.Series:
from narwhals._arrow.utils import list_contains

array: Incomplete = partition.array
result = pd.arrays.ArrowExtensionArray(list_contains(array._pa_array, item))
return pd.Series(result, index=partition.index, name=partition.name)
def _apply_to_partition(
partition: pd.Series,
version: Version,
function: Callable[[PandasLikeSeries], PandasLikeSeries],
) -> pd.Series:
series = PandasLikeSeries(
partition, implementation=Implementation.PANDAS, version=version
)
return function(series).native


class DaskExprListNamespace(LazyExprNamespace["DaskExpr"], ListNamespace["DaskExpr"]):
def contains(self, item: NonNestedLiteral) -> DaskExpr:
def func(expr: dx.Series) -> dx.Series:
if not is_dtype_pyarrow(expr.dtype):
def _map_partitions(
self, function: Callable[[PandasLikeSeries], PandasLikeSeries]
) -> DaskExpr:
version = self.compliant._version

def call(native: dx.Series) -> dx.Series:
if not is_dtype_pyarrow(native.dtype):
msg = "Only pyarrow-backed lists are supported for Dask."
raise NotImplementedError(msg)
return expr.map_partitions(
_list_contains_partition, item, meta=(expr.name, "bool[pyarrow]")
# An explicit `meta` skips Dask's inference, which would wrap errors (e.g.
# `InvalidOperationError` for a mismatched item) in a `ValueError`.
meta = _apply_to_partition(native._meta, version, function)
return native.map_partitions(
_apply_to_partition, version, function, meta=meta
)

return self.compliant._with_callable(func)
return self.compliant._with_callable(call)

def len(self) -> DaskExpr:
return self._map_partitions(lambda series: series.list.len())

def contains(self, item: NonNestedLiteral) -> DaskExpr:
return self._map_partitions(lambda series: series.list.contains(item))

def get(self, index: int) -> DaskExpr:
return self._map_partitions(lambda series: series.list.get(index))

def min(self) -> DaskExpr:
return self._map_partitions(lambda series: series.list.min())

def max(self) -> DaskExpr:
return self._map_partitions(lambda series: series.list.max())

def mean(self) -> DaskExpr:
return self._map_partitions(lambda series: series.list.mean())

def median(self) -> DaskExpr:
return self._map_partitions(lambda series: series.list.median())

def sum(self) -> DaskExpr:
return self._map_partitions(lambda series: series.list.sum())

def sort(self, *, descending: bool, nulls_last: bool) -> DaskExpr:
return self._map_partitions(
lambda series: series.list.sort(descending=descending, nulls_last=nulls_last)
)

len = not_implemented()
unique = not_implemented()
get = not_implemented()
min = not_implemented()
max = not_implemented()
mean = not_implemented()
median = not_implemented()
sum = not_implemented()
sort = not_implemented()
4 changes: 4 additions & 0 deletions src/narwhals/_pandas_like/series_list.py
Original file line number Diff line number Diff line change
Expand Up @@ -48,6 +48,10 @@ def contains(self, item: NonNestedLiteral) -> PandasLikeSeries:

def get(self, index: int) -> PandasLikeSeries:
result = self.native.list[index]
implementation, backend_version = self.implementation, self.backend_version
if implementation.is_pandas() and backend_version < (3, 0): # pragma: no cover
# `result` is a new object so it's safe to do this inplace.
result.index = self.native.index
result.name = self.native.name
return self.with_native(result)

Expand Down
10 changes: 5 additions & 5 deletions tests/expr_and_series/list/get_test.py
Original file line number Diff line number Diff line change
Expand Up @@ -8,17 +8,17 @@
import narwhals as nw
from tests.utils import PANDAS_VERSION, Constructor, ConstructorEager, assert_equal_data

data = {"a": [[1, 2], [None, 3], [None], None]}
data = {"a": [[1, 2], [None, 3], [None], None, [4]]}


@pytest.mark.parametrize(("index", "expected"), [(0, {"a": [1, None, None, None]})])
@pytest.mark.parametrize(("index", "expected"), [(0, {"a": [1, None, None, None, 4]})])
def test_get_expr(
request: pytest.FixtureRequest, constructor: Constructor, index: int, expected: Any
) -> None:
if any(backend in str(constructor) for backend in ("dask", "cudf")):
if "cudf" in str(constructor):
request.applymarker(pytest.mark.xfail)

if "pandas" in str(constructor):
if any(backend in str(constructor) for backend in ("pandas", "dask")):
if PANDAS_VERSION < (2, 2):
pytest.skip()
pytest.importorskip("pyarrow")
Expand All @@ -30,7 +30,7 @@ def test_get_expr(
assert_equal_data(result, expected)


@pytest.mark.parametrize(("index", "expected"), [(0, {"a": [1, None, None, None]})])
@pytest.mark.parametrize(("index", "expected"), [(0, {"a": [1, None, None, None, 4]})])
def test_get_series(
request: pytest.FixtureRequest,
constructor_eager: ConstructorEager,
Expand Down
4 changes: 2 additions & 2 deletions tests/expr_and_series/list/len_test.py
Original file line number Diff line number Diff line change
Expand Up @@ -10,10 +10,10 @@


def test_len_expr(request: pytest.FixtureRequest, constructor: Constructor) -> None:
if any(backend in str(constructor) for backend in ("dask", "cudf")):
if "cudf" in str(constructor):
request.applymarker(pytest.mark.xfail)

if "pandas" in str(constructor):
if any(backend in str(constructor) for backend in ("pandas", "dask")):
if PANDAS_VERSION < (2, 2):
pytest.skip()
pytest.importorskip("pyarrow")
Expand Down
4 changes: 2 additions & 2 deletions tests/expr_and_series/list/max_test.py
Original file line number Diff line number Diff line change
Expand Up @@ -15,9 +15,9 @@


def test_max_expr(request: pytest.FixtureRequest, constructor: Constructor) -> None:
if any(backend in str(constructor) for backend in ("dask", "cudf")):
if "cudf" in str(constructor):
request.applymarker(pytest.mark.xfail)
if "pandas" in str(constructor):
if any(backend in str(constructor) for backend in ("pandas", "dask")):
if PANDAS_VERSION < (2, 2):
pytest.skip()
pytest.importorskip("pyarrow")
Expand Down
4 changes: 2 additions & 2 deletions tests/expr_and_series/list/mean_test.py
Original file line number Diff line number Diff line change
Expand Up @@ -15,9 +15,9 @@


def test_mean_expr(request: pytest.FixtureRequest, constructor: Constructor) -> None:
if any(backend in str(constructor) for backend in ("dask", "cudf")):
if "cudf" in str(constructor):
request.applymarker(pytest.mark.xfail)
if "pandas" in str(constructor):
if any(backend in str(constructor) for backend in ("pandas", "dask")):
if PANDAS_VERSION < (2, 2):
pytest.skip()
pytest.importorskip("pyarrow")
Expand Down
6 changes: 3 additions & 3 deletions tests/expr_and_series/list/median_test.py
Original file line number Diff line number Diff line change
Expand Up @@ -17,11 +17,11 @@


def test_median_expr(request: pytest.FixtureRequest, constructor: Constructor) -> None:
if any(backend in str(constructor) for backend in ("dask", "cudf")) or (
if "cudf" in str(constructor) or (
"polars" in str(constructor) and POLARS_VERSION < (0, 20, 7)
):
request.applymarker(pytest.mark.xfail)
if "pandas" in str(constructor):
if any(backend in str(constructor) for backend in ("pandas", "dask")):
if PANDAS_VERSION < (2, 2):
pytest.skip()
pytest.importorskip("pyarrow")
Expand All @@ -37,7 +37,7 @@ def test_median_expr(request: pytest.FixtureRequest, constructor: Constructor) -
)
if any(
backend in str(constructor)
for backend in ("pandas", "pyarrow", "pandas[pyarrow]")
for backend in ("pandas", "pyarrow", "pandas[pyarrow]", "dask")
):
# there is a mismatch as pyarrow uses an approximate median
assert_equal_data(result, {"a": expected_pyarrow})
Expand Down
4 changes: 2 additions & 2 deletions tests/expr_and_series/list/min_test.py
Original file line number Diff line number Diff line change
Expand Up @@ -15,9 +15,9 @@


def test_min_expr(request: pytest.FixtureRequest, constructor: Constructor) -> None:
if any(backend in str(constructor) for backend in ("dask", "cudf")):
if "cudf" in str(constructor):
request.applymarker(pytest.mark.xfail)
if "pandas" in str(constructor):
if any(backend in str(constructor) for backend in ("pandas", "dask")):
if PANDAS_VERSION < (2, 2):
pytest.skip()
pytest.importorskip("pyarrow")
Expand Down
4 changes: 2 additions & 2 deletions tests/expr_and_series/list/sort_test.py
Original file line number Diff line number Diff line change
Expand Up @@ -60,14 +60,14 @@ def test_sort_expr_args(
nulls_last: bool, # noqa: FBT001
expected: list[Any],
) -> None:
if any(backend in str(constructor) for backend in ("dask", "cudf")):
if "cudf" in str(constructor):
request.applymarker(pytest.mark.xfail)
if "ibis" in str(constructor) and descending:
# https://github.com/ibis-project/ibis/issues/11735
request.applymarker(pytest.mark.xfail)
if "polars" in str(constructor) and POLARS_VERSION < (0, 20, 5):
pytest.skip()
if "pandas" in str(constructor):
if any(backend in str(constructor) for backend in ("pandas", "dask")):
if PANDAS_VERSION < (2, 2):
pytest.skip()
pytest.importorskip("pyarrow")
Expand Down
4 changes: 2 additions & 2 deletions tests/expr_and_series/list/sum_test.py
Original file line number Diff line number Diff line change
Expand Up @@ -15,9 +15,9 @@


def test_sum_expr(request: pytest.FixtureRequest, constructor: Constructor) -> None:
if any(backend in str(constructor) for backend in ("dask", "cudf")):
if "cudf" in str(constructor):
request.applymarker(pytest.mark.xfail)
if "pandas" in str(constructor):
if any(backend in str(constructor) for backend in ("pandas", "dask")):
if PANDAS_VERSION < (2, 2):
pytest.skip()
pytest.importorskip("pyarrow")
Expand Down
Loading