Skip to content
Closed
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
32 changes: 26 additions & 6 deletions python/pyspark/pandas/numpy_compat.py
Original file line number Diff line number Diff line change
Expand Up @@ -18,6 +18,7 @@

import numpy as np

from pyspark.loose_version import LooseVersion
from pyspark.sql import Column, functions as F
from pyspark.sql.pandas.functions import pandas_udf
from pyspark.sql.types import DoubleType, BooleanType
Expand Down Expand Up @@ -203,6 +204,29 @@ def _floor_divide_func(c1: Column, c2: Column) -> Column:
)


# NumPy 2.3.0 changed how fmax/fmin break a signed-zero tie: for equal operands
# (for example +0.0 and -0.0) it returns the first operand, while older versions
# returned the second. Track the installed NumPy so the result keeps the matching
# sign of zero.
_tie_returns_first_operand = LooseVersion(np.__version__) >= LooseVersion("2.3.0")


def _fmax_func(c1: Column, c2: Column) -> Column:
tie = c1 if _tie_returns_first_operand else c2
return (
F.when(F.isnan(c1.cast("double")), c2)
.when(F.isnan(c2.cast("double")), c1)
.when(c1 == c2, tie)
.otherwise(F.greatest(c1, c2))
.cast("double")
)


def _fmin_func(c1: Column, c2: Column) -> Column:
tie = c1 if _tie_returns_first_operand else c2
return F.when(c1 == c2, tie).otherwise(F.least(c1, c2)).cast("double")


binary_np_spark_mappings = {
"arctan2": F.atan2,
"bitwise_and": lambda c1, c2: c1.bitwiseAND(c2),
Expand All @@ -215,12 +239,8 @@ def _floor_divide_func(c1: Column, c2: Column) -> Column:
# np.floor_divide dispatches to the pandas-on-Spark floordiv dunder operation
# before this registry is consulted, so this mapping is not used for that case.
"floor_divide": _floor_divide_func,
"fmax": lambda c1, c2: F.when(F.isnan(c1.cast("double")), c2)
.when(F.isnan(c2.cast("double")), c1)
.when(c1 == c2, c1)
.otherwise(F.greatest(c1, c2))
.cast("double"),
"fmin": lambda c1, c2: F.when(c1 == c2, c1).otherwise(F.least(c1, c2)).cast("double"),
"fmax": _fmax_func,
"fmin": _fmin_func,
"fmod": _fmod_func,
"gcd": pandas_udf(lambda s1, s2: np.gcd(s1, s2), DoubleType()), # type: ignore[call-overload]
"heaviside": lambda c1, c2: F.when(
Expand Down