diff --git a/narwhals/stable/v1/__init__.py b/narwhals/stable/v1/__init__.py index 33531480c4..0ad44bc27c 100644 --- a/narwhals/stable/v1/__init__.py +++ b/narwhals/stable/v1/__init__.py @@ -51,7 +51,7 @@ from narwhals.translate import _from_native_impl from narwhals.translate import get_native_namespace as nw_get_native_namespace from narwhals.translate import to_native -from narwhals.translate import to_py_scalar as nw_to_py_scalar +from narwhals.translate import to_py_scalar from narwhals.typing import IntoDataFrameT from narwhals.typing import IntoFrameT from narwhals.typing import IntoSeriesT @@ -956,28 +956,6 @@ def wrapper(*args: Any, **kwargs: Any) -> Any: return decorator(func) -def to_py_scalar(scalar: Any) -> Any: - """If a scalar is not Python native, converts it to Python native. - - Raises: - ValueError: If the object is not convertible to a scalar. - - Examples: - >>> import narwhals.stable.v1 as nw - >>> import pandas as pd - >>> df = nw.from_native(pd.DataFrame({"a": [1, 2, 3]})) - >>> nw.to_py_scalar(df["a"].item(0)) - 1 - >>> import pyarrow as pa - >>> df = nw.from_native(pa.table({"a": [1, 2, 3]})) - >>> nw.to_py_scalar(df["a"].item(0)) - 1 - >>> nw.to_py_scalar(1) - 1 - """ - return _stableify(nw_to_py_scalar(scalar)) - - def all() -> Expr: """ Instantiate an expression representing all columns. diff --git a/narwhals/translate.py b/narwhals/translate.py index a1b0e23239..d439890c6a 100644 --- a/narwhals/translate.py +++ b/narwhals/translate.py @@ -1,6 +1,5 @@ from __future__ import annotations -import numbers from datetime import datetime from datetime import timedelta from functools import wraps @@ -46,6 +45,15 @@ T = TypeVar("T") +NON_TEMPORAL_SCALAR_TYPES = ( + bool, + bytes, + str, + int, + float, + complex, +) + @overload def to_native( @@ -843,16 +851,10 @@ def to_py_scalar(scalar_like: Any) -> Any: >>> nw.to_py_scalar(1) 1 """ - - pa = get_pyarrow() - if pa and isinstance(scalar_like, pa.Scalar): - return scalar_like.as_py() - - cupy = get_cupy() - if ( # pragma: no cover - cupy and isinstance(scalar_like, cupy.ndarray) and scalar_like.size == 1 - ): - return scalar_like.item() + if scalar_like is None: + return None + if isinstance(scalar_like, NON_TEMPORAL_SCALAR_TYPES): + return scalar_like np = get_numpy() if np and np.isscalar(scalar_like) and hasattr(scalar_like, "item"): @@ -863,21 +865,30 @@ def to_py_scalar(scalar_like: Any) -> Any: return scalar_like.to_pydatetime() if pd and isinstance(scalar_like, pd.Timedelta): return scalar_like.to_pytimedelta() - - all_scalar_types = ( - int, - float, - complex, - bool, - bytes, - str, - datetime, - timedelta, - numbers.Number, - ) - if isinstance(scalar_like, all_scalar_types): + if pd and pd.api.types.is_scalar(scalar_like): + try: + is_na = pd.isna(scalar_like) + except Exception: # pragma: no cover # noqa: BLE001, S110 + pass + else: + if is_na: + return None + + # pd.Timestamp and pd.Timedelta subclass datetime and timedelta, + # so we need to check this separately + if isinstance(scalar_like, (datetime, timedelta)): return scalar_like + pa = get_pyarrow() + if pa and isinstance(scalar_like, pa.Scalar): + return scalar_like.as_py() + + cupy = get_cupy() + if ( # pragma: no cover + cupy and isinstance(scalar_like, cupy.ndarray) and scalar_like.size == 1 + ): + return scalar_like.item() + msg = ( f"Expected object convertible to a scalar, found {type(scalar_like)}. " "Please report a bug to https://github.com/narwhals-dev/narwhals/issues" diff --git a/tests/stable_api_test.py b/tests/stable_api_test.py index c1b2f14043..c0a97f1547 100644 --- a/tests/stable_api_test.py +++ b/tests/stable_api_test.py @@ -87,6 +87,10 @@ def test_stable_api_docstrings() -> None: nw_doc = getattr(nw, item).__doc__ if item == "from_native": v1_doc = v1_doc.replace("native_dataframe", "native_object") + if item in {"to_py_scalar"}: + # We don't overwrite the docstring for these + assert v1_doc == nw_doc + continue assert ( v1_doc.replace("import narwhals.stable.v1 as nw", "import narwhals as nw") == nw_doc diff --git a/tests/translate/to_py_scalar_test.py b/tests/translate/to_py_scalar_test.py index 53c3df7385..3519b5e878 100644 --- a/tests/translate/to_py_scalar_test.py +++ b/tests/translate/to_py_scalar_test.py @@ -20,7 +20,9 @@ ("input_value", "expected"), [ (1, 1), + (np.int64(1), 1), (1.0, 1.0), + (None, None), ("a", "a"), (True, True), (b"a", b"a"),