View Javadoc
1   /*
2    * Licensed to the Apache Software Foundation (ASF) under one or more
3    * contributor license agreements.  See the NOTICE file distributed with
4    * this work for additional information regarding copyright ownership.
5    * The ASF licenses this file to You under the Apache License, Version 2.0
6    * (the "License"); you may not use this file except in compliance with
7    * the License.  You may obtain a copy of the License at
8    *
9    *      https://www.apache.org/licenses/LICENSE-2.0
10   *
11   * Unless required by applicable law or agreed to in writing, software
12   * distributed under the License is distributed on an "AS IS" BASIS,
13   * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
14   * See the License for the specific language governing permissions and
15   * limitations under the License.
16   */
17  package org.apache.commons.statistics.distribution;
18  
19  import java.util.stream.IntStream;
20  import org.junit.jupiter.api.Assertions;
21  import org.junit.jupiter.api.Test;
22  
23  /**
24   * Test cases for AbstractDiscreteDistribution default implementations.
25   */
26  class AbstractDiscreteDistributionTest {
27      private final DiceDistribution diceDistribution = new DiceDistribution();
28  
29      @Test
30      void testInverseCumulativeProbabilityMethod() {
31          // This must be consistent with its own cumulative probability
32          final double[] p = IntStream.rangeClosed(1, 6).mapToDouble(diceDistribution::cumulativeProbability).toArray();
33          Assertions.assertEquals(1.0, p[5], "Incorrect cumulative probability at upper bound");
34          Assertions.assertEquals(1, diceDistribution.inverseCumulativeProbability(0));
35          for (int i = 0; i < 6; i++) {
36              final int x = i + 1;
37              Assertions.assertEquals(x, diceDistribution.inverseCumulativeProbability(Math.nextDown(p[i])));
38              Assertions.assertEquals(x, diceDistribution.inverseCumulativeProbability(p[i]));
39              if (x < 6) {
40                  Assertions.assertEquals(x + 1, diceDistribution.inverseCumulativeProbability(Math.nextUp(p[i])));
41              }
42          }
43      }
44  
45      @Test
46      void testInverseSurvivalProbabilityMethod() {
47          // This must be consistent with its own survival probability
48          final double[] p = IntStream.rangeClosed(1, 6).mapToDouble(diceDistribution::survivalProbability).toArray();
49          Assertions.assertEquals(0.0, p[5], "Incorrect survival probability at upper bound");
50          Assertions.assertEquals(1, diceDistribution.inverseSurvivalProbability(1));
51          for (int i = 0; i < 6; i++) {
52              final int x = i + 1;
53              Assertions.assertEquals(x, diceDistribution.inverseSurvivalProbability(Math.nextUp(p[i])));
54              Assertions.assertEquals(x, diceDistribution.inverseSurvivalProbability(p[i]));
55              if (x < 6) {
56                  Assertions.assertEquals(x + 1, diceDistribution.inverseSurvivalProbability(Math.nextDown(p[i])));
57              }
58          }
59      }
60  
61      @Test
62      void testCumulativeProbabilitiesSingleArguments() {
63          final double p = diceDistribution.probability(1);
64          for (int i = 1; i <= 6; i++) {
65              Assertions.assertEquals(p * i,
66                      diceDistribution.cumulativeProbability(i), Math.ulp(p * i));
67          }
68          Assertions.assertEquals(0.0, diceDistribution.cumulativeProbability(0));
69          Assertions.assertEquals(1.0, diceDistribution.cumulativeProbability(7));
70      }
71  
72      @Test
73      void testProbabilitiesRangeArguments() {
74          final double p = diceDistribution.probability(1);
75          int lower = 0;
76          int upper = 6;
77          for (int i = 0; i < 2; i++) {
78              // cum(0,6) = p(0 < X <= 6) = 1, cum(1,5) = 4/6, cum(2,4) = 2/6
79              Assertions.assertEquals(1 - p * 2 * i,
80                      diceDistribution.probability(lower, upper), 1E-12);
81              lower++;
82              upper--;
83          }
84          for (int i = 0; i < 6; i++) {
85              Assertions.assertEquals(p, diceDistribution.probability(i, i + 1), 1E-12);
86          }
87      }
88  
89      @Test
90      void testInverseCumulativeProbabilityExtremes() {
91          // Require a lower bound of MIN_VALUE and the cumulative probability
92          // at that bound to be lower/higher than the argument cumulative probability.
93          // Use a uniform distribution so that the default search is supported
94          // for the inverse.
95          final DiscreteDistribution dist = new AbstractDiscreteDistribution() {
96              @Override
97              public double probability(int x) {
98                  throw new AssertionError();
99              }
100             @Override
101             public double cumulativeProbability(int x) {
102                 final int y = x - Integer.MIN_VALUE;
103                 if (y == 0) {
104                     return 0.25;
105                 } else if (y == 1) {
106                     return 0.5;
107                 } else if (y == 2) {
108                     return 0.75;
109                 } else {
110                     return 1.0;
111                 }
112             }
113             @Override
114             public double getMean() {
115                 return 1.5 + Integer.MIN_VALUE;
116             }
117             @Override
118             public double getVariance() {
119                 // n = upper - lower + 1, the variance is (n^2 - 1) / 12
120                 return 15.0 / 12;
121             }
122             @Override
123             public int getSupportLowerBound() {
124                 return Integer.MIN_VALUE;
125             }
126             @Override
127             public int getSupportUpperBound() {
128                 return Integer.MIN_VALUE + 3;
129             }
130         };
131         Assertions.assertEquals(dist.getSupportLowerBound(), dist.inverseCumulativeProbability(0.0));
132         Assertions.assertEquals(Integer.MIN_VALUE, dist.inverseCumulativeProbability(0.05));
133         Assertions.assertEquals(Integer.MIN_VALUE + 1, dist.inverseCumulativeProbability(0.35));
134         Assertions.assertEquals(Integer.MIN_VALUE + 2, dist.inverseCumulativeProbability(0.55));
135         Assertions.assertEquals(dist.getSupportUpperBound(), dist.inverseCumulativeProbability(1.0));
136 
137         Assertions.assertEquals(dist.getSupportLowerBound(), dist.inverseSurvivalProbability(1.0));
138         Assertions.assertEquals(Integer.MIN_VALUE, dist.inverseSurvivalProbability(0.95));
139         Assertions.assertEquals(Integer.MIN_VALUE + 1, dist.inverseSurvivalProbability(0.65));
140         Assertions.assertEquals(Integer.MIN_VALUE + 2, dist.inverseSurvivalProbability(0.45));
141         Assertions.assertEquals(dist.getSupportUpperBound(), dist.inverseSurvivalProbability(0.0));
142     }
143 
144     @Test
145     void testInverseCumulativeProbabilityWithNoMean() {
146         // A NaN mean will invalidate the Chebyshev inequality
147         // to prevent bracketing
148         final DiscreteDistribution dist = new AbstractDiscreteDistribution() {
149             @Override
150             public double probability(int x) {
151                 throw new AssertionError();
152             }
153             @Override
154             public double cumulativeProbability(int x) {
155                 if (x == 0) {
156                     return 0.25;
157                 } else if (x == 1) {
158                     return 0.5;
159                 } else if (x == 2) {
160                     return 0.75;
161                 } else {
162                     return 1.0;
163                 }
164             }
165             @Override
166             public double getMean() {
167                 return Double.NaN;
168             }
169             @Override
170             public double getVariance() {
171                 return Double.NaN;
172             }
173             @Override
174             public int getSupportLowerBound() {
175                 return 0;
176             }
177             @Override
178             public int getSupportUpperBound() {
179                 return 3;
180             }
181         };
182         Assertions.assertEquals(dist.getSupportLowerBound(), dist.inverseCumulativeProbability(0.0));
183         Assertions.assertEquals(0, dist.inverseCumulativeProbability(0.05));
184         Assertions.assertEquals(1, dist.inverseCumulativeProbability(0.35));
185         Assertions.assertEquals(2, dist.inverseCumulativeProbability(0.55));
186         Assertions.assertEquals(dist.getSupportUpperBound(), dist.inverseCumulativeProbability(1.0));
187 
188         Assertions.assertEquals(dist.getSupportLowerBound(), dist.inverseSurvivalProbability(1.0));
189         Assertions.assertEquals(0, dist.inverseSurvivalProbability(0.95));
190         Assertions.assertEquals(1, dist.inverseSurvivalProbability(0.65));
191         Assertions.assertEquals(2, dist.inverseSurvivalProbability(0.45));
192         Assertions.assertEquals(dist.getSupportUpperBound(), dist.inverseSurvivalProbability(0.0));
193     }
194 
195     /**
196      * Simple distribution modeling a 6-sided die
197      */
198     class DiceDistribution extends AbstractDiscreteDistribution {
199         private final double p = 1d / 6d;
200 
201         @Override
202         public double probability(int x) {
203             if (x < 1 || x > 6) {
204                 return 0;
205             } else {
206                 return p;
207             }
208         }
209 
210         @Override
211         public double cumulativeProbability(int x) {
212             if (x < 1) {
213                 return 0;
214             } else if (x >= 6) {
215                 return 1;
216             } else {
217                 return x / 6d;
218             }
219         }
220 
221         @Override
222         public double getMean() {
223             return 3.5;
224         }
225 
226         @Override
227         public double getVariance() {
228             return 70 / 24;  // E(X^2) - E(X)^2
229         }
230 
231         @Override
232         public int getSupportLowerBound() {
233             return 1;
234         }
235 
236         @Override
237         public int getSupportUpperBound() {
238             return 6;
239         }
240     }
241 }