Skip to content
Closed
Show file tree
Hide file tree
Changes from all 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
3 changes: 2 additions & 1 deletion fairlearn/preprocessing/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -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"]
88 changes: 88 additions & 0 deletions fairlearn/preprocessing/_optimized_preprocessing.py
Original file line number Diff line number Diff line change
@@ -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


29 changes: 29 additions & 0 deletions test/unit/preprocessing/test_optimized_preprocessing.py
Original file line number Diff line number Diff line change
@@ -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)