From c068c8308d97e193ab922a1e2b3b98a57ef1898f Mon Sep 17 00:00:00 2001 From: David Shorten Date: Fri, 25 Feb 2022 18:49:12 +1100 Subject: [PATCH] improving jittered surrogates --- ...erEntropyCalculatorSpikingIntegration.java | 45 ++++++++++++++----- tester.py | 8 ++++ 2 files changed, 41 insertions(+), 12 deletions(-) diff --git a/java/source/infodynamics/measures/spiking/integration/TransferEntropyCalculatorSpikingIntegration.java b/java/source/infodynamics/measures/spiking/integration/TransferEntropyCalculatorSpikingIntegration.java index ede39a4..653bcc9 100644 --- a/java/source/infodynamics/measures/spiking/integration/TransferEntropyCalculatorSpikingIntegration.java +++ b/java/source/infodynamics/measures/spiking/integration/TransferEntropyCalculatorSpikingIntegration.java @@ -119,6 +119,15 @@ public class TransferEntropyCalculatorSpikingIntegration implements TransferEntr */ protected double noiseLevel = (double) 1e-8; + /** + * Whether to use the jittered sampling approach. Useful for bursty spike trains. Explained in methods section of + * doi.org/10.1101/2021.06.29.450432 + */ + protected boolean jitteredSamplesForSurrogates = false; + public static final String DO_JITTERED_SAMPLING_PROP_NAME = "DO_JITTERED_SAMPLING"; + protected double jitteredSamplingNoiseLevel = 1; + public static final String JITTERED_SAMPLING_NOISE_LEVEL = "JITTERED_SAMPLING_NOISE_LEVEL"; + /** * Stores whether we are in debug mode */ @@ -233,6 +242,10 @@ public class TransferEntropyCalculatorSpikingIntegration implements TransferEntr } } else if (propertyName.equalsIgnoreCase(KNNS_PROP_NAME)) { Knns = Integer.parseInt(propertyValue); + } else if (propertyName.equalsIgnoreCase(DO_JITTERED_SAMPLING_PROP_NAME)) { + jitteredSamplesForSurrogates = Boolean.parseBoolean(propertyValue); + } else if (propertyName.equalsIgnoreCase(JITTERED_SAMPLING_NOISE_LEVEL)) { + jitteredSamplingNoiseLevel = Double.parseDouble(propertyValue); } else if (propertyName.equalsIgnoreCase(PROP_K_PERM)) { kPerm = Integer.parseInt(propertyValue); } else if (propertyName.equalsIgnoreCase(PROP_ADD_NOISE)) { @@ -397,7 +410,7 @@ public class TransferEntropyCalculatorSpikingIntegration implements TransferEntr } processEventsFromSpikingTimeSeries(sourceSpikeTimes, destSpikeTimes, conditionalSpikeTimes, conditioningEmbeddingsFromSpikes, jointEmbeddingsFromSpikes, conditioningEmbeddingsFromSamples, jointEmbeddingsFromSamples, - numSamplesMultiplier); + numSamplesMultiplier, false); } @@ -588,18 +601,18 @@ public class TransferEntropyCalculatorSpikingIntegration implements TransferEntr } protected double[] generateRandomSampleTimes(double[] sourceSpikeTimes, double[] destSpikeTimes, double[][] conditionalSpikeTimes, - double actualNumSamplesMultiplier, int firstTargetIndexOfEmbedding) { + double actualNumSamplesMultiplier, int firstTargetIndexOfEmbedding, boolean doJitteredSampling) { double sampleLowerBound = destSpikeTimes[firstTargetIndexOfEmbedding]; double sampleUpperBound = destSpikeTimes[destSpikeTimes.length - 1]; int num_samples = (int) Math.round(actualNumSamplesMultiplier * (destSpikeTimes.length - firstTargetIndexOfEmbedding + 1)); double[] randomSampleTimes = new double[num_samples]; Random rand = new Random(); - boolean doCellCulture = true; - if (doCellCulture) { + if (doJitteredSampling) { + //System.out.println("jittering " + jitteredSamplingNoiseLevel); for (int i = 0; i < randomSampleTimes.length; i++) { randomSampleTimes[i] = destSpikeTimes[firstTargetIndexOfEmbedding + (i % (destSpikeTimes.length - firstTargetIndexOfEmbedding - 1))] - + 200 * (rand.nextDouble() - 0.5); + + jitteredSamplingNoiseLevel * (rand.nextDouble() - 0.5); //randomSampleTimes[i] = -1.0; if ((randomSampleTimes[i] > sampleUpperBound) || (randomSampleTimes[i] < sampleLowerBound)) { randomSampleTimes[i] = sampleLowerBound + rand.nextDouble() * (sampleUpperBound - sampleLowerBound); @@ -622,12 +635,13 @@ public class TransferEntropyCalculatorSpikingIntegration implements TransferEntr protected void processEventsFromSpikingTimeSeries(double[] sourceSpikeTimes, double[] destSpikeTimes, double[][] conditionalSpikeTimes, Vector conditioningEmbeddingsFromSpikes, Vector jointEmbeddingsFromSpikes, Vector conditioningEmbeddingsFromSamples, Vector jointEmbeddingsFromSamples, - double actualNumSamplesMultiplier) + double actualNumSamplesMultiplier, boolean doJitteredSampling) throws Exception { int firstTargetIndexOfEmbedding = getFirstDestIndex(sourceSpikeTimes, destSpikeTimes, conditionalSpikeTimes, true); double[] randomSampleTimes = generateRandomSampleTimes(sourceSpikeTimes, destSpikeTimes, conditionalSpikeTimes, - actualNumSamplesMultiplier, firstTargetIndexOfEmbedding); + actualNumSamplesMultiplier, firstTargetIndexOfEmbedding, + doJitteredSampling); makeEmbeddingsAtPoints(destSpikeTimes, firstTargetIndexOfEmbedding, sourceSpikeTimes, destSpikeTimes, conditionalSpikeTimes, conditioningEmbeddingsFromSpikes, jointEmbeddingsFromSpikes); @@ -637,12 +651,13 @@ public class TransferEntropyCalculatorSpikingIntegration implements TransferEntr protected void processEventsFromSpikingTimeSeries(double[] sourceSpikeTimes, double[] destSpikeTimes, double[][] conditionalSpikeTimes, Vector conditioningEmbeddingsFromSamples, Vector jointEmbeddingsFromSamples, - double actualNumSamplesMultiplier) + double actualNumSamplesMultiplier, boolean doJitteredSampling) throws Exception { int firstTargetIndexOfEmbedding = getFirstDestIndex(sourceSpikeTimes, destSpikeTimes, conditionalSpikeTimes, false); double[] randomSampleTimes = generateRandomSampleTimes(sourceSpikeTimes, destSpikeTimes, conditionalSpikeTimes, - actualNumSamplesMultiplier, firstTargetIndexOfEmbedding); + actualNumSamplesMultiplier, firstTargetIndexOfEmbedding, + doJitteredSampling); makeEmbeddingsAtPoints(randomSampleTimes, 0, sourceSpikeTimes, destSpikeTimes, conditionalSpikeTimes, conditioningEmbeddingsFromSamples, jointEmbeddingsFromSamples); } @@ -850,9 +865,15 @@ public class TransferEntropyCalculatorSpikingIntegration implements TransferEntr } else { conditionalSpikeTimes = new double[][] {}; } - processEventsFromSpikingTimeSeries(sourceSpikeTimes, destSpikeTimes, conditionalSpikeTimes, - resampledConditioningEmbeddingsFromSamples, resampledJointEmbeddingsFromSamples, - surrogateNumSamplesMultiplier); + if (jitteredSamplesForSurrogates) { + processEventsFromSpikingTimeSeries(sourceSpikeTimes, destSpikeTimes, conditionalSpikeTimes, + resampledConditioningEmbeddingsFromSamples, resampledJointEmbeddingsFromSamples, + surrogateNumSamplesMultiplier, true); + } else { + processEventsFromSpikingTimeSeries(sourceSpikeTimes, destSpikeTimes, conditionalSpikeTimes, + resampledConditioningEmbeddingsFromSamples, resampledJointEmbeddingsFromSamples, + surrogateNumSamplesMultiplier, false); + } } // Convert the vectors to arrays so that they can be put in the trees double[][] arrayedResampledConditioningEmbeddingsFromSamples = new double[resampledConditioningEmbeddingsFromSamples.size()][numDestPastIntervals]; diff --git a/tester.py b/tester.py index bd8265b..1caac7d 100755 --- a/tester.py +++ b/tester.py @@ -83,6 +83,8 @@ teCalc.setProperty("knns", "4") print("Independent Poisson Processes") teCalc.setProperty("DEST_PAST_INTERVALS", "1,2") teCalc.setProperty("SOURCE_PAST_INTERVALS", "1,2") +teCalc.setProperty("DO_JITTERED_SAMPLING", "true") +teCalc.setProperty("JITTERED_SAMPLING_NOISE_LEVEL", "0") teCalc.appendConditionalIntervals(JArray(JInt, 1)([1, 2])) teCalc.appendConditionalIntervals(JArray(JInt, 1)([1, 2])) teCalc.setProperty("NORM_TYPE", "MAX_NORM") @@ -112,6 +114,8 @@ print("Noisy copy zero TE") #teCalc.appendConditionalIntervals(JArray(JInt, 1)([1])) teCalc.setProperty("DEST_PAST_INTERVALS", "1") teCalc.setProperty("SOURCE_PAST_INTERVALS", "1") +teCalc.setProperty("DO_JITTERED_SAMPLING", "true") +teCalc.setProperty("JITTERED_SAMPLING_NOISE_LEVEL", "0") #teCalc.setProperty("NORM_TYPE", "MAX_NORM") results_noisy_zero = np.zeros(NUM_REPS) @@ -143,6 +147,8 @@ print("Noisy copy non-zero TE") teCalc.appendConditionalIntervals(JArray(JInt, 1)([1])) teCalc.setProperty("DEST_PAST_INTERVALS", "1,2") teCalc.setProperty("SOURCE_PAST_INTERVALS", "1") +teCalc.setProperty("DO_JITTERED_SAMPLING", "true") +teCalc.setProperty("JITTERED_SAMPLING_NOISE_LEVEL", "0") #teCalc.setProperty("NORM_TYPE", "MAX_NORM") results_noisy_non_zero = np.zeros(NUM_REPS) @@ -171,6 +177,8 @@ teCalc = teCalcClass() teCalc.setProperty("knns", "4") teCalc.setProperty("DEST_PAST_INTERVALS", "1,2") teCalc.setProperty("SOURCE_PAST_INTERVALS", "1") +teCalc.setProperty("DO_JITTERED_SAMPLING", "true") +teCalc.setProperty("JITTERED_SAMPLING_NOISE_LEVEL", "0") #teCalc.setProperty("NUM_SAMPLES_MULTIPLIER", "1") #teCalc.setProperty("NORM_TYPE", "MAX_NORM")