Skip to content

feat: support list.contains for PyArrow, pandas and Dask - #4001

Merged
FBruzzesi merged 12 commits into
narwhals-dev:mainfrom
jonasdedden:upstream/list-contains-backends
Oct 3, 2026
Merged

FBruzzesi merged 12 commits into
narwhals-dev:mainfrom
jonasdedden:upstream/list-contains-backends

Conversation

@jonasdedden

@jonasdedden jonasdedden commented Sep 27, 2026 •

Copy link
Copy Markdown
Contributor

Description

list.contains wasn't implemented for PyArrow and pandas-like backends, and Dask had no list support at all, so the tests xfailed there:

backend before after
PyArrow not implemented supported
pandas, Modin (pyarrow-backed lists) not implemented supported
Dask no list namespace or List cast list.contains, List cast
cuDF not implemented unchanged

pandas, Modin and Dask need pandas >= 2.2 and pyarrow, as for the other list methods.

PyArrow and pandas

The check keeps a running count of matches over the flattened values and reads it off at each list's bounds, so it stays linear without a group-by or sort. pandas reuses it for pyarrow-backed lists. ns per list element, 20M elements:

list length 10 1k 100k 1M
contains(-1) 4.4 3.3 3.3 3.3
contains(None) 4.0 3.0 3.0 2.9

Dask

  • list.contains runs the same check on each partition, so it needs pyarrow-backed lists. The other list methods stay unimplemented.
  • Casting to List maps to pandas' Arrow list dtype, as for pandas.
  • The Dask test constructor now keeps list columns as lists. dd.from_dict turned them into strings like '[2, 2, 3, None, None]' (Dask's convert-string). As a result, two assert_frame_equal tests with nested data now pass on Dask, so their xfails are gone.

What type of PR is this? (check all applicable)

  • 💾 Refactor
  • ✨ Feature
  • 🐛 Bug Fix
  • 🔧 Optimization
  • 📝 Documentation
  • ✅ Test
  • 🐳 Other

AI assistance

  • No AI tools were used for this PR.
  • AI tools were used.

Checklist

  • Code follows style guide (ruff)
  • Tests added
  • Documented the changes (N/A, the API completeness tables are generated)
  • If this is your first PR to narwhals, attach a screenshot of pytest passing locally (not CI):

Tested the list and assert_frame_equal tests on pandas 2.2-3.0, PyArrow 13-25 and Dask 2024.10-2026.8 (older pandas skips, as in CI's old-version jobs), plus the whole suite with the Dask constructor. make test-full-coverage passes locally, apart from PySpark tests that need a JVM. Modin is covered by CI's Modin job.

Counts matches with a running sum over the flat values and reads it off at
each list's bounds, so it stays linear without a group-by or sort. pandas
reuses it for pyarrow-backed lists, like the other list aggregations.
Add a Dask list namespace whose `contains` runs the PyArrow helper per
partition, and map `List` casts to pandas' Arrow list dtype. The Dask test
constructor now keeps list columns as lists, as `from_dict` turned them into
strings.
@jonasdedden
jonasdedden force-pushed the upstream/list-contains-backends branch from dfd2b29 to ba10a96 Compare September 28, 2026 16:11
@jonasdedden

Copy link
Copy Markdown
Contributor Author

With commit b97ffee this PR now should in principle be compatible with #3915 's requirements.

Showing off some pyspark runs:

PySpark test runs

Local runs with JDK 17 (openjdk 17.0.20.1), on commit b97ffee4.

PySpark 4.2.0 (Python 3.14, project env):

JAVA_HOME=/path/to/jdk-17 uv run pytest tests/expr_and_series/list/contains_test.py \
  --runslow --constructors=pyspark -p no:randomly -v
tests/expr_and_series/list/contains_test.py::test_contains_expr[pyspark] PASSED
tests/expr_and_series/list/contains_test.py::test_contains_no_match_with_null_elements_expr[pyspark] PASSED
tests/expr_and_series/list/contains_test.py::test_contains_none_expr[pyspark] PASSED
tests/expr_and_series/list/contains_test.py::test_contains_series[NOTSET] SKIPPED
tests/expr_and_series/list/contains_test.py::test_contains_none_inner_dtypes_expr[pyspark-values0-dtype0] PASSED
tests/expr_and_series/list/contains_test.py::test_contains_none_inner_dtypes_expr[pyspark-values1-dtype1] PASSED
tests/expr_and_series/list/contains_test.py::test_contains_inner_dtypes_expr[pyspark-x-y-dtype0]  PASSED
tests/expr_and_series/list/contains_test.py::test_contains_inner_dtypes_expr[pyspark-True-False-dtype1] PASSED
tests/expr_and_series/list/contains_test.py::test_contains_inner_dtypes_expr[pyspark-x2-y2-dtype2]PASSED
tests/expr_and_series/list/contains_test.py::test_contains_inner_dtypes_expr[pyspark-x3-y3-dtype3]PASSED
tests/expr_and_series/list/contains_test.py::test_contains_inner_dtypes_expr[pyspark-1.5-nan-dtype4] PASSED
tests/expr_and_series/list/contains_test.py::test_contains_chunked_pyarrow PASSED                 
tests/expr_and_series/list/contains_test.py::test_contains_none_single_empty_list_expr[pyspark] PASSED
tests/expr_and_series/list/contains_test.py::test_contains_none_non_orderable_inner_type_pyspark PASSED
======================== 13 passed, 1 skipped in 7.01s =========================

PySpark 3.5.0 (Python 3.11, isolated env):

JAVA_HOME=/path/to/jdk-17 uv run --isolated --python 3.11 --group tests \
  --with 'pyspark==3.5.0' --with 'pandas<2.3' --with 'numpy<2' --with pyarrow --with polars \
  pytest tests/expr_and_series/list/contains_test.py --runslow --constructors=pyspark -p no:randomly -v
tests/expr_and_series/list/contains_test.py::test_contains_expr[pyspark] PASSED
tests/expr_and_series/list/contains_test.py::test_contains_no_match_with_null_elements_expr[pyspark] PASSED
tests/expr_and_series/list/contains_test.py::test_contains_none_expr[pyspark] PASSED
tests/expr_and_series/list/contains_test.py::test_contains_series[NOTSET] SKIPPED
tests/expr_and_series/list/contains_test.py::test_contains_none_inner_dtypes_expr[pyspark-values0-dtype0] PASSED
tests/expr_and_series/list/contains_test.py::test_contains_none_inner_dtypes_expr[pyspark-values1-dtype1] PASSED
tests/expr_and_series/list/contains_test.py::test_contains_inner_dtypes_expr[pyspark-x-y-dtype0] PASSED
tests/expr_and_series/list/contains_test.py::test_contains_inner_dtypes_expr[pyspark-True-False-dtype1] PASSED
tests/expr_and_series/list/contains_test.py::test_contains_inner_dtypes_expr[pyspark-x2-y2-dtype2] PASSED
tests/expr_and_series/list/contains_test.py::test_contains_inner_dtypes_expr[pyspark-x3-y3-dtype3] PASSED
tests/expr_and_series/list/contains_test.py::test_contains_inner_dtypes_expr[pyspark-1.5-nan-dtype4] PASSED
tests/expr_and_series/list/contains_test.py::test_contains_chunked_pyarrow PASSED
tests/expr_and_series/list/contains_test.py::test_contains_none_single_empty_list_expr[pyspark] PASSED
tests/expr_and_series/list/contains_test.py::test_contains_none_non_orderable_inner_type_pyspark PASSED
======================== 13 passed, 1 skipped in 6.18s =========================

@jonasdedden

jonasdedden commented Sep 28, 2026 •

Copy link
Copy Markdown
Contributor Author

Regarding list.contains in pyarrow in general: Of course it would be much nicer if they just could have a nice kernel upstream which does this entirely in C++. For this there seems to be already an issue, although stale: apache/arrow#45167. I try to bring a solution upstream with PR apache/arrow#51611.

I mentioned the slightly complex workaround approach used here in another adjacent issue: apache/arrow#47118 (comment)

@FBruzzesi FBruzzesi left a comment

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Thanks @jonasdedden left a few nitpicks

Comment thread tests/expr_and_series/list/contains_test.py Outdated
Comment thread tests/expr_and_series/list/contains_test.py Outdated
Comment thread tests/conftest.py Outdated
Comment thread src/narwhals/_arrow/utils.py Outdated
Comment thread src/narwhals/_dask/expr_list.py
@FBruzzesi

Copy link
Copy Markdown
Member

Thanks @jonasdedden I am happy with the changes. Only doubt I have is regarding Dask. There is a conversation (#3690) regarding moving that into a dedicated plugin. I really really appreciate the fact that we are able to support list namespace, yet I am not sure if it's worth it for the time being to add new specific features. @MarcoGorelli what are your thoughts here?

@jonasdedden

Copy link
Copy Markdown
Contributor Author

@FBruzzesi I worked a bit further on benchmarking various different ideas, as I believe the current one was still not the optimum.

Have a look at this table:

List length, how often the item matches native PR #4001 (before) flatten + is_in pa_offsets_cumsum pa_narrow_cumsum pa_narrow_blocked np_reduceat np_searchsorted np_cumsum np_hybrid
3 elements, rare (1e-5) 51 ms, 10 MB 219 ms, 613 MB 429 ms, 546 MB 174 ms, 413 MB 150 ms, 113 MB 140 ms, 9 MB 138 ms, 306 MB 48 ms, 156 MB 186 ms, 156 MB 49 ms, 158 MB
3 elements, common (30%) 51 ms, 10 MB 217 ms, 613 MB 2,108 ms, 1,274 MB 183 ms, 413 MB 151 ms, 113 MB 144 ms, 9 MB 222 ms, 306 MB 427 ms, 396 MB 188 ms, 156 MB 157 ms, 16 MB
1,000 elements, rare 20 ms, 6 MB 135 ms, 413 MB 202 ms, 413 MB 126 ms, 413 MB 137 ms, 213 MB 167 ms, 9 MB 27 ms, 57 MB 29 ms, 56 MB 112 ms, 256 MB 29 ms, 56 MB
1,000 elements, common 20 ms, 6 MB 136 ms, 413 MB 520 ms, 535 MB 135 ms, 413 MB 141 ms, 213 MB 151 ms, 9 MB 26 ms, 57 MB 286 ms, 366 MB 110 ms, 256 MB 26 ms, 57 MB
  • "native" corresponds to a native pyarrow compute kernel implemented in GH-33295: [C++][Python] Add list_contains compute function apache/arrow#51611
  • The various pa_* and np_* variants can be found in the details block at the end of this comment
  • If only using pyarrow compute functions is allowed, pa_narrow_blocked seems to have similar performance to the previous implementation of this PR, but its RAM load is 2 orders of magnitude cheaper
  • If numpy methods are allowed in the _arrow dir, there are solutions which get really really close to native performance

=> For now in 2d5a039, I swapped the cumsum for pa_narrow_blocked: same results on all list.contains tests, ~9 MB extra memory instead of 413–613 MB.

Benchmark script with all variants
"""Benchmark `list.contains` implementations for PyArrow lists.

`native` needs `pc.list_contains` (apache/arrow#51611). Each (case, variant) runs in
a fresh process; memory is the peak above the input, Arrow pool plus NumPy.
"""

from __future__ import annotations

import hashlib
import json
import math
import subprocess
import sys
import time
import tracemalloc

import numpy as np
import pyarrow as pa
import pyarrow.compute as pc

N_VALUES = 50_000_000
CASES = [
    ("3 elements, rare (1e-5)", 3,
    ("3 elements, common (30%)", 3, 0.3),
    ("1,000 elements, rare", 1000,
    ("1,000 elements, common", 1000, 0.3),
]
ITEM = 1
REPEATS = 5
BLOCK_VALUES = 1 << 21


def make_input(list_length, match_
    rng = np.random.default_rng(42)
    values = rng.integers(2, 1 <<
    values[rng.random(N_VALUES) < match_rate] = ITEM
    offsets = np.arange(0, N_VALUE.int32)
    return pa.chunked_array([pa.ListArray.from_arrays(offsets, values)])


# -- shared helpers -----------------------------------


def per_chunk(one):
    def run(array, item):
        return pa.chunked_array([one(arr, item) for arr in array.chunks], pa.bool_())

    run.__doc__ = one.__doc__
    return run


def hits_of(values, item):
    """Like Polars, and unlike `pcNaN matches NaN."""
    if item is None:
        hits = pc.is_null(values)
    elif isinstance(item, float) and math.isnan(item):
        hits = pc.is_nan(values)
    else:
        hits = pc.equal(values, pa
    return hits.fill_null(False) if hits.null_count else hits


def values_and_offsets(arr):
    offsets = arr.offsets
    first = offsets[0].as_py()
    values = arr.values.slice(first, offsets[-1].as_py() - first)
    if first:
        offsets = pc.subtract(offsets, pa.scalar(first, offsets.type))
    return values, offsets


def with_list_nulls(arr, found):
    return found if arr.null_count_valid(), found, None)


def unpack(hits):
    bits = np.frombuffer(hits.buff
    return np.unpackbits(bits, bitorder="little", count=hits.offset + len(hits))[
        hits.offset :
    ].view(bool)


# -- variants -------------------------------------------


def native(array, item):
    return pc.list_contains(array,


@per_chunk
def pr4001(arr, item):
    """narwhals#4001 before this commit: cumsum over list lengths and matches."""
    lengths = pc.list_value_length
    ends = pc.cumulative_sum(lengths, skip_nulls=True)
    starts = pc.subtract(ends, lengths)
    hits = hits_of(pc.list_flatten(arr), item)
    hits = pa.concat_arrays([pa.array([False]), hits])
    running = pc.cumulative_sum(hits.cast(lengths.type))
    return pc.greater(running.take(ends), running.take(starts))


@per_chunk
def flatten_is_in(arr, item):
    """apache/arrow#47118: flatten + is_in, mapped back to lists via parent indices."""
    hits = pc.is_in(pc.list_flatten(arr), value_set=pa.array([item]))
    parents = pc.filter(pc.list_parent_indices(arr), hits)
    rows = pa.array(np.arange(len(arr), dtype=np.int64))
    return with_list_nulls(arr, pc.is_in(rows, value_set=pc.unique(parents)))


@per_chunk
def pa_offsets_cumsum(arr, item):
    """pr4001, with list bounds taken from the offsets."""
    values, offsets = values_and_offsets(arr)
    hits = pa.concat_arrays([pa.array([False]), hits_of(values, item)])
    running = pc.cumulative_sum(hits.cast(offsets.type))
    ends, starts = offsets.slice(1), offsets.slice(0, len(arr))
    return with_list_nulls(arr, pc.greater(running.take(ends), running.take(starts)))


def narrow(arr, item):
    """Count matches modulo 2^k, exact within lists shorter than 2^k."""
    values, offsets = values_and_offsets(arr)
    ends, starts = offsets.slice(1), offsets.slice(0, len(arr))
    longest = pc.max(pc.subtract(ends, starts)).as_py() or 0
    count_type = (
        pa.uint8() if longest < 1 << 8
        else pa.uint16() if longest < 1 << 16
        else pa.uint32() if longest < 1 << 32
        else pa.uint64()
    )
    hits = pa.concat_arrays([pa.array([False]), hits_of(values, item)])
    running = pc.cumulative_sum(hits.cast(count_type))
    in_list = pc.subtract(running.take(ends), running.take(starts))
    return with_list_nulls(arr, pc.not_equal(in_list, pa.scalar(0, count_type)))


pa_narrow_cumsum = per_chunk(narrow)


@per_chunk
def pa_narrow_blocked(arr, item):
    """pa_narrow_cumsum over zero-copy slices of about 2M values (this commit)."""
    n = len(arr)
    n_values = arr.offsets[-1].as_py() - arr.offsets[0].as_py()
    step = max(1, n * BLOCK_VALUES // max(1, n_values))
    blocks = [narrow(arr.slice(i, step), item) for i in range(0, n, step)]
    return pa.concat_arrays(blocks) if blocks else pa.array([], pa.bool_())


def found_by_reduceat(hits, offsets):
    starts = offsets[:-1]
    nonempty = offsets[1:] > starts
    found = np.zeros(len(starts), bool)
    if nonempty.any():
        found[nonempty] = np.logical_or.reduceat(unpack(hits), starts[nonempty])
    return found


def found_by_searchsorted(hits, offsets):
    found = np.zeros(len(offsets) - 1, bool)
    positions = np.flatnonzero(unpack(hits))
    found[np.searchsorted(offsets, positions, side="right") - 1] = True
    return found


def found_by_cumsum(hits, offsets):
    longest = int((offsets[1:] - offsets[:-1]).max(initial=0))
    count_type = np.uint8 if longest < 1 << 8 else np.uint16 if longest < 1 << 16 else np.uint32
    running = np.zeros(len(hits) + 1, count_type)
    np.cumsum(unpack(hits), dtype=count_type, out=running[1:])
    return running[offsets[1:]] != running[offsets[:-1]]


def numpy_variant(found_by):
    @per_chunk
    def run(arr, item):
        values, offsets = values_and_offsets(arr)
        found = found_by(hits_of(values, item), offsets.to_numpy())
        return with_list_nulls(arr, pa.array(found))

    return run


np_reduceat = numpy_variant(found_by_reduceat)
np_searchsorted = numpy_variant(found_by_searchsorted)
np_cumsum = numpy_variant(found_by_cumsum)


@per_chunk
def np_hybrid(arr, item):
    """searchsorted for few matches, reduceat for long lists, else pa_narrow_blocked."""
    values, offsets = values_and_offsets(arr)
    hits = hits_of(values, item)
    if hits.true_count * 64 < len(hits):
        found = found_by_searchsorted(hits, offsets.to_numpy())
    elif len(hits) > 64 * len(arr):
        found = found_by_reduceat(hits, offsets.to_numpy())
    else:
        return pa_narrow_blocked(pa.chunked_array([arr]), item).chunk(0)
    return with_list_nulls(arr, pa.array(found))


VARIANTS = {
    "native": native,
    "PR #4001": pr4001,
    "flatten + is_in": flatten_is_in,
    "pa_offsets_cumsum": pa_offsets_cumsum,
    "pa_narrow_cumsum": pa_narrow_cumsum,
    "pa_narrow_blocked": pa_narrow_blocked,
    "np_reduceat": np_reduceat,
    "np_searchsorted": np_searchsorted,
    "np_cumsum": np_cumsum,
    "np_hybrid": np_hybrid,
}
if not hasattr(pc, "list_contains"):
    del VARIANTS["native"]


# -- harness ------------------------------------------------------------------------


def digest(result):
    codes = result.cast(pa.int8()).fill_null(-1).to_numpy()
    return hashlib.sha256(codes.tobytes()).hexdigest()


def run_one(case, name):
    array = make_input(*CASES[case][1:])
    func = VARIANTS[name]
    pool = pa.default_memory_pool()
    baseline = pool.bytes_allocated()
    times = []
    for _ in range(REPEATS):
        start = time.perf_counter()
        func(array, ITEM)
        times.append(time.perf_counter() - start)
    tracemalloc.start()
    result = func(array, ITEM)
    numpy_peak = tracemalloc.get_traced_memory()[1]
    tracemalloc.stop()
    extra = pool.max_memory() - baseline + numpy_peak
    return {"ms": min(times) * 1e3, "mb": extra / 1e6, "digest": digest(result)}


def main():
    if len(sys.argv) == 3:
        print(json.dumps(run_one(int(sys.argv[1]), sys.argv[2])))
        return
    print("| List length, how often the item matches | "
          + " | ".join(f"`{n}`" if "_" in n else n for n in VARIANTS) + " |")
    print("|---|" + "---|" * len(VARIANTS))
    for case, (label, _, _) in enumerate(CASES):
        outs = {
            name: json.loads(subprocess.run(
                [sys.executable, __file__, str(case), name],
                capture_output=True, text=True, check=True).stdout)
            for name in VARIANTS
        }
        assert len({o["digest"] for o in outs.values()}) == 1, "results differ"
        cells = [f"{o['ms']:,.0f} ms, {o['mb']:,.0f} MB" for o in outs.values()]
        print(f"| {label} | " + " | ".join(cells) + " |", flush=True)


if __name__ == "__main__":
    main()

@MarcoGorelli

Copy link
Copy Markdown
Member

thanks both!

i guess the dask code is simple enough that we can port this over if/when we make a plugin

@FBruzzesi

Copy link
Copy Markdown
Member

@jonasdedden I opened a PR against your fork (see jonasdedden#1) to suggest some further possible improvements and some minor variable renaming for (better?) readability. Feel free to push back on that. The perf improvement on pyarrow was fully suggested by Opus 5.5 while reviewing your implementation here.

chore: review follow-ups for `list.contains`
@jonasdedden

Copy link
Copy Markdown
Contributor Author

@FBruzzesi after the merge of your branch a lot of CI failures appeared 😬 Feel free to commit on this branch for fixes, if you want

@FBruzzesi FBruzzesi left a comment

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Thanks @jonasdedden for iterating so much of this one, hopefully pyarrow will eventually implement the feature natively and we can simplify this whole workaround, although I am quite happy with the final result already 🙏🏼

@FBruzzesi
FBruzzesi merged commit 0971017 into narwhals-dev:main Oct 3, 2026
41 checks passed
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

enhancement New feature or request nested data `list`, `struct`, etc

Projects

None yet

Development

Successfully merging this pull request may close these issues.

3 participants