mirror of https://github.com/jlizier/jidt
Implementing multiple data sets in the nearest neighbour searches, to facilitate dynamic correlation exclusion only for points within the same data set
This commit is contained in:
parent
a554361de9
commit
938dd7c4fb
|
|
@ -113,6 +113,20 @@ public class KdTree extends NearestNeighbourSearcher {
|
||||||
this(new int[] {data[0].length}, new double[][][] {data});
|
this(new int[] {data[0].length}, new double[][][] {data});
|
||||||
}
|
}
|
||||||
|
|
||||||
|
/**
|
||||||
|
* Construct the k-d tree from a set of double[][] data.
|
||||||
|
*
|
||||||
|
* @param data a double[][] 2D data set, first indexed
|
||||||
|
* by time, second index by variable number.
|
||||||
|
* @param observationSetIndices array indicating for each sample which
|
||||||
|
* observation set is came from (only used for dynamic correlation exclusion)
|
||||||
|
* @param observationTimePoints array indicating for each sample which
|
||||||
|
* time index it had in the observation set it came from
|
||||||
|
*/
|
||||||
|
public KdTree(double[][] data, int[] observationSetIndices, int[] observationTimePoints) {
|
||||||
|
this(new int[] {data[0].length}, new double[][][] {data}, observationSetIndices, observationTimePoints);
|
||||||
|
}
|
||||||
|
|
||||||
/**
|
/**
|
||||||
* Construct the k-d tree from a <b>set</b> of double[][] data,
|
* Construct the k-d tree from a <b>set</b> of double[][] data,
|
||||||
* considered jointly.
|
* considered jointly.
|
||||||
|
|
@ -126,7 +140,29 @@ public class KdTree extends NearestNeighbourSearcher {
|
||||||
* within this data set)
|
* within this data set)
|
||||||
*/
|
*/
|
||||||
public KdTree(int[] dimensions, double[][][] data) {
|
public KdTree(int[] dimensions, double[][][] data) {
|
||||||
|
this(dimensions, data, null, null);
|
||||||
|
}
|
||||||
|
|
||||||
|
/**
|
||||||
|
* Construct the k-d tree from a <b>set</b> of double[][] data,
|
||||||
|
* considered jointly.
|
||||||
|
*
|
||||||
|
* @param dimensions an array of dimensions for each
|
||||||
|
* of the 2D data sets.
|
||||||
|
* @param data an array of double[][] 2D data sets
|
||||||
|
* for each data[i]
|
||||||
|
* (where i is the main variable number within data, then
|
||||||
|
* after that the first index is sample number, second is dimension
|
||||||
|
* within this data set)
|
||||||
|
* @param observationSetIndices array indicating for each sample which
|
||||||
|
* observation set is came from (only used for dynamic correlation exclusion).
|
||||||
|
* null means only a single observation set used
|
||||||
|
* @param observationTimePoints array indicating for each sample which
|
||||||
|
* time index it had in the observation set it came from
|
||||||
|
* null means only a single observation set used
|
||||||
|
*/
|
||||||
|
public KdTree(int[] dimensions, double[][][] data, int[] observationSetIndices, int[] observationTimePoints) {
|
||||||
|
|
||||||
this.originalDataSets = data;
|
this.originalDataSets = data;
|
||||||
int numObservations = data[0].length;
|
int numObservations = data[0].length;
|
||||||
|
|
||||||
|
|
@ -183,6 +219,18 @@ public class KdTree extends NearestNeighbourSearcher {
|
||||||
rootNode = constructKdTree(0, 0, numObservations, masterSortedArrayIndices);
|
rootNode = constructKdTree(0, 0, numObservations, masterSortedArrayIndices);
|
||||||
// And destroy the temporary storage of sorted array indices:
|
// And destroy the temporary storage of sorted array indices:
|
||||||
masterSortedArrayIndices = null;
|
masterSortedArrayIndices = null;
|
||||||
|
|
||||||
|
if (observationSetIndices == null) {
|
||||||
|
// observationSetIndices and observationTimePoints are
|
||||||
|
// not supplied, so by default we assume only the one observation set:
|
||||||
|
observationSetIndices = new int[numObservations];
|
||||||
|
observationTimePoints = new int[numObservations];
|
||||||
|
for (int n = 0; n < numObservations; n++) {
|
||||||
|
observationTimePoints[n] = n;
|
||||||
|
}
|
||||||
|
}
|
||||||
|
this.observationSetIndices = observationSetIndices;
|
||||||
|
this.observationTimePoints = observationTimePoints;
|
||||||
}
|
}
|
||||||
|
|
||||||
/**
|
/**
|
||||||
|
|
@ -667,7 +715,7 @@ public class KdTree extends NearestNeighbourSearcher {
|
||||||
* @param sampleIndex sample index in the data to find a nearest neighbour
|
* @param sampleIndex sample index in the data to find a nearest neighbour
|
||||||
* for
|
* for
|
||||||
* @param dynCorrExclTime size of dynamic correlation exclusion time window
|
* @param dynCorrExclTime size of dynamic correlation exclusion time window
|
||||||
* on either side of sampleIndex. 0 means exclude only sampleIndex itself.
|
* on either side of sampleIndex in the same observation set. 0 means exclude only sampleIndex itself.
|
||||||
* @param node node to start searching from in the kd-tree. Cannot be null
|
* @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 level which level we're currently at in the tree
|
||||||
* @param currentKBest a PriorityQueue of NeighbourNodeData objects
|
* @param currentKBest a PriorityQueue of NeighbourNodeData objects
|
||||||
|
|
@ -700,8 +748,8 @@ public class KdTree extends NearestNeighbourSearcher {
|
||||||
// (will not throw an Exception if the PQ is empty)
|
// (will not throw an Exception if the PQ is empty)
|
||||||
NeighbourNodeData furthestCached = currentKBest.peek();
|
NeighbourNodeData furthestCached = currentKBest.peek();
|
||||||
|
|
||||||
if (((node.indexOfThisPoint - sampleIndex > dynCorrExclTime)
|
if (((observationSetIndices[node.indexOfThisPoint] != observationSetIndices[sampleIndex])
|
||||||
|| (node.indexOfThisPoint - sampleIndex < -dynCorrExclTime)) &&
|
|| (Math.abs(observationTimePoints[node.indexOfThisPoint] - observationTimePoints[sampleIndex]) > dynCorrExclTime)) &&
|
||||||
((currentKBest.size() < K) || (absDistOnThisDim < furthestCached.distance))) {
|
((currentKBest.size() < K) || (absDistOnThisDim < furthestCached.distance))) {
|
||||||
// Preliminary check says we need to compute the full distance
|
// Preliminary check says we need to compute the full distance
|
||||||
// to use or at least to check if it should be
|
// to use or at least to check if it should be
|
||||||
|
|
@ -1073,7 +1121,8 @@ public class KdTree extends NearestNeighbourSearcher {
|
||||||
absDistOnThisDim = distOnThisDim * distOnThisDim;
|
absDistOnThisDim = distOnThisDim * distOnThisDim;
|
||||||
}
|
}
|
||||||
|
|
||||||
if ((Math.abs(node.indexOfThisPoint - sampleIndex) > dynCorrExclTime) &&
|
if (((observationSetIndices[node.indexOfThisPoint] != observationSetIndices[sampleIndex])
|
||||||
|
|| (Math.abs(observationTimePoints[node.indexOfThisPoint] - observationTimePoints[sampleIndex]) > dynCorrExclTime)) &&
|
||||||
((absDistOnThisDim < r) ||
|
((absDistOnThisDim < r) ||
|
||||||
( allowEqualToR && (absDistOnThisDim == r)))) {
|
( allowEqualToR && (absDistOnThisDim == r)))) {
|
||||||
// Preliminary check says we need to compute the full distance
|
// Preliminary check says we need to compute the full distance
|
||||||
|
|
@ -1591,7 +1640,8 @@ public class KdTree extends NearestNeighbourSearcher {
|
||||||
absDistOnThisDim = distOnThisDim * distOnThisDim;
|
absDistOnThisDim = distOnThisDim * distOnThisDim;
|
||||||
}
|
}
|
||||||
|
|
||||||
if ((Math.abs(node.indexOfThisPoint - sampleIndex) > dynCorrExclTime) &&
|
if (((observationSetIndices[node.indexOfThisPoint] != observationSetIndices[sampleIndex]) ||
|
||||||
|
(Math.abs(observationTimePoints[node.indexOfThisPoint] - observationTimePoints[sampleIndex]) > dynCorrExclTime)) &&
|
||||||
((absDistOnThisDim < r) ||
|
((absDistOnThisDim < r) ||
|
||||||
( allowEqualToR && (absDistOnThisDim == r)))) {
|
( allowEqualToR && (absDistOnThisDim == r)))) {
|
||||||
// Preliminary check says we need to compute the full distance
|
// Preliminary check says we need to compute the full distance
|
||||||
|
|
@ -1738,7 +1788,8 @@ public class KdTree extends NearestNeighbourSearcher {
|
||||||
absDistOnThisDim = distOnThisDim * distOnThisDim;
|
absDistOnThisDim = distOnThisDim * distOnThisDim;
|
||||||
}
|
}
|
||||||
|
|
||||||
if ((Math.abs(node.indexOfThisPoint - sampleIndex) > dynCorrExclTime) &&
|
if (((observationSetIndices[node.indexOfThisPoint] != observationSetIndices[sampleIndex]) ||
|
||||||
|
(Math.abs(observationTimePoints[node.indexOfThisPoint] - observationTimePoints[sampleIndex]) > dynCorrExclTime)) &&
|
||||||
((absDistOnThisDim < r) ||
|
((absDistOnThisDim < r) ||
|
||||||
( allowEqualToR && (absDistOnThisDim == r)))) {
|
( allowEqualToR && (absDistOnThisDim == r)))) {
|
||||||
// Preliminary check says we need to compute the full distance
|
// Preliminary check says we need to compute the full distance
|
||||||
|
|
@ -1912,7 +1963,8 @@ public class KdTree extends NearestNeighbourSearcher {
|
||||||
}
|
}
|
||||||
|
|
||||||
if (testResultsForGivenVariable[node.indexOfThisPoint] &&
|
if (testResultsForGivenVariable[node.indexOfThisPoint] &&
|
||||||
(Math.abs(node.indexOfThisPoint - sampleIndex) > dynCorrExclTime) &&
|
((observationSetIndices[node.indexOfThisPoint] != observationSetIndices[sampleIndex]) ||
|
||||||
|
(Math.abs(observationTimePoints[node.indexOfThisPoint] - observationTimePoints[sampleIndex]) > dynCorrExclTime)) &&
|
||||||
((absDistOnThisDim < r) ||
|
((absDistOnThisDim < r) ||
|
||||||
( allowEqualToR && (absDistOnThisDim == r)))) {
|
( allowEqualToR && (absDistOnThisDim == r)))) {
|
||||||
// Preliminary check says we need to compute the full distance
|
// Preliminary check says we need to compute the full distance
|
||||||
|
|
@ -2574,7 +2626,8 @@ public class KdTree extends NearestNeighbourSearcher {
|
||||||
absDistOnThisDim = distOnThisDim * distOnThisDim;
|
absDistOnThisDim = distOnThisDim * distOnThisDim;
|
||||||
}
|
}
|
||||||
|
|
||||||
if ((Math.abs(node.indexOfThisPoint - sampleIndex) > dynCorrExclTime) &&
|
if (((observationSetIndices[node.indexOfThisPoint] != observationSetIndices[sampleIndex]) ||
|
||||||
|
(Math.abs(observationTimePoints[node.indexOfThisPoint] - observationTimePoints[sampleIndex]) > dynCorrExclTime)) &&
|
||||||
((absDistOnThisDim < rs[variableNumber]) ||
|
((absDistOnThisDim < rs[variableNumber]) ||
|
||||||
( allowEqualToR && (absDistOnThisDim == rs[variableNumber])))) {
|
( allowEqualToR && (absDistOnThisDim == rs[variableNumber])))) {
|
||||||
// Preliminary check says we need to compute the full distance
|
// Preliminary check says we need to compute the full distance
|
||||||
|
|
|
||||||
|
|
@ -40,6 +40,19 @@ public abstract class NearestNeighbourSearcher {
|
||||||
*/
|
*/
|
||||||
protected int normTypeToUse = EuclideanUtils.NORM_MAX_NORM;
|
protected int normTypeToUse = EuclideanUtils.NORM_MAX_NORM;
|
||||||
|
|
||||||
|
/**
|
||||||
|
* observationSetIndices is an array indicating for each sample which
|
||||||
|
* observation set is came from (only used for dynamic correlation exclusion).
|
||||||
|
* null means only a single observation set used
|
||||||
|
*/
|
||||||
|
protected int[] observationSetIndices;
|
||||||
|
/**
|
||||||
|
* observationTimePoints is an array indicating for each sample which
|
||||||
|
* time index it had in the observation set it came from
|
||||||
|
* null means only a single observation set used
|
||||||
|
*/
|
||||||
|
protected int[] observationTimePoints;
|
||||||
|
|
||||||
/**
|
/**
|
||||||
* Factory method to construct the searcher from a set of double[][] data.
|
* Factory method to construct the searcher from a set of double[][] data.
|
||||||
* This will return a {@link KdTree} or if the data is univaraite
|
* This will return a {@link KdTree} or if the data is univaraite
|
||||||
|
|
@ -51,17 +64,37 @@ public abstract class NearestNeighbourSearcher {
|
||||||
public static NearestNeighbourSearcher create(double[][] data)
|
public static NearestNeighbourSearcher create(double[][] data)
|
||||||
throws Exception {
|
throws Exception {
|
||||||
|
|
||||||
|
return create(data, null, null);
|
||||||
|
}
|
||||||
|
|
||||||
|
/**
|
||||||
|
* 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.
|
||||||
|
* @param observationSetIndices array indicating for each sample which
|
||||||
|
* observation set is came from (only used for dynamic correlation exclusion)
|
||||||
|
* @param observationTimePoints array indicating for each sample which
|
||||||
|
* time index it had in the observation set it came from
|
||||||
|
*/
|
||||||
|
public static NearestNeighbourSearcher create(double[][] data,
|
||||||
|
int[] observationSetIndices, int[] observationTimePoints)
|
||||||
|
throws Exception {
|
||||||
|
|
||||||
if ((data == null) || (data[0].length == 0)) {
|
if ((data == null) || (data[0].length == 0)) {
|
||||||
// We have null data:
|
// We have null data:
|
||||||
return null;
|
return null;
|
||||||
} else if (data[0].length == 1) {
|
} else if (data[0].length == 1) {
|
||||||
// We have univariate data:
|
// We have univariate data:
|
||||||
return new UnivariateNearestNeighbourSearcher(MatrixUtils.selectColumn(data, 0));
|
return new UnivariateNearestNeighbourSearcher(MatrixUtils.selectColumn(data, 0),
|
||||||
|
observationSetIndices, observationTimePoints);
|
||||||
} else {
|
} else {
|
||||||
return new KdTree(data);
|
return new KdTree(data, observationSetIndices, observationTimePoints);
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
/**
|
/**
|
||||||
* Factory method to construct the searcher from a set of double[][][] data.
|
* Factory method to construct the searcher from a set of double[][][] data.
|
||||||
*
|
*
|
||||||
|
|
@ -71,11 +104,29 @@ public abstract class NearestNeighbourSearcher {
|
||||||
public static NearestNeighbourSearcher create(int[] dimensions, double[][][] data)
|
public static NearestNeighbourSearcher create(int[] dimensions, double[][][] data)
|
||||||
throws Exception {
|
throws Exception {
|
||||||
|
|
||||||
|
return create(dimensions, data, null, null);
|
||||||
|
}
|
||||||
|
|
||||||
|
/**
|
||||||
|
* Factory method to construct the searcher from a set of double[][][] data.
|
||||||
|
*
|
||||||
|
* @param data an array of double[][] 2D data sets, first indexed
|
||||||
|
* by time, second index by variable number.
|
||||||
|
* @param observationSetIndices array indicating for each sample which
|
||||||
|
* observation set is came from (only used for dynamic correlation exclusion)
|
||||||
|
* @param observationTimePoints array indicating for each sample which
|
||||||
|
* time index it had in the observation set it came from
|
||||||
|
*/
|
||||||
|
public static NearestNeighbourSearcher create(int[] dimensions, double[][][] data,
|
||||||
|
int[] observationSetIndices, int[] observationTimePoints)
|
||||||
|
throws Exception {
|
||||||
|
|
||||||
if ((dimensions.length == 1) && (dimensions[0] == 1)) {
|
if ((dimensions.length == 1) && (dimensions[0] == 1)) {
|
||||||
// We have univariate data:
|
// We have univariate data:
|
||||||
return new UnivariateNearestNeighbourSearcher(MatrixUtils.selectColumn(data[0], 0));
|
return new UnivariateNearestNeighbourSearcher(MatrixUtils.selectColumn(data[0], 0),
|
||||||
|
observationSetIndices, observationTimePoints);
|
||||||
} else {
|
} else {
|
||||||
return new KdTree(dimensions, data);
|
return new KdTree(dimensions, data, observationSetIndices, observationTimePoints);
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -63,15 +63,25 @@ public class UnivariateNearestNeighbourSearcher extends NearestNeighbourSearcher
|
||||||
protected double[] sortedValues = null;
|
protected double[] sortedValues = null;
|
||||||
|
|
||||||
public UnivariateNearestNeighbourSearcher(double[][] data) throws Exception {
|
public UnivariateNearestNeighbourSearcher(double[][] data) throws Exception {
|
||||||
|
this(data, null, null);
|
||||||
|
}
|
||||||
|
|
||||||
|
public UnivariateNearestNeighbourSearcher(double[][] data,
|
||||||
|
int[] observationSetIndices, int[] observationTimePoints) throws Exception {
|
||||||
// Ideally we would not call the constructor until after the following check,
|
// Ideally we would not call the constructor until after the following check,
|
||||||
// but the constructor must come first in Java.
|
// but the constructor must come first in Java.
|
||||||
this(MatrixUtils.selectColumn(data, 0));
|
this(MatrixUtils.selectColumn(data, 0), observationSetIndices, observationTimePoints);
|
||||||
if (data[0].length != 1) {
|
if (data[0].length != 1) {
|
||||||
throw new Exception("Cannot define UnivariateNearestNeighbourSearcher for multivariate data");
|
throw new Exception("Cannot define UnivariateNearestNeighbourSearcher for multivariate data");
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
public UnivariateNearestNeighbourSearcher(double[] data) throws Exception {
|
public UnivariateNearestNeighbourSearcher(double[] data) throws Exception {
|
||||||
|
this(data, null, null);
|
||||||
|
}
|
||||||
|
|
||||||
|
public UnivariateNearestNeighbourSearcher(double[] data,
|
||||||
|
int[] observationSetIndices, int[] observationTimePoints) throws Exception {
|
||||||
this.originalDataSet = data;
|
this.originalDataSet = data;
|
||||||
numObservations = data.length;
|
numObservations = data.length;
|
||||||
if (numObservations <= 1) {
|
if (numObservations <= 1) {
|
||||||
|
|
@ -102,6 +112,18 @@ public class UnivariateNearestNeighbourSearcher extends NearestNeighbourSearcher
|
||||||
for (int i = 0; i < numObservations; i++) {
|
for (int i = 0; i < numObservations; i++) {
|
||||||
sortedValues[i] = originalDataSet[sortedArrayIndices[i]];
|
sortedValues[i] = originalDataSet[sortedArrayIndices[i]];
|
||||||
}
|
}
|
||||||
|
|
||||||
|
if (observationSetIndices == null) {
|
||||||
|
// observationSetIndices and observationTimePoints are
|
||||||
|
// not supplied, so by default we assume only the one observation set:
|
||||||
|
observationSetIndices = new int[numObservations];
|
||||||
|
observationTimePoints = new int[numObservations];
|
||||||
|
for (int n = 0; n < numObservations; n++) {
|
||||||
|
observationTimePoints[n] = n;
|
||||||
|
}
|
||||||
|
}
|
||||||
|
this.observationSetIndices = observationSetIndices;
|
||||||
|
this.observationTimePoints = observationTimePoints;
|
||||||
}
|
}
|
||||||
|
|
||||||
/**
|
/**
|
||||||
|
|
@ -252,7 +274,8 @@ public class UnivariateNearestNeighbourSearcher extends NearestNeighbourSearcher
|
||||||
lowerCandidate >= 0;
|
lowerCandidate >= 0;
|
||||||
lowerCandidate--) {
|
lowerCandidate--) {
|
||||||
indexOfLowerCandidate = sortedArrayIndices[lowerCandidate];
|
indexOfLowerCandidate = sortedArrayIndices[lowerCandidate];
|
||||||
if (Math.abs(sampleIndex - indexOfLowerCandidate) > dynCorrExclTime) {
|
if ((observationSetIndices[sampleIndex] != observationSetIndices[indexOfLowerCandidate]) ||
|
||||||
|
(Math.abs(observationTimePoints[sampleIndex] - observationTimePoints[indexOfLowerCandidate]) > dynCorrExclTime)) {
|
||||||
// This sample is outside the dynamic correlation exclusion window
|
// This sample is outside the dynamic correlation exclusion window
|
||||||
break;
|
break;
|
||||||
}
|
}
|
||||||
|
|
@ -266,7 +289,8 @@ public class UnivariateNearestNeighbourSearcher extends NearestNeighbourSearcher
|
||||||
upperCandidate <= numObservations - 1;
|
upperCandidate <= numObservations - 1;
|
||||||
upperCandidate++) {
|
upperCandidate++) {
|
||||||
indexOfUpperCandidate = sortedArrayIndices[upperCandidate];
|
indexOfUpperCandidate = sortedArrayIndices[upperCandidate];
|
||||||
if (Math.abs(sampleIndex - indexOfUpperCandidate) > dynCorrExclTime) {
|
if ((observationSetIndices[sampleIndex] != observationSetIndices[indexOfUpperCandidate]) ||
|
||||||
|
(Math.abs(observationTimePoints[sampleIndex] - observationTimePoints[indexOfUpperCandidate]) > dynCorrExclTime)) {
|
||||||
// This sample is outside the dynamic correlation exclusion window
|
// This sample is outside the dynamic correlation exclusion window
|
||||||
break;
|
break;
|
||||||
}
|
}
|
||||||
|
|
@ -301,7 +325,8 @@ public class UnivariateNearestNeighbourSearcher extends NearestNeighbourSearcher
|
||||||
upperCandidate <= numObservations - 1;
|
upperCandidate <= numObservations - 1;
|
||||||
upperCandidate++) {
|
upperCandidate++) {
|
||||||
indexOfUpperCandidate = sortedArrayIndices[upperCandidate];
|
indexOfUpperCandidate = sortedArrayIndices[upperCandidate];
|
||||||
if (Math.abs(sampleIndex - indexOfUpperCandidate) > dynCorrExclTime) {
|
if ((observationSetIndices[sampleIndex] != observationSetIndices[indexOfUpperCandidate]) ||
|
||||||
|
(Math.abs(observationTimePoints[sampleIndex] - observationTimePoints[indexOfUpperCandidate]) > dynCorrExclTime)) {
|
||||||
// This sample is outside the dynamic correlation exclusion window
|
// This sample is outside the dynamic correlation exclusion window
|
||||||
break;
|
break;
|
||||||
}
|
}
|
||||||
|
|
@ -317,7 +342,8 @@ public class UnivariateNearestNeighbourSearcher extends NearestNeighbourSearcher
|
||||||
lowerCandidate >= 0;
|
lowerCandidate >= 0;
|
||||||
lowerCandidate--) {
|
lowerCandidate--) {
|
||||||
indexOfLowerCandidate = sortedArrayIndices[lowerCandidate];
|
indexOfLowerCandidate = sortedArrayIndices[lowerCandidate];
|
||||||
if (Math.abs(sampleIndex - indexOfLowerCandidate) > dynCorrExclTime) {
|
if ((observationSetIndices[sampleIndex] != observationSetIndices[indexOfLowerCandidate]) ||
|
||||||
|
(Math.abs(observationTimePoints[sampleIndex] - observationTimePoints[indexOfLowerCandidate]) > dynCorrExclTime)) {
|
||||||
// This sample is outside the dynamic correlation exclusion window
|
// This sample is outside the dynamic correlation exclusion window
|
||||||
break;
|
break;
|
||||||
}
|
}
|
||||||
|
|
@ -474,7 +500,8 @@ public class UnivariateNearestNeighbourSearcher extends NearestNeighbourSearcher
|
||||||
int indexInSortedArray = indicesInSortedArray[sampleIndex];
|
int indexInSortedArray = indicesInSortedArray[sampleIndex];
|
||||||
// Check the points with smaller data values first:
|
// Check the points with smaller data values first:
|
||||||
for (int i = indexInSortedArray - 1; i >= 0; i--) {
|
for (int i = indexInSortedArray - 1; i >= 0; i--) {
|
||||||
if (Math.abs(sampleIndex - sortedArrayIndices[i]) <= dynCorrExclTime) {
|
if ((observationSetIndices[sampleIndex] == observationSetIndices[sortedArrayIndices[i]]) &&
|
||||||
|
(Math.abs(observationTimePoints[sampleIndex] - observationTimePoints[sortedArrayIndices[i]]) <= dynCorrExclTime)) {
|
||||||
// Can't count this point, but keep checking:
|
// Can't count this point, but keep checking:
|
||||||
continue;
|
continue;
|
||||||
}
|
}
|
||||||
|
|
@ -490,7 +517,8 @@ public class UnivariateNearestNeighbourSearcher extends NearestNeighbourSearcher
|
||||||
}
|
}
|
||||||
// Next check the points with larger data values:
|
// Next check the points with larger data values:
|
||||||
for (int i = indexInSortedArray + 1; i < numObservations; i++) {
|
for (int i = indexInSortedArray + 1; i < numObservations; i++) {
|
||||||
if (Math.abs(sampleIndex - sortedArrayIndices[i]) <= dynCorrExclTime) {
|
if ((observationSetIndices[sampleIndex] == observationSetIndices[sortedArrayIndices[i]]) &&
|
||||||
|
(Math.abs(observationTimePoints[sampleIndex] - observationTimePoints[sortedArrayIndices[i]]) <= dynCorrExclTime)) {
|
||||||
// Can't count this point, but keep checking:
|
// Can't count this point, but keep checking:
|
||||||
continue;
|
continue;
|
||||||
}
|
}
|
||||||
|
|
@ -747,7 +775,8 @@ public class UnivariateNearestNeighbourSearcher extends NearestNeighbourSearcher
|
||||||
int indexInSortedArray = indicesInSortedArray[sampleIndex];
|
int indexInSortedArray = indicesInSortedArray[sampleIndex];
|
||||||
// Check the points with smaller data values first:
|
// Check the points with smaller data values first:
|
||||||
for (int i = indexInSortedArray - 1; i >= 0; i--) {
|
for (int i = indexInSortedArray - 1; i >= 0; i--) {
|
||||||
if (Math.abs(sampleIndex - sortedArrayIndices[i]) <= dynCorrExclTime) {
|
if ((observationSetIndices[sampleIndex] == observationSetIndices[sortedArrayIndices[i]]) &&
|
||||||
|
(Math.abs(observationTimePoints[sampleIndex] - observationTimePoints[sortedArrayIndices[i]]) <= dynCorrExclTime)) {
|
||||||
// Can't count this point, but keep checking:
|
// Can't count this point, but keep checking:
|
||||||
continue;
|
continue;
|
||||||
}
|
}
|
||||||
|
|
@ -764,7 +793,8 @@ public class UnivariateNearestNeighbourSearcher extends NearestNeighbourSearcher
|
||||||
}
|
}
|
||||||
// Next check the points with larger data values:
|
// Next check the points with larger data values:
|
||||||
for (int i = indexInSortedArray + 1; i < numObservations; i++) {
|
for (int i = indexInSortedArray + 1; i < numObservations; i++) {
|
||||||
if (Math.abs(sampleIndex - sortedArrayIndices[i]) <= dynCorrExclTime) {
|
if ((observationSetIndices[sampleIndex] == observationSetIndices[sortedArrayIndices[i]]) &&
|
||||||
|
(Math.abs(observationTimePoints[sampleIndex] - observationTimePoints[sortedArrayIndices[i]]) <= dynCorrExclTime)) {
|
||||||
// Can't count this point, but keep checking:
|
// Can't count this point, but keep checking:
|
||||||
continue;
|
continue;
|
||||||
}
|
}
|
||||||
|
|
@ -793,7 +823,8 @@ public class UnivariateNearestNeighbourSearcher extends NearestNeighbourSearcher
|
||||||
int indexInSortedArray = indicesInSortedArray[sampleIndex];
|
int indexInSortedArray = indicesInSortedArray[sampleIndex];
|
||||||
// Check the points with smaller data values first:
|
// Check the points with smaller data values first:
|
||||||
for (int i = indexInSortedArray - 1; i >= 0; i--) {
|
for (int i = indexInSortedArray - 1; i >= 0; i--) {
|
||||||
if (Math.abs(sampleIndex - sortedArrayIndices[i]) <= dynCorrExclTime) {
|
if ((observationSetIndices[sampleIndex] == observationSetIndices[sortedArrayIndices[i]]) &&
|
||||||
|
(Math.abs(observationTimePoints[sampleIndex] - observationTimePoints[sortedArrayIndices[i]]) <= dynCorrExclTime)) {
|
||||||
// Can't count this point, but keep checking:
|
// Can't count this point, but keep checking:
|
||||||
continue;
|
continue;
|
||||||
}
|
}
|
||||||
|
|
@ -812,7 +843,8 @@ public class UnivariateNearestNeighbourSearcher extends NearestNeighbourSearcher
|
||||||
}
|
}
|
||||||
// Next check the points with larger data values:
|
// Next check the points with larger data values:
|
||||||
for (int i = indexInSortedArray + 1; i < numObservations; i++) {
|
for (int i = indexInSortedArray + 1; i < numObservations; i++) {
|
||||||
if (Math.abs(sampleIndex - sortedArrayIndices[i]) <= dynCorrExclTime) {
|
if ((observationSetIndices[sampleIndex] == observationSetIndices[sortedArrayIndices[i]]) &&
|
||||||
|
(Math.abs(observationTimePoints[sampleIndex] - observationTimePoints[sortedArrayIndices[i]]) <= dynCorrExclTime)) {
|
||||||
// Can't count this point, but keep checking:
|
// Can't count this point, but keep checking:
|
||||||
continue;
|
continue;
|
||||||
}
|
}
|
||||||
|
|
@ -868,7 +900,8 @@ public class UnivariateNearestNeighbourSearcher extends NearestNeighbourSearcher
|
||||||
// Check the points with smaller data values first:
|
// Check the points with smaller data values first:
|
||||||
for (int i = indexInSortedArray - 1; i >= 0; i--) {
|
for (int i = indexInSortedArray - 1; i >= 0; i--) {
|
||||||
if (!additionalCriteria[sortedArrayIndices[i]] ||
|
if (!additionalCriteria[sortedArrayIndices[i]] ||
|
||||||
(Math.abs(sampleIndex - sortedArrayIndices[i]) <= dynCorrExclTime)) {
|
((observationSetIndices[sampleIndex] == observationSetIndices[sortedArrayIndices[i]]) &&
|
||||||
|
(Math.abs(observationTimePoints[sampleIndex] - observationTimePoints[sortedArrayIndices[i]]) <= dynCorrExclTime))) {
|
||||||
// Can't count this point, but keep checking:
|
// Can't count this point, but keep checking:
|
||||||
continue;
|
continue;
|
||||||
}
|
}
|
||||||
|
|
@ -889,7 +922,8 @@ public class UnivariateNearestNeighbourSearcher extends NearestNeighbourSearcher
|
||||||
// Next check the points with larger data values:
|
// Next check the points with larger data values:
|
||||||
for (int i = indexInSortedArray + 1; i < numObservations; i++) {
|
for (int i = indexInSortedArray + 1; i < numObservations; i++) {
|
||||||
if (!additionalCriteria[sortedArrayIndices[i]] ||
|
if (!additionalCriteria[sortedArrayIndices[i]] ||
|
||||||
(Math.abs(sampleIndex - sortedArrayIndices[i]) <= dynCorrExclTime)) {
|
((observationSetIndices[sampleIndex] == observationSetIndices[sortedArrayIndices[i]]) &&
|
||||||
|
(Math.abs(observationTimePoints[sampleIndex] - observationTimePoints[sortedArrayIndices[i]]) <= dynCorrExclTime))) {
|
||||||
// Can't count this point, but keep checking:
|
// Can't count this point, but keep checking:
|
||||||
continue;
|
continue;
|
||||||
}
|
}
|
||||||
|
|
|
||||||
|
|
@ -422,6 +422,72 @@ public class KdTreeTest extends TestCase {
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
public void testFindKNearestNeighboursWithExclusionWindowAndMultiDataSets() throws Exception {
|
||||||
|
int dimension = 4;
|
||||||
|
int numSamplesPerSet = 400;
|
||||||
|
int numSets = 5;
|
||||||
|
int exclusionWindow = 50;
|
||||||
|
|
||||||
|
for (int K = 1; K < 5; K++) {
|
||||||
|
double[][] data = rg.generateNormalData(numSamplesPerSet*numSets, dimension, 0, 1);
|
||||||
|
int[] obsSetIds = new int[numSamplesPerSet*numSets];
|
||||||
|
int[] timeIndicesInSets = new int[numSamplesPerSet*numSets];
|
||||||
|
int ti = 0;
|
||||||
|
for (int s = 0; s < numSets; s++) {
|
||||||
|
for (int s2 = 0; s2 < numSamplesPerSet; s2++) {
|
||||||
|
obsSetIds[ti] = s;
|
||||||
|
timeIndicesInSets[ti] = s2;
|
||||||
|
ti++;
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
long startTime = Calendar.getInstance().getTimeInMillis();
|
||||||
|
KdTree kdTree = new KdTree(data, obsSetIds, timeIndicesInSets);
|
||||||
|
long endTimeTree = Calendar.getInstance().getTimeInMillis();
|
||||||
|
System.out.printf("Tree of %d points for %d NNs constructed in: %.3f sec\n",
|
||||||
|
data.length, K, ((double) (endTimeTree - startTime)/1000.0));
|
||||||
|
|
||||||
|
EuclideanUtils normCalculator = new EuclideanUtils(EuclideanUtils.NORM_MAX_NORM);
|
||||||
|
startTime = Calendar.getInstance().getTimeInMillis();
|
||||||
|
for (int t = 0; t < data.length; t++) {
|
||||||
|
PriorityQueue<NeighbourNodeData> nnPQ =
|
||||||
|
kdTree.findKNearestNeighbours(K, t, exclusionWindow);
|
||||||
|
assertTrue(nnPQ.size() == K);
|
||||||
|
// Now find the K nearest neighbours with a naive all-pairs comparison
|
||||||
|
double[][] distancesAndIndices = new double[data.length][2];
|
||||||
|
for (int t2 = 0; t2 < data.length; t2++) {
|
||||||
|
boolean inDifferentSets = ((t / numSamplesPerSet) != (t2 / numSamplesPerSet));
|
||||||
|
if (inDifferentSets || (Math.abs(t2 - t) > exclusionWindow)) {
|
||||||
|
distancesAndIndices[t2][0] = normCalculator.norm(data[t], data[t2]);
|
||||||
|
} else {
|
||||||
|
distancesAndIndices[t2][0] = Double.POSITIVE_INFINITY;
|
||||||
|
}
|
||||||
|
distancesAndIndices[t2][1] = t2;
|
||||||
|
}
|
||||||
|
int[] timeStepsOfKthMins =
|
||||||
|
MatrixUtils.kMinIndices(distancesAndIndices, 0, K);
|
||||||
|
for (int i = 0; i < K; i++) {
|
||||||
|
// Check that the ith nearest neighbour matches for each method.
|
||||||
|
// Note that these two method provide a different sorting order
|
||||||
|
NeighbourNodeData nnData = nnPQ.poll();
|
||||||
|
if (timeStepsOfKthMins[K - 1 - i] != nnData.sampleIndex) {
|
||||||
|
// We have an error:
|
||||||
|
System.out.printf("Erroneous match between indices %d (expected) " +
|
||||||
|
" and %d\n", timeStepsOfKthMins[K - 1 - i], nnData.sampleIndex);
|
||||||
|
}
|
||||||
|
assertEquals(timeStepsOfKthMins[K - 1 - i], nnData.sampleIndex);
|
||||||
|
// And check that none of the nearest neighbours were within the window
|
||||||
|
// and from the same data set
|
||||||
|
boolean inDifferentSets = ((t / numSamplesPerSet) != (nnData.sampleIndex / numSamplesPerSet));
|
||||||
|
assertTrue(inDifferentSets || (Math.abs(nnData.sampleIndex - t) > exclusionWindow));
|
||||||
|
}
|
||||||
|
}
|
||||||
|
long endTimeValidate = Calendar.getInstance().getTimeInMillis();
|
||||||
|
System.out.printf("All %d nearest neighbours found in: %.3f sec\n",
|
||||||
|
K, ((double) (endTimeValidate - startTime)/1000.0));
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
public void testFindKNearestNeighboursForSeparateArrays() throws Exception {
|
public void testFindKNearestNeighboursForSeparateArrays() throws Exception {
|
||||||
int variables = 3;
|
int variables = 3;
|
||||||
int dimensionsPerVariable = 3;
|
int dimensionsPerVariable = 3;
|
||||||
|
|
|
||||||
|
|
@ -625,4 +625,66 @@ public class UnivariateNearestNeighbourTest extends TestCase {
|
||||||
K, ((double) (endTimeValidate - nnEndTime)/1000.0));
|
K, ((double) (endTimeValidate - nnEndTime)/1000.0));
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
public void testFindKNearestNeighboursWithExclusionWindowAndMultiDataSets() throws Exception {
|
||||||
|
int numSamplesPerSet = 400;
|
||||||
|
int numSets = 5;
|
||||||
|
int exclusionWindow = 50;
|
||||||
|
|
||||||
|
for (int K = 1; K < 5; K++) {
|
||||||
|
double[] data = rg.generateNormalData(numSamplesPerSet*numSets, 0, 1);
|
||||||
|
int[] obsSetIds = new int[numSamplesPerSet*numSets];
|
||||||
|
int[] timeIndicesInSets = new int[numSamplesPerSet*numSets];
|
||||||
|
int ti = 0;
|
||||||
|
for (int s = 0; s < numSets; s++) {
|
||||||
|
for (int s2 = 0; s2 < numSamplesPerSet; s2++) {
|
||||||
|
obsSetIds[ti] = s;
|
||||||
|
timeIndicesInSets[ti] = s2;
|
||||||
|
ti++;
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
long startTime = Calendar.getInstance().getTimeInMillis();
|
||||||
|
UnivariateNearestNeighbourSearcher searcher = new UnivariateNearestNeighbourSearcher(data, obsSetIds, timeIndicesInSets);
|
||||||
|
long endTimeTree = Calendar.getInstance().getTimeInMillis();
|
||||||
|
System.out.printf("Searcher of %d points for %d NNs constructed in: %.3f sec\n",
|
||||||
|
data.length, K, ((double) (endTimeTree - startTime)/1000.0));
|
||||||
|
|
||||||
|
startTime = Calendar.getInstance().getTimeInMillis();
|
||||||
|
for (int t = 0; t < data.length; t++) {
|
||||||
|
PriorityQueue<NeighbourNodeData> nnPQ =
|
||||||
|
searcher.findKNearestNeighbours(K, t, exclusionWindow);
|
||||||
|
assertTrue(nnPQ.size() == K);
|
||||||
|
// Now find the K nearest neighbours with a naive all-pairs comparison
|
||||||
|
double[][] distancesAndIndices = new double[data.length][2];
|
||||||
|
for (int t2 = 0; t2 < data.length; t2++) {
|
||||||
|
boolean inDifferentSets = ((t / numSamplesPerSet) != (t2 / numSamplesPerSet));
|
||||||
|
// If we weren't catering for different sample sets, it would run like this (so if you run this it will lead to an error):
|
||||||
|
// if ((Math.abs(t2 - t) > exclusionWindow)) {
|
||||||
|
if (inDifferentSets || (Math.abs(t2 - t) > exclusionWindow)) {
|
||||||
|
distancesAndIndices[t2][0] = Math.abs(data[t] - data[t2]);
|
||||||
|
} else {
|
||||||
|
distancesAndIndices[t2][0] = Double.POSITIVE_INFINITY;
|
||||||
|
}
|
||||||
|
distancesAndIndices[t2][1] = t2;
|
||||||
|
}
|
||||||
|
int[] timeStepsOfKthMins =
|
||||||
|
MatrixUtils.kMinIndices(distancesAndIndices, 0, K);
|
||||||
|
for (int i = 0; i < K; i++) {
|
||||||
|
// Check that the ith nearest neighbour matches for each method.
|
||||||
|
// Note that these two method provide a different sorting order
|
||||||
|
NeighbourNodeData nnData = nnPQ.poll();
|
||||||
|
if (timeStepsOfKthMins[K - 1 - i] != nnData.sampleIndex) {
|
||||||
|
// We have an error:
|
||||||
|
System.out.printf("Erroneous match between indices %d (expected) " +
|
||||||
|
" and %d\n", timeStepsOfKthMins[K - 1 - i], nnData.sampleIndex);
|
||||||
|
}
|
||||||
|
assertEquals(timeStepsOfKthMins[K - 1 - i], nnData.sampleIndex);
|
||||||
|
}
|
||||||
|
}
|
||||||
|
long endTimeValidate = Calendar.getInstance().getTimeInMillis();
|
||||||
|
System.out.printf("All %d nearest neighbours found in: %.3f sec\n",
|
||||||
|
K, ((double) (endTimeValidate - startTime)/1000.0));
|
||||||
|
}
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|
|
||||||
Loading…
Reference in New Issue