1#ifndef TATAMI_STATS_SKIP_NAN_RSS_HPP
2#define TATAMI_STATS_SKIP_NAN_RSS_HPP
15#include "sanisizer/sanisizer.hpp"
32template<
typename Output_ =
double>
54template<
typename Output_,
typename Count_>
78template<
typename Value_,
typename Index_,
typename Output_,
typename Count_>
80 const auto dim = (row ? mat.
nrow() : mat.
ncol());
81 const auto otherdim = (row ? mat.
ncol() : mat.
nrow());
93 for (Index_ x = 0; x < l; ++x) {
94 auto out = ext->fetch(vbuffer.data(), NULL);
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;
110 for (Index_ x = 0; x < l; ++x) {
111 auto out = ext->fetch(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;
123template<
typename Value_,
typename Index_,
typename Output_,
typename Count_>
125 const auto dim = (row ? mat.
nrow() : mat.
ncol());
126 const auto otherdim = (row ? mat.
ncol() : mat.
nrow());
129 std::fill_n(output.rss, dim, 0);
130 std::fill_n(output.count, dim, 0);
132 std::fill_n(output.mean, dim, opt.mean_placeholder);
135 std::fill_n(output.mean, dim, 0);
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;
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));
153 std::optional<std::vector<Output_> > cur_rss, cur_mean;
154 std::optional<std::vector<Count_> > cur_count;
158 rss_ptr = output.rss;
159 mean_ptr = output.mean;
160 count_ptr = output.count;
165 mean_ptr = cur_mean->data();
167 count_ptr = cur_count->data();
169 rss_ptr = output.rss;
172 rss_ptr = cur_rss->data();
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];
198 for (Index_ d = 0; d < dim; ++d) {
199 auto& unskipped_total = count_ptr[d];
200 unskipped_total = l - unskipped_total;
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)) {
221 (*all_partial_mean)[thread] = std::move(cur_mean);
222 (*all_partial_count)[thread] = std::move(cur_count);
224 (*all_partial_rss)[thread - 1] = std::move(cur_rss);
227 }, otherdim, opt.num_threads);
233 const auto& ap_mean = *all_partial_mean;
234 const auto& ap_rss = *all_partial_rss;
235 const auto& ap_count = *all_partial_count;
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];
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) {
251 const auto mult =
static_cast<Output_
>(cur_count[d]) /
static_cast<Output_
>(output.count[d]);
252 output.mean[d] += cur_mean[d] * mult;
258 for (
int u = 0; u < nused; ++u) {
259 const auto& cur_count = *(ap_count[u]);
260 const auto& cur_mean = *(ap_mean[u]);
262 for (Index_ d = 0; d < dim; ++d) {
266 const auto& cur_rss = *(ap_rss[u - 1]);
267 for (Index_ d = 0; d < dim; ++d) {
274 for (Index_ d = 0; d < dim; ++ d) {
275 if (output.count[d] == 0) {
276 output.mean[d] = opt.mean_placeholder;
302template<
typename Value_,
typename Index_,
typename Output_,
typename Count_>
305 rss_direct(row, mat, output, opt);
307 rss_running(row, mat, output, opt);
318template<
typename Output_,
typename Count_>
355template<
typename Output_ =
double,
typename Count_,
typename Value_,
typename Index_>
358 const auto dim = (row ? mat.
nrow() : mat.
ncol());
361#ifdef TATAMI_STATS_TEST_DIRTY
366#ifdef TATAMI_STATS_TEST_DIRTY
371#ifdef TATAMI_STATS_TEST_DIRTY
378 buffers.
rss = output.
rss.data();
381 rss(row, mat, buffers, opt);
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