diff --git a/java/source/infodynamics/utils/KdTree.java b/java/source/infodynamics/utils/KdTree.java
index b049a2d..1076af7 100755
--- a/java/source/infodynamics/utils/KdTree.java
+++ b/java/source/infodynamics/utils/KdTree.java
@@ -3128,6 +3128,136 @@ public class KdTree extends NearestNeighbourSearcher {
return count;
}
+ @Override
+ public int countPointsWithinR(int sampleIndex, double r,
+ boolean allowEqualToR, boolean[] additionalCriteria, int[] remapping) {
+ if (rootNode == null) {
+ return 0;
+ }
+ return countPointsWithinR(sampleIndex, rootNode, 0, r, allowEqualToR,
+ additionalCriteria, remapping);
+ }
+
+ /**
+ * Count the number of points within radius r of a given sample (sampleIndex),
+ * in the tree rooted at node (which is at the specified level in the tree),
+ * subject to an additional criteria,
+ * and the search points are reindexed according to the remapping specified in remapping.
+ *
+ * Nearest neighbour function to compare to r is a max norm between the
+ * high-level variables, with norm for each variable being the specified norm.
+ * (If {@link EuclideanUtils#NORM_EUCLIDEAN} was selected, then the supplied
+ * r should be the required Euclidean norm squared, since we switch it
+ * to {@link EuclideanUtils#NORM_EUCLIDEAN_SQUARED} internally).
+ *
+ * @param sampleIndex sample index in the data to find a nearest neighbour
+ * for
+ * @param node node to start searching from in the kd-tree. Cannot be null
+ * @param level which level we're currently at in the tree
+ * @param r radius within which to count points
+ * @param allowEqualToR if true, then count points at radius r also,
+ * otherwise only those strictly within r
+ * @param additionalCriteria array of booleans. Only count a point if it
+ * is within r and is true in additionalCrtieria.
+ * @param remapping array of time indices with which to remap the search points
+ * onto the same time index space as the additionalCriteria (this will
+ * apply to the supplied sampleIndex as well as other search points)
+ * @return count of points within r
+ */
+ protected int countPointsWithinR(int sampleIndex,
+ KdTreeNode node, int level, double r, boolean allowEqualToR,
+ boolean[] additionalCriteria, int[] remapping) {
+
+ int count = 0;
+
+ // Point to the correct array for the data at this level
+ int currentDim = level % totalDimensions;
+ double[][] data = dimensionToArray[currentDim];
+ int actualDim = dimensionToArrayIndex[currentDim];
+
+ // Check the distance on this particular dimension
+ double distOnThisDim = data[sampleIndex][actualDim] -
+ data[node.indexOfThisPoint][actualDim];
+
+ double absDistOnThisDim;
+ if (normTypeToUse == EuclideanUtils.NORM_MAX_NORM) {
+ absDistOnThisDim = (distOnThisDim > 0) ? distOnThisDim : - distOnThisDim;
+ } else {
+ // norm type is EuclideanUtils#NORM_EUCLIDEAN_SQUARED
+ // Track the square distance
+ absDistOnThisDim = distOnThisDim * distOnThisDim;
+ }
+
+ if ((node.indexOfThisPoint != sampleIndex) &&
+ ((absDistOnThisDim < r) ||
+ ( allowEqualToR && (absDistOnThisDim == r)))) {
+ // Preliminary check says we need to compute the full distance
+ // to use or at least to check if it should be counted.
+ if (additionalCriteria[remapping[node.indexOfThisPoint]]) {
+ // The additional criteria was ok, so continue:
+ boolean withinBounds = true;
+ for (int v = 0; v < originalDataSets.length; v++) {
+ // For each of our separate (multivariate) variables,
+ // compute the (specified) norm in that variable's space:
+ double distForVariableV;
+ // Distance calculation terminates early with Double.POSITIVE_INFINITY
+ // if it is clearly larger than r:
+ distForVariableV = normWithAbort(
+ originalDataSets[v][sampleIndex],
+ originalDataSets[v][node.indexOfThisPoint],
+ r, normTypeToUse);
+ if ((distForVariableV >= r) &&
+ !(allowEqualToR && (distForVariableV == r))) {
+ // We don't fit on this dimension, no point
+ // checking the others:
+ withinBounds = false;
+ break;
+ }
+ }
+ if (withinBounds) {
+ // This node gets counted
+ count++;
+ }
+ }
+ }
+
+ KdTreeNode closestSubTree = null;
+ KdTreeNode furthestSubTree = null;
+ // And translate this to which subtree is closer
+ if (distOnThisDim < 0) {
+ // We need to search the left tree
+ closestSubTree = node.leftTree;
+ furthestSubTree = node.rightTree;
+ } else {
+ // We need to search the right tree
+ closestSubTree = node.rightTree;
+ furthestSubTree = node.leftTree;
+ }
+ // Update the search on that subtree
+ if (closestSubTree != null) {
+ count += countPointsWithinR(sampleIndex, closestSubTree,
+ level + 1, r, allowEqualToR, additionalCriteria, remapping);
+ }
+ if ((absDistOnThisDim < r) ||
+ ( allowEqualToR && (distOnThisDim < 0) && (absDistOnThisDim == r))) {
+ // It's possible we could have a node within (or on) r
+ // in the other branch as well, so search there too.
+ // (Note: we only check furthest subtree in the == case
+ // when it's allowed
+ // *if* it's the right subtree, as only the right sub-tree
+ // can have node with distance in this coordinate *equal* to
+ // that of the current node -- left subtree must be strictly
+ // less than the coordinate of the current node, so
+ // distance to any of those points could not be equal.)
+ if (furthestSubTree != null) {
+ count += countPointsWithinR(sampleIndex, furthestSubTree,
+ level + 1, r, allowEqualToR, additionalCriteria, remapping);
+ }
+ }
+
+ return count;
+ }
+
@Override
public int countPointsWithinR(double[][] sampleVectors, double r,
boolean allowEqualToR, boolean[] additionalCriteria) {
diff --git a/java/source/infodynamics/utils/NearestNeighbourSearcher.java b/java/source/infodynamics/utils/NearestNeighbourSearcher.java
index 9d9f216..5f66dcb 100755
--- a/java/source/infodynamics/utils/NearestNeighbourSearcher.java
+++ b/java/source/infodynamics/utils/NearestNeighbourSearcher.java
@@ -42,6 +42,8 @@ public abstract class NearestNeighbourSearcher {
/**
* Factory method to construct the searcher from a set of double[][] data.
+ * This will return a {@link KdTree} or if the data is univaraite
+ * (i.e. only one column) a {@link UnivariateNearestNeighbourSearcher}
*
* @param data a double[][] 2D data set, first indexed
* by time, second index by variable number.
@@ -546,6 +548,29 @@ public abstract class NearestNeighbourSearcher {
public abstract int countPointsWithinR(int sampleIndex, double r,
boolean allowEqualToR, boolean[] additionalCriteria);
+ /**
+ * As per {@link #countPointsWithinR(int, double, boolean)}
+ * however each point is subject to also meeting the additional
+ * criteria of being true in additionalCriteria,
+ * and the search points are reindexed according to the remapping
+ * specified in remapping
+ *
+ * @param sampleIndex sample index in the data to find a nearest neighbour
+ * for (already remapped if required)
+ * @param r radius within which to count points
+ * @param allowEqualToR if true, then count points at radius r also,
+ * otherwise only those strictly within r
+ * @param additionalCriteria array of booleans. Only count a point if it
+ * is within r and is true in additionalCrtieria.
+ * @param remapping array of time indices with which to remap the search points
+ * onto the same time index space as the additionalCriteria (this will
+ * apply to the supplied sampleIndex as well as other search points)
+ * @return the count of points within r.
+ */
+ public abstract int countPointsWithinR(int sampleIndex, double r,
+ boolean allowEqualToR, boolean[] additionalCriteria,
+ int[] remapping);
+
/**
* As per {@link #countPointsWithinR(double[][], double, boolean)}
* however each point is subject to also meeting the additional
diff --git a/java/source/infodynamics/utils/UnivariateNearestNeighbourSearcher.java b/java/source/infodynamics/utils/UnivariateNearestNeighbourSearcher.java
index 73e7b58..3c3d3db 100755
--- a/java/source/infodynamics/utils/UnivariateNearestNeighbourSearcher.java
+++ b/java/source/infodynamics/utils/UnivariateNearestNeighbourSearcher.java
@@ -1000,6 +1000,60 @@ public class UnivariateNearestNeighbourSearcher extends NearestNeighbourSearcher
return count;
}
+ @Override
+ public int countPointsWithinR(int sampleIndex, double r, boolean allowEqualToR,
+ boolean[] additionalCriteria, int[] remapping) {
+
+ int count = 0;
+ // Find where this node sits in the sorted array:
+ int indexInSortedArray = indicesInSortedArray[sampleIndex];
+ // Check the points with smaller data values first:
+ for (int i = indexInSortedArray - 1; i >= 0; i--) {
+ if (!additionalCriteria[remapping[sortedArrayIndices[i]]]) {
+ // This point failed the additional criteria,
+ // so skip checking it.
+ // We check this first, even though it's the norm
+ // that really determines whether we need to continue
+ // checking or not. In some circumstances this may be
+ // slower, but for our main application - KSG conditional MI -
+ // this should be faster.
+ continue;
+ }
+ double theNorm = norm(originalDataSet[sampleIndex],
+ originalDataSet[sortedArrayIndices[i]], normTypeToUse);
+ if ((allowEqualToR && (theNorm <= r)) ||
+ (!allowEqualToR && (theNorm < r))) {
+ count++;
+ continue;
+ }
+ // Else no point checking further points
+ break;
+ }
+ // Next check the points with larger data values:
+ for (int i = indexInSortedArray + 1; i < numObservations; i++) {
+ if (!additionalCriteria[remapping[sortedArrayIndices[i]]]) {
+ // This point failed the additional criteria,
+ // so skip checking it.
+ // We check this first, even though it's the norm
+ // that really determines whether we need to continue
+ // checking or not. In some circumstances this may be
+ // slower, but for our main application - KSG conditional MI -
+ // this should be faster.
+ continue;
+ }
+ double theNorm = norm(originalDataSet[sampleIndex],
+ originalDataSet[sortedArrayIndices[i]], normTypeToUse);
+ if ((allowEqualToR && (theNorm <= r)) ||
+ (!allowEqualToR && (theNorm < r))) {
+ count++;
+ continue;
+ }
+ // Else no point checking further points
+ break;
+ }
+ return count;
+ }
+
/**
* As per {@link #countPointsWithinR(int, double, boolean, boolean[])}
* only for a specific data point.
diff --git a/java/unittests/infodynamics/utils/KdTreeTest.java b/java/unittests/infodynamics/utils/KdTreeTest.java
index e4c28f7..a02a571 100755
--- a/java/unittests/infodynamics/utils/KdTreeTest.java
+++ b/java/unittests/infodynamics/utils/KdTreeTest.java
@@ -1396,6 +1396,95 @@ public class KdTreeTest extends TestCase {
}
}
+ public void testAdditionalCriteriaWithRemapping() throws Exception {
+ // The following will include a little unit testing
+ // of UnivariateNearestNeighbourSearcher
+ // for the data0 in both cases, and the remaining data in
+ // the last:
+ validateAdditionalCriteriaWithRemapping(3, 1, false);
+ validateAdditionalCriteriaWithRemapping(2, 1, false);
+ // and now these will only test KdTree:
+ validateAdditionalCriteriaWithRemapping(3, 3, true);
+ validateAdditionalCriteriaWithRemapping(3, 3, false);
+ }
+
+ public void validateAdditionalCriteriaWithRemapping(int variables,
+ int dimensionsPerVariable, boolean allowEqualsR) throws Exception {
+ int samples = 1000;
+ double[][][] data = new double[variables][][];
+ double[][][] dataAfter1Original = new double[variables-1][][];
+ for (int v = 0; v < variables; v++) {
+ data[v] = rg.generateNormalData(samples, dimensionsPerVariable, 0, 1);
+ if (v > 0) {
+ dataAfter1Original[v-1] = data[v];
+ }
+ }
+ double[][] data1 = data[0];
+ // Now shuffle dataAfter1 inside the data variable:
+ RandomGenerator rg = new RandomGenerator();
+ int[][] newOrderings = rg.generateRandomPerturbations(
+ samples, 1);
+ for (int v = 1; v < variables; v++) {
+ data[v] = MatrixUtils.extractSelectedTimePointsReusingArrays(data[v], newOrderings[0]);
+ }
+ // And store the reverse time index mapping (of the original time index for each data point in the 1 surrodate):
+ int[][] reverseMapping = new int[1][samples];
+ for (int n = 0; n < samples; n++) {
+ reverseMapping[0][newOrderings[0][n]] = n;
+ }
+
+ long startTime = Calendar.getInstance().getTimeInMillis();
+ int[] dimensions = new int[variables];
+ for (int v = 0; v < variables; v++) {
+ dimensions[v] = dimensionsPerVariable;
+ }
+ int[] dimensionsAfter1 = new int[variables-1];
+ for (int v = 1; v < variables; v++) {
+ dimensionsAfter1[v-1] = dimensionsPerVariable;
+ }
+ KdTree kdTree = new KdTree(dimensions, data);
+ NearestNeighbourSearcher nnSearcherExceptVar1 =
+ NearestNeighbourSearcher.create(dimensionsAfter1, dataAfter1Original);
+ NearestNeighbourSearcher nnSearcherVar1 =
+ NearestNeighbourSearcher.create(data1);
+ long endTimeTree = Calendar.getInstance().getTimeInMillis();
+ System.out.printf("Additional criteria test: Tree of %d points constructed in : %.3f sec\n",
+ samples, ((double) (endTimeTree - startTime)/1000.0));
+
+ double[] r1 = {1.0};
+ boolean[] isWithinR = new boolean[samples];
+ int[] indicesWithinR = new int[samples+1];
+
+ for (int i = 0; i < r1.length; i++) {
+ for (int t = 0; t < samples; t++) {
+ // Establish our ground truth, including the shuffled variables after 1:
+ int count = kdTree.countPointsWithinR(t, r1[i], allowEqualsR);
+ // Alternative if we need to debug:
+ // Collection coll = kdTree.findPointsWithinR(t, r1[i], allowEqualsR);
+ // int count = coll.size();
+
+ // Now search in variable 1 first, then use that
+ // to help the others:
+ nnSearcherVar1.findPointsWithinR(t, r1[i], allowEqualsR,
+ isWithinR, indicesWithinR);
+ int countUsingVar1 = kdTree.countPointsWithinR(t, r1[i],
+ allowEqualsR, 0, isWithinR);
+ assertEquals(count, countUsingVar1);
+
+ // And then search the kdTree which does not have variable 1, unshuffled but providing the reverse mapping:
+ int countAdditionalCriteria =
+ nnSearcherExceptVar1.countPointsWithinR(newOrderings[0][t], r1[i],
+ allowEqualsR, isWithinR, reverseMapping[0]);
+ assertEquals(count, countAdditionalCriteria);
+
+ // Finally, reset our boolean array:
+ for (int t2 = 0; indicesWithinR[t2] != -1; t2++) {
+ isWithinR[indicesWithinR[t2]] = false;
+ }
+ }
+ }
+ }
+
public void testPreviouslyTestedVariableMethodMultipleRs() throws Exception {
for (int vars = 2; vars <= 3; vars++) {