Conversation
Compute the target mean per bin and per category directly instead of through a Pipeline of discretiser and MeanEncoders, reusing the fitted discretiser's bin edges and labels. Fit and _predict accept pandas and polars dataframes, and pandas integer column names now work. Co-Authored-By: Claude Opus 5 <noreply@anthropic.com>
solegalli
force-pushed
the
narwhals-prediction-base
branch
from
September 19, 2026 09:33
9055577 to
6d284b4
Compare
This was referenced Sep 19, 2026
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
Sign up for free
to join this conversation on GitHub.
Already have an account?
Sign in to comment
Add this suggestion to a batch that can be applied as a single commit.This suggestion is invalid because no changes were made to the code.Suggestions cannot be applied while the pull request is closed.Suggestions cannot be applied while viewing a subset of changes.Only one suggestion per line can be applied in a batch.Add this suggestion to a batch that can be applied as a single commit.Applying suggestions on deleted lines is not supported.You must change the existing code in this line in order to create a valid suggestion.Outdated suggestions cannot be applied.This suggestion has been applied or marked resolved.Suggestions cannot be applied from pending reviews.Suggestions cannot be applied on multi-line comments.Suggestions cannot be applied while the pull request is queued to merge.Suggestion cannot be applied right now. Please check back later.
Summary
Migrates
BaseTargetMeanEstimator(feature_engine/_prediction/base_predictor.py) to narwhals.fitand_predictnow accept pandas and polars dataframes.TargetMeanClassifierandTargetMeanRegressorare not changed here; they get their own PRs stacked on this one.Before this PR, the base built a
Pipelineof a discretiser (return_boundaries=True) and one or twoMeanEncoders. Every numerical value became an interval string, and the encoder grouped the target by those strings. On polars,fitran but_predictfailed atDataFrame.mean(axis=1). On pandas, integer column names failed infit.The new implementation computes the means directly:
EqualWidthDiscretiser/EqualFrequencyDiscretiseris fitted to getbinner_dict_(reused as is). Its_digitizeturns values into integer bin codes. The target mean per code is computed, and the keys ofencoder_dict_are the same interval labels the discretiser writes (_format_bin_labels). A per-bin array of means (_bin_means, NaN for bins that had no training rows) makes_predicta numpy lookup:means[codes]._predict: sums the encoded columns in numpy and divides by the number of variables. It raises the same "NaN values were introduced" error for unseen categories and for values that fall in bins that were empty in the training set.nw_X = check_X_y(X, y); the nativeXgoes to the variable, NaN and inf helpers; the target is paired withadd_target_to_X._pipeline,_make_*_pipelineand_transformare removed._transformwas only used by_predict.strategychecks the type before the membership test, andbinsmust be a positive integer (see "Pre-existing issues").Benchmarks
Median of 7 runs, versions run in alternation, times in ms. Data:
numstandard-normal columns,catstring columns with 20 categories, continuous target,bins=5,equal_width. "old" is the pipeline on the base branch. Its polars_predictfails, so that column shows "-".On categorical variables, most of the remaining fit time is spent in
find_categorical_and_numerical_variables, which is shared and not changed here (see "Pre-existing issues").How each step was chosen, timed per variable, isolated, median of repeats:
y.groupby(codes).mean()2.53 ms; narwhals group_by on pandas 3.05;np.bincount2.31 (10k: 0.10 / 0.75 / 0.05)bincountis about 8% faster but differs from the old means in the last bits (see "Needs decision")bincount124.7 (2M: 320 vs 336; 5M: 810 vs 904)groupby(observed=True, dropna=False)vs narwhals group_by: 10k 0.22 vs 0.48; 100k 1.9 vs 2.0; 500k 9.5 vs 9.2 (20 categories)dropna=Falseis safe after the NaN check and is about 40% faster than the default_digitize5.8 ms;np.searchsorted5.5; polarssearch_sorted5.1 (2M: 23.5 / 22.1 / 20.6)_digitizeon both backends, to reuse the discretiser. A polars-only branch would save about 10% of this stepfactorize(use_na_sentinel=False)+ lookup 8.7 ms;.map(dict)12.5; narwhalsreplace_strict14.1replace_strict2.6 ms vs polars-native 2.3DataFrame.mean(axis=1))Behaviour
pandas: I compared outputs of
BaseTargetMeanEstimator,TargetMeanClassifierandTargetMeanRegressor(fitattributes,_predict,predict,predict_probaand error messages) with the base branch on 44 cases. The cases covered both strategies, 3/5/10 bins, binary and continuous targets, list/array/bool/int targets, reordered columns, constant columns, bins left empty in training, unseen categories, out-of-range values, NaN and inf at fit and at predict, a wrong number of columns, datetime columns, category dtype with unused categories, object columns holding integers, more bins than rows, a custom index and nullable pandas dtypes.binner_dict_,encoder_dict_values, predictions and error messages are bit-identical. The differences:fitraisedInvalidIntoExprErrorbefore.encoder_dict_changed (dict equality is unchanged). Numerical variables now list their intervals in bin order. Before, the order was alphabetical by label string, for example'(-inf, 9.8]', '(19.6, 29.4]', ..., '(9.8, 19.6]'. Categorical variables are sorted on pandas; before, they were sometimes sorted by count, depending on which pipeline branch ran. On polars the order followsgroup_by.polars: gives the same values as pandas on all cases (rtol 1e-12). Group-by means on polars can differ from pandas in the last bit, so the tests use
pytest.approx.Tests
tests/test_prediction/test_base_predictor.py, following the conventions. It starts with init errors (bins, strategy) andtest_init_param_assignment. The fit and predict tests run on both backends throughmake_df: attributes per strategy, predictions, list and array targets, numerical-only and categorical-only data, variable selection with a datetime column, reordered columns, out-of-range values, a constant variable, unseen categories, bins that were empty in training, NaN/None and inf at fit and predict, the column count check andNotFittedError. Three tests are pandas-only: integer column names, a custom index and category dtype.test_check_estimator_prediction.py:test_attributes_upon_fittingchecked_pipeline.named_steps. It now checks_discretiser.tests/test_prediction: 58 passed / 0 failed before; 122 passed / 0 failed after.tests/test_selection(imports the estimators): 177 failed / 251 passed before and after, with the same list of failures.test_target_mean_selection.pyalready fails on the base branch, inSelectByTargetEncodingitself ('list' object has no attribute 'to_list').flake8 feature_engine testsis clean.mypy feature_enginereports the same 2 errors as the base branch.Docs: the prediction module has no user guide page, and the base class is private, so only the docstring changed. The user-facing examples belong in the classifier and regressor PRs.
Needs decision
np.bincount: it is about 8% faster on numerical-only fits. It changes the means in the last bits (relative difference up to about 1e-13) because pandas uses a compensated sum. I kept pandas groupby so the output stays identical.Pre-existing issues
bins=0gave a misleading NaN error at fit, and negative bins raisedIndexError. They now raise at init withbins must be a positive integer. Got {bins} instead.Before, the message said "bins must be an integer".find_categorical_and_numerical_variablesscans the full values of every string column to check whether they parse as datetimes or numbers. It takes about 60% offiton categorical data (for example, 315 ms of 583 ms at 500k rows with 10 categorical variables). This is shared code outside this PR.Objectcolumns, such as the output of a discretiser withreturn_object=True, can be fitted but_predictfails inreplace_strict.MeanEncoder.transformfails the same way. On the base branch,fitalready rejected them.