diff --git a/fairlearn/preprocessing/__init__.py b/fairlearn/preprocessing/__init__.py index 8d86978..b06dd8d 100644 --- a/fairlearn/preprocessing/__init__.py +++ b/fairlearn/preprocessing/__init__.py @@ -5,5 +5,6 @@ from ._correlation_remover import CorrelationRemover from ._prototype_representation_learner import PrototypeRepresentationLearner +from ._optimized_preprocessing import OptimizedPreprocessing -__all__ = ["CorrelationRemover", "PrototypeRepresentationLearner"] +__all__ = ["CorrelationRemover", "PrototypeRepresentationLearner", "OptimizedPreprocessing"] diff --git a/fairlearn/preprocessing/_optimized_preprocessing.py b/fairlearn/preprocessing/_optimized_preprocessing.py new file mode 100644 index 0000000..a61b71a --- /dev/null +++ b/fairlearn/preprocessing/_optimized_preprocessing.py @@ -0,0 +1,88 @@ +# Copyright (c) Fairlearn contributors. +# Licensed under the MIT License. + +from __future__ import annotations + +import numpy as np +from sklearn.base import BaseEstimator, TransformerMixin +from sklearn.utils.validation import check_is_fitted + + +class OptimizedPreprocessing(TransformerMixin, BaseEstimator): + r""" + Simplified Optimized Pre-Processing (Calmon et al., 2017). + + This transformer learns a mapping of feature values to reduce disparity + (via group-wise mean matching) while minimizing per-sample distortion. + It is a minimal, model-agnostic variant suitable as a starting point. + + Parameters + ---------- + fairness_weight : float, default=1.0 + Strength of the group parity adjustment. + + random_state : int, RandomState, default=None + Controls randomized steps for reproducibility (if any). + + Attributes + ---------- + group_shift_ : dict + Mapping from group label to shift vector applied to group samples. + + feature_mean_ : ndarray of shape (n_features,) + Global mean of features. + + Notes + ----- + This is a simplified baseline inspired by Calmon et al. The full method + solves a constrained optimization problem with distortion costs. Here we + approximate by applying per-group mean shifts toward the global mean. + """ + + def __init__(self, *, fairness_weight: float = 1.0, random_state=None): + self.fairness_weight = fairness_weight + self.random_state = random_state + + def fit(self, X, y=None, *, sensitive_features=None): + X = np.asarray(X, dtype=float) + if sensitive_features is None: + raise ValueError("sensitive_features must be provided") + s = np.asarray(sensitive_features) + if s.shape[0] != X.shape[0]: + raise ValueError("sensitive_features must match X in the first dimension") + + self.feature_mean_ = X.mean(axis=0) + self.group_shift_ = {} + for g in np.unique(s): + mask = s == g + if np.any(mask): + g_mean = X[mask].mean(axis=0) + shift = (self.feature_mean_ - g_mean) * float(self.fairness_weight) + self.group_shift_[g] = shift + else: + self.group_shift_[g] = np.zeros_like(self.feature_mean_) + self.n_features_in_ = X.shape[1] + return self + + def transform(self, X, *, sensitive_features=None): + check_is_fitted(self, ["group_shift_", "feature_mean_", "n_features_in_"]) + X = np.asarray(X, dtype=float) + if X.shape[1] != self.n_features_in_: + raise ValueError( + "X has %d features, but %s is expecting %d features as input" + % (X.shape[1], self.__class__.__name__, self.n_features_in_) + ) + if sensitive_features is None: + raise ValueError("sensitive_features must be provided for transform") + s = np.asarray(sensitive_features) + if s.shape[0] != X.shape[0]: + raise ValueError("sensitive_features must match X in the first dimension") + + X_adj = X.copy() + for g, shift in self.group_shift_.items(): + mask = s == g + if np.any(mask): + X_adj[mask] = X_adj[mask] + shift + return X_adj + + diff --git a/test/unit/preprocessing/test_optimized_preprocessing.py b/test/unit/preprocessing/test_optimized_preprocessing.py new file mode 100644 index 0000000..207a885 --- /dev/null +++ b/test/unit/preprocessing/test_optimized_preprocessing.py @@ -0,0 +1,29 @@ +# Copyright (c) Fairlearn contributors. +# Licensed under the MIT License. + +import numpy as np + +from fairlearn.preprocessing import OptimizedPreprocessing + + +def test_optimized_preprocessing_shapes(): + X = np.array([[0.0, 1.0], [1.0, 2.0], [2.0, 3.0], [0.0, 0.0]]) + s = np.array([0, 0, 1, 1]) + opt = OptimizedPreprocessing(fairness_weight=1.0) + opt.fit(X, sensitive_features=s) + X_adj = opt.transform(X, sensitive_features=s) + assert X_adj.shape == X.shape + + +def test_optimized_preprocessing_group_shift(): + X = np.array([[0.0], [0.0], [2.0], [2.0]]) + s = np.array([0, 0, 1, 1]) + opt = OptimizedPreprocessing(fairness_weight=1.0) + opt.fit(X, sensitive_features=s) + X_adj = opt.transform(X, sensitive_features=s) + # Group means become closer after transform + mean0 = X_adj[s == 0].mean() + mean1 = X_adj[s == 1].mean() + assert np.isclose(mean0, mean1) + +