Skip to content

Exploring support for eager and lazy dataframe semantics #335

Description

@aaraney
from __future__ import annotations

import typing

import narwhals as nw
import numpy as np

if typing.TYPE_CHECKING:
    import numpy.typing as npt
    import narwhals.typing as nwt


def mean_error(
    y_true: npt.ArrayLike,
    y_pred: npt.ArrayLike,
    power: float = 1.0,
    root: bool = False,
) -> float:
    # Compute mean error
    _mean_error = np.sum(np.abs(np.subtract(y_true, y_pred)) ** power) / len(y_true)

    # Return mean error, optionally return root mean error
    if root:
        return np.sqrt(_mean_error)
    return _mean_error


def mean_error_narwals_eager(
    y_true: nwt.IntoSeriesT,
    y_pred: nwt.IntoSeriesT,
    power: float = 1.0,
    root: bool = False,
) -> float:
    true, pred = (
        nw.from_native(y_true, series_only=True),
        nw.from_native(y_pred, series_only=True),
    )
    _mean_error = ((true - pred).abs() ** power).sum() / true.len()

    # Return mean error, optionally return root mean error
    if root:
        return np.sqrt(_mean_error)
    return _mean_error


def mean_error_narwals_lazy(
    y_true: nw.Expr,
    y_pred: nw.Expr,
    power: float = 1.0,
    root: bool = False,
) -> nw.Expr:
    _mean_error = ((y_true - y_pred).abs() ** power).sum() / y_true.len()

    # Return mean error, optionally return root mean error
    if root:
        return _mean_error.sqrt()
    return _mean_error


if typing.TYPE_CHECKING:
    import pandas as pd
    import polars as pl

    IntoSeriesT = typing.TypeVar("IntoSeriesT", pd.Series, pl.Series, nw.Series)


@typing.overload
def mean_error_narwals_overload(
    y_true: IntoSeriesT,
    y_pred: IntoSeriesT,
    power: float = 1.0,
    root: bool = False,
) -> float: ...


@typing.overload
def mean_error_narwals_overload(
    y_true: nwt.Expr,
    y_pred: nwt.Expr,
    power: float = 1.0,
    root: bool = False,
) -> nw.Expr: ...


def mean_error_narwals_overload(
    y_true: IntoSeriesT | nwt.Expr,
    y_pred: IntoSeriesT | nwt.Expr,
    power: float = 1.0,
    root: bool = False,
) -> float | nw.Expr:
    if isinstance(y_true, nw.Expr):
        true, pred = y_true, y_pred
    else:
        true, pred = (
            nw.from_native(y_true, series_only=True),
            nw.from_native(y_pred, series_only=True),
        )

    _mean_error = ((true - pred).abs() ** power).sum() / true.len()

    # Return mean error, optionally return root mean error
    if root:
        if isinstance(_mean_error, nw.Expr):
            return _mean_error.sqrt()
        return np.sqrt(_mean_error)
    return _mean_error


def p(x: object) -> None:
    print(x)


import pandas as pd
import polars as pl

true = np.array([1, 2, 3], dtype="float64")
pred = np.array([1, 2, 3], dtype="float64")

pd_df = pd.DataFrame({"true": true, "pred": pred})
pl_df = pl.DataFrame({"true": true, "pred": pred})
pl_lf = pl.LazyFrame({"true": true, "pred": pred})
nw_df = nw.from_native(pd_df)

p(mean_error(pd_df["true"], pd_df["pred"]))

p(
    mean_error_narwals_eager(pd_df["true"], pd_df["pred"]),
)
p(
    mean_error_narwals_eager(pl_df["true"], pl_df["pred"]),
)
p(
    mean_error_narwals_eager(nw_df["true"], nw_df["pred"]),
)
try:
    p(
        mean_error_narwals_eager(
            pd_df["true"], pl_df["pred"]
        )  # we want this to type error
    )
except TypeError as e:
    p(e)

try:
    p(
        mean_error_narwals_eager(true, pred),  # should typing error
    )
except TypeError as e:
    p(e)

p(
    mean_error_narwals_lazy(nw.col("true"), nw.col("pred")),
)
# mean_error_narwals_lazy("true", "pred")  # likely want this too

p(
    mean_error_narwals_overload(pd_df["true"], pd_df["pred"]),
)
p(
    mean_error_narwals_overload(pl_df["true"], pl_df["pred"]),
)
p(
    mean_error_narwals_overload(nw_df["true"], nw_df["pred"]),
)

p(
    nw.from_native(pl_lf)
    .select(
        mean_error_narwals_overload(nw.col("true"), nw.col("pred")),
    )
    .to_native()
    .collect()
)

Activity

Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Metadata

Metadata

Assignees

No one assigned

    Labels

    No labels
    No labels

    Type

    No type

    Projects

    No projects

      Milestone

      No milestone

      Relationships

      None yet

      Development

      No branches or pull requests

      Issue actions