public class CrossValidatorModel extends Model<CrossValidatorModel> implements MLWritable
param: bestModel The best model selected from k-fold cross validation.
param: avgMetrics Average cross-validation metrics for each paramMap in
CrossValidator.estimatorParamMaps, in the corresponding order.
| Modifier and Type | Method and Description |
|---|---|
double[] |
avgMetrics() |
java.lang.Object |
bestModel() |
CrossValidatorModel |
copy(ParamMap extra)
Creates a copy of this instance with the same UID and some extra params.
|
Param<Estimator<?>> |
estimator()
param for the estimator to be validated
|
Param<ParamMap[]> |
estimatorParamMaps()
param for estimator param maps
|
Param<Evaluator> |
evaluator()
param for the evaluator used to select hyper-parameters that maximize the validated metric
|
Estimator<?> |
getEstimator() |
ParamMap[] |
getEstimatorParamMaps() |
Evaluator |
getEvaluator() |
int |
getNumFolds() |
static CrossValidatorModel |
load(java.lang.String path) |
IntParam |
numFolds()
Param for number of folds for cross validation.
|
static MLReader<CrossValidatorModel> |
read() |
DataFrame |
transform(DataFrame dataset)
Transforms the input dataset.
|
StructType |
transformSchema(StructType schema)
:: DeveloperApi ::
|
java.lang.String |
uid()
An immutable unique ID for the object and its derivatives.
|
void |
validateParams() |
MLWriter |
write()
Returns an
MLWriter instance for this ML instance. |
transform, transform, transformtransformSchemaclone, equals, finalize, getClass, hashCode, notify, notifyAll, toString, wait, wait, waitclear, copyValues, defaultCopy, defaultParamMap, explainParam, explainParams, extractParamMap, extractParamMap, get, getDefault, getOrDefault, getParam, hasDefault, hasParam, isDefined, isSet, paramMap, params, set, set, set, setDefault, setDefault, shouldOwntoStringsaveinitializeIfNecessary, initializeLogging, isTraceEnabled, log_, log, logDebug, logDebug, logError, logError, logInfo, logInfo, logName, logTrace, logTrace, logWarning, logWarningpublic static MLReader<CrossValidatorModel> read()
public static CrossValidatorModel load(java.lang.String path)
public java.lang.String uid()
Identifiableuid in interface Identifiablepublic java.lang.Object bestModel()
public double[] avgMetrics()
public void validateParams()
validateParams in interface Paramspublic DataFrame transform(DataFrame dataset)
Transformertransform in class Transformerdataset - (undocumented)public StructType transformSchema(StructType schema)
PipelineStageDerives the output schema from the input schema.
transformSchema in class PipelineStageschema - (undocumented)public CrossValidatorModel copy(ParamMap extra)
Paramscopy in interface Paramscopy in class Model<CrossValidatorModel>extra - (undocumented)defaultCopy()public MLWriter write()
MLWritableMLWriter instance for this ML instance.write in interface MLWritablepublic IntParam numFolds()
public int getNumFolds()
public Param<Estimator<?>> estimator()
public Estimator<?> getEstimator()
public Param<ParamMap[]> estimatorParamMaps()
public ParamMap[] getEstimatorParamMaps()
public Param<Evaluator> evaluator()
public Evaluator getEvaluator()