From 5f87e2018b6156a6d064e670190b2fe5f123617b Mon Sep 17 00:00:00 2001 From: Joseph Lizier Date: Mon, 9 Oct 2023 13:14:29 +1100 Subject: [PATCH] MI common estimators utilising underlying nearest neighbour searchers holding different data sets separately (for dynamic correlation exclusion) --- .../MutualInfoCalculatorMultiVariate.java | 33 ++++ .../MutualInfoMultiVariateCommon.java | 141 +++++++++++++++--- .../kraskov/MutualInfoMultiVariateTester.java | 74 +++++++++ 3 files changed, 224 insertions(+), 24 deletions(-) diff --git a/java/source/infodynamics/measures/continuous/MutualInfoCalculatorMultiVariate.java b/java/source/infodynamics/measures/continuous/MutualInfoCalculatorMultiVariate.java index 666b199..8818840 100755 --- a/java/source/infodynamics/measures/continuous/MutualInfoCalculatorMultiVariate.java +++ b/java/source/infodynamics/measures/continuous/MutualInfoCalculatorMultiVariate.java @@ -92,6 +92,25 @@ public interface MutualInfoCalculatorMultiVariate */ public static final String PROP_ADD_NOISE = "NOISE_LEVEL_TO_ADD"; + /** + *

As per {@link #addObservations(double[][], double[][])}; + * but also includes parameters to track which observation set + * the samples came from. Intended to only be used by other + * estimator classes here and not by users directly.

+ * + * @param source multivariate observations for variable 1 + * (first index is time or observation index, second is variable number) + * @param destination multivariate observations for variable 2 + * (first index is time or observation index, second is variable number) + * Length must match source, and their indices must correspond. + * @param observationSetIndexToUse which set of observations these came fmor + * @param startTimeIndex which was the first time index of these + * samples within that observation set. + * @throws Exception + */ + public void addObservationsTrackObservationIDs(double[][] source, double[][] destination, + int observationSetIndexToUse, int startTimeIndex) throws Exception; + /** * Compute the mutual information if the observations of the * first variable (source) @@ -130,4 +149,18 @@ public interface MutualInfoCalculatorMultiVariate * @throws Exception for invalid property values */ public String getProperty(String propertyName) throws Exception; + + /** + * Retrieve an array indicating which observation set each sample came from + * + * @return array of integers + */ + public int[] getObservationSetIndices(); + + /** + * Retrieve an array indicating which time index within its observation set that sample came from + * + * @return array of integers + */ + public int[] getObservationTimePoints(); } diff --git a/java/source/infodynamics/measures/continuous/MutualInfoMultiVariateCommon.java b/java/source/infodynamics/measures/continuous/MutualInfoMultiVariateCommon.java index ec45d02..3229b1d 100755 --- a/java/source/infodynamics/measures/continuous/MutualInfoMultiVariateCommon.java +++ b/java/source/infodynamics/measures/continuous/MutualInfoMultiVariateCommon.java @@ -23,6 +23,7 @@ import infodynamics.utils.EmpiricalMeasurementDistribution; import infodynamics.utils.MatrixUtils; import infodynamics.utils.RandomGenerator; +import java.util.Arrays; import java.util.Iterator; import java.util.Random; import java.util.Vector; @@ -86,6 +87,16 @@ public abstract class MutualInfoMultiVariateCommon implements */ protected double[] destStdsBeforeNorm; + /** + * Track which observation set each sample came from + */ + protected int[] observationSetIndices; + + /** + * Track which sample index within an observation set that each sample came from + */ + protected int[] observationTimePoints; + /** * Total number of observations supplied. * Only valid after {@link #finaliseAddObservations()} is called. @@ -127,6 +138,22 @@ public abstract class MutualInfoMultiVariateCommon implements * type calls */ protected Vector vectorOfDestinationObservations; + /** + * Tracks separate (time-series) observation sets + * we are taking samples from + */ + protected int observationSetIndex = 0; + /** + * Storage for which observation set each + * block of samples comes from + */ + protected Vector vectorOfObservationSetIndices; + /** + * Storage for start time point for the observation + * set within its block of samples + */ + protected Vector vectorOfObservationStartTimePoints; + /** * Whether the user has supplied more than one (disjoint) set of samples */ @@ -163,6 +190,11 @@ public abstract class MutualInfoMultiVariateCommon implements sourceStdsBeforeNorm = null; destMeansBeforeNorm = null; destStdsBeforeNorm = null; + observationSetIndices = null; + observationTimePoints = null; + observationSetIndex = 0; + vectorOfObservationSetIndices = null; + vectorOfObservationStartTimePoints = null; addedMoreThanOneObservationSet = false; } @@ -296,14 +328,54 @@ public abstract class MutualInfoMultiVariateCommon implements finaliseAddObservations(); } + @Override + public void setObservations(double[][] source, double[][] destination, + boolean[] sourceValid, boolean[] destValid) throws Exception { + startAddObservations(); + addObservations(source, destination, sourceValid, destValid); + finaliseAddObservations(); + } + + @Override + public void setObservations(double[] source, double[] destination, + boolean[] sourceValid, boolean[] destValid) throws Exception { + + if ((dimensionsDest != 1) || (dimensionsSource != 1)) { + throw new Exception("The number of source and dest dimensions (having been initialised to " + + dimensionsSource + " and " + dimensionsDest + ") can only be 1 when " + + "the univariate addObservations(double[],double[]) and " + + "setObservations(double[],double[]) methods are called"); + } + setObservations(MatrixUtils.reshape(source, source.length, 1), + MatrixUtils.reshape(destination, destination.length, 1), + sourceValid, destValid); + } + + @Override + public void setObservations(double[][] source, double[][] destination, + boolean[][] sourceValid, boolean[][] destValid) throws Exception { + + boolean[] allSourceValid = MatrixUtils.andRows(sourceValid); + boolean[] allDestValid = MatrixUtils.andRows(destValid); + setObservations(source, destination, allSourceValid, allDestValid); + } + @Override public void startAddObservations() { vectorOfSourceObservations = new Vector(); vectorOfDestinationObservations = new Vector(); + vectorOfObservationSetIndices = new Vector(); + vectorOfObservationStartTimePoints = new Vector(); } @Override public void addObservations(double[][] source, double[][] destination) throws Exception { + // Use the current observationSetIndex and increment for next use: + addObservationsTrackObservationIDs(source, destination, observationSetIndex++, 0); + } + + public void addObservationsTrackObservationIDs(double[][] source, double[][] destination, + int observationSetIndexToUse, int startTimeIndex) throws Exception { if (vectorOfSourceObservations == null) { // startAddObservations was not called first throw new RuntimeException("User did not call startAddObservations before addObservations"); @@ -326,6 +398,8 @@ public abstract class MutualInfoMultiVariateCommon implements } vectorOfSourceObservations.add(source); vectorOfDestinationObservations.add(destination); + vectorOfObservationSetIndices.add(observationSetIndexToUse); + vectorOfObservationStartTimePoints.add(startTimeIndex); } /** @@ -391,8 +465,14 @@ public abstract class MutualInfoMultiVariateCommon implements MatrixUtils.reshape(destination, destination.length, 1)); } + @Override public void addObservations(double[][] source, double[][] destination, int startTime, int numTimeSteps) throws Exception { + addObservations(source, destination, startTime, numTimeSteps, observationSetIndex++); + } + + protected void addObservations(double[][] source, double[][] destination, + int startTime, int numTimeSteps, int observationSetIndexToUse) throws Exception { if (vectorOfSourceObservations == null) { // startAddObservations was not called first throw new RuntimeException("User did not call startAddObservations before addObservations"); @@ -403,12 +483,12 @@ public abstract class MutualInfoMultiVariateCommon implements } double[][] sourceToAdd = new double[numTimeSteps][]; System.arraycopy(source, startTime, sourceToAdd, 0, numTimeSteps); - vectorOfSourceObservations.add(sourceToAdd); double[][] destToAdd = new double[numTimeSteps][]; System.arraycopy(destination, startTime, destToAdd, 0, numTimeSteps); - vectorOfDestinationObservations.add(destToAdd); + addObservationsTrackObservationIDs(sourceToAdd, destToAdd, observationSetIndexToUse, startTime); } + @Override public void addObservations(double[] source, double[] destination, int startTime, int numTimeSteps) throws Exception { @@ -423,8 +503,14 @@ public abstract class MutualInfoMultiVariateCommon implements startTime, numTimeSteps); } - - public void setObservations(double[][] source, double[][] destination, + public void addObservations(double[] source, double[] destination, + boolean[] sourceValid, boolean[] destValid) throws Exception { + addObservations(MatrixUtils.reshape(source, source.length, 1), + MatrixUtils.reshape(destination, destination.length, 1), + sourceValid, destValid); + } + + public void addObservations(double[][] source, double[][] destination, boolean[] sourceValid, boolean[] destValid) throws Exception { Vector startAndEndTimePairs = computeStartAndEndTimePairs(sourceValid, destValid); @@ -434,32 +520,18 @@ public abstract class MutualInfoMultiVariateCommon implements for (int[] timePair : startAndEndTimePairs) { int startTime = timePair[0]; int endTime = timePair[1]; - addObservations(source, destination, startTime, endTime - startTime + 1); + addObservations(source, destination, startTime, endTime - startTime + 1, observationSetIndex); } + observationSetIndex++; finaliseAddObservations(); } - @Override - public void setObservations(double[] source, double[] destination, - boolean[] sourceValid, boolean[] destValid) throws Exception { - - if ((dimensionsDest != 1) || (dimensionsSource != 1)) { - throw new Exception("The number of source and dest dimensions (having been initialised to " + - dimensionsSource + " and " + dimensionsDest + ") can only be 1 when " + - "the univariate addObservations(double[],double[]) and " + - "setObservations(double[],double[]) methods are called"); - } - setObservations(MatrixUtils.reshape(source, source.length, 1), - MatrixUtils.reshape(destination, destination.length, 1), - sourceValid, destValid); - } - - public void setObservations(double[][] source, double[][] destination, + public void addObservations(double[][] source, double[][] destination, boolean[][] sourceValid, boolean[][] destValid) throws Exception { - + boolean[] allSourceValid = MatrixUtils.andRows(sourceValid); boolean[] allDestValid = MatrixUtils.andRows(destValid); - setObservations(source, destination, allSourceValid, allDestValid); + addObservations(source, destination, allSourceValid, allDestValid); } /** @@ -482,11 +554,15 @@ public abstract class MutualInfoMultiVariateCommon implements } destObservations = new double[totalObservations][dimensionsDest]; sourceObservations = new double[totalObservations][dimensionsSource]; + observationSetIndices = new int[totalObservations]; + observationTimePoints = new int[totalObservations]; // Construct the joint vectors from the given observations // (removing redundant data which is outside any timeDiff) int startObservation = 0; Iterator iterator = vectorOfDestinationObservations.iterator(); + Iterator iteratorObsSetIndices = vectorOfObservationSetIndices.iterator(); + Iterator iteratorObsStartTimePoints = vectorOfObservationStartTimePoints.iterator(); for (double[][] source : vectorOfSourceObservations) { double[][] destination = iterator.next(); // Copy the data from these given observations into our master @@ -497,7 +573,14 @@ public abstract class MutualInfoMultiVariateCommon implements MatrixUtils.arrayCopy(destination, timeDiff, 0, destObservations, startObservation, 0, destination.length - timeDiff, dimensionsDest); - startObservation += destination.length - timeDiff; + int numNewObservations = destination.length - timeDiff; + // And update which observation set and time index each sample came from: + Arrays.fill(observationSetIndices, startObservation, startObservation + numNewObservations, iteratorObsSetIndices.next()); + int firstTimeSampleId = iteratorObsStartTimePoints.next(); // This is the first sample of the source; storing the time index for destination below. + for (int i = 0; i < numNewObservations; i++) { + observationTimePoints[startObservation + i] = firstTimeSampleId + timeDiff + i; + } + startObservation += numNewObservations; } if (vectorOfSourceObservations.size() > 1) { addedMoreThanOneObservationSet = true; @@ -834,4 +917,14 @@ public abstract class MutualInfoMultiVariateCommon implements } return startAndEndTimePairs; } + + @Override + public int[] getObservationSetIndices() { + return observationSetIndices; + } + + @Override + public int[] getObservationTimePoints() { + return observationTimePoints; + } } diff --git a/java/unittests/infodynamics/measures/continuous/kraskov/MutualInfoMultiVariateTester.java b/java/unittests/infodynamics/measures/continuous/kraskov/MutualInfoMultiVariateTester.java index bed843f..0ab5fae 100644 --- a/java/unittests/infodynamics/measures/continuous/kraskov/MutualInfoMultiVariateTester.java +++ b/java/unittests/infodynamics/measures/continuous/kraskov/MutualInfoMultiVariateTester.java @@ -21,6 +21,7 @@ package infodynamics.measures.continuous.kraskov; import infodynamics.utils.ArrayFileReader; import infodynamics.utils.MathsUtils; import infodynamics.utils.MatrixUtils; +import infodynamics.utils.RandomGenerator; public class MutualInfoMultiVariateTester extends infodynamics.measures.continuous.MutualInfoMultiVariateAbstractTester { @@ -652,5 +653,78 @@ public class MutualInfoMultiVariateTester assertEquals(expected_H_X_given_Y, conditionalEnt, 0.02); } + + /** + * Test that observationSetIndices and observationStartTimePoints are written properly + * + * @throws Exception + */ + public void testObservationSetIndices() throws Exception { + + int dimensions = 1; + int timeSteps = 100; + + MutualInfoCalculatorMultiVariateKraskov miCalc = getNewCalc(1); + miCalc.initialise(dimensions, dimensions); + + // generate some random data + RandomGenerator rg = new RandomGenerator(); + double[][] sourceData = rg.generateNormalData(timeSteps, dimensions, + 0, 1); + double[][] destData = rg.generateNormalData(timeSteps, dimensions, + 0, 1); + + // First check that for a simple single observation set everything works: + miCalc.setObservations(sourceData, destData); + + int[] observationSetIds = miCalc.getObservationSetIndices(); + int[] timeSeriesIndices = miCalc.getObservationTimePoints(); + assert(observationSetIds.length == timeSteps); + for (int t = 0; t < timeSteps; t++) { + assertEquals(observationSetIds[t], 0); + assertEquals(timeSeriesIndices[t], t); + } + + // Now add the same one twice: + miCalc.initialise(dimensions, dimensions); + miCalc.startAddObservations(); + miCalc.addObservations(sourceData, destData); + miCalc.addObservations(sourceData, destData); + miCalc.finaliseAddObservations(); + observationSetIds = miCalc.getObservationSetIndices(); + timeSeriesIndices = miCalc.getObservationTimePoints(); + assert(observationSetIds.length == 2*timeSteps); + for (int t = 0; t < timeSteps; t++) { + assertEquals(observationSetIds[t], 0); + assertEquals(timeSeriesIndices[t], t); + } + for (int t = 0; t < timeSteps; t++) { + assertEquals(observationSetIds[timeSteps + t], 1); + assertEquals(timeSeriesIndices[timeSteps + t], t); + } + + // Now add NUM_SEGMENTS randomly chosen segments: + int NUM_SEGMENTS = 10; + int maxLength = 10; + int[] startPoints = rg.generateRandomInts(NUM_SEGMENTS, timeSteps - maxLength); + int[] lengthsMinus1 = rg.generateRandomInts(NUM_SEGMENTS, maxLength - 1); // ensures we don't add segments of length 0 + miCalc.initialise(dimensions, dimensions); + miCalc.startAddObservations(); + for (int r = 0; r < NUM_SEGMENTS; r++) { + miCalc.addObservations(sourceData, destData, startPoints[r], lengthsMinus1[r]+1); + } + miCalc.finaliseAddObservations(); + observationSetIds = miCalc.getObservationSetIndices(); + timeSeriesIndices = miCalc.getObservationTimePoints(); + assert(observationSetIds.length == MatrixUtils.sum(lengthsMinus1) + NUM_SEGMENTS); + int t = 0; + for (int r = 0; r < NUM_SEGMENTS; r++) { + for (int i = 0; i < lengthsMinus1[r]+1; i++) { + assertEquals(observationSetIds[t], r); + assertEquals(timeSeriesIndices[t], startPoints[r] + i); + t++; + } + } + } }