tatami_stats
Matrix statistics for tatami
Loading...
Searching...
No Matches
rss.hpp
Go to the documentation of this file.
1#ifndef TATAMI_STATS_SKIP_NAN_RSS_HPP
2#define TATAMI_STATS_SKIP_NAN_RSS_HPP
3
4#include "../utils.hpp"
5
6#include <vector>
7#include <cmath>
8#include <numeric>
9#include <limits>
10#include <algorithm>
11#include <cstddef>
12#include <optional>
13
14#include "tatami/tatami.hpp"
15#include "sanisizer/sanisizer.hpp"
17
24namespace tatami_stats {
25
26namespace skip_nan {
27
32template<typename Output_ = double>
46
54template<typename Output_, typename Count_>
55struct RssBuffers {
60 Output_* mean;
61
66 Output_* rss;
67
72 Count_* count;
73};
74
78template<typename Value_, typename Index_, typename Output_, typename Count_>
79void rss_direct(bool row, const tatami::Matrix<Value_, Index_>& mat, RssBuffers<Output_, Count_>& output, const RssOptions<Output_>& opt) {
80 const auto dim = (row ? mat.nrow() : mat.ncol());
81 const auto otherdim = (row ? mat.ncol() : mat.nrow());
84
85 if (mat.sparse()) {
86 tatami::Options topt;
87 topt.sparse_extract_index = false;
88
89 tatami::parallelize([&](int, Index_ s, Index_ l) -> void {
90 auto ext = tatami::consecutive_extractor<true>(mat, row, s, l, topt);
93 for (Index_ x = 0; x < l; ++x) {
94 auto out = ext->fetch(vbuffer.data(), NULL);
95 tatami::copy_n(out.value, out.number, vbuffer.data());
96 const auto new_number = shift_nans(vbuffer.data(), out.number);
97 const Index_ new_total = otherdim - (out.number - new_number);
98 const auto res = quickstats::rss(new_total, new_number, vbuffer.data(), work, ropt);
99 output.mean[x + s] = res.mean;
100 output.rss[x + s] = res.rss;
101 output.count[x + s] = new_total;
102 }
103 }, dim, opt.num_threads);
104
105 } else {
106 tatami::parallelize([&](int, Index_ s, Index_ l) -> void {
107 auto ext = tatami::consecutive_extractor<false>(mat, row, s, l);
110 for (Index_ x = 0; x < l; ++x) {
111 auto out = ext->fetch(buffer.data());
112 tatami::copy_n(out, otherdim, buffer.data());
113 const auto new_total = shift_nans(buffer.data(), otherdim);
114 const auto res = quickstats::rss(new_total, buffer.data(), work, ropt);
115 output.mean[x + s] = res.mean;
116 output.rss[x + s] = res.rss;
117 output.count[x + s] = new_total;
118 }
119 }, dim, opt.num_threads);
120 }
121}
122
123template<typename Value_, typename Index_, typename Output_, typename Count_>
124void rss_running(bool row, const tatami::Matrix<Value_, Index_>& mat, RssBuffers<Output_, Count_>& output, const RssOptions<Output_>& opt) {
125 const auto dim = (row ? mat.nrow() : mat.ncol());
126 const auto otherdim = (row ? mat.ncol() : mat.nrow());
127 const bool is_sparse = mat.is_sparse();
128
129 std::fill_n(output.rss, dim, 0);
130 std::fill_n(output.count, dim, 0);
131 if (otherdim == 0) {
132 std::fill_n(output.mean, dim, opt.mean_placeholder);
133 return;
134 } else {
135 std::fill_n(output.mean, dim, 0);
136 }
137
138 assert(opt.num_threads > 0);
139 const bool do_parallel = opt.num_threads > 1;
140 std::optional<std::vector<std::optional<std::vector<Output_> > > > all_partial_mean, all_partial_rss;
141 std::optional<std::vector<std::optional<std::vector<Count_> > > > all_partial_count;
142 if (do_parallel) {
143 // -1, as we'll repurpose the output buffers to store the output of the first thread.
144 all_partial_rss.emplace(sanisizer::cast<I<decltype(all_partial_rss->size())> >(opt.num_threads - 1));
145 all_partial_mean.emplace(sanisizer::cast<I<decltype(all_partial_mean->size())> >(opt.num_threads));
146 all_partial_count.emplace(sanisizer::cast<I<decltype(all_partial_mean->size())> >(opt.num_threads));
147 }
148
149 const int nused = tatami::parallelize([&](int thread, Index_ s, Index_ l) -> void {
150 Output_* rss_ptr;
151 Output_* mean_ptr;
152 Count_* count_ptr;
153 std::optional<std::vector<Output_> > cur_rss, cur_mean;
154 std::optional<std::vector<Count_> > cur_count;
155
156 if (!do_parallel) {
157 // Storing mean and RSS directly in the output vector to cut down two allocations if we're not working in parallel.
158 rss_ptr = output.rss;
159 mean_ptr = output.mean;
160 count_ptr = output.count;
161 } else {
162 // Storing the partial RSS directly in the output vector to save ourselves an allocation if we're in the first thread.
163 // We can't do the same for the mean, though, as we need to keep the partial mean and the global mean separate for the reduction.
164 cur_mean.emplace(tatami::cast_Index_to_container_size<std::vector<Output_> >(dim));
165 mean_ptr = cur_mean->data();
166 cur_count.emplace(tatami::cast_Index_to_container_size<std::vector<Count_> >(dim));
167 count_ptr = cur_count->data();
168 if (thread == 0) {
169 rss_ptr = output.rss;
170 } else {
171 cur_rss.emplace(tatami::cast_Index_to_container_size<std::vector<Output_> >(dim));
172 rss_ptr = cur_rss->data();
173 }
174 }
175
176 if (is_sparse) {
177 tatami::Options topt;
178 topt.sparse_ordered_index = false;
179 auto ext = tatami::consecutive_extractor<true>(mat, !row, s, l, topt);
183
184 for (Index_ x = 0; x < l; ++x) {
185 auto out = ext->fetch(vbuffer.data(), ibuffer.data());
186 for (Index_ i = 0; i < out.number; ++i) {
187 const auto d = out.index[i];
188 const auto val = out.value[i];
189 if (!std::isnan(val)) {
190 auto& nnz = nonzeros[d];
191 quickstats::update_rss(mean_ptr[d], rss_ptr[d], val, ++nnz); // increment is safe as 'nnz + 1 <= l' fits in an Index_;
192 } else {
193 ++count_ptr[d];
194 }
195 }
196 }
197
198 for (Index_ d = 0; d < dim; ++d) {
199 auto& unskipped_total = count_ptr[d];
200 unskipped_total = l - unskipped_total; // could be zero, so the update with zeros needs to be safe.
201 quickstats::update_rss_with_zeros(mean_ptr[d], rss_ptr[d], static_cast<Count_>(unskipped_total - nonzeros[d]), unskipped_total);
202 }
203
204 } else {
205 auto ext = tatami::consecutive_extractor<false>(mat, !row, s, l);
207
208 for (Index_ x = 0; x < l; ++x) {
209 auto out = ext->fetch(buffer.data());
210 for (Index_ d = 0; d < dim; ++d) {
211 const auto val = out[d];
212 if (!std::isnan(val)) {
213 quickstats::update_rss(mean_ptr[d], rss_ptr[d], val, ++count_ptr[d]); // increment is safe as 'count_ptr[d] + 1 <= l' fits in an index.
214 }
215 }
216 }
217 }
218
219 // Moving results to the main containers.
220 if (do_parallel) {
221 (*all_partial_mean)[thread] = std::move(cur_mean);
222 (*all_partial_count)[thread] = std::move(cur_count);
223 if (thread > 0) {
224 (*all_partial_rss)[thread - 1] = std::move(cur_rss);
225 }
226 }
227 }, otherdim, opt.num_threads);
228 assert(nused > 0);
229
230 // Don't check nused > 1, as it's possible for do_parallel = true with nused = 1 if not all threads are used.
231 // This would cause us to leave output.mean and output.rss empty.
232 if (do_parallel) {
233 const auto& ap_mean = *all_partial_mean;
234 const auto& ap_rss = *all_partial_rss;
235 const auto& ap_count = *all_partial_count;
236
237 // Computing the global total.
238 for (int u = 0; u < nused; ++u) {
239 const auto& cur_count = *(ap_count[u]);
240 for (Index_ d = 0; d < dim; ++d) {
241 output.count[d] += cur_count[d];
242 }
243 }
244
245 // Computing the global mean from its components.
246 for (int u = 0; u < nused; ++u) {
247 const auto& cur_count = *(ap_count[u]);
248 const auto& cur_mean = *(ap_mean[u]);
249 for (Index_ d = 0; d < dim; ++d) {
250 if (cur_count[d] > 0) { // protect against NaN means at a count of 0.
251 const auto mult = static_cast<Output_>(cur_count[d]) / static_cast<Output_>(output.count[d]);
252 output.mean[d] += cur_mean[d] * mult;
253 }
254 }
255 }
256
257 // Combining the RSS. This time, we need to use the safe version as we don't know whether all elements were skipped in a thread.
258 for (int u = 0; u < nused; ++u) {
259 const auto& cur_count = *(ap_count[u]);
260 const auto& cur_mean = *(ap_mean[u]);
261 if (u == 0) {
262 for (Index_ d = 0; d < dim; ++d) {
263 output.rss[d] = quickstats::recenter_rss(cur_count[d], output.rss[d], cur_mean[d], output.mean[d]);
264 }
265 } else {
266 const auto& cur_rss = *(ap_rss[u - 1]);
267 for (Index_ d = 0; d < dim; ++d) {
268 output.rss[d] += quickstats::recenter_rss(cur_count[d], cur_rss[d], cur_mean[d], output.mean[d]);
269 }
270 }
271 }
272 }
273
274 for (Index_ d = 0; d < dim; ++ d) {
275 if (output.count[d] == 0) {
276 output.mean[d] = opt.mean_placeholder;
277 }
278 }
279}
302template<typename Value_, typename Index_, typename Output_, typename Count_>
304 if (mat.prefer_rows() == row) {
305 rss_direct(row, mat, output, opt);
306 } else {
307 rss_running(row, mat, output, opt);
308 }
309}
310
318template<typename Output_, typename Count_>
319struct RssResult {
324 std::vector<Output_> mean;
325
330 std::vector<Output_> rss;
331
336 std::vector<Count_> count;
337};
338
355template<typename Output_ = double, typename Count_, typename Value_, typename Index_>
358 const auto dim = (row ? mat.nrow() : mat.ncol());
359
361#ifdef TATAMI_STATS_TEST_DIRTY
362 , -1
363#endif
364 );
366#ifdef TATAMI_STATS_TEST_DIRTY
367 , -1
368#endif
369 );
371#ifdef TATAMI_STATS_TEST_DIRTY
372 , -1
373#endif
374 );
375
377 buffers.mean = output.mean.data();
378 buffers.rss = output.rss.data();
379 buffers.count = output.count.data();
380
381 rss(row, mat, buffers, opt);
382 return output;
383}
384
385}
386
387}
388
389#endif
virtual Index_ ncol() const=0
virtual Index_ nrow() const=0
virtual bool prefer_rows() const=0
virtual bool is_sparse() const=0
virtual std::unique_ptr< MyopicSparseExtractor< Value_, Index_ > > sparse(bool row, const Options &opt) const=0
RssResult< Output_ > rss(const std::size_t num_total, const std::size_t num_non_zero, const Input_ *const ptr, RssWorkspace< Output_ > &work, const RssOptions< Output_ > &options)
Float_ recenter_rss(const Count_ num_total, const Float_ old_rss, const Float_ old_mean, const Float_ new_mean)
constexpr Value_ nan_if_available_else_zero()
void update_rss_with_zeros(Output_ &mean, Output_ &rss, const Count_ num_zeros, const Count_ num_total)
void update_rss(Output_ &mean, Output_ &rss, const Input_ value, const Count_ num_total)
void rss(bool row, const tatami::Matrix< Value_, Index_ > &mat, RssBuffers< Output_, Count_ > &output, const RssOptions< Output_ > &opt)
Definition rss.hpp:303
Functions to compute statistics from a tatami::Matrix.
Definition count.hpp:20
void resize_container_to_Index_size(Container_ &container, const Index_ x, Args_ &&... args)
int parallelize(Function_ fun, const Index_ tasks, const int workers)
I< decltype(std::declval< Container_ >().size())> cast_Index_to_container_size(const Index_ x)
Value_ * copy_n(const Value_ *const input, const Size_ n, Value_ *const output)
Container_ create_container_of_Index_size(const Index_ x, Args_ &&... args)
auto consecutive_extractor(const Matrix< Value_, Index_ > &matrix, const bool row, const Index_ iter_start, const Index_ iter_length, Args_ &&... args)
bool sparse_extract_index
bool sparse_ordered_index
Result buffers for skip_nan::rss().
Definition rss.hpp:55
Output_ * mean
Definition rss.hpp:60
Output_ * rss
Definition rss.hpp:66
Count_ * count
Definition rss.hpp:72
Options for skip_nan::rss().
Definition rss.hpp:33
Output_ mean_placeholder
Definition rss.hpp:44
int num_threads
Definition rss.hpp:38
Results of skip_nan::rss().
Definition rss.hpp:319
std::vector< Count_ > count
Definition rss.hpp:336
std::vector< Output_ > mean
Definition rss.hpp:324
std::vector< Output_ > rss
Definition rss.hpp:330