Implementation and unit test for KSG MI calculation with new samples (both algorithms 1 and 2)

This commit is contained in:
Joseph Lizier 2020-10-30 16:21:10 +11:00
parent 0606f8ab39
commit 3222b5a629
5 changed files with 353 additions and 30 deletions

View File

@ -322,7 +322,7 @@ public abstract class MutualInfoCalculatorMultiVariateKraskov
public double computeAverageLocalOfObservations() throws Exception {
// Compute the MI
double startTime = Calendar.getInstance().getTimeInMillis();
lastAverage = computeFromObservations(false)[0];
lastAverage = computeFromObservations(false, null)[0];
miComputed = true;
if (debug) {
Calendar rightNow2 = Calendar.getInstance();
@ -355,7 +355,7 @@ public abstract class MutualInfoCalculatorMultiVariateKraskov
// Generate a new re-ordered data2
destObservations = MatrixUtils.extractSelectedTimePointsReusingArrays(originalData2, reordering);
// Compute the MI
double newMI = computeFromObservations(false)[0];
double newMI = computeFromObservations(false, null)[0];
// restore original variables:
destObservations = originalData2;
@ -384,21 +384,32 @@ public abstract class MutualInfoCalculatorMultiVariateKraskov
* @throws Exception
*/
public double[] computeLocalOfPreviousObservations() throws Exception {
double[] localValues = computeFromObservations(true);
double[] localValues = computeFromObservations(true, null);
lastAverage = MatrixUtils.mean(localValues);
miComputed = true;
return localValues;
}
/**
* This method, specified in {@link MutualInfoCalculatorMultiVariate}
* is not implemented yet here.
*/
@Override
public double[] computeLocalUsingPreviousObservations(double[][] states1, double[][] states2) throws Exception {
// TODO If implemented, will need to incorporate any time difference here.
// Will also need to handle normalisation of the incoming data
// appropriately
throw new Exception("Local method not implemented yet");
// Do normalisation of the incoming data if required:
double[][] states1ToUse, states2ToUse;
if (normalise) {
states1ToUse = MatrixUtils.normaliseIntoNewArray(states1, sourceMeansBeforeNorm, sourceStdsBeforeNorm, 0, states1.length-timeDiff);
states2ToUse = MatrixUtils.normaliseIntoNewArray(states2, destMeansBeforeNorm, destStdsBeforeNorm, timeDiff, states2.length-timeDiff);
} else {
if (timeDiff > 0) {
states1ToUse = MatrixUtils.selectRows(states1, 0, states1.length-timeDiff);
states2ToUse = MatrixUtils.selectRows(states2, 0, states2.length-timeDiff);
} else {
states1ToUse = states1;
states2ToUse = states2;
}
}
// And call the algorithm:
double[] localValues = computeFromObservations(true,
new double[][][]{states1ToUse, states2ToUse});
return localValues;
}
/**
@ -415,37 +426,52 @@ public abstract class MutualInfoCalculatorMultiVariateKraskov
*
* @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 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
protected double[] computeFromObservations(boolean returnLocals, double[][][] newObservations) throws Exception {
int N = sourceObservations.length; // number of observations for the PDFs
double[] returnValues = null;
if (useGPU) {
// How many time points are we averaging over?
int numTimePointsToComputeFor = (newObservations == null) ?
N : newObservations[0].length;
if (useGPU && (newObservations == null)) {
System.out.println("Cannot use GPU for estimation based on new observations -- falling back to CPU calculation...");
}
if (useGPU && (newObservations == null)) {
returnValues = gpuComputeFromObservations(0, N, returnLocals);
} else if (numThreads == 1) {
// Single-threaded implementation:
ensureKdTreesConstructed();
returnValues = partialComputeFromObservations(0, N, returnLocals);
if (newObservations == null) {
returnValues = partialComputeFromObservations(0, numTimePointsToComputeFor, returnLocals);
} else {
returnValues = partialComputeFromNewObservations(0, numTimePointsToComputeFor,
newObservations[0], newObservations[1], returnLocals);
}
} else {
// We're going multithreaded:
ensureKdTreesConstructed();
if (returnLocals) {
// We're computing local MI
returnValues = new double[N];
returnValues = new double[numTimePointsToComputeFor];
} else {
// We're computing average MI
returnValues = new double[MiKraskovThreadRunner.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 MI with %d threads (%d timesteps each, plus %d residual)\n",
numThreads, lTimesteps, res);
@ -459,7 +485,7 @@ public abstract class MutualInfoCalculatorMultiVariateKraskov
System.out.println(t + ".Thread: from " + startTime +
" to " + (startTime + numTimesteps)); // Trace Message
}
runners[t] = new MiKraskovThreadRunner(this, startTime, numTimesteps, returnLocals);
runners[t] = new MiKraskovThreadRunner(this, startTime, numTimesteps, newObservations, returnLocals);
tCalculators[t] = new Thread(runners[t]);
tCalculators[t].start();
}
@ -488,18 +514,21 @@ public abstract class MutualInfoCalculatorMultiVariateKraskov
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;
double averageDiGammas = returnValues[MiKraskovThreadRunner.INDEX_SUM_DIGAMMAS] / (double) numTimePointsToComputeFor;
double avNx = returnValues[MiKraskovThreadRunner.INDEX_SUM_NX] / (double) numTimePointsToComputeFor;
double avNy = returnValues[MiKraskovThreadRunner.INDEX_SUM_NY] / (double) numTimePointsToComputeFor;
if (debug) {
System.out.println(String.format("Average n_x=%.3f, Average n_y=%.3f", avNx, avNy));
}
// Use digamma(N) normally, unless we're looking at new observations:
double digammaNToUse = (newObservations == null) ? digammaN : MathsUtils.digamma(totalObservations+1);
// Finalise the average result, depending on which algorithm we are implementing:
if (isAlgorithm1) {
return new double[] { digammaK - averageDiGammas + digammaN };
return new double[] { digammaK - averageDiGammas + digammaNToUse };
} else {
return new double[] { digammaK - (1.0 / (double)k) - averageDiGammas + digammaN };
return new double[] { digammaK - (1.0 / (double)k) - averageDiGammas + digammaNToUse };
}
}
}
@ -528,6 +557,37 @@ public abstract class MutualInfoCalculatorMultiVariateKraskov
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)
*
* <p>The method returns:<ol>
* <li>for average MIs (returnLocals == false), the relevant sums of digamma(n_x+1), digamma(n_y+1)
* for a partial set of the observations</li>
* <li>for local MIs (returnLocals == true), the array of local MI values</li>
* </ol>
*
* @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 2
* @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[] partialComputeFromNewObservations(
int startTimePoint, int numTimePoints,
double[][] newVar1Observations, double[][] newVar2Observations,
boolean returnLocals) throws Exception;
/**
* Protected method to be used internally for GPU implementations.
* This method serves the same purpose as partialComputeFromObservations,
@ -791,6 +851,7 @@ public abstract class MutualInfoCalculatorMultiVariateKraskov
protected MutualInfoCalculatorMultiVariateKraskov miCalc;
protected int myStartTimePoint;
protected int numberOfTimePoints;
protected double[][][] newObservations;
protected boolean computeLocals;
protected double[] returnValues = null;
@ -804,11 +865,13 @@ public abstract class MutualInfoCalculatorMultiVariateKraskov
public MiKraskovThreadRunner(
MutualInfoCalculatorMultiVariateKraskov miCalc,
int myStartTimePoint, int numberOfTimePoints,
double[][][] newObservations,
boolean computeLocals) {
this.miCalc = miCalc;
this.myStartTimePoint = myStartTimePoint;
this.numberOfTimePoints = numberOfTimePoints;
this.computeLocals = computeLocals;
this.newObservations = newObservations;
}
/**
@ -831,8 +894,17 @@ public abstract class MutualInfoCalculatorMultiVariateKraskov
*/
public void run() {
try {
returnValues = miCalc.partialComputeFromObservations(
myStartTimePoint, numberOfTimePoints, computeLocals);
if (newObservations == null) {
// Computing on existing observations
returnValues = miCalc.partialComputeFromObservations(
myStartTimePoint, numberOfTimePoints, computeLocals);
} else {
// Computing on new observations
returnValues = miCalc.partialComputeFromNewObservations(
myStartTimePoint, numberOfTimePoints,
newObservations[0], newObservations[1],
computeLocals);
}
} catch (Exception e) {
// Store the exception for later retrieval
problem = e;

View File

@ -23,9 +23,7 @@ import java.util.PriorityQueue;
import infodynamics.measures.continuous.MutualInfoCalculatorMultiVariate;
import infodynamics.utils.MathsUtils;
import infodynamics.utils.MatrixUtils;
import infodynamics.utils.NeighbourNodeData;
import infodynamics.utils.EuclideanUtils;
/**
* <p>Computes the differential mutual information of two given multivariate sets of
@ -176,6 +174,72 @@ public class MutualInfoCalculatorMultiVariateKraskov1
return neighbourCounts;
}
@Override
protected double[] partialComputeFromNewObservations(int startTimePoint, int numTimePoints,
double[][] newVar1Observations, double[][] newVar2Observations, boolean returnLocals) throws Exception {
double startTime = Calendar.getInstance().getTimeInMillis();
double[] localMi = null;
if (returnLocals) {
localMi = new double[numTimePoints];
}
// 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 by
// finding the kth closest neighbour for point t:
PriorityQueue<NeighbourNodeData> nnPQ =
kdTreeJoint.findKNearestNeighbours(k,
new double[][] {newVar1Observations[t], newVar2Observations[t]});
// First element in the PQ is the kth NN,
// and epsilon = kthNnData.distance
NeighbourNodeData kthNnData = nnPQ.poll();
// Count the number of points whose x distance is less
// than eps, and whose y distance is less than
// epsilon = kthNnData.distance
int n_x = nnSearcherSource.countPointsWithinR(
new double[][] {newVar1Observations[t]},
kthNnData.distance, false);
int n_y = nnSearcherDest.countPointsWithinR(
new double[][] {newVar2Observations[t]},
kthNnData.distance, false);
sumNx += n_x;
sumNy += n_y;
// And take the digammas:
double digammaNxPlusOne = MathsUtils.digamma(n_x+1);
double digammaNyPlusOne = MathsUtils.digamma(n_y+1);
sumDiGammas += digammaNxPlusOne + digammaNyPlusOne;
if (returnLocals) {
// For new observations we're taking the probability counts over an extra point (no self-exclusion)
// so we don't use digamma(N) but digamma(N+1)
localMi[t-startTimePoint] = digammaK - digammaNxPlusOne - digammaNyPlusOne + MathsUtils.digamma(totalObservations+1);
}
}
if (debug) {
Calendar rightNow2 = Calendar.getInstance();
long endTime = rightNow2.getTimeInMillis();
System.out.println("Subset " + startTimePoint + ":" +
(startTimePoint + numTimePoints) + " Calculation time: " +
((endTime - startTime)/1000.0) + " sec" );
}
// Select what to return:
if (returnLocals) {
return localMi;
} else {
return new double[] {sumDiGammas, sumNx, sumNy};
}
}
}

View File

@ -130,4 +130,84 @@ public class MutualInfoCalculatorMultiVariateKraskov2
}
}
@Override
protected double[] partialComputeFromNewObservations(int startTimePoint, int numTimePoints,
double[][] newVar1Observations, double[][] newVar2Observations, boolean returnLocals) throws Exception {
double startTime = Calendar.getInstance().getTimeInMillis();
double[] localMi = null;
if (returnLocals) {
localMi = new double[numTimePoints];
}
// Constants:
double invK = 1.0 / (double)k;
// 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 by
// finding the kth closest neighbours for point t:
PriorityQueue<NeighbourNodeData> nnPQ =
kdTreeJoint.findKNearestNeighbours(k,
new double[][] {newVar1Observations[t], newVar2Observations[t]});
// Find eps_{x,y} as the maximum x and y norms amongst this set:
double eps_x = 0.0;
double eps_y = 0.0;
for (int j = 0; j < k; j++) {
// Take the furthest remaining of the nearest neighbours from the PQ:
NeighbourNodeData nnData = nnPQ.poll();
if (nnData.norms[0] > eps_x) {
eps_x = nnData.norms[0];
}
if (nnData.norms[1] > eps_y) {
eps_y = nnData.norms[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 = nnSearcherSource.countPointsWithinR(
new double[][] {newVar1Observations[t]},
eps_x, true);
int n_y = nnSearcherDest.countPointsWithinR(
new double[][] {newVar2Observations[t]},
eps_y, true);
sumNx += n_x;
sumNy += n_y;
// And take the digammas:
double digammaNx = MathsUtils.digamma(n_x);
double digammaNy = MathsUtils.digamma(n_y);
sumDiGammas += digammaNx + digammaNy;
if (returnLocals) {
// For new observations we're taking the probability counts over an extra point (no self-exclusion)
// so we don't use digamma(N) but digamma(N+1)
localMi[t-startTimePoint] = digammaK - invK - digammaNx - digammaNy + MathsUtils.digamma(totalObservations+1);
}
}
if (debug) {
Calendar rightNow2 = Calendar.getInstance();
long endTime = rightNow2.getTimeInMillis();
System.out.println("Subset " + startTimePoint + ":" +
(startTimePoint + numTimePoints) + " Calculation time: " +
((endTime - startTime)/1000.0) + " sec" );
}
// Select what to return:
if (returnLocals) {
return localMi;
} else {
return new double[] {sumDiGammas, sumNx, sumNy};
}
}
}

View File

@ -2554,8 +2554,17 @@ public class MatrixUtils {
* @param matrix 2D matrix of doubles
*/
public static double[][] normaliseIntoNewArray(double[][] matrix, double[] means, double[] stds) {
return normaliseIntoNewArray(matrix, means, stds, 0, matrix.length);
}
/**
* Normalises the elements along each column of the matrix
*
* @param matrix 2D matrix of doubles
*/
public static double[][] normaliseIntoNewArray(double[][] matrix, double[] means, double[] stds, int startRow, int rows) {
double[][] newMatrix = new double[matrix.length][matrix[0].length];
for (int r = 0; r < newMatrix.length; r++) {
for (int r = startRow; r < startRow + rows; r++) {
for (int c = 0; c < newMatrix[r].length; c++) {
newMatrix[r][c] = matrix[r][c] - means[c];
if (!Double.isInfinite(1.0 / stds[c])) {

View File

@ -19,6 +19,7 @@
package infodynamics.measures.continuous.kraskov;
import infodynamics.utils.ArrayFileReader;
import infodynamics.utils.MathsUtils;
import infodynamics.utils.MatrixUtils;
public class MutualInfoMultiVariateTester
@ -462,5 +463,102 @@ public class MutualInfoMultiVariateTester
}
/**
* Unit test for MI on new observations.
* We can test this against the calculator itself. If we send in the original data set as new observations,
* we can recreate the neighbour counts (plus one) by setting K to 1 larger (to account for the data point itself), and
* account for the change in bias.
*
* @throws Exception
*/
public void testMultivariateCondMIForNewObservations() throws Exception {
ArrayFileReader afr = new ArrayFileReader("demos/data/4ColsPairedOneStepNoisyDependence-1.txt");
double[][] data = afr.getDouble2DMatrix();
// RandomGenerator rg = new RandomGenerator();
//double[][] data = rg.generateNormalData(50, 4, 0, 1);
// Use various Kraskov k nearest neighbours parameter
int[] kNNs = {4, 10, 15};
System.out.println("Kraskov MI testing new Observations:");
for (int alg = 1; alg < 3; alg++) {
for (int ki = 0; ki < kNNs.length; ki++) {
MutualInfoCalculatorMultiVariateKraskov miCalc = getNewCalc(alg);
MutualInfoCalculatorMultiVariateKraskov miCalcForNew = getNewCalc(alg);
// Let it normalise by default
// And no noise addition to protect the integrity of our neighbour counts under both techniques here:
miCalc.setProperty(MutualInfoCalculatorMultiVariateKraskov.PROP_ADD_NOISE, "0");
miCalcForNew.setProperty(MutualInfoCalculatorMultiVariateKraskov.PROP_ADD_NOISE, "0");
double[][] var1 = MatrixUtils.selectColumns(data, new int[] {0});
double[][] var2 = MatrixUtils.selectColumns(data, new int[] {1});
// Compute MI(0;1|2,3) :
miCalc.setProperty(
MutualInfoCalculatorMultiVariateKraskov.PROP_K,
Integer.toString(kNNs[ki]));
System.out.println("Main calc normalisation is " + miCalc.getProperty(MutualInfoCalculatorMultiVariateKraskov.PROP_NORMALISE));
miCalc.initialise(var1[0].length, var2[0].length);
miCalc.setObservations(var1, var2);
@SuppressWarnings("unused")
double miAverage = miCalc.computeAverageLocalOfObservations();
// Now compute as new observations:
miCalcForNew.setProperty(
MutualInfoCalculatorMultiVariateKraskov.PROP_K,
Integer.toString(kNNs[ki] + 1)); // Using K = K + 1
// condMiCalcForNew.setProperty(
// MutualInfoCalculatorMultiVariateKraskov.PROP_NUM_THREADS,
// "1");
System.out.println("New obs calc normalisation is " + miCalcForNew.getProperty(MutualInfoCalculatorMultiVariateKraskov.PROP_NORMALISE));
miCalcForNew.initialise(var1[0].length, var2[0].length);
miCalcForNew.setObservations(var1, var2);
// condMiCalc.setDebug(true);
//condMiCalcForNew.setDebug(true);
double[] newLocals = miCalcForNew.computeLocalUsingPreviousObservations(var1, var2);
//condMiCalcForNew.setDebug(false);
@SuppressWarnings("unused")
double averageFromNewObservations = MatrixUtils.mean(newLocals);
// We can't check this directly, so test each point individually:
for (int t = 0; t < data.length; t++) {
double[] originalNeighbourCounts = miCalc.partialComputeFromObservations(t, 1, false);
// Need to normalise the data before passing it in here -- this is
// what is happening inside computeLocalUsingPreviousObservations above
double[] newObsNeighbourCounts = miCalcForNew.partialComputeFromNewObservations(
t, 1,
MatrixUtils.normaliseIntoNewArray(var1),
MatrixUtils.normaliseIntoNewArray(var2), false);
// Now check each return count in the array:
if (originalNeighbourCounts[1] != newObsNeighbourCounts[1] - 1) {
System.out.println("Assertion failure for t=" + t + ": expected " + originalNeighbourCounts[1] +
" from original, plus 1, but got " + newObsNeighbourCounts[1]);
System.out.print("Actual raw data was: ");
MatrixUtils.printArray(System.out, data[0]);
}
assertEquals(originalNeighbourCounts[1], newObsNeighbourCounts[1] - 1); // Nx should be 1 higher
assertEquals(originalNeighbourCounts[2], newObsNeighbourCounts[2] - 1); // Ny should be 1 higher
// Now check the local value at each point using these verified counts:
double newLocalValue;
if (alg == 1) {
newLocalValue = miCalcForNew.digammaK -
MathsUtils.digamma((int) newObsNeighbourCounts[1] + 1) -
MathsUtils.digamma((int) newObsNeighbourCounts[2] + 1) +
MathsUtils.digamma(miCalcForNew.getNumObservations() + 1); // correct digammaN for new samples
} else {
newLocalValue = miCalcForNew.digammaK -
(double) 1 / (double) miCalcForNew.k -
MathsUtils.digamma((int) newObsNeighbourCounts[1]) -
MathsUtils.digamma((int) newObsNeighbourCounts[2]) +
MathsUtils.digamma(miCalcForNew.getNumObservations() + 1); // correct digammaN for new samples
}
if (Math.abs(newLocalValue - newLocals[t]) > 0.00000001) {
System.out.printf("t=%d: Assertion failed: computed local was %.5f, local from nn counts was %.5f\n",
t, newLocals[t], newLocalValue);
}
assertEquals(newLocalValue, newLocals[t], 0.00000001);
}
}
}
}
}