use of org.apache.ignite.ml.optimization.updatecalculators.SimpleGDUpdateCalculator in project ignite by apache.
the class OneVsRestTrainerTest method testUpdate.
/**
*/
@Test
public void testUpdate() {
Map<Integer, double[]> cacheMock = new HashMap<>();
for (int i = 0; i < twoLinearlySeparableClasses.length; i++) cacheMock.put(i, twoLinearlySeparableClasses[i]);
LogisticRegressionSGDTrainer binaryTrainer = new LogisticRegressionSGDTrainer().withUpdatesStgy(new UpdatesStrategy<>(new SimpleGDUpdateCalculator(0.2), SimpleGDParameterUpdate.SUM_LOCAL, SimpleGDParameterUpdate.AVG)).withMaxIterations(1000).withLocIterations(10).withBatchSize(100).withSeed(123L);
OneVsRestTrainer<LogisticRegressionModel> trainer = new OneVsRestTrainer<>(binaryTrainer);
Vectorizer<Integer, double[], Integer, Double> vectorizer = new DoubleArrayVectorizer<Integer>().labeled(Vectorizer.LabelCoordinate.FIRST);
MultiClassModel originalMdl = trainer.fit(cacheMock, parts, vectorizer);
MultiClassModel updatedOnSameDS = trainer.update(originalMdl, cacheMock, parts, vectorizer);
MultiClassModel updatedOnEmptyDS = trainer.update(originalMdl, new HashMap<>(), parts, vectorizer);
List<Vector> vectors = Arrays.asList(VectorUtils.of(-100, 0), VectorUtils.of(100, 0));
for (Vector vec : vectors) {
TestUtils.assertEquals(originalMdl.predict(vec), updatedOnSameDS.predict(vec), PRECISION);
TestUtils.assertEquals(originalMdl.predict(vec), updatedOnEmptyDS.predict(vec), PRECISION);
}
}
Aggregations