From ad6cbbebd82388b808ae59348c56db669d56257f Mon Sep 17 00:00:00 2001 From: Jurgen Van Gael Date: Mon, 26 Apr 2010 06:18:15 +0800 Subject: [PATCH] Added CumulativeDistribution for the StudentT distribution. Added unit tests for the StudentT distribution. Added sampling methods for the normal gamma distribution. --- .../Distributions/Continuous/Normal.cs | 14 ++- .../Distributions/Continuous/StudentT.cs | 23 +++- .../Distributions/Multivariate/NormalGamma.cs | 116 ++++++------------ .../CommonDistributionTests.cs | 67 +++++----- .../Continuous/StudentTTests.cs | 20 +++ .../Multivariate/NormalGammaTests.cs | 88 +++++++++++++ 6 files changed, 213 insertions(+), 115 deletions(-) diff --git a/src/Numerics/Distributions/Continuous/Normal.cs b/src/Numerics/Distributions/Continuous/Normal.cs index e8364fcd..344d58d4 100644 --- a/src/Numerics/Distributions/Continuous/Normal.cs +++ b/src/Numerics/Distributions/Continuous/Normal.cs @@ -328,6 +328,18 @@ namespace MathNet.Numerics.Distributions return DensityLn(_mean, _stdDev, x); } + /// + /// Computes the cumulative distribution function of the normal distribution. + /// + /// The mean of the normal distribution. + /// The standard deviation of the normal distribution. + /// The location at which to compute the cumulative density. + /// the cumulative density at . + internal static double CumulativeDistribution(double mean, double sdev, double x) + { + return 0.5 * (1.0 + SpecialFunctions.Erf((x - mean) / (sdev * Constants.Sqrt2))); + } + /// /// Computes the cumulative distribution function of the normal distribution. /// @@ -335,7 +347,7 @@ namespace MathNet.Numerics.Distributions /// the cumulative density at . public double CumulativeDistribution(double x) { - return 0.5 * (1.0 + SpecialFunctions.Erf((x - _mean) / (_stdDev * Math.Sqrt(2.0)))); + return CumulativeDistribution(_mean, _stdDev, x); } /// diff --git a/src/Numerics/Distributions/Continuous/StudentT.cs b/src/Numerics/Distributions/Continuous/StudentT.cs index 07d413b2..1a30c051 100644 --- a/src/Numerics/Distributions/Continuous/StudentT.cs +++ b/src/Numerics/Distributions/Continuous/StudentT.cs @@ -333,6 +333,7 @@ namespace MathNet.Numerics.Distributions /// the density at . public double Density(double x) { + // TODO JVG we can probably do a better job for Cauchy special case if (Double.IsPositiveInfinity(_dof)) { return Normal.Density(_location, Math.Sqrt(_scale), x); @@ -355,6 +356,7 @@ namespace MathNet.Numerics.Distributions /// the log density at . public double DensityLn(double x) { + // TODO JVG we can probably do a better job for Cauchy special case if (Double.IsPositiveInfinity(_dof)) { return Normal.DensityLn(_location, Math.Sqrt(_scale), x); @@ -377,8 +379,25 @@ namespace MathNet.Numerics.Distributions /// the cumulative density at . public double CumulativeDistribution(double x) { - throw new NotImplementedException(); - // TODO Jurgen: once this is implemented; enable the StudentT stuff in commondistributiontests. + // TODO JVG we can probably do a better job for Cauchy special case + if (Double.IsPositiveInfinity(_dof)) + { + return Normal.CumulativeDistribution(_location, _scale, x); + } + else + { + double k = (x - _location) / _scale; + double h = _dof / (_dof + k * k); + double ib = 0.5 * SpecialFunctions.BetaRegularized(_dof / 2.0, 0.5, h); + if (x <= _location) + { + return ib; + } + else + { + return 1.0 - ib; + } + } } /// diff --git a/src/Numerics/Distributions/Multivariate/NormalGamma.cs b/src/Numerics/Distributions/Multivariate/NormalGamma.cs index b7f70bfe..1da563ac 100644 --- a/src/Numerics/Distributions/Multivariate/NormalGamma.cs +++ b/src/Numerics/Distributions/Multivariate/NormalGamma.cs @@ -280,44 +280,6 @@ namespace MathNet.Numerics.Distributions /* - /// - /// The mode of the distribution. - /// - /// - public MeanPrecisionPair Mode - { - get - { - if (Double.IsPositiveInfinity(_precisionInvScale)) - { - return new MeanPrecisionPair(_meanLocation, _precisionShape); - } - else - { - return new MeanPrecisionPair(_meanLocation, _precisionShape / _precisionInvScale); - } - } - } - - /// - /// The median of the distribution. - /// - /// - public MeanPrecisionPair Median - { - get - { - if (Double.IsPositiveInfinity(_precisionInvScale)) - { - return new MeanPrecisionPair(_meanLocation, _precisionShape); - } - else - { - return new MeanPrecisionPair(_meanLocation, _precisionShape / _precisionInvScale); - } - } - } - /// /// Evaluates the probability density function for a NormalGamma distribution. /// @@ -382,40 +344,43 @@ namespace MathNet.Numerics.Distributions return (_precisionShape - 0.5) * System.Math.Log(prec) + _precisionShape * System.Math.Log(_precisionInvScale) + e - Math.Constants.LogSqrt2Pi - Math.SpecialFunctions.GammaLn(_precisionShape); } - } + }*/ /// - /// Samples a NormalGamma distributed random variable. + /// Generates a sample from the NormalGamma distribution. /// - /// A random number from this distribution. + /// a sample from the distribution. public MeanPrecisionPair Sample() { - return NormalGamma.Sample(RandomNumberGenerator, _meanLocation, _meanScale, _precisionShape, _precisionInvScale); + return NormalGamma.Sample(RandomSource, _meanLocation, _meanScale, _precisionShape, _precisionInvScale); } /// - /// Samples an array of NormalGamma distributed random variables. + /// Generates a sequence of samples from the NormalGamma distribution /// - /// The number of variables needed. - /// An array of random numbers from this distribution. - public MeanPrecisionPair[] Sample(int size) + /// a sequence of samples from the distribution. + public IEnumerable Samples() { - return NormalGamma.Sample(RandomNumberGenerator, size, _meanLocation, _meanScale, _precisionShape, _precisionInvScale); + while (true) + { + yield return NormalGamma.Sample(RandomSource, _meanLocation, _meanScale, _precisionShape, _precisionInvScale); + } } /// - /// Samples an array of NormalGamma distributed random variables. + /// Generates a sample from the NormalGamma distribution. /// /// The random number generator to use. /// The location of the mean. /// The scale of the mean. /// The shape of the precision. /// The inverse scale of the precision. + /// a sample from the distribution. public static MeanPrecisionPair Sample(System.Random rnd, double meanLocation, double meanScale, double precShape, double precInvScale) { - if (Control.CheckDistributionParameters) + if (Control.CheckDistributionParameters && !IsValidParameterSet(meanLocation, meanScale, precShape, precInvScale)) { - CheckParameters(meanScale, precShape, precInvScale); + throw new ArgumentOutOfRangeException(Resources.InvalidDistributionParameters); } MeanPrecisionPair mp = new MeanPrecisionPair(); @@ -437,61 +402,54 @@ namespace MathNet.Numerics.Distributions } else { - mp.Mean = Normal.Sample(rnd, meanLocation, System.Math.Sqrt(meanScale / mp.Precision)); + mp.Mean = Normal.Sample(rnd, meanLocation, System.Math.Sqrt(1.0 / (meanScale * mp.Precision))); } return mp; } /// - /// Samples an array of NormalGamma distributed random variables. + /// Generates a sequence of samples from the NormalGamma distribution /// /// The random number generator to use. - /// The number of variables needed. /// The location of the mean. /// The scale of the mean. /// The shape of the precision. /// The inverse scale of the precision. - public static MeanPrecisionPair[] Sample(System.Random rnd, int n, double meanLocation, double meanScale, double precShape, double precInvScale) + /// a sequence of samples from the distribution. + public static IEnumerable Samples(System.Random rnd, double meanLocation, double meanScale, double precShape, double precInvScale) { - if (Control.CheckDistributionParameters) + if (Control.CheckDistributionParameters && !IsValidParameterSet(meanLocation, meanScale, precShape, precInvScale)) { - CheckParameters(meanScale, precShape, precInvScale); + throw new ArgumentOutOfRangeException(Resources.InvalidDistributionParameters); } - // First sample all the precisions independently. - double[] precs = null; - if (Double.IsPositiveInfinity(precInvScale)) + while (true) { - precs = new double[n]; - for (int i = 0; i < n; i++) + MeanPrecisionPair mp = new MeanPrecisionPair(); + + // Sample the precision. + if (Double.IsPositiveInfinity(precInvScale)) { - precs[i] = precShape; + mp.Precision = precShape; + } + else + { + mp.Precision = Gamma.Sample(rnd, precShape, precInvScale); } - } - else - { - precs = Gamma.Sample(rnd, n, precShape, precInvScale); - } - - // Construct all the mean precision pairs. - MeanPrecisionPair[] arr = new MeanPrecisionPair[n]; - // Conditionally sample all the mean. - for (int i = 0; i < n; i++) - { - arr[i].Precision = precs[i]; + // Sample the mean. if (meanScale == 0.0) { - arr[i].Mean = meanLocation; + mp.Mean = meanLocation; } else { - arr[i].Mean = Normal.Sample(rnd, meanLocation, System.Math.Sqrt(meanScale / precs[i])); + mp.Mean = Normal.Sample(rnd, meanLocation, System.Math.Sqrt(1.0 / (meanScale * mp.Precision))); } - } - return arr; - }*/ + yield return mp; + } + } } } \ No newline at end of file diff --git a/src/UnitTests/DistributionTests/CommonDistributionTests.cs b/src/UnitTests/DistributionTests/CommonDistributionTests.cs index c3eebe91..a267efd4 100644 --- a/src/UnitTests/DistributionTests/CommonDistributionTests.cs +++ b/src/UnitTests/DistributionTests/CommonDistributionTests.cs @@ -36,15 +36,19 @@ namespace MathNet.Numerics.UnitTests.DistributionTests using MathNet.Numerics.Statistics; using MathNet.Numerics.Distributions; + /// + /// This class will perform various tests on discrete and continuous univariate distributions. The multivariate distributions + /// will implement these respective tests in their local unit test classes as they do not adhere to the same interfaces. + /// [TestFixture] public class CommonDistributionTests { // The number of samples we want. - private int numberOfTestSamples = 100000; + public static int NumberOfTestSamples = 10000000; // The accuracy of the histograms. - private double sampleAccuracy = 0.01; + public static double SampleAccuracy = 0.01; // The number of buckets to use to test against the cdf. - private int numberOfBuckets = 100; + public static int NumberOfBuckets = 100; // The list of discrete distributions which we test. private List discreteDistributions; // The list of continuous distributions which we test. @@ -66,7 +70,7 @@ namespace MathNet.Numerics.UnitTests.DistributionTests continuousDistributions.Add(new Normal(0.0, 1.0)); continuousDistributions.Add(new Weibull(1.0, 1.0)); continuousDistributions.Add(new LogNormal(1.0, 1.0)); - //continuousDistributions.Add(new StudentT(0.0, 1.0, 3.0)); + continuousDistributions.Add(new StudentT(0.0, 1.0, 5.0)); } [Test] @@ -114,11 +118,14 @@ namespace MathNet.Numerics.UnitTests.DistributionTests } } + /// + /// Test the method which samples only one variable at a time. + /// [Test] [MultipleAsserts] public void SampleFollowsCorrectDistribution() { - Random rnd = new MersenneTwister(); + Random rnd = new MersenneTwister(1); // The test samples from the distributions, builds a histogram and checks // whether the histogram follows the CDF. @@ -126,81 +133,75 @@ namespace MathNet.Numerics.UnitTests.DistributionTests { dd.RandomSource = rnd; - double[] samples = new double[numberOfTestSamples]; - for (int i = 0; i < numberOfTestSamples; i++) + double[] samples = new double[NumberOfTestSamples]; + for (int i = 0; i < NumberOfTestSamples; i++) { samples[i] = (double) dd.Sample(); } - var histogram = new Histogram(samples, numberOfBuckets); - for (int i = 0; i < numberOfBuckets; i++) - { - var bucket = histogram[i]; - double empiricalProbability = bucket.Count / (double)numberOfTestSamples; - double realProbability = dd.CumulativeDistribution(bucket.UpperBound) - - dd.CumulativeDistribution(bucket.LowerBound); - Assert.LessThan(Math.Abs(empiricalProbability - realProbability), sampleAccuracy, dd.ToString()); - } } foreach (var cd in continuousDistributions) { cd.RandomSource = rnd; - double[] samples = new double[numberOfTestSamples]; - for (int i = 0; i < numberOfTestSamples; i++) + double[] samples = new double[NumberOfTestSamples]; + for (int i = 0; i < NumberOfTestSamples; i++) { samples[i] = cd.Sample(); } - var histogram = new Histogram(samples, numberOfBuckets); - for (int i = 0; i < numberOfBuckets; i++) + var histogram = new Histogram(samples, NumberOfBuckets); + for (int i = 0; i < NumberOfBuckets; i++) { var bucket = histogram[i]; - double empiricalProbability = bucket.Count / (double)numberOfTestSamples; + double empiricalProbability = bucket.Count / (double)NumberOfTestSamples; double realProbability = cd.CumulativeDistribution(bucket.UpperBound) - cd.CumulativeDistribution(bucket.LowerBound); - Assert.LessThan(Math.Abs(empiricalProbability - realProbability), sampleAccuracy, cd.ToString()); + Assert.LessThan(Math.Abs(empiricalProbability - realProbability), SampleAccuracy, cd.ToString()); } } } + /// + /// Test the method which samples a sequence of variables. + /// [Test] [MultipleAsserts] public void SamplesFollowsCorrectDistribution() { - Random rnd = new MersenneTwister(); + Random rnd = new MersenneTwister(1); // The test samples from the distributions, builds a histogram and checks // whether the histogram follows the CDF. foreach (var dd in discreteDistributions) { dd.RandomSource = rnd; - var samples = dd.Samples().Take(numberOfTestSamples).Select(x => (double)x); + var samples = dd.Samples().Take(NumberOfTestSamples).Select(x => (double)x); - var histogram = new Histogram(samples, numberOfBuckets); - for (int i = 0; i < numberOfBuckets; i++) + var histogram = new Histogram(samples, NumberOfBuckets); + for (int i = 0; i < NumberOfBuckets; i++) { var bucket = histogram[i]; - double empiricalProbability = bucket.Count / (double)numberOfTestSamples; + double empiricalProbability = bucket.Count / (double)NumberOfTestSamples; double realProbability = dd.CumulativeDistribution(bucket.UpperBound) - dd.CumulativeDistribution(bucket.LowerBound); - Assert.LessThan(Math.Abs(empiricalProbability - realProbability), sampleAccuracy, dd.ToString()); + Assert.LessThan(Math.Abs(empiricalProbability - realProbability), SampleAccuracy, dd.ToString()); } } foreach (var cd in continuousDistributions) { cd.RandomSource = rnd; - var samples = cd.Samples().Take(numberOfTestSamples); + var samples = cd.Samples().Take(NumberOfTestSamples); - var histogram = new Histogram(samples, numberOfBuckets); - for (int i = 0; i < numberOfBuckets; i++) + var histogram = new Histogram(samples, NumberOfBuckets); + for (int i = 0; i < NumberOfBuckets; i++) { var bucket = histogram[i]; - double empiricalProbability = bucket.Count / (double)numberOfTestSamples; + double empiricalProbability = bucket.Count / (double)NumberOfTestSamples; double realProbability = cd.CumulativeDistribution(bucket.UpperBound) - cd.CumulativeDistribution(bucket.LowerBound); - Assert.LessThan(Math.Abs(empiricalProbability - realProbability), sampleAccuracy, cd.ToString()); + Assert.LessThan(Math.Abs(empiricalProbability - realProbability), SampleAccuracy, cd.ToString()); } } } diff --git a/src/UnitTests/DistributionTests/Continuous/StudentTTests.cs b/src/UnitTests/DistributionTests/Continuous/StudentTTests.cs index 7b4dabb5..be234788 100644 --- a/src/UnitTests/DistributionTests/Continuous/StudentTTests.cs +++ b/src/UnitTests/DistributionTests/Continuous/StudentTTests.cs @@ -344,5 +344,25 @@ namespace MathNet.Numerics.UnitTests.DistributionTests var ied = n.Samples(); var e = ied.Take(5).ToArray(); } + + [Test] + [Row(0.0, 1.0, 1.0, 0.0, 0.5)] + [Row(0.0, 1.0, 1.0, 1.0, 0.75)] + [Row(0.0, 1.0, 1.0, -1.0, 0.25)] + [Row(0.0, 1.0, 1.0, 2.0, 0.852416382349567)] + [Row(0.0, 1.0, 1.0, -2.0, 0.147583617650433)] + [Row(0.0, 1.0, 2.0, 0.0, 0.5)] + [Row(0.0, 1.0, 2.0, 1.0, 0.788675134594813)] + [Row(0.0, 1.0, 2.0, -1.0, 0.211324865405187)] + [Row(0.0, 1.0, 2.0, 2.0, 0.908248290463863)] + [Row(0.0, 1.0, 2.0, -2.0, 0.091751709536137)] + [Row(0.0, 1.0, Double.PositiveInfinity, 0.0, 0.5)] + [Row(0.0, 1.0, Double.PositiveInfinity, 1.0, 0.841344746068543)] + [Row(0.0, 1.0, Double.PositiveInfinity, 2.0, 0.977249868051821)] + public void ValidateCumulativeDistribution(double location, double scale, double dof, double x, double c) + { + var n = new StudentT(location, scale, dof); + AssertHelpers.AlmostEqual(c, n.CumulativeDistribution(x), 13); + } } } diff --git a/src/UnitTests/DistributionTests/Multivariate/NormalGammaTests.cs b/src/UnitTests/DistributionTests/Multivariate/NormalGammaTests.cs index db3a8883..bbb6605e 100644 --- a/src/UnitTests/DistributionTests/Multivariate/NormalGammaTests.cs +++ b/src/UnitTests/DistributionTests/Multivariate/NormalGammaTests.cs @@ -31,6 +31,8 @@ namespace MathNet.Numerics.UnitTests.DistributionTests using System; using System.Linq; using MbUnit.Framework; + using MathNet.Numerics.Random; + using MathNet.Numerics.Statistics; using MathNet.Numerics.Distributions; [TestFixture] @@ -199,5 +201,91 @@ namespace MathNet.Numerics.UnitTests.DistributionTests NormalGamma ng = new NormalGamma(0.0, 1.0, 1.0, 1.0); ng.RandomSource = new Random(); } + + /// + /// Test the method which samples one variable at a time. + /// + [Test] + public void SampleFollowsCorrectDistribution() + { + Random rnd = new MersenneTwister(); + var cd = new NormalGamma(1.0, 4.0, 3.0, 3.5); + + // Sample from the distribution. + MeanPrecisionPair[] samples = new MeanPrecisionPair[CommonDistributionTests.NumberOfTestSamples]; + for (int i = 0; i < CommonDistributionTests.NumberOfTestSamples; i++) + { + samples[i] = cd.Sample(); + } + + // Extract the mean and precisions. + var means = samples.Select(mp => mp.Mean); + var precs = samples.Select(mp => mp.Precision); + var meanMarginal = cd.MeanMarginal(); + var precMarginal = cd.PrecisionMarginal(); + + // Check the mean distribution. + var histogram = new Histogram(means, CommonDistributionTests.NumberOfBuckets); + for (int i = 0; i < CommonDistributionTests.NumberOfBuckets; i++) + { + var bucket = histogram[i]; + double empiricalProbability = bucket.Count / (double)CommonDistributionTests.NumberOfTestSamples; + double realProbability = meanMarginal.CumulativeDistribution(bucket.UpperBound) + - meanMarginal.CumulativeDistribution(bucket.LowerBound); + Assert.LessThan(Math.Abs(empiricalProbability - realProbability), CommonDistributionTests.SampleAccuracy, cd.ToString()); + } + + // Check the precision distribution. + histogram = new Histogram(precs, CommonDistributionTests.NumberOfBuckets); + for (int i = 0; i < CommonDistributionTests.NumberOfBuckets; i++) + { + var bucket = histogram[i]; + double empiricalProbability = bucket.Count / (double)CommonDistributionTests.NumberOfTestSamples; + double realProbability = precMarginal.CumulativeDistribution(bucket.UpperBound) + - precMarginal.CumulativeDistribution(bucket.LowerBound); + Assert.LessThan(Math.Abs(empiricalProbability - realProbability), CommonDistributionTests.SampleAccuracy, cd.ToString()); + } + } + + /// + /// Test the method which samples a sequence of variables. + /// + [Test] + public void SamplesFollowsCorrectDistribution() + { + Random rnd = new MersenneTwister(); + var cd = new NormalGamma(1.0, 4.0, 3.0, 3.5); + + // Sample from the distribution. + var samples = cd.Samples().Take(CommonDistributionTests.NumberOfTestSamples).ToArray(); + + // Extract the mean and precisions. + var means = samples.Select(mp => mp.Mean); + var precs = samples.Select(mp => mp.Precision); + var meanMarginal = cd.MeanMarginal(); + var precMarginal = cd.PrecisionMarginal(); + + // Check the mean distribution. + var histogram = new Histogram(means, CommonDistributionTests.NumberOfBuckets); + for (int i = 0; i < CommonDistributionTests.NumberOfBuckets; i++) + { + var bucket = histogram[i]; + double empiricalProbability = bucket.Count / (double)CommonDistributionTests.NumberOfTestSamples; + double realProbability = meanMarginal.CumulativeDistribution(bucket.UpperBound) + - meanMarginal.CumulativeDistribution(bucket.LowerBound); + Assert.LessThan(Math.Abs(empiricalProbability - realProbability), CommonDistributionTests.SampleAccuracy, cd.ToString()); + } + + // Check the precision distribution. + histogram = new Histogram(precs, CommonDistributionTests.NumberOfBuckets); + for (int i = 0; i < CommonDistributionTests.NumberOfBuckets; i++) + { + var bucket = histogram[i]; + double empiricalProbability = bucket.Count / (double)CommonDistributionTests.NumberOfTestSamples; + double realProbability = precMarginal.CumulativeDistribution(bucket.UpperBound) + - precMarginal.CumulativeDistribution(bucket.LowerBound); + Assert.LessThan(Math.Abs(empiricalProbability - realProbability), CommonDistributionTests.SampleAccuracy, cd.ToString()); + } + } } } \ No newline at end of file