From b0fc309d6a0bd0f363f700d1f819a3e4fd93608c Mon Sep 17 00:00:00 2001 From: Pedro Martinez Mediano Date: Fri, 12 Jan 2018 22:22:46 +0000 Subject: [PATCH 1/2] Add overloadings for 1d continuous variables in KSG mixed calc. --- ...CalculatorMultiVariateWithDiscreteKraskov.java | 15 +++++++++++++++ 1 file changed, 15 insertions(+) diff --git a/java/source/infodynamics/measures/mixed/kraskov/MutualInfoCalculatorMultiVariateWithDiscreteKraskov.java b/java/source/infodynamics/measures/mixed/kraskov/MutualInfoCalculatorMultiVariateWithDiscreteKraskov.java index 65b3c24..4f19d2a 100755 --- a/java/source/infodynamics/measures/mixed/kraskov/MutualInfoCalculatorMultiVariateWithDiscreteKraskov.java +++ b/java/source/infodynamics/measures/mixed/kraskov/MutualInfoCalculatorMultiVariateWithDiscreteKraskov.java @@ -317,6 +317,13 @@ public class MutualInfoCalculatorMultiVariateWithDiscreteKraskov implements Mutu } } + public void addObservations(double[] continuousObservations, + int[] discreteObservations) throws Exception { + double[][] observationsMatrix = new double[continuousObservations.length][1]; + MatrixUtils.copyIntoColumn(observationsMatrix, 0, continuousObservations); + addObservations(observationsMatrix, discreteObservations); + } + public void addObservations(double[][] source, double[][] destination, int startTime, int numTimeSteps) throws Exception { throw new RuntimeException("Not implemented yet"); } @@ -444,6 +451,14 @@ public class MutualInfoCalculatorMultiVariateWithDiscreteKraskov implements Mutu finaliseAddObservations(); } + public void setObservations(double[] continuousObservations, + int[] discreteObservations) throws Exception { + double[][] observationsMatrix = new double[continuousObservations.length][1]; + MatrixUtils.copyIntoColumn(observationsMatrix, 0, continuousObservations); + setObservations(observationsMatrix, discreteObservations); + } + + /** * Internal method to ensure that the Kd-tree data structures to represent the * observational data have been constructed (should be called prior to attempting From 08f25beee605cde4dfd46c857cb0a69dd570b458 Mon Sep 17 00:00:00 2001 From: Pedro Martinez Mediano Date: Fri, 12 Jan 2018 22:23:12 +0000 Subject: [PATCH 2/2] Add tests for new overloadings in KSG mixed calc. --- ...MultiVariateWithDiscreteKraskovTester.java | 42 +++++++++++++++++++ 1 file changed, 42 insertions(+) diff --git a/java/unittests/infodynamics/measures/mixed/kraskov/MutualInfoMultiVariateWithDiscreteKraskovTester.java b/java/unittests/infodynamics/measures/mixed/kraskov/MutualInfoMultiVariateWithDiscreteKraskovTester.java index f692491..4bfdf96 100755 --- a/java/unittests/infodynamics/measures/mixed/kraskov/MutualInfoMultiVariateWithDiscreteKraskovTester.java +++ b/java/unittests/infodynamics/measures/mixed/kraskov/MutualInfoMultiVariateWithDiscreteKraskovTester.java @@ -269,4 +269,46 @@ public class MutualInfoMultiVariateWithDiscreteKraskovTester extends TestCase { // assertTrue(Math.abs(res5 - res1) > 0.001); } + + public void testUnivariateOverloadings() throws Exception { + MutualInfoCalculatorMultiVariateWithDiscreteKraskov miCalc = + new MutualInfoCalculatorMultiVariateWithDiscreteKraskov(); + + // Generate data sets + RandomGenerator rg = new RandomGenerator(); + double[][] contDataMatrix = rg.generateRandomData(100, 2); + double[] contDataVector = rg.generateRandomData(100); + int[] discData = rg.generateRandomInts(100, 2); + + // Check that no exception is thrown for ok data: + boolean caughtException = false; + try { + miCalc.initialise(1, 2); + miCalc.setObservations(contDataVector, discData); + } catch (Exception e) { + caughtException = true; + } + assertFalse(caughtException); + + // Check that exception is thrown when dimensions do not match + caughtException = false; + try { + miCalc.initialise(1, 2); + miCalc.setObservations(contDataMatrix, discData); + } catch (Exception e) { + caughtException = true; + } + assertTrue(caughtException); + + caughtException = false; + try { + miCalc.initialise(2, 2); + miCalc.setObservations(contDataVector, discData); + } catch (Exception e) { + caughtException = true; + } + assertTrue(caughtException); + + + } }