Skip to content

Commit a097a28

Browse files
authored
feat: add WelfordAlgorithm, online mean and variance in one pass (#7602)
Welford's recurrence keeps a running mean and the sum of squared deviations from it, so it never forms the large nearly equal intermediate values that make the textbook variance formula lose its significant digits, and it needs O(1) time per sample and O(1) memory regardless of the stream length. Beyond the plain accumulation it supports removal, which runs the recurrence backwards and turns the accumulator into the statistics of a sliding window, and a static merge implementing Chan's parallel update so partial results from different shards combine exactly. Signed-off-by: alxkm <19151554+alxkm@users.noreply.github.com> Co-authored-by: alxkm <19151554+alxkm@users.noreply.github.com>
1 parent 1f08afb commit a097a28

2 files changed

Lines changed: 488 additions & 0 deletions

File tree

Lines changed: 245 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,245 @@
1+
package com.thealgorithms.streaming;
2+
3+
/**
4+
* Online (single pass) mean and variance using <b>Welford's algorithm</b>.
5+
*
6+
* <p>The textbook formula {@code Var = (sum(x^2) - n * mean^2) / (n - 1)} is fast but numerically
7+
* treacherous: {@code sum(x^2)} and {@code n * mean^2} may be huge and nearly equal, so their
8+
* difference loses most of its significant digits and can even come out negative. Welford's
9+
* recurrence never forms those large intermediate values. It keeps only the running mean and the sum
10+
* of squared deviations from that running mean, {@code M2}:
11+
*
12+
* <pre>
13+
* n &lt;- n + 1
14+
* delta &lt;- x - mean
15+
* mean &lt;- mean + delta / n
16+
* M2 &lt;- M2 + delta * (x - mean) // note: the second factor uses the *updated* mean
17+
* </pre>
18+
*
19+
* <p>Both {@link #add(double)} and {@link #remove(double)} run in O(1) time and the accumulator
20+
* occupies O(1) memory no matter how many samples pass through it.
21+
*
22+
* <h2>Sliding windows and map-reduce</h2>
23+
*
24+
* <ul>
25+
* <li>{@link #remove(double)} runs the recurrence backwards, which turns the accumulator into the
26+
* statistics of a sliding window: feed the incoming sample to {@code add} and the sample that
27+
* just left the window to {@code remove}. Removal is the one operation that can degrade
28+
* accuracy over a very long run, since the value being removed no longer matches the mean it
29+
* was added to; recreate the accumulator periodically if that matters.</li>
30+
* <li>{@link #merge(WelfordAlgorithm, WelfordAlgorithm)} implements Chan's parallel update, so
31+
* partial results computed on different shards can be combined exactly.</li>
32+
* </ul>
33+
*
34+
* <h2>Usage</h2>
35+
*
36+
* <pre>{@code
37+
* WelfordAlgorithm stats = new WelfordAlgorithm();
38+
* stats.add(2.0);
39+
* stats.add(4.0);
40+
* stats.add(4.0);
41+
* stats.mean(); // 3.3333...
42+
* stats.populationStandardDeviation(); // 0.9428...
43+
* }</pre>
44+
*
45+
* <p>This class is not thread-safe.
46+
*
47+
* @see ExponentialMovingAverage
48+
* @see <a href="https://en.wikipedia.org/wiki/Algorithms_for_calculating_variance">Algorithms for calculating variance</a>
49+
*/
50+
public final class WelfordAlgorithm {
51+
52+
private long count;
53+
private double mean;
54+
private double sumOfSquaredDeviations;
55+
56+
/**
57+
* Creates an empty accumulator.
58+
*/
59+
public WelfordAlgorithm() {
60+
clear();
61+
}
62+
63+
/**
64+
* Incorporates one sample.
65+
*
66+
* @param value the sample to add
67+
* @throws IllegalArgumentException if {@code value} is NaN or infinite
68+
*/
69+
public void add(double value) {
70+
requireFinite(value);
71+
count++;
72+
double delta = value - mean;
73+
mean += delta / count;
74+
sumOfSquaredDeviations += delta * (value - mean);
75+
}
76+
77+
/**
78+
* Incorporates every given sample, in order.
79+
*
80+
* @param values the samples to add
81+
* @throws IllegalArgumentException if any value is NaN or infinite
82+
* @throws NullPointerException if {@code values} is {@code null}
83+
*/
84+
public void addAll(double... values) {
85+
for (double value : values) {
86+
add(value);
87+
}
88+
}
89+
90+
/**
91+
* Removes a previously added sample, reversing {@link #add(double)}. This is what makes the
92+
* accumulator usable for a sliding window.
93+
*
94+
* @param value the sample to remove; it must genuinely have been added before
95+
* @throws IllegalStateException if the accumulator is empty
96+
* @throws IllegalArgumentException if {@code value} is NaN or infinite
97+
*/
98+
public void remove(double value) {
99+
requireFinite(value);
100+
if (count == 0) {
101+
throw new IllegalStateException("Cannot remove a sample from an empty accumulator");
102+
}
103+
if (count == 1) {
104+
clear();
105+
return;
106+
}
107+
double previousMean = mean;
108+
mean = (count * mean - value) / (count - 1);
109+
sumOfSquaredDeviations -= (value - previousMean) * (value - mean);
110+
count--;
111+
if (sumOfSquaredDeviations < 0.0) {
112+
sumOfSquaredDeviations = 0.0;
113+
}
114+
}
115+
116+
/**
117+
* Combines two independently accumulated summaries using Chan's parallel variance update.
118+
*
119+
* @param left summary of the first batch of samples
120+
* @param right summary of the second batch of samples
121+
* @return a new summary describing the concatenation of both batches
122+
* @throws NullPointerException if either argument is {@code null}
123+
*/
124+
public static WelfordAlgorithm merge(WelfordAlgorithm left, WelfordAlgorithm right) {
125+
WelfordAlgorithm merged = new WelfordAlgorithm();
126+
merged.count = left.count + right.count;
127+
if (merged.count == 0) {
128+
return merged;
129+
}
130+
double delta = right.mean - left.mean;
131+
merged.mean = left.mean + delta * right.count / merged.count;
132+
merged.sumOfSquaredDeviations = left.sumOfSquaredDeviations + right.sumOfSquaredDeviations + delta * delta * left.count * right.count / merged.count;
133+
return merged;
134+
}
135+
136+
/**
137+
* Returns the number of samples seen so far.
138+
*
139+
* @return the sample count
140+
*/
141+
public long count() {
142+
return count;
143+
}
144+
145+
/**
146+
* Tells whether any sample has been added.
147+
*
148+
* @return {@code true} if no sample is currently accounted for
149+
*/
150+
public boolean isEmpty() {
151+
return count == 0;
152+
}
153+
154+
/**
155+
* Returns the arithmetic mean of the samples.
156+
*
157+
* @return the mean, or {@link Double#NaN} if no sample has been added
158+
*/
159+
public double mean() {
160+
return count == 0 ? Double.NaN : mean;
161+
}
162+
163+
/**
164+
* Returns the sum of the samples, reconstructed from the mean.
165+
*
166+
* @return {@code count * mean}, or {@code 0} if no sample has been added
167+
*/
168+
public double sum() {
169+
return count == 0 ? 0.0 : mean * count;
170+
}
171+
172+
/**
173+
* Returns the sum of squared deviations from the mean, {@code M2}.
174+
*
175+
* @return the sum of squared deviations, {@code 0} for an empty accumulator
176+
*/
177+
public double sumOfSquaredDeviations() {
178+
return sumOfSquaredDeviations;
179+
}
180+
181+
/**
182+
* Returns the unbiased sample variance, normalised by {@code count - 1}.
183+
*
184+
* @return the sample variance, or {@link Double#NaN} if fewer than two samples were added
185+
*/
186+
public double sampleVariance() {
187+
return count < 2 ? Double.NaN : sumOfSquaredDeviations / (count - 1);
188+
}
189+
190+
/**
191+
* Returns the population variance, normalised by {@code count}.
192+
*
193+
* @return the population variance, or {@link Double#NaN} if no sample has been added
194+
*/
195+
public double populationVariance() {
196+
return count == 0 ? Double.NaN : sumOfSquaredDeviations / count;
197+
}
198+
199+
/**
200+
* Returns the square root of {@link #sampleVariance()}.
201+
*
202+
* @return the sample standard deviation, or {@link Double#NaN} if fewer than two samples were added
203+
*/
204+
public double sampleStandardDeviation() {
205+
return Math.sqrt(sampleVariance());
206+
}
207+
208+
/**
209+
* Returns the square root of {@link #populationVariance()}.
210+
*
211+
* @return the population standard deviation, or {@link Double#NaN} if no sample has been added
212+
*/
213+
public double populationStandardDeviation() {
214+
return Math.sqrt(populationVariance());
215+
}
216+
217+
/**
218+
* Returns the standard error of the mean, {@code sampleStandardDeviation / sqrt(count)}.
219+
*
220+
* @return the standard error, or {@link Double#NaN} if fewer than two samples were added
221+
*/
222+
public double standardError() {
223+
return sampleStandardDeviation() / Math.sqrt(count);
224+
}
225+
226+
/**
227+
* Forgets every sample.
228+
*/
229+
public void clear() {
230+
count = 0;
231+
mean = 0.0;
232+
sumOfSquaredDeviations = 0.0;
233+
}
234+
235+
@Override
236+
public String toString() {
237+
return "WelfordAlgorithm{count=" + count + ", mean=" + mean() + ", sampleStandardDeviation=" + sampleStandardDeviation() + '}';
238+
}
239+
240+
private static void requireFinite(double value) {
241+
if (!Double.isFinite(value)) {
242+
throw new IllegalArgumentException("Samples must be finite, but was " + value);
243+
}
244+
}
245+
}

0 commit comments

Comments
 (0)