From c8d17e9522c45e83c64609a9520708f943718b27 Mon Sep 17 00:00:00 2001 From: David Shorten Date: Mon, 23 Aug 2021 17:50:34 +1000 Subject: [PATCH] functionality allowing specific embedding intervals --- ...erEntropyCalculatorSpikingIntegration.java | 275 ++++++++++-------- tester.py | 41 ++- 2 files changed, 178 insertions(+), 138 deletions(-) diff --git a/java/source/infodynamics/measures/spiking/integration/TransferEntropyCalculatorSpikingIntegration.java b/java/source/infodynamics/measures/spiking/integration/TransferEntropyCalculatorSpikingIntegration.java index 89b6a38..6397961 100644 --- a/java/source/infodynamics/measures/spiking/integration/TransferEntropyCalculatorSpikingIntegration.java +++ b/java/source/infodynamics/measures/spiking/integration/TransferEntropyCalculatorSpikingIntegration.java @@ -7,6 +7,7 @@ import java.util.PriorityQueue; import java.util.Random; import java.util.Vector; import java.util.ArrayList; +import java.lang.Math; //import infodynamics.measures.continuous.kraskov.EuclideanUtils; import infodynamics.measures.spiking.TransferEntropyCalculatorSpiking; @@ -44,23 +45,28 @@ import infodynamics.utils.ParsedProperties; public class TransferEntropyCalculatorSpikingIntegration implements TransferEntropyCalculatorSpiking { /** - * Number of past destination interspike intervals to consider (akin to embedding length) + * The past destination interspike intervals to consider + * (and associated property name and convenience length variable) + * The code assumes that the first interval is numbered 1, the next is numbered 2, etc. + * It also assumes that the intervals are sorted. The setter method performs sorting to ensure this. */ - protected int k = 1; + public static final String DEST_PAST_INTERVALS_PROP_NAME = "DEST_PAST_INTERVALS"; + protected int[] destPastIntervals = new int[] {}; + protected int numDestPastIntervals = 0; /** - * Number of past source interspike intervals to consider (akin to embedding length) + * As above but for the sources */ - protected int l = 1; + public static final String SOURCE_PAST_INTERVALS_PROP_NAME = "SOURCE_PAST_INTERVALS"; + protected int[] sourcePastIntervals = new int[] {}; + protected int numSourcePastIntervals = 0; /** - * Property name for number of interspike intervals for the conditional variables + * As above but for the conditioning processes. There is no property name for this variable, + * due to there not currently being a method for converting strings to 2d arrays. Instead, + * a separate setter method is implemented. */ - public static final String COND_EMBED_LENGTHS_PROP_NAME = "COND_EMBED_LENGTHS"; - /** - * Array of history interspike interval embedding lengths for the conditional variables. - * Can be an empty array or null if there are no conditional variables. - */ - protected int[] condEmbedDims = new int[] {}; + protected Vector vectorOfCondPastIntervals = new Vector(); + protected int numCondPastIntervals = 0; /** * Number of nearest neighbours to search for in the full joint space @@ -123,13 +129,13 @@ public class TransferEntropyCalculatorSpikingIntegration implements TransferEntr * of the number of target spikes. */ public static final String PROP_SAMPLE_MULTIPLIER = "NUM_SAMPLES_MULTIPLIER"; - protected double numSamplesMultiplier = 1.0; + protected double numSamplesMultiplier = 2.0; /** * Property name for the number of random sample points to use in the construction of the surrogates as a multiple * of the number of target spikes. */ public static final String PROP_SURROGATE_SAMPLE_MULTIPLIER = "SURROGATE_NUM_SAMPLES_MULTIPLIER"; - protected double surrogateNumSamplesMultiplier = 1.0; + protected double surrogateNumSamplesMultiplier = 2.0; /** * Property for the number of nearest neighbours to use in the construction of the surrogates @@ -161,7 +167,7 @@ public class TransferEntropyCalculatorSpikingIntegration implements TransferEntr */ @Override public void initialise() throws Exception { - initialise(k, l); + initialise(0, 0); } /* @@ -173,7 +179,7 @@ public class TransferEntropyCalculatorSpikingIntegration implements TransferEntr */ @Override public void initialise(int k) throws Exception { - initialise(k, this.l); + initialise(0, 0); } /* @@ -185,11 +191,6 @@ public class TransferEntropyCalculatorSpikingIntegration implements TransferEntr */ @Override public void initialise(int k, int l) throws Exception { - if ((k < 1) || (l < 1)) { - throw new Exception("Zero history length not supported"); - } - this.k = k; - this.l = l; vectorOfSourceSpikeTimes = null; vectorOfDestinationSpikeTimes = null; } @@ -204,28 +205,32 @@ public class TransferEntropyCalculatorSpikingIntegration implements TransferEntr @Override public void setProperty(String propertyName, String propertyValue) throws Exception { boolean propertySet = true; - if (propertyName.equalsIgnoreCase(K_PROP_NAME)) { - int kTemp = Integer.parseInt(propertyValue); - if (kTemp < 1) { - throw new Exception ("Invalid k value less than 1."); + if (propertyName.equalsIgnoreCase(DEST_PAST_INTERVALS_PROP_NAME)) { + if (propertyValue.length() == 0) { + destPastIntervals = new int[] {}; } else { - k = kTemp; - } - } else if (propertyName.equalsIgnoreCase(L_PROP_NAME)) { - int lTemp = Integer.parseInt(propertyValue); - if (lTemp < 1) { - throw new Exception ("Invalid l value less than 1."); - } else { - l = lTemp; - } - } else if (propertyName.equalsIgnoreCase(COND_EMBED_LENGTHS_PROP_NAME)) { - int[] condEmbedDimsTemp = ParsedProperties.parseStringArrayOfInts(propertyValue); - for (int dim : condEmbedDimsTemp) { - if (dim < 1) { - throw new Exception ("Invalid conditional embedding value less than 1."); + int[] destPastIntervalsTemp = ParsedProperties.parseStringArrayOfInts(propertyValue); + for (int interval : destPastIntervalsTemp) { + if (interval < 1) { + throw new Exception ("Invalid interval number less than 1."); + } } + destPastIntervals = destPastIntervalsTemp; + Arrays.sort(destPastIntervals); + } + } else if (propertyName.equalsIgnoreCase(SOURCE_PAST_INTERVALS_PROP_NAME)) { + if (propertyValue.length() == 0) { + sourcePastIntervals = new int[] {}; + } else { + int[] sourcePastIntervalsTemp = ParsedProperties.parseStringArrayOfInts(propertyValue); + for (int interval : sourcePastIntervalsTemp) { + if (interval < 1) { + throw new Exception ("Invalid interval number less than 1."); + } + } + sourcePastIntervals = sourcePastIntervalsTemp; + Arrays.sort(sourcePastIntervals); } - condEmbedDims = condEmbedDimsTemp; } else if (propertyName.equalsIgnoreCase(KNNS_PROP_NAME)) { Knns = Integer.parseInt(propertyValue); } else if (propertyName.equalsIgnoreCase(PROP_K_PERM)) { @@ -238,7 +243,6 @@ public class TransferEntropyCalculatorSpikingIntegration implements TransferEntr addNoise = true; noiseLevel = Double.parseDouble(propertyValue); } - } else if (propertyName.equalsIgnoreCase(PROP_SAMPLE_MULTIPLIER)) { double tempNumSamplesMultiplier = Double.parseDouble(propertyValue); if (tempNumSamplesMultiplier <= 0) { @@ -274,11 +278,7 @@ public class TransferEntropyCalculatorSpikingIntegration implements TransferEntr */ @Override public String getProperty(String propertyName) throws Exception { - if (propertyName.equalsIgnoreCase(K_PROP_NAME)) { - return Integer.toString(k); - } else if (propertyName.equalsIgnoreCase(L_PROP_NAME)) { - return Integer.toString(l); - } else if (propertyName.equalsIgnoreCase(KNNS_PROP_NAME)) { + if (propertyName.equalsIgnoreCase(KNNS_PROP_NAME)) { return Integer.toString(Knns); } else if (propertyName.equalsIgnoreCase(PROP_ADD_NOISE)) { return Double.toString(noiseLevel); @@ -290,6 +290,21 @@ public class TransferEntropyCalculatorSpikingIntegration implements TransferEntr } } + + public void appendConditionalIntervals(int[] intervals) throws Exception{ + for (int interval : intervals) { + if (interval < 1) { + throw new Exception ("Invalid interval number less than 1."); + } + } + Arrays.sort(intervals); + vectorOfCondPastIntervals.add(intervals); + } + + public void clearConditionalIntervals() { + vectorOfCondPastIntervals = new Vector(); + } + /* * (non-Javadoc) * @@ -351,6 +366,14 @@ public class TransferEntropyCalculatorSpikingIntegration implements TransferEntr @Override public void finaliseAddObservations() throws Exception { + // Set these conveniance variables as they are used quite a bit later on + numDestPastIntervals = destPastIntervals.length; + numSourcePastIntervals = sourcePastIntervals.length; + numCondPastIntervals = 0; + for (int[] intervals : vectorOfCondPastIntervals) { + numCondPastIntervals += intervals.length; + } + conditioningEmbeddingsFromSpikes = new Vector(); jointEmbeddingsFromSpikes = new Vector(); conditioningEmbeddingsFromSamples = new Vector(); @@ -377,15 +400,19 @@ public class TransferEntropyCalculatorSpikingIntegration implements TransferEntr numSamplesMultiplier); } + + // Convert the vectors to arrays so that they can be put in the trees - double[][] arrayedTargetEmbeddingsFromSpikes = new double[conditioningEmbeddingsFromSpikes.size()][k]; - double[][] arrayedJointEmbeddingsFromSpikes = new double[conditioningEmbeddingsFromSpikes.size()][k + l]; + double[][] arrayedTargetEmbeddingsFromSpikes = new double[conditioningEmbeddingsFromSpikes.size()][numDestPastIntervals + numCondPastIntervals]; + double[][] arrayedJointEmbeddingsFromSpikes = new double[conditioningEmbeddingsFromSpikes.size()][numDestPastIntervals + + numCondPastIntervals + numSourcePastIntervals]; for (int i = 0; i < conditioningEmbeddingsFromSpikes.size(); i++) { arrayedTargetEmbeddingsFromSpikes[i] = conditioningEmbeddingsFromSpikes.elementAt(i); arrayedJointEmbeddingsFromSpikes[i] = jointEmbeddingsFromSpikes.elementAt(i); } - double[][] arrayedTargetEmbeddingsFromSamples = new double[conditioningEmbeddingsFromSamples.size()][k]; - double[][] arrayedJointEmbeddingsFromSamples = new double[conditioningEmbeddingsFromSamples.size()][k + l]; + double[][] arrayedTargetEmbeddingsFromSamples = new double[conditioningEmbeddingsFromSamples.size()][numDestPastIntervals + numCondPastIntervals]; + double[][] arrayedJointEmbeddingsFromSamples = new double[conditioningEmbeddingsFromSamples.size()][numDestPastIntervals + + numCondPastIntervals + numSourcePastIntervals]; for (int i = 0; i < conditioningEmbeddingsFromSamples.size(); i++) { arrayedTargetEmbeddingsFromSamples[i] = conditioningEmbeddingsFromSamples.elementAt(i); arrayedJointEmbeddingsFromSamples[i] = jointEmbeddingsFromSamples.elementAt(i); @@ -410,15 +437,16 @@ public class TransferEntropyCalculatorSpikingIntegration implements TransferEntr Random random = new Random(); + // Initialise the starting points of all the tracking variables int embeddingPointIndex = indexOfFirstPointToUse; - int mostRecentDestIndex = k; - int mostRecentSourceIndex = l; - int[] mostRecentConditioningIndices = Arrays.copyOf(condEmbedDims, condEmbedDims.length); - int totalLengthOfConditioningEmbeddings = 0; - for (int i = 0; i < condEmbedDims.length; i++) { - totalLengthOfConditioningEmbeddings += condEmbedDims[i]; + int mostRecentDestIndex = destPastIntervals[destPastIntervals.length - 1]; + int mostRecentSourceIndex = sourcePastIntervals[sourcePastIntervals.length - 1]; + int[] mostRecentConditioningIndices = new int[vectorOfCondPastIntervals.size()]; + for (int i = 0; i < vectorOfCondPastIntervals.size(); i++) { + mostRecentConditioningIndices[i] = vectorOfCondPastIntervals.elementAt(i).length; } + // Loop through the points at which embeddings need to be made for (; embeddingPointIndex < pointsAtWhichToMakeEmbeddings.length; embeddingPointIndex++) { @@ -432,8 +460,7 @@ public class TransferEntropyCalculatorSpikingIntegration implements TransferEntr } // Do the same for the most recent source index while (mostRecentSourceIndex < (sourceSpikeTimes.length - 1)) { - if (sourceSpikeTimes[mostRecentSourceIndex - + 1] < pointsAtWhichToMakeEmbeddings[embeddingPointIndex]) { + if (sourceSpikeTimes[mostRecentSourceIndex + 1] < pointsAtWhichToMakeEmbeddings[embeddingPointIndex]) { mostRecentSourceIndex++; } else { break; @@ -451,53 +478,70 @@ public class TransferEntropyCalculatorSpikingIntegration implements TransferEntr } - double[] conditioningPast = new double[k + totalLengthOfConditioningEmbeddings]; - double[] jointPast = new double[k + totalLengthOfConditioningEmbeddings + l]; + double[] conditioningPast = new double[numDestPastIntervals + numCondPastIntervals]; + double[] jointPast = new double[numDestPastIntervals + numCondPastIntervals + numSourcePastIntervals]; // Add the embedding intervals from the target process - conditioningPast[0] = pointsAtWhichToMakeEmbeddings[embeddingPointIndex] - destSpikeTimes[mostRecentDestIndex]; - jointPast[0] = pointsAtWhichToMakeEmbeddings[embeddingPointIndex] - - destSpikeTimes[mostRecentDestIndex]; - for (int i = 1; i < k; i++) { - conditioningPast[i] = destSpikeTimes[mostRecentDestIndex - i + 1] - - destSpikeTimes[mostRecentDestIndex - i]; - jointPast[i] = destSpikeTimes[mostRecentDestIndex - i + 1] - - destSpikeTimes[mostRecentDestIndex - i]; + for (int i = 0; i < destPastIntervals.length; i++) { + // Case where we are inserting an interval from an observation point back to the most recent event in the target process + if (destPastIntervals[i] == 1) { + conditioningPast[i] = pointsAtWhichToMakeEmbeddings[embeddingPointIndex] - destSpikeTimes[mostRecentDestIndex]; + jointPast[i] = pointsAtWhichToMakeEmbeddings[embeddingPointIndex] + - destSpikeTimes[mostRecentDestIndex]; + // Case where we are inserting an inter-event intervvl from the target process + } else { + conditioningPast[i] = destSpikeTimes[mostRecentDestIndex - destPastIntervals[i] + 2] + - destSpikeTimes[mostRecentDestIndex - destPastIntervals[i] + 1]; + jointPast[i] = destSpikeTimes[mostRecentDestIndex - destPastIntervals[i] + 2] + - destSpikeTimes[mostRecentDestIndex - destPastIntervals[i] + 1]; + } } // Add the embeding intervals from the conditional processes - int indexOfNextEmbeddingInterval = k; - for (int i = 0; i < condEmbedDims.length; i++) { - conditioningPast[indexOfNextEmbeddingInterval] = - pointsAtWhichToMakeEmbeddings[embeddingPointIndex] - conditionalSpikeTimes[i][mostRecentConditioningIndices[i]]; - jointPast[indexOfNextEmbeddingInterval] = - pointsAtWhichToMakeEmbeddings[embeddingPointIndex] - conditionalSpikeTimes[i][mostRecentConditioningIndices[i]]; - indexOfNextEmbeddingInterval += 1; - for (int j = 1; j < condEmbedDims[i]; j++) { - conditioningPast[indexOfNextEmbeddingInterval] = - conditionalSpikeTimes[i][mostRecentConditioningIndices[i] - j + 1] - - conditionalSpikeTimes[i][mostRecentConditioningIndices[i] - j]; - jointPast[indexOfNextEmbeddingInterval] = - conditionalSpikeTimes[i][mostRecentConditioningIndices[i] - j + 1] - - conditionalSpikeTimes[i][mostRecentConditioningIndices[i] - j]; - indexOfNextEmbeddingInterval += 1; + int indexOfNextEmbeddingInterval = numDestPastIntervals; + for (int i = 0; i < vectorOfCondPastIntervals.size(); i++) { + for (int j = 0; j < vectorOfCondPastIntervals.elementAt(i).length; j++) { + // Case where we are inserting an interval from an observation point back to the most recent event in the conditioning process + if (vectorOfCondPastIntervals.elementAt(i)[j] == 1) { + conditioningPast[indexOfNextEmbeddingInterval] = pointsAtWhichToMakeEmbeddings[embeddingPointIndex] + - conditionalSpikeTimes[i][mostRecentConditioningIndices[i]]; + jointPast[indexOfNextEmbeddingInterval] = pointsAtWhichToMakeEmbeddings[embeddingPointIndex] + - conditionalSpikeTimes[i][mostRecentConditioningIndices[i]]; + // Case where we are inserting an inter-event interval from the conditioning process + } else { + // Convenience variable + int intervalNumber = vectorOfCondPastIntervals.elementAt(i)[j]; + conditioningPast[indexOfNextEmbeddingInterval] = conditionalSpikeTimes[i][mostRecentConditioningIndices[i] - intervalNumber + 2] + - conditionalSpikeTimes[i][mostRecentConditioningIndices[i] - intervalNumber + 1]; + jointPast[indexOfNextEmbeddingInterval] = conditionalSpikeTimes[i][mostRecentConditioningIndices[i] - intervalNumber + 2] + - conditionalSpikeTimes[i][mostRecentConditioningIndices[i] - intervalNumber + 1]; + } + indexOfNextEmbeddingInterval++; } } // Add the embedding intervals from the source process (this only gets added to the joint embeddings) - jointPast[k + totalLengthOfConditioningEmbeddings] = pointsAtWhichToMakeEmbeddings[embeddingPointIndex] - - sourceSpikeTimes[mostRecentSourceIndex]; - for (int i = 1; i < l; i++) { - jointPast[k + totalLengthOfConditioningEmbeddings + i] = sourceSpikeTimes[mostRecentSourceIndex - i + 1] - - sourceSpikeTimes[mostRecentSourceIndex - i]; + for (int i = 0; i < sourcePastIntervals.length; i++) { + // Case where we are inserting an interval from an observation point back to the most recent event in the source process + if (sourcePastIntervals[i] == 1) { + jointPast[indexOfNextEmbeddingInterval] = pointsAtWhichToMakeEmbeddings[embeddingPointIndex] + - sourceSpikeTimes[mostRecentSourceIndex]; + // Case where we are inserting an inter-event interval from the source process + } else { + jointPast[indexOfNextEmbeddingInterval] = sourceSpikeTimes[mostRecentSourceIndex - sourcePastIntervals[i] + 2] + - sourceSpikeTimes[mostRecentSourceIndex - sourcePastIntervals[i] + 1]; + } + indexOfNextEmbeddingInterval++; } // Add Gaussian noise, if necessary if (addNoise) { for (int i = 0; i < conditioningPast.length; i++) { + //conditioningPast[i] = Math.exp(-conditioningPast[i]); conditioningPast[i] += random.nextGaussian() * noiseLevel; } for (int i = 0; i < jointPast.length; i++) { + //jointPast[i] = Math.exp(-jointPast[i]); jointPast[i] += random.nextGaussian() * noiseLevel; } } @@ -516,33 +560,36 @@ public class TransferEntropyCalculatorSpikingIntegration implements TransferEntr Arrays.sort(sourceSpikeTimes); Arrays.sort(destSpikeTimes); - int firstTargetIndexOfEMbedding = k; - while (destSpikeTimes[firstTargetIndexOfEMbedding] <= sourceSpikeTimes[l - 1]) { - firstTargetIndexOfEMbedding++; + int firstTargetIndexOfEmbedding = destPastIntervals[destPastIntervals.length - 1]; + + int furthestInterval = sourcePastIntervals[sourcePastIntervals.length - 1]; + while (destSpikeTimes[firstTargetIndexOfEmbedding] <= sourceSpikeTimes[furthestInterval - 1]) { + firstTargetIndexOfEmbedding++; } - if (conditionalSpikeTimes.length != condEmbedDims.length) { + if (conditionalSpikeTimes.length != vectorOfCondPastIntervals.size()) { throw new Exception("Number of conditional embedding lengths does not match the number of conditional processes"); - } + } for (int i = 0; i < conditionalSpikeTimes.length; i++) { - while (destSpikeTimes[firstTargetIndexOfEMbedding] <= conditionalSpikeTimes[i][condEmbedDims[i]]) { - firstTargetIndexOfEMbedding++; + furthestInterval = vectorOfCondPastIntervals.elementAt(i)[vectorOfCondPastIntervals.elementAt(i).length - 1]; + while (destSpikeTimes[firstTargetIndexOfEmbedding] <= conditionalSpikeTimes[i][furthestInterval - 1]) { + firstTargetIndexOfEmbedding++; } } // We don't want to reset these lengths when resampling for surrogates if (setProcessTimeLengths) { - processTimeLengths.add(destSpikeTimes[destSpikeTimes.length - 1] - destSpikeTimes[firstTargetIndexOfEMbedding]); + processTimeLengths.add(destSpikeTimes[destSpikeTimes.length - 1] - destSpikeTimes[firstTargetIndexOfEmbedding]); } - return firstTargetIndexOfEMbedding; + return firstTargetIndexOfEmbedding; } protected double[] generateRandomSampleTimes(double[] sourceSpikeTimes, double[] destSpikeTimes, double[][] conditionalSpikeTimes, - double actualNumSamplesMultiplier, int firstTargetIndexOfEMbedding) { + double actualNumSamplesMultiplier, int firstTargetIndexOfEmbedding) { - double sampleLowerBound = destSpikeTimes[firstTargetIndexOfEMbedding]; + double sampleLowerBound = destSpikeTimes[firstTargetIndexOfEmbedding]; double sampleUpperBound = destSpikeTimes[destSpikeTimes.length - 1]; - int num_samples = (int) Math.round(actualNumSamplesMultiplier * (destSpikeTimes.length - firstTargetIndexOfEMbedding + 1)); + int num_samples = (int) Math.round(actualNumSamplesMultiplier * (destSpikeTimes.length - firstTargetIndexOfEmbedding + 1)); double[] randomSampleTimes = new double[num_samples]; Random rand = new Random(); for (int i = 0; i < randomSampleTimes.length; i++) { @@ -559,26 +606,24 @@ public class TransferEntropyCalculatorSpikingIntegration implements TransferEntr double actualNumSamplesMultiplier) throws Exception { - int firstTargetIndexOfEMbedding = getFirstDestIndex(sourceSpikeTimes, destSpikeTimes, conditionalSpikeTimes, true); - double[] randomSampleTimes = generateRandomSampleTimes(sourceSpikeTimes, destSpikeTimes, conditionalSpikeTimes, actualNumSamplesMultiplier, firstTargetIndexOfEMbedding); + int firstTargetIndexOfEmbedding = getFirstDestIndex(sourceSpikeTimes, destSpikeTimes, conditionalSpikeTimes, true); + double[] randomSampleTimes = generateRandomSampleTimes(sourceSpikeTimes, destSpikeTimes, conditionalSpikeTimes, + actualNumSamplesMultiplier, firstTargetIndexOfEmbedding); - makeEmbeddingsAtPoints(destSpikeTimes, firstTargetIndexOfEMbedding, sourceSpikeTimes, destSpikeTimes, conditionalSpikeTimes, + makeEmbeddingsAtPoints(destSpikeTimes, firstTargetIndexOfEmbedding, sourceSpikeTimes, destSpikeTimes, conditionalSpikeTimes, conditioningEmbeddingsFromSpikes, jointEmbeddingsFromSpikes); makeEmbeddingsAtPoints(randomSampleTimes, 0, sourceSpikeTimes, destSpikeTimes, conditionalSpikeTimes, conditioningEmbeddingsFromSamples, jointEmbeddingsFromSamples); } - /* - * Method to do the - */ protected void processEventsFromSpikingTimeSeries(double[] sourceSpikeTimes, double[] destSpikeTimes, double[][] conditionalSpikeTimes, Vector conditioningEmbeddingsFromSamples, Vector jointEmbeddingsFromSamples, double actualNumSamplesMultiplier) throws Exception { - int firstTargetIndexOfEMbedding = getFirstDestIndex(sourceSpikeTimes, destSpikeTimes, conditionalSpikeTimes, false); - double[] randomSampleTimes = generateRandomSampleTimes(sourceSpikeTimes, destSpikeTimes, conditionalSpikeTimes, actualNumSamplesMultiplier, firstTargetIndexOfEMbedding); - + int firstTargetIndexOfEmbedding = getFirstDestIndex(sourceSpikeTimes, destSpikeTimes, conditionalSpikeTimes, false); + double[] randomSampleTimes = generateRandomSampleTimes(sourceSpikeTimes, destSpikeTimes, conditionalSpikeTimes, + actualNumSamplesMultiplier, firstTargetIndexOfEmbedding); makeEmbeddingsAtPoints(randomSampleTimes, 0, sourceSpikeTimes, destSpikeTimes, conditionalSpikeTimes, conditioningEmbeddingsFromSamples, jointEmbeddingsFromSamples); } @@ -732,9 +777,9 @@ public class TransferEntropyCalculatorSpikingIntegration implements TransferEntr } currentSum += (MathsUtils.digamma(kJointSpikes) - MathsUtils.digamma(kJointSamples) + - ((k + l) * (-Math.log(radiusJointSpikes) + Math.log(radiusJointSamples))) - + ((numDestPastIntervals + numCondPastIntervals + numSourcePastIntervals) * (-Math.log(radiusJointSpikes) + Math.log(radiusJointSamples))) - MathsUtils.digamma(kConditioningSpikes) + MathsUtils.digamma(kConditioningSamples) + - + (k * (Math.log(radiusConditioningSpikes) - Math.log(radiusConditioningSamples)))); + + ((numDestPastIntervals + numCondPastIntervals) * (Math.log(radiusConditioningSpikes) - Math.log(radiusConditioningSamples)))); if (Double.isNaN(currentSum)) { throw new Exception("NaNs in TE clac"); } @@ -786,9 +831,8 @@ public class TransferEntropyCalculatorSpikingIntegration implements TransferEntr resampledConditioningEmbeddingsFromSamples, resampledJointEmbeddingsFromSamples, surrogateNumSamplesMultiplier); } - // Convert the vectors to arrays so that they can be put in the trees - double[][] arrayedResampledConditioningEmbeddingsFromSamples = new double[resampledConditioningEmbeddingsFromSamples.size()][k]; + double[][] arrayedResampledConditioningEmbeddingsFromSamples = new double[resampledConditioningEmbeddingsFromSamples.size()][numDestPastIntervals]; for (int i = 0; i < resampledConditioningEmbeddingsFromSamples.size(); i++) { arrayedResampledConditioningEmbeddingsFromSamples[i] = resampledConditioningEmbeddingsFromSamples.elementAt(i); } @@ -818,9 +862,9 @@ public class TransferEntropyCalculatorSpikingIntegration implements TransferEntr } usedIndices.add(chosenIndex); int embeddingLength = conditionallyPermutedJointEmbeddingsFromSpikes.elementAt(i).length; - for(int j = 0; j < l; j++) { - conditionallyPermutedJointEmbeddingsFromSpikes.elementAt(i)[embeddingLength - l + j] = - resampledJointEmbeddingsFromSamples.elementAt(chosenIndex)[embeddingLength - l + j]; + for(int j = 0; j < numSourcePastIntervals; j++) { + conditionallyPermutedJointEmbeddingsFromSpikes.elementAt(i)[embeddingLength - numSourcePastIntervals + j] = + resampledJointEmbeddingsFromSamples.elementAt(chosenIndex)[embeddingLength - numSourcePastIntervals + j]; } } @@ -831,7 +875,6 @@ public class TransferEntropyCalculatorSpikingIntegration implements TransferEntr } KdTree conditionallyPermutedKdTreeJointFromSpikes = new KdTree(arrayedConditionallyPermutedJointEmbeddingsFromSpikes); conditionallyPermutedKdTreeJointFromSpikes.setNormType(normType); - surrogateTEValues[permutationNumber] = computeAverageLocalOfObservations(conditionallyPermutedKdTreeJointFromSpikes, conditionallyPermutedJointEmbeddingsFromSpikes); } diff --git a/tester.py b/tester.py index cae3631..bd8265b 100755 --- a/tester.py +++ b/tester.py @@ -30,6 +30,7 @@ import numpy as np NUM_REPS = 2 NUM_SPIKES = int(2e3) NUM_OBSERVATIONS = 2 +NUM_SURROGATES = 10 # Params for canonical example generation RATE_Y = 1.0 @@ -80,9 +81,10 @@ teCalcClass = JPackage("infodynamics.measures.spiking.integration").TransferEntr teCalc = teCalcClass() teCalc.setProperty("knns", "4") print("Independent Poisson Processes") -teCalc.setProperty("COND_EMBED_LENGTHS", "2,2") -teCalc.setProperty("k_HISTORY", "2") -teCalc.setProperty("l_HISTORY", "2") +teCalc.setProperty("DEST_PAST_INTERVALS", "1,2") +teCalc.setProperty("SOURCE_PAST_INTERVALS", "1,2") +teCalc.appendConditionalIntervals(JArray(JInt, 1)([1, 2])) +teCalc.appendConditionalIntervals(JArray(JInt, 1)([1, 2])) teCalc.setProperty("NORM_TYPE", "MAX_NORM") results_poisson = np.zeros(NUM_REPS) @@ -99,23 +101,17 @@ for i in range(NUM_REPS): teCalc.finaliseAddObservations(); result = teCalc.computeAverageLocalOfObservations() print("TE result %.4f nats" % (result,)) - sig = teCalc.computeSignificance(10, result) + sig = teCalc.computeSignificance(NUM_SURROGATES, result) print(sig.pValue) results_poisson[i] = result print("Summary: mean ", np.mean(results_poisson), " std dev ", np.std(results_poisson)) - - - - - - teCalc = teCalcClass() teCalc.setProperty("knns", "4") print("Noisy copy zero TE") -teCalc.setProperty("COND_EMBED_LENGTHS", "1") -teCalc.setProperty("k_HISTORY", "2") -teCalc.setProperty("l_HISTORY", "1") +#teCalc.appendConditionalIntervals(JArray(JInt, 1)([1])) +teCalc.setProperty("DEST_PAST_INTERVALS", "1") +teCalc.setProperty("SOURCE_PAST_INTERVALS", "1") #teCalc.setProperty("NORM_TYPE", "MAX_NORM") results_noisy_zero = np.zeros(NUM_REPS) @@ -129,11 +125,12 @@ for i in range(NUM_REPS): sourceArray.sort() destArray = condArray[0, :] + 0.5 + 0.05 * np.random.normal(size = condArray.shape[1]) destArray.sort() - teCalc.addObservations(JArray(JDouble, 1)(sourceArray), JArray(JDouble, 1)(destArray), JArray(JDouble, 2)(condArray)) + #teCalc.addObservations(JArray(JDouble, 1)(sourceArray), JArray(JDouble, 1)(destArray), JArray(JDouble, 2)(condArray)) + teCalc.addObservations(JArray(JDouble, 1)(sourceArray), JArray(JDouble, 1)(destArray)) teCalc.finaliseAddObservations(); result = teCalc.computeAverageLocalOfObservations() print("TE result %.4f nats" % (result,)) - sig = teCalc.computeSignificance(10, result) + sig = teCalc.computeSignificance(NUM_SURROGATES, result) print(sig.pValue) results_noisy_zero[i] = result print("Summary: mean ", np.mean(results_noisy_zero), " std dev ", np.std(results_noisy_zero)) @@ -143,9 +140,9 @@ print("Summary: mean ", np.mean(results_noisy_zero), " std dev ", np.std(results teCalc = teCalcClass() teCalc.setProperty("knns", "4") print("Noisy copy non-zero TE") -teCalc.setProperty("COND_EMBED_LENGTHS", "1") -teCalc.setProperty("k_HISTORY", "2") -teCalc.setProperty("l_HISTORY", "1") +teCalc.appendConditionalIntervals(JArray(JInt, 1)([1])) +teCalc.setProperty("DEST_PAST_INTERVALS", "1,2") +teCalc.setProperty("SOURCE_PAST_INTERVALS", "1") #teCalc.setProperty("NORM_TYPE", "MAX_NORM") results_noisy_non_zero = np.zeros(NUM_REPS) @@ -164,7 +161,7 @@ for i in range(NUM_REPS): teCalc.finaliseAddObservations(); result = teCalc.computeAverageLocalOfObservations() print("TE result %.4f nats" % (result,)) - sig = teCalc.computeSignificance(10, result) + sig = teCalc.computeSignificance(NUM_SURROGATES, result) print(sig.pValue) results_noisy_non_zero[i] = result print("Summary: mean ", np.mean(results_noisy_non_zero), " std dev ", np.std(results_noisy_zero)) @@ -172,8 +169,8 @@ print("Summary: mean ", np.mean(results_noisy_non_zero), " std dev ", np.std(res print("Canonical example") teCalc = teCalcClass() teCalc.setProperty("knns", "4") -teCalc.setProperty("k_HISTORY", "1") -teCalc.setProperty("l_HISTORY", "2") +teCalc.setProperty("DEST_PAST_INTERVALS", "1,2") +teCalc.setProperty("SOURCE_PAST_INTERVALS", "1") #teCalc.setProperty("NUM_SAMPLES_MULTIPLIER", "1") #teCalc.setProperty("NORM_TYPE", "MAX_NORM") @@ -184,6 +181,6 @@ for i in range(NUM_REPS): result = teCalc.computeAverageLocalOfObservations() results_canonical[i] = result print("TE result %.4f nats" % (result,)) - sig = teCalc.computeSignificance(10, result) + sig = teCalc.computeSignificance(NUM_SURROGATES, result) print(sig.pValue) print("Summary: mean ", np.mean(results_canonical), " std dev ", np.std(results_canonical))