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