Class LibLinearClassificationTrainer
java.lang.Object
org.tribuo.common.liblinear.LibLinearTrainer<Label>
org.tribuo.classification.liblinear.LibLinearClassificationTrainer
- All Implemented Interfaces:
com.oracle.labs.mlrg.olcut.config.Configurable,com.oracle.labs.mlrg.olcut.provenance.Provenancable<org.tribuo.provenance.TrainerProvenance>,WeightedLabels,org.tribuo.Trainer<Label>
public class LibLinearClassificationTrainer
extends LibLinearTrainer<Label>
implements WeightedLabels
A
Trainer which wraps a liblinear-java classifier trainer.
See:
Fan RE, Chang KW, Hsieh CJ, Wang XR, Lin CJ. "LIBLINEAR: A library for Large Linear Classification" Journal of Machine Learning Research, 2008.and for the original algorithm:
Cortes C, Vapnik V. "Support-Vector Networks" Machine Learning, 1995.
-
Field Summary
Fields inherited from class org.tribuo.common.liblinear.LibLinearTrainer
cost, epsilon, libLinearParams, maxIterations, seed, terminationCriterion, trainerTypeFields inherited from interface org.tribuo.Trainer
DEFAULT_SEED, INCREMENT_INVOCATION_COUNT -
Constructor Summary
ConstructorsConstructorDescriptionCreates a trainer using the default values (LinearClassificationType.LinearType.L2R_L2LOSS_SVC_DUAL, 1, 0.1,Trainer.DEFAULT_SEED).LibLinearClassificationTrainer(LinearClassificationType trainerType, double cost, double terminationCriterion) Creates a trainer for a LibLinearClassificationModel.LibLinearClassificationTrainer(LinearClassificationType trainerType, double cost, int maxIterations, double terminationCriterion) Creates a trainer for a LibLinear modelLibLinearClassificationTrainer(LinearClassificationType trainerType, double cost, int maxIterations, double terminationCriterion, long seed) Creates a trainer for a LibLinear model -
Method Summary
Modifier and TypeMethodDescriptionprotected LibLinearModel<Label> createModel(org.tribuo.provenance.ModelProvenance provenance, org.tribuo.ImmutableFeatureMap featureIDMap, org.tribuo.ImmutableOutputInfo<Label> outputIDInfo, List<de.bwaldvogel.liblinear.Model> models) protected com.oracle.labs.mlrg.olcut.util.Pair<de.bwaldvogel.liblinear.FeatureNode[][], double[][]> extractData(org.tribuo.Dataset<Label> data, org.tribuo.ImmutableOutputInfo<Label> outputInfo, org.tribuo.ImmutableFeatureMap featureMap) voidUsed by the OLCUT configuration system, and should not be called by external code.voidsetLabelWeights(Map<Label, Float> weights) protected de.bwaldvogel.liblinear.ParametersetupParameters(org.tribuo.ImmutableOutputInfo<Label> labelIDMap) protected List<de.bwaldvogel.liblinear.Model> trainModels(de.bwaldvogel.liblinear.Parameter curParams, int numFeatures, de.bwaldvogel.liblinear.FeatureNode[][] features, double[][] outputs) Methods inherited from class org.tribuo.common.liblinear.LibLinearTrainer
exampleToNodes, getInvocationCount, getProvenance, setInvocationCount, toString, train, train, train
-
Constructor Details
-
LibLinearClassificationTrainer
public LibLinearClassificationTrainer()Creates a trainer using the default values (LinearClassificationType.LinearType.L2R_L2LOSS_SVC_DUAL, 1, 0.1,Trainer.DEFAULT_SEED). -
LibLinearClassificationTrainer
public LibLinearClassificationTrainer(LinearClassificationType trainerType, double cost, double terminationCriterion) Creates a trainer for a LibLinearClassificationModel.Uses
Trainer.DEFAULT_SEEDas the RNG seed. Sets maxIterations to 1000.- Parameters:
trainerType- Loss function and optimisation method combination.cost- Cost penalty for each incorrectly classified training point.terminationCriterion- How close does the optimisation function need to be before terminating that subproblem (usually set to 0.1).
-
LibLinearClassificationTrainer
public LibLinearClassificationTrainer(LinearClassificationType trainerType, double cost, int maxIterations, double terminationCriterion) Creates a trainer for a LibLinear modelUses
Trainer.DEFAULT_SEEDas the RNG seed.- Parameters:
trainerType- Loss function and optimisation method combination.cost- Cost penalty for each incorrectly classified training point.maxIterations- The maximum number of dataset iterations.terminationCriterion- How close does the optimisation function need to be before terminating that subproblem (usually set to 0.1).
-
LibLinearClassificationTrainer
public LibLinearClassificationTrainer(LinearClassificationType trainerType, double cost, int maxIterations, double terminationCriterion, long seed) Creates a trainer for a LibLinear model- Parameters:
trainerType- Loss function and optimisation method combination.cost- Cost penalty for each incorrectly classified training point.maxIterations- The maximum number of dataset iterations.terminationCriterion- How close does the optimisation function need to be before terminating that subproblem (usually set to 0.1).seed- The RNG seed.
-
-
Method Details
-
postConfig
Used by the OLCUT configuration system, and should not be called by external code.- Specified by:
postConfigin interfacecom.oracle.labs.mlrg.olcut.config.Configurable- Overrides:
postConfigin classLibLinearTrainer<Label>
-
trainModels
protected List<de.bwaldvogel.liblinear.Model> trainModels(de.bwaldvogel.liblinear.Parameter curParams, int numFeatures, de.bwaldvogel.liblinear.FeatureNode[][] features, double[][] outputs) - Specified by:
trainModelsin classLibLinearTrainer<Label>
-
createModel
protected LibLinearModel<Label> createModel(org.tribuo.provenance.ModelProvenance provenance, org.tribuo.ImmutableFeatureMap featureIDMap, org.tribuo.ImmutableOutputInfo<Label> outputIDInfo, List<de.bwaldvogel.liblinear.Model> models) - Specified by:
createModelin classLibLinearTrainer<Label>
-
extractData
protected com.oracle.labs.mlrg.olcut.util.Pair<de.bwaldvogel.liblinear.FeatureNode[][], double[][]> extractData(org.tribuo.Dataset<Label> data, org.tribuo.ImmutableOutputInfo<Label> outputInfo, org.tribuo.ImmutableFeatureMap featureMap) - Specified by:
extractDatain classLibLinearTrainer<Label>
-
setupParameters
protected de.bwaldvogel.liblinear.Parameter setupParameters(org.tribuo.ImmutableOutputInfo<Label> labelIDMap) - Overrides:
setupParametersin classLibLinearTrainer<Label>
-
setLabelWeights
- Specified by:
setLabelWeightsin interfaceWeightedLabels
-