1#ifndef TATAMI_STATS_RSS_HPP
2#define TATAMI_STATS_RSS_HPP
14#include "sanisizer/sanisizer.hpp"
29template<
typename Output_ =
double>
49template<
typename Output_>
67template<
typename Value_,
typename Index_,
typename Output_>
69 const auto dim = (row ? mat.
nrow() : mat.
ncol());
70 const auto otherdim = (row ? mat.
ncol() : mat.
nrow());
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;
95 for (Index_ x = 0; x < l; ++x) {
96 auto out = ext->fetch(buffer.data());
98 output.
mean[x + s] = res.mean;
99 output.
rss[x + s] = res.rss;
105template<
typename Value_,
typename Index_,
typename Output_>
107 const auto dim = (row ? mat.
nrow() : mat.
ncol());
108 const auto otherdim = (row ? mat.
ncol() : mat.
nrow());
110 std::fill_n(output.mean, dim, opt.mean_placeholder);
111 std::fill_n(output.rss, dim, 0);
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;
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));
129 std::fill_n(output.mean, dim, 0);
131 std::fill_n(output.rss, dim, 0);
137 std::optional<std::vector<Output_> > cur_rss, cur_mean;
141 rss_ptr = output.rss;
142 mean_ptr = output.mean;
147 mean_ptr = cur_mean->data();
149 rss_ptr = output.rss;
152 rss_ptr = cur_rss->data();
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];
173 for (Index_ d = 0; d < dim; ++d) {
182 for (Index_ x = 0; x < l; ++x) {
183 auto out = ext->fetch(buffer.data());
184 for (Index_ d = 0; d < dim; ++d) {
191 (*all_partial_count)[thread] = l;
192 (*all_partial_mean)[thread] = std::move(cur_mean);
194 (*all_partial_rss)[thread - 1] = std::move(cur_rss);
197 }, otherdim, opt.num_threads);
203 const auto& ap_count = *all_partial_count;
204 const auto& ap_mean = *all_partial_mean;
205 const auto& ap_rss = *all_partial_rss;
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]);
212 for (Index_ d = 0; d < dim; ++d) {
213 output.mean[d] = cur_mean[d] * mult;
216 for (Index_ d = 0; d < dim; ++d) {
217 output.mean[d] += cur_mean[d] * mult;
224 for (
int u = 0; u < nused; ++u) {
225 const auto cur_count = ap_count[u];
226 const auto& cur_mean = *(ap_mean[u]);
228 for (Index_ d = 0; d < dim; ++d) {
232 const auto& cur_rss = *(ap_rss[u - 1]);
233 for (Index_ d = 0; d < dim; ++d) {
260template<
typename Value_,
typename Index_,
typename Output_>
263 rss_direct(row, mat, output, opt);
265 rss_running(row, mat, output, opt);
274template<
typename Output_>
303template<
typename Output_ =
double,
typename Value_,
typename Index_>
306 const auto dim = (row ? mat.
nrow() : mat.
ncol());
308#ifdef TATAMI_STATS_TEST_DIRTY
313#ifdef TATAMI_STATS_TEST_DIRTY
320 buffers.
rss = output.
rss.data();
322 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_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