Source code for pyspark.sql.udf

#
# 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)


[docs] 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]
[docs] 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
[docs] 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
[docs] 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
[docs] 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)
[docs] 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()