1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
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
25
26 class AbstractDiscreteDistributionTest {
27 private final DiceDistribution diceDistribution = new DiceDistribution();
28
29 @Test
30 void testInverseCumulativeProbabilityMethod() {
31
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
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
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
92
93
94
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
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
147
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
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;
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 }