@@ -57,18 +57,25 @@ public SparkMLPipeline(Map<FieldDefinition, Feature<DataType>> featurers, ModelC
5757 public void buildPipeline (Map <FieldDefinition , Feature <DataType >> featurers , ModelColumnHelper columnHelper ) {
5858 pipelineStage = new ArrayList <>();
5959
60- featureCreators = new SparkFeatureCreators (featurers , columnHelper );
60+ featureCreators = getFeatureCreators (featurers );
6161 pipelineStage .addAll (featureCreators .getTransformers ());
6262
6363 pipelineStage .add (getAssembler ());
6464 pipelineStage .add (getPolyExpansion ());
6565 pipelineStage .add (getLR ());
6666
67- vve = new VectorValueExtractor (ColName .PROBABILITY_COL , ColName .SCORE_COL );
68- columnHelper .getColumnsAdded ().add (ColName .PROBABILITY_COL );
69- columnHelper .getColumnsAdded ().add (ColName .RAW_PREDICTION );
67+ //vve is not used in all cases, but since we dont have a different flow for prediction
68+ //creating it here so it gets registered
69+ vve = getVVE ();
70+
71+
7072 }
7173
74+ protected SparkFeatureCreators getFeatureCreators (Map <FieldDefinition , Feature <DataType >> featurers ){
75+ featureCreators = new SparkFeatureCreators (featurers , columnHelper );
76+ return featureCreators ;
77+ }
78+
7279 protected VectorAssembler getAssembler () {
7380 VectorAssembler assembler = new VectorAssembler ();
7481 assembler .setInputCols (columnHelper .getColumnsAdded ().toArray (new String [0 ]))
@@ -123,11 +130,18 @@ public ZFrame<Dataset<Row>, Row, Column> fit(ZFrame<Dataset<Row>, Row, Column> i
123130 public ZFrame <Dataset <Row >, Row , Column > predict (ZFrame <Dataset <Row >, Row , Column > data ) {
124131 LOG .info ("threshold while predicting is " + lr .getThreshold ());
125132 Dataset <Row > predictWithFeatures = transformer .transform (data .df ());
126- predictWithFeatures = vve .transform (predictWithFeatures );
133+ predictWithFeatures = getVVE () .transform (predictWithFeatures );
127134 LOG .debug ("Return schema is " + predictWithFeatures .schema ());
128135 return new SparkFrame (predictWithFeatures );
129136 }
130137
138+ public VectorValueExtractor getVVE () {
139+ vve = new VectorValueExtractor (ColName .PROBABILITY_COL , ColName .SCORE_COL );
140+ columnHelper .getColumnsAdded ().add (ColName .PROBABILITY_COL );
141+ columnHelper .getColumnsAdded ().add (ColName .RAW_PREDICTION );
142+ return vve ;
143+ }
144+
131145 public void register (SparkSession session ) {
132146 featureCreators .register (session );
133147 vve .register (session );
@@ -141,7 +155,5 @@ public void save(String path) throws IOException {
141155 ((CrossValidatorModel ) transformer ).write ().overwrite ().save (path );
142156 }
143157
144- public List <SparkTransformer > getFeatureCreators () {
145- return featureCreators .getTransformers ();
146- }
158+
147159}
0 commit comments