Skip to content
Merged
Show file tree
Hide file tree
Changes from 2 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
1 change: 1 addition & 0 deletions dask_ml/metrics/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -6,6 +6,7 @@
)
from .regression import ( # noqa
mean_absolute_error,
mean_absolute_percentage_error,
mean_squared_error,
mean_squared_log_error,
r2_score,
Expand Down
27 changes: 27 additions & 0 deletions dask_ml/metrics/regression.py
Original file line number Diff line number Diff line change
Expand Up @@ -81,6 +81,33 @@ def mean_absolute_error(
return result


@derived_from(sklearn.metrics)
def mean_absolute_percentage_error(
y_true: ArrayLike,
y_pred: ArrayLike,
sample_weight: Optional[ArrayLike] = None,
multioutput: Optional[str] = "uniform_average",
compute: bool = True,
) -> ArrayLike:
_check_sample_weight(sample_weight)
epsilon = np.finfo(np.float64).eps
mape = abs(y_pred - y_true) / da.maximum(y_true, epsilon)
output_errors = mape.mean(axis=0)

if isinstance(multioutput, str) or multioutput is None:
if multioutput == "raw_values":
if compute:
return output_errors.compute()
else:
return output_errors
else:
raise ValueError("Weighted 'multioutput' not supported.")
result = output_errors.mean()
if compute:
result = result.compute()
return result


@derived_from(sklearn.metrics)
def r2_score(
y_true: ArrayLike,
Expand Down
10 changes: 9 additions & 1 deletion tests/metrics/test_regression.py
Original file line number Diff line number Diff line change
Expand Up @@ -7,12 +7,20 @@
import dask_ml.metrics


@pytest.fixture(params=["mean_squared_error", "mean_absolute_error", "r2_score"])
@pytest.fixture(
params=[
"mean_squared_error",
"mean_absolute_error",
"mean_absolute_percentage_error",
"r2_score",
]
)
def metric_pairs(request):
"""Pairs of (dask-ml, sklearn) regression metrics.

* mean_squared_error
* mean_absolute_error
* mean_absolute_percentage_error
* r2_score
"""
return (
Expand Down