tatami_stats
Matrix statistics for tatami
Loading...
Searching...
No Matches
rss.hpp
Go to the documentation of this file.
1#ifndef TATAMI_STATS_RSS_HPP
2#define TATAMI_STATS_RSS_HPP
3
4#include "utils.hpp"
5
6#include <vector>
7#include <cmath>
8#include <numeric>
9#include <algorithm>
10#include <cstddef>
11#include <optional>
12
13#include "tatami/tatami.hpp"
14#include "sanisizer/sanisizer.hpp"
16
23namespace tatami_stats {
24
29template<typename Output_ = double>
43
49template<typename Output_>
50struct RssBuffers {
55 Output_* mean;
56
61 Output_* rss;
62};
63
67template<typename Value_, typename Index_, typename Output_>
68void rss_direct(bool row, const tatami::Matrix<Value_, Index_>& mat, RssBuffers<Output_>& output, const RssOptions<Output_>& opt) {
69 const auto dim = (row ? mat.nrow() : mat.ncol());
70 const auto otherdim = (row ? mat.ncol() : mat.nrow());
73
74 if (mat.sparse()) {
75 tatami::Options topt;
76 topt.sparse_extract_index = false;
77
78 tatami::parallelize([&](int, Index_ s, Index_ l) -> void {
79 auto ext = tatami::consecutive_extractor<true>(mat, row, s, l, topt);
82 for (Index_ x = 0; x < l; ++x) {
83 auto out = ext->fetch(vbuffer.data(), NULL);
84 const auto res = quickstats::rss(otherdim, out.number, out.value, work, ropt);
85 output.mean[x + s] = res.mean;
86 output.rss[x + s] = res.rss;
87 }
88 }, dim, opt.num_threads);
89
90 } else {
91 tatami::parallelize([&](int, Index_ s, Index_ l) -> void {
92 auto ext = tatami::consecutive_extractor<false>(mat, row, s, l);
95 for (Index_ x = 0; x < l; ++x) {
96 auto out = ext->fetch(buffer.data());
97 const auto res = quickstats::rss(otherdim, out, work, ropt);
98 output.mean[x + s] = res.mean;
99 output.rss[x + s] = res.rss;
100 }
101 }, dim, opt.num_threads);
102 }
103}
104
105template<typename Value_, typename Index_, typename Output_>
106void rss_running(bool row, const tatami::Matrix<Value_, Index_>& mat, RssBuffers<Output_>& output, const RssOptions<Output_>& opt) {
107 const auto dim = (row ? mat.nrow() : mat.ncol());
108 const auto otherdim = (row ? mat.ncol() : mat.nrow());
109 if (otherdim == 0) {
110 std::fill_n(output.mean, dim, opt.mean_placeholder);
111 std::fill_n(output.rss, dim, 0);
112 return;
113 }
114
115 assert(opt.num_threads > 0);
116 const bool do_parallel = opt.num_threads > 1;
117 std::optional<std::vector<std::optional<std::vector<Output_> > > > all_partial_mean, all_partial_rss;
118 std::optional<std::vector<Index_> > all_partial_count;
119 if (do_parallel) {
120 // -1, as we'll repurpose the RSS output buffer to store the partial RSS of the first thread.
121 all_partial_rss.emplace(sanisizer::cast<I<decltype(all_partial_rss->size())> >(opt.num_threads - 1));
122 all_partial_mean.emplace(sanisizer::cast<I<decltype(all_partial_mean->size())> >(opt.num_threads));
123 all_partial_count.emplace(sanisizer::cast<I<decltype(all_partial_count->size())> >(opt.num_threads));
124 }
125
126 // We overwrite any existing mean value in the array in the do_parallel=true situation.
127 // So, the initial value doesn't need to be zero.
128 if (!do_parallel) {
129 std::fill_n(output.mean, dim, 0);
130 }
131 std::fill_n(output.rss, dim, 0);
132
133 const bool is_sparse = mat.is_sparse();
134 const int nused = tatami::parallelize([&](int thread, Index_ s, Index_ l) -> void {
135 Output_* rss_ptr;
136 Output_* mean_ptr;
137 std::optional<std::vector<Output_> > cur_rss, cur_mean;
138
139 if (!do_parallel) {
140 // Storing mean and RSS directly in the output vector to cut down two allocations if we're not working in parallel.
141 rss_ptr = output.rss;
142 mean_ptr = output.mean;
143 } else {
144 // Storing the partial RSS directly in the output vector to save ourselves an allocation if we're in the first thread.
145 // 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.
146 cur_mean.emplace(tatami::cast_Index_to_container_size<std::vector<Output_> >(dim));
147 mean_ptr = cur_mean->data();
148 if (thread == 0) {
149 rss_ptr = output.rss;
150 } else {
151 cur_rss.emplace(tatami::cast_Index_to_container_size<std::vector<Output_> >(dim));
152 rss_ptr = cur_rss->data();
153 }
154 }
155
156 if (is_sparse) {
157 tatami::Options topt;
158 topt.sparse_ordered_index = false;
159 auto ext = tatami::consecutive_extractor<true>(mat, !row, s, l, topt);
163
164 for (Index_ x = 0; x < l; ++x) {
165 auto out = ext->fetch(vbuffer.data(), ibuffer.data());
166 for (Index_ i = 0; i < out.number; ++i) {
167 const auto d = out.index[i];
168 auto& nnz = nonzeros[d];
169 quickstats::update_rss(mean_ptr[d], rss_ptr[d], out.value[i], ++nnz); // increment is safe as 'nnz + 1 <= l' fits in an Index_;
170 }
171 }
172
173 for (Index_ d = 0; d < dim; ++d) {
174 // otherdim > 0 is guaranteed, so we can use the unsafe version.
175 quickstats::update_rss_with_zeros_unsafe(mean_ptr[d], rss_ptr[d], static_cast<Index_>(l - nonzeros[d]), l);
176 }
177
178 } else {
179 auto ext = tatami::consecutive_extractor<false>(mat, !row, s, l);
181
182 for (Index_ x = 0; x < l; ++x) {
183 auto out = ext->fetch(buffer.data());
184 for (Index_ d = 0; d < dim; ++d) {
185 quickstats::update_rss(mean_ptr[d], rss_ptr[d], out[d], x + 1); // increment is safe as ' x + 1 <= l' fits in an Index_.
186 }
187 }
188 }
189
190 if (do_parallel) {
191 (*all_partial_count)[thread] = l;
192 (*all_partial_mean)[thread] = std::move(cur_mean);
193 if (thread > 0) {
194 (*all_partial_rss)[thread - 1] = std::move(cur_rss);
195 }
196 }
197 }, otherdim, opt.num_threads);
198 assert(nused > 0);
199
200 // Don't check nused > 1, as it's possible for do_parallel = true with nused = 1 if not all threads are used.
201 // This would cause us to leave output.mean and output.rss empty.
202 if (do_parallel) {
203 const auto& ap_count = *all_partial_count;
204 const auto& ap_mean = *all_partial_mean;
205 const auto& ap_rss = *all_partial_rss;
206
207 // Computing the global mean. All ap_count is positive so we don't have to worry about cur_mean[d] being NaN.
208 for (int u = 0; u < nused; ++u) {
209 const Output_ mult = static_cast<Output_>(ap_count[u]) / static_cast<Output_>(otherdim);
210 const auto& cur_mean = *(ap_mean[u]);
211 if (u == 0) {
212 for (Index_ d = 0; d < dim; ++d) {
213 output.mean[d] = cur_mean[d] * mult;
214 }
215 } else {
216 for (Index_ d = 0; d < dim; ++d) {
217 output.mean[d] += cur_mean[d] * mult;
218 }
219 }
220 }
221
222 // Combining the RSS. We can use recenter_rss_unsafe() as we are guaranteed that cur_count > 0,
223 // as parallelize() will only ever split into non-empty ranges if those ranges are used.
224 for (int u = 0; u < nused; ++u) {
225 const auto cur_count = ap_count[u];
226 const auto& cur_mean = *(ap_mean[u]);
227 if (u == 0) {
228 for (Index_ d = 0; d < dim; ++d) {
229 output.rss[d] = quickstats::recenter_rss_unsafe(cur_count, output.rss[d], cur_mean[d], output.mean[d]);
230 }
231 } else {
232 const auto& cur_rss = *(ap_rss[u - 1]);
233 for (Index_ d = 0; d < dim; ++d) {
234 output.rss[d] += quickstats::recenter_rss_unsafe(cur_count, cur_rss[d], cur_mean[d], output.mean[d]);
235 }
236 }
237 }
238 }
239}
260template<typename Value_, typename Index_, typename Output_>
261void rss(bool row, const tatami::Matrix<Value_, Index_>& mat, RssBuffers<Output_>& output, const RssOptions<Output_>& opt) {
262 if (mat.prefer_rows() == row) {
263 rss_direct(row, mat, output, opt);
264 } else {
265 rss_running(row, mat, output, opt);
266 }
267}
268
274template<typename Output_>
275struct RssResult {
280 std::vector<Output_> mean;
281
286 std::vector<Output_> rss;
287};
288
303template<typename Output_ = double, typename Value_, typename Index_>
305 RssResult<Output_> output;
306 const auto dim = (row ? mat.nrow() : mat.ncol());
308#ifdef TATAMI_STATS_TEST_DIRTY
309 , -1
310#endif
311 );
313#ifdef TATAMI_STATS_TEST_DIRTY
314 , -1
315#endif
316 );
317
318 RssBuffers<Output_> buffers;
319 buffers.mean = output.mean.data();
320 buffers.rss = output.rss.data();
321
322 rss(row, mat, buffers, opt);
323 return output;
324}
325
326}
327
328#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_unsafe(const Count_ num_total, const Float_ old_rss, const Float_ old_mean, const Float_ new_mean)
void update_rss_with_zeros_unsafe(Output_ &mean, Output_ &rss, const Count_ num_zeros, const Count_ num_total)
constexpr Value_ nan_if_available_else_zero()
void update_rss(Output_ &mean, Output_ &rss, const Input_ value, const Count_ num_total)
Functions to compute statistics from a tatami::Matrix.
Definition count.hpp:20
void rss(bool row, const tatami::Matrix< Value_, Index_ > &mat, RssBuffers< Output_ > &output, const RssOptions< Output_ > &opt)
Definition rss.hpp:261
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)
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 rss().
Definition rss.hpp:50
Output_ * mean
Definition rss.hpp:55
Output_ * rss
Definition rss.hpp:61
Options for rss().
Definition rss.hpp:30
Output_ mean_placeholder
Definition rss.hpp:41
int num_threads
Definition rss.hpp:35
Results of rss().
Definition rss.hpp:275
std::vector< Output_ > rss
Definition rss.hpp:286
std::vector< Output_ > mean
Definition rss.hpp:280