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:

      + *
    1. for (returnLocals == false), an array of size 1, + * containing the average MI
    2. + *
    3. for local MIs (returnLocals == true), the array of local MI values
    4. + *
    + * + * @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:

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