diff --git a/java/source/infodynamics/measures/continuous/ConditionalMutualInfoMultiVariateCommon.java b/java/source/infodynamics/measures/continuous/ConditionalMutualInfoMultiVariateCommon.java index 6f11173..3c6fc0e 100755 --- a/java/source/infodynamics/measures/continuous/ConditionalMutualInfoMultiVariateCommon.java +++ b/java/source/infodynamics/measures/continuous/ConditionalMutualInfoMultiVariateCommon.java @@ -134,6 +134,39 @@ public abstract class ConditionalMutualInfoMultiVariateCommon implements public static final String PROP_NORMALISE = "NORMALISE"; protected boolean normalise = true; + /** + * Cache for the means of each dimension in variable 1, in case we need to normalise + * new observations later + */ + protected double[] var1Means = null; + /** + * Cache for the standard deviations of each dimension in variable 1, in case we need to normalise + * new observations later + */ + protected double[] var1Stds = null; + /** + * Cache for the means of each dimension in variable 2, in case we need to normalise + * new observations later + */ + protected double[] var2Means = null; + /** + * Cache for the standard deviations of each dimension in variable 2, in case we need to normalise + * new observations later + */ + protected double[] var2Stds = null; + /** + * Cache for the means of each dimension in conditional variable, in case we need to normalise + * new observations later + */ + protected double[] condMeans = null; + /** + * Cache for the standard deviations of each dimension in conditional variable, in case we need to normalise + * new observations later + */ + protected double[] condStds = null; + + + @Override public void initialise(int var1Dimensions, int var2Dimensions, int condDimensions) { dimensionsVar1 = var1Dimensions; @@ -146,6 +179,12 @@ public abstract class ConditionalMutualInfoMultiVariateCommon implements var2Observations = null; condObservations = null; addedMoreThanOneObservationSet = false; + var1Means = null; + var1Stds = null; + var2Means = null; + var2Stds = null; + condMeans = null; + condStds = null; } /** @@ -159,7 +198,8 @@ public abstract class ConditionalMutualInfoMultiVariateCommon implements *
  • {@link #PROP_NORMALISE} - whether to normalise the individual * variables to mean 0, standard deviation 1 * (true by default, except for child class - * {@link infodynamics.measures.continuous.gaussian.ConditionalMutualInfoCalculatorMultiVariateGaussian})
  • + * {@link infodynamics.measures.continuous.gaussian.ConditionalMutualInfoCalculatorMultiVariateGaussian} + * for which this property is false and cannot be altered) * * *

    Unknown property values are ignored.

    @@ -208,7 +248,8 @@ public abstract class ConditionalMutualInfoMultiVariateCommon implements // startAddObservations was not called first throw new RuntimeException("User did not call startAddObservations before addObservations"); } - if ((var1.length != var2.length) || (var1.length != cond.length)) { + if ((var1.length != var2.length) || + ((dimensionsCond != 0) && (var1.length != cond.length))) { throw new Exception(String.format("Observation vector lengths (%d, %d and %d) must match!", var1.length, var2.length, cond.length)); } @@ -220,7 +261,7 @@ public abstract class ConditionalMutualInfoMultiVariateCommon implements throw new Exception("Number of joint variables in var2 data " + "does not match the initialised value"); } - if (cond[0].length != dimensionsCond) { + if ((dimensionsCond != 0) && (cond[0].length != dimensionsCond)) { throw new Exception("Number of joint variables in cond data " + "does not match the initialised value"); } @@ -244,8 +285,11 @@ public abstract class ConditionalMutualInfoMultiVariateCommon implements System.arraycopy(var1, startTime, var1ToAdd, 0, numTimeSteps); double[][] var2ToAdd = new double[numTimeSteps][]; System.arraycopy(var2, startTime, var2ToAdd, 0, numTimeSteps); - double[][] condToAdd = new double[numTimeSteps][]; - System.arraycopy(cond, startTime, condToAdd, 0, numTimeSteps); + double[][] condToAdd = null; + if (dimensionsCond != 0) { + condToAdd = new double[numTimeSteps][]; + System.arraycopy(cond, startTime, condToAdd, 0, numTimeSteps); + } addObservations(var1ToAdd, var2ToAdd, condToAdd); } @@ -317,17 +361,27 @@ public abstract class ConditionalMutualInfoMultiVariateCommon implements MatrixUtils.arrayCopy(var2, 0, 0, var2Observations, startObservation, 0, var2.length, dimensionsVar2); - MatrixUtils.arrayCopy(cond, 0, 0, + if (dimensionsCond != 0) { + MatrixUtils.arrayCopy(cond, 0, 0, condObservations, startObservation, 0, cond.length, dimensionsCond); + } // else we can do nothing there startObservation += var2.length; } // Normalise the data if required if (normalise) { - MatrixUtils.normalise(var1Observations); - MatrixUtils.normalise(var2Observations); - MatrixUtils.normalise(condObservations); + var1Means = MatrixUtils.means(var1Observations); + var1Stds = MatrixUtils.stdDevs(var1Observations, var1Means); + MatrixUtils.normalise(var1Observations, var1Means, var1Stds); + var2Means = MatrixUtils.means(var2Observations); + var2Stds = MatrixUtils.stdDevs(var2Observations, var2Means); + MatrixUtils.normalise(var2Observations, var2Means, var2Stds); + if (dimensionsCond != 0) { + condMeans = MatrixUtils.means(condObservations); + condStds = MatrixUtils.stdDevs(condObservations, condMeans); + MatrixUtils.normalise(condObservations, condMeans, condStds); + } } // We don't need to keep the vectors of observation sets anymore: diff --git a/java/source/infodynamics/measures/continuous/gaussian/ConditionalMutualInfoCalculatorMultiVariateGaussian.java b/java/source/infodynamics/measures/continuous/gaussian/ConditionalMutualInfoCalculatorMultiVariateGaussian.java index c108cb9..8a7aed7 100755 --- a/java/source/infodynamics/measures/continuous/gaussian/ConditionalMutualInfoCalculatorMultiVariateGaussian.java +++ b/java/source/infodynamics/measures/continuous/gaussian/ConditionalMutualInfoCalculatorMultiVariateGaussian.java @@ -158,6 +158,22 @@ public class ConditionalMutualInfoCalculatorMultiVariateGaussian var2IndicesInCovariance = null; } + @Override + public void setProperty(String propertyName, String propertyValue) { + if (propertyName.equalsIgnoreCase(PROP_NORMALISE)) { + // We do not allow the user to alter this property away from false, + // otherwise we would need to alter the functionality of the + // local calculations, etc. + return; + } + super.setProperty(propertyName, propertyValue); + } + + @Override + protected Object clone() throws CloneNotSupportedException { + return super.clone(); + } + /** * @throws Exception if the observation variables are not linearly independent * (leading to a non-positive definite covariance matrix). diff --git a/java/source/infodynamics/measures/continuous/kraskov/ConditionalMutualInfoCalculatorMultiVariateKraskov.java b/java/source/infodynamics/measures/continuous/kraskov/ConditionalMutualInfoCalculatorMultiVariateKraskov.java index 99f48e9..7a1644e 100755 --- a/java/source/infodynamics/measures/continuous/kraskov/ConditionalMutualInfoCalculatorMultiVariateKraskov.java +++ b/java/source/infodynamics/measures/continuous/kraskov/ConditionalMutualInfoCalculatorMultiVariateKraskov.java @@ -46,7 +46,8 @@ import infodynamics.utils.UnivariateNearestNeighbourSearcher; * actually implement the two KSG algorithms.

    * *

    Crucially, the calculation is performed by examining - * neighbours in the full joint space (as specified by Frenzel and Pompe) + * neighbours in the full joint space (as specified by Frenzel and Pompe, + * and later by Vejmelka and Paluš, and Vlachos and Kugiumtzis) * rather than two MI calculators.

    * *

    Usage is as per the paradigm outlined for {@link ConditionalMutualInfoCalculatorMultiVariate}, @@ -69,6 +70,12 @@ import infodynamics.utils.UnivariateNearestNeighbourSearcher; *

  • Frenzel and Pompe, * "Partial Mutual Information for Coupling Analysis of Multivariate Time Series", * Physical Review Letters, 99, p. 204101+ (2007).
  • + *
  • Vejmelka and Paluš, + * "Inferring the directionality of coupling with conditional mutual information", + * Physical Review E, 77, 026214, (2008)
  • + *
  • I. Vlachos and D. Kugiumtzis, + * "Nonuniform state-space reconstruction and coupling detection", + * Physical Review E, 82, 016207, (2010)
  • *
  • Kraskov, A., Stoegbauer, H., Grassberger, P., * "Estimating mutual information", * Physical Review E 69, (2004) 066138.
  • @@ -331,6 +338,7 @@ public abstract class ConditionalMutualInfoCalculatorMultiVariateKraskov var2Observations[r][c] += random.nextGaussian()*noiseLevel; } + // This next loop will only execute if dimensionsCond > 0 for (int c = 0; c < dimensionsCond; c++) { condObservations[r][c] += random.nextGaussian()*noiseLevel; @@ -351,7 +359,7 @@ public abstract class ConditionalMutualInfoCalculatorMultiVariateKraskov public double computeAverageLocalOfObservations() throws Exception { // Compute the conditional MI double startTime = Calendar.getInstance().getTimeInMillis(); - lastAverage = computeFromObservations(false)[0]; + lastAverage = computeFromObservations(false, null)[0]; condMiComputed = true; if (debug) { Calendar rightNow2 = Calendar.getInstance(); @@ -395,7 +403,7 @@ public abstract class ConditionalMutualInfoCalculatorMultiVariateKraskov var2Observations = MatrixUtils.extractSelectedTimePointsReusingArrays(originalData, reordering); } // Compute the conditional MI - double newCondMI = computeFromObservations(false)[0]; + double newCondMI = computeFromObservations(false, null)[0]; // restore original data kdTreeJoint = originalKdTreeJoint; if (variableToReorder == 1) { @@ -412,7 +420,7 @@ public abstract class ConditionalMutualInfoCalculatorMultiVariateKraskov @Override public double[] computeLocalOfPreviousObservations() throws Exception { - double[] localValues = computeFromObservations(true); + double[] localValues = computeFromObservations(true, null); lastAverage = MatrixUtils.mean(localValues); condMiComputed = true; return localValues; @@ -424,10 +432,25 @@ public abstract class ConditionalMutualInfoCalculatorMultiVariateKraskov @Override public double[] computeLocalUsingPreviousObservations(double[][] states1, double[][] states2, double[][] condStates) throws Exception { - // If implemented, will need to incorporate any normalisation here - // (normalising the incoming data the same way the previously - // supplied observations were normalised). - throw new Exception("Local method not implemented yet"); + // Do normalisation of the incoming data if required: + double[][] states1ToUse, states2ToUse, condStatesToUse; + if (normalise) { + states1ToUse = MatrixUtils.normaliseIntoNewArray(states1, var1Means, var1Stds); + states2ToUse = MatrixUtils.normaliseIntoNewArray(states2, var2Means, var2Stds); + if (dimensionsCond != 0) { + condStatesToUse = MatrixUtils.normaliseIntoNewArray(condStates, condMeans, condStds); + } else { + condStatesToUse = null; + } + } else { + states1ToUse = states1; + states2ToUse = states2; + condStatesToUse = condStates; + } + // And call the algorithm: + double[] localValues = computeFromObservations(true, + new double[][][]{states1ToUse, states2ToUse, condStatesToUse}); + return localValues; } /** @@ -444,10 +467,12 @@ public abstract class ConditionalMutualInfoCalculatorMultiVariateKraskov * * @param returnLocals whether to return an array or local values, or else * sums of these values + * @param newObservations set to null for computing for the observation set for the PDF, or pass in a new set + * of observations to compute the average/locals for (using the existing observations to construct the PDF) * @return either the average conditional MI, or array of local conditional MI value, in nats not bits * @throws Exception */ - protected double[] computeFromObservations(boolean returnLocals) throws Exception { + protected double[] computeFromObservations(boolean returnLocals, double[][][] newObservations) throws Exception { int N = var1Observations.length; // number of observations double[] returnValues = null; @@ -488,28 +513,36 @@ public abstract class ConditionalMutualInfoCalculatorMultiVariateKraskov uniNNSearcherVar2 = new UnivariateNearestNeighbourSearcher(var2Observations); } } - if (nnSearcherConditional == null) { + if ((nnSearcherConditional == null) && (dimensionsCond > 0)) { nnSearcherConditional = NearestNeighbourSearcher.create(condObservations); nnSearcherConditional.setNormType(normType); } + // How many time points are we averaging over? + int numTimePointsToComputeFor = (newObservations == null) ? + N : newObservations[0].length; + if (numThreads == 1) { // Single-threaded implementation: - returnValues = partialComputeFromObservations(0, N, returnLocals); - + if (newObservations == null) { + returnValues = partialComputeFromObservations(0, numTimePointsToComputeFor, returnLocals); + } else { + returnValues = partialComputeFromNewObservations(0, numTimePointsToComputeFor, + newObservations[0], newObservations[1], newObservations[2], returnLocals); + } } else { // We're going multithreaded: if (returnLocals) { // We're computing local conditional MI - returnValues = new double[N]; + returnValues = new double[numTimePointsToComputeFor]; } else { // We're computing average conditional MI returnValues = new double[CondMiKraskovThreadRunner.RETURN_ARRAY_LENGTH]; } // Distribute the observations to the threads for the parallel processing - int lTimesteps = N / numThreads; // each thread gets the same amount of data - int res = N % numThreads; // the first thread gets the residual data + int lTimesteps = numTimePointsToComputeFor / numThreads; // each thread gets the same amount of data + int res = numTimePointsToComputeFor % numThreads; // the first thread gets the residual data if (debug) { System.out.printf("Computing Kraskov conditional MI with %d threads (%d timesteps each, plus %d residual)\n", numThreads, lTimesteps, res); @@ -523,7 +556,7 @@ public abstract class ConditionalMutualInfoCalculatorMultiVariateKraskov System.out.println(t + ".Thread: from " + startTime + " to " + (startTime + numTimesteps)); // Trace Message } - runners[t] = new CondMiKraskovThreadRunner(this, startTime, numTimesteps, returnLocals); + runners[t] = new CondMiKraskovThreadRunner(this, startTime, numTimesteps, newObservations, returnLocals); tCalculators[t] = new Thread(runners[t]); tCalculators[t].start(); } @@ -552,10 +585,10 @@ public abstract class ConditionalMutualInfoCalculatorMultiVariateKraskov return returnValues; } else { // Average out the components for the final equation(s) and for debugging: - double averageDiGammas = returnValues[CondMiKraskovThreadRunner.INDEX_SUM_DIGAMMAS] / (double) N; - double avNxz = returnValues[CondMiKraskovThreadRunner.INDEX_SUM_NXZ] / (double) N; - double avNyz = returnValues[CondMiKraskovThreadRunner.INDEX_SUM_NYZ] / (double) N; - double avNz = returnValues[CondMiKraskovThreadRunner.INDEX_SUM_NZ] / (double) N; + double averageDiGammas = returnValues[CondMiKraskovThreadRunner.INDEX_SUM_DIGAMMAS] / (double) numTimePointsToComputeFor; + double avNxz = returnValues[CondMiKraskovThreadRunner.INDEX_SUM_NXZ] / (double) numTimePointsToComputeFor; + double avNyz = returnValues[CondMiKraskovThreadRunner.INDEX_SUM_NYZ] / (double) numTimePointsToComputeFor; + double avNz = returnValues[CondMiKraskovThreadRunner.INDEX_SUM_NZ] / (double) numTimePointsToComputeFor; if (debug) { System.out.printf("=%.3f, =%.3f, =%.3f\n", avNxz, avNyz, avNz); @@ -573,17 +606,27 @@ public abstract class ConditionalMutualInfoCalculatorMultiVariateKraskov // Algorithm 2: // We also retrieve the sums of inverses for debugging purposes: double averageInverseCountInJointXZ = - returnValues[CondMiKraskovThreadRunner.INDEX_SUM_INV_NXZ] / (double) N; + returnValues[CondMiKraskovThreadRunner.INDEX_SUM_INV_NXZ] / (double) numTimePointsToComputeFor; double averageInverseCountInJointYZ = - returnValues[CondMiKraskovThreadRunner.INDEX_SUM_INV_NYZ] / (double) N; - double averageMeasure = MathsUtils.digamma(k) - (2.0 / (double)k) + averageDiGammas + + returnValues[CondMiKraskovThreadRunner.INDEX_SUM_INV_NYZ] / (double) numTimePointsToComputeFor; + 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; + } + double averageMeasure = MathsUtils.digamma(k) - inverseKTerm + averageDiGammas + averageInverseCountInJointXZ + averageInverseCountInJointYZ; if (debug) { - System.out.printf("Av = digamma(k)=%.3f + =%.3f +=%.3f - 2/k=%.3f = %.3f" + + System.out.printf("Av = digamma(k)=%.3f + =%.3f +=%.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:

      + *
    1. 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
    2. + *
    3. for local conditional MIs (returnLocals == true), the array of local conditional MI values
    4. + *
    + * + * @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"); + } }