Repository navigation
feat: support list.contains for PyArrow, pandas and Dask - #4001
Conversation
d43e3d4 to
9b1d230
Compare
9b1d230 to
dfd2b29
Compare
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.
…lve its peak memory
dfd2b29 to
ba10a96
Compare
|
With commit b97ffee this PR now should in principle be compatible with #3915 's requirements. Showing off some pyspark runs: PySpark test runsLocal runs with JDK 17 ( 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 -vPySpark 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 |
|
Regarding I mentioned the slightly complex workaround approach used here in another adjacent issue: apache/arrow#47118 (comment) |
FBruzzesi
left a comment
There was a problem hiding this comment.
Thanks @jonasdedden left a few nitpicks
|
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? |
|
@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:
=> For now in 2d5a039, I swapped the cumsum for 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() |
|
thanks both! i guess the dask code is simple enough that we can port this over if/when we make a plugin |
|
@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`
|
@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
left a comment
There was a problem hiding this comment.
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 🙏🏼

Description
list.containswasn't implemented for PyArrow and pandas-like backends, and Dask had no list support at all, so the tests xfailed there:listnamespace orListcastlist.contains,Listcastpandas, 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:
contains(-1)contains(None)Dask
list.containsruns the same check on each partition, so it needs pyarrow-backed lists. The other list methods stay unimplemented.Listmaps to pandas' Arrow list dtype, as for pandas.dd.from_dictturned them into strings like'[2, 2, 3, None, None]'(Dask'sconvert-string). As a result, twoassert_frame_equaltests with nested data now pass on Dask, so their xfails are gone.What type of PR is this? (check all applicable)
AI assistance
Checklist
pytestpassing locally (not CI):Tested the list and
assert_frame_equaltests 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-coveragepasses locally, apart from PySpark tests that need a JVM. Modin is covered by CI's Modin job.