quickstats
Quickly compute simple statistics
Loading...
Searching...
No Matches
pairwise_sum.hpp
Go to the documentation of this file.
1#ifndef QUICKSTATS_PAIRWISE_SUM_HPP
2#define QUICKSTATS_PAIRWISE_SUM_HPP
3
4#include <vector>
5#include <cstddef>
6#include <optional>
7#include <cassert>
8#include <array>
9#include <stdexcept>
10
16namespace quickstats {
17
23template<typename Output_ = double>
28 struct State {
29 State(const std::size_t left_end) : left_end(left_end) {}
30 std::size_t left_end;
31 std::optional<Output_> left_sum;
32 };
33 std::vector<State> states;
37};
38
48 std::size_t max_sum_length = 128;
49};
50
54// Mathematically equivalent to std::accumulate but reorders summations for greater instruction-level parallelism.
55// See performance tests in https://github.com/tatami-inc/test-multiplication/tree/master/other/accumulators.
56template<typename Output_, std::size_t width_>
57Output_ recursive_sum(std::array<Output_, width_>& dots) {
58 if constexpr(width_ == 1) {
59 return dots[0];
60 } else if constexpr(width_ == 2) {
61 return dots[0] + dots[1];
62 } else {
63 constexpr auto half_width = width_ / 2;
64 std::array<Output_, half_width> tmp;
65 for (std::size_t s = 0; s < half_width; ++s) { // Increase potential for vectorization.
66 tmp[s] = dots[s] + dots[s + half_width];
67 }
68 if constexpr(width_ % 2 == 1) {
69 return recursive_sum(tmp) + dots[width_ - 1];
70 } else {
71 return recursive_sum(tmp);
72 }
73 }
74}
101template<std::size_t accumulators_ = 4, class Input_, typename Output_>
102Output_ pairwise_sum_abstract(const std::size_t num_total, Input_ input, PairwiseSumWorkspace<Output_>& work, const PairwiseSumOptions& options) {
103 static_assert(accumulators_ > 0);
104
105 work.states.clear();
106 if (num_total < accumulators_) {
107 if constexpr(accumulators_ == 1) {
108 return 0;
109 } else {
110 Output_ out = 0;
111 for (std::size_t i = 0; i < num_total; ++i) {
112 out += input(i);
113 }
114 return out;
115 }
116 }
117
118 const std::size_t limit = options.max_sum_length;
119 if (accumulators_ > limit / 2) {
120 throw std::runtime_error("'max_sum_length' must be greater than '2 * accumulators_'");
121 }
122
123 std::size_t start = 0, end = num_total;
124 Output_ out = 0;
125 while (1) {
126 const std::size_t len = end - start;
127 if (len > limit) {
128 work.states.emplace_back(end);
129 end = start + len / 2;
130 continue;
131 }
132
133 // 'start + accumulators_ <= end' should always be true.
134 //
135 // Let's start by considering the initial case where 'num_total <= limit'.
136 // As we already know that 'num_total >= accumulators_', we get 'start + accumulators_ == accumulators_ <= num_total == end'.
137 //
138 // Alright, what about the left-hand-side of the recursion?
139 // We defined 'end = start + len / 2', and we already know that 'len > limit' and 'limit / 2 >= accumulators_';
140 // hence, we know that that 'start + len / 2 >= start + accumulators_'.
141 //
142 // The right-hand-side of the recursion is easier as we know its length is greater than or equal to 'len / 2' (as integer division is truncating).
143 // So, if the LHS fulfills the requirement, then the RHS must definitely fulfill it.
144 assert(start + accumulators_ <= end);
145
146 Output_ tmp;
147 if constexpr(accumulators_ == 1) {
148 tmp = input(start); // We know that start < end, so we can skip one addition.
149 for (std::size_t i = start + 1; i < end; ++i) {
150 tmp += input(i);
151 }
152
153 } else {
154 // This accumulator logic was originally implemented in https://github.com/tatami-inc/tatami_mult.
155 // We added peeling as we can guarantee that we have enough observations and thus can omit the conditional.
156 std::array<Output_, accumulators_> partials;
157 for (std::size_t a = 0; a < accumulators_; ++a) { // peeling the first loop as we know that start + accumulators_ <= end.
158 partials[a] = input(start + a);
159 }
160
161 const std::size_t num_cycles = len / accumulators_;
162 const std::size_t remainder = len % accumulators_;
163 for (std::size_t c = 1; c < num_cycles; ++c) {
164 for (std::size_t a = 0; a < accumulators_; ++a) {
165 const std::size_t idx = start + c * accumulators_ + a;
166 partials[a] += input(idx);
167 }
168 }
169
170 // Technically, we could structure the splits to reduce the number of calls to the epilogue loops.
171 // However, this would weaken the symmetry of the splitting and compromise the precision improvements.
172 tmp = 0;
173 for (std::size_t i = 0; i < remainder; ++i) {
174 const std::size_t idx = start + num_cycles * accumulators_ + i;
175 tmp += input(idx);
176 }
177
178 tmp += recursive_sum(partials);
179 }
180
181 start = end;
182 while (work.states.size() && work.states.back().left_sum.has_value()) {
183 tmp += *(work.states.back().left_sum);
184 start = work.states.back().left_end;
185 work.states.pop_back();
186 }
187
188 if (work.states.empty()) {
189 out = tmp;
190 break;
191 }
192
193 work.states.back().left_sum = tmp;
194 end = work.states.back().left_end;
195 }
196
197 return out;
198}
199
214template<std::size_t accumulators_ = 4, typename Input_, typename Output_>
215Output_ pairwise_sum(const std::size_t num_total, const Input_* const ptr, PairwiseSumWorkspace<Output_>& work, const PairwiseSumOptions& options) {
217 num_total,
218 [&](const std::size_t i) -> auto {
219 return ptr[i];
220 },
221 work,
222 options
223 );
224}
225
226/* COMMENTS:
227 * I tried to write a multi-threaded version of this where each direct summation was submitted to a separate worker until all workers were occupied,
228 * and then added the results once they became available from each worker.
229 * This worksharing is fine-grained but imposes a high cost for inter-thread communication relative to the summation for small `limit_`.
230 * As a consequence, the performance of this multi-threaded version is worse than its serial counterpart.
231 *
232 * I could have implemented alternative approaches that involve less communication but require more memory.
233 * For example, we could split elements into the subarrays ahead of time, distribute the summations to threads once, and then sum the results once all workers are done.
234 * This greatly reduces the cross-talk between threads but requires an extra allocation to store the results.
235 *
236 * TBH, the easiest and most performant approach to parallelization is to just split your input array into one subarray per worker,
237 * perform the sum within each worker for that subarray, and then add the sums afterwards.
238 * This won't give exactly the same result as serial execution but we've crossed that bridge already.
239 * (If exact results are required, we can split it into ceil(log2(num_workers)) subarrays,
240 * which allows us to follow the same halving as pairwise_sum() to get the exact same result at the cost of suboptimal worksharing.)
241 *
242 * In any case, summation is already so fast that I don't think we need to spend a lot of effort in thinking about parallelization.
243 * Especially given that, in real applications, the other threads will typically be occupied elsewhere.
244 * Indeed, we don't deal with parallelization in other parts of this library, so it would be odd to implement it here.
245 *
246 * I also have a sneaking suspicion that the serial code is already pseudo-parallelized via out-of-order execution,
247 * where the next summation starts before the first one has ended.
248 * I say this because pairwise_sum() somehow manages to be slightly faster than std::accumulate() in our R bindings.
249 */
250
254// For back-compatibility.
255template<std::size_t limit_ = 128, std::size_t accumulators_ = 4, class Input_, class Modifier_, typename Output_>
256Output_ pairwise_sum(const std::size_t num_total, const Input_* const ptr, Modifier_ mod, PairwiseSumWorkspace<Output_>& work) {
258 num_total,
259 [&](const std::size_t i) -> auto {
260 return mod(i, ptr[i]);
261 },
262 work,
263 [&]{
264 PairwiseSumOptions opt;
265 opt.max_sum_length = limit_;
266 return opt;
267 }()
268 );
269}
270
271template<std::size_t limit_ = 128, std::size_t accumulators_ = 4, typename Input_, typename Output_>
272Output_ pairwise_sum(const std::size_t num_total, const Input_* const ptr, PairwiseSumWorkspace<Output_>& work) {
274 num_total,
275 ptr,
276 work,
277 [&]{
278 PairwiseSumOptions opt;
279 opt.max_sum_length = limit_;
280 return opt;
281 }()
282 );
283}
288}
289
290#endif
Quickly compute simple statistics.
Definition mad.hpp:15
Output_ pairwise_sum_abstract(const std::size_t num_total, Input_ input, PairwiseSumWorkspace< Output_ > &work, const PairwiseSumOptions &options)
Definition pairwise_sum.hpp:102
Output_ pairwise_sum(const std::size_t num_total, const Input_ *const ptr, PairwiseSumWorkspace< Output_ > &work, const PairwiseSumOptions &options)
Definition pairwise_sum.hpp:215
Options for pairwise_sum() and pairwise_sum_abstract().
Definition pairwise_sum.hpp:42
std::size_t max_sum_length
Definition pairwise_sum.hpp:48
Re-usable workspace for pairwise_sum() and pairwise_sum_abstract().
Definition pairwise_sum.hpp:24