=%.3f - $d/k=%.3f = %.3f" +
" (<1/n_yz>=%.3f, <1/n_xz>=%.3f)\n",
MathsUtils.digamma(k), averageDiGammas,
averageInverseCountInJointXZ + averageInverseCountInJointYZ,
- (2.0 / (double)k), averageMeasure,
+ dimensionsCond > 0 ? 2 : 1,
+ inverseKTerm, averageMeasure,
averageInverseCountInJointYZ, averageInverseCountInJointXZ);
}
double[] result = new double[1];
@@ -619,6 +662,40 @@ public abstract class ConditionalMutualInfoCalculatorMultiVariateKraskov
protected abstract double[] partialComputeFromObservations(
int startTimePoint, int numTimePoints, boolean returnLocals) throws Exception;
+ /**
+ * Protected method to be used internally for threaded implementations.
+ * This method implements the guts of each Kraskov algorithm, computing the number of
+ * nearest neighbours in each dimension for a sub-set of the data points.
+ * It is intended to be called by one thread to work on that specific
+ * sub-set of the data.
+ * In particular, this method differs from {@link #partialComputeFromObservations(int, int, boolean)}
+ * because it operates on a new set of observations (using the old set of observations for
+ * constructing the search spaces and PDFs)
+ *
+ * The method returns:
+ * - for average conditional MIs (returnLocals == false), the relevant sums of digamma(n_{xz}), digamma(n_{yz})
+ * and digamma(n_z)
+ * for a partial set of the observations
+ * - for local conditional MIs (returnLocals == true), the array of local conditional MI values
+ *
+ *
+ * @param startTimePoint start time for the partial set we examine
+ * @param numTimePoints number of time points (including startTimePoint to examine)
+ * @param newVar1Observations new time series of observations for variable 1
+ * @param newVar2Observations new time series of observations for variable 1
+ * @param newCondObservations new time series of observations for variable 1
+ * @param returnLocals whether to return an array or local values, or else
+ * sums of these values
+ * @return an array of the relevant sum of digamma(n_xz+1) and digamma(n_yz+1) and digamma(n_z), then
+ * sum of n_xz, n_yz, n_z and for algorithm 2 only, sum of 1/n_xz and 1/n_yz
+ * (these latter five are only for debugging purposes).
+ * @throws Exception
+ */
+ protected abstract double[] partialComputeFromNewObservations(
+ int startTimePoint, int numTimePoints,
+ double[][] newVar1Observations, double[][] newVar2Observations, double[][] newCondObservations,
+ boolean returnLocals) throws Exception;
+
/**
* Private class to handle multi-threading of the Kraskov algorithms.
* Each instance calls partialComputeFromObservations()
@@ -633,6 +710,7 @@ public abstract class ConditionalMutualInfoCalculatorMultiVariateKraskov
protected ConditionalMutualInfoCalculatorMultiVariateKraskov condMiCalc;
protected int myStartTimePoint;
protected int numberOfTimePoints;
+ protected double[][][] newObservations;
protected boolean computeLocals;
protected double[] returnValues = null;
@@ -649,11 +727,13 @@ public abstract class ConditionalMutualInfoCalculatorMultiVariateKraskov
public CondMiKraskovThreadRunner(
ConditionalMutualInfoCalculatorMultiVariateKraskov condMiCalc,
int myStartTimePoint, int numberOfTimePoints,
+ double[][][] newObservations,
boolean computeLocals) {
this.condMiCalc = condMiCalc;
this.myStartTimePoint = myStartTimePoint;
this.numberOfTimePoints = numberOfTimePoints;
this.computeLocals = computeLocals;
+ this.newObservations = newObservations;
}
/**
@@ -676,7 +756,16 @@ public abstract class ConditionalMutualInfoCalculatorMultiVariateKraskov
*/
public void run() {
try {
- returnValues = condMiCalc.partialComputeFromObservations(myStartTimePoint, numberOfTimePoints, computeLocals);
+ if (newObservations == null) {
+ // Computing on existing observations
+ returnValues = condMiCalc.partialComputeFromObservations(myStartTimePoint, numberOfTimePoints, computeLocals);
+ } else {
+ // Computing on new observations
+ returnValues = condMiCalc.partialComputeFromNewObservations(
+ myStartTimePoint, numberOfTimePoints,
+ newObservations[0], newObservations[1],
+ newObservations[2], computeLocals);
+ }
} catch (Exception e) {
// Store the exception for later retrieval
problem = e;
diff --git a/java/source/infodynamics/measures/continuous/kraskov/ConditionalMutualInfoCalculatorMultiVariateKraskov1.java b/java/source/infodynamics/measures/continuous/kraskov/ConditionalMutualInfoCalculatorMultiVariateKraskov1.java
index bfbfa5a..677bc43 100755
--- a/java/source/infodynamics/measures/continuous/kraskov/ConditionalMutualInfoCalculatorMultiVariateKraskov1.java
+++ b/java/source/infodynamics/measures/continuous/kraskov/ConditionalMutualInfoCalculatorMultiVariateKraskov1.java
@@ -18,12 +18,18 @@
package infodynamics.measures.continuous.kraskov;
+import java.util.Arrays;
import java.util.Calendar;
import java.util.PriorityQueue;
import infodynamics.measures.continuous.ConditionalMutualInfoCalculatorMultiVariate;
+import infodynamics.utils.EmpiricalMeasurementDistribution;
+import infodynamics.utils.FirstIndexComparatorDouble;
+import infodynamics.utils.KdTree;
import infodynamics.utils.MathsUtils;
+import infodynamics.utils.MatrixUtils;
import infodynamics.utils.NeighbourNodeData;
+import infodynamics.utils.UnivariateNearestNeighbourSearcher;
/**
* Computes the differential conditional mutual information of two multivariate
@@ -99,7 +105,20 @@ public class ConditionalMutualInfoCalculatorMultiVariateKraskov1
// First element in the PQ is the kth NN,
// and epsilon = kthNnData.distance
NeighbourNodeData kthNnData = nnPQ.poll();
-
+ /*
+ if (debug) {
+ System.out.print("t = " + t + " : for data point: ");
+ MatrixUtils.printArray(System.out, var1Observations[t]);
+ MatrixUtils.printArray(System.out, var2Observations[t]);
+ if (dimensionsCond > 0) {
+ MatrixUtils.printArray(System.out, condObservations[t]);
+ } else {
+ System.out.print("[]");
+ }
+ System.out.printf("t=%d: K=%d NNs found at range %.5f (point %d)\n", t, k, kthNnData.distance, kthNnData.sampleIndex);
+ }
+ */
+
// Now count the points in the conditional space, and
// the var1-conditional and var2-conditional spaces.
// We have 3 coded options for how to do this:
@@ -112,7 +131,7 @@ public class ConditionalMutualInfoCalculatorMultiVariateKraskov1
int n_yz = kdTreeVar2Conditional.countPointsStrictlyWithinR(
t, kthNnData.distance, dynCorrExclTime);
int n_z = nnSearcherConditional.countPointsStrictlyWithinR(
- t, kthNnData.distance, dynCorrExclTime);
+ t, kthNnData.distance, dynCorrExclTime);
*/ // end option A
/* Option B --
@@ -151,8 +170,10 @@ public class ConditionalMutualInfoCalculatorMultiVariateKraskov1
if (debug) {
methodStartTime = Calendar.getInstance().getTimeInMillis();
}
- nnSearcherConditional.findPointsWithinR(t, kthNnData.distance, dynCorrExclTime,
+ if (dimensionsCond > 0) {
+ nnSearcherConditional.findPointsWithinR(t, kthNnData.distance, dynCorrExclTime,
false, isWithinRForConditionals, indicesWithinRForConditionals);
+ }
if (debug) {
conditionalTime += Calendar.getInstance().getTimeInMillis() -
methodStartTime;
@@ -160,15 +181,26 @@ public class ConditionalMutualInfoCalculatorMultiVariateKraskov1
}
// 2. Then compute n_xz and n_yz harnessing our knowledge of
// which points qualified for the conditional already:
- // Don't need to supply dynCorrExclTime in the following, because only
+ // Don't need to supply dynCorrExclTime in most of the following, because only
// points outside of it have been included in 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 (dimensionsCond == 0) {
+ // We're really only counting whether the x space qualifies
+ if (dimensionsVar1 > 1) {
+ n_xz = kdTreeVar1Conditional.countPointsStrictlyWithinR(
+ t, kthNnData.distance, dynCorrExclTime);
+ } else {
+ n_xz = uniNNSearcherVar1.countPointsWithinR(t, kthNnData.distance,
+ dynCorrExclTime, false);
+ }
+ } else {
+ 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() -
@@ -176,12 +208,23 @@ public class ConditionalMutualInfoCalculatorMultiVariateKraskov1
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 (dimensionsCond == 0) {
+ // We're really only counting whether the y space qualifies
+ if (dimensionsVar2 > 1) {
+ n_yz = kdTreeVar2Conditional.countPointsStrictlyWithinR(
+ t, kthNnData.distance, dynCorrExclTime);
+ } else {
+ n_yz = uniNNSearcherVar2.countPointsWithinR(t, kthNnData.distance,
+ dynCorrExclTime, false);
+ }
+ } else {
+ 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() -
@@ -189,8 +232,15 @@ public class ConditionalMutualInfoCalculatorMultiVariateKraskov1
}
// 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;
+ if (dimensionsCond == 0) {
+ n_z = totalObservations - 1; // - 1 to remove the point itself.
+ // Note: This doesn't respect the dynamic correlation exclusion, but
+ // does align with a mutual information calculation here (to give the digamma(N) term).
+ // No need to reset the boolean array
+ } else {
+ for (n_z = 0; indicesWithinRForConditionals[n_z] != -1; n_z++) {
+ isWithinRForConditionals[indicesWithinRForConditionals[n_z]] = false;
+ }
}
// end option C
@@ -206,11 +256,16 @@ public class ConditionalMutualInfoCalculatorMultiVariateKraskov1
if (returnLocals) {
localCondMi[t-startTimePoint] = digammaK - digammaNxzPlusOne - digammaNyzPlusOne + digammaNzPlusOne;
if (debug) {
- System.out.printf("t=%d, n_xz=%d, n_yz=%d, n_z=%d, local=%.4f," +
+ System.out.printf("t=%d, r=%.5f, n_xz=%d, n_yz=%d, n_z=%d, local=%.4f," +
" digamma(n_xz+1)=%.5f, digamma(n_yz+1)=%.5f, digamma(n_z+1)=%.5f, \n",
- t, n_xz, n_yz, n_z, localCondMi[t-startTimePoint],
+ t, kthNnData.distance, n_xz, n_yz, n_z, localCondMi[t-startTimePoint],
digammaNxzPlusOne, digammaNyzPlusOne, digammaNzPlusOne);
}
+ } else if (debug) {
+ System.out.printf("t=%d, n_xz=%d, n_yz=%d, n_z=%d," +
+ " digamma(n_xz+1)=%.5f, digamma(n_yz+1)=%.5f, digamma(n_z+1)=%.5f, \n",
+ t, n_xz, n_yz, n_z,
+ digammaNxzPlusOne, digammaNyzPlusOne, digammaNzPlusOne);
}
}
@@ -244,4 +299,214 @@ public class ConditionalMutualInfoCalculatorMultiVariateKraskov1
return results;
}
}
+
+ /**
+ * Compute the local conditional mutual information values for a new set of
+ * observations, where the PDFs are constructed from the observations added earlier.
+ *
+ * @param startTimePoint
+ * @param numTimePoints
+ * @param newStates1
+ * @param newStates2
+ * @param newCondStates
+ * @return
+ * @throws Exception
+ */
+ @Override
+ protected double[] partialComputeFromNewObservations(int startTimePoint,
+ int numTimePoints, double[][] newStates1,
+ double[][] newStates2, double[][] newCondStates, boolean returnLocals) throws Exception {
+
+ double startTime = Calendar.getInstance().getTimeInMillis();
+
+ double[] localCondMi = null;
+ if (returnLocals) {
+ localCondMi = new double[numTimePoints];
+ }
+
+ // Count the average number of points within eps_xz and eps_yz and eps_z of each point
+ double sumDiGammas = 0;
+ double sumNxz = 0;
+ 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];
+
+ 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,
+ new double[][] {newStates1[t], newStates2[t], newCondStates[t]});
+ knnTime += Calendar.getInstance().getTimeInMillis() -
+ methodStartTime;
+ // First element in the PQ is the kth NN,
+ // and epsilon = kthNnData.distance
+ NeighbourNodeData kthNnData = nnPQ.poll();
+ if (debug) {
+ System.out.print("t = " + t + " : for data point: ");
+ MatrixUtils.printArray(System.out, newStates1[t]);
+ MatrixUtils.printArray(System.out, newStates2[t]);
+ if (dimensionsCond > 0) {
+ MatrixUtils.printArray(System.out, newCondStates[t]);
+ } else {
+ System.out.print("[]");
+ }
+ System.out.printf("t=%d : K=%d NNs found at range %.5f (point %d)\n", t, k, kthNnData.distance, kthNnData.sampleIndex);
+ }
+
+ // Now count the points in the conditional space, and
+ // the var1-conditional and var2-conditional spaces.
+ // We use Option C determined above to do this --
+ // 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:
+ if (debug) {
+ methodStartTime = Calendar.getInstance().getTimeInMillis();
+ }
+ if (dimensionsCond > 0) {
+ nnSearcherConditional.findPointsWithinR(kthNnData.distance,
+ new double[][] {newCondStates[t]},
+ 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:
+ // Don't need to supply dynCorrExclTime in the following, because
+ // it's not relevant for the new points
+ int n_xz;
+ if (dimensionsCond == 0) {
+ // We're really only counting whether the x space qualifies
+ if (dimensionsVar1 > 1) {
+ n_xz = kdTreeVar1Conditional.countPointsWithinR(
+ new double[][] {newStates1[t], newCondStates[t]},
+ kthNnData.distance, false);
+ } else {
+ n_xz = uniNNSearcherVar1.countPointsWithinR(
+ new double[][] {newStates1[t]}, kthNnData.distance, false);
+ }
+ } else {
+ if (dimensionsVar1 > 1) {
+ n_xz = kdTreeVar1Conditional.countPointsWithinR(kthNnData.distance,
+ new double[][] {newStates1[t], newCondStates[t]},
+ false, 1, isWithinRForConditionals);
+ } else { // Generally faster to search only the marginal space if it is univariate
+ n_xz = uniNNSearcherVar1.countPointsWithinR(
+ new double[][] {newStates1[t]}, kthNnData.distance,
+ false, isWithinRForConditionals);
+ }
+ }
+
+ if (debug) {
+ conditionalXTime += Calendar.getInstance().getTimeInMillis() -
+ methodStartTime;
+ methodStartTime = Calendar.getInstance().getTimeInMillis();
+ }
+ int n_yz;
+ if (dimensionsCond == 0) {
+ // We're really only counting whether the y space qualifies
+ if (dimensionsVar2 > 1) {
+ n_yz = kdTreeVar2Conditional.countPointsWithinR(
+ new double[][] {newStates2[t], newCondStates[t]},
+ kthNnData.distance, false);
+ } else {
+ n_yz = uniNNSearcherVar2.countPointsWithinR(
+ new double[][] {newStates2[t]},
+ kthNnData.distance, false);
+ }
+ } else {
+ if (dimensionsVar2 > 1) {
+ n_yz = kdTreeVar2Conditional.countPointsWithinR(kthNnData.distance,
+ new double[][] {newStates2[t], newCondStates[t]},
+ false, 1, isWithinRForConditionals);
+ } else { // Generally faster to search only the marginal space if it is univariate
+ n_yz = uniNNSearcherVar2.countPointsWithinR(
+ new double[][] {newStates2[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;
+ if (dimensionsCond == 0) {
+ n_z = totalObservations - 1; // - 1 to remove the point itself.
+ // no need to reset the boolean array
+ } else {
+ for (n_z = 0; indicesWithinRForConditionals[n_z] != -1; n_z++) {
+ isWithinRForConditionals[indicesWithinRForConditionals[n_z]] = false;
+ }
+ }
+ // end option C
+
+ sumNxz += n_xz;
+ sumNyz += n_yz;
+ sumNz += n_z;
+ // And take the digammas:
+ double digammaNxzPlusOne = MathsUtils.digamma(n_xz+1);
+ double digammaNyzPlusOne = MathsUtils.digamma(n_yz+1);
+ double digammaNzPlusOne = MathsUtils.digamma(n_z+1);
+ sumDiGammas += digammaNzPlusOne - digammaNxzPlusOne - digammaNyzPlusOne;
+
+ if (returnLocals) {
+ localCondMi[t-startTimePoint] = digammaK - digammaNxzPlusOne - digammaNyzPlusOne + digammaNzPlusOne;
+ if (debug) {
+ System.out.printf("t=%d, n_xz=%d, n_yz=%d, n_z=%d, local=%.4f," +
+ " digamma(n_xz+1)=%.5f, digamma(n_yz+1)=%.5f, digamma(n_z+1)=%.5f, \n",
+ t, n_xz, n_yz, n_z, localCondMi[t-startTimePoint],
+ digammaNxzPlusOne, digammaNyzPlusOne, digammaNzPlusOne);
+ }
+ } else if (debug) {
+ double localValue = digammaK - digammaNxzPlusOne - digammaNyzPlusOne + digammaNzPlusOne;
+ System.out.printf("t=%d, n_xz=%d, n_yz=%d, n_z=%d, local=%.4f," +
+ " digamma(n_xz+1)=%.5f, digamma(n_yz+1)=%.5f, digamma(n_z+1)=%.5f, \n",
+ t, n_xz, n_yz, n_z, localValue,
+ digammaNxzPlusOne, digammaNyzPlusOne, digammaNzPlusOne);
+ }
+ }
+
+ if (debug) {
+ Calendar rightNow2 = Calendar.getInstance();
+ long endTime = rightNow2.getTimeInMillis();
+ 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));
+ System.out.printf("%d:%d -- Returning: %.4f, %.4f, %.4f, %.4f\n",
+ startTimePoint, (startTimePoint + numTimePoints),
+ sumDiGammas, sumNxz, sumNyz, sumNz);
+ }
+
+ // Select what to return:
+ if (returnLocals) {
+ return localCondMi;
+ } else {
+ // Pad return array with two values, to allow compatibility in
+ // return length with algorithm 2
+ double[] results = new double[6];
+ results[0] = sumDiGammas;
+ results[1] = sumNxz;
+ results[2] = sumNyz;
+ results[3] = sumNz;
+ return results;
+ }
+ }
+
}
diff --git a/java/source/infodynamics/measures/continuous/kraskov/ConditionalMutualInfoCalculatorMultiVariateKraskov2.java b/java/source/infodynamics/measures/continuous/kraskov/ConditionalMutualInfoCalculatorMultiVariateKraskov2.java
index cc3da9b..e508f37 100755
--- a/java/source/infodynamics/measures/continuous/kraskov/ConditionalMutualInfoCalculatorMultiVariateKraskov2.java
+++ b/java/source/infodynamics/measures/continuous/kraskov/ConditionalMutualInfoCalculatorMultiVariateKraskov2.java
@@ -79,7 +79,15 @@ public class ConditionalMutualInfoCalculatorMultiVariateKraskov2
}
// Constants:
- double twoOnK = 2.0 / (double) k;
+ double inverseKTerm;
+ if (dimensionsCond > 0) {
+ // We will add in the 2/K term as usual
+ inverseKTerm = 2.0 / (double) k;
+ } else {
+ // We will only add in 1/K term since without a conditional there is
+ // technically one less variable in the full joint space
+ inverseKTerm = 1.0 / (double) k;
+ }
// Count the average number of points within eps_xz, eps_yz and eps_z
double sumDiGammas = 0;
@@ -136,36 +144,68 @@ public class ConditionalMutualInfoCalculatorMultiVariateKraskov2
// 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, dynCorrExclTime,
+ if (dimensionsCond > 0) {
+ nnSearcherConditional.findPointsWithinR(t, eps_z, dynCorrExclTime,
true, isWithinRForConditionals, indicesWithinRForConditionals);
+ }
// 2. Then compute n_xz and n_yz harnessing our knowledge of
// which points qualified for the conditional already:
// Don't need to supply dynCorrExclTime in the following, because only
// points outside of it have been included in isWithinRForConditionals
int n_xz;
- if (dimensionsVar1 > 1) {
- // Check only the x variable against eps_x, use existing results for z
- n_xz = kdTreeVar1Conditional.countPointsWithinRs(t,
- new double[] {eps_x, eps_z},
- 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);
+ if (dimensionsCond == 0) {
+ // We're really only counting whether the x space qualifies
+ if (dimensionsVar1 > 1) {
+ n_xz = kdTreeVar1Conditional.countPointsWithinOrOnR(
+ t, eps_x, dynCorrExclTime); // No need to pass eps_z in
+ } else { // Generally faster to search only the marginal space if it is univariate
+ n_xz = uniNNSearcherVar1.countPointsWithinOrOnR(t, eps_x,
+ dynCorrExclTime);
+ }
+ } else {
+ if (dimensionsVar1 > 1) {
+ // Check only the x variable against eps_x, use existing results for z
+ n_xz = kdTreeVar1Conditional.countPointsWithinRs(t,
+ new double[] {eps_x, eps_z},
+ 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.countPointsWithinRs(t,
- new double[] {eps_y, eps_z},
- 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);
+ if (dimensionsCond == 0) {
+ if (dimensionsVar2 > 1) {
+ // Check only the y variable against eps_y, use existing results for z
+ n_yz = kdTreeVar2Conditional.countPointsWithinOrOnR(
+ t, eps_y, dynCorrExclTime); // No need to pass eps_z in
+ } else { // Generally faster to search only the marginal space if it is univariate
+ n_yz = uniNNSearcherVar2.countPointsWithinOrOnR(t, eps_y,
+ dynCorrExclTime);
+ }
+ } else {
+ // We're really only counting whether the y space qualifies
+ if (dimensionsVar2 > 1) {
+ // Check only the y variable against eps_y, use existing results for z
+ n_yz = kdTreeVar2Conditional.countPointsWithinRs(t,
+ new double[] {eps_y, eps_z},
+ 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;
+ if (dimensionsCond == 0) {
+ n_z = totalObservations; // Dropping - 1 to remove the point itself, and
+ // Note: This doesn't respect the dynamic correlation exclusion, but this
+ // does align with a mutual information calculation here (to give the digamma(N) term).
+ // No need to reset the boolean array
+ } else {
+ for (n_z = 0; indicesWithinRForConditionals[n_z] != -1; n_z++) {
+ isWithinRForConditionals[indicesWithinRForConditionals[n_z]] = false;
+ }
}
// end option C
@@ -177,15 +217,25 @@ public class ConditionalMutualInfoCalculatorMultiVariateKraskov2
double digammaNxz = MathsUtils.digamma(n_xz);
double digammaNyz = MathsUtils.digamma(n_yz);
double digammaNz = MathsUtils.digamma(n_z);
- double invN_xz = 1.0/(double) n_xz;
- double invN_yz = 1.0/(double) n_yz;
+ double invN_xz, invN_yz;
+ if (dimensionsCond > 0) {
+ // We will add in the inverse neighbour counts as usual
+ invN_xz = 1.0/(double) n_xz;
+ invN_yz = 1.0/(double) n_yz;
+ } else {
+ // We do not add in the inverse neighbour counts, because they
+ // are technically coming from a single variable neighbour search,
+ // and so there is a 0/n_xz and 0/n_yz contribution
+ invN_xz = 0.0;
+ invN_yz = 0.0;
+ }
sumInverseCountInJointXZ += invN_xz;
sumInverseCountInJointYZ += invN_yz;
double contributionDigammas = digammaNz - digammaNxz - digammaNyz;
sumDiGammas += contributionDigammas;
if (returnLocals) {
- localCondMi[t-startTimePoint] = digammaK - twoOnK + contributionDigammas + invN_xz + invN_yz;
+ localCondMi[t-startTimePoint] = digammaK - inverseKTerm + contributionDigammas + invN_xz + invN_yz;
if (debug) {
// Only tracking this for debugging purposes:
System.out.printf("t=%d, n_xz=%d, n_yz=%d, n_z=%d, 1/n_yz=%.3f, 1/n_xz=%.3f, local=%.4f\n",
@@ -216,4 +266,13 @@ public class ConditionalMutualInfoCalculatorMultiVariateKraskov2
return returnValues;
}
}
+
+ @Override
+ protected double[] partialComputeFromNewObservations(int startTimePoint,
+ int numTimePoints, double[][] newVar1Observations,
+ double[][] newVar2Observations, double[][] newCondObservations,
+ boolean returnLocals) throws Exception {
+ // TODO implement me
+ throw new RuntimeException("Not implemented yet");
+ }
}