103 static_assert(accumulators_ > 0);
106 if (num_total < accumulators_) {
107 if constexpr(accumulators_ == 1) {
111 for (std::size_t i = 0; i < num_total; ++i) {
119 if (accumulators_ > limit / 2) {
120 throw std::runtime_error(
"'max_sum_length' must be greater than '2 * accumulators_'");
123 std::size_t start = 0, end = num_total;
126 const std::size_t len = end - start;
128 work.states.emplace_back(end);
129 end = start + len / 2;
144 assert(start + accumulators_ <= end);
147 if constexpr(accumulators_ == 1) {
149 for (std::size_t i = start + 1; i < end; ++i) {
156 std::array<Output_, accumulators_> partials;
157 for (std::size_t a = 0; a < accumulators_; ++a) {
158 partials[a] = input(start + a);
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);
173 for (std::size_t i = 0; i < remainder; ++i) {
174 const std::size_t idx = start + num_cycles * accumulators_ + i;
178 tmp += recursive_sum(partials);
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();
188 if (work.states.empty()) {
193 work.states.back().left_sum = tmp;
194 end = work.states.back().left_end;
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