diff --git a/java/source/infodynamics/measures/continuous/kraskov/ConditionalMutualInfoCalculatorMultiVariateKraskov.java b/java/source/infodynamics/measures/continuous/kraskov/ConditionalMutualInfoCalculatorMultiVariateKraskov.java index 09fa30d..6783c49 100755 --- a/java/source/infodynamics/measures/continuous/kraskov/ConditionalMutualInfoCalculatorMultiVariateKraskov.java +++ b/java/source/infodynamics/measures/continuous/kraskov/ConditionalMutualInfoCalculatorMultiVariateKraskov.java @@ -28,6 +28,7 @@ import infodynamics.utils.KdTree; import infodynamics.utils.MathsUtils; import infodynamics.utils.MatrixUtils; import infodynamics.utils.NearestNeighbourSearcher; +import infodynamics.utils.UnivariateNearestNeighbourSearcher; /** *

Computes the differential conditional mutual information of two multivariate @@ -148,11 +149,21 @@ public abstract class ConditionalMutualInfoCalculatorMultiVariateKraskov * representing the (var1,conditional) space */ protected KdTree kdTreeVar1Conditional; + /** + * protected univariate neighbour searcher data structure (for fast nearest neighbour searches) + * representing the (var1) space; used only if var1 is univariate + */ + protected UnivariateNearestNeighbourSearcher uniNNSearcherVar1; /** * protected k-d tree data structure (for fast nearest neighbour searches) * representing the (var2,conditional) space */ protected KdTree kdTreeVar2Conditional; + /** + * protected univariate neighbour searcher data structure (for fast nearest neighbour searches) + * representing the (var2) space; used only if var2 is univariate + */ + protected UnivariateNearestNeighbourSearcher uniNNSearcherVar2; /** * protected data structure (for fast nearest neighbour searches) * representing the conditional space. @@ -179,6 +190,8 @@ public abstract class ConditionalMutualInfoCalculatorMultiVariateKraskov kdTreeVar1Conditional = null; kdTreeVar2Conditional = null; nnSearcherConditional = null; + uniNNSearcherVar1 = null; + uniNNSearcherVar2 = null; super.initialise(dimensions1, dimensions2, dimensionsCond); } @@ -314,14 +327,18 @@ public abstract class ConditionalMutualInfoCalculatorMultiVariateKraskov KdTree originalKdTreeJoint = kdTreeJoint; kdTreeJoint = null; // So that it is rebuilt for the new ordering KdTree originalKdTreeVar1Conditional = kdTreeVar1Conditional; + UnivariateNearestNeighbourSearcher originalUniNNSearcherVar1 = uniNNSearcherVar1; KdTree originalKdTreeVar2Conditional = kdTreeVar2Conditional; + UnivariateNearestNeighbourSearcher originalUniNNSearcherVar2 = uniNNSearcherVar2; if (variableToReorder == 1) { originalData = var1Observations; - kdTreeVar1Conditional = null; // So that it is rebuilt for the new ordering + kdTreeVar1Conditional = null; // So that it is rebuilt for the new ordering if required + uniNNSearcherVar1 = null; // So that it is rebuilt for the new ordering if required var1Observations = MatrixUtils.extractSelectedTimePointsReusingArrays(originalData, reordering); } else { originalData = var2Observations; kdTreeVar2Conditional = null; // So that it is rebuilt for the new ordering + uniNNSearcherVar2 = null; // So that it is rebuilt for the new ordering if required var2Observations = MatrixUtils.extractSelectedTimePointsReusingArrays(originalData, reordering); } // Compute the conditional MI @@ -331,9 +348,11 @@ public abstract class ConditionalMutualInfoCalculatorMultiVariateKraskov if (variableToReorder == 1) { var1Observations = originalData; kdTreeVar1Conditional = originalKdTreeVar1Conditional; + uniNNSearcherVar1 = originalUniNNSearcherVar1; } else { var2Observations = originalData; kdTreeVar2Conditional = originalKdTreeVar2Conditional; + uniNNSearcherVar2 = originalUniNNSearcherVar2; } return newCondMI; } @@ -392,17 +411,29 @@ public abstract class ConditionalMutualInfoCalculatorMultiVariateKraskov new double[][][] {var1Observations, var2Observations, condObservations}); kdTreeJoint.setNormType(normType); } - if (kdTreeVar1Conditional == null) { - kdTreeVar1Conditional = new KdTree( - new int[] {dimensionsVar1, dimensionsCond}, - new double[][][] {var1Observations, condObservations}); - kdTreeVar1Conditional.setNormType(normType); + if (dimensionsVar1 > 1) { + if (kdTreeVar1Conditional == null) { + kdTreeVar1Conditional = new KdTree( + new int[] {dimensionsVar1, dimensionsCond}, + new double[][][] {var1Observations, condObservations}); + kdTreeVar1Conditional.setNormType(normType); + } + } else { // Univariate variable 1, so we'll search its space alone as this is faster + if (uniNNSearcherVar1 == null) { + uniNNSearcherVar1 = new UnivariateNearestNeighbourSearcher(var1Observations); + } } - if (kdTreeVar2Conditional == null) { - kdTreeVar2Conditional = new KdTree( - new int[] {dimensionsVar2, dimensionsCond}, - new double[][][] {var2Observations, condObservations}); - kdTreeVar2Conditional.setNormType(normType); + if (dimensionsVar2 > 1) { + if (kdTreeVar2Conditional == null) { + kdTreeVar2Conditional = new KdTree( + new int[] {dimensionsVar2, dimensionsCond}, + new double[][][] {var2Observations, condObservations}); + kdTreeVar2Conditional.setNormType(normType); + } + } else { // Univariate variable 2, so we'll search its space alone as this is faster + if (uniNNSearcherVar2 == null) { + uniNNSearcherVar2 = new UnivariateNearestNeighbourSearcher(var2Observations); + } } if (nnSearcherConditional == null) { nnSearcherConditional = NearestNeighbourSearcher.create(condObservations); diff --git a/java/source/infodynamics/measures/continuous/kraskov/ConditionalMutualInfoCalculatorMultiVariateKraskov1.java b/java/source/infodynamics/measures/continuous/kraskov/ConditionalMutualInfoCalculatorMultiVariateKraskov1.java index 7669270..22105d6 100755 --- a/java/source/infodynamics/measures/continuous/kraskov/ConditionalMutualInfoCalculatorMultiVariateKraskov1.java +++ b/java/source/infodynamics/measures/continuous/kraskov/ConditionalMutualInfoCalculatorMultiVariateKraskov1.java @@ -81,6 +81,9 @@ public class ConditionalMutualInfoCalculatorMultiVariateKraskov1 double sumNyz = 0; double sumNz = 0; + long knnTime = 0, conditionalTime = 0, + conditionalXTime = 0, conditionalYTime = 0; + // Arrays used for fast searching on conditionals with a marginal: boolean[] isWithinRForConditionals = new boolean[totalObservations]; int[] indicesWithinRForConditionals = new int[totalObservations+1]; @@ -88,8 +91,11 @@ public class ConditionalMutualInfoCalculatorMultiVariateKraskov1 for (int t = startTimePoint; t < startTimePoint + numTimePoints; t++) { // Compute eps for this time step by // finding the kth closest neighbour for point t: + long methodStartTime = Calendar.getInstance().getTimeInMillis(); PriorityQueue nnPQ = kdTreeJoint.findKNearestNeighbours(k, t); + knnTime += Calendar.getInstance().getTimeInMillis() - + methodStartTime; // First element in the PQ is the kth NN, // and epsilon = kthNnData.distance NeighbourNodeData kthNnData = nnPQ.poll(); @@ -99,6 +105,8 @@ public class ConditionalMutualInfoCalculatorMultiVariateKraskov1 // We have 3 coded options for how to do this: /* Option A -- straightforward way using each k-d tree separately: + * To use this, need to construct kdTreeVar1Conditional and + * kdTreeVar2Conditional regardless of dimensionsVar1 and 2. int n_xz = kdTreeVar1Conditional.countPointsStrictlyWithinR( t, kthNnData.distance); int n_yz = kdTreeVar2Conditional.countPointsStrictlyWithinR( @@ -140,14 +148,43 @@ public class ConditionalMutualInfoCalculatorMultiVariateKraskov1 // the knowledge of which points made this cut to speed up the searching // in the conditional-marginal spaces: // 1. Identify the n_z points within the conditional boundaries: + if (debug) { + methodStartTime = Calendar.getInstance().getTimeInMillis(); + } nnSearcherConditional.findPointsWithinR(t, kthNnData.distance, false, isWithinRForConditionals, indicesWithinRForConditionals); + if (debug) { + conditionalTime += Calendar.getInstance().getTimeInMillis() - + methodStartTime; + methodStartTime = Calendar.getInstance().getTimeInMillis(); + } // 2. Then compute n_xz and n_yz harnessing our knowledge of // which points qualified for the conditional already: - int n_xz = kdTreeVar1Conditional.countPointsWithinR(t, kthNnData.distance, - false, 1, isWithinRForConditionals); - int n_yz = kdTreeVar2Conditional.countPointsWithinR(t, kthNnData.distance, - false, 1, isWithinRForConditionals); + int n_xz; + if (dimensionsVar1 > 1) { + n_xz = kdTreeVar1Conditional.countPointsWithinR(t, kthNnData.distance, + false, 1, isWithinRForConditionals); + } else { // Generally faster to search only the marginal space if it is univariate + n_xz = uniNNSearcherVar1.countPointsWithinR(t, kthNnData.distance, + false, isWithinRForConditionals); + } + if (debug) { + conditionalXTime += Calendar.getInstance().getTimeInMillis() - + methodStartTime; + methodStartTime = Calendar.getInstance().getTimeInMillis(); + } + int n_yz; + if (dimensionsVar2 > 1) { + n_yz = kdTreeVar2Conditional.countPointsWithinR(t, kthNnData.distance, + false, 1, isWithinRForConditionals); + } else { // Generally faster to search only the marginal space if it is univariate + n_yz = uniNNSearcherVar2.countPointsWithinR(t, kthNnData.distance, + false, isWithinRForConditionals); + } + if (debug) { + conditionalYTime += Calendar.getInstance().getTimeInMillis() - + methodStartTime; + } // 3. Finally, reset our boolean array for its next use while we count n_z: int n_z; for (n_z = 0; indicesWithinRForConditionals[n_z] != -1; n_z++) { @@ -179,6 +216,11 @@ public class ConditionalMutualInfoCalculatorMultiVariateKraskov1 System.out.println("Subset " + startTimePoint + ":" + (startTimePoint + numTimePoints) + " Calculation time: " + ((endTime - startTime)/1000.0) + " sec" ); + System.out.println("Total exec times for: "); + System.out.println("\tknn search: " + (knnTime/1000.0)); + System.out.println("\tz search: " + (conditionalTime/1000.0)); + System.out.println("\tzx search: " + (conditionalXTime/1000.0)); + System.out.println("\tzy search: " + (conditionalYTime/1000.0)); } // Select what to return: diff --git a/java/source/infodynamics/measures/continuous/kraskov/ConditionalMutualInfoCalculatorMultiVariateKraskov2.java b/java/source/infodynamics/measures/continuous/kraskov/ConditionalMutualInfoCalculatorMultiVariateKraskov2.java index 34c8a28..06fd952 100755 --- a/java/source/infodynamics/measures/continuous/kraskov/ConditionalMutualInfoCalculatorMultiVariateKraskov2.java +++ b/java/source/infodynamics/measures/continuous/kraskov/ConditionalMutualInfoCalculatorMultiVariateKraskov2.java @@ -89,6 +89,10 @@ public class ConditionalMutualInfoCalculatorMultiVariateKraskov2 double sumInverseCountInJointYZ = 0; double sumInverseCountInJointXZ = 0; + // Arrays used for fast searching on conditionals with a marginal: + boolean[] isWithinRForConditionals = new boolean[totalObservations]; + int[] indicesWithinRForConditionals = new int[totalObservations+1]; + for (int t = startTimePoint; t < startTimePoint + numTimePoints; t++) { // Compute eps_x and eps_y and eps_z for this time step by // finding the kth closest neighbours for point t: @@ -117,12 +121,50 @@ public class ConditionalMutualInfoCalculatorMultiVariateKraskov2 // than or equal to eps_z, and whose x and z distances are less // than or equal to eps_z and eps_x, and whose y and z distance are less // than or equal to eps_z and eps_y: - int n_xz = kdTreeVar1Conditional.countPointsWithinOrOnRs( + + /* Option A -- straightforward way using each k-d tree separately: + int n_xz = nnSearcherVar1.countPointsWithinOrOnRs( t, new double[] {eps_x, eps_z}); - int n_yz = kdTreeVar2Conditional.countPointsWithinOrOnRs( + int n_yz = nnSearcherVar2.countPointsWithinOrOnRs( t, new double[] {eps_y, eps_z}); int n_z = nnSearcherConditional.countPointsWithinOrOnR( t, eps_z); + */ + + // Option C -- + // Identify the points satisfying the conditional criteria, then use + // the knowledge of which points made this cut to speed up the searching + // in the conditional-marginal spaces: + // 1. Identify the n_z points within the conditional boundaries: + nnSearcherConditional.findPointsWithinR(t, eps_z, + true, isWithinRForConditionals, indicesWithinRForConditionals); + // 2. Then compute n_xz and n_yz harnessing our knowledge of + // which points qualified for the conditional already: + int n_xz; + if (dimensionsVar1 > 1) { + // Check only the x variable against eps_x, use existing results for z + n_xz = kdTreeVar1Conditional.countPointsWithinR(t, eps_x, + true, 1, isWithinRForConditionals); + } else { // Generally faster to search only the marginal space if it is univariate + n_xz = uniNNSearcherVar1.countPointsWithinR(t, eps_x, + true, isWithinRForConditionals); + } + int n_yz; + if (dimensionsVar2 > 1) { + // Check only the y variable against eps_y, use existing results for z + n_yz = kdTreeVar2Conditional.countPointsWithinR(t, eps_y, + true, 1, isWithinRForConditionals); + } else { // Generally faster to search only the marginal space if it is univariate + n_yz = uniNNSearcherVar2.countPointsWithinR(t, eps_y, + true, isWithinRForConditionals); + } + // 3. Finally, reset our boolean array for its next use while we count n_z: + int n_z; + for (n_z = 0; indicesWithinRForConditionals[n_z] != -1; n_z++) { + isWithinRForConditionals[indicesWithinRForConditionals[n_z]] = false; + } + // end option C + sumNxz += n_xz; sumNyz += n_yz; diff --git a/java/source/infodynamics/utils/UnivariateNearestNeighbourSearcher.java b/java/source/infodynamics/utils/UnivariateNearestNeighbourSearcher.java index 3874b14..fdfba7f 100755 --- a/java/source/infodynamics/utils/UnivariateNearestNeighbourSearcher.java +++ b/java/source/infodynamics/utils/UnivariateNearestNeighbourSearcher.java @@ -55,6 +55,15 @@ public class UnivariateNearestNeighbourSearcher extends NearestNeighbourSearcher */ protected int[] indicesInSortedArray = null; + public UnivariateNearestNeighbourSearcher(double[][] data) throws Exception { + // Ideally we would not call the constructor until after the following check, + // but the constructor must come first in Java. + this(MatrixUtils.selectColumn(data, 0)); + if (data[0].length != 1) { + throw new Exception("Cannot define UnivariateNearestNeighbourSearcher for multivariate data"); + } + } + public UnivariateNearestNeighbourSearcher(double[] data) throws Exception { this.originalDataSet = data; numObservations = data.length;