diff --git a/.github/actions/rust-test/action.yaml b/.github/actions/rust-test/action.yaml index 446b39db967..ab5af0b68ca 100644 --- a/.github/actions/rust-test/action.yaml +++ b/.github/actions/rust-test/action.yaml @@ -23,12 +23,28 @@ runs: steps: # Note: cargo fmt check is now handled by the lint job that gates this workflow + - name: Set up embedded Python for native UDF tests + shell: bash + run: | + apt-get update + apt-get install -y --no-install-recommends python3 python3-dev python3-venv + python3 -m venv /tmp/comet-rust-python + /tmp/comet-rust-python/bin/pip install "pyarrow>=14" cloudpickle + echo "PYO3_PYTHON=/tmp/comet-rust-python/bin/python" >> "$GITHUB_ENV" + echo "PYTHONPATH=$(/tmp/comet-rust-python/bin/python -c 'import sysconfig; print(sysconfig.get_path("purelib"))')" >> "$GITHUB_ENV" + - name: Check Cargo clippy shell: bash run: | cd native cargo clippy --color=never --all-targets --workspace -- -D warnings + - name: Check native Python UDF feature + shell: bash + run: | + cd native + cargo clippy --color=never -p datafusion-comet --all-targets --features python-udf -- -D warnings + - name: Check compilation shell: bash run: | @@ -83,6 +99,13 @@ runs: export LD_LIBRARY_PATH=${JAVA_HOME}/lib/server:${LD_LIBRARY_PATH} RUST_BACKTRACE=1 cargo nextest run + - name: Run native Python UDF tests + shell: bash + run: | + cd native + export LD_LIBRARY_PATH=${JAVA_HOME}/lib/server:${LD_LIBRARY_PATH} + RUST_BACKTRACE=1 cargo nextest run -p datafusion-comet --lib --features python-udf python_udf + # The steps above lint and test the accounting allocator over the system allocator. Lint and # test it over jemalloc too, since that is a different allocator backend and nothing else in CI # builds it. diff --git a/.github/workflows/pr_build_linux.yml b/.github/workflows/pr_build_linux.yml index d6d487eb146..a026d9e513b 100644 --- a/.github/workflows/pr_build_linux.yml +++ b/.github/workflows/pr_build_linux.yml @@ -524,6 +524,7 @@ jobs: org.apache.comet.exec.CometJoinSuite org.apache.comet.exec.CometTypedDatasetSuite org.apache.spark.sql.comet.CometMapInBatchSuite + org.apache.spark.sql.comet.CometArrowPythonUdfSuite org.apache.spark.sql.execution.python.CometArrowPythonRunnerSuite org.apache.comet.CometNativeSuite org.apache.comet.CometConfSuite diff --git a/.github/workflows/pr_build_macos.yml b/.github/workflows/pr_build_macos.yml index bf6a86b8f42..15c1fd68fef 100644 --- a/.github/workflows/pr_build_macos.yml +++ b/.github/workflows/pr_build_macos.yml @@ -228,6 +228,7 @@ jobs: org.apache.comet.exec.CometJoinSuite org.apache.comet.exec.CometTypedDatasetSuite org.apache.spark.sql.comet.CometMapInBatchSuite + org.apache.spark.sql.comet.CometArrowPythonUdfSuite org.apache.spark.sql.execution.python.CometArrowPythonRunnerSuite org.apache.comet.CometNativeSuite org.apache.comet.CometConfSuite diff --git a/.github/workflows/pyarrow_udf_test.yml b/.github/workflows/pyarrow_udf_test.yml index ce3737e0d8a..f29caa052ff 100644 --- a/.github/workflows/pyarrow_udf_test.yml +++ b/.github/workflows/pyarrow_udf_test.yml @@ -44,12 +44,15 @@ jobs: - name: Spark 4.0 maven_profiles: "-Pspark-4.0 -Pscala-2.13" pyspark: "4.0.4" + native_python_udf: false - name: Spark 4.1 maven_profiles: "-Pspark-4.1" pyspark: "4.1.3" + native_python_udf: true - name: Spark 4.2 maven_profiles: "-Pspark-4.2" pyspark: "4.2.0" + native_python_udf: true container: # Pinned to the Debian 12 (bookworm) base so the system `python3` is 3.11. The default # `amd64/rust` image is Debian 13 (trixie) which ships Python 3.13 and no python3.11 apt @@ -77,19 +80,34 @@ jobs: restore-keys: | ${{ runner.os }}-java-maven- - - name: Build Comet (debug, ${{ matrix.name }} / Scala 2.13) - run: | - cd native && cargo build - cd .. && ./mvnw -B install -DskipTests ${{ matrix.maven_profiles }} - - name: Install Python 3.11 and pip run: | apt-get update - apt-get install -y --no-install-recommends python3 python3-venv python3-pip + apt-get install -y --no-install-recommends python3 python3-dev python3-venv python3-pip python3 -m venv /tmp/venv /tmp/venv/bin/pip install --upgrade pip /tmp/venv/bin/pip install "pyspark==${{ matrix.pyspark }}" "pyarrow>=14" pandas pytest + - name: Build Comet (debug, ${{ matrix.name }} / Scala 2.13) + env: + PYO3_PYTHON: /tmp/venv/bin/python + run: | + if [ "${{ matrix.native_python_udf }}" = "true" ]; then + (cd native && cargo build --features python-udf) + else + (cd native && cargo build) + fi + ./mvnw -B install -DskipTests ${{ matrix.maven_profiles }} + + - name: Run native Arrow UDF Scala suite + if: matrix.native_python_udf + env: + PYSPARK_PYTHON: /tmp/venv/bin/python + PYTHONPATH: /tmp/venv/lib/python3.11/site-packages + run: | + ./mvnw -B test ${{ matrix.maven_profiles }} \ + -Dsuites=org.apache.spark.sql.comet.CometArrowPythonUdfSuite + - name: Run PyArrow UDF pytest env: # Spark launches Python workers in a fresh subprocess and looks up `python3` @@ -98,11 +116,17 @@ jobs: # ModuleNotFoundError. PYSPARK_PYTHON: /tmp/venv/bin/python PYSPARK_DRIVER_PYTHON: /tmp/venv/bin/python + # The embedded interpreter does not inherit PySpark worker sys.path setup. + PYTHONPATH: /tmp/venv/lib/python3.11/site-packages run: | /tmp/venv/bin/python -m pytest -v \ spark/src/test/resources/pyspark/test_pyarrow_udf.py /tmp/venv/bin/python -m pytest -v \ spark/src/test/resources/pyspark/test_pyarrow_udf_dictionary_shuffle.py + if [ "${{ matrix.native_python_udf }}" = "true" ]; then + /tmp/venv/bin/python -m pytest -v \ + spark/src/test/resources/pyspark/test_native_arrow_udf_files.py + fi /tmp/venv/bin/python -m pytest -v \ spark/src/test/resources/pyspark/test_pyarrow_udf_fuzz.py diff --git a/dev/ci/compute-changes.py b/dev/ci/compute-changes.py index 4d62020bb83..d5d851dc9ff 100644 --- a/dev/ci/compute-changes.py +++ b/dev/ci/compute-changes.py @@ -136,11 +136,16 @@ ], # A real Python worker against each Spark 4.x Arrow runner. The list is # deliberately narrow: the suite builds Comet three times, once per Spark - # version, and only the map-in-batch wiring can change its verdict. + # version, and covers map-in-batch wiring and the native scalar Arrow UDF. "pyarrow_udf": [ "pom.xml", "common/pom.xml", "native/shuffle/src/spark_unsafe/row.rs", + "native/core/src/execution/python_udf.rs", + "native/core/src/execution/operators/arrow_python_udf.rs", + "spark/src/main/spark-4.1+/org/apache/spark/sql/comet/CometArrowEvalPythonExec.scala", + "spark/src/main/spark-4.1+/org/apache/spark/sql/comet/shims/ShimCometArrowEvalPythonExec.scala", + "spark/src/test/spark-4.1+/org/apache/spark/sql/comet/CometArrowPythonUdfSuite.scala", "spark/pom.xml", "spark/src/main/java/org/apache/comet/vector/**", "spark/src/main/java/org/apache/spark/sql/comet/execution/shuffle/SpillWriter.java", @@ -161,6 +166,7 @@ "spark/src/main/spark-4.x/org/apache/spark/sql/execution/python/CometArrowPythonRunnerBase.scala", "spark/src/test/resources/pyspark/conftest.py", "spark/src/test/resources/pyspark/test_pyarrow_udf.py", + "spark/src/test/resources/pyspark/test_native_arrow_udf_files.py", "spark/src/test/resources/pyspark/test_pyarrow_udf_fuzz.py", "spark/src/test/resources/pyspark/test_pyarrow_udf_dictionary_shuffle.py", "spark/src/test/spark-3.5/org/apache/spark/sql/comet/CometMapInBatchSuite.scala", diff --git a/docs/source/contributor-guide/memory_management.md b/docs/source/contributor-guide/memory_management.md index 2a48baa3f18..ca016d993ec 100644 --- a/docs/source/contributor-guide/memory_management.md +++ b/docs/source/contributor-guide/memory_management.md @@ -439,11 +439,14 @@ diverge for several structural reasons: allocation counters (`native_allocated`, `jemalloc_allocated`) see it. In a default build the C dependencies are libzstd (`zstd-sys`, behind the Parquet `zstd` codec), libhdfs (`hdfs-sys`, pulled in by the default `hdfs-opendal` feature), and the TLS stack used for cloud object stores - (`aws-lc-sys`). Building with the `jemalloc` or `mimalloc` feature adds the allocator itself - (`tikv-jemalloc-sys`, `libmimalloc-sys`). It is worth knowing which dependencies are _not_ C, - because several names suggest otherwise: the other Parquet codecs are pure Rust in this build, - `snap` for Snappy, `lz4_flex` for LZ4 and `zlib-rs` for gzip, as is `libbz2-rs-sys` despite its - name, so those allocations do pass through `GlobalAlloc` and are counted. + (`aws-lc-sys`). With the `python-udf` feature, the embedded Python interpreter and PyArrow also + allocate outside `GlobalAlloc`; their memory is absent from the executor's native `allocated` + figure, and Python working allocations are not reserved in Comet's pool. Building with the + `jemalloc` or `mimalloc` feature adds the allocator itself (`tikv-jemalloc-sys`, + `libmimalloc-sys`). It is worth knowing which dependencies are _not_ C, because several names + suggest otherwise: the other Parquet codecs are pure Rust in this build, `snap` for Snappy, + `lz4_flex` for LZ4 and `zlib-rs` for gzip, as is `libbz2-rs-sys` despite its name, so those + allocations do pass through `GlobalAlloc` and are counted. - **Batches in flight across the FFI boundary.** Reservations stop at the operator that made them. Imported JVM batches are reserved only while a reserving operator holds them, and exported native batches have usually been released by the time the JVM receives them yet stay resident until the diff --git a/docs/source/user-guide/latest/operators.md b/docs/source/user-guide/latest/operators.md index 0155ebb8dae..5cd79614705 100644 --- a/docs/source/user-guide/latest/operators.md +++ b/docs/source/user-guide/latest/operators.md @@ -141,10 +141,11 @@ natively on `BroadcastHashJoinExec` and `ShuffledHashJoinExec`. Existence sort-m ## Python and UDF -| Operator | Status | Notes | -| -------------------------------------------------- | ------ | -------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------- | -| `MapInArrowExec`, `MapInPandasExec` | ⚠️ | Spark 4.0 and later. Experimental, disabled by default (`spark.comet.exec.pyarrowUDF.enabled`). See [PyArrow UDF Acceleration](pyarrow-udfs.md). | -| `ArrowEvalPythonExec`, `FlatMapGroupsInPandasExec` | 🔜 | Scalar `@pandas_udf` ([#5386](https://github.com/apache/datafusion-comet/issues/5386)) and grouped `applyInPandas` ([#5123](https://github.com/apache/datafusion-comet/issues/5123)) fall back to Spark. | +| Operator | Status | Notes | +| -------------------------------------------------------------- | ------ | -------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------- | +| `ArrowEvalPythonExec` for scalar `@arrow_udf` (Spark 4.1+) | ⚠️ | Experimental native PyO3 path, requires the `python-udf` feature and `spark.comet.exec.nativeArrowPythonUDF.enabled`. See [PyArrow UDF Acceleration](pyarrow-udfs.md). | +| `MapInArrowExec`, `MapInPandasExec` | ⚠️ | Spark 4.0 and later. Experimental, disabled by default (`spark.comet.exec.pyarrowUDF.enabled`). See [PyArrow UDF Acceleration](pyarrow-udfs.md). | +| Other `ArrowEvalPythonExec` types, `FlatMapGroupsInPandasExec` | 🔜 | Scalar `@pandas_udf` ([#5386](https://github.com/apache/datafusion-comet/issues/5386)) and grouped `applyInPandas` ([#5123](https://github.com/apache/datafusion-comet/issues/5123)) fall back to Spark. | ## See also diff --git a/docs/source/user-guide/latest/pyarrow-udfs.md b/docs/source/user-guide/latest/pyarrow-udfs.md index d3f405fc909..449e91ac64a 100644 --- a/docs/source/user-guide/latest/pyarrow-udfs.md +++ b/docs/source/user-guide/latest/pyarrow-udfs.md @@ -80,6 +80,84 @@ spark.comet.exec.pyarrowUDF.enabled=true The default is `false` while the feature stabilizes. +### Native scalar Arrow UDFs (Spark 4.1+) + +Scalar `@arrow_udf` in Spark 4.1 and later can run inside Comet's Rust execution pipeline when +the native library is built with the `python-udf` Cargo feature and this separate option is enabled: + +``` +spark.comet.exec.nativeArrowPythonUDF.enabled=true +``` + +Build the native library with a Python interpreter that matches the major and minor version used +by the executors' PySpark workers. That interpreter needs its development headers (for example, +`python3-dev` on Debian and Ubuntu) and a shared `libpython` (`libpython3.x.so` on Linux). A Python +built with pyenv may need `PYTHON_CONFIGURE_OPTS="--enable-shared"` when it is installed. For example: + +```sh +PYO3_PYTHON=/path/to/python make release COMET_FEATURES=python-udf +``` + +Install the matching shared `libpython` on every executor and make it discoverable by the dynamic +linker. On Linux, the JVM loads `libcomet` with local symbols, so Comet promotes the already-loaded +shared `libpython` to the global namespace before importing Python extensions such as PyArrow. A +statically linked Python cannot provide those symbols this way, and importing PyArrow can fail with +an undefined-symbol error. A library built with `python-udf` depends on `libpython` as soon as +`libcomet` is loaded: if an executor cannot find it, **Comet itself fails to load**, even when a +query does not use an Arrow UDF. + +Comet passes each argument as a `pyarrow.Array` through the Arrow C Data Interface, invokes the +pickled Python function with PyO3, and appends the result array to the input batch. It checks the +result length and safely casts it to the declared return type, matching Spark's scalar Arrow UDF +serializer. It splits larger input batches according to +`spark.sql.execution.arrow.maxRecordsPerBatch` and +`spark.sql.execution.arrow.maxBytesPerBatch`. A native worker is created per partition. + +Each partition unpickles its own callable. Imported modules and their global state are shared by +concurrent tasks in the executor's embedded Python interpreter. + +The executor's embedded Python must be able to import `pyspark`, `pyarrow`, and the user's Python +modules. Spark serializes the callable and a PySpark return type with `pyspark.cloudpickle`, so +`pyspark` is required even when the callable itself only uses PyArrow. Build the `python-udf` +feature against the same Python major/minor version used by PySpark workers. Install those packages +into that Python environment and ensure the executor process can find them through the embedded +interpreter's `sys.path` (for example, by setting `PYTHONPATH` before launching the executor). +`PYSPARK_PYTHON` selects the external worker executable; it does not select or configure the +embedded interpreter. The worker-only `pyspark.zip` path is not automatically added to it. The +feature and config are disabled by default. Without either, `ArrowEvalPythonExec` stays on Spark's +normal path. + +The initial native path accepts scalar `@arrow_udf` calls with regular or named arguments and +multiple independent UDFs in one `ArrowEvalPythonExec`. Chained Python UDFs, broadcast variables, +Python includes, per-function environment overrides other than Spark's default +`PYTHONHASHSEED=0`, and `spark.sql.execution.arrow.useLargeVarTypes=true` stay on Spark's path. +UDFs that capture a PySpark accumulator also stay on Spark's worker path so their task updates +reach the driver. +Queries also stay on Spark's Python worker path when the Spark context has files added through +`addPyFile` or `addFile`, because the embedded interpreter does not receive Spark's per-task file +setup. +The embedded interpreter starts with the same default hash seed as Spark's Python workers. +Iterator Arrow UDFs, +ordinary `udf(..., useArrow=True)`, scalar pandas UDFs, and `mapInArrow` are separate execution +types; `mapInArrow` retains the columnar runner described above. + +The native path accepts boolean, byte, short, integer, long, float, double, plain string, binary, +decimal, date, and timestamp without time zone. Other input or result types and an +enabled `spark.sql.pyspark.udf.profiler` stay on Spark's path. Spark labels `TimestampType` with the +session time zone, while Comet uses UTC; nested Arrow field names can also differ. This allow-list +keeps types with unverified Arrow schemas on Spark's path. `TimeType` also stays on Spark's path: +Spark 4.2's Arrow UDF row converter rejects it even though PySpark can describe its Arrow type. + +The embedded interpreter is shared by tasks. Pure Python code contends on its global interpreter +lock, so multiple partitions may be slower than Spark's separate Python workers; PyArrow kernels +that release the lock can still run concurrently. Python execution stays synchronous on JVM input +paths and hands off other async tasks when it runs on a Tokio worker. `pyspark.TaskContext.get()` +returns `None` inside a native UDF. A native extension crash or `os._exit` terminates the executor +process. The embedded Python interpreter and PyArrow allocate outside Comet's memory pool. Those +allocations are also absent from the executor's `Comet native memory usage: allocated` figure and +are not limited by `spark.executor.pyspark.memory`. Budget them in executor memory overhead in +addition to the [memory log estimate](tuning/memory.md#sizing-the-overhead-from-the-memory-usage-log). + ### Relationship to Spark's PySpark Arrow conversion conf `spark.comet.exec.pyarrowUDF.enabled` is **not** the same as PySpark's @@ -91,12 +169,14 @@ worker. Both confs can be set independently. ## Supported APIs -| PySpark API | Spark Plan Node | Supported | -| -------------------------------- | --------------------------- | --------- | -| `df.mapInArrow(func, schema)` | `MapInArrowExec` | Yes | -| `df.mapInPandas(func, schema)` | `MapInPandasExec` | Yes | -| `@pandas_udf` (scalar) | `ArrowEvalPythonExec` | Not yet | -| `df.applyInPandas(func, schema)` | `FlatMapGroupsInPandasExec` | Not yet | +| PySpark API | Spark Plan Node | Supported | +| -------------------------------- | --------------------------- | ------------------------ | +| `df.mapInArrow(func, schema)` | `MapInArrowExec` | Yes | +| `df.mapInPandas(func, schema)` | `MapInPandasExec` | Yes | +| scalar `@arrow_udf` (Spark 4.1+) | `ArrowEvalPythonExec` | Experimental native path | +| `udf(..., useArrow=True)` | `ArrowEvalPythonExec` | Not yet | +| `@pandas_udf` (scalar) | `ArrowEvalPythonExec` | Not yet | +| `df.applyInPandas(func, schema)` | `FlatMapGroupsInPandasExec` | Not yet | ## Example @@ -177,8 +257,9 @@ on the unoptimized path. ## Limitations -- The optimization currently applies only to `mapInArrow` and `mapInPandas`. Scalar pandas UDFs - (`@pandas_udf`) and grouped operations (`applyInPandas`) are not yet supported. +- The columnar Python runner applies to `mapInArrow` and `mapInPandas`. The separate native path + applies to scalar `@arrow_udf` on Spark 4.1+. Scalar pandas UDFs (`@pandas_udf`) and grouped + operations (`applyInPandas`) are not yet supported. - The optimization requires Arrow data on the input side. If a shuffle sits between the upstream Comet operator and the Python UDF, use Comet's columnar shuffle for the optimization to apply. Both the `jvm` and `native` shuffle modes can feed `CometMapInBatch`. Set diff --git a/docs/source/user-guide/latest/tuning/memory.md b/docs/source/user-guide/latest/tuning/memory.md index 1f6ea98a2be..2dbfad509f8 100644 --- a/docs/source/user-guide/latest/tuning/memory.md +++ b/docs/source/user-guide/latest/tuning/memory.md @@ -142,13 +142,15 @@ one line every 10 seconds for the whole executor: Comet native memory usage: allocated 5412.3 MiB, reserved 3890.0 MiB (16 native plans, 8 memory pools); JVM Arrow allocated 310.4 MiB, 96.2 MiB of it imported from native ``` -- `allocated` is the memory that Comet's native code has allocated and not yet freed, whether or not - a pool tracks it. -- `reserved` is the part that Comet's memory pools have reserved from Spark's off-heap memory. It is - charged against `spark.memory.offHeap.size`, so the container already has room for it. A pool - sometimes has to track memory that Spark could not grant, such as a spilled batch read back from - disk while the off-heap memory is full. `reserved` leaves that memory out, since nothing charges it - against `spark.memory.offHeap.size`. +- `allocated` is the memory allocated through Comet's Rust global allocator and not yet freed, + whether or not a pool tracks it. It excludes allocations made by native libraries outside that + allocator, including the embedded Python interpreter and PyArrow when native Arrow UDFs are + enabled. +- `reserved` is the part that Comet's memory pools track. It is charged against + `spark.memory.offHeap.size`, so the container already has room for it. A pool sometimes has to + track memory that Spark could not grant, such as a spilled batch read back from disk while the + off-heap memory is full. `reserved` leaves that memory out, since nothing charges it against + `spark.memory.offHeap.size`. - `JVM Arrow allocated` is the Arrow memory Comet holds on the JVM side, such as batches read from Comet's in-memory cache, broadcast data, and batches exchanged with native code or Python workers. The part imported from native was allocated by Comet's native code, so `allocated` already counts @@ -174,7 +176,11 @@ alongside the JVM's own non-heap memory. To size the overhead from it: non-heap memory, and add the most untracked memory seen on any executor. 3. Add a margin on top. The log can miss the true peak between samples, and none of the figures includes the allocator's fragmentation and retained pages, or memory allocated by native C - libraries such as zstd. + libraries such as zstd. With native Arrow UDFs, budget the embedded Python interpreter and + PyArrow allocations in addition to the untracked memory calculated from the log. A PyArrow + result buffer later imported into the JVM can appear in the `imported from native` figure that + the estimate subtracts, even though Rust's `allocated` does not count it. Use the executor's + peak resident memory to size this additional overhead. For example, a 16 GiB executor derives an overhead of 1638 MiB. If the line above has the most untracked memory in its log, that is 5412.3 + (310.4 - 96.2) - 3890.0 = 1736.5 MiB. The overhead diff --git a/native/Cargo.lock b/native/Cargo.lock index e777e80d00c..d1c41b85d7b 100644 --- a/native/Cargo.lock +++ b/native/Cargo.lock @@ -2019,6 +2019,7 @@ dependencies = [ "itertools 0.15.0", "jni 0.22.4", "lazy_static", + "libc", "libloading", "log", "log4rs", @@ -2034,6 +2035,7 @@ dependencies = [ "pprof", "procfs", "prost", + "pyo3", "rand 0.10.2", "reqsign-aws-v4", "reqsign-core", @@ -5370,6 +5372,64 @@ dependencies = [ "prost", ] +[[package]] +name = "pyo3" +version = "0.28.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "91fd8e38a3b50ed1167fb981cd6fd60147e091784c427b8f7183a7ee32c31c12" +dependencies = [ + "libc", + "once_cell", + "portable-atomic", + "pyo3-build-config", + "pyo3-ffi", + "pyo3-macros", +] + +[[package]] +name = "pyo3-build-config" +version = "0.28.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "e368e7ddfdeb98c9bca7f8383be1648fd84ab466bf2bc015e94008db6d35611e" +dependencies = [ + "target-lexicon", +] + +[[package]] +name = "pyo3-ffi" +version = "0.28.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "7f29e10af80b1f7ccaf7f69eace800a03ecd13e883acfacc1e5d0988605f651e" +dependencies = [ + "libc", + "pyo3-build-config", +] + +[[package]] +name = "pyo3-macros" +version = "0.28.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "df6e520eff47c45997d2fc7dd8214b25dd1310918bbb2642156ef66a67f29813" +dependencies = [ + "proc-macro2", + "pyo3-macros-backend", + "quote", + "syn 2.0.119", +] + +[[package]] +name = "pyo3-macros-backend" +version = "0.28.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "c4cdc218d835738f81c2338f822078af45b4afdf8b2e33cbb5916f108b813acb" +dependencies = [ + "heck", + "proc-macro2", + "pyo3-build-config", + "quote", + "syn 2.0.119", +] + [[package]] name = "quad-rand" version = "0.2.3" @@ -6579,6 +6639,12 @@ version = "0.2.0" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "7b2093cf4c8eb1e67749a6762251bc9cd836b6fc171623bd0a9d324d37af2417" +[[package]] +name = "target-lexicon" +version = "0.13.5" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "adb6935a6f5c20170eeceb1a3835a49e12e19d792f6dd344ccc76a985ca5a6ca" + [[package]] name = "tempfile" version = "3.27.0" diff --git a/native/core/Cargo.toml b/native/core/Cargo.toml index 2de67226df0..45d4779089b 100644 --- a/native/core/Cargo.toml +++ b/native/core/Cargo.toml @@ -36,6 +36,8 @@ publish = false [dependencies] arrow = { workspace = true } +pyo3 = { version = "0.28.3", optional = true, features = ["auto-initialize"] } +libc = { version = "0.2", optional = true } base64 = "0.23.0" chrono-tz = { workspace = true } bytes = { workspace = true } @@ -123,6 +125,7 @@ reqsign-aws-v4 = "3" [features] backtrace = ["datafusion/backtrace"] default = ["hdfs-opendal"] +python-udf = ["dep:pyo3", "dep:libc"] contrib-lance = ["dep:comet-contrib-lance"] hdfs-opendal = ["opendal/services-hdfs", "object_store_opendal", "hdfs-sys"] jemalloc = ["tikv-jemallocator", "tikv-jemalloc-ctl"] diff --git a/native/core/src/execution/jni_api.rs b/native/core/src/execution/jni_api.rs index 361d3847da7..4e13b364cd1 100644 --- a/native/core/src/execution/jni_api.rs +++ b/native/core/src/execution/jni_api.rs @@ -490,6 +490,7 @@ fn op_name(op: &OpStruct) -> &'static str { OpStruct::RangeScan(_) => "RangeScan", OpStruct::ContribScan(_) => "ContribScan", OpStruct::WindowGroupLimit(_) => "WindowGroupLimit", + OpStruct::ArrowPythonUdf(_) => "ArrowPythonUdf", } } diff --git a/native/core/src/execution/mod.rs b/native/core/src/execution/mod.rs index 9a2e54783b4..12019ab8a47 100644 --- a/native/core/src/execution/mod.rs +++ b/native/core/src/execution/mod.rs @@ -24,6 +24,8 @@ pub(crate) mod merge_as_partial; pub(crate) mod metrics; pub mod operators; pub(crate) mod planner; +#[cfg(feature = "python-udf")] +pub mod python_udf; pub mod serde; pub use datafusion_comet_shuffle as shuffle; mod memory_pools; diff --git a/native/core/src/execution/operators/arrow_python_udf.rs b/native/core/src/execution/operators/arrow_python_udf.rs new file mode 100644 index 00000000000..77f077b5148 --- /dev/null +++ b/native/core/src/execution/operators/arrow_python_udf.rs @@ -0,0 +1,441 @@ +// Licensed to the Apache Software Foundation (ASF) under one +// or more contributor license agreements. See the NOTICE file +// distributed with this work for additional information +// regarding copyright ownership. The ASF licenses this file +// to you under the Apache License, Version 2.0 (the +// "License"); you may not use this file except in compliance +// with the License. You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, +// software distributed under the License is distributed on an +// "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY +// KIND, either express or implied. See the License for the +// specific language governing permissions and limitations +// under the License. + +use std::fmt::Formatter; +use std::sync::Arc; + +use arrow::array::{ArrayRef, BinaryArray, RecordBatch, StringArray}; +use arrow::datatypes::{DataType, Field, Schema, SchemaRef, TimeUnit}; +use datafusion::common::tree_node::TreeNodeRecursion; +use datafusion::common::{exec_err, Result}; +use datafusion::execution::TaskContext; +use datafusion::physical_expr::{EquivalenceProperties, PhysicalExpr}; +use datafusion::physical_plan::execution_plan::EmissionType; +use datafusion::physical_plan::stream::RecordBatchStreamAdapter; +use datafusion::physical_plan::{ + apply_expression_roots, DisplayAs, DisplayFormatType, ExecutionPlan, ExecutionPlanProperties, + PlanProperties, SendableRecordBatchStream, +}; +use futures::stream; +use futures::StreamExt; + +use crate::execution::python_udf::ArrowPythonUdf; + +#[derive(Debug, Clone)] +pub struct ArrowPythonUdfSpec { + pub command: Vec, + pub args: Vec>, + pub arg_names: Vec, + pub return_type: DataType, + pub return_name: String, + pub python_version: String, +} + +/// Evaluates scalar PyArrow UDFs inside the native pipeline. Workers are +/// instantiated in `execute`, once per partition. Python module state remains +/// shared by every task in the executor's embedded interpreter. +#[derive(Debug)] +pub struct ArrowPythonUdfExec { + child: Arc, + specs: Vec, + max_records_per_batch: usize, + max_bytes_per_batch: usize, + schema: SchemaRef, + cache: Arc, +} + +impl ArrowPythonUdfExec { + pub fn try_new( + child: Arc, + specs: Vec, + max_records_per_batch: i32, + max_bytes_per_batch: i64, + ) -> Result { + if specs.is_empty() { + return exec_err!("ArrowPythonUdfExec requires at least one UDF"); + } + let mut fields: Vec = child + .schema() + .fields() + .iter() + .map(|f| f.as_ref().clone()) + .collect(); + for spec in &specs { + if spec.args.len() != spec.arg_names.len() { + return exec_err!("ArrowPythonUdf argument names are not aligned with arguments"); + } + for arg in &spec.args { + arg.data_type(&child.schema())?; + } + fields.push(Field::new( + &spec.return_name, + spec.return_type.clone(), + true, + )); + } + let schema = Arc::new(Schema::new(fields)); + let cache = Arc::new(PlanProperties::new( + EquivalenceProperties::new(Arc::clone(&schema)), + child.output_partitioning().clone(), + EmissionType::Incremental, + child.boundedness(), + )); + Ok(Self { + child, + specs, + max_records_per_batch: max_records_per_batch.max(0) as usize, + max_bytes_per_batch: max_bytes_per_batch.max(0) as usize, + schema, + cache, + }) + } + + fn evaluate_args( + specs: &[ArrowPythonUdfSpec], + batch: &RecordBatch, + ) -> Result>> { + specs + .iter() + .map(|spec| { + spec.args + .iter() + .map(|arg| arg.evaluate(batch)?.into_array(batch.num_rows())) + .collect::>>() + }) + .collect() + } + + // Spark's row-based Arrow writer checks its buffer size after each row, so + // the row that reaches the byte limit remains in that batch. The native + // path uses logical Arrow buffer sizes for its verified scalar types. + fn input_bytes(args: &[Vec], offset: usize, length: usize) -> Result { + if length == 0 { + return Ok(0); + } + let mut bytes = 0usize; + for array in args.iter().flatten() { + let value_bytes = match array.data_type() { + DataType::Boolean => length.div_ceil(8), + DataType::Int8 | DataType::UInt8 => length, + DataType::Int16 | DataType::UInt16 => length.saturating_mul(2), + DataType::Int32 | DataType::UInt32 | DataType::Float32 | DataType::Date32 => { + length.saturating_mul(4) + } + DataType::Int64 + | DataType::UInt64 + | DataType::Float64 + | DataType::Timestamp(TimeUnit::Microsecond, None) => length.saturating_mul(8), + DataType::Decimal128(_, _) => length.saturating_mul(16), + DataType::Utf8 => { + let values = array + .as_any() + .downcast_ref::() + .ok_or_else(|| { + datafusion::error::DataFusionError::Execution( + "Arrow UDF string argument has an unexpected array type" + .to_string(), + ) + })?; + let offsets = values.value_offsets(); + (offsets[offset + length] - offsets[offset]) as usize + + (length + 1).saturating_mul(4) + } + DataType::Binary => { + let values = array + .as_any() + .downcast_ref::() + .ok_or_else(|| { + datafusion::error::DataFusionError::Execution( + "Arrow UDF binary argument has an unexpected array type" + .to_string(), + ) + })?; + let offsets = values.value_offsets(); + (offsets[offset + length] - offsets[offset]) as usize + + (length + 1).saturating_mul(4) + } + other => return exec_err!("Unsupported Arrow UDF argument type: {other}"), + }; + bytes = bytes.saturating_add(value_bytes); + // Arrow Java's getBufferSizeFor counts the validity bitmap even + // when every value is non-null. + bytes = bytes.saturating_add(length.div_ceil(8)); + } + Ok(bytes) + } + + fn next_batch_length( + args: &[Vec], + offset: usize, + remaining: usize, + max_records: usize, + max_bytes: usize, + ) -> Result { + let limit = if max_records == 0 { + remaining + } else { + remaining.min(max_records) + }; + if limit == 0 || max_bytes == 0 || args.iter().all(Vec::is_empty) { + return Ok(limit); + } + if Self::input_bytes(args, offset, limit)? < max_bytes { + return Ok(limit); + } + let (mut low, mut high) = (1, limit); + while low < high { + let middle = low + (high - low) / 2; + if Self::input_bytes(args, offset, middle)? >= max_bytes { + high = middle; + } else { + low = middle + 1; + } + } + Ok(low) + } + + fn evaluate_batch( + specs: &[ArrowPythonUdfSpec], + workers: &[ArrowPythonUdf], + schema: SchemaRef, + batch: &RecordBatch, + args: &[Vec], + offset: usize, + length: usize, + ) -> Result { + let mut columns = batch.slice(offset, length).columns().to_vec(); + for ((spec, worker), function_args) in specs.iter().zip(workers).zip(args) { + let sliced_args: Vec<_> = function_args + .iter() + .map(|array| array.slice(offset, length)) + .collect(); + columns.push(worker.evaluate_named(&sliced_args, &spec.arg_names, length)?); + } + Ok(RecordBatch::try_new(schema, columns)?) + } +} + +impl DisplayAs for ArrowPythonUdfExec { + fn fmt_as(&self, t: DisplayFormatType, f: &mut Formatter) -> std::fmt::Result { + match t { + DisplayFormatType::Default + | DisplayFormatType::Verbose + | DisplayFormatType::TreeRender => { + write!(f, "CometArrowPythonUdfExec: {} UDF(s)", self.specs.len()) + } + } + } +} + +impl ExecutionPlan for ArrowPythonUdfExec { + fn name(&self) -> &str { + "CometArrowPythonUdfExec" + } + + fn schema(&self) -> SchemaRef { + Arc::clone(&self.schema) + } + + fn properties(&self) -> &Arc { + &self.cache + } + + fn children(&self) -> Vec<&Arc> { + vec![&self.child] + } + + fn apply_expressions( + &self, + f: &mut dyn FnMut(&Arc) -> Result, + ) -> Result { + apply_expression_roots(self.specs.iter().flat_map(|spec| spec.args.iter()), f) + } + + fn with_new_children( + self: Arc, + children: Vec>, + ) -> Result> { + if children.len() != 1 { + return exec_err!("ArrowPythonUdfExec requires exactly one child"); + } + Ok(Arc::new(Self::try_new( + Arc::clone(&children[0]), + self.specs.clone(), + self.max_records_per_batch as i32, + self.max_bytes_per_batch as i64, + )?)) + } + + fn execute( + &self, + partition: usize, + context: Arc, + ) -> Result { + let input = self.child.execute(partition, context)?; + let workers: Vec<_> = self + .specs + .iter() + .map(|spec| { + ArrowPythonUdf::from_command( + &spec.command, + spec.return_type.clone(), + true, + true, + &spec.python_version, + ) + }) + .collect::>()?; + let workers = Arc::new(workers); + let specs = Arc::new(self.specs.clone()); + let schema = Arc::clone(&self.schema); + let max_records_per_batch = self.max_records_per_batch; + let max_bytes_per_batch = self.max_bytes_per_batch; + let stream = input.flat_map(move |batch| { + let workers = Arc::clone(&workers); + let specs = Arc::clone(&specs); + let schema = Arc::clone(&schema); + let (batch, args, mut error) = match batch { + // Spark's Arrow writer does not invoke a scalar UDF for an empty input + // batch. Native scans may still emit one, so skip it before evaluating + // arguments or entering Python. + Ok(batch) if batch.num_rows() == 0 => (None, None, None), + Ok(batch) => { + match tokio::task::block_in_place(|| Self::evaluate_args(&specs, &batch)) { + Ok(args) => (Some(batch), Some(args), None), + Err(error) => (None, None, Some(error)), + } + } + Err(error) => (None, None, Some(error)), + }; + let mut offset = 0; + let mut failed = false; + // RecordBatch::slice shares Arrow buffers. Produce one result per poll + // so the remaining slices do not pin a second set of output batches. + stream::iter(std::iter::from_fn(move || { + if let Some(error) = error.take() { + return Some(Err(error)); + } + if failed { + return None; + } + let batch = batch.as_ref()?; + let args = args.as_ref()?; + if offset == batch.num_rows() { + return None; + } + let length = match Self::next_batch_length( + args, + offset, + batch.num_rows() - offset, + max_records_per_batch, + max_bytes_per_batch, + ) { + Ok(length) => length, + Err(error) => { + failed = true; + return Some(Err(error)); + } + }; + // Keep the JVM scan path synchronous so its Pending loop does not spin while + // Python runs. On a tokio worker, this hands its other tasks to another worker. + let result = tokio::task::block_in_place(|| { + Self::evaluate_batch( + &specs, + &workers, + Arc::clone(&schema), + batch, + args, + offset, + length, + ) + }); + offset += length; + Some(result) + })) + }); + Ok(Box::pin(RecordBatchStreamAdapter::new( + Arc::clone(&self.schema), + stream, + ))) + } +} + +#[cfg(test)] +mod tests { + use std::sync::Arc; + + use arrow::array::{ArrayRef, Int64Array, StringArray}; + + use super::ArrowPythonUdfExec; + + #[test] + fn byte_limit_splits_fixed_width_udf_arguments() { + let values: ArrayRef = Arc::new(Int64Array::from(vec![1, 2, 3, 4])); + let args = vec![vec![values]]; + assert_eq!( + ArrowPythonUdfExec::next_batch_length(&args, 0, 4, 10_000, 16).unwrap(), + 2 + ); + assert_eq!(ArrowPythonUdfExec::input_bytes(&args, 0, 2).unwrap(), 17); + assert_eq!( + ArrowPythonUdfExec::next_batch_length(&args, 0, 4, 10_000, 9).unwrap(), + 1 + ); + assert_eq!( + ArrowPythonUdfExec::next_batch_length(&args, 2, 2, 10_000, 16).unwrap(), + 2 + ); + assert_eq!( + ArrowPythonUdfExec::next_batch_length(&args, 0, 4, 1, 16).unwrap(), + 1 + ); + assert_eq!( + ArrowPythonUdfExec::next_batch_length(&args, 0, 4, 10_000, 0).unwrap(), + 4 + ); + } + + #[test] + fn byte_limit_counts_all_arguments_and_keeps_oversized_row() { + let values: ArrayRef = Arc::new(Int64Array::from(vec![1, 2, 3, 4])); + let args = vec![vec![Arc::clone(&values), values]]; + assert_eq!( + ArrowPythonUdfExec::next_batch_length(&args, 0, 4, 10_000, 16).unwrap(), + 1 + ); + assert_eq!( + ArrowPythonUdfExec::next_batch_length(&args, 0, 4, 10_000, 8).unwrap(), + 1 + ); + } + + #[test] + fn byte_limit_counts_variable_width_values_and_offsets() { + let values: ArrayRef = Arc::new(StringArray::from(vec!["a", "bb", "ccc", "d"])); + let args = vec![vec![values]]; + // Two values use 3 data bytes, 3 four-byte offsets, and 1 validity byte. + assert_eq!(ArrowPythonUdfExec::input_bytes(&args, 0, 2).unwrap(), 16); + assert_eq!( + ArrowPythonUdfExec::next_batch_length(&args, 0, 4, 10_000, 15).unwrap(), + 2 + ); + assert_eq!( + ArrowPythonUdfExec::next_batch_length(&args, 2, 2, 10_000, 15).unwrap(), + 2 + ); + } +} diff --git a/native/core/src/execution/operators/mod.rs b/native/core/src/execution/operators/mod.rs index a15a528f90b..1e51808bf29 100644 --- a/native/core/src/execution/operators/mod.rs +++ b/native/core/src/execution/operators/mod.rs @@ -22,6 +22,10 @@ pub use crate::errors::ExecutionError; pub use iceberg_scan::*; pub use scan::*; +#[cfg(feature = "python-udf")] +mod arrow_python_udf; +#[cfg(feature = "python-udf")] +pub use arrow_python_udf::{ArrowPythonUdfExec, ArrowPythonUdfSpec}; mod dynamic_filter; pub(crate) use dynamic_filter::{DynamicFilterJoinExec, TopKReaderFilterExec}; pub(crate) mod iceberg_common; diff --git a/native/core/src/execution/planner.rs b/native/core/src/execution/planner.rs index 99a52fb562f..4110bcf5d69 100644 --- a/native/core/src/execution/planner.rs +++ b/native/core/src/execution/planner.rs @@ -38,6 +38,8 @@ use crate::execution::operators::init_csv_datasource_exec; use crate::execution::operators::DynamicFilterJoinExec; use crate::execution::operators::IcebergScanExec; use crate::execution::operators::TopKReaderFilterExec; +#[cfg(feature = "python-udf")] +use crate::execution::operators::{ArrowPythonUdfExec, ArrowPythonUdfSpec}; use crate::execution::{ operators::{ ExecutionError, MergeActionContext, MergeInstructionExec, MergeRowsExec, ScanExec, @@ -1474,6 +1476,56 @@ impl PhysicalPlanner { // Fall back to the original monolithic match for other operators let children = &spark_plan.children; match spark_plan.op_struct.as_ref().unwrap() { + #[cfg(feature = "python-udf")] + OpStruct::ArrowPythonUdf(udf) => { + if children.len() != 1 { + return Err(ExecutionError::GeneralError( + "ArrowPythonUdf requires one child".to_string(), + )); + } + let (scans, shuffle_scans, child) = + self.create_plan(&children[0], inputs, partition_count)?; + let specs = udf + .functions + .iter() + .map(|function| { + let args = function + .args + .iter() + .map(|arg| self.create_expr(arg, child.schema())) + .collect::, _>>()?; + let return_type = to_arrow_datatype(function.return_type.as_ref().ok_or_else(|| { + ExecutionError::GeneralError( + "ArrowPythonUdf missing return type".to_string(), + ) + })?); + Ok(ArrowPythonUdfSpec { + command: function.command.clone(), + args, + arg_names: function.arg_names.clone(), + return_type, + return_name: function.return_name.clone(), + python_version: function.python_version.clone(), + }) + }) + .collect::, ExecutionError>>()?; + let native_plan = Arc::new(ArrowPythonUdfExec::try_new( + Arc::clone(&child.native_plan), + specs, + udf.max_records_per_batch, + udf.max_bytes_per_batch, + )?); + Ok(( + scans, + shuffle_scans, + Arc::new(SparkPlan::new(spark_plan.plan_id, native_plan, vec![child])), + )) + } + #[cfg(not(feature = "python-udf"))] + OpStruct::ArrowPythonUdf(_) => Err(ExecutionError::GeneralError( + "Native Arrow Python UDF support was not compiled in; rebuild with --features python-udf" + .to_string(), + )), OpStruct::Filter(filter) => { assert_eq!(children.len(), 1); let (scans, shuffle_scans, child) = diff --git a/native/core/src/execution/planner/operator_registry.rs b/native/core/src/execution/planner/operator_registry.rs index 70ed309ebeb..e1dcb4a8a1d 100644 --- a/native/core/src/execution/planner/operator_registry.rs +++ b/native/core/src/execution/planner/operator_registry.rs @@ -187,5 +187,6 @@ fn get_operator_type(spark_operator: &Operator) -> Option { // so the supports-mixed-codegen check skips it. OpStruct::ContribScan(_) => None, OpStruct::WindowGroupLimit(_) => None, // Not yet in OperatorType enum + OpStruct::ArrowPythonUdf(_) => None, // Built by the planner behind the python-udf feature } } diff --git a/native/core/src/execution/python_udf.rs b/native/core/src/execution/python_udf.rs new file mode 100644 index 00000000000..c346b68ab2c --- /dev/null +++ b/native/core/src/execution/python_udf.rs @@ -0,0 +1,453 @@ +// Licensed to the Apache Software Foundation (ASF) under one +// or more contributor license agreements. See the NOTICE file +// distributed with this work for additional information +// regarding copyright ownership. The ASF licenses this file +// to you under the Apache License, Version 2.0 (the +// "License"); you may not use this file except in compliance +// with the License. You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, +// software distributed under the License is distributed on an +// "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY +// KIND, either express or implied. See the License for the +// specific language governing permissions and limitations +// under the License. + +//! In-process bridge for Spark 4.1+ scalar Arrow UDFs. Each instance owns one +//! unpickled Python callable and must be created for one Spark task/partition. +//! The public API deliberately deals in Arrow arrays; the physical operator is +//! responsible for evaluating Catalyst arguments and preserving input columns. + +use arrow::array::{make_array, Array, ArrayRef}; +use arrow::datatypes::DataType; +use arrow::error::{ArrowError, Result}; +use arrow::ffi::{from_ffi, FFI_ArrowArray, FFI_ArrowSchema}; +use pyo3::ffi::Py_uintptr_t; +use pyo3::prelude::*; +use pyo3::types::{PyBytes, PyTuple}; + +fn initialize_python() -> Result<()> { + use std::ffi::CStr; + use std::sync::OnceLock; + + static RESULT: OnceLock> = OnceLock::new(); + RESULT + .get_or_init(|| { + // Spark sets PYTHONHASHSEED=0 on its Python workers by default. + // Match that seed before any Python object is created in the embedded + // interpreter, without changing the JVM process environment. + // SAFETY: OnceLock serializes initialization by Comet. No other Comet + // code accesses the Python C API before this function returns. + unsafe { + if pyo3::ffi::Py_IsInitialized() != 0 { + return Ok(()); + } + let mut config = std::mem::MaybeUninit::::uninit(); + pyo3::ffi::PyConfig_InitPythonConfig(config.as_mut_ptr()); + let mut config = config.assume_init(); + config.install_signal_handlers = 0; + config.use_hash_seed = 1; + config.hash_seed = 0; + let status = pyo3::ffi::Py_InitializeFromConfig(&config); + let error = if pyo3::ffi::PyStatus_Exception(status) != 0 { + if status.err_msg.is_null() { + "Python interpreter initialization failed".to_string() + } else { + CStr::from_ptr(status.err_msg) + .to_string_lossy() + .into_owned() + } + } else { + String::new() + }; + pyo3::ffi::PyConfig_Clear(&mut config); + if !error.is_empty() { + return Err(error); + } + pyo3::ffi::PyEval_SaveThread(); + Ok(()) + } + }) + .clone() + .map_err(ArrowError::ComputeError) +} + +#[cfg(target_os = "linux")] +fn make_python_symbols_global() -> Result<()> { + use std::ffi::CStr; + use std::sync::OnceLock; + + static RESULT: OnceLock> = OnceLock::new(); + RESULT + .get_or_init(|| { + // The JVM loads libcomet with RTLD_LOCAL. Its libpython dependency is + // local too, but CPython extension modules resolve Python C API + // symbols from the global namespace when they are imported. + let mut info = std::mem::MaybeUninit::::uninit(); + // SAFETY: Py_Initialize is a linked function address and info is + // writable storage for dladdr's result. + if unsafe { + libc::dladdr( + pyo3::ffi::Py_Initialize as *const () as *const libc::c_void, + info.as_mut_ptr(), + ) + } == 0 + { + return Err("cannot locate the linked Python library".to_string()); + } + // SAFETY: dladdr initialized info on success and dli_fname is a + // null-terminated path valid for the duration of this call. + let info = unsafe { info.assume_init() }; + if info.dli_fname.is_null() { + return Err("linked Python library has no path".to_string()); + } + let path = unsafe { CStr::from_ptr(info.dli_fname) }; + // RTLD_NOLOAD promotes the already-loaded libpython rather than + // loading a second copy with separate interpreter state. Keep the + // handle for the executor lifetime so its symbols remain global. + // SAFETY: path points to a valid C string returned by dladdr. + if unsafe { + libc::dlopen( + path.as_ptr(), + libc::RTLD_NOW | libc::RTLD_GLOBAL | libc::RTLD_NOLOAD, + ) + } + .is_null() + { + // SAFETY: dlerror returns a null-terminated message, if any. + let error = unsafe { libc::dlerror() }; + let detail = if error.is_null() { + "unknown dynamic loader error".to_string() + } else { + unsafe { CStr::from_ptr(error) } + .to_string_lossy() + .into_owned() + }; + return Err(format!("cannot expose Python C API symbols: {detail}")); + } + Ok(()) + }) + .clone() + .map_err(ArrowError::ComputeError) +} + +#[cfg(not(target_os = "linux"))] +fn make_python_symbols_global() -> Result<()> { + Ok(()) +} + +/// A scalar Arrow UDF loaded from Spark's pickled `(function, returnType)` command. +/// Spark serializes the return type for its worker; Comet uses the separately +/// serialized Arrow type from the physical plan instead. +pub struct ArrowPythonUdf { + callable: Py, + return_type: DataType, + allow_cast: bool, + safe_cast: bool, +} + +impl ArrowPythonUdf { + pub fn from_command( + command: &[u8], + return_type: DataType, + allow_cast: bool, + safe_cast: bool, + python_version: &str, + ) -> Result { + make_python_symbols_global()?; + initialize_python()?; + Python::attach(|py| { + if !python_version.is_empty() { + let info = py + .import("sys") + .map_err(python_error)? + .getattr("version_info") + .map_err(python_error)?; + let major: u8 = info + .get_item(0) + .map_err(python_error)? + .extract() + .map_err(python_error)?; + let minor: u8 = info + .get_item(1) + .map_err(python_error)? + .extract() + .map_err(python_error)?; + let actual = format!("{major}.{minor}"); + if actual != python_version { + return Err(ArrowError::ComputeError(format!( + "Arrow UDF requires Python {python_version}, embedded interpreter is {actual}" + ))); + } + } + let pickle = py.import("pickle").map_err(python_error)?; + let loaded = pickle + .call_method1("loads", (PyBytes::new(py, command),)) + .map_err(python_error)?; + let tuple = loaded.cast::().map_err(python_error)?; + if tuple.len() != 2 { + return Err(ArrowError::ComputeError(format!( + "Arrow UDF command must contain (function, returnType), got {} items", + tuple.len() + ))); + } + let callable = tuple.get_item(0).map_err(python_error)?; + if !callable.is_callable() { + return Err(ArrowError::ComputeError( + "Arrow UDF command does not contain a callable".to_string(), + )); + } + Ok(Self { + callable: callable.unbind(), + return_type, + allow_cast, + safe_cast, + }) + }) + } + + /// Evaluate one Arrow batch, with the same row count for every argument. + /// Python receives and returns `pyarrow.Array` objects via the Arrow C Data + /// interface; no row conversion or Arrow IPC serialization occurs here. + pub fn evaluate(&self, args: &[ArrayRef], num_rows: usize) -> Result { + let names = vec![String::new(); args.len()]; + self.evaluate_named(args, &names, num_rows) + } + + pub fn evaluate_named( + &self, + args: &[ArrayRef], + names: &[String], + num_rows: usize, + ) -> Result { + if args.len() != names.len() { + return Err(ArrowError::ComputeError( + "Arrow UDF argument names are not aligned with arguments".to_string(), + )); + } + for (index, arg) in args.iter().enumerate() { + if arg.len() != num_rows { + return Err(ArrowError::ComputeError(format!( + "Arrow UDF argument {index} has {} rows, expected {num_rows}", + arg.len() + ))); + } + } + + Python::attach(|py| { + let pa = py.import("pyarrow").map_err(python_error)?; + let array_class = pa.getattr("Array").map_err(python_error)?; + let mut py_args = Vec::with_capacity(args.len()); + for arg in args { + let data = arg.to_data(); + // PyArrow takes ownership of these C Data structs and clears their + // release callbacks, so the pointed-to storage must be writable. + let mut ffi_array = FFI_ArrowArray::new(&data); + let mut ffi_schema = FFI_ArrowSchema::try_from(data.data_type())?; + let py_arg = array_class + .call_method1( + "_import_from_c", + ( + &raw mut ffi_array as Py_uintptr_t, + &raw mut ffi_schema as Py_uintptr_t, + ), + ) + .map_err(python_error)?; + py_args.push(py_arg); + } + + let kwargs = pyo3::types::PyDict::new(py); + let mut positional = Vec::new(); + for (arg, name) in py_args.into_iter().zip(names) { + if name.is_empty() { + positional.push(arg); + } else { + kwargs.set_item(name, arg).map_err(python_error)?; + } + } + let result = self + .callable + .bind(py) + .call( + PyTuple::new(py, positional).map_err(python_error)?, + Some(&kwargs), + ) + .map_err(python_error)?; + if !result.is_instance(&array_class).map_err(python_error)? { + return Err(ArrowError::ComputeError( + "Arrow UDF must return a pyarrow.Array".to_string(), + )); + } + let result_len = result.len().map_err(python_error)?; + if result_len != num_rows { + return Err(ArrowError::ComputeError(format!( + "Arrow UDF returned {result_len} rows, expected {num_rows}" + ))); + } + + let mut ffi_return_type = FFI_ArrowSchema::try_from(&self.return_type)?; + let expected_type = pa + .getattr("DataType") + .map_err(python_error)? + .call_method1( + "_import_from_c", + (&raw mut ffi_return_type as Py_uintptr_t,), + ) + .map_err(python_error)?; + let actual_type = result.getattr("type").map_err(python_error)?; + let typed_result = if actual_type.eq(&expected_type).map_err(python_error)? { + result + } else if self.allow_cast { + let kwargs = pyo3::types::PyDict::new(py); + kwargs + .set_item("safe", self.safe_cast) + .map_err(python_error)?; + result + .call_method("cast", (expected_type,), Some(&kwargs)) + .map_err(python_error)? + } else { + return Err(ArrowError::ComputeError(format!( + "Arrow UDF returned type {}, expected {}", + actual_type.str().map_err(python_error)?, + expected_type.str().map_err(python_error)? + ))); + }; + + let mut out_array = FFI_ArrowArray::empty(); + let mut out_schema = FFI_ArrowSchema::empty(); + typed_result + .call_method1( + "_export_to_c", + ( + &raw mut out_array as Py_uintptr_t, + &raw mut out_schema as Py_uintptr_t, + ), + ) + .map_err(python_error)?; + // SAFETY: PyArrow filled both C Data structs and transferred ownership + // of the array to `out_array`; Arrow validates the schema and buffers. + let data = unsafe { from_ffi(out_array, &out_schema) }?; + if data.data_type() != &self.return_type { + return Err(ArrowError::ComputeError(format!( + "Arrow UDF returned type {}, expected {}", + data.data_type(), + self.return_type + ))); + } + Ok(make_array(data)) + }) + } +} + +fn python_error(error: impl std::fmt::Display) -> ArrowError { + ArrowError::ComputeError(format!("Arrow UDF Python error: {error}")) +} + +#[cfg(test)] +mod tests { + use super::*; + use arrow::array::{Int32Array, Int64Array}; + use std::sync::Arc; + + fn pickled_command(py: Python<'_>, module: &str, function: &str) -> Vec { + let pickle = py.import("pickle").unwrap(); + let module = py.import(module).unwrap(); + let callable = module.getattr(function).unwrap(); + pickle + .call_method1("dumps", ((callable.unbind(), py.None()),)) + .unwrap() + .extract() + .unwrap() + } + + #[test] + fn evaluates_arrow_arrays_and_nulls() { + initialize_python().unwrap(); + Python::attach(|py| { + let command = pickled_command(py, "pyarrow.compute", "negate"); + let udf = + ArrowPythonUdf::from_command(&command, DataType::Int64, false, true, "").unwrap(); + let input: ArrayRef = Arc::new(Int64Array::from(vec![Some(1), None, Some(3)])); + let result = udf.evaluate(&[input], 3).unwrap(); + let expected = Int64Array::from(vec![Some(-1), None, Some(-3)]); + assert_eq!(result.as_ref(), &expected); + }); + } + + #[test] + fn rejects_wrong_length_and_non_array_result() { + initialize_python().unwrap(); + Python::attach(|py| { + let command = pickled_command(py, "builtins", "len"); + let udf = + ArrowPythonUdf::from_command(&command, DataType::Int64, false, true, "").unwrap(); + let input: ArrayRef = Arc::new(Int64Array::from(vec![1, 2])); + assert!(udf.evaluate(&[Arc::clone(&input)], 1).is_err()); + assert!(udf + .evaluate(&[input], 2) + .unwrap_err() + .to_string() + .contains("pyarrow.Array")); + }); + } + + #[test] + fn supports_named_arguments_and_safe_cast() { + initialize_python().unwrap(); + Python::attach(|py| { + let callable = py + .eval( + c"lambda *, x, y: __import__('pyarrow.compute', fromlist=['add']).add(x, y)", + None, + None, + ) + .unwrap(); + let command: Vec = py + .import("cloudpickle") + .unwrap() + .call_method1("dumps", ((callable.unbind(), py.None()),)) + .unwrap() + .extract() + .unwrap(); + let input: ArrayRef = Arc::new(Int64Array::from(vec![1, 2])); + let names = vec!["x".to_string(), "y".to_string()]; + let strict = + ArrowPythonUdf::from_command(&command, DataType::Int32, false, true, "").unwrap(); + assert!(strict + .evaluate_named(&[Arc::clone(&input), Arc::clone(&input)], &names, 2) + .is_err()); + let cast = + ArrowPythonUdf::from_command(&command, DataType::Int32, true, true, "").unwrap(); + let result = cast + .evaluate_named(&[Arc::clone(&input), input], &names, 2) + .unwrap(); + assert_eq!(result.as_ref(), &Int32Array::from(vec![2, 4])); + }); + } + + #[test] + fn rejects_python_result_with_wrong_length() { + initialize_python().unwrap(); + Python::attach(|py| { + let command = pickled_command(py, "pyarrow.compute", "drop_null"); + let udf = + ArrowPythonUdf::from_command(&command, DataType::Int64, true, true, "").unwrap(); + let input: ArrayRef = Arc::new(Int64Array::from(vec![Some(1), None])); + assert!(udf + .evaluate(&[input], 2) + .unwrap_err() + .to_string() + .contains("returned 1 rows, expected 2")); + }); + } + + #[test] + fn rejects_mismatched_python_version() { + let error = ArrowPythonUdf::from_command(&[], DataType::Int64, true, true, "0.0") + .err() + .unwrap(); + assert!(error.to_string().contains("requires Python 0.0")); + } +} diff --git a/native/core/src/lib.rs b/native/core/src/lib.rs index b712a65c115..1ae731aee7e 100644 --- a/native/core/src/lib.rs +++ b/native/core/src/lib.rs @@ -178,6 +178,14 @@ fn logger_config(log_conf_path: &str, log_level: &str) -> CometResult { } } +#[no_mangle] +pub extern "system" fn Java_org_apache_comet_NativeBase_nativeSupportsPythonUdf( + _: EnvUnowned, + _: JClass, +) -> jni::sys::jboolean { + cfg!(feature = "python-udf") +} + #[no_mangle] /// Releases the global Tokio runtime used by Comet native execution. pub extern "system" fn Java_org_apache_comet_NativeBase_release(_e: EnvUnowned, _class: JClass) { diff --git a/native/proto/src/proto/operator.proto b/native/proto/src/proto/operator.proto index 527461f47a5..1ecc4c32bc4 100644 --- a/native/proto/src/proto/operator.proto +++ b/native/proto/src/proto/operator.proto @@ -70,6 +70,7 @@ message Operator { IcebergWrite iceberg_write = 120; RangeScan range_scan = 121; MergeRows merge_rows = 122; + ArrowPythonUdf arrow_python_udf = 123; // Extension point for optional, out-of-tree contrib scans (Delta, Lance, ...). The concrete // scan message (e.g. `DeltaScan`) is packed into this envelope on the JVM side and dispatched // by `type_url` on the native side. Using a single permanent field -- rather than a new oneof @@ -801,6 +802,29 @@ message Projection { repeated spark.spark_expression.Expr project_list = 1; } +// Spark 4.1 scalar @arrow_udf evaluation. The input columns pass through, and +// one result column is appended for each function in order. +message ArrowPythonUdf { + repeated ArrowPythonFunction functions = 1; + // Spark's Arrow writer splits input batches before calling a scalar Arrow UDF. + // A non-positive value disables the limit. + int32 max_records_per_batch = 2; + // Spark also stops a batch once its Arrow input buffers reach this size. + int64 max_bytes_per_batch = 3; +} + +message ArrowPythonFunction { + // Spark's pickled (function, returnType) command. Broadcast-backed commands + // are rejected by the JVM planner until broadcast support is implemented. + bytes command = 1; + repeated spark.spark_expression.Expr args = 2; + spark.spark_expression.DataType return_type = 3; + string return_name = 4; + // Empty string for a positional argument, otherwise Python keyword name. + repeated string arg_names = 5; + string python_version = 6; +} + message Filter { spark.spark_expression.Expr predicate = 1; } diff --git a/spark/src/main/java/org/apache/comet/NativeBase.java b/spark/src/main/java/org/apache/comet/NativeBase.java index ab4eee856ba..6189b107505 100644 --- a/spark/src/main/java/org/apache/comet/NativeBase.java +++ b/spark/src/main/java/org/apache/comet/NativeBase.java @@ -341,6 +341,20 @@ private static String resourceName() { */ static native void init(String logConfPath, String logLevel); + private static native boolean nativeSupportsPythonUdf(); + + /** Whether this native library was built with the in-process Python UDF feature. */ + public static boolean supportsPythonUdf() { + if (!loaded) { + return false; + } + try { + return nativeSupportsPythonUdf(); + } catch (UnsatisfiedLinkError ignored) { + return false; + } + } + /** Release native resources through JNI */ static native void release(); diff --git a/spark/src/main/scala/org/apache/comet/CometConf.scala b/spark/src/main/scala/org/apache/comet/CometConf.scala index 09167333deb..71d3953aa78 100644 --- a/spark/src/main/scala/org/apache/comet/CometConf.scala +++ b/spark/src/main/scala/org/apache/comet/CometConf.scala @@ -455,6 +455,16 @@ object CometConf extends ShimCometConf { .booleanConf .createWithDefault(false) + val COMET_NATIVE_ARROW_PYTHON_UDF_ENABLED: ConfigEntry[Boolean] = + conf("spark.comet.exec.nativeArrowPythonUDF.enabled") + .category(CATEGORY_EXEC) + .doc( + "Experimental: execute Spark 4.1 scalar @arrow_udf functions inside the Comet native " + + "pipeline via PyO3. Requires a native library built with the python-udf Cargo feature " + + "and a compatible Python/PyArrow installation on every executor.") + .booleanConf + .createWithDefault(false) + val COMET_TRACING_ENABLED: ConfigEntry[Boolean] = conf("spark.comet.tracing.enabled") .category(CATEGORY_TUNING) .doc(s"Enable fine-grained tracing of events and memory usage. $TRACING_GUIDE.") diff --git a/spark/src/main/scala/org/apache/comet/rules/CometExecRule.scala b/spark/src/main/scala/org/apache/comet/rules/CometExecRule.scala index c16d1fc8f8f..3d36758f874 100644 --- a/spark/src/main/scala/org/apache/comet/rules/CometExecRule.scala +++ b/spark/src/main/scala/org/apache/comet/rules/CometExecRule.scala @@ -34,7 +34,7 @@ import org.apache.spark.sql.catalyst.util.sideBySide import org.apache.spark.sql.comet._ import org.apache.spark.sql.comet.execution.arrow.ArrowCachedBatchSerializer import org.apache.spark.sql.comet.execution.shuffle.{CometColumnarShuffle, CometNativeShuffle, CometShuffleExchangeExec} -import org.apache.spark.sql.comet.shims.{ShimCometEmptyRelation, ShimCometOneRowRelation} +import org.apache.spark.sql.comet.shims.{ShimCometArrowEvalPythonExec, ShimCometEmptyRelation, ShimCometOneRowRelation} import org.apache.spark.sql.comet.util.Utils import org.apache.spark.sql.execution._ import org.apache.spark.sql.execution.adaptive.{AdaptiveSparkPlanExec, AQEShuffleReadExec, BroadcastQueryStageExec, LogicalQueryStage, QueryStageExec, ShuffleQueryStageExec} @@ -126,6 +126,7 @@ object CometExecRule { classOf[WindowExec] -> CometWindowExec) ++ // EmptyRelationExec was introduced in Spark 4.0. ShimCometEmptyRelation.emptyRelationClass.map(_ -> CometEmptyRelationExec) ++ + ShimCometArrowEvalPythonExec.nativeOperator ++ // WindowGroupLimitExec exists only on Spark 3.5+; the shim returns None on 3.4. ShimCometWindowGroupLimit.windowGroupLimitClass.map(_ -> CometWindowGroupLimitExec) ++ // MergeRowsExec exists only on Spark 3.5+; the shim is empty on 3.4. diff --git a/spark/src/main/spark-3.x/org/apache/spark/sql/comet/shims/ShimCometArrowEvalPythonExec.scala b/spark/src/main/spark-3.x/org/apache/spark/sql/comet/shims/ShimCometArrowEvalPythonExec.scala new file mode 100644 index 00000000000..cc7491a0816 --- /dev/null +++ b/spark/src/main/spark-3.x/org/apache/spark/sql/comet/shims/ShimCometArrowEvalPythonExec.scala @@ -0,0 +1,28 @@ +/* + * Licensed to the Apache Software Foundation (ASF) under one + * or more contributor license agreements. See the NOTICE file + * distributed with this work for additional information + * regarding copyright ownership. The ASF licenses this file + * to you under the Apache License, Version 2.0 (the + * "License"); you may not use this file except in compliance + * with the License. You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, + * software distributed under the License is distributed on an + * "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY + * KIND, either express or implied. See the License for the + * specific language governing permissions and limitations + * under the License. + */ + +package org.apache.spark.sql.comet.shims + +import org.apache.spark.sql.execution.SparkPlan + +import org.apache.comet.serde.CometOperatorSerde + +object ShimCometArrowEvalPythonExec { + def nativeOperator: Option[(Class[_ <: SparkPlan], CometOperatorSerde[_])] = None +} diff --git a/spark/src/main/spark-4.0/org/apache/spark/sql/comet/shims/ShimCometArrowEvalPythonExec.scala b/spark/src/main/spark-4.0/org/apache/spark/sql/comet/shims/ShimCometArrowEvalPythonExec.scala new file mode 100644 index 00000000000..cc7491a0816 --- /dev/null +++ b/spark/src/main/spark-4.0/org/apache/spark/sql/comet/shims/ShimCometArrowEvalPythonExec.scala @@ -0,0 +1,28 @@ +/* + * Licensed to the Apache Software Foundation (ASF) under one + * or more contributor license agreements. See the NOTICE file + * distributed with this work for additional information + * regarding copyright ownership. The ASF licenses this file + * to you under the Apache License, Version 2.0 (the + * "License"); you may not use this file except in compliance + * with the License. You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, + * software distributed under the License is distributed on an + * "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY + * KIND, either express or implied. See the License for the + * specific language governing permissions and limitations + * under the License. + */ + +package org.apache.spark.sql.comet.shims + +import org.apache.spark.sql.execution.SparkPlan + +import org.apache.comet.serde.CometOperatorSerde + +object ShimCometArrowEvalPythonExec { + def nativeOperator: Option[(Class[_ <: SparkPlan], CometOperatorSerde[_])] = None +} diff --git a/spark/src/main/spark-4.1+/org/apache/spark/sql/comet/CometArrowEvalPythonExec.scala b/spark/src/main/spark-4.1+/org/apache/spark/sql/comet/CometArrowEvalPythonExec.scala new file mode 100644 index 00000000000..6eb40873215 --- /dev/null +++ b/spark/src/main/spark-4.1+/org/apache/spark/sql/comet/CometArrowEvalPythonExec.scala @@ -0,0 +1,215 @@ +/* + * Licensed to the Apache Software Foundation (ASF) under one + * or more contributor license agreements. See the NOTICE file + * distributed with this work for additional information + * regarding copyright ownership. The ASF licenses this file + * to you under the Apache License, Version 2.0 (the + * "License"); you may not use this file except in compliance + * with the License. You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, + * software distributed under the License is distributed on an + * "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY + * KIND, either express or implied. See the License for the + * specific language governing permissions and limitations + * under the License. + */ + +package org.apache.spark.sql.comet + +import java.nio.charset.StandardCharsets + +import scala.jdk.CollectionConverters._ + +import org.apache.spark.api.python.PythonEvalType +import org.apache.spark.sql.SparkSession +import org.apache.spark.sql.catalyst.expressions.{Attribute, AttributeSet, Expression, NamedArgumentExpression, NamedExpression, PythonUDF} +import org.apache.spark.sql.execution.{PartitioningPreservingUnaryExecNode, SparkPlan} +import org.apache.spark.sql.execution.python.ArrowEvalPythonExec +import org.apache.spark.sql.types.{BinaryType, BooleanType, ByteType, DataType, DateType, DecimalType, DoubleType, FloatType, IntegerType, LongType, ShortType, StringType, TimestampNTZType} + +import com.google.common.base.Objects +import com.google.protobuf.ByteString + +import org.apache.comet.{CometConf, ConfigEntry, NativeBase} +import org.apache.comet.CometSparkSessionExtensions.withFallbackReason +import org.apache.comet.serde.{CometOperatorSerde, Compatible, OperatorOuterClass, QueryPlanSerde, SupportLevel, Unsupported} +import org.apache.comet.serde.OperatorOuterClass.Operator + +/** Native execution for Spark 4.1+ scalar `@arrow_udf` functions. */ +object CometArrowEvalPythonExec extends CometOperatorSerde[ArrowEvalPythonExec] { + + // SparkContext adds this entry even when the user has not configured a Python + // environment. Keep other overrides on Spark's worker path. + private def hasUnsupportedEnvironment(env: java.util.Map[String, String]): Boolean = + env != null && env.asScala.exists { case (key, value) => + key != "PYTHONHASHSEED" || value != "0" + } + + // PySpark's Accumulator.__reduce__ serializes a reference to + // pyspark.accumulators._deserialize_accumulator. Spark's worker forwards its + // task-local updates to the JVM when the task finishes; embedded Python does + // not have that worker protocol. A match may also come from a harmless string + // in the pickle, in which case Spark's worker path is the safe choice. + private def hasSerializedAccumulator(command: Seq[Byte]): Boolean = + new String(command.toArray, StandardCharsets.ISO_8859_1).contains("pyspark.accumulators") + + private def hasCompatibleArrowSchema(dataType: DataType): Boolean = dataType match { + case _: BooleanType | _: ByteType | _: ShortType | _: IntegerType | _: LongType | + _: FloatType | _: DoubleType | _: BinaryType | _: DateType | _: DecimalType | + _: TimestampNTZType => + true + // Spark's Arrow conversion accepts plain strings. Collated and constrained strings + // may carry semantics that are not represented by Comet's Utf8 Arrow type. + case s: StringType if s == StringType => true + case _ => false + } + + override def enabledConfig: Option[ConfigEntry[Boolean]] = + Some(CometConf.COMET_NATIVE_ARROW_PYTHON_UDF_ENABLED) + + override def getSupportLevel(op: ArrowEvalPythonExec): SupportLevel = { + if (!NativeBase.supportsPythonUdf()) { + return Unsupported(Some("Native library lacks the python-udf feature")) + } + if (op.evalType != PythonEvalType.SQL_SCALAR_ARROW_UDF) { + return Unsupported(Some("Only scalar @arrow_udf is supported")) + } + if (op.udfs.isEmpty || op.udfs.length != op.resultAttrs.length) { + return Unsupported(Some("Arrow UDF functions and result attributes do not match")) + } + if (op.conf.arrowUseLargeVarTypes) { + return Unsupported(Some("Arrow UDF large variable types are not supported in-process")) + } + if (op.conf.pythonUDFProfiler.nonEmpty) { + return Unsupported(Some("Arrow UDF profiling is not supported in-process")) + } + if (SparkSession.active.sparkContext.listFiles().nonEmpty) { + return Unsupported(Some("Spark-added files are not supported in-process")) + } + if (op.udfs.exists(_.children.exists(expr => !hasCompatibleArrowSchema(expr.dataType))) || + op.resultAttrs.exists(attr => !hasCompatibleArrowSchema(attr.dataType))) { + return Unsupported(Some("Arrow UDF type is outside the verified native Arrow schema set")) + } + op.udfs.collectFirst { + case udf if udf.func.broadcastVars != null && !udf.func.broadcastVars.isEmpty => + "Arrow UDF broadcast variables are not supported in-process" + case udf if udf.func.pythonIncludes != null && !udf.func.pythonIncludes.isEmpty => + "Arrow UDF Python includes are not supported in-process" + case udf if hasSerializedAccumulator(udf.func.command) => + "Arrow UDF accumulators are not supported in-process" + case udf if hasUnsupportedEnvironment(udf.func.envVars) => + "Arrow UDF Python environment overrides are not supported in-process" + case udf if udf.children.exists(_.find(_.isInstanceOf[PythonUDF]).nonEmpty) => + "Chained Arrow UDFs are not supported in-process" + } match { + case Some(reason) => Unsupported(Some(reason)) + case None => Compatible(None) + } + } + + override def convert( + op: ArrowEvalPythonExec, + builder: Operator.Builder, + childOp: Operator*): Option[Operator] = { + if (childOp.length != 1) { + withFallbackReason(op, "Arrow UDF requires one native child") + return None + } + + val functions = op.udfs.zip(op.resultAttrs).map { case (udf, attr) => + val args: Seq[(Expression, String)] = udf.children.map { + case NamedArgumentExpression(key, value) => (value, key) + case other => (other, "") + } + val argProtos = args.map { case (expr, _) => + QueryPlanSerde.exprToProto(expr, op.child.output) + } + val returnType = QueryPlanSerde.serializeDataType(attr.dataType) + if (argProtos.exists(_.isEmpty) || returnType.isEmpty) { + None + } else { + Some( + OperatorOuterClass.ArrowPythonFunction + .newBuilder() + .setCommand(ByteString.copyFrom(udf.func.command.toArray)) + .addAllArgs(argProtos.map(_.get).asJava) + .addAllArgNames(args.map(_._2).asJava) + .setReturnType(returnType.get) + .setReturnName(attr.name) + .setPythonVersion(udf.func.pythonVer) + .build()) + } + } + if (functions.exists(_.isEmpty)) { + withFallbackReason(op, "Arrow UDF argument or return type cannot be serialized") + None + } else { + val native = OperatorOuterClass.ArrowPythonUdf + .newBuilder() + .addAllFunctions(functions.map(_.get).asJava) + .setMaxRecordsPerBatch(op.conf.arrowMaxRecordsPerBatch) + .setMaxBytesPerBatch(op.conf.arrowMaxBytesPerBatch) + Some(builder.setArrowPythonUdf(native).build()) + } + } + + override def createExec(nativeOp: Operator, op: ArrowEvalPythonExec): CometNativeExec = + CometArrowEvalPythonExec( + nativeOp, + op, + op.output, + op.resultAttrs, + op.udfs, + op.child, + SerializedPlan(None)) +} + +case class CometArrowEvalPythonExec( + override val nativeOp: Operator, + override val originalPlan: SparkPlan, + override val output: Seq[Attribute], + resultAttrs: Seq[Attribute], + udfs: Seq[PythonUDF], + child: SparkPlan, + override val serializedPlanOpt: SerializedPlan) + extends CometUnaryExec + with PartitioningPreservingUnaryExecNode { + + override def producedAttributes: AttributeSet = AttributeSet(resultAttrs) + + // Never render nativeOp: it contains the pickled Python command and can also + // contain scan credentials in its child operators. + override def stringArgs: Iterator[Any] = Iterator(output, resultAttrs, child) + + override def equals(obj: Any): Boolean = obj match { + case other: CometArrowEvalPythonExec => + output == other.output && + resultAttrs == other.resultAttrs && + udfs == other.udfs && + child == other.child && + nativeOp.getArrowPythonUdf.getMaxRecordsPerBatch == + other.nativeOp.getArrowPythonUdf.getMaxRecordsPerBatch && + nativeOp.getArrowPythonUdf.getMaxBytesPerBatch == + other.nativeOp.getArrowPythonUdf.getMaxBytesPerBatch && + serializedPlanOpt == other.serializedPlanOpt + case _ => false + } + + override def hashCode(): Int = + Objects.hashCode( + output, + resultAttrs, + udfs, + child, + nativeOp.getArrowPythonUdf.getMaxRecordsPerBatch: java.lang.Integer, + nativeOp.getArrowPythonUdf.getMaxBytesPerBatch: java.lang.Long, + serializedPlanOpt) + + override protected def outputExpressions: Seq[NamedExpression] = output + + override protected def withNewChildInternal(newChild: SparkPlan): SparkPlan = + copy(child = newChild) +} diff --git a/spark/src/main/spark-4.1+/org/apache/spark/sql/comet/shims/ShimCometArrowEvalPythonExec.scala b/spark/src/main/spark-4.1+/org/apache/spark/sql/comet/shims/ShimCometArrowEvalPythonExec.scala new file mode 100644 index 00000000000..47acfb96f8e --- /dev/null +++ b/spark/src/main/spark-4.1+/org/apache/spark/sql/comet/shims/ShimCometArrowEvalPythonExec.scala @@ -0,0 +1,31 @@ +/* + * Licensed to the Apache Software Foundation (ASF) under one + * or more contributor license agreements. See the NOTICE file + * distributed with this work for additional information + * regarding copyright ownership. The ASF licenses this file + * to you under the Apache License, Version 2.0 (the + * "License"); you may not use this file except in compliance + * with the License. You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, + * software distributed under the License is distributed on an + * "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY + * KIND, either express or implied. See the License for the + * specific language governing permissions and limitations + * under the License. + */ + +package org.apache.spark.sql.comet.shims + +import org.apache.spark.sql.comet.CometArrowEvalPythonExec +import org.apache.spark.sql.execution.SparkPlan +import org.apache.spark.sql.execution.python.ArrowEvalPythonExec + +import org.apache.comet.serde.CometOperatorSerde + +object ShimCometArrowEvalPythonExec { + def nativeOperator: Option[(Class[_ <: SparkPlan], CometOperatorSerde[_])] = + Some(classOf[ArrowEvalPythonExec] -> CometArrowEvalPythonExec) +} diff --git a/spark/src/test/resources/pyspark/test_native_arrow_udf_files.py b/spark/src/test/resources/pyspark/test_native_arrow_udf_files.py new file mode 100644 index 00000000000..3bb68189a04 --- /dev/null +++ b/spark/src/test/resources/pyspark/test_native_arrow_udf_files.py @@ -0,0 +1,92 @@ +#!/usr/bin/env python3 +# Licensed to the Apache Software Foundation (ASF) under one +# or more contributor license agreements. See the NOTICE file +# distributed with this work for additional information +# regarding copyright ownership. The ASF licenses this file +# to you under the Apache License, Version 2.0 (the +# "License"); you may not use this file except in compliance +# with the License. You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, +# software distributed under the License is distributed on an +# "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY +# KIND, either express or implied. See the License for the +# specific language governing permissions and limitations +# under the License. + +"""Verify that Spark-distributed files keep scalar Arrow UDFs on Python workers.""" + +import os + +import pyarrow as pa +import pytest +from pyspark import SparkFiles +from pyspark.sql import SparkSession +from pyspark.sql.pandas.functions import arrow_udf + +from conftest import resolve_comet_jar + + +@pytest.fixture +def spark(): + jar = resolve_comet_jar() + os.environ["PYSPARK_SUBMIT_ARGS"] = ( + f"--jars {jar} --driver-class-path {jar} pyspark-shell" + ) + session = ( + SparkSession.builder.master("local[2]") + .appName("comet-native-arrow-udf-files") + .config("spark.plugins", "org.apache.spark.CometPlugin") + .config("spark.comet.enabled", "true") + .config("spark.comet.exec.enabled", "true") + .config("spark.comet.shuffle.enabled", "false") + .config("spark.comet.exec.nativeArrowPythonUDF.enabled", "true") + .config("spark.sql.adaptive.enabled", "false") + .config("spark.memory.offHeap.enabled", "true") + .config("spark.memory.offHeap.size", "2g") + .getOrCreate() + ) + try: + yield session + finally: + session.stop() + + +@pytest.mark.parametrize("file_kind", ["python", "data"]) +def test_spark_added_files_use_python_worker(spark, tmp_path, file_kind): + assert spark.sparkContext._jsc.sc().listFiles().size() == 0 + if file_kind == "python": + helper = tmp_path / "comet_arrow_udf_helper.py" + helper.write_text("def offset():\n return 7\n") + spark.sparkContext.addPyFile(str(helper)) + + @arrow_udf("long") + def add_offset(values): + import comet_arrow_udf_helper + + return pa.array( + [value + comet_arrow_udf_helper.offset() for value in values.to_pylist()], + type=pa.int64(), + ) + + else: + data = tmp_path / "comet_arrow_udf_offset.txt" + data.write_text("7\n") + spark.sparkContext.addFile(str(data)) + + @arrow_udf("long") + def add_offset(values): + with open(SparkFiles.get("comet_arrow_udf_offset.txt")) as source: + offset = int(source.read()) + return pa.array( + [value + offset for value in values.to_pylist()], type=pa.int64() + ) + + assert spark.sparkContext._jsc.sc().listFiles().size() == 1 + result = spark.range(2).select(add_offset("id")) + plan = result._jdf.queryExecution().executedPlan().toString() + assert "ArrowEvalPython" in plan + assert "CometArrowEvalPython" not in plan + assert [row[0] for row in result.collect()] == [7, 8] diff --git a/spark/src/test/resources/pyspark/test_pyarrow_udf.py b/spark/src/test/resources/pyspark/test_pyarrow_udf.py index 683045743e3..5bcd2fac507 100644 --- a/spark/src/test/resources/pyspark/test_pyarrow_udf.py +++ b/spark/src/test/resources/pyspark/test_pyarrow_udf.py @@ -112,6 +112,174 @@ def _assert_plan_matches_mode( ) +def test_scalar_arrow_udf_uses_native_path_and_spark_batch_limit(spark): + # This must use PySpark's real UDF wrapper: it populates the default + # PYTHONHASHSEED entry that a hand-built SimplePythonFunction omits. + from pyspark.sql.pandas import functions as pandas_functions + + if not hasattr(pandas_functions, "arrow_udf"): + pytest.skip("scalar arrow_udf requires Spark 4.1 or later") + + @pandas_functions.arrow_udf("long") + def batch_length(values): + return pa.array([len(values)] * len(values), type=pa.int64()) + + @pandas_functions.arrow_udf("long") + def string_hash(values): + return pa.array([hash(value) for value in values.to_pylist()], type=pa.int64()) + + previous_comet_enabled = spark.conf.get("spark.comet.enabled") + spark.conf.set("spark.sql.adaptive.enabled", "false") + spark.conf.set("spark.sql.execution.arrow.maxRecordsPerBatch", "2") + spark.conf.set("spark.comet.sparkToColumnar.enabled", "true") + try: + # Spark 4.2's columnar Python input does not apply the row limit. + # Disable Comet so the reference uses Spark's row-based Arrow writer. + spark.conf.set("spark.comet.enabled", "false") + spark_source = spark.range(1, 5, 1, 1) + spark_reference = spark_source.select(batch_length("id")) + assert "Comet" not in _executed_plan(spark_reference) + spark_rows = spark_reference.collect() + + spark.conf.set("spark.comet.enabled", "true") + spark.conf.set("spark.comet.exec.nativeArrowPythonUDF.enabled", "true") + source = spark.range(1, 5, 1, 1) + result = source.select(batch_length("id")) + assert "CometArrowEvalPython" in _executed_plan(result) + assert result.collect() == spark_rows + assert [row[0] for row in spark_rows] == [2, 2, 2, 2] + + spark.conf.set("spark.comet.enabled", "false") + spark_strings = spark.range(1, 5, 1, 1).selectExpr("cast(id as string) as value") + spark_hashes = spark_strings.select(string_hash("value")).collect() + spark.conf.set("spark.comet.enabled", "true") + strings = source.selectExpr("cast(id as string) as value") + native_hashes = strings.select(string_hash("value")) + assert "CometArrowEvalPython" in _executed_plan(native_hashes) + assert native_hashes.collect() == spark_hashes + finally: + spark.conf.set("spark.comet.exec.nativeArrowPythonUDF.enabled", "false") + spark.conf.set("spark.comet.enabled", previous_comet_enabled) + spark.conf.unset("spark.sql.execution.arrow.maxRecordsPerBatch") + spark.conf.unset("spark.comet.sparkToColumnar.enabled") + spark.conf.unset("spark.sql.adaptive.enabled") + + +def test_scalar_arrow_udf_respects_spark_byte_batch_limit(spark): + from pyspark.sql.pandas import functions as pandas_functions + + if not hasattr(pandas_functions, "arrow_udf"): + pytest.skip("scalar arrow_udf requires Spark 4.1 or later") + + @pandas_functions.arrow_udf("long") + def checked_batch_size(values): + if values.nbytes > 16: + raise ValueError(f"configured byte cap exceeded: {values.nbytes} > 16") + return pa.array([values.nbytes] * len(values), type=pa.int64()) + + previous_comet_enabled = spark.conf.get("spark.comet.enabled") + spark.conf.set("spark.sql.adaptive.enabled", "false") + spark.conf.set("spark.sql.execution.arrow.maxRecordsPerBatch", "10000") + spark.conf.set("spark.sql.execution.arrow.maxBytesPerBatch", "16") + spark.conf.set("spark.comet.sparkToColumnar.enabled", "true") + try: + spark.conf.set("spark.comet.enabled", "false") + spark_reference = spark.range(1, 5, 1, 1).select(checked_batch_size("id")) + assert "Comet" not in _executed_plan(spark_reference) + spark_rows = spark_reference.collect() + + spark.conf.set("spark.comet.enabled", "true") + spark.conf.set("spark.comet.exec.nativeArrowPythonUDF.enabled", "true") + native_rows = spark.range(1, 5, 1, 1).select(checked_batch_size("id")) + assert "CometArrowEvalPython" in _executed_plan(native_rows) + assert native_rows.collect() == spark_rows + assert [row[0] for row in spark_rows] == [16, 16, 16, 16] + finally: + spark.conf.set("spark.comet.exec.nativeArrowPythonUDF.enabled", "false") + spark.conf.set("spark.comet.enabled", previous_comet_enabled) + spark.conf.unset("spark.sql.execution.arrow.maxRecordsPerBatch") + spark.conf.unset("spark.sql.execution.arrow.maxBytesPerBatch") + spark.conf.unset("spark.comet.sparkToColumnar.enabled") + spark.conf.unset("spark.sql.adaptive.enabled") + + +def test_scalar_arrow_udf_accumulator_uses_spark_worker(spark): + from pyspark.sql.pandas import functions as pandas_functions + + if not hasattr(pandas_functions, "arrow_udf"): + pytest.skip("scalar arrow_udf requires Spark 4.1 or later") + + def counting_udf(counter): + @pandas_functions.arrow_udf("long") + def count_rows(values): + counter.add(len(values)) + return values + + return count_rows + + previous_comet_enabled = spark.conf.get("spark.comet.enabled") + spark.conf.set("spark.sql.adaptive.enabled", "false") + spark.conf.set("spark.comet.sparkToColumnar.enabled", "true") + try: + spark.conf.set("spark.comet.enabled", "false") + spark_counter = spark.sparkContext.accumulator(0) + spark_result = spark.range(4, numPartitions=2).select( + counting_udf(spark_counter)("id") + ) + expected = spark_result.collect() + assert spark_counter.value == 4 + + spark.conf.set("spark.comet.enabled", "true") + spark.conf.set("spark.comet.exec.nativeArrowPythonUDF.enabled", "true") + comet_counter = spark.sparkContext.accumulator(0) + result = spark.range(4, numPartitions=2).select( + counting_udf(comet_counter)("id") + ) + plan = _executed_plan(result) + assert "ArrowEvalPython" in plan + assert "CometArrowEvalPython" not in plan + assert result.collect() == expected + assert comet_counter.value == 4 + finally: + spark.conf.set("spark.comet.exec.nativeArrowPythonUDF.enabled", "false") + spark.conf.set("spark.comet.enabled", previous_comet_enabled) + spark.conf.unset("spark.comet.sparkToColumnar.enabled") + spark.conf.unset("spark.sql.adaptive.enabled") + + +def test_scalar_arrow_udf_skips_empty_input_batches(spark): + from pyspark.sql.pandas import functions as pandas_functions + + if not hasattr(pandas_functions, "arrow_udf"): + pytest.skip("scalar arrow_udf requires Spark 4.1 or later") + + @pandas_functions.arrow_udf("long") + def reject_empty(values): + if len(values) == 0: + raise ValueError("scalar Arrow UDF received an empty batch") + return values + + previous_comet_enabled = spark.conf.get("spark.comet.enabled") + spark.conf.set("spark.sql.adaptive.enabled", "false") + spark.conf.set("spark.comet.sparkToColumnar.enabled", "true") + try: + spark.conf.set("spark.comet.enabled", "false") + spark_reference = spark.range(1, 5, 1, 1).sample(False, 0.01, 42) + assert spark_reference.select(reject_empty("id")).collect() == [] + + spark.conf.set("spark.comet.enabled", "true") + spark.conf.set("spark.comet.exec.nativeArrowPythonUDF.enabled", "true") + native_source = spark.range(1, 5, 1, 1).sample(False, 0.01, 42) + native_result = native_source.select(reject_empty("id")) + assert "CometArrowEvalPython" in _executed_plan(native_result) + assert native_result.collect() == [] + finally: + spark.conf.set("spark.comet.exec.nativeArrowPythonUDF.enabled", "false") + spark.conf.set("spark.comet.enabled", previous_comet_enabled) + spark.conf.unset("spark.comet.sparkToColumnar.enabled") + spark.conf.unset("spark.sql.adaptive.enabled") + + def test_map_in_arrow_doubles_value(spark, tmp_path, accelerated): data = [(i, float(i * 1.5), f"name_{i}") for i in range(100)] src = str(tmp_path / "src.parquet") diff --git a/spark/src/test/scala/org/apache/spark/sql/benchmark/CometBenchmarkBase.scala b/spark/src/test/scala/org/apache/spark/sql/benchmark/CometBenchmarkBase.scala index ea0cdb690f1..c6ea9db91f2 100644 --- a/spark/src/test/scala/org/apache/spark/sql/benchmark/CometBenchmarkBase.scala +++ b/spark/src/test/scala/org/apache/spark/sql/benchmark/CometBenchmarkBase.scala @@ -50,7 +50,7 @@ trait CometBenchmarkBase val conf = new SparkConf() .setAppName("CometReadBenchmark") // Since `spark.master` always exists, overrides this value - .set("spark.master", "local[1]") + .set("spark.master", sys.env.getOrElse("COMET_BENCHMARK_MASTER", "local[1]")) .setIfMissing("spark.driver.memory", "3g") .setIfMissing("spark.executor.memory", "3g") // Use Comet's shuffle manager so operators that require Comet shuffle can diff --git a/spark/src/test/spark-4.1+/org/apache/spark/sql/benchmark/CometArrowPythonUdfBenchmark.scala b/spark/src/test/spark-4.1+/org/apache/spark/sql/benchmark/CometArrowPythonUdfBenchmark.scala new file mode 100644 index 00000000000..707dae5d908 --- /dev/null +++ b/spark/src/test/spark-4.1+/org/apache/spark/sql/benchmark/CometArrowPythonUdfBenchmark.scala @@ -0,0 +1,133 @@ +/* + * Licensed to the Apache Software Foundation (ASF) under one + * or more contributor license agreements. See the NOTICE file + * distributed with this work for additional information + * regarding copyright ownership. The ASF licenses this file + * to you under the Apache License, Version 2.0 (the + * "License"); you may not use this file except in compliance + * with the License. You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, + * software distributed under the License is distributed on an + * "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY + * KIND, either express or implied. See the License for the + * specific language governing permissions and limitations + * under the License. + */ + +package org.apache.spark.sql.benchmark + +import java.util.{Base64, Collections} + +import scala.sys.process._ + +import org.apache.spark.api.python.{PythonEvalType, SimplePythonFunction} +import org.apache.spark.sql.execution.python.UserDefinedPythonFunction +import org.apache.spark.sql.functions.sum +import org.apache.spark.sql.internal.SQLConf +import org.apache.spark.sql.types.LongType + +import org.apache.comet.{CometConf, NativeBase} + +/** + * End-to-end benchmark of Spark and native scalar Arrow UDF execution. Run with `make + * benchmark-org.apache.spark.sql.benchmark.CometArrowPythonUdfBenchmark PROFILES=-Pspark-4.1 + * COMET_FEATURES=python-udf`, with PYO3_PYTHON set for the native build. Arguments are rows, + * warmups, iterations, partitions, and mode (`arrow` or `python`). The Python mode performs a + * per-element Python loop to expose GIL contention. Set `COMET_BENCHMARK_MASTER=local[16]` and + * request at least 17 partitions for a concurrency comparison. Set PYSPARK_PYTHON and PYTHONPATH + * for the embedded Python environment. + */ +object CometArrowPythonUdfBenchmark extends CometBenchmarkBase { + override def runCometBenchmark(args: Array[String]): Unit = { + require(NativeBase.supportsPythonUdf(), "native library was built without python-udf") + + val rows = args.headOption.map(_.toLong).getOrElse(1000000L) + val warmups = args.lift(1).map(_.toInt).getOrElse(2) + val iterations = args.lift(2).map(_.toInt).getOrElse(5) + val partitions = args.lift(3).map(_.toInt).getOrElse(2) + val mode = args.lift(4).getOrElse("arrow") + require(Set("arrow", "python").contains(mode), s"Unknown benchmark mode: $mode") + val python = sys.env.getOrElse("PYSPARK_PYTHON", "python3") + val callable = + if (mode == "arrow") "pc.negate" else "lambda a: pa.array([-x.as_py() for x in a])" + val code = + "import base64, pyspark.cloudpickle as cloudpickle, pyarrow as pa, " + + "pyarrow.compute as pc; " + + "from pyspark.sql.types import LongType; " + + s"print(base64.b64encode(cloudpickle.dumps(($callable, LongType()))).decode())" + val command = Base64.getDecoder.decode(Seq(python, "-c", code).!!.trim) + val pythonVersion = + Seq(python, "-c", "import sys; print('%d.%d' % sys.version_info[:2])").!!.trim + val function = SimplePythonFunction( + command, + Collections.emptyMap[String, String](), + Collections.emptyList[String](), + python, + pythonVersion, + Collections.emptyList(), + null) + val udf = UserDefinedPythonFunction( + "negate_arrow", + function, + LongType, + PythonEvalType.SQL_SCALAR_ARROW_UDF, + udfDeterministic = true) + + val samples = + scala.collection.mutable.Map(false -> Vector.empty[Double], true -> Vector.empty[Double]) + val expected = -rows * (rows - 1L) / 2L + + val configs = Seq( + CometConf.COMET_ENABLED.key -> "true", + CometConf.COMET_EXEC_ENABLED.key -> "true", + CometConf.COMET_ONHEAP_ENABLED.key -> "true", + CometConf.COMET_SPARK_TO_ARROW_ENABLED.key -> "true", + SQLConf.ADAPTIVE_EXECUTION_ENABLED.key -> "false", + SQLConf.ARROW_EXECUTION_MAX_RECORDS_PER_BATCH.key -> "10000") + withSQLConf(configs: _*) { + println( + s"ARROW_UDF_BENCHMARK spark=${spark.version} rows=$rows warmups=$warmups " + + s"iterations=$iterations partitions=$partitions mode=$mode") + for (iteration <- 0 until warmups + iterations; native <- Seq(false, true)) { + withSQLConf(CometConf.COMET_NATIVE_ARROW_PYTHON_UDF_ENABLED.key -> native.toString) { + val source = spark.range(0L, rows, 1L, partitions) + val df = source.select(udf(source.col("id")).as("value")).agg(sum("value")) + val plan = df.queryExecution.executedPlan.toString() + if (native) { + assert(plan.contains("CometArrowEvalPython"), plan) + } else { + assert(plan.contains("ArrowEvalPython"), plan) + assert(!plan.contains("CometArrowEvalPython"), plan) + } + if (iteration == 0) { + println(s"ARROW_UDF_PLAN native=$native\n$plan") + } + val start = System.nanoTime() + val actual = df.collect().head.getLong(0) + val seconds = (System.nanoTime() - start).toDouble / 1e9 + assert(actual == expected, s"native=$native: $actual != $expected") + if (iteration >= warmups) { + samples(native) :+= seconds + println( + s"ARROW_UDF_SAMPLE native=$native iteration=${iteration - warmups} seconds=$seconds") + } + } + } + } + + def median(values: Seq[Double]): Double = { + val sorted = values.sorted + (sorted((sorted.size - 1) / 2) + sorted(sorted.size / 2)) / 2.0 + } + if (iterations > 0) { + val sparkMedian = median(samples(false)) + val nativeMedian = median(samples(true)) + println( + s"ARROW_UDF_RESULT spark_median_seconds=$sparkMedian " + + s"native_median_seconds=$nativeMedian speedup=${sparkMedian / nativeMedian}") + } + } +} diff --git a/spark/src/test/spark-4.1+/org/apache/spark/sql/comet/CometArrowPythonUdfSuite.scala b/spark/src/test/spark-4.1+/org/apache/spark/sql/comet/CometArrowPythonUdfSuite.scala new file mode 100644 index 00000000000..524fdc6ef52 --- /dev/null +++ b/spark/src/test/spark-4.1+/org/apache/spark/sql/comet/CometArrowPythonUdfSuite.scala @@ -0,0 +1,333 @@ +/* + * Licensed to the Apache Software Foundation (ASF) under one + * or more contributor license agreements. See the NOTICE file + * distributed with this work for additional information + * regarding copyright ownership. The ASF licenses this file + * to you under the Apache License, Version 2.0 (the + * "License"); you may not use this file except in compliance + * with the License. You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, + * software distributed under the License is distributed on an + * "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY + * KIND, either express or implied. See the License for the + * specific language governing permissions and limitations + * under the License. + */ + +package org.apache.spark.sql.comet + +import java.util.{Base64, Collections} + +import scala.sys.process._ + +import org.apache.spark.api.python.{PythonEvalType, SimplePythonFunction} +import org.apache.spark.sql.{CometTestBase, Row} +import org.apache.spark.sql.execution.python.UserDefinedPythonFunction +import org.apache.spark.sql.functions.{array, expr, lit, map, struct, when} +import org.apache.spark.sql.internal.SQLConf +import org.apache.spark.sql.types.{ArrayType, BinaryType, BooleanType, ByteType, CalendarIntervalType, DataType, DateType, DecimalType, DoubleType, FloatType, IntegerType, LongType, MapType, ShortType, StringType, StructField, StructType, TimestampNTZType, TimestampType, TimeType, VariantType, YearMonthIntervalType} + +import org.apache.comet.{CometConf, NativeBase} + +class CometArrowPythonUdfSuite extends CometTestBase { + + test("scalar Arrow UDF falls back when the native feature is unavailable") { + assume(!NativeBase.supportsPythonUdf()) + + val function = SimplePythonFunction( + Array.emptyByteArray, + Collections.emptyMap[String, String](), + Collections.emptyList[String](), + "python3", + "3.13", + Collections.emptyList(), + null) + val udf = UserDefinedPythonFunction( + "arrow_udf", + function, + LongType, + PythonEvalType.SQL_SCALAR_ARROW_UDF, + udfDeterministic = true) + + withSQLConf(CometConf.COMET_NATIVE_ARROW_PYTHON_UDF_ENABLED.key -> "true") { + val source = spark.range(1, 2) + val plan = source.select(udf(source.col("id"))).queryExecution.executedPlan + assert(plan.collect { case _: CometArrowEvalPythonExec => true }.isEmpty) + } + } + + test("scalar Arrow UDF executes in the native pipeline") { + assume(NativeBase.supportsPythonUdf(), "native library was built without python-udf") + + val python = sys.env.getOrElse("PYSPARK_PYTHON", "python3") + val code = + "import base64, pyspark.cloudpickle as cloudpickle, pyarrow.compute as pc; " + + "from pyspark.sql.types import LongType; " + + "print(base64.b64encode(cloudpickle.dumps((pc.negate, LongType()))).decode())" + val command = Base64.getDecoder.decode(Seq(python, "-c", code).!!.trim) + val pythonVersion = + Seq(python, "-c", "import sys; print('%d.%d' % sys.version_info[:2])").!!.trim + val function = SimplePythonFunction( + command, + Collections.emptyMap[String, String](), + Collections.emptyList[String](), + python, + pythonVersion, + Collections.emptyList(), + null) + val udf = UserDefinedPythonFunction( + "negate_arrow", + function, + LongType, + PythonEvalType.SQL_SCALAR_ARROW_UDF, + udfDeterministic = true) + + withSQLConf( + CometConf.COMET_NATIVE_ARROW_PYTHON_UDF_ENABLED.key -> "true", + SQLConf.ADAPTIVE_EXECUTION_ENABLED.key -> "false") { + val source = spark.range(1, 5) + val df = source.select(udf(source.col("id"))) + assert(df.queryExecution.executedPlan.collect { case _: CometArrowEvalPythonExec => + true + }.nonEmpty) + checkAnswer(df, Seq(Row(-1L), Row(-2L), Row(-3L), Row(-4L))) + + val twoResults = + source.select(udf(source.col("id")).as("first"), udf(source.col("id") + 1L).as("second")) + val nativeUdfs = twoResults.queryExecution.executedPlan.collect { + case op: CometArrowEvalPythonExec => op + } + assert(nativeUdfs.exists(_.nativeOp.getArrowPythonUdf.getFunctionsCount == 2)) + checkAnswer(twoResults, Seq(Row(-1L, -2L), Row(-2L, -3L), Row(-3L, -4L), Row(-4L, -5L))) + + val withSubquery = spark.range(4).select(udf(expr("id + (SELECT max(id) FROM range(8))"))) + val subqueryPlan = withSubquery.queryExecution.executedPlan + assert(subqueryPlan.collect { case _: CometArrowEvalPythonExec => true }.nonEmpty) + checkAnswer(withSubquery, Seq(Row(-7L), Row(-8L), Row(-9L), Row(-10L))) + } + + withSQLConf(CometConf.COMET_NATIVE_ARROW_PYTHON_UDF_ENABLED.key -> "false") { + val source = spark.range(1, 2) + val plan = source.select(udf(source.col("id"))).queryExecution.executedPlan + assert(plan.collect { case _: CometArrowEvalPythonExec => true }.isEmpty) + } + + withSQLConf( + CometConf.COMET_NATIVE_ARROW_PYTHON_UDF_ENABLED.key -> "true", + SQLConf.ARROW_EXECUTION_USE_LARGE_VAR_TYPES.key -> "true") { + val source = spark.range(1, 2) + val plan = source.select(udf(source.col("id"))).queryExecution.executedPlan + assert(plan.collect { case _: CometArrowEvalPythonExec => true }.isEmpty) + } + } + + test("native Arrow UDF plan hides commands and compares UDF identity without plan IDs") { + assume(NativeBase.supportsPythonUdf(), "native library was built without python-udf") + + val secret = "private_arrow_udf_command" + val function = SimplePythonFunction( + secret.getBytes(java.nio.charset.StandardCharsets.UTF_8), + Collections.emptyMap[String, String](), + Collections.emptyList[String](), + "python3", + "3.13", + Collections.emptyList(), + null) + val udf = UserDefinedPythonFunction( + "secret_arrow", + function, + LongType, + PythonEvalType.SQL_SCALAR_ARROW_UDF, + udfDeterministic = true) + + withSQLConf( + CometConf.COMET_NATIVE_ARROW_PYTHON_UDF_ENABLED.key -> "true", + SQLConf.ADAPTIVE_EXECUTION_ENABLED.key -> "false") { + val source = spark.range(1, 2) + val plan = source.select(udf(source.col("id"))).queryExecution.executedPlan + val native = plan.collectFirst { case op: CometArrowEvalPythonExec => op }.get + assert(!plan.treeString.contains(secret)) + assert(!native.toString.contains(secret)) + + val differentPlanId = native.copy(nativeOp = + native.nativeOp.toBuilder.setPlanId(native.nativeOp.getPlanId + 1).build()) + assert(native == differentPlanId) + assert(native.hashCode() == differentPlanId.hashCode()) + + val otherFunction = SimplePythonFunction( + "different_arrow_udf_command".getBytes(java.nio.charset.StandardCharsets.UTF_8), + Collections.emptyMap[String, String](), + Collections.emptyList[String](), + "python3", + "3.13", + Collections.emptyList(), + null) + val otherUdf = UserDefinedPythonFunction( + "secret_arrow", + otherFunction, + LongType, + PythonEvalType.SQL_SCALAR_ARROW_UDF, + udfDeterministic = true) + val otherPlan = source.select(otherUdf(source.col("id"))).queryExecution.executedPlan + val otherNative = otherPlan.collectFirst { case op: CometArrowEvalPythonExec => op }.get + assert(native.copy(udfs = otherNative.udfs) != native) + + val differentBatchSize = native.nativeOp.toBuilder + differentBatchSize.getArrowPythonUdfBuilder.setMaxRecordsPerBatch( + native.nativeOp.getArrowPythonUdf.getMaxRecordsPerBatch + 1) + assert(native.copy(nativeOp = differentBatchSize.build()) != native) + + val sameFunctionPlan = source.select(udf(source.col("id"))).queryExecution.executedPlan + val sameFunctionNative = + sameFunctionPlan.collectFirst { case op: CometArrowEvalPythonExec => op }.get + assert(native.udfs != sameFunctionNative.udfs) + assert(native.sameResult(sameFunctionNative)) + } + } + + test("native Arrow UDF preserves every accepted scalar type and nulls") { + assume(NativeBase.supportsPythonUdf(), "native library was built without python-udf") + + val python = sys.env.getOrElse("PYSPARK_PYTHON", "python3") + val pythonVersion = + Seq(python, "-c", "import sys; print('%d.%d' % sys.version_info[:2])").!!.trim + val cases: Seq[(DataType, String, String)] = Seq( + (BooleanType, "BooleanType()", "true"), + (ByteType, "ByteType()", "7"), + (ShortType, "ShortType()", "123"), + (IntegerType, "IntegerType()", "1234"), + (LongType, "LongType()", "12345"), + (FloatType, "FloatType()", "1.25"), + (DoubleType, "DoubleType()", "1.25"), + (StringType, "StringType()", "hello"), + (BinaryType, "BinaryType()", "hello"), + (DecimalType(12, 2), "DecimalType(12, 2)", "12.34"), + (DateType, "DateType()", "2024-01-02"), + (TimestampNTZType, "TimestampNTZType()", "2024-01-02 03:04:05.123456")) + + withSQLConf(SQLConf.ADAPTIVE_EXECUTION_ENABLED.key -> "false") { + val source = spark.range(2) + cases.foreach { case (dataType, pythonType, value) => + val code = + "import base64, pyspark.cloudpickle as cloudpickle, pyarrow as pa; " + + "from pyspark.sql.types import *; " + + s"print(base64.b64encode(cloudpickle.dumps((lambda a: a, $pythonType))).decode()); " + + "print(base64.b64encode(cloudpickle.dumps((" + + "lambda a: pa.array([str(a.type)] * len(a)), StringType()))).decode())" + val commands = + Seq(python, "-c", code).!!.trim.linesIterator.map(Base64.getDecoder.decode).toSeq + def arrowUdf(name: String, command: Array[Byte], returnType: DataType) = { + val function = SimplePythonFunction( + command, + Collections.emptyMap[String, String](), + Collections.emptyList[String](), + python, + pythonVersion, + Collections.emptyList(), + null) + UserDefinedPythonFunction( + name, + function, + returnType, + PythonEvalType.SQL_SCALAR_ARROW_UDF, + udfDeterministic = true) + } + val identity = arrowUdf("identity_arrow", commands.head, dataType) + val describeType = arrowUdf("arrow_input_type", commands(1), StringType) + val input = source.select( + when(source.col("id") === 1L, lit(null).cast(dataType)) + .otherwise(lit(value).cast(dataType)) + .as("value")) + val expected = + withSQLConf(CometConf.COMET_NATIVE_ARROW_PYTHON_UDF_ENABLED.key -> "false") { + input + .select(identity(input.col("value")), describeType(input.col("value"))) + .collect() + .toSeq + } + withSQLConf(CometConf.COMET_NATIVE_ARROW_PYTHON_UDF_ENABLED.key -> "true") { + val df = input.select(identity(input.col("value")), describeType(input.col("value"))) + assert( + df.queryExecution.executedPlan.collect { case _: CometArrowEvalPythonExec => + true + }.nonEmpty, + s"Native Arrow UDF was not selected for $dataType") + checkAnswer(df, expected) + } + } + } + } + + test("native Arrow UDF falls back for schemas and options with different Spark semantics") { + assume(NativeBase.supportsPythonUdf(), "native library was built without python-udf") + + val function = SimplePythonFunction( + Array.emptyByteArray, + Collections.emptyMap[String, String](), + Collections.emptyList[String](), + "python3", + "3.11", + Collections.emptyList(), + null) + def arrowUdf(returnType: org.apache.spark.sql.types.DataType) = UserDefinedPythonFunction( + "planning_only", + function, + returnType, + PythonEvalType.SQL_SCALAR_ARROW_UDF, + udfDeterministic = true) + + withSQLConf( + CometConf.COMET_NATIVE_ARROW_PYTHON_UDF_ENABLED.key -> "true", + SQLConf.ADAPTIVE_EXECUTION_ENABLED.key -> "false", + SQLConf.SESSION_LOCAL_TIMEZONE.key -> "America/Los_Angeles") { + val source = spark.range(1) + val timestamp = source.col("id").cast(TimestampType) + val plans = Seq( + source.select(arrowUdf(LongType)(timestamp)), + source.select(arrowUdf(LongType)(array(source.col("id")))), + source.select(arrowUdf(LongType)(map(source.col("id"), source.col("id")))), + source.select(arrowUdf(LongType)(struct(source.col("id")))), + source.select(arrowUdf(ArrayType(LongType))(source.col("id"))), + source.select(arrowUdf(MapType(LongType, LongType))(source.col("id"))), + source.select( + arrowUdf(StructType(Seq(StructField("value", LongType))))(source.col("id"))), + source.select(arrowUdf(TimestampType)(source.col("id"))), + source.select(arrowUdf(TimeType(6))(source.col("id"))), + source.select(arrowUdf(VariantType)(source.col("id"))), + source.select(arrowUdf(YearMonthIntervalType())(source.col("id"))), + source.select(arrowUdf(CalendarIntervalType)(source.col("id")))) + plans.foreach { df => + val plan = df.queryExecution.executedPlan + assert(plan.collect { case _: CometArrowEvalPythonExec => true }.isEmpty) + } + withSQLConf("spark.sql.pyspark.udf.profiler" -> "perf") { + val plan = source.select(arrowUdf(LongType)(source.col("id"))).queryExecution.executedPlan + assert(plan.collect { case _: CometArrowEvalPythonExec => true }.isEmpty) + } + + Seq( + Collections.singletonMap("PYTHONHASHSEED", "123"), + Collections.singletonMap("CUSTOM_PYTHON_SETTING", "value")).foreach { env => + val withEnvironment = SimplePythonFunction( + Array.emptyByteArray, + env, + Collections.emptyList[String](), + "python3", + "3.11", + Collections.emptyList(), + null) + val udf = UserDefinedPythonFunction( + "planning_only", + withEnvironment, + LongType, + PythonEvalType.SQL_SCALAR_ARROW_UDF, + udfDeterministic = true) + val plan = source.select(udf(source.col("id"))).queryExecution.executedPlan + assert(plan.collect { case _: CometArrowEvalPythonExec => true }.isEmpty) + } + } + } +}