diff --git a/java/source/infodynamics/measures/continuous/kraskov/ConditionalMutualInfoCalculatorMultiVariateKraskov.java b/java/source/infodynamics/measures/continuous/kraskov/ConditionalMutualInfoCalculatorMultiVariateKraskov.java index da3601d..2c2490b 100755 --- a/java/source/infodynamics/measures/continuous/kraskov/ConditionalMutualInfoCalculatorMultiVariateKraskov.java +++ b/java/source/infodynamics/measures/continuous/kraskov/ConditionalMutualInfoCalculatorMultiVariateKraskov.java @@ -960,6 +960,22 @@ public abstract class ConditionalMutualInfoCalculatorMultiVariateKraskov public double[] kNNDistancesForNewSamples(int startTimePoint, int numTimePoints, double[][] newStates1, double[][] newStates2, double[][] newCondStates) throws Exception { + // Do normalisation of the incoming data if required: + double[][] states1ToUse, states2ToUse, condStatesToUse; + if (normalise) { + states1ToUse = MatrixUtils.normaliseIntoNewArray(newStates1, var1Means, var1Stds); + states2ToUse = MatrixUtils.normaliseIntoNewArray(newStates2, var2Means, var2Stds); + if (dimensionsCond != 0) { + condStatesToUse = MatrixUtils.normaliseIntoNewArray(newCondStates, condMeans, condStds); + } else { + condStatesToUse = null; + } + } else { + states1ToUse = newStates1; + states2ToUse = newStates2; + condStatesToUse = newCondStates; + } + double[] kNNdistances = new double[numTimePoints]; for (int t = startTimePoint; t < startTimePoint + numTimePoints; t++) { @@ -967,7 +983,7 @@ public abstract class ConditionalMutualInfoCalculatorMultiVariateKraskov // finding the kth closest neighbour for the new sample: PriorityQueue nnPQ = kdTreeJoint.findKNearestNeighbours(k, - new double[][] {newStates1[t], newStates2[t], newCondStates[t]}); + new double[][] {states1ToUse[t], states2ToUse[t], condStatesToUse[t]}); // First element in the PQ is the kth NN, // and epsilon = kthNnData.distance NeighbourNodeData kthNnData = nnPQ.poll(); @@ -997,7 +1013,7 @@ public abstract class ConditionalMutualInfoCalculatorMultiVariateKraskov "the univariate kNNDistancesForNewSamples(int, int, double[],double[],double[][]) " + "method is called"); } - return computeLocalUsingPreviousObservations( + return kNNDistancesForNewSamples(startTimePoint, numTimePoints, MatrixUtils.reshape(newStates1, newStates1.length, 1), MatrixUtils.reshape(newStates2, newStates2.length, 1), newCondStates); @@ -1023,7 +1039,7 @@ public abstract class ConditionalMutualInfoCalculatorMultiVariateKraskov "the univariate kNNDistancesForNewSamples(int, int, double[],double[],double[]) " + "method is called"); } - return computeLocalUsingPreviousObservations( + return kNNDistancesForNewSamples(startTimePoint, numTimePoints, MatrixUtils.reshape(newStates1, newStates1.length, 1), MatrixUtils.reshape(newStates2, newStates2.length, 1), MatrixUtils.reshape(newCondStates, newCondStates.length, 1)); diff --git a/java/source/infodynamics/measures/continuous/kraskov/MutualInfoCalculatorMultiVariateKraskov.java b/java/source/infodynamics/measures/continuous/kraskov/MutualInfoCalculatorMultiVariateKraskov.java index 1ada2df..04ecd98 100755 --- a/java/source/infodynamics/measures/continuous/kraskov/MutualInfoCalculatorMultiVariateKraskov.java +++ b/java/source/infodynamics/measures/continuous/kraskov/MutualInfoCalculatorMultiVariateKraskov.java @@ -1158,6 +1158,20 @@ public abstract class MutualInfoCalculatorMultiVariateKraskov public double[] kNNDistancesForNewSamples(int startTimePoint, int numTimePoints, double[][] newVar1Observations, double[][] newVar2Observations) throws Exception { + double[][] states1ToUse, states2ToUse; + if (normalise) { + states1ToUse = MatrixUtils.normaliseIntoNewArray(newVar1Observations, sourceMeansBeforeNorm, sourceStdsBeforeNorm, 0, newVar1Observations.length-timeDiff); + states2ToUse = MatrixUtils.normaliseIntoNewArray(newVar2Observations, destMeansBeforeNorm, destStdsBeforeNorm, timeDiff, newVar2Observations.length-timeDiff); + } else { + if (timeDiff > 0) { + states1ToUse = MatrixUtils.selectRows(newVar1Observations, 0, newVar1Observations.length-timeDiff); + states2ToUse = MatrixUtils.selectRows(newVar2Observations, 0, newVar2Observations.length-timeDiff); + } else { + states1ToUse = newVar1Observations; + states2ToUse = newVar2Observations; + } + } + double[] kNNdistances = new double[numTimePoints]; for (int t = startTimePoint; t < startTimePoint + numTimePoints; t++) { @@ -1165,7 +1179,7 @@ public abstract class MutualInfoCalculatorMultiVariateKraskov // finding the kth closest neighbour for the new sample: PriorityQueue nnPQ = kdTreeJoint.findKNearestNeighbours(k, - new double[][] {newVar1Observations[t], newVar2Observations[t]}); + new double[][] {states1ToUse[t], states2ToUse[t]}); // First element in the PQ is the kth NN, // and epsilon = kthNnData.distance NeighbourNodeData kthNnData = nnPQ.poll(); @@ -1195,7 +1209,7 @@ public abstract class MutualInfoCalculatorMultiVariateKraskov "the univariate kNNDistancesForNewSamples(int, int, double[],double[]) " + "method is called"); } - return computeLocalUsingPreviousObservations( + return kNNDistancesForNewSamples(startTimePoint, numTimePoints, MatrixUtils.reshape(newVar1Observations, newVar1Observations.length, 1), MatrixUtils.reshape(newVar2Observations, newVar2Observations.length, 1)); }