| # |
| # 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. |
| # |
| """ |
| User-defined function related classes and functions |
| """ |
| |
| import functools |
| import inspect |
| import sys |
| import warnings |
| from typing import TYPE_CHECKING, Any, Callable, Optional, Union, cast |
| |
| from pyspark.errors import PySparkNotImplementedError, PySparkRuntimeError, PySparkTypeError |
| from pyspark.sql.column import Column |
| from pyspark.sql.pandas.types import to_arrow_type |
| from pyspark.sql.pandas.utils import require_minimum_pandas_version, require_minimum_pyarrow_version |
| from pyspark.sql.types import ( |
| DataType, |
| StringType, |
| StructType, |
| _parse_datatype_string, |
| ) |
| from pyspark.sql.utils import get_active_spark_context |
| from pyspark.util import PythonEvalType |
| |
| if TYPE_CHECKING: |
| from py4j.java_gateway import JavaObject |
| |
| from pyspark.core.context import SparkContext |
| from pyspark.sql._typing import ColumnOrName, DataTypeOrString, UserDefinedFunctionLike |
| from pyspark.sql.session import SparkSession |
| |
| __all__ = ["UDFRegistration"] |
| |
| |
| def _wrap_function( |
| sc: "SparkContext", func: Callable[..., Any], returnType: Optional[DataType] = None |
| ) -> "JavaObject": |
| from pyspark.core.rdd import _prepare_for_python_RDD |
| |
| command: Any |
| if returnType is None: |
| command = func |
| else: |
| command = (func, returnType) |
| pickled_command, broadcast_vars, env, includes = _prepare_for_python_RDD(sc, command) |
| assert sc._jvm is not None |
| return sc._jvm.SimplePythonFunction( |
| bytearray(pickled_command), |
| env, |
| includes, |
| sc.pythonExec, |
| sc.pythonVer, |
| broadcast_vars, |
| sc._javaAccumulator, |
| ) |
| |
| |
| def _create_udf( |
| f: Callable[..., Any], |
| returnType: "DataTypeOrString", |
| evalType: int, |
| name: Optional[str] = None, |
| deterministic: bool = True, |
| bufferSchema: Optional[StructType] = None, |
| ) -> "UserDefinedFunctionLike": |
| """Create a regular(non-Arrow-optimized) Python UDF.""" |
| # Set the name of the UserDefinedFunction object to be the name of function f |
| udf_obj = UserDefinedFunction( |
| f, |
| returnType=returnType, |
| name=name, |
| evalType=evalType, |
| deterministic=deterministic, |
| bufferSchema=bufferSchema, |
| ) |
| return udf_obj._wrapped() |
| |
| |
| def _create_py_udf( |
| f: Callable[..., Any], |
| returnType: "DataTypeOrString", |
| useArrow: Optional[bool] = None, |
| ) -> "UserDefinedFunctionLike": |
| """Create a regular/Arrow-optimized Python UDF.""" |
| # The tables in python/pyspark/sql/tests/udf_type_tests show the results when the type coercion |
| # in Arrow is needed, that is, when the user-specified return type(SQL Type) of the UDF and the |
| # actual instance(Python Value(Type)) that the UDF returns are different. |
| # Arrow and Pickle have different type coercion rules, so a UDF might have a different result |
| # with/without Arrow optimization. That's the main reason the Arrow optimization for Python |
| # UDFs is disabled by default. |
| |
| is_arrow_enabled = False |
| if useArrow is None: |
| from pyspark.sql import SparkSession |
| |
| session = SparkSession._instantiatedSession |
| is_arrow_enabled = ( |
| False |
| if session is None |
| else session.conf.get("spark.sql.execution.pythonUDF.arrow.enabled") == "true" |
| ) |
| else: |
| is_arrow_enabled = useArrow |
| |
| if is_arrow_enabled: |
| try: |
| require_minimum_pandas_version() |
| require_minimum_pyarrow_version() |
| except ImportError: |
| is_arrow_enabled = False |
| warnings.warn( |
| "Arrow optimization failed to enable because PyArrow or Pandas is not installed. " |
| "Falling back to a non-Arrow-optimized UDF.", |
| RuntimeWarning, |
| ) |
| |
| eval_type: Optional[int] = None |
| if useArrow is None: |
| # If the user doesn't explicitly set useArrow |
| from pyspark.sql.pandas.typehints import infer_eval_type_for_udf |
| |
| try: |
| # Try to infer the eval type from type hints |
| eval_type = infer_eval_type_for_udf(f) |
| except Exception: |
| warnings.warn("Cannot infer the eval type from type hints. ", UserWarning) |
| |
| if eval_type is None: |
| if is_arrow_enabled: |
| # Arrow optimized Python UDF |
| eval_type = PythonEvalType.SQL_ARROW_BATCHED_UDF |
| else: |
| # Fallback to Regular Python UDF |
| eval_type = PythonEvalType.SQL_BATCHED_UDF |
| |
| return _create_udf(f, returnType, eval_type) |
| |
| |
| class UserDefinedFunction: |
| """ |
| User defined function in Python |
| |
| .. versionadded:: 1.3 |
| |
| Notes |
| ----- |
| The constructor of this class is not supposed to be directly called. |
| Use :meth:`pyspark.sql.functions.udf` or :meth:`pyspark.sql.functions.pandas_udf` |
| to create this instance. |
| """ |
| |
| def __init__( |
| self, |
| func: Callable[..., Any], |
| returnType: "DataTypeOrString" = StringType(), |
| name: Optional[str] = None, |
| evalType: int = PythonEvalType.SQL_BATCHED_UDF, |
| deterministic: bool = True, |
| bufferSchema: Optional[StructType] = None, |
| ): |
| if not callable(func): |
| raise PySparkTypeError( |
| errorClass="NOT_EXPECTED_TYPE", |
| messageParameters={ |
| "expected_type": "callable", |
| "arg_name": "func", |
| "arg_type": type(func).__name__, |
| }, |
| ) |
| |
| if not isinstance(returnType, (DataType, str)): |
| raise PySparkTypeError( |
| errorClass="NOT_EXPECTED_TYPE", |
| messageParameters={ |
| "expected_type": "DataType or str", |
| "arg_name": "returnType", |
| "arg_type": type(returnType).__name__, |
| }, |
| ) |
| |
| if not isinstance(evalType, int): |
| raise PySparkTypeError( |
| errorClass="NOT_EXPECTED_TYPE", |
| messageParameters={ |
| "expected_type": "int", |
| "arg_name": "evalType", |
| "arg_type": type(evalType).__name__, |
| }, |
| ) |
| |
| self.func = func |
| self._returnType = returnType |
| # Stores UserDefinedPythonFunctions jobj, once initialized |
| self._returnType_placeholder: Optional[DataType] = None |
| self._judf_placeholder = None |
| self._name = name or ( |
| func.__name__ if hasattr(func, "__name__") else func.__class__.__name__ |
| ) |
| self.evalType = evalType |
| self.deterministic = deterministic |
| # Schema of the intermediate aggregation buffer, set only for an incremental Python |
| # aggregator (see :class:`pyspark.sql.aggregator.Aggregator`); ``None`` otherwise. It is a |
| # first-class field so it survives reconstruction paths such as ``_wrapped()``, |
| # ``asNondeterministic()`` and ``spark.udf.register``, and is threaded to the JVM in |
| # ``_create_judf`` so ``PythonAggregate`` can plan the two-stage aggregation. |
| self.bufferSchema = bufferSchema |
| # Extract Python UDF details if transpilation is enabled. |
| self.transpiled: list = [] |
| self._transpiled_param_names: list[str] = [] |
| # Per-option input-type categories ("numeric"/"string" per public param), |
| # parallel to ``self.transpiled``; the JVM picks the option matching the |
| # actual column types or falls back to interpreted Python. |
| self._transpiled_input_categories: list = [] |
| # When we have a transpiled rewrite, ``__call__`` resolves any |
| # user-supplied kwargs against this positional parameter list so |
| # the JVM-side ``_udf_param_N`` substitution sees the inputs in |
| # the right order. Empty list when transpilation didn't happen. |
| from pyspark.sql import SparkSession |
| |
| session = SparkSession._instantiatedSession |
| |
| # A nondeterministic UDF must not be transpiled: replacing it with a plain |
| # Catalyst expression would let the optimizer fold/reorder/duplicate it, |
| # discarding the nondeterminism barrier. (asNondeterministic() also clears |
| # any options set here, for the udf(f).asNondeterministic() ordering.) |
| # Conf values are compared case-insensitively: `SET conf=True` stores |
| # the literal "True", which would otherwise silently disable |
| # transpilation (or mis-trigger the ANSI warning below). |
| # |
| # Each conf read is a JVM roundtrip, so keep the default construction |
| # path cheap: the experimental gate is only read for deterministic |
| # batched UDFs (the only shape we transpile), and the ANSI conf is only |
| # read once the gate is known to be on. When ``default`` is given it is |
| # passed through to ``RuntimeConfig.get`` so construction never depends |
| # on the JVM having the (experimental) conf registered -- e.g. a newer |
| # Python client against an older driver. No default is passed for |
| # ``spark.sql.ansi.enabled``: its registered default is dynamic |
| # (environment-driven) and must be respected when the key is unset. |
| def _conf_is_true(key: str, default: Optional[str] = None) -> bool: |
| if session is None: |
| return False |
| if default is None: |
| value = session.conf.get(key) |
| else: |
| value = session.conf.get(key, default) |
| return value is not None and value.lower() == "true" |
| |
| try: |
| transpile_enabled = ( |
| deterministic |
| and evalType == PythonEvalType.SQL_BATCHED_UDF |
| and _conf_is_true("spark.sql.experimental.optimizer.transpilePyUDFs", "false") |
| ) |
| # Transpilation only attempts to reproduce ANSI-mode Spark SQL |
| # semantics (no silent integer overflow, divide-by-zero raises, |
| # etc.). Running it against non-ANSI Spark would balloon the test |
| # matrix we'd have to maintain to verify Python-vs-SQL equivalence, |
| # so we gate on ANSI here and warn the user instead of trying to |
| # transpile in a mode we don't claim to support yet. |
| if transpile_enabled and not _conf_is_true("spark.sql.ansi.enabled"): |
| warnings.warn( |
| "Python UDF transpilation " |
| "(spark.sql.experimental.optimizer.transpilePyUDFs) is only " |
| "supported when ANSI mode is enabled " |
| "(spark.sql.ansi.enabled=true). Skipping transpilation for " |
| f"{func} -- enable ANSI mode or set transpilePyUDFs=false to " |
| "silence this warning.", |
| RuntimeWarning, |
| ) |
| transpile_enabled = False |
| if transpile_enabled and session: |
| # Import only if needed, also avoid circular import loops. |
| from pyspark.sql.transpile import _transpile_func |
| |
| # ``self.returnType`` parses (and caches) the declared return |
| # type; the transpiler needs the parsed form to decide whether |
| # the final Cast to it can resolve at all. The parse is reused |
| # later by ``_create_judf``, so this adds no extra JVM work. |
| ( |
| self.transpiled, |
| errors, |
| self._transpiled_param_names, |
| self._transpiled_input_categories, |
| ) = _transpile_func(session, func, self.returnType) |
| if not self.transpiled: |
| detail = f": {errors}" if errors else "" |
| warnings.warn(f"Unable to transpile UDF {func}{detail}") |
| except Exception as e: |
| # An inability to transpile must never break a working UDF -- fall |
| # back to interpreted Python execution and surface the failure as a |
| # warning so users can opt to investigate without losing their |
| # query. The conf reads above are included: a session whose JVM |
| # cannot answer them should degrade to "no transpilation", not |
| # break UDF definition. |
| warnings.warn(f"Exception transpiling UDF {func}: {e}") |
| self.transpiled = [] |
| self._transpiled_param_names = [] |
| self._transpiled_input_categories = [] |
| |
| @staticmethod |
| def _check_return_type(returnType: DataType, evalType: int) -> None: |
| if evalType == PythonEvalType.SQL_ARROW_BATCHED_UDF: |
| try: |
| to_arrow_type(returnType, timezone="UTC") |
| except TypeError: |
| raise PySparkNotImplementedError( |
| errorClass="NOT_IMPLEMENTED", |
| messageParameters={ |
| "feature": f"Invalid return type with Arrow-optimized Python UDF: " |
| f"{returnType}" |
| }, |
| ) |
| elif ( |
| evalType == PythonEvalType.SQL_SCALAR_PANDAS_UDF |
| or evalType == PythonEvalType.SQL_SCALAR_PANDAS_ITER_UDF |
| ): |
| try: |
| to_arrow_type(returnType, timezone="UTC") |
| except TypeError: |
| raise PySparkNotImplementedError( |
| errorClass="NOT_IMPLEMENTED", |
| messageParameters={ |
| "feature": f"Invalid return type with scalar Pandas UDFs: {returnType}" |
| }, |
| ) |
| elif ( |
| evalType == PythonEvalType.SQL_SCALAR_ARROW_UDF |
| or evalType == PythonEvalType.SQL_SCALAR_ARROW_ITER_UDF |
| ): |
| try: |
| to_arrow_type(returnType, timezone="UTC") |
| except TypeError: |
| raise PySparkNotImplementedError( |
| errorClass="NOT_IMPLEMENTED", |
| messageParameters={ |
| "feature": f"Invalid return type with scalar Arrow UDFs: {returnType}" |
| }, |
| ) |
| elif ( |
| evalType == PythonEvalType.SQL_GROUPED_MAP_PANDAS_UDF |
| or evalType == PythonEvalType.SQL_GROUPED_MAP_PANDAS_ITER_UDF |
| or evalType == PythonEvalType.SQL_GROUPED_MAP_PANDAS_UDF_WITH_STATE |
| ): |
| if isinstance(returnType, StructType): |
| try: |
| to_arrow_type(returnType, timezone="UTC") |
| except TypeError: |
| raise PySparkNotImplementedError( |
| errorClass="NOT_IMPLEMENTED", |
| messageParameters={ |
| "feature": f"Invalid return type with grouped map Pandas UDFs or " |
| f"at groupby.applyInPandas(WithState): {returnType}" |
| }, |
| ) |
| else: |
| raise PySparkTypeError( |
| errorClass="INVALID_RETURN_TYPE_FOR_PANDAS_UDF", |
| messageParameters={ |
| "eval_type": "SQL_GROUPED_MAP_PANDAS_UDF or " |
| "SQL_GROUPED_MAP_PANDAS_ITER_UDF or " |
| "SQL_GROUPED_MAP_PANDAS_UDF_WITH_STATE", |
| "return_type": str(returnType), |
| }, |
| ) |
| elif ( |
| evalType == PythonEvalType.SQL_MAP_PANDAS_ITER_UDF |
| or evalType == PythonEvalType.SQL_MAP_ARROW_ITER_UDF |
| ): |
| if isinstance(returnType, StructType): |
| try: |
| to_arrow_type(returnType, timezone="UTC") |
| except TypeError: |
| raise PySparkNotImplementedError( |
| errorClass="NOT_IMPLEMENTED", |
| messageParameters={ |
| "feature": f"Invalid return type in mapInPandas: {returnType}" |
| }, |
| ) |
| else: |
| raise PySparkTypeError( |
| errorClass="INVALID_RETURN_TYPE_FOR_PANDAS_UDF", |
| messageParameters={ |
| "eval_type": "SQL_MAP_PANDAS_ITER_UDF or SQL_MAP_ARROW_ITER_UDF", |
| "return_type": str(returnType), |
| }, |
| ) |
| elif ( |
| evalType == PythonEvalType.SQL_GROUPED_MAP_ARROW_UDF |
| or evalType == PythonEvalType.SQL_GROUPED_MAP_ARROW_ITER_UDF |
| ): |
| if isinstance(returnType, StructType): |
| try: |
| to_arrow_type(returnType, timezone="UTC") |
| except TypeError: |
| raise PySparkNotImplementedError( |
| errorClass="NOT_IMPLEMENTED", |
| messageParameters={ |
| "feature": "Invalid return type with grouped map Arrow UDFs or " |
| f"at groupby.applyInArrow: {returnType}" |
| }, |
| ) |
| else: |
| raise PySparkTypeError( |
| errorClass="INVALID_RETURN_TYPE_FOR_ARROW_UDF", |
| messageParameters={ |
| "eval_type": "SQL_GROUPED_MAP_ARROW_UDF or SQL_GROUPED_MAP_ARROW_ITER_UDF", |
| "return_type": str(returnType), |
| }, |
| ) |
| elif evalType == PythonEvalType.SQL_COGROUPED_MAP_PANDAS_UDF: |
| if isinstance(returnType, StructType): |
| try: |
| to_arrow_type(returnType, timezone="UTC") |
| except TypeError: |
| raise PySparkNotImplementedError( |
| errorClass="NOT_IMPLEMENTED", |
| messageParameters={ |
| "feature": f"Invalid return type in cogroup.applyInPandas: {returnType}" |
| }, |
| ) |
| else: |
| raise PySparkTypeError( |
| errorClass="INVALID_RETURN_TYPE_FOR_PANDAS_UDF", |
| messageParameters={ |
| "eval_type": "SQL_COGROUPED_MAP_PANDAS_UDF", |
| "return_type": str(returnType), |
| }, |
| ) |
| elif evalType == PythonEvalType.SQL_COGROUPED_MAP_ARROW_UDF: |
| if isinstance(returnType, StructType): |
| try: |
| to_arrow_type(returnType, timezone="UTC") |
| except TypeError: |
| raise PySparkNotImplementedError( |
| errorClass="NOT_IMPLEMENTED", |
| messageParameters={ |
| "feature": f"Invalid return type in cogroup.applyInArrow: {returnType}" |
| }, |
| ) |
| else: |
| raise PySparkTypeError( |
| errorClass="INVALID_RETURN_TYPE_FOR_ARROW_UDF", |
| messageParameters={ |
| "eval_type": "SQL_COGROUPED_MAP_ARROW_UDF", |
| "return_type": str(returnType), |
| }, |
| ) |
| elif evalType == PythonEvalType.SQL_GROUPED_AGG_PANDAS_UDF: |
| try: |
| # StructType is not yet allowed as a return type, explicitly check here to fail fast |
| if isinstance(returnType, StructType): |
| raise PySparkNotImplementedError( |
| errorClass="NOT_IMPLEMENTED", |
| messageParameters={ |
| "feature": f"Invalid return type with grouped aggregate Pandas UDFs: " |
| f"{returnType}" |
| }, |
| ) |
| to_arrow_type(returnType, timezone="UTC") |
| except TypeError: |
| raise PySparkNotImplementedError( |
| errorClass="NOT_IMPLEMENTED", |
| messageParameters={ |
| "feature": f"Invalid return type with grouped aggregate Pandas UDFs: " |
| f"{returnType}" |
| }, |
| ) |
| elif evalType == PythonEvalType.SQL_GROUPED_AGG_ARROW_UDF: |
| try: |
| # Different from SQL_GROUPED_AGG_PANDAS_UDF, StructType is allowed here |
| to_arrow_type(returnType, timezone="UTC") |
| except TypeError: |
| raise PySparkNotImplementedError( |
| errorClass="NOT_IMPLEMENTED", |
| messageParameters={ |
| "feature": f"Invalid return type with grouped aggregate Arrow UDFs: " |
| f"{returnType}" |
| }, |
| ) |
| |
| @property |
| def returnType(self) -> DataType: |
| # Make sure this is called after SparkContext is initialized. |
| # ``_parse_datatype_string`` accesses to JVM for parsing a DDL formatted string. |
| if self._returnType_placeholder is None: |
| if isinstance(self._returnType, DataType): |
| self._returnType_placeholder = self._returnType |
| else: |
| self._returnType_placeholder = _parse_datatype_string(self._returnType) |
| |
| UserDefinedFunction._check_return_type(self._returnType_placeholder, self.evalType) |
| return self._returnType_placeholder |
| |
| @property |
| def _judf(self) -> "JavaObject": |
| # It is possible that concurrent access, to newly created UDF, |
| # will initialize multiple UserDefinedPythonFunctions. |
| # This is unlikely, doesn't affect correctness, |
| # and should have a minimal performance impact. |
| if self._judf_placeholder is None: |
| self._judf_placeholder = self._create_judf(self.func) |
| return self._judf_placeholder |
| |
| def _create_judf( |
| self, func: Callable[..., Any], include_transpiled: bool = True |
| ) -> "JavaObject": |
| from pyspark.sql import SparkSession |
| from pyspark.sql.classic.column import _to_java_column_opt |
| |
| spark = SparkSession._getActiveSessionOrCreate() |
| sc = spark.sparkContext |
| |
| wrapped_func = _wrap_function(sc, func, self.returnType) |
| jdt = spark._jsparkSession.parseDataType(self.returnType.json()) |
| assert sc._jvm is not None |
| transpiled = self.transpiled if include_transpiled else [] |
| input_categories = self._transpiled_input_categories if include_transpiled else [] |
| # Incremental Python aggregators additionally carry the intermediate buffer schema, which |
| # the JVM needs at planning time to build the two-stage aggregation (see PythonAggregate). |
| # Everyone else passes ``None`` here, which Py4J maps to the JVM ``null`` the ``bufferType`` |
| # parameter already defaults to. |
| jbuf = ( |
| spark._jsparkSession.parseDataType(self.bufferSchema.json()) |
| if self.bufferSchema is not None |
| else None |
| ) |
| judf = getattr(sc._jvm, "org.apache.spark.sql.execution.python.UserDefinedPythonFunction")( |
| self._name, |
| wrapped_func, |
| jdt, |
| self.evalType, |
| self.deterministic, |
| map(_to_java_column_opt, transpiled), |
| input_categories, |
| jbuf, |
| ) |
| return judf |
| |
| def __call__(self, *args: "ColumnOrName", **kwargs: "ColumnOrName") -> Column: |
| from pyspark.sql.classic.column import _to_java_column, _to_seq |
| |
| sc = get_active_spark_context() |
| |
| # Transpilation rewrites the UDF into a Catalyst expression that |
| # references its inputs positionally via ``_udf_param_N`` (see |
| # ``UserDefinedPythonFunction.builder.resolveUDFParams``). If the |
| # caller used kwargs, the JVM-side substitution would otherwise |
| # splice ``NamedArgumentExpression`` wrappers into the rewritten |
| # tree (and into nested function calls like ``isnotnull``, which |
| # rejects named arguments). Resolve kwargs to positional here |
| # using the parameter list captured at transpilation time so the |
| # rewritten expression sees plain column refs in declared order. |
| if kwargs and self.transpiled and self._transpiled_param_names: |
| params = self._transpiled_param_names |
| ordered: list = list(args) |
| remaining_kwargs = dict(kwargs) |
| for pname in params[len(args) :]: |
| if pname in remaining_kwargs: |
| ordered.append(remaining_kwargs.pop(pname)) |
| else: |
| # Caller didn't supply this param positionally or by |
| # name -- bail out of the rewrite and let the regular |
| # JVM-side path raise a user-facing error. |
| break |
| else: |
| if not remaining_kwargs: |
| args = tuple(ordered) |
| kwargs = {} |
| |
| assert sc._jvm is not None |
| jcols = [_to_java_column(arg) for arg in args] + [ |
| sc._jvm.PythonSQLUtils.namedArgumentExpression(key, _to_java_column(value)) |
| for key, value in kwargs.items() |
| ] |
| |
| profiler_enabled = sc._conf.get("spark.python.profile", "false") == "true" |
| memory_profiler_enabled = sc._conf.get("spark.python.profile.memory", "false") == "true" |
| |
| if profiler_enabled or memory_profiler_enabled: |
| # Profiling is not supported for incremental Python aggregators. Their ``self.func`` is |
| # an ``Aggregator`` object, not a plain function: the profiler wrappers below would |
| # replace it with a function the worker cannot drive (it has no ``zero``/``reduce``/ |
| # ``bufferSchema``), and the memory profiler's ``inspect.getsourcelines(f.__code__)`` |
| # fails on the driver because an ``Aggregator`` instance has no ``__code__``. |
| if self.evalType == PythonEvalType.SQL_GROUPED_AGG_ARROW_INCREMENTAL_FINAL_UDF: |
| warnings.warn( |
| "Profiling incremental Python aggregators is not supported.", |
| UserWarning, |
| ) |
| judf = self._judf |
| return Column(judf.apply(_to_seq(sc, jcols))) |
| |
| # Disable profiling Pandas UDFs with iterators as input/output. |
| if self.evalType in [ |
| PythonEvalType.SQL_SCALAR_PANDAS_ITER_UDF, |
| PythonEvalType.SQL_SCALAR_ARROW_ITER_UDF, |
| PythonEvalType.SQL_MAP_PANDAS_ITER_UDF, |
| PythonEvalType.SQL_MAP_ARROW_ITER_UDF, |
| PythonEvalType.SQL_GROUPED_AGG_ARROW_ITER_UDF, |
| PythonEvalType.SQL_GROUPED_AGG_PANDAS_ITER_UDF, |
| ]: |
| warnings.warn( |
| "Profiling UDFs with iterators input/output is not supported.", |
| UserWarning, |
| ) |
| judf = self._judf |
| return Column(judf.apply(_to_seq(sc, jcols))) |
| |
| # Disallow enabling two profilers at the same time. |
| if profiler_enabled and memory_profiler_enabled: |
| # When both profilers are enabled, they interfere with each other, |
| # that makes the result profile misleading. |
| raise PySparkRuntimeError( |
| errorClass="CANNOT_SET_TOGETHER", |
| messageParameters={ |
| "arg_list": "'spark.python.profile' and " |
| "'spark.python.profile.memory' configuration" |
| }, |
| ) |
| elif profiler_enabled: |
| f = self.func |
| profiler = sc.profiler_collector.new_udf_profiler(sc) |
| |
| @functools.wraps(f) |
| def func(*args: Any, **kwargs: Any) -> Any: |
| assert profiler is not None |
| return profiler.profile(f, *args, **kwargs) |
| |
| func.__signature__ = inspect.signature(f) # type: ignore[attr-defined] |
| # Profiling requires the Python function to actually execute, |
| # and the transpiled path never runs it (it also produces a |
| # TranspiledPythonUDF, which has no resultId for the profiler |
| # to key on). Build this call's judf without transpiled options. |
| judf = self._create_judf(func, include_transpiled=False) |
| jUDFExpr = judf.builderWithColumns(_to_seq(sc, jcols)) |
| jPythonUDF = judf.fromUDFExpr(jUDFExpr) |
| id = jUDFExpr.resultId().id() |
| sc.profiler_collector.add_profiler(id, profiler) |
| else: # memory_profiler_enabled |
| f = self.func |
| memory_profiler = sc.profiler_collector.new_memory_profiler(sc) |
| sub_lines, start_line = inspect.getsourcelines(f.__code__) |
| |
| @functools.wraps(f) |
| def func(*args: Any, **kwargs: Any) -> Any: |
| assert memory_profiler is not None |
| return memory_profiler.profile( |
| sub_lines, # type: ignore[arg-type] |
| start_line, |
| f, |
| *args, |
| **kwargs, |
| ) |
| |
| func.__signature__ = inspect.signature(f) # type: ignore[attr-defined] |
| # See the profiler branch above: no transpiled options while |
| # profiling, since only the interpreted path runs the function. |
| judf = self._create_judf(func, include_transpiled=False) |
| jUDFExpr = judf.builderWithColumns(_to_seq(sc, jcols)) |
| jPythonUDF = judf.fromUDFExpr(jUDFExpr) |
| id = jUDFExpr.resultId().id() |
| sc.profiler_collector.add_profiler(id, memory_profiler) |
| else: |
| judf = self._judf |
| jPythonUDF = judf.apply(_to_seq(sc, jcols)) |
| return Column(jPythonUDF) |
| |
| # This function is for improving the online help system in the interactive interpreter. |
| # For example, the built-in help / pydoc.help. It wraps the UDF with the docstring and |
| # argument annotation. (See: SPARK-19161) |
| def _wrapped(self) -> "UserDefinedFunctionLike": |
| """ |
| Wrap this udf with a function and attach docstring from func |
| """ |
| |
| # It is possible for a callable instance without __name__ attribute or/and |
| # __module__ attribute to be wrapped here. For example, functools.partial. In this case, |
| # we should avoid wrapping the attributes from the wrapped function to the wrapper |
| # function. So, we take out these attribute names from the default names to set and |
| # then manually assign it after being wrapped. |
| assignments = tuple( |
| a for a in functools.WRAPPER_ASSIGNMENTS if a != "__name__" and a != "__module__" |
| ) |
| |
| @functools.wraps(self.func, assigned=assignments) |
| def wrapper(*args: "ColumnOrName", **kwargs: "ColumnOrName") -> Column: |
| return self(*args, **kwargs) |
| |
| wrapper.__name__ = self._name |
| wrapper.__module__ = ( |
| self.func.__module__ |
| if hasattr(self.func, "__module__") |
| else self.func.__class__.__module__ |
| ) |
| |
| wrapper.func = self.func # type: ignore[attr-defined] |
| wrapper.returnType = self.returnType # type: ignore[attr-defined] |
| wrapper.evalType = self.evalType # type: ignore[attr-defined] |
| wrapper.deterministic = self.deterministic # type: ignore[attr-defined] |
| wrapper.bufferSchema = self.bufferSchema # type: ignore[attr-defined] |
| wrapper.asNondeterministic = functools.wraps( # type: ignore[attr-defined] |
| self.asNondeterministic |
| )(lambda: self.asNondeterministic()._wrapped()) |
| wrapper._unwrapped = self # type: ignore[attr-defined] |
| return wrapper # type: ignore[return-value] |
| |
| def asNondeterministic(self) -> "UserDefinedFunction": |
| """ |
| Updates UserDefinedFunction to nondeterministic. |
| |
| .. versionadded:: 2.3 |
| """ |
| # Here, we explicitly clean the cache to create a JVM UDF instance |
| # with 'deterministic' updated. See SPARK-23233. |
| self._judf_placeholder = None |
| self.deterministic = False |
| # A transpiled rewrite replaces the (now nondeterministic) Python UDF |
| # with a plain Catalyst expression, which the optimizer is free to |
| # fold, reorder, or duplicate -- discarding the nondeterminism barrier |
| # the caller just asked for. Drop any transpiled options so a |
| # nondeterministic UDF always runs as interpreted Python. |
| self.transpiled = [] |
| self._transpiled_param_names = [] |
| self._transpiled_input_categories = [] |
| return self |
| |
| |
| class UDFRegistration: |
| """ |
| Wrapper for user-defined function registration. This instance can be accessed by |
| :attr:`spark.udf` or :attr:`sqlContext.udf`. |
| |
| .. versionadded:: 1.3.1 |
| """ |
| |
| def __init__(self, sparkSession: "SparkSession"): |
| self.sparkSession = sparkSession |
| |
| def register( |
| self, |
| name: str, |
| f: Union[Callable[..., Any], "UserDefinedFunctionLike"], |
| returnType: Optional["DataTypeOrString"] = None, |
| ) -> "UserDefinedFunctionLike": |
| """Register a Python function (including lambda function) or a user-defined function |
| as a SQL function. |
| |
| .. versionadded:: 1.3.1 |
| |
| .. versionchanged:: 3.4.0 |
| Supports Spark Connect. |
| |
| Parameters |
| ---------- |
| name : str, |
| name of the user-defined function in SQL statements. |
| f : function, :meth:`pyspark.sql.functions.udf` or :meth:`pyspark.sql.functions.pandas_udf` |
| a Python function, or a user-defined function. The user-defined function can |
| be either row-at-a-time or vectorized. See :meth:`pyspark.sql.functions.udf` and |
| :meth:`pyspark.sql.functions.pandas_udf`. |
| returnType : :class:`pyspark.sql.types.DataType` or str, optional |
| the return type of the registered user-defined function. The value can |
| be either a :class:`pyspark.sql.types.DataType` object or a DDL-formatted type string. |
| `returnType` can be optionally specified when `f` is a Python function but not |
| when `f` is a user-defined function. Please see the examples below. |
| |
| Returns |
| ------- |
| function |
| a user-defined function |
| |
| Notes |
| ----- |
| To register a nondeterministic Python function, users need to first build |
| a nondeterministic user-defined function for the Python function and then register it |
| as a SQL function. |
| |
| Examples |
| -------- |
| 1. When `f` is a Python function: |
| |
| `returnType` defaults to string type and can be optionally specified. The produced |
| object must match the specified type. In this case, this API works as if |
| `register(name, f, returnType=StringType())`. |
| |
| >>> strlen = spark.udf.register("stringLengthString", lambda x: len(x)) |
| >>> spark.sql("SELECT stringLengthString('test')").collect() |
| [Row(stringLengthString(test)='4')] |
| |
| >>> spark.sql("SELECT 'foo' AS text").select(strlen("text")).collect() |
| [Row(stringLengthString(text)='3')] |
| |
| >>> from pyspark.sql.types import IntegerType |
| >>> _ = spark.udf.register("stringLengthInt", lambda x: len(x), IntegerType()) |
| >>> spark.sql("SELECT stringLengthInt('test')").collect() |
| [Row(stringLengthInt(test)=4)] |
| |
| >>> from pyspark.sql.types import IntegerType |
| >>> _ = spark.udf.register("stringLengthInt", lambda x: len(x), IntegerType()) |
| >>> spark.sql("SELECT stringLengthInt('test')").collect() |
| [Row(stringLengthInt(test)=4)] |
| |
| 2. When `f` is a user-defined function (from Spark 2.3.0): |
| |
| Spark uses the return type of the given user-defined function as the return type of |
| the registered user-defined function. `returnType` should not be specified. |
| In this case, this API works as if `register(name, f)`. |
| |
| >>> from pyspark.sql.types import IntegerType |
| >>> from pyspark.sql.functions import udf |
| >>> slen = udf(lambda s: len(s), IntegerType()) |
| >>> _ = spark.udf.register("slen", slen) |
| >>> spark.sql("SELECT slen('test')").collect() |
| [Row(slen(test)=4)] |
| |
| >>> import random |
| >>> from pyspark.sql.functions import udf |
| >>> from pyspark.sql.types import IntegerType |
| >>> random_udf = udf(lambda: random.randint(0, 100), IntegerType()).asNondeterministic() |
| >>> new_random_udf = spark.udf.register("random_udf", random_udf) |
| >>> spark.sql("SELECT random_udf()").collect() # doctest: +SKIP |
| [Row(random_udf()=82)] |
| |
| >>> import pandas as pd |
| >>> from pyspark.sql.functions import pandas_udf |
| >>> @pandas_udf("integer") |
| ... def add_one(s: pd.Series) -> pd.Series: |
| ... return s + 1 |
| ... |
| >>> _ = spark.udf.register("add_one", add_one) |
| >>> spark.sql("SELECT add_one(id) FROM range(3)").collect() |
| [Row(add_one(id)=1), Row(add_one(id)=2), Row(add_one(id)=3)] |
| |
| >>> @pandas_udf("integer") |
| ... def sum_udf(v: pd.Series) -> int: |
| ... return v.sum() |
| ... |
| >>> _ = spark.udf.register("sum_udf", sum_udf) |
| >>> q = "SELECT sum_udf(v1) FROM VALUES (3, 0), (2, 0), (1, 1) tbl(v1, v2) GROUP BY v2" |
| >>> spark.sql(q).sort("sum_udf(v1)").collect() |
| [Row(sum_udf(v1)=1), Row(sum_udf(v1)=5)] |
| """ |
| |
| # This is to check whether the input function is from a user-defined function or |
| # Python function. |
| if hasattr(f, "asNondeterministic"): |
| if returnType is not None: |
| raise PySparkTypeError( |
| errorClass="CANNOT_SPECIFY_RETURN_TYPE_FOR_UDF", |
| messageParameters={"arg_name": "f", "return_type": str(returnType)}, |
| ) |
| f = cast("UserDefinedFunctionLike", f) |
| if f.evalType not in [ |
| PythonEvalType.SQL_BATCHED_UDF, |
| PythonEvalType.SQL_ARROW_BATCHED_UDF, |
| PythonEvalType.SQL_SCALAR_PANDAS_UDF, |
| PythonEvalType.SQL_SCALAR_ARROW_UDF, |
| PythonEvalType.SQL_SCALAR_PANDAS_ITER_UDF, |
| PythonEvalType.SQL_SCALAR_ARROW_ITER_UDF, |
| PythonEvalType.SQL_GROUPED_AGG_PANDAS_UDF, |
| PythonEvalType.SQL_GROUPED_AGG_ARROW_UDF, |
| PythonEvalType.SQL_GROUPED_AGG_PANDAS_ITER_UDF, |
| PythonEvalType.SQL_GROUPED_AGG_ARROW_ITER_UDF, |
| PythonEvalType.SQL_GROUPED_AGG_ARROW_INCREMENTAL_FINAL_UDF, |
| ]: |
| raise PySparkTypeError( |
| errorClass="INVALID_UDF_EVAL_TYPE", |
| messageParameters={ |
| "eval_type": "SQL_BATCHED_UDF, SQL_ARROW_BATCHED_UDF, " |
| "SQL_SCALAR_PANDAS_UDF, SQL_SCALAR_ARROW_UDF, " |
| "SQL_SCALAR_PANDAS_ITER_UDF, SQL_SCALAR_ARROW_ITER_UDF, " |
| "SQL_GROUPED_AGG_PANDAS_UDF, SQL_GROUPED_AGG_ARROW_UDF, " |
| "SQL_GROUPED_AGG_PANDAS_ITER_UDF, SQL_GROUPED_AGG_ARROW_ITER_UDF " |
| "or SQL_GROUPED_AGG_ARROW_INCREMENTAL_FINAL_UDF" |
| }, |
| ) |
| source_udf = _create_udf( |
| f.func, |
| returnType=f.returnType, |
| name=name, |
| evalType=f.evalType, |
| deterministic=f.deterministic, |
| # Preserve the incremental aggregator's buffer schema (None for other UDFs). |
| bufferSchema=getattr(f, "bufferSchema", None), |
| ) |
| register_udf = source_udf._unwrapped # type: ignore[attr-defined] |
| return_udf = register_udf |
| else: |
| if returnType is None: |
| returnType = StringType() |
| return_udf = _create_udf( |
| f, returnType=returnType, evalType=PythonEvalType.SQL_BATCHED_UDF, name=name |
| ) |
| register_udf = return_udf._unwrapped # type: ignore[attr-defined] |
| self.sparkSession._jsparkSession.udf().registerPython(name, register_udf._judf) |
| return return_udf |
| |
| def registerJavaFunction( |
| self, |
| name: str, |
| javaClassName: str, |
| returnType: Optional["DataTypeOrString"] = None, |
| ) -> None: |
| """Register a Java user-defined function as a SQL function. |
| |
| In addition to a name and the function itself, the return type can be optionally specified. |
| When the return type is not specified we would infer it via reflection. |
| |
| .. versionadded:: 2.3.0 |
| |
| .. versionchanged:: 3.4.0 |
| Supports Spark Connect. |
| |
| Parameters |
| ---------- |
| name : str |
| name of the user-defined function |
| javaClassName : str |
| fully qualified name of java class |
| returnType : :class:`pyspark.sql.types.DataType` or str, optional |
| the return type of the registered Java function. The value can be either |
| a :class:`pyspark.sql.types.DataType` object or a DDL-formatted type string. |
| |
| Examples |
| -------- |
| >>> from pyspark.sql.types import IntegerType |
| >>> spark.udf.registerJavaFunction( |
| ... "javaStringLength", "test.org.apache.spark.sql.JavaStringLength", IntegerType()) |
| ... # doctest: +SKIP |
| >>> spark.sql("SELECT javaStringLength('test')").collect() # doctest: +SKIP |
| [Row(javaStringLength(test)=4)] |
| |
| >>> spark.udf.registerJavaFunction( |
| ... "javaStringLength2", "test.org.apache.spark.sql.JavaStringLength") |
| ... # doctest: +SKIP |
| >>> spark.sql("SELECT javaStringLength2('test')").collect() # doctest: +SKIP |
| [Row(javaStringLength2(test)=4)] |
| |
| >>> spark.udf.registerJavaFunction( |
| ... "javaStringLength3", "test.org.apache.spark.sql.JavaStringLength", "integer") |
| ... # doctest: +SKIP |
| >>> spark.sql("SELECT javaStringLength3('test')").collect() # doctest: +SKIP |
| [Row(javaStringLength3(test)=4)] |
| """ |
| |
| jdt = None |
| if returnType is not None: |
| if not isinstance(returnType, DataType): |
| returnType = _parse_datatype_string(returnType) |
| jdt = self.sparkSession._jsparkSession.parseDataType(returnType.json()) |
| self.sparkSession._jsparkSession.udf().registerJava(name, javaClassName, jdt) |
| |
| def registerJavaUDAF(self, name: str, javaClassName: str) -> None: |
| """Register a Java user-defined aggregate function as a SQL function. |
| |
| .. versionadded:: 2.3.0 |
| |
| .. versionchanged:: 3.4.0 |
| Supports Spark Connect. |
| |
| name : str |
| name of the user-defined aggregate function |
| javaClassName : str |
| fully qualified name of java class |
| |
| Examples |
| -------- |
| >>> spark.udf.registerJavaUDAF("javaUDAF", "test.org.apache.spark.sql.MyDoubleAvg") |
| ... # doctest: +SKIP |
| >>> df = spark.createDataFrame([(1, "a"),(2, "b"), (3, "a")],["id", "name"]) |
| >>> df.createOrReplaceTempView("df") |
| >>> q = "SELECT name, javaUDAF(id) as avg from df group by name order by name desc" |
| >>> spark.sql(q).collect() # doctest: +SKIP |
| [Row(name='b', avg=102.0), Row(name='a', avg=102.0)] |
| """ |
| |
| self.sparkSession._jsparkSession.udf().registerJavaUDAF(name, javaClassName) |
| |
| |
| def _test() -> None: |
| import doctest |
| |
| import pyspark.sql.udf |
| from pyspark.sql import SparkSession |
| from pyspark.testing.utils import have_pandas, have_pyarrow |
| |
| globs = pyspark.sql.udf.__dict__.copy() |
| |
| if not have_pandas or not have_pyarrow: |
| del pyspark.sql.udf.UDFRegistration.register.__doc__ |
| |
| spark = SparkSession.builder.master("local[4]").appName("sql.udf tests").getOrCreate() |
| globs["spark"] = spark |
| failure_count, test_count = doctest.testmod( |
| pyspark.sql.udf, globs=globs, optionflags=doctest.ELLIPSIS | doctest.NORMALIZE_WHITESPACE |
| ) |
| spark.stop() |
| if failure_count: |
| sys.exit(-1) |
| |
| |
| if __name__ == "__main__": |
| _test() |