feat: opt-in pipeline fit optimisations, native bucketize, and sampled fitting support with caching - #69
Conversation
Plans are too slow to materialise - this is our attempt to speed it up
Wrap the moments aggregation in StandardScale, SingleFeatureArrayStandardScale and ConditionalStandardScale estimators in a guarded persist/unpersist so the array-size probe and the aggregation reuse a materialised result instead of re-scanning the upstream lineage twice. Repair the incomplete persist edit in ConditionalStandardScale._fit. Add checkpointInterval / pruneInputColumns coverage to the pipeline tests and a checkpoint directory to the spark_session fixture. Surface estimator fit errors as RuntimeError chained from the original exception. Co-Authored-By: Claude Opus 4.7 <noreply@anthropic.com>
Merge latest version
georyetti
left a comment
There was a problem hiding this comment.
Mostly minor changes to pipeline logic. Also bucketize changes I assume should not be here. Lastly you need to run a uv lock.
| """ | ||
| super().__init__(stages=stages) | ||
| kwargs = self._input_kwargs | ||
| super().__init__() |
There was a problem hiding this comment.
The new super call drops the stages=stages, was this intended?
| kwargs = self._input_kwargs | ||
| super().__init__() | ||
| self._setDefault( | ||
| checkpointInterval=0, |
There was a problem hiding this comment.
Default to None instead of 0?
| for param in stage.params: | ||
| if not (param.name.endswith("Col") or param.name.endswith("Cols")): | ||
| continue | ||
| if not stage.isDefined(param): | ||
| continue | ||
| value = stage.getOrDefault(param) | ||
| if isinstance(value, str): | ||
| required_input_columns.add(value) | ||
| elif isinstance(value, (list, tuple)): | ||
| required_input_columns.update( | ||
| item for item in value if isinstance(item, str) | ||
| ) |
There was a problem hiding this comment.
This feels more complex than I think it needs to be. Two things:
- Why do we need to iterate through the stage.params at all if we have already added the inputs using
get_layer_inputs_outputs - Even if we do need, we can just check for presence of
inputColorinputColsas a stage always has strictly one of these defined.if stage.hasParam("inputCol") and stage.isDefined("inputCol"): value = stage.getInputCol()and after check input cols. But this is really whatget_layer_inputs_outputsdoes.
There was a problem hiding this comment.
You're right that inputCol/inputCols are already covered by get_layer_inputs_outputs but the sweep isn't for those.
It's there for the other column params a stage reads that aren't its input col: maskCols/relevanceCol on ConditionalStandardScaleEstimator, and queryIdCol on the listwise transformers. Those get read at fit time, so if we don't keep them, pruneInputColumns=True drops them and the fit blows up with "column not found". There are regression tests covering exactly that.
If the suffix-matching feels too hacky, I'm happy to swap it for a small get_fit_input_columns() hook on the handful of stages that need it instead. Let me know what you'd prefer.
| """ | ||
| required_input_columns = self.collect_required_input_columns(stages) | ||
| columns_to_keep = [c for c in dataset.columns if c in required_input_columns] | ||
| if columns_to_keep and len(columns_to_keep) < len(dataset.columns): |
There was a problem hiding this comment.
Pedantic but do we need that second condition? If we are inside this function then we are pruning, and columns_to_keep is always a subset. So I would just check it's not empty and otherwise select
| fit. 0 (or None) disables checkpointing. | ||
| :returns: KamaeSparkPipeline object with checkpointInterval set. | ||
| """ | ||
| return self._set(checkpointInterval=value) |
There was a problem hiding this comment.
We should error on checkpoint interval being negative here. Personally I think we should error on 0 too and treat None as the no checkpoint behaviour
| :returns: KamaeSparkPipeline object with params set. | ||
| """ | ||
| kwargs = self._input_kwargs | ||
| return self._set(**kwargs) |
There was a problem hiding this comment.
This set params does not use the setter methods at all. So any validation in them will not be respected when the user passes arguments to the init as opposed to using the setter method. If you check how I defined the setParams for the estimator and transformer you can see I use the setter method.
| return None | ||
| # We add 1 because we want to reserve the 0 index for mask/padding. | ||
| return bisect_right(splits, value) + 1 | ||
| def bucketize(value: Column) -> Column: |
There was a problem hiding this comment.
Same to assume that all this bucket logic changes to this transformer are not meant to be in this PR?
There was a problem hiding this comment.
I was looking for inefficiencies in the library... I was going to patch more but only bucketize was really impacted in a way I could speed up simply. We don't really use it but figured it was a good improvement so may as well include.
There was a problem hiding this comment.
Hmmm ok if you want to keep it I will have more comments, will add them then you can decide.
| dependencies = [ | ||
| "pyspark>=3.4.0,<4.0.0", | ||
| "pandas>=1.3.4,<3.0.0", | ||
| "pyarrow>=4.0.0", |
There was a problem hiding this comment.
Adding new dependencies needs a uv lock pls
There was a problem hiding this comment.
Lock updated
- Restore stages=stages in __init__ super call - checkpointInterval defaults to None; reject non-positive via setter - Route setParams through setter methods so validation runs - Drop redundant length check in prune_unused_input_columns - Regenerate uv.lock to include pyarrow (required by pandas_udf) Retains the aux-column sweep in collect_required_input_columns: it is load-bearing for pruning correctness (maskCols/relevanceCol/queryIdCol are not returned by get_layer_inputs_outputs) and defended by regression tests. Co-Authored-By: Claude Opus 4.7 <noreply@anthropic.com>
georyetti
left a comment
There was a problem hiding this comment.
Effectively fine, just revert the type hint for bucketize back to tf layer as its still tf only
| for param in stage.params: | ||
| if not (param.name.endswith("Col") or param.name.endswith("Cols")): | ||
| continue | ||
| if not stage.isDefined(param): | ||
| continue | ||
| value = stage.getOrDefault(param) | ||
| if isinstance(value, str): | ||
| required_input_columns.add(value) | ||
| elif isinstance(value, (list, tuple)): | ||
| required_input_columns.update( | ||
| item for item in value if isinstance(item, str) | ||
| ) |
| return None | ||
| # We add 1 because we want to reserve the 0 index for mask/padding. | ||
| return bisect_right(splits, value) + 1 | ||
| def bucketize(value: Column) -> Column: |
There was a problem hiding this comment.
Hmmm ok if you want to keep it I will have more comments, will add them then you can decide.
| ) | ||
|
|
||
| def get_keras_layer(self) -> tf.keras.layers.Layer: | ||
| def get_keras_layer(self) -> keras.layers.Layer: |
There was a problem hiding this comment.
Why are we making this type hint as keras layer? The layer still uses tf only functions so we should revert this back
…s.Layer BucketizeLayer is TensorFlow-only, so the tf-specific return type is accurate. Co-Authored-By: Claude Opus 4.7 <noreply@anthropic.com>
…peline Adds a default-off boolean pipeline param that, at the first estimator-fit boundary, projects the working frame to the columns still read downstream and persists (MEMORY_AND_DISK) that narrow frame once, reused by all subsequent estimators. This collapses repeated full scans of a wide input across independent sibling estimators into a single populating scan plus in-RAM reuse. When both cacheIntermediateData and cacheEstimatorInput are enabled, cacheEstimatorInput takes precedence (with a warning) since it is a strictly narrower cache and the intermediate cache would evict it. Co-Authored-By: Claude Opus 4.7 <noreply@anthropic.com>
georyetti
left a comment
There was a problem hiding this comment.
Sorry some minor comments, I don't think its particularly blocking so can approve if you have good reason for it.
| if cached_dataset is not None: | ||
| cached_dataset.unpersist() | ||
| if estimator_input_cache is not None: | ||
| estimator_input_cache.unpersist() |
There was a problem hiding this comment.
Got a little confused here as we have 3 dataframes:
- dataset
- cached_dataset
- new_cached
And we set the first 2 to the 3rd, but then only unpersist the 2nd...?
Can we simplify this to just persist dataset?
| estimator_input_cache = dataset.select(*keep_columns).persist( | ||
| StorageLevel.MEMORY_AND_DISK | ||
| ) | ||
| dataset = estimator_input_cache |
There was a problem hiding this comment.
Similar here, what is the use of estimator_input_cache here if we just set to dataset and then have to remember to unpersist()? Can we not just reuse the dataset variable and unpersist() just that at the end?
…line Adds a default-None float param (0, 1] that draws a single sample of the input up-front, persists (MEMORY_AND_DISK) and materialises it once, and fits every estimator from that shared sample with each estimator's own sampleFraction temporarily disabled and restored afterwards. This collapses the N independent per-estimator Bernoulli scans of a wide source (which spilled and GC-thrashed at scale) into a single populating scan plus in-RAM reuse. fitSampleSeed makes the sample reproducible. fitSampleFraction is incompatible with cacheIntermediateData/cacheEstimatorInput (which persist frames it is designed to avoid), so enabling it warns and disables them. It only computes correct statistics for sample-robust estimators (mean/std/quantiles); a runtime warning documents that vocabulary builders, min/max scalers and distinct counts need exact statistics. Also addresses review feedback on the cache bookkeeping: since the two cache strategies are mutually exclusive, collapse the separate cached_dataset / estimator_input_cache / new_cached handles into a single persisted_frame kept distinct from dataset (which model.transform reassigns), with one unpersist. Co-Authored-By: Claude Opus 4.7 <noreply@anthropic.com>
|
Pushed a new commit ( 1. New opt-in param: This is a fit-time optimisation for the case that motivated the caching work: a pipeline of several independent estimators (each Insight: the fit only needs a sample. So when
The returned model is unchanged — transformers still apply to the full dataset at transform time; only fit-time stats come from the sample. Caveats (also in the param docstring + a runtime warning):
Tests cover: Note: this is correctness-tested, but I have not benchmarked the actual speedup at prod scale — the single-scan behaviour is proven, the wall-clock win is expected but unmeasured. 2. Addressed your two latest comments on the cache bookkeeping Since |
Replace the all-estimators sampling behaviour of fitSampleFraction with a per-estimator opt-in boolean (useFitSample) on SampleFractionParams. Only estimators with useFitSample=True fit on the shared pipeline sample; all others fit on the full input, so vocabulary builders and min/max scalers that need exact/global statistics stay correct. Warn on the two conflicting configurations: an estimator that sets both useFitSample=True and its own sampleFraction (shared sample wins), and useFitSample=True with no pipeline fitSampleFraction (no-op). Co-Authored-By: Claude Opus 4.7 <noreply@anthropic.com>
… load The scalar pandas_udf path let Arrow deliver Spark NULLs as NaN/pd.NA for numeric series, bypassing the `is None` null/OOV guards in element funcs. Restore Python None before mapping, guarded by hasnans so the null-free fast path (the common case) keeps its speedup. Also persist the pipeline-level fit params (checkpointInterval, cache/prune flags, fitSampleFraction/Seed) in the pipeline writer's metadata and restore them in the reader, so non-default values survive a save/load round-trip instead of silently resetting to defaults. Co-Authored-By: Claude Opus 4.7 <noreply@anthropic.com>
georyetti
left a comment
There was a problem hiding this comment.
Just one minor nitpick, we have kamae classes for reading and writing metadata/params, could we use them. We added these due to slowdowns in how metadata get written on databricks
| for p in self.instance.params | ||
| if p.name != "stages" and self.instance.isSet(p) | ||
| } | ||
| DefaultParamsWriter.saveMetadata( |
There was a problem hiding this comment.
Can we use the kamae classes we have here: KamaeDefaultParamsWriter
| @@ -754,7 +831,15 @@ def load(self, path: str) -> KamaeSparkPipeline: | |||
| """ | |||
| metadata = DefaultParamsReader.loadMetadata(path, self.sc) | |||
There was a problem hiding this comment.
Sorry spotted we don't use the reader here too. I know its not your change but could you use here too? KamaeDefaultParamsReader
Swap DefaultParamsReader/DefaultParamsWriter for the Kamae variants so pipeline metadata read/write uses the Databricks fast-path workaround the rest of kamae already relies on, keeping the write and read sides consistent. Co-Authored-By: Claude Opus 4.7 <noreply@anthropic.com>
Sbranikas
left a comment
There was a problem hiding this comment.
Just a minor question on the unpersist().
| if sample_dataset is not None: | ||
| sample_dataset = model.transform(sample_dataset) | ||
| if persisted_frame is not None: | ||
| persisted_frame.unpersist() |
There was a problem hiding this comment.
Just a small question, but right now only sampled_dataset gets cleaned up on failure (via the outer _fit's finally) — persisted_frame doesn't, right?
Description
Provide a short description of the PR changes.
The below checklists come from the docs page on adding new transformers here
Keras Layer Checklist
Verify that:
_callmethod has been implemented in the new layer.compatible_dtypesproperty is defined in the new layer.@tf.keras.utils.register_keras_serializable(package=kamae.__name__).name,input_dtype, andoutput_dtypeas arguments to the constructor and that this is passed to the super constructor.get_configmethod.layersdirectory.Spark Transformer/Estimator Checklist
Verify that:
__init__andsetParamsmethods.Paramsclass here.compatible_dtypesproperty has been implemented to specify the input/output data types that my transformer/estimator supports.get_tf_layermethod.transformers/estimatorsdirectory.Finally, please verify that: