diff --git a/java/source/infodynamics/measures/continuous/kraskov/MutualInfoCalculatorMultiVariateKraskov.java b/java/source/infodynamics/measures/continuous/kraskov/MutualInfoCalculatorMultiVariateKraskov.java
index 5195e05..7384583 100755
--- a/java/source/infodynamics/measures/continuous/kraskov/MutualInfoCalculatorMultiVariateKraskov.java
+++ b/java/source/infodynamics/measures/continuous/kraskov/MutualInfoCalculatorMultiVariateKraskov.java
@@ -23,6 +23,7 @@ import java.util.Random;
import infodynamics.measures.continuous.MutualInfoCalculatorMultiVariate;
import infodynamics.measures.continuous.MutualInfoMultiVariateCommon;
import infodynamics.utils.EuclideanUtils;
+import infodynamics.utils.MathsUtils;
import infodynamics.utils.MatrixUtils;
/**
@@ -57,6 +58,7 @@ import infodynamics.utils.MatrixUtils;
*
* @author Joseph Lizier (email,
* www)
+ * @author Ipek Özdemir
*/
public abstract class MutualInfoCalculatorMultiVariateKraskov
extends MutualInfoMultiVariateCommon
@@ -71,24 +73,6 @@ public abstract class MutualInfoCalculatorMultiVariateKraskov
* Calculator for the norm between data points
*/
protected EuclideanUtils normCalculator;
- /**
- * Cache for the norms between x (source) points
- */
- protected double[][] xNorms;
- /**
- * Cache for the norms between x (dest) points
- */
- protected double[][] yNorms;
- /**
- * Whether we cache the norms each time (making reordering very quick).
- * (Should only be set to false for testing)
- */
- public static boolean tryKeepAllPairsNorms = true;
- /**
- * An upper limit on the number of samples for which
- * we will cache the norms between data points.
- */
- public static int MAX_DATA_SIZE_FOR_KEEP_ALL_PAIRS_NORM = 2000;
/**
* Property name for the number of K nearest neighbours used in
@@ -112,6 +96,16 @@ public abstract class MutualInfoCalculatorMultiVariateKraskov
* added to the data (default is 0).
*/
public static final String PROP_ADD_NOISE = "NOISE_LEVEL_TO_ADD";
+ /**
+ * Property name for the number of parallel threads to use in the
+ * computation
+ */
+ public static final String PROP_NUM_THREADS = "NUM_THREADS";
+ /**
+ * Valid property value for {@link #PROP_NUM_THREADS} to indicate
+ * that all available processors should be used.
+ */
+ public static final String USE_ALL_THREADS = "USE_ALL";
/**
* Whether to normalise the incoming data
@@ -125,7 +119,15 @@ public abstract class MutualInfoCalculatorMultiVariateKraskov
* Amount of random Gaussian noise to add to the incoming data
*/
protected double noiseLevel = 0.0;
-
+ /**
+ * Number of parallel threads to use in the computation
+ */
+ protected int numThreads = 1;
+ /**
+ * Private variable to record which algorithm this instance is implementing
+ */
+ protected boolean isAlgorithm1 = false;
+
/**
* Construct an instance of the KSG MI calculator
*/
@@ -134,12 +136,6 @@ public abstract class MutualInfoCalculatorMultiVariateKraskov
normCalculator = new EuclideanUtils(EuclideanUtils.NORM_MAX_NORM);
}
- public void initialise(int sourceDimensions, int destDimensions) {
- super.initialise(sourceDimensions, destDimensions);
- xNorms = null;
- yNorms = null;
- }
-
/**
* Sets properties for the KSG MI calculator.
* New property values are not guaranteed to take effect until the next call
@@ -163,6 +159,10 @@ public abstract class MutualInfoCalculatorMultiVariateKraskov
* so can be considered as a number of standard deviations of the data.
* (Recommended by Kraskov. MILCA uses 1e-8; but adds in
* a random amount of noise in [0,noiseLevel) ). Default 0.
+ *
{@link #PROP_NUM_THREADS} -- the integer number of parallel threads
+ * to use in the computation. Can be passed as a string "USE_ALL"
+ * to use all available processors on the machine.
+ * Default is 1 for single-threaded.
* any valid properties for {@link MutualInfoMultiVariateCommon#setProperty(String, String)}.
*
*
@@ -183,6 +183,12 @@ public abstract class MutualInfoCalculatorMultiVariateKraskov
} else if (propertyName.equalsIgnoreCase(PROP_ADD_NOISE)) {
addNoise = true;
noiseLevel = Double.parseDouble(propertyValue);
+ } else if (propertyName.equalsIgnoreCase(PROP_NUM_THREADS)) {
+ if (propertyValue.equalsIgnoreCase(USE_ALL_THREADS)) {
+ numThreads = Runtime.getRuntime().availableProcessors();
+ } else { // otherwise the user has passed in an integer:
+ numThreads = Integer.parseInt(propertyValue);
+ }
} else {
// No property was set here
propertySet = false;
@@ -228,37 +234,42 @@ public abstract class MutualInfoCalculatorMultiVariateKraskov
}
/**
- * Utility function to compute the norms between each pair of points in each marginal time series
- *
- */
- protected void computeNorms() {
- int N = sourceObservations.length; // number of observations
-
- xNorms = new double[N][N];
- yNorms = new double[N][N];
- for (int t = 0; t < N; t++) {
- // Compute the norms from t to all other time points
- double[][] xyNormsForT = normCalculator.computeNorms(sourceObservations, destObservations, t);
- for (int t2 = 0; t2 < N; t2++) {
- xNorms[t][t2] = xyNormsForT[t2][0];
- yNorms[t][t2] = xyNormsForT[t2][1];
- }
- }
- }
-
- /**
- * Compute the average MI from the previously supplied observations.
+ * {@inheritDoc}
*
* @return the average MI in nats (not bits!)
*/
- public abstract double computeAverageLocalOfObservations() throws Exception;
+ public double computeAverageLocalOfObservations() throws Exception {
+ // Compute the MI
+ lastAverage = computeFromObservations(false)[0];
+ miComputed = true;
+ return lastAverage;
+ }
/**
+ * {@inheritDoc}
+ *
* @return the MI under the new ordering, in nats (not bits!).
* Returns NaN if any of the determinants are zero
* (because this will make the denominator of the log 0).
*/
- public abstract double computeAverageLocalOfObservations(int[] reordering) throws Exception;
+ public double computeAverageLocalOfObservations(int[] reordering) throws Exception {
+ double[][] originalData2 = destObservations;
+ if (reordering != null) {
+ // Generate a new re-ordered data2
+ destObservations = MatrixUtils.extractSelectedTimePointsReusingArrays(originalData2, reordering);
+ }
+ // Compute the MI
+ double newMI = computeFromObservations(false)[0];
+ // restore data2
+ destObservations = originalData2;
+ if (reordering == null) {
+ // Only keep this average value if it was for
+ // the original data:
+ miComputed = true;
+ lastAverage = newMI;
+ }
+ return newMI;
+ }
/**
* Computes the local values of the MI,
@@ -278,7 +289,12 @@ public abstract class MutualInfoCalculatorMultiVariateKraskov
* @return the "time-series" of local MIs in bits
* @throws Exception
*/
- public abstract double[] computeLocalOfPreviousObservations() throws Exception;
+ public double[] computeLocalOfPreviousObservations() throws Exception {
+ double[] localValues = computeFromObservations(true);
+ lastAverage = MatrixUtils.mean(localValues);
+ miComputed = true;
+ return localValues;
+ }
/**
* This method, specified in {@link MutualInfoCalculatorMultiVariate}
@@ -291,6 +307,190 @@ public abstract class MutualInfoCalculatorMultiVariateKraskov
throw new Exception("Local method not implemented yet");
}
+ /**
+ * This protected method handles the multiple threads which
+ * computes either the average or local MI (over parts of the total
+ * observations), computing the x and y
+ * distances between all tuples in time.
+ *
+ *
The method returns:
+ * - for (returnLocals == false), an array of size 1,
+ * containing the average MI
+ * - for local MIs (returnLocals == true), the array of local MI values
+ *
+ *
+ * @param returnLocals whether to return an array or local values, or else
+ * sums of these values
+ * @return either the average MI, or array of local MI value, in nats not bits
+ * @throws Exception
+ */
+ protected double[] computeFromObservations(boolean returnLocals) throws Exception {
+ int N = sourceObservations.length; // number of observations
+
+ double[] returnValues = null;
+
+ if (numThreads == 1) {
+ // Single-threaded implementation:
+ returnValues = partialComputeFromObservations(0, N, returnLocals);
+
+ } else {
+ // We're going multithreaded:
+ if (returnLocals) {
+ // We're computing local MI
+ returnValues = new double[N];
+ } else {
+ // We're computing average MI
+ returnValues = new double[3];
+ }
+
+ // 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
+ if (debug) {
+ System.out.printf("Computing Kraskov MI alg1 with %d threads (%d timesteps each, plus %d residual)\n",
+ numThreads, lTimesteps, res);
+ }
+ Thread[] tCalculators = new Thread[numThreads];
+ MiKraskovThreadRunner[] runners = new MiKraskovThreadRunner[numThreads];
+ for (int t = 0; t < numThreads; t++) {
+ int startTime = (t == 0) ? 0 : lTimesteps * t + res;
+ int numTimesteps = (t == 0) ? lTimesteps + res : lTimesteps;
+ if (debug) {
+ System.out.println(t + ".Thread: from " + startTime +
+ " to " + (startTime + numTimesteps)); // Trace Message
+ }
+ runners[t] = new MiKraskovThreadRunner(this, startTime, numTimesteps, returnLocals);
+ tCalculators[t] = new Thread(runners[t]);
+ tCalculators[t].start();
+ }
+
+ // Here, we should wait for the termination of the all threads
+ // and collect their results
+ for (int t = 0; t < numThreads; t++) {
+ if (tCalculators[t] != null) { // TODO Ipek: can you comment on why we're checking for null here?
+ tCalculators[t].join();
+ }
+ // Now we add in the data from this completed thread:
+ if (returnLocals) {
+ // We're computing local MI; copy these local values
+ // into the full array of locals
+ System.arraycopy(runners[t].getReturnValues(), 0,
+ returnValues, runners[t].myStartTimePoint, runners[t].numberOfTimePoints);
+ } else {
+ // We're computing the average MI, keep the running sums of digammas and counts
+ MatrixUtils.addInPlace(returnValues, runners[t].getReturnValues());
+ }
+ }
+ }
+
+ // Finalise the results:
+ if (returnLocals) {
+ return returnValues;
+ } else {
+ // Compute the average number of points within eps_x and eps_y
+ double averageDiGammas = returnValues[MiKraskovThreadRunner.INDEX_SUM_DIGAMMAS] / (double) N;
+ double avNx = returnValues[MiKraskovThreadRunner.INDEX_SUM_NX] / (double) N;
+ double avNy = returnValues[MiKraskovThreadRunner.INDEX_SUM_NY] / (double) N;
+ if (debug) {
+ System.out.println(String.format("Average n_x=%.3f, Average n_y=%.3f", avNx, avNy));
+ }
+
+ // Finalise the average result, depending on which algorithm we are implementing:
+ if (isAlgorithm1) {
+ return new double[] { MathsUtils.digamma(k) - averageDiGammas + MathsUtils.digamma(N)};
+ } else {
+ return new double[] { MathsUtils.digamma(k) - (1.0 / (double)k) - averageDiGammas + MathsUtils.digamma(N)};
+ }
+ }
+ }
+
+ /**
+ * 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.
+ *
+ * The method returns:
+ * - for average MIs (returnLocals == false), the relevant sums of digamma(n_x+1) and digamma(n_y+1)
+ * for a partial set of the observations
+ * - for local MIs (returnLocals == true), the array of local MI values
+ *
+ *
+ * @param startTimePoint start time for the partial set we examine
+ * @param numTimePoints number of time points (including startTimePoint to examine)
+ * @param returnLocals whether to return an array or local values, or else
+ * sums of these values
+ * @return an array of sum of digamma(n_x+1) and digamma(n_y+1), then
+ * sum of n_x and finally sum of n_y (these latter two are for debugging purposes).
+ * @throws Exception
+ */
+ protected abstract double[] partialComputeFromObservations(
+ int startTimePoint, int numTimePoints, boolean returnLocals) throws Exception;
+
+ /**
+ * Private class to handle multi-threading of the Kraskov algorithms.
+ * Each instance calls partialComputeFromObservations()
+ * to compute nearest neighbours for a part of the data.
+ *
+ *
+ * @author Joseph Lizier (email,
+ * www)
+ * @author Ipek Özdemir
+ */
+ private class MiKraskovThreadRunner implements Runnable {
+ protected MutualInfoCalculatorMultiVariateKraskov miCalc;
+ protected int myStartTimePoint;
+ protected int numberOfTimePoints;
+ protected boolean computeLocals;
+
+ protected double[] returnValues = null;
+ protected Exception problem = null;
+
+ public static final int INDEX_SUM_DIGAMMAS = 0;
+ public static final int INDEX_SUM_NX = 1;
+ public static final int INDEX_SUM_NY = 2;
+
+ public MiKraskovThreadRunner(
+ MutualInfoCalculatorMultiVariateKraskov miCalc,
+ int myStartTimePoint, int numberOfTimePoints,
+ boolean computeLocals) {
+ this.miCalc = miCalc;
+ this.myStartTimePoint = myStartTimePoint;
+ this.numberOfTimePoints = numberOfTimePoints;
+ this.computeLocals = computeLocals;
+ }
+
+ /**
+ * Return the values from this part of the data,
+ * or throw any exception that was encountered by the
+ * thread.
+ *
+ * @return an exception previously encountered by this thread.
+ * @throws Exception
+ */
+ public double[] getReturnValues() throws Exception {
+ if (problem != null) {
+ throw problem;
+ }
+ return returnValues;
+ }
+
+ /**
+ * Start the thread for the given parameters
+ */
+ public void run() {
+ try {
+ returnValues = miCalc.partialComputeFromObservations(myStartTimePoint, numberOfTimePoints, computeLocals);
+ } catch (Exception e) {
+ // Store the exception for later retrieval
+ problem = e;
+ return;
+ }
+ }
+ }
+ // end class MiKraskovThreadRunner
+
/**
* Utility function used for debugging, printing digamma constants
*
diff --git a/java/source/infodynamics/measures/continuous/kraskov/MutualInfoCalculatorMultiVariateKraskov1.java b/java/source/infodynamics/measures/continuous/kraskov/MutualInfoCalculatorMultiVariateKraskov1.java
index 10f9a7f..2a68bfd 100755
--- a/java/source/infodynamics/measures/continuous/kraskov/MutualInfoCalculatorMultiVariateKraskov1.java
+++ b/java/source/infodynamics/measures/continuous/kraskov/MutualInfoCalculatorMultiVariateKraskov1.java
@@ -43,6 +43,7 @@ import infodynamics.utils.MatrixUtils;
*
* @author Joseph Lizier (email,
* www)
+ * @author Ipek Özdemir
*/
public class MutualInfoCalculatorMultiVariateKraskov1
extends MutualInfoCalculatorMultiVariateKraskov {
@@ -53,188 +54,36 @@ public class MutualInfoCalculatorMultiVariateKraskov1
*/
protected static final double CUTOFF_MULTIPLIER = 1.5;
- @Override
- public double computeAverageLocalOfObservations() throws Exception {
- return computeAverageLocalOfObservations(null);
+ public MutualInfoCalculatorMultiVariateKraskov1() {
+ super();
+ isAlgorithm1 = true;
}
@Override
- public double computeAverageLocalOfObservations(int[] reordering) throws Exception {
- if (!tryKeepAllPairsNorms || (sourceObservations.length > MAX_DATA_SIZE_FOR_KEEP_ALL_PAIRS_NORM)) {
- double[][] originalData2 = destObservations;
- if (reordering != null) {
- // Generate a new re-ordered data2
- destObservations = MatrixUtils.extractSelectedTimePointsReusingArrays(originalData2, reordering);
- }
- // Compute the MI
- double newMI = computeAverageLocalOfObservationsWhileComputingDistances();
- // restore data2
- destObservations = originalData2;
- return newMI;
- }
+ protected double[] partialComputeFromObservations(
+ int startTimePoint, int numTimePoints, boolean returnLocals) throws Exception {
- if (xNorms == null) {
- computeNorms();
- }
int N = sourceObservations.length; // number of observations
int cutoffForKthMinLinear = (int) (CUTOFF_MULTIPLIER * Math.log(N) / Math.log(2.0));
- // Count the average number of points within eps_x and eps_y
- double averageDiGammas = 0;
- double avNx = 0;
- double avNy = 0;
-
- for (int t = 0; t < N; t++) {
- // Compute eps for this time step:
- // using x and y norms to all neighbours
- // (note that norm of point t to itself will be set to infinity).
-
- int tForY = (reordering == null) ? t : reordering[t];
-
- double[] jointNorm = new double[N];
- for (int t2 = 0; t2 < N; t2++) {
- int t2ForY = (reordering == null) ? t2 : reordering[t2];
- jointNorm[t2] = Math.max(xNorms[t][t2], yNorms[tForY][t2ForY]);
- }
- // Then find the kth closest neighbour, using a heuristic to
- // select whether to keep the k mins only or to do a sort.
- double epsilon = 0.0;
- if (k <= cutoffForKthMinLinear) {
- // just do a linear search for the minimum
- epsilon = MatrixUtils.kthMin(jointNorm, k);
- } else {
- // Sort the array of joint norms first
- java.util.Arrays.sort(jointNorm);
- // And find the distance to it's kth closest neighbour
- // (we subtract one since the array is indexed from zero)
- epsilon = jointNorm[k-1];
- }
-
- // Count the number of points whose x distance is less
- // than eps, and whose y distance is less than eps
- int n_x = 0;
- int n_y = 0;
- for (int t2 = 0; t2 < N; t2++) {
- if (xNorms[t][t2] < epsilon) {
- n_x++;
- }
- int t2ForY = (reordering == null) ? t2 : reordering[t2];
- if (yNorms[tForY][t2ForY] < epsilon) {
- n_y++;
- }
- }
- avNx += n_x;
- avNy += n_y;
- // And take the digamma before adding into the
- // average:
- averageDiGammas += MathsUtils.digamma(n_x+1) + MathsUtils.digamma(n_y+1);
+ double[] localMi = null;
+ if (returnLocals) {
+ localMi = new double[numTimePoints];
}
- averageDiGammas /= (double) N;
- if (debug) {
- avNx /= (double)N;
- avNy /= (double)N;
- System.out.println(String.format("Average n_x=%.3f, Average n_y=%.3f", avNx, avNy));
- }
-
- double average = MathsUtils.digamma(k) - averageDiGammas + MathsUtils.digamma(N);
- miComputed = true;
- if (reordering == null) {
- lastAverage = average;
- }
- return average;
- }
-
- /**
- * This method correctly computes the average MI, but recomputes the x and y
- * distances between all tuples in time.
- * Kept here for cases where we have too many observations
- * to keep the norm between all pairs, and for testing purposes.
- *
- * @see #computeAverageLocalOfObservations()
- * @return average MI value in nats not bits
- * @throws Exception
- */
- public double computeAverageLocalOfObservationsWhileComputingDistances() throws Exception {
- int N = sourceObservations.length; // number of observations
- int cutoffForKthMinLinear = (int) (CUTOFF_MULTIPLIER * Math.log(N) / Math.log(2.0));
-
- // Count the average number of points within eps_x and eps_y
- double averageDiGammas = 0;
- double avNx = 0;
- double avNy = 0;
-
- for (int t = 0; t < N; t++) {
- // Compute eps for this time step:
- // First get x and y norms to all neighbours
- // (note that norm of point t to itself will be set to infinity).
- double[][] xyNorms = normCalculator.computeNorms(sourceObservations, destObservations, t);
- double[] jointNorm = new double[N];
- for (int t2 = 0; t2 < N; t2++) {
- jointNorm[t2] = Math.max(xyNorms[t2][0], xyNorms[t2][1]);
- }
- // Then find the kth closest neighbour, using a heuristic to
- // select whether to keep the k mins only or to do a sort.
- double epsilon = 0.0;
- if (k <= cutoffForKthMinLinear) {
- // just do a linear search for the minimum
- epsilon = MatrixUtils.kthMin(jointNorm, k);
- } else {
- // Sort the array of joint norms first
- java.util.Arrays.sort(jointNorm);
- // And find the distance to it's kth closest neighbour
- // (we subtract one since the array is indexed from zero)
- epsilon = jointNorm[k-1];
- }
-
- // Count the number of points whose x distance is less
- // than eps, and whose y distance is less than eps
- int n_x = 0;
- int n_y = 0;
- for (int t2 = 0; t2 < N; t2++) {
- if (xyNorms[t2][0] < epsilon) {
- n_x++;
- }
- if (xyNorms[t2][1] < epsilon) {
- n_y++;
- }
- }
- avNx += n_x;
- avNy += n_y;
- // And take the digamma before adding into the
- // average:
- averageDiGammas += MathsUtils.digamma(n_x+1) + MathsUtils.digamma(n_y+1);
- }
- averageDiGammas /= (double) N;
- if (debug) {
- avNx /= (double)N;
- avNy /= (double)N;
- System.out.println(String.format("Average n_x=%.3f, Average n_y=%.3f", avNx, avNy));
- }
-
- lastAverage = MathsUtils.digamma(k) - averageDiGammas + MathsUtils.digamma(N);
- miComputed = true;
- return lastAverage;
- }
-
- @Override
- public double[] computeLocalOfPreviousObservations() throws Exception {
- int N = sourceObservations.length; // number of observations
- int cutoffForKthMinLinear = (int) (CUTOFF_MULTIPLIER * Math.log(N) / Math.log(2.0));
- double[] localMi = new double[N];
// Constants:
double digammaK = MathsUtils.digamma(k);
double digammaN = MathsUtils.digamma(N);
- // Count the average number of points within eps_x and eps_y
- double averageDiGammas = 0;
- double avNx = 0;
- double avNy = 0;
-
- for (int t = 0; t < N; t++) {
+ // Count the average number of points within eps_x and eps_y of each point
+ double sumDiGammas = 0;
+ double sumNx = 0;
+ double sumNy = 0;
+
+ for (int t = startTimePoint; t < startTimePoint + numTimePoints; t++) {
// Compute eps for this time step:
// First get x and y norms to all neighbours
- // (note that norm of point t to itself will be set to infinity.
+ // (note that norm of point t to itself will be set to infinity).
double[][] xyNorms = normCalculator.computeNorms(sourceObservations, destObservations, t);
double[] jointNorm = new double[N];
for (int t2 = 0; t2 < N; t2++) {
@@ -249,7 +98,7 @@ public class MutualInfoCalculatorMultiVariateKraskov1
} else {
// Sort the array of joint norms first
java.util.Arrays.sort(jointNorm);
- // And find the distance to it's kth closest neighbour
+ // And find the distance to its kth closest neighbour
// (we subtract one since the array is indexed from zero)
epsilon = jointNorm[k-1];
}
@@ -266,29 +115,26 @@ public class MutualInfoCalculatorMultiVariateKraskov1
n_y++;
}
}
- // And take the digamma:
+ sumNx += n_x;
+ sumNy += n_y;
+ // And take the digammas:
double digammaNxPlusOne = MathsUtils.digamma(n_x+1);
double digammaNyPlusOne = MathsUtils.digamma(n_y+1);
-
- localMi[t] = digammaK - digammaNxPlusOne - digammaNyPlusOne + digammaN;
-
- avNx += n_x;
- avNy += n_y;
- // And keep track of the average
- averageDiGammas += digammaNxPlusOne + digammaNyPlusOne;
- }
- averageDiGammas /= (double) N;
- if (debug) {
- avNx /= (double)N;
- avNy /= (double)N;
- System.out.println(String.format("Average n_x=%.3f, Average n_y=%.3f", avNx, avNy));
+ sumDiGammas += digammaNxPlusOne + digammaNyPlusOne;
+
+ if (returnLocals) {
+ localMi[t-startTimePoint] = digammaK - digammaNxPlusOne - digammaNyPlusOne + digammaN;
+ }
}
- lastAverage = digammaK - averageDiGammas + digammaN;
- miComputed = true;
- return localMi;
+ // Select what to return:
+ if (returnLocals) {
+ return localMi;
+ } else {
+ return new double[] {sumDiGammas, sumNx, sumNy};
+ }
}
-
+
@Override
public String printConstants(int N) throws Exception {
String constants = String.format("digamma(k=%d)=%.3e + digamma(N=%d)=%.3e => %.3e",
diff --git a/java/source/infodynamics/measures/continuous/kraskov/MutualInfoCalculatorMultiVariateKraskov2.java b/java/source/infodynamics/measures/continuous/kraskov/MutualInfoCalculatorMultiVariateKraskov2.java
index 3a7ca6b..4f6ae2a 100755
--- a/java/source/infodynamics/measures/continuous/kraskov/MutualInfoCalculatorMultiVariateKraskov2.java
+++ b/java/source/infodynamics/measures/continuous/kraskov/MutualInfoCalculatorMultiVariateKraskov2.java
@@ -44,6 +44,7 @@ import infodynamics.utils.MatrixUtils;
*
* @author Joseph Lizier (email,
* www)
+ * @author Ipek Özdemir
*/
public class MutualInfoCalculatorMultiVariateKraskov2
extends MutualInfoCalculatorMultiVariateKraskov {
@@ -57,309 +58,33 @@ public class MutualInfoCalculatorMultiVariateKraskov2
*/
protected static final double CUTOFF_MULTIPLIER = 1.5;
- public double computeAverageLocalOfObservations(int[] reordering) throws Exception {
- if (!tryKeepAllPairsNorms || (sourceObservations.length > MAX_DATA_SIZE_FOR_KEEP_ALL_PAIRS_NORM)) {
- double[][] originalData2 = destObservations;
- // Generate a new re-ordered data2
- destObservations = MatrixUtils.extractSelectedTimePointsReusingArrays(originalData2, reordering);
- // Compute the MI
- double newMI = computeAverageLocalOfObservationsWhileComputingDistances();
- // restore data2
- destObservations = originalData2;
- return newMI;
- }
-
- // Otherwise we will use the norms we've already computed, and use a "virtual"
- // reordered data2.
-
- if (xNorms == null) {
- computeNorms();
- }
- int N = sourceObservations.length; // number of observations
- int cutoffForKthMinLinear = (int) (CUTOFF_MULTIPLIER * Math.log(N) / Math.log(2.0));
-
- // Count the average number of points within eps_x and eps_y
- double averageDiGammas = 0;
- double avNx = 0;
- double avNy = 0;
-
- for (int t = 0; t < N; t++) {
- // Compute eps_x and eps_y for this time step:
- // First get x and y norms to all neighbours
- // (note that norm of point t to itself will be set to infinity).
-
- int tForY = reordering[t];
-
- double[][] jointNorm = new double[N][2];
- for (int t2 = 0; t2 < N; t2++) {
- int t2ForY = reordering[t2];
- jointNorm[t2][JOINT_NORM_VAL_COLUMN] = Math.max(xNorms[t][t2], yNorms[tForY][t2ForY]);
- // And store the time step for back reference after the
- // array is sorted.
- jointNorm[t2][JOINT_NORM_TIMESTEP_COLUMN] = t2;
- }
- // Then find the k closest neighbours:
- double eps_x = 0.0;
- double eps_y = 0.0;
- int[] timeStepsOfKthMins = null;
- if (k <= cutoffForKthMinLinear) {
- // just do a linear search for the minimum epsilon value
- timeStepsOfKthMins = MatrixUtils.kMinIndices(jointNorm, JOINT_NORM_VAL_COLUMN, k);
- } else {
- // Sort the array of joint norms
- java.util.Arrays.sort(jointNorm, FirstIndexComparatorDouble.getInstance());
- // and now we have the closest k points.
- timeStepsOfKthMins = new int[k];
- for (int j = 0; j < k; j++) {
- timeStepsOfKthMins[j] = (int) jointNorm[j][JOINT_NORM_TIMESTEP_COLUMN];
- }
- }
- // and now we have the closest k points.
- // Find eps_{x,y} as the maximum x and y norms amongst this set:
- for (int j = 0; j < k; j++) {
- int timeStepOfJthPoint = timeStepsOfKthMins[j];
- if (xNorms[t][timeStepOfJthPoint] > eps_x) {
- eps_x = xNorms[t][timeStepOfJthPoint];
- }
- if (yNorms[tForY][reordering[timeStepOfJthPoint]] > eps_y) {
- eps_y = yNorms[tForY][reordering[timeStepOfJthPoint]];
- }
- }
-
- // Count the number of points whose x distance is less
- // than or equal to eps_x, and whose y distance is less
- // than or equal to eps_y
- int n_x = 0;
- int n_y = 0;
- for (int t2 = 0; t2 < N; t2++) {
- if (xNorms[t][t2] <= eps_x) {
- n_x++;
- }
- if (yNorms[tForY][reordering[t2]] <= eps_y) {
- n_y++;
- }
- }
- avNx += n_x;
- avNy += n_y;
- // And take the digamma before adding into the
- // average:
- averageDiGammas += MathsUtils.digamma(n_x) + MathsUtils.digamma(n_y);
- }
- averageDiGammas /= (double) N;
- if (debug) {
- avNx /= (double)N;
- avNy /= (double)N;
- System.out.println(String.format("Average n_x=%.3f, Average n_y=%.3f", avNx, avNy));
- }
-
- double average = MathsUtils.digamma(k) - 1.0/(double)k - averageDiGammas + MathsUtils.digamma(N);
- miComputed = true;
- if (reordering == null) {
- lastAverage = average;
- }
- return average;
+ public MutualInfoCalculatorMultiVariateKraskov2() {
+ super();
+ isAlgorithm1 = false;
}
- public double computeAverageLocalOfObservations() throws Exception {
- if (!tryKeepAllPairsNorms || (sourceObservations.length > MAX_DATA_SIZE_FOR_KEEP_ALL_PAIRS_NORM)) {
- return computeAverageLocalOfObservationsWhileComputingDistances();
- }
+ protected double[] partialComputeFromObservations(
+ int startTimePoint, int numTimePoints, boolean returnLocals) throws Exception {
- if (xNorms == null) {
- computeNorms();
- }
int N = sourceObservations.length; // number of observations
int cutoffForKthMinLinear = (int) (CUTOFF_MULTIPLIER * Math.log(N) / Math.log(2.0));
-
- // Count the average number of points within eps_x and eps_y
- double averageDiGammas = 0;
- double avNx = 0;
- double avNy = 0;
- for (int t = 0; t < N; t++) {
- // Compute eps_x and eps_y for this time step:
- // using x and y norms to all neighbours
- // (note that norm of point t to itself will be set to infinity).
-
- double[][] jointNorm = new double[N][2];
- for (int t2 = 0; t2 < N; t2++) {
- jointNorm[t2][JOINT_NORM_VAL_COLUMN] = Math.max(xNorms[t][t2], yNorms[t][t2]);
- // And store the time step for back reference after the
- // array is sorted.
- jointNorm[t2][JOINT_NORM_TIMESTEP_COLUMN] = t2;
- }
- // Then find the k closest neighbours:
- double eps_x = 0.0;
- double eps_y = 0.0;
- int[] timeStepsOfKthMins = null;
- if (k <= cutoffForKthMinLinear) {
- // just do a linear search for the minimum epsilon value
- timeStepsOfKthMins = MatrixUtils.kMinIndices(jointNorm, JOINT_NORM_VAL_COLUMN, k);
- } else {
- // Sort the array of joint norms
- java.util.Arrays.sort(jointNorm, FirstIndexComparatorDouble.getInstance());
- // and now we have the closest k points.
- timeStepsOfKthMins = new int[k];
- for (int j = 0; j < k; j++) {
- timeStepsOfKthMins[j] = (int) jointNorm[j][JOINT_NORM_TIMESTEP_COLUMN];
- }
- }
- // and now we have the closest k points.
- // Find eps_{x,y} as the maximum x and y norms amongst this set:
- for (int j = 0; j < k; j++) {
- int timeStepOfJthPoint = timeStepsOfKthMins[j];
- if (xNorms[t][timeStepOfJthPoint] > eps_x) {
- eps_x = xNorms[t][timeStepOfJthPoint];
- }
- if (yNorms[t][timeStepOfJthPoint] > eps_y) {
- eps_y = yNorms[t][timeStepOfJthPoint];
- }
- }
-
- // Count the number of points whose x distance is less
- // than or equal to eps_x, and whose y distance is less
- // than or equal to eps_y
- int n_x = 0;
- int n_y = 0;
- for (int t2 = 0; t2 < N; t2++) {
- if (xNorms[t][t2] <= eps_x) {
- n_x++;
- }
- if (yNorms[t][t2] <= eps_y) {
- n_y++;
- }
- }
- avNx += n_x;
- avNy += n_y;
- // And take the digamma before adding into the
- // average:
- averageDiGammas += MathsUtils.digamma(n_x) + MathsUtils.digamma(n_y);
- // if (debug) {
- // System.out.printf("n=%d, \n");
- // }
+ double[] localMi = null;
+ if (returnLocals) {
+ localMi = new double[numTimePoints];
}
- averageDiGammas /= (double) N;
- lastAverage = MathsUtils.digamma(k) - 1.0/(double)k - averageDiGammas + MathsUtils.digamma(N);
- miComputed = true;
- if (debug) {
- avNx /= (double)N;
- avNy /= (double)N;
- System.out.printf("Average n_x=%.3f, Average n_y=%.3f", avNx, avNy);
- System.out.printf("psi(k=%d)=%.4f - 1/k=%.4f - averageDiGammas=%.4f -psi(N)=%.4f => %.4f\n",
- k, MathsUtils.digamma(k), 1.0/(double)k, averageDiGammas, MathsUtils.digamma(N), lastAverage);
- }
-
- return lastAverage;
- }
-
- /**
- * This method correctly computes the average local MI, but recomputes the x and y
- * distances between all tuples in time.
- * Kept here for cases where we have too many observations
- * to keep the norm between all pairs, and for testing purposes.
- *
- * @see #computeAverageLocalOfObservations()
- * @return
- * @throws Exception
- */
- public double computeAverageLocalOfObservationsWhileComputingDistances() throws Exception {
- int N = sourceObservations.length; // number of observations
- int cutoffForKthMinLinear = (int) (CUTOFF_MULTIPLIER * Math.log(N) / Math.log(2.0));
-
- // Count the average number of points within eps_x and eps_y
- double averageDiGammas = 0;
- double avNx = 0;
- double avNy = 0;
-
- for (int t = 0; t < N; t++) {
- // Compute eps_x and eps_y for this time step:
- // First get x and y norms to all neighbours
- // (note that norm of point t to itself will be set to infinity).
- double[][] xyNorms = normCalculator.computeNorms(sourceObservations, destObservations, t);
- double[][] jointNorm = new double[N][2];
- for (int t2 = 0; t2 < N; t2++) {
- jointNorm[t2][JOINT_NORM_VAL_COLUMN] = Math.max(xyNorms[t2][0], xyNorms[t2][1]);
- // And store the time step for back reference after the
- // array is sorted.
- jointNorm[t2][JOINT_NORM_TIMESTEP_COLUMN] = t2;
- }
- // Then find the k closest neighbours:
- double eps_x = 0.0;
- double eps_y = 0.0;
- int[] timeStepsOfKthMins = null;
- if (k <= cutoffForKthMinLinear) {
- // just do a linear search for the minimum epsilon value
- timeStepsOfKthMins = MatrixUtils.kMinIndices(jointNorm, JOINT_NORM_VAL_COLUMN, k);
- } else {
- // Sort the array of joint norms
- java.util.Arrays.sort(jointNorm, FirstIndexComparatorDouble.getInstance());
- // and now we have the closest k points.
- timeStepsOfKthMins = new int[k];
- for (int j = 0; j < k; j++) {
- timeStepsOfKthMins[j] = (int) jointNorm[j][JOINT_NORM_TIMESTEP_COLUMN];
- }
- }
- // and now we have the closest k points.
- // Find eps_{x,y} as the maximum x and y norms amongst this set:
- for (int j = 0; j < k; j++) {
- int timeStepOfJthPoint = timeStepsOfKthMins[j];
- if (xyNorms[timeStepOfJthPoint][0] > eps_x) {
- eps_x = xyNorms[timeStepOfJthPoint][0];
- }
- if (xyNorms[timeStepOfJthPoint][1] > eps_y) {
- eps_y = xyNorms[timeStepOfJthPoint][1];
- }
- }
-
- // Count the number of points whose x distance is less
- // than or equal to eps_x, and whose y distance is less
- // than or equal to eps_y
- int n_x = 0;
- int n_y = 0;
- for (int t2 = 0; t2 < N; t2++) {
- if (xyNorms[t2][0] <= eps_x) {
- n_x++;
- }
- if (xyNorms[t2][1] <= eps_y) {
- n_y++;
- }
- }
- avNx += n_x;
- avNy += n_y;
- // And take the digamma before adding into the
- // average:
- averageDiGammas += MathsUtils.digamma(n_x) + MathsUtils.digamma(n_y);
- }
- averageDiGammas /= (double) N;
- lastAverage = MathsUtils.digamma(k) - 1.0/(double)k - averageDiGammas + MathsUtils.digamma(N);
- miComputed = true;
- if (debug) {
- avNx /= (double)N;
- avNy /= (double)N;
- System.out.printf("Average n_x=%.3f, Average n_y=%.3f\n", avNx, avNy);
- System.out.printf("psi(k=%d)=%.4f - 1/k=%.4f - averageDiGammas=%.4f + psi(N)=%.4f => %.4f\n",
- k, MathsUtils.digamma(k), 1.0/(double)k, averageDiGammas, MathsUtils.digamma(N), lastAverage);
- }
-
- return lastAverage;
- }
-
- public double[] computeLocalOfPreviousObservations() throws Exception {
- int N = sourceObservations.length; // number of observations
- int cutoffForKthMinLinear = (int) (CUTOFF_MULTIPLIER * Math.log(N) / Math.log(2.0));
- double[] localMi = new double[N];
// Constants:
double digammaK = MathsUtils.digamma(k);
double invK = 1.0 / (double)k;
double digammaN = MathsUtils.digamma(N);
- // Count the average number of points within eps_x and eps_y
- double averageDiGammas = 0;
- double avNx = 0;
- double avNy = 0;
-
- for (int t = 0; t < N; t++) {
+ // Count the average number of points within eps_x and eps_y of each point
+ double sumDiGammas = 0;
+ double sumNx = 0;
+ double sumNy = 0;
+
+ for (int t = startTimePoint; t < startTimePoint + numTimePoints; t++) {
// Compute eps_x and eps_y for this time step:
// First get x and y norms to all neighbours
// (note that norm of point t to itself will be set to infinity).
@@ -371,7 +96,8 @@ public class MutualInfoCalculatorMultiVariateKraskov2
// array is sorted.
jointNorm[t2][JOINT_NORM_TIMESTEP_COLUMN] = t2;
}
- // Then find the k closest neighbours:
+ // Then find the k closest neighbours, using a heuristic to
+ // select whether to keep the k mins only or to do a sort.
double eps_x = 0.0;
double eps_y = 0.0;
int[] timeStepsOfKthMins = null;
@@ -412,26 +138,23 @@ public class MutualInfoCalculatorMultiVariateKraskov2
n_y++;
}
}
- avNx += n_x;
- avNy += n_y;
- // And take the digamma:
+ sumNx += n_x;
+ sumNy += n_y;
+ // And take the digammas:
double digammaNx = MathsUtils.digamma(n_x);
double digammaNy = MathsUtils.digamma(n_y);
-
- localMi[t] = digammaK - invK - digammaNx - digammaNy + digammaN;
+ sumDiGammas += digammaNx + digammaNy;
- averageDiGammas += digammaNx + digammaNy;
+ if (returnLocals) {
+ localMi[t-startTimePoint] = digammaK - invK - digammaNx - digammaNy + digammaN;
+ }
}
- averageDiGammas /= (double) N;
- if (debug) {
- avNx /= (double)N;
- avNy /= (double)N;
- System.out.println(String.format("Average n_x=%.3f, Average n_y=%.3f", avNx, avNy));
+ // Select what to return:
+ if (returnLocals) {
+ return localMi;
+ } else {
+ return new double[] {sumDiGammas, sumNx, sumNy};
}
-
- lastAverage = digammaK - invK - averageDiGammas + digammaN;
- miComputed = true;
- return localMi;
}
public String printConstants(int N) throws Exception {
diff --git a/java/source/infodynamics/utils/MathsUtils.java b/java/source/infodynamics/utils/MathsUtils.java
index 0000234..b3a7eec 100755
--- a/java/source/infodynamics/utils/MathsUtils.java
+++ b/java/source/infodynamics/utils/MathsUtils.java
@@ -32,8 +32,8 @@ public class MathsUtils {
private static final double EULER_MASCHERONI_CONSTANT = 0.5772156;
private static int highestDigammaArgCalced = 0;
- private static final int NUM_STORED_DIGAMMAS = 10000;
- private static double[] storedDigammas;
+ private static final int NUM_STORED_DIGAMMAS = 10000; // commons.math to handle beyond this
+ private static double[] storedDigammas = new double[0];
/**
* Returns the integer result of base^power
@@ -272,31 +272,42 @@ public class MathsUtils {
if (d < 1) {
return Double.NaN;
}
- if (storedDigammas == null) {
- // allocate space to store our results
- storedDigammas = new double[NUM_STORED_DIGAMMAS];
- storedDigammas[0] = Double.NaN;
- storedDigammas[1] = -EULER_MASCHERONI_CONSTANT;
- highestDigammaArgCalced = 1;
+ if (storedDigammas.length == 0) {
+ synchronized(storedDigammas) { // Ensure no race condition here
+ // We do two checks on whether the storage has been
+ // created so that the first is very fast (without
+ // requiring synchronization), and the second
+ // ensures no race condition.
+ if (storedDigammas.length == 0) {
+ // Using length == 0 as proxy to null, since
+ // we can't synchronize on a null object
+
+ // allocate space to store our results
+ storedDigammas = new double[NUM_STORED_DIGAMMAS];
+ storedDigammas[0] = Double.NaN;
+ storedDigammas[1] = -EULER_MASCHERONI_CONSTANT;
+ highestDigammaArgCalced = 1;
+ }
+ }
}
if (d <= highestDigammaArgCalced) {
// We've already calculated this one
return storedDigammas[d];
}
- // else need to calculate it
+ if (d >= NUM_STORED_DIGAMMAS) {
+ // Don't bother updating our storage,
+ // directly use commons.math:
+ return Gamma.digamma(d);
+ }
+ // Else we'll calculate it and update the storage:
double result = storedDigammas[highestDigammaArgCalced];
for (int n = highestDigammaArgCalced + 1; n <= d; n++) {
result += 1.0 / (double) (n-1);
- if (n < NUM_STORED_DIGAMMAS) {
- storedDigammas[n] = result;
- }
+ // n must be < NUM_STORED_DIGAMMAS by earlier if statement on d
+ storedDigammas[n] = result;
}
- if (d < NUM_STORED_DIGAMMAS) {
- highestDigammaArgCalced = d;
- } else {
- highestDigammaArgCalced = NUM_STORED_DIGAMMAS - 1;
- }
- return result;
+ highestDigammaArgCalced = d;
+ return result;
}
/**