diff --git a/src/Numerics/Statistics/ArrayStatistics.cs b/src/Numerics/Statistics/ArrayStatistics.cs index 92deba98..76e8856a 100644 --- a/src/Numerics/Statistics/ArrayStatistics.cs +++ b/src/Numerics/Statistics/ArrayStatistics.cs @@ -29,6 +29,7 @@ // using System; +using MathNet.Numerics.Properties; namespace MathNet.Numerics.Statistics { @@ -171,6 +172,56 @@ namespace MathNet.Numerics.Statistics return Math.Sqrt(PopulationVariance(population)); } + /// + /// Estimates the unbiased population covariance from the provided two sample arrays. + /// On a dataset of size N will use an N-1 normalizer (Bessel's correction). + /// Returns NaN if data has less than two entries or if any entry is NaN. + /// + /// First sample array. + /// Second sample array. + public static double Covariance(double[] samples1, double[] samples2) + { + if (samples1 == null) throw new ArgumentNullException("samples1"); + if (samples2 == null) throw new ArgumentNullException("samples2"); + if (samples1.Length != samples2.Length) throw new ArgumentException(Resources.ArgumentVectorsSameLength); + if (samples1.Length <= 1) return double.NaN; + + var mean1 = Mean(samples1); + var mean2 = Mean(samples2); + + var covariance = 0.0; + for (int i = 0; i < samples1.Length; i++) + { + covariance += (samples1[i] - mean1)*(samples2[i] - mean2); + } + return covariance/(samples1.Length - 1); + } + + /// + /// Evaluates the population covariance from the full population provided as two arrays. + /// On a dataset of size N will use an N normalizer and would thus be biased if applied to a subset. + /// Returns NaN if data is empty or if any entry is NaN. + /// + /// First sample array. + /// Second sample array. + public static double PopulationCovariance(double[] samples1, double[] samples2) + { + if (samples1 == null) throw new ArgumentNullException("samples1"); + if (samples2 == null) throw new ArgumentNullException("samples2"); + if (samples1.Length != samples2.Length) throw new ArgumentException(Resources.ArgumentVectorsSameLength); + if (samples1.Length == 0) return double.NaN; + + var mean1 = Mean(samples1); + var mean2 = Mean(samples2); + + var covariance = 0.0; + for (int i = 0; i < samples1.Length; i++) + { + covariance += (samples1[i] - mean1) * (samples2[i] - mean2); + } + return covariance/samples1.Length; + } + /// /// Returns the order statistic (order 1..N) from the unsorted data array. /// WARNING: Works inplace and can thus causes the data array to be reordered. diff --git a/src/Numerics/Statistics/Statistics.cs b/src/Numerics/Statistics/Statistics.cs index f9c80bc3..e91445a0 100644 --- a/src/Numerics/Statistics/Statistics.cs +++ b/src/Numerics/Statistics/Statistics.cs @@ -227,6 +227,68 @@ namespace MathNet.Numerics.Statistics return StreamingStatistics.PopulationStandardDeviation(population.Where(d => d.HasValue).Select(d => d.Value)); } + /// + /// Estimates the unbiased population covariance from the provided samples. + /// On a dataset of size N will use an N-1 normalizer (Bessel's correction). + /// Returns NaN if data has less than two entries or if any entry is NaN. + /// + /// A subset of samples, sampled from the full population. + /// A subset of samples, sampled from the full population. + public static double Covariance(this IEnumerable samples1, IEnumerable samples2) + { + var array1 = samples1 as double[]; + var array2 = samples2 as double[]; + return array1 != null && array2 != null + ? ArrayStatistics.Covariance(array1, array2) + : StreamingStatistics.Covariance(samples1, samples2); + } + + /// + /// Estimates the unbiased population covariance from the provided samples. + /// On a dataset of size N will use an N-1 normalizer (Bessel's correction). + /// Returns NaN if data has less than two entries or if any entry is NaN. + /// Null-entries are ignored. + /// + /// A subset of samples, sampled from the full population. + /// A subset of samples, sampled from the full population. + public static double Covariance(this IEnumerable samples1, IEnumerable samples2) + { + if (samples1 == null) throw new ArgumentNullException("samples1"); + if (samples2 == null) throw new ArgumentNullException("samples2"); + return StreamingStatistics.Covariance(samples1.Where(d => d.HasValue).Select(d => d.Value), samples2.Where(d => d.HasValue).Select(d => d.Value)); + } + + /// + /// Evaluates the population covariance from the provided full populations. + /// On a dataset of size N will use an N normalizer and would thus be biased if applied to a subset. + /// Returns NaN if data is empty or if any entry is NaN. + /// + /// The full population data. + /// The full population data. + public static double PopulationCovariance(this IEnumerable population1, IEnumerable population2) + { + var array1 = population1 as double[]; + var array2 = population2 as double[]; + return array1 != null && array2 != null + ? ArrayStatistics.PopulationCovariance(array1, array2) + : StreamingStatistics.PopulationCovariance(population1, population2); + } + + /// + /// Evaluates the population covariance from the provided full populations. + /// On a dataset of size N will use an N normalize and would thus be biased if applied to a subsetr. + /// Returns NaN if data is empty or if any entry is NaN. + /// Null-entries are ignored. + /// + /// The full population data. + /// The full population data. + public static double PopulationCovariance(this IEnumerable population1, IEnumerable population2) + { + if (population1 == null) throw new ArgumentNullException("population1"); + if (population2 == null) throw new ArgumentNullException("population2"); + return StreamingStatistics.PopulationCovariance(population1.Where(d => d.HasValue).Select(d => d.Value), population2.Where(d => d.HasValue).Select(d => d.Value)); + } + /// /// Estimates the sample median from the provided samples (R8). /// diff --git a/src/Numerics/Statistics/StreamingStatistics.cs b/src/Numerics/Statistics/StreamingStatistics.cs index 2170cb73..23e87116 100644 --- a/src/Numerics/Statistics/StreamingStatistics.cs +++ b/src/Numerics/Statistics/StreamingStatistics.cs @@ -30,6 +30,7 @@ using System; using System.Collections.Generic; +using MathNet.Numerics.Properties; namespace MathNet.Numerics.Statistics { @@ -193,5 +194,95 @@ namespace MathNet.Numerics.Statistics { return Math.Sqrt(PopulationVariance(population)); } + + /// + /// Estimates the unbiased population covariance from the provided two sample enumerable sequences, in a single pass without memoization. + /// On a dataset of size N will use an N-1 normalizer (Bessel's correction). + /// Returns NaN if data has less than two entries or if any entry is NaN. + /// + /// First sample stream. + /// Second sample stream. + public static double Covariance(IEnumerable samples1, IEnumerable samples2) + { + if (samples1 == null) throw new ArgumentNullException("samples1"); + if (samples2 == null) throw new ArgumentNullException("samples2"); + + // https://en.wikipedia.org/wiki/Algorithms_for_calculating_variance + + var n = 0; + var mean1 = 0.0; + var mean2 = 0.0; + var comoment = 0.0; + + using (var s1 = samples1.GetEnumerator()) + using (var s2 = samples2.GetEnumerator()) + { + while (s1.MoveNext()) + { + if (!s2.MoveNext()) + { + throw new ArgumentException(Resources.ArgumentVectorsSameLength); + } + + var mean2Prev = mean2; + n++; + mean1 += (s1.Current - mean1)/n; + mean2 += (s2.Current - mean2)/n; + comoment += (s1.Current - mean1)*(s2.Current - mean2Prev); + + } + if (s2.MoveNext()) + { + throw new ArgumentException(Resources.ArgumentVectorsSameLength); + } + } + + return n > 1 ? comoment/(n - 1) : double.NaN; + } + + /// + /// Evaluates the population covariance from the full population provided as two enumerable sequences, in a single pass without memoization. + /// On a dataset of size N will use an N normalizer and would thus be biased if applied to a subset. + /// Returns NaN if data is empty or if any entry is NaN. + /// + /// First sample stream. + /// Second sample stream. + public static double PopulationCovariance(IEnumerable samples1, IEnumerable samples2) + { + if (samples1 == null) throw new ArgumentNullException("samples1"); + if (samples2 == null) throw new ArgumentNullException("samples2"); + + // https://en.wikipedia.org/wiki/Algorithms_for_calculating_variance + + var n = 0; + var mean1 = 0.0; + var mean2 = 0.0; + var comoment = 0.0; + + using (var s1 = samples1.GetEnumerator()) + using (var s2 = samples2.GetEnumerator()) + { + while (s1.MoveNext()) + { + if (!s2.MoveNext()) + { + throw new ArgumentException(Resources.ArgumentVectorsSameLength); + } + + var mean2Prev = mean2; + n++; + mean1 += (s1.Current - mean1) / n; + mean2 += (s2.Current - mean2) / n; + comoment += (s1.Current - mean1) * (s2.Current - mean2Prev); + + } + if (s2.MoveNext()) + { + throw new ArgumentException(Resources.ArgumentVectorsSameLength); + } + } + + return comoment/n; + } } } diff --git a/src/UnitTests/StatisticsTests/CorrelationTests.cs b/src/UnitTests/StatisticsTests/CorrelationTests.cs index 504fed58..a333acb7 100644 --- a/src/UnitTests/StatisticsTests/CorrelationTests.cs +++ b/src/UnitTests/StatisticsTests/CorrelationTests.cs @@ -74,6 +74,21 @@ namespace MathNet.Numerics.UnitTests.StatisticsTests AssertHelpers.AlmostEqual(-0.029470861580726, corr, 13); } + /// + /// Pearson correlation test. + /// + [Test] + public void PearsonCorrelationConsistentWithCovariance() + { + var dataA = _data["lottery"].Data.Take(200); + var dataB = _data["lew"].Data.Take(200); + + var direct = Correlation.Pearson(dataA, dataB); + var covariance = dataA.Covariance(dataB)/(dataA.StandardDeviation()*dataB.StandardDeviation()); + + AssertHelpers.AlmostEqual(covariance, direct, 13); + } + /// /// Pearson correlation test fail. /// diff --git a/src/UnitTests/StatisticsTests/StatisticsTests.cs b/src/UnitTests/StatisticsTests/StatisticsTests.cs index 82592bef..a9f51ed3 100644 --- a/src/UnitTests/StatisticsTests/StatisticsTests.cs +++ b/src/UnitTests/StatisticsTests/StatisticsTests.cs @@ -69,6 +69,8 @@ namespace MathNet.Numerics.UnitTests.StatisticsTests Assert.Throws(() => Statistics.StandardDeviation(data)); Assert.Throws(() => Statistics.PopulationVariance(data)); Assert.Throws(() => Statistics.PopulationStandardDeviation(data)); + Assert.Throws(() => Statistics.Covariance(data, data)); + Assert.Throws(() => Statistics.PopulationCovariance(data, data)); Assert.Throws(() => SortedArrayStatistics.Minimum(data)); Assert.Throws(() => SortedArrayStatistics.Maximum(data)); @@ -91,6 +93,8 @@ namespace MathNet.Numerics.UnitTests.StatisticsTests Assert.Throws(() => ArrayStatistics.StandardDeviation(data)); Assert.Throws(() => ArrayStatistics.PopulationVariance(data)); Assert.Throws(() => ArrayStatistics.PopulationStandardDeviation(data)); + Assert.Throws(() => ArrayStatistics.Covariance(data, data)); + Assert.Throws(() => ArrayStatistics.PopulationCovariance(data, data)); Assert.Throws(() => ArrayStatistics.MedianInplace(data)); Assert.Throws(() => ArrayStatistics.QuantileInplace(data, 0.3)); @@ -101,6 +105,8 @@ namespace MathNet.Numerics.UnitTests.StatisticsTests Assert.Throws(() => StreamingStatistics.StandardDeviation(data)); Assert.Throws(() => StreamingStatistics.PopulationVariance(data)); Assert.Throws(() => StreamingStatistics.PopulationStandardDeviation(data)); + Assert.Throws(() => StreamingStatistics.Covariance(data, data)); + Assert.Throws(() => StreamingStatistics.PopulationCovariance(data, data)); } [Test] @@ -117,6 +123,8 @@ namespace MathNet.Numerics.UnitTests.StatisticsTests Assert.DoesNotThrow(() => Statistics.StandardDeviation(data)); Assert.DoesNotThrow(() => Statistics.PopulationVariance(data)); Assert.DoesNotThrow(() => Statistics.PopulationStandardDeviation(data)); + Assert.DoesNotThrow(() => Statistics.Covariance(data, data)); + Assert.DoesNotThrow(() => Statistics.PopulationCovariance(data, data)); Assert.DoesNotThrow(() => SortedArrayStatistics.Minimum(data)); Assert.DoesNotThrow(() => SortedArrayStatistics.Maximum(data)); @@ -139,6 +147,8 @@ namespace MathNet.Numerics.UnitTests.StatisticsTests Assert.DoesNotThrow(() => ArrayStatistics.StandardDeviation(data)); Assert.DoesNotThrow(() => ArrayStatistics.PopulationVariance(data)); Assert.DoesNotThrow(() => ArrayStatistics.PopulationStandardDeviation(data)); + Assert.DoesNotThrow(() => ArrayStatistics.Covariance(data, data)); + Assert.DoesNotThrow(() => ArrayStatistics.PopulationCovariance(data, data)); Assert.DoesNotThrow(() => ArrayStatistics.MedianInplace(data)); Assert.DoesNotThrow(() => ArrayStatistics.QuantileInplace(data, 0.3)); @@ -149,6 +159,8 @@ namespace MathNet.Numerics.UnitTests.StatisticsTests Assert.DoesNotThrow(() => StreamingStatistics.StandardDeviation(data)); Assert.DoesNotThrow(() => StreamingStatistics.PopulationVariance(data)); Assert.DoesNotThrow(() => StreamingStatistics.PopulationStandardDeviation(data)); + Assert.DoesNotThrow(() => StreamingStatistics.Covariance(data, data)); + Assert.DoesNotThrow(() => StreamingStatistics.PopulationCovariance(data, data)); } [TestCase("lottery")] @@ -563,6 +575,64 @@ namespace MathNet.Numerics.UnitTests.StatisticsTests AssertHelpers.AlmostEqual(2d, StreamingStatistics.StandardDeviation(gaussian.Samples().Take(10000)), 2); } + [TestCase("lottery")] + [TestCase("lew")] + [TestCase("mavro")] + [TestCase("michelso")] + [TestCase("numacc1")] + public void CovarianceConsistentWithVariance(string dataSet) + { + var data = _data[dataSet]; + AssertHelpers.AlmostEqual(Statistics.Variance(data.Data), Statistics.Covariance(data.Data, data.Data), 10); + AssertHelpers.AlmostEqual(ArrayStatistics.Variance(data.Data), ArrayStatistics.Covariance(data.Data, data.Data), 10); + AssertHelpers.AlmostEqual(StreamingStatistics.Variance(data.Data), StreamingStatistics.Covariance(data.Data, data.Data), 10); + } + + [TestCase("lottery")] + [TestCase("lew")] + [TestCase("mavro")] + [TestCase("michelso")] + [TestCase("numacc1")] + public void PopulationCovarianceConsistentWithPopulationVariance(string dataSet) + { + var data = _data[dataSet]; + AssertHelpers.AlmostEqual(Statistics.PopulationVariance(data.Data), Statistics.PopulationCovariance(data.Data, data.Data), 10); + AssertHelpers.AlmostEqual(ArrayStatistics.PopulationVariance(data.Data), ArrayStatistics.PopulationCovariance(data.Data, data.Data), 10); + AssertHelpers.AlmostEqual(StreamingStatistics.PopulationVariance(data.Data), StreamingStatistics.PopulationCovariance(data.Data, data.Data), 10); + } + + [Test] + public void CovarianceIsSymmetric() + { + var dataA = _data["lottery"].Data.Take(200); + var dataB = _data["lew"].Data.Take(200); + + AssertHelpers.AlmostEqual(Statistics.Covariance(dataA, dataB), Statistics.Covariance(dataB, dataA), 12); + AssertHelpers.AlmostEqual(StreamingStatistics.Covariance(dataA, dataB), StreamingStatistics.Covariance(dataB, dataA), 12); + AssertHelpers.AlmostEqual(ArrayStatistics.Covariance(dataA.ToArray(), dataB.ToArray()), ArrayStatistics.Covariance(dataB.ToArray(), dataA.ToArray()), 12); + + AssertHelpers.AlmostEqual(Statistics.PopulationCovariance(dataA, dataB), Statistics.PopulationCovariance(dataB, dataA), 12); + AssertHelpers.AlmostEqual(StreamingStatistics.PopulationCovariance(dataA, dataB), StreamingStatistics.PopulationCovariance(dataB, dataA), 12); + AssertHelpers.AlmostEqual(ArrayStatistics.PopulationCovariance(dataA.ToArray(), dataB.ToArray()), ArrayStatistics.PopulationCovariance(dataB.ToArray(), dataA.ToArray()), 12); + } + + [TestCase("lottery")] + [TestCase("lew")] + [TestCase("mavro")] + [TestCase("michelso")] + [TestCase("numacc1")] + public void ArrayStatisticsConsistentWithStreamimgStatistics(string dataSet) + { + var data = _data[dataSet]; + AssertHelpers.AlmostEqual(ArrayStatistics.Mean(data.Data), StreamingStatistics.Mean(data.Data), 10); + AssertHelpers.AlmostEqual(ArrayStatistics.Variance(data.Data), StreamingStatistics.Variance(data.Data), 10); + AssertHelpers.AlmostEqual(ArrayStatistics.StandardDeviation(data.Data), StreamingStatistics.StandardDeviation(data.Data), 10); + AssertHelpers.AlmostEqual(ArrayStatistics.PopulationVariance(data.Data), StreamingStatistics.PopulationVariance(data.Data), 10); + AssertHelpers.AlmostEqual(ArrayStatistics.PopulationStandardDeviation(data.Data), StreamingStatistics.PopulationStandardDeviation(data.Data), 10); + AssertHelpers.AlmostEqual(ArrayStatistics.Covariance(data.Data, data.Data), StreamingStatistics.Covariance(data.Data, data.Data), 10); + AssertHelpers.AlmostEqual(ArrayStatistics.PopulationCovariance(data.Data, data.Data), StreamingStatistics.PopulationCovariance(data.Data, data.Data), 10); + } + [Test] public void MinimumOfEmptyMustBeNaN() {