welford_mean.c 5.8 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258
  1. #include "welford_mean.h"
  2. #include <math.h>
  3. #include <stdint.h>
  4. static inline float kahanSummation(float sum_previous, float input,
  5. float *accumulator) {
  6. const float y = input - *accumulator;
  7. const float t = sum_previous + y;
  8. *accumulator = (t - sum_previous) - y;
  9. return t;
  10. }
  11. void welford_mean_reset(WelfordMean_t *wm) {
  12. if (wm) {
  13. wm->_mean = 0.0f;
  14. wm->_M2 = 0.0f;
  15. wm->_mean_accum = 0.0f;
  16. wm->_M2_accum = 0.0f;
  17. wm->_count = 0;
  18. }
  19. }
  20. bool welford_mean_valid(WelfordMean_t *wm) {
  21. if (wm == NULL) {
  22. return false;
  23. }
  24. return wm->_count > 2;
  25. }
  26. bool welford_mean_updata(WelfordMean_t *wm, float new_val) {
  27. if (wm == NULL) {
  28. return false;
  29. }
  30. if (wm->_count == 0) {
  31. welford_mean_reset(wm);
  32. wm->_count = 1;
  33. wm->_mean = new_val;
  34. return false;
  35. } else if (wm->_count == UINT16_MAX) {
  36. // count overflow
  37. // reset count, but maintain mean and variance
  38. wm->_M2 = wm->_M2 / wm->_count;
  39. wm->_M2_accum = 0;
  40. wm->_count = 1;
  41. } else {
  42. wm->_count++;
  43. }
  44. // mean accumulates the mean of the entire dataset
  45. // delta can be very small compared to the mean, use algorithm to minimise
  46. // numerical error
  47. float delta = new_val - wm->_mean;
  48. float mean_change = delta / wm->_count;
  49. wm->_mean = kahanSummation(wm->_mean, mean_change, &wm->_mean_accum);
  50. // M2 aggregates the squared distance from the mean
  51. // count aggregates the number of samples seen so far
  52. const float M2_change = delta * (new_val - wm->_mean);
  53. wm->_M2 = kahanSummation(wm->_M2, M2_change, &wm->_M2_accum);
  54. // protect against floating point precision causing negative variances
  55. wm->_M2 = wm->_M2 > 0 ? wm->_M2 : 0;
  56. if (!isfinite(wm->_mean) || !isfinite(wm->_M2)) {
  57. welford_mean_reset(wm);
  58. return false;
  59. }
  60. return welford_mean_valid(wm);
  61. }
  62. int welford_mean_get_count(WelfordMean_t *wm) {
  63. if (wm) {
  64. return wm->_count;
  65. } else {
  66. return 0;
  67. }
  68. }
  69. float welford_mean_get_mean(WelfordMean_t *wm) {
  70. if (wm) {
  71. return wm->_mean;
  72. } else {
  73. return NAN;
  74. }
  75. }
  76. float welford_mean_get_variance(WelfordMean_t *wm) {
  77. if (wm) {
  78. return wm->_M2 / (wm->_count - 1);
  79. } else {
  80. return NAN;
  81. }
  82. }
  83. float welford_mean_get_stddev(WelfordMean_t *wm) {
  84. if (wm) {
  85. return sqrtf(welford_mean_get_variance(wm));
  86. } else {
  87. return NAN;
  88. }
  89. }
  90. bool welford_mean_vector3f_valid(WelfordMeanVector3f_t *wm) {
  91. if (wm == NULL) {
  92. return false;
  93. }
  94. return wm->_count > 2;
  95. }
  96. void welford_mean_vector3f_reset(WelfordMeanVector3f_t *wm) {
  97. if (wm) {
  98. for (int i = 0; i < 3; i++) {
  99. wm->_mean[i] = 0.0f;
  100. wm->_mean_accum[i] = 0.0f;
  101. for (int j = 0; j < 3; j++) {
  102. wm->_M2[i][j] = 0.0f;
  103. wm->_M2_accum[i][j] = 0.0f;
  104. }
  105. }
  106. wm->_count = 0;
  107. }
  108. }
  109. bool welford_mean_vector3f_updata(WelfordMeanVector3f_t *wm, float new_val[3]) {
  110. if (wm == NULL || new_val == NULL) {
  111. return false;
  112. }
  113. if (wm->_count == 0) {
  114. welford_mean_vector3f_reset(wm);
  115. wm->_count = 1;
  116. for (int i = 0; i < 3; i++) {
  117. wm->_mean[i] = new_val[i];
  118. }
  119. return false;
  120. } else if (wm->_count == UINT16_MAX) {
  121. // count overflow
  122. // reset count, but maintain mean and variance
  123. for (int i = 0; i < 3; i++) {
  124. for (int j = 0; j < 3; j++) {
  125. wm->_M2[i][j] = wm->_M2[i][j] / wm->_count;
  126. wm->_M2_accum[i][j] = 0;
  127. }
  128. }
  129. wm->_count = 1;
  130. } else {
  131. wm->_count++;
  132. }
  133. // mean
  134. // accumulates the mean of the entire dataset
  135. // use Kahan summation because delta can be very small compared to the mean
  136. float delta[3];
  137. for (int i = 0; i < 3; i++) {
  138. delta[i] = new_val[i] - wm->_mean[i];
  139. }
  140. float y[3];
  141. float t[3];
  142. for (int i = 0; i < 3; i++) {
  143. y[i] = delta[i] - wm->_mean_accum[i];
  144. t[i] = wm->_mean[i] + y[i];
  145. wm->_mean_accum[i] = (t[i] - wm->_mean[i]) - y[i];
  146. wm->_mean[i] = t[i];
  147. }
  148. for (int i = 0; i < 3; ++i) {
  149. if (!isfinite(wm->_mean[i])) {
  150. welford_mean_vector3f_reset(wm);
  151. return false;
  152. }
  153. }
  154. // covariance
  155. // Kahan summation (upper triangle only)
  156. // eg C(x,y) += dx * (y - mean_y)
  157. float m2_change[3][3];
  158. for (size_t r = 0; r < 3; r++) {
  159. for (size_t c = r; c < 3; c++) {
  160. m2_change[r][c] = delta[r] * (new_val[c] - wm->_mean[c]);
  161. }
  162. }
  163. for (size_t r = 0; r < 3; r++) {
  164. for (size_t c = r; c < 3; c++) {
  165. float y = m2_change[r][c] - wm->_M2_accum[r][c];
  166. float t = wm->_M2[r][c] + y;
  167. wm->_M2_accum[r][c] = (t - wm->_M2[r][c]) - y;
  168. wm->_M2[r][c] = t;
  169. }
  170. // protect against floating point precision causing negative variances
  171. if (wm->_M2[r][r] < 0) {
  172. wm->_M2[r][r] = 0;
  173. }
  174. }
  175. // make symmetric
  176. for (size_t r = 0; r < 3; r++) {
  177. for (size_t c = r + 1; c < 3; c++) {
  178. wm->_M2[c][r] = wm->_M2[r][c];
  179. }
  180. }
  181. for (size_t r = 0; r < 3; r++) {
  182. for (size_t c = 0; c < 3; c++) {
  183. if (!isfinite(wm->_M2[r][c])) {
  184. welford_mean_vector3f_reset(wm);
  185. return false;
  186. }
  187. }
  188. }
  189. return welford_mean_vector3f_valid(wm);
  190. }
  191. int welford_mean_vector3f_get_count(WelfordMeanVector3f_t *wm) {
  192. if (wm == NULL) {
  193. return 0;
  194. }
  195. return wm->_count;
  196. }
  197. bool welford_mean_vector3f_get_mean(WelfordMeanVector3f_t *wm, float mean[3]) {
  198. if (wm == NULL || mean == NULL) {
  199. return false;
  200. }
  201. for (int i = 0; i < 3; i++) {
  202. mean[i] = wm->_mean[i];
  203. }
  204. return true;
  205. }
  206. bool welford_mean_vector3f_get_variance(WelfordMeanVector3f_t *wm,
  207. float variance[3]) {
  208. if (wm == NULL || variance == NULL) {
  209. return false;
  210. }
  211. for (int i = 0; i < 3; i++) {
  212. variance[i] = wm->_M2[i][i] / (wm->_count - 1);
  213. }
  214. return true;
  215. }
  216. float welford_mean_vector3f_get_covariance(WelfordMeanVector3f_t *wm, int x,
  217. int y) {
  218. if (wm == NULL || x < 0 || x > 2 || y < 0 || y > 2) {
  219. return NAN;
  220. }
  221. return wm->_M2[x][y] / (wm->_count - 1);
  222. }