From ce1a3b4b3199c055d2e22694e0e6cdafa5d2d821 Mon Sep 17 00:00:00 2001 From: "joseph.lizier" Date: Sat, 6 Jul 2013 13:54:02 +0000 Subject: [PATCH] Extended fix to local cond MI values for handling linear redundancies amongst variables properly in ConditionalMICalculatorMultivariateGaussian - this resolves Issue 16 --- ...ualInfoCalculatorMultiVariateGaussian.java | 53 +++++++++++--- .../infodynamics/utils/MatrixUtils.java | 73 +++++++++++++++++++ 2 files changed, 114 insertions(+), 12 deletions(-) diff --git a/java/source/infodynamics/measures/continuous/gaussian/ConditionalMutualInfoCalculatorMultiVariateGaussian.java b/java/source/infodynamics/measures/continuous/gaussian/ConditionalMutualInfoCalculatorMultiVariateGaussian.java index 796ba92..b2710ca 100755 --- a/java/source/infodynamics/measures/continuous/gaussian/ConditionalMutualInfoCalculatorMultiVariateGaussian.java +++ b/java/source/infodynamics/measures/continuous/gaussian/ConditionalMutualInfoCalculatorMultiVariateGaussian.java @@ -577,13 +577,31 @@ public class ConditionalMutualInfoCalculatorMultiVariateGaussian // Simple way: // detCovariance = MatrixUtils.determinantSymmPosDefMatrix(covariance); // Using cached Cholesky decomposition: - detCovariance = MatrixUtils.determinantViaCholeskyResult(L); - if (detCovariance == 0) { - throw new Exception("Covariance matrix is not positive definite"); - } - det1cCovariance = MatrixUtils.determinantViaCholeskyResult(L_1c); - det2cCovariance = MatrixUtils.determinantViaCholeskyResult(L_2c); + + // Should always have a valid L_cc (since we can reduce it down to one variable: detccCovariance = MatrixUtils.determinantViaCholeskyResult(L_cc); + if (L_1c == null) { + // Variable 1 is fully linearly redundant with conditional, so + // we will have zero conditional MI: + return MatrixUtils.constantArray(newVar2Obs.length, 0); + } else { + det1cCovariance = MatrixUtils.determinantViaCholeskyResult(L_1c); + if (L_2c == null) { + // Variable 2 is fully linearly redundant with conditional, so + // we will have zero conditional MI: + return MatrixUtils.constantArray(newVar2Obs.length, 0); + } else { + det2cCovariance = MatrixUtils.determinantViaCholeskyResult(L_2c); + if (L == null) { + // There is a linear dependence amongst variables 1 and 2 given the + // conditional which did not exist for either with the conditional alone, + // so conditional MI diverges: + return MatrixUtils.constantArray(newVar2Obs.length, Double.POSITIVE_INFINITY); + } else { + detCovariance = MatrixUtils.determinantViaCholeskyResult(L); + } + } + } } // Now we are clear to take the matrix inverse (via Cholesky decomposition, @@ -597,23 +615,34 @@ public class ConditionalMutualInfoCalculatorMultiVariateGaussian double[][] invCondCovariance = MatrixUtils.solveViaCholeskyResult(L_cc, MatrixUtils.identityMatrix(L_cc.length)); - double[] var1Means = MatrixUtils.select(means, 0, dimensionsVar1); - double[] var2Means = MatrixUtils.select(means, dimensionsVar1, dimensionsVar2); - double[] condMeans = MatrixUtils.select(means, dimensionsVar1 + dimensionsVar2, dimensionsCond); + // Now, only use the means from the subsets of linearly independent variables: + // double[] var1Means = MatrixUtils.select(means, 0, dimensionsVar1); + double[] var1Means = MatrixUtils.select(means, var1IndicesInCovariance); + // double[] var2Means = MatrixUtils.select(means, dimensionsVar1, dimensionsVar2); + double[] var2Means = MatrixUtils.select(means, var2IndicesInCovariance); + // double[] condMeans = MatrixUtils.select(means, dimensionsVar1 + dimensionsVar2, dimensionsCond); + double[] condMeans = MatrixUtils.select(means, condIndicesInCovariance); int lengthOfReturnArray; lengthOfReturnArray = newVar2Obs.length; double[] localValues = new double[lengthOfReturnArray]; + int[] var2IndicesSelected = MatrixUtils.subtract(var2IndicesInCovariance, dimensionsVar1); + int[] condIndicesSelected = MatrixUtils.subtract(condIndicesInCovariance, dimensionsVar1 + dimensionsVar2); for (int t = 0; t < newVar2Obs.length; t++) { double[] var1DeviationsFromMean = - MatrixUtils.subtract(newVar1Obs[t], + MatrixUtils.subtract( + MatrixUtils.select(newVar1Obs[t], var1IndicesInCovariance), var1Means); double[] var2DeviationsFromMean = - MatrixUtils.subtract(newVar2Obs[t], var2Means); + MatrixUtils.subtract( + MatrixUtils.select(newVar2Obs[t], var2IndicesSelected), + var2Means); double[] condDeviationsFromMean = - MatrixUtils.subtract(newCondObs[t], condMeans); + MatrixUtils.subtract( + MatrixUtils.select(newCondObs[t], condIndicesSelected), + condMeans); double[] var1CondDeviationsFromMean = MatrixUtils.append(var1DeviationsFromMean, condDeviationsFromMean); diff --git a/java/source/infodynamics/utils/MatrixUtils.java b/java/source/infodynamics/utils/MatrixUtils.java index 26f73c2..91b0c41 100755 --- a/java/source/infodynamics/utils/MatrixUtils.java +++ b/java/source/infodynamics/utils/MatrixUtils.java @@ -44,6 +44,19 @@ public class MatrixUtils { return array; } + /** + * Return an array with the given value at every index + * + * @param length length of array + * @param value value for every element of the arry + * @return + */ + public static double[] constantArray(int length, double value) { + double[] array = new double[length]; + Arrays.fill(array, value); + return array; + } + public static double sum(double[] input) { double total = 0; for (int i = 0; i < input.length; i++) { @@ -599,6 +612,21 @@ public class MatrixUtils { return returnValues; } + /** + * Subtracts a constant value from all items in an array + * + * @param array + * @param value + * @return array - constant value + */ + public static double[] subtract(double[] array, double value) throws Exception { + double[] returnValues = new double[array.length]; + for (int i = 0; i < returnValues.length; i++) { + returnValues[i] = array[i] - value; + } + return returnValues; + } + /** * Subtracts second array from the first, overwriting the * values in first @@ -641,6 +669,21 @@ public class MatrixUtils { return returnValues; } + /** + * Subtracts a constant value from all items in an array + * + * @param array + * @param value + * @return array - constant value + */ + public static int[] subtract(int[] array, int value) throws Exception { + int[] returnValues = new int[array.length]; + for (int i = 0; i < returnValues.length; i++) { + returnValues[i] = array[i] - value; + } + return returnValues; + } + /** * Return the matrix product A x B * @@ -1150,6 +1193,21 @@ public class MatrixUtils { return returnData; } + /** + * Select out part of an array. + * + * @param data + * @param indices which array indices to pull out + * @return + */ + public static double[] select(double[] data, int[] indices) { + double[] returnData = new double[indices.length]; + for (int i = 0; i < indices.length; i++) { + returnData[i] = data[indices[i]]; + } + return returnData; + } + /** * Select out part of an array. * @@ -1164,6 +1222,21 @@ public class MatrixUtils { return returnData; } + /** + * Select out part of an array. + * + * @param data + * @param indices which array indices to pull out + * @return + */ + public static int[] select(int[] data, int[] indices) { + int[] returnData = new int[indices.length]; + for (int i = 0; i < indices.length; i++) { + returnData[i] = data[indices[i]]; + } + return returnData; + } + public static int[] selectColumn(int matrix[][], int columnNo) { int[] column = new int[matrix.length]; for (int r = 0; r < matrix.length; r++) {