Skip to content

Commit 620c521

Browse files
committed
Refactored H2OEstimator#encodeLabel(List<String>, SkLearnEncoder) method
1 parent b5ece23 commit 620c521

1 file changed

Lines changed: 21 additions & 4 deletions

File tree

pmml-sklearn-h2o/src/main/java/h2o/estimators/H2OEstimator.java

Lines changed: 21 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -26,6 +26,7 @@
2626
import java.util.List;
2727

2828
import hex.genmodel.MojoModel;
29+
import hex.genmodel.algos.glm.GlmOrdinalMojoModel;
2930
import org.dmg.pmml.DataField;
3031
import org.dmg.pmml.DataType;
3132
import org.dmg.pmml.MiningFunction;
@@ -38,7 +39,9 @@
3839
import org.jpmml.converter.FeatureList;
3940
import org.jpmml.converter.FeatureUtil;
4041
import org.jpmml.converter.Label;
42+
import org.jpmml.converter.OrdinalLabel;
4143
import org.jpmml.converter.PMMLEncoder;
44+
import org.jpmml.converter.ScalarLabelUtil;
4245
import org.jpmml.converter.Schema;
4346
import org.jpmml.converter.TypeUtil;
4447
import org.jpmml.h2o.Converter;
@@ -131,6 +134,7 @@ public boolean hasProbabilityDistribution(){
131134
@Override
132135
public Label encodeLabel(List<String> names, SkLearnEncoder encoder){
133136
String estimatorType = getEstimatorType();
137+
MojoModel mojoModel = getMojoModel();
134138

135139
ClassDictUtil.checkSize(1, names);
136140

@@ -141,24 +145,37 @@ public Label encodeLabel(List<String> names, SkLearnEncoder encoder){
141145
{
142146
List<?> categories = getClasses();
143147

148+
OpType opType = OpType.CATEGORICAL;
144149
DataType dataType = TypeUtil.getDataType(categories, DataType.STRING);
145150

151+
// XXX
152+
if(mojoModel instanceof GlmOrdinalMojoModel){
153+
opType = OpType.ORDINAL;
154+
} // End if
155+
146156
if(name != null){
147-
DataField dataField = encoder.createDataField(name, OpType.CATEGORICAL, dataType, categories);
157+
DataField dataField = encoder.createDataField(name, opType, dataType, categories);
148158

149-
return new CategoricalLabel(dataField);
159+
return ScalarLabelUtil.createScalarLabel(dataField);
150160
} else
151161

152162
{
153-
return new CategoricalLabel(dataType, categories);
163+
switch(opType){
164+
case CATEGORICAL:
165+
return new CategoricalLabel(dataType, categories);
166+
case ORDINAL:
167+
return new OrdinalLabel(dataType, categories);
168+
default:
169+
throw new IllegalArgumentException();
170+
}
154171
}
155172
}
156173
case "regressor":
157174
{
158175
if(name != null){
159176
DataField dataField = encoder.createDataField(name, OpType.CONTINUOUS, DataType.DOUBLE);
160177

161-
return new ContinuousLabel(dataField);
178+
return ScalarLabelUtil.createScalarLabel(dataField);
162179
} else
163180

164181
{

0 commit comments

Comments
 (0)