Browse Source

Distributions: Categorical impl of mean, variance, stddev and median #187

provider
Christoph Ruegg 13 years ago
parent
commit
2689e5c989
  1. 26
      src/Numerics/Distributions/Categorical.cs
  2. 14
      src/UnitTests/DistributionTests/Discrete/CategoricalTests.cs

26
src/Numerics/Distributions/Categorical.cs

@ -48,6 +48,9 @@ namespace MathNet.Numerics.Distributions
/// does not have to be normalized and sum to 1. The reason is that some vectors can't be exactly normalized /// does not have to be normalized and sum to 1. The reason is that some vectors can't be exactly normalized
/// to sum to 1 in floating point representation. /// to sum to 1 in floating point representation.
/// </remarks> /// </remarks>
/// <remarks>
/// Support: 0..k where k = length(probability mass array)-1
/// </remarks>
public class Categorical : IDiscreteDistribution public class Categorical : IDiscreteDistribution
{ {
System.Random _random; System.Random _random;
@ -215,14 +218,14 @@ namespace MathNet.Numerics.Distributions
{ {
get get
{ {
// Mean = E[X] = Sum(x * p(x), x=0..N-1)
// where f(x) is the probability mass function, and N is the number of categories.
var sum = 0.0; var sum = 0.0;
// Mean = Sum(x * f(x), x=0..N)
// where f(x) is the probability mass function, and N is the maximum value.
for (int i = 0; i < _pmfNormalized.Length; i++) for (int i = 0; i < _pmfNormalized.Length; i++)
{ {
sum += i * _pmfNormalized[i]; sum += i * _pmfNormalized[i];
} }
return sum; return sum;
} }
} }
@ -232,7 +235,7 @@ namespace MathNet.Numerics.Distributions
/// </summary> /// </summary>
public double StdDev public double StdDev
{ {
get { return _pmfNormalized.StandardDeviation(); } get { return Math.Sqrt(Variance); }
} }
/// <summary> /// <summary>
@ -240,7 +243,18 @@ namespace MathNet.Numerics.Distributions
/// </summary> /// </summary>
public double Variance public double Variance
{ {
get { return _pmfNormalized.Variance(); } get
{
// Variance = E[(X-E[X])^2] = E[X^2] - (E[X])^2 = Sum(p(x) * (x - E[X])^2), x=0..N-1)
var m = Mean;
var sum = 0.0;
for (int i = 0; i < _pmfNormalized.Length; i++)
{
var r = i - m;
sum += r*r*_pmfNormalized[i];
}
return sum;
}
} }
/// <summary> /// <summary>
@ -290,7 +304,7 @@ namespace MathNet.Numerics.Distributions
/// </summary> /// </summary>
public double Median public double Median
{ {
get { return _pmfNormalized.Median(); } get { return InverseCumulativeDistribution(0.5); }
} }
/// <summary> /// <summary>

14
src/UnitTests/DistributionTests/Discrete/CategoricalTests.cs

@ -171,8 +171,7 @@ namespace MathNet.Numerics.UnitTests.DistributionTests.Discrete
[TestCase(new double[] { 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1 }, 5)] [TestCase(new double[] { 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1 }, 5)]
public void ValidateMean(double[] p, double mean) public void ValidateMean(double[] p, double mean)
{ {
var n = new Categorical(p); Assert.That(new Categorical(p).Mean, Is.EqualTo(mean).Within(1e-14));
AssertHelpers.AlmostEqual(mean, n.Mean, 14);
} }
/// <summary> /// <summary>
@ -187,8 +186,7 @@ namespace MathNet.Numerics.UnitTests.DistributionTests.Discrete
[TestCase(new double[] { 1, 0, 1 }, 1)] [TestCase(new double[] { 1, 0, 1 }, 1)]
public void ValidateStdDev(double[] p, double stdDev) public void ValidateStdDev(double[] p, double stdDev)
{ {
var n = new Categorical(p); Assert.That(new Categorical(p).StdDev, Is.EqualTo(stdDev).Within(1e-14));
AssertHelpers.AlmostEqual(stdDev, n.StdDev, 14);
} }
/// <summary> /// <summary>
@ -203,8 +201,7 @@ namespace MathNet.Numerics.UnitTests.DistributionTests.Discrete
[TestCase(new double[] { 1, 0, 1 }, 1)] [TestCase(new double[] { 1, 0, 1 }, 1)]
public void ValidateVariance(double[] p, double variance) public void ValidateVariance(double[] p, double variance)
{ {
var n = new Categorical(p); Assert.That(new Categorical(p).Variance, Is.EqualTo(variance).Within(1e-14));
AssertHelpers.AlmostEqual(variance, n.Variance, 14);
} }
/// <summary> /// <summary>
@ -219,13 +216,12 @@ namespace MathNet.Numerics.UnitTests.DistributionTests.Discrete
// P(X < 5) = (1+2+6+3+2)/29 = 14/29 < 0.5. // P(X < 5) = (1+2+6+3+2)/29 = 14/29 < 0.5.
// P(X <= 5) = 19/29 > 0.5. // P(X <= 5) = 19/29 > 0.5.
[TestCase(new double[] { 1, 2, 6, 3, 2, 5, 1, 1, 0, 1, 7 }, 5)] [TestCase(new double[] { 1, 2, 6, 3, 2, 5, 1, 1, 0, 1, 7 }, 5)]
// TODO: Find out the expected behavour of Median in ambiguous cases like the following: // TODO: Find out the expected behavior of Median in ambiguous cases like the following:
//[TestCase(new double[] { 0, 0.5, 0.5 }, ???)] //[TestCase(new double[] { 0, 0.5, 0.5 }, ???)]
//[TestCase(new double[] { 1, 0, 1 }, ???)] //[TestCase(new double[] { 1, 0, 1 }, ???)]
public void ValidateMedian(double[] p, int median) public void ValidateMedian(double[] p, int median)
{ {
var n = new Categorical(p); Assert.That(new Categorical(p).Median, Is.EqualTo(median).Within(1e-14));
Assert.AreEqual(median, n.Median);
} }
/// <summary> /// <summary>

Loading…
Cancel
Save