Skip to content

fix: align list.contains null handling with Polars on lazy backends - #3990

Merged
FBruzzesi merged 4 commits into
narwhals-dev:mainfrom
jonasdedden:upstream/spark-list-contains-nulls
Sep 26, 2026
Merged

FBruzzesi merged 4 commits into
narwhals-dev:mainfrom
jonasdedden:upstream/spark-list-contains-nulls

Conversation

@jonasdedden

@jonasdedden jonasdedden commented Sep 26, 2026 •

Copy link
Copy Markdown
Contributor

Description

list.contains disagreed with Polars around nulls on the lazy backends:

case Polars before
[1, None], .contains(2) False PySpark: null (Spark's array_contains uses three-valued logic)
[1, None], .contains(None) True DuckDB, Ibis, SQLFrame: null for every row; PySpark: raises AnalysisException
any list, .contains(None) on Polars < 1.24 one value per row a single null

This PR aligns all of them with Polars, using a native check where the backend has one.

A null list still returns null on every backend.

Surfaced in review of #3915, whose numeric coercion cases hit this on PySpark. This should be merged first; #3915 then merges on top without a rebase.

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

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

Related issues

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, aligns PySpark with the other backends)

  • If this is your first PR to narwhals, attach a screenshot of pytest passing locally (not CI):

    PYTEST_ADDOPTS="--numprocesses=logical" \
    make run-ci DEPS="--extra pandas --extra dask --group core-tests --group sklearn --group plugins" \
    CMD="pytest tests --cov=src --cov=tests --runslow --constructors=pandas,pandas[nullable],pandas[pyarrow],pyarrow,polars[eager],polars[lazy],dask,duckdb,sqlframe"

Note: the full-suite command above was not run. Targeted run instead: pytest tests/expr_and_series/list --constructors=pyspark,sqlframe,duckdb,ibis,polars[eager],polars[lazy],pandas,pyarrow gives 167 passed, 19 xfailed (PySpark 4.2.0). The new test fails on pyspark without the fix. PySpark Connect and the 3.5.0 minimum were not run locally; array_compact has been available since Spark 3.4.

@jonasdedden jonasdedden changed the title fix(spark): return \False\ from \list.contains\ on no match when the list holds nulls fix(spark): return False from list.contains on no match when the list holds nulls Sep 26, 2026
@FBruzzesi FBruzzesi added fix pyspark Issue is related to pyspark backend labels Sep 26, 2026

@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.

Hey @jonasdedden thanks for taking care of this. As I saw array_compact I had two questions:

  1. What happens if item=None?
  2. That seem like an expensive op to run, are there going to be performance regressions?

I double checked with Opus 5.5, and it's proposing an alternative workaround keeping array_contains as is and only turning its null into False when the list itself is not null gives the same results at no extra cost:

return F.coalesce(
    F.array_contains(expr, F.lit(item)),
    F.when(expr.isNotNull(), F.lit(False)),
)

which takes care of both questions.

Claude generated benchmark for sqlframe and pyspark

"""Benchmark spark-like `list.contains` variants on PySpark and SQLFrame.

Usage: uv run --with="pandas==2.3" --with="pyspark==4.2.0" --with="sqlframe==4.4.0" --with="duckdb==1.5.5" t.py [pyspark|sqlframe|all]
"""

from __future__ import annotations

import statistics
import sys
import tempfile
import time
from typing import Any, Callable

REPEATS = 10
# (rows, list length); every 5th element is null, and the item (-1) is never found.
CASES = [(5_000_000, 10), (1_000_000, 100)]


def variants(F: Any) -> dict[str, Callable[[Any], Any]]:
    item = F.lit(-1)
    return {
        "main": lambda a: F.array_contains(a, item),
        "array_compact": lambda a: F.array_contains(F.array_compact(a), item),
        "coalesce": lambda a: F.coalesce(
            F.array_contains(a, item), F.when(a.isNotNull(), F.lit(False))
        ),
    }


def timeit(run: Callable[[], object]) -> float:
    run()
    times = []
    for _ in range(REPEATS):
        start = time.perf_counter()
        run()
        times.append(time.perf_counter() - start)
    return statistics.median(times)


def bench_pyspark() -> None:
    from pyspark.sql import SparkSession, functions as F

    spark = (
        SparkSession.builder.master("local[4]")
        .config("spark.ui.enabled", "false")
        .config("spark.driver.memory", "6g")
        .getOrCreate()
    )
    spark.sparkContext.setLogLevel("ERROR")
    with tempfile.TemporaryDirectory() as tmp:
        for rows, length in CASES:
            path = f"{tmp}/data_{length}.parquet"
            spark.range(rows).select(
                F.transform(
                    F.sequence(F.lit(0), F.lit(length - 1)),
                    lambda x: F.when(x % 5 == 0, None).otherwise(x + F.col("id") % 3),
                ).alias("a")
            ).write.parquet(path)
            df = spark.read.parquet(path)
            for name, fn in variants(F).items():
                # A parquet source and a noop sink force every row to be evaluated;
                # `cache()` + `collect()` of an aggregate gave misleading timings.
                q = df.select(fn(F.col("a")).alias("r"))
                counts = {r["r"]: r["count"] for r in q.groupBy("r").count().collect()}
                elapsed = timeit(lambda: q.write.format("noop").mode("overwrite").save())
                print(f"pyspark  rows={rows:>9,} len={length:>3} {name:<14} {elapsed:.3f}s  {counts}")


def bench_sqlframe() -> None:
    import duckdb
    from sqlframe.duckdb import DuckDBSession, functions as F

    con = duckdb.connect()
    for rows, length in CASES:
        con.execute(
            f"CREATE OR REPLACE TABLE t AS SELECT list_transform(range({length}), "
            f"x -> CASE WHEN x % 5 = 0 THEN NULL ELSE x + i % 3 END) AS a "
            f"FROM range({rows}) r(i)"
        )
        df = DuckDBSession(conn=con).table("t")
        for name, fn in variants(F).items():
            q = df.select(F.sum(fn(F.col("a")).cast("int")).alias("s"))
            result = q.collect()[0][0]
            elapsed = timeit(q.collect)
            print(f"sqlframe rows={rows:>9,} len={length:>3} {name:<14} {elapsed:.3f}s  sum={result}")


if __name__ == "__main__":
    target = sys.argv[1] if len(sys.argv) > 1 else "all"
    if target in {"sqlframe", "all"}:
        bench_sqlframe()
    if target in {"pyspark", "all"}:
        bench_pyspark()

Median of 5 runs. Every 5th element is null, and the item is never found (the worst case for the bug):

backend rows list len main array_compact (this PR) coalesce
PySpark 5M 10 0.162s 1.212s 0.135s
PySpark 1M 100 0.178s 2.378s 0.183s
SQLFrame 5M 10 0.008s 0.053s 0.009s
SQLFrame 1M 100 0.015s 0.098s 0.015s

On top of this, it would be nice if you can add a test with item=None, which we don't have yet 🫥

…compact` in Spark

Use `coalesce(array_contains(...), when(list is not null, False))` for Spark
instead of compacting the list first, which was up to ~10x slower. Handle
`item=None` explicitly on Spark-like, DuckDB and Ibis, where native
`contains` returned null (or raised on PySpark) instead of matching Polars.

Co-Authored-By: Claude Opus 5.5 <noreply@anthropic.com>
Claude-Session: https://claude.ai/code/session_01Ba9JftUrcct4rBVY5qsB1C
@jonasdedden

Copy link
Copy Markdown
Contributor Author

@FBruzzesi thank you so much for your investigation!! I had an interactive look with Claude and I think it implemented your suggestion quite well.

@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.

Thank you for checking and resolving the issue. I’ve carried out a similar review, taking performance into account. I believe there are still some potential performance issues - Added inline suggestions to address them. Local runtimes:

backend variant null first null last no nulls
duckdb PR (list_count) 0.034s 0.034s 0.043s
duckdb array_position 0.009s 0.014s 0.011s
pyspark PR (array_compact) 2.562s 2.537s 2.542s
pyspark exists 0.199s 2.480s 2.491s

Comment thread src/narwhals/_duckdb/expr_list.py Outdated
Comment thread src/narwhals/_spark_like/expr_list.py Outdated
Comment thread tests/expr_and_series/list/contains_test.py Outdated
… old Polars

Use `array_position(..., NULL)` on DuckDB>=1.3 and `exists` on PySpark,
which can short-circuit instead of scanning the whole list. Handle
`list.contains(None)` on Polars<1.24, which returned null.

Co-Authored-By: Claude Opus 5.5 <noreply@anthropic.com>
Claude-Session: https://claude.ai/code/session_01Ba9JftUrcct4rBVY5qsB1C
@FBruzzesi FBruzzesi removed the pyspark Issue is related to pyspark backend label Sep 26, 2026
…rk and Ibis

On PySpark, `exists` only short-circuits when a null comes early; with a late
or no null it is as slow as `array_compact` (lambda evaluation). Sorting
(`sort_array` places nulls first) is ~4-7x faster in that case. On Ibis,
filtering for nulls is cheaper than filtering for non-nulls and comparing
lengths. Extend the `item=None` test with duplicate elements.

Co-Authored-By: Claude Opus 5.5 <noreply@anthropic.com>
Claude-Session: https://claude.ai/code/session_01Ba9JftUrcct4rBVY5qsB1C
@FBruzzesi FBruzzesi changed the title fix(spark): return False from list.contains on no match when the list holds nulls fix: align list.contains null handling with Polars on lazy backends Sep 26, 2026

@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 a ton @jonasdedden - I really like where we landed after these iterations, especially the last commit is a really good improvements I didn't know about

@FBruzzesi
FBruzzesi merged commit dbc4cdd into narwhals-dev:main Sep 26, 2026
44 checks passed
@jonasdedden

jonasdedden commented Sep 26, 2026 •

Copy link
Copy Markdown
Contributor Author

Ah, it maybe wasn't the best idea to have this merged already, this was still in progress 😔

I think the last commit was just a hallucination, and mostly based on really stupid test data (basically already sorted data, where Spark's TimSort is returning results way faster).

I'm still running benchmarks, but now it seems:

The results confirm the suspicion, show sorted-data had flattered sort_array, and reveal array_intersect outperforming both. Now testing the remaining case: a null at the very start where exists can exit early.

I.e. the "sort values and test for first one" probably is garbage, performing abysmal on very large data (think of O(n log n) vs. O(n), and highly aggregate computation [distributed sorting] vs. embarrassingly parallel single data passes [exist]).

@jonasdedden

Copy link
Copy Markdown
Contributor Author

I'd probably recommend reverting the last commit and reopening this PR. Will have a look tomorrow

@FBruzzesi

Copy link
Copy Markdown
Member

Thanks for raising this, I should have run my own benchmark 🙈 Reverting now

@FBruzzesi

Copy link
Copy Markdown
Member

@jonasdedden I don't think I have a way to re-open the PR directly. You can open a new one once you think it's ready!

@jonasdedden

Copy link
Copy Markdown
Contributor Author

Reply: (commit is f7b5af4 but unfortunately PR is merged already)

Yes, the sorting fix got worse and worse with list length, and my earlier benchmark hid that. My test lists were already in ascending order, which Spark's sort handles in one pass. On random data it falls behind exists from about 10k elements per list and is 2.7× slower at 1M. I've replaced it with size(array_intersect(a, [NULL])) > 0, which does one pass and runs no lambda. It's the fastest or near-fastest option at every length, and it's pushed to the PR branch as f7b5af42.

Nanoseconds per element, 20M elements in total, PySpark 4.2

Data Variant 10 100 1k 10k 100k 1M
random, null last array_contains (one-pass reference) 11.3 8.6 6.9 6.3 6.6 5.9
exists 53.4 49.3 49.7 45.4 49.8 43.4
sort_array 22.8 29.7 38.2 57.0 78.6 113.0
array_intersect 15.3 9.9 7.9 7.2 7.8 6.2
sorted, no null sort_array 21.5 10.6 10.3 14.5 16.3 13.0

Holding total elements fixed means a one-pass method shows a flat cost per element.

  • Sorting grows with list length: its cost per element rises steadily, matching n log n. On sorted input it stays flat; that's what my earlier benchmark measured.

  • exists does scale linearly, but at about 45–53 ns per element, roughly 7× the cost of array_contains. The cost comes from Spark evaluating its lambda for each element, not from the complexity.

  • array_intersect is linear, runs no lambda, and stays within about 1.5× of array_contains at every length. With random data and no null the picture is the same (8.3 vs 44.1 ns per element at 1M).

  • Null in first position, the best case for exists because it stops early:

    Length exists array_intersect sort_array
    100 13.5 13.3 35.9
    10k 10.2 10.7 62.0
    1M 7.5 11.8 119.1

    Reading the data dominates there, so array_intersect gives up little even in this case.

@jonasdedden

Copy link
Copy Markdown
Contributor Author

Feel free to do things, if you want, but I'll have a actual look earliest tomorrow. Sorry for the hassle!

Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

Projects

None yet

Development

Successfully merging this pull request may close these issues.

3 participants