1#ifndef TATAMI_STATS_GROUP_RSS_HPP
2#define TATAMI_STATS_GROUP_RSS_HPP
15#include "sanisizer/sanisizer.hpp"
17#include "jiwoo/jiwoo.hpp"
31template<
typename Output_ =
double>
51template<
typename Output_>
58 std::vector<Output_*>
mean;
65 std::vector<Output_*>
rss;
71template<
typename Group_,
typename Count_,
typename Output_,
typename Index_>
72void group_rss_finish_means(
73 const Group_ num_groups,
74 const Count_*
const group_size,
75 std::vector<Output_>& means,
77 const std::vector<Output_*>& output_means,
78 const Output_ placeholder
80 for (Group_ g = 0; g < num_groups; ++g) {
82 means[g] /= group_size[g];
84 means[g] = placeholder;
86 output_means[g][i] = means[g];
90template<
typename Value_,
typename Index_,
typename Group_,
typename Count_,
typename Output_>
94 const Group_*
const group,
95 const Group_ num_groups,
96 const Count_*
const group_size,
97 const GroupRssBuffers<Output_>& output,
98 const GroupRssOptions<Output_>& opt
100 const auto dim = (row ? mat.
nrow() : mat.
ncol());
101 const auto otherdim = (row ? mat.
ncol() : mat.
nrow());
108 auto cur_means = sanisizer::create<std::vector<Output_> >(num_groups);
109 auto cur_rss = sanisizer::create<std::vector<Output_> >(num_groups);
110 auto cur_non_zeros = sanisizer::create<std::vector<Index_> >(num_groups);
112 for (Index_ x = 0; x < l; ++x) {
113 auto range = ext->fetch(vbuffer.data(), ibuffer.data());
116 for (Index_ i = 0; i <
range.number; ++i) {
117 const auto g = group[
range.index[i]];
118 cur_means[g] +=
range.value[i];
121 group_rss_finish_means(num_groups, group_size, cur_means,
static_cast<Index_
>(s + x), output.mean, opt.mean_placeholder);
124 for (Index_ i = 0; i <
range.number; ++i) {
125 const auto g = group[
range.index[i]];
126 const auto delta =
range.value[i] - cur_means[g];
127 cur_rss[g] += delta * delta;
129 for (Group_ g = 0; g < num_groups; ++g) {
130 if (group_size[g] > 0) {
131 const Output_ my_rss = cur_rss[g] + cur_means[g] * cur_means[g] * (group_size[g] - cur_non_zeros[g]);
132 output.rss[g][s + x] = my_rss;
134 output.rss[g][s + x] = 0;
138 std::fill(cur_means.begin(), cur_means.end(), 0);
139 std::fill(cur_rss.begin(), cur_rss.end(), 0);
140 std::fill(cur_non_zeros.begin(), cur_non_zeros.end(), 0);
142 }, dim, opt.num_threads);
148 auto cur_means = sanisizer::create<std::vector<Output_> >(num_groups);
149 auto cur_rss = sanisizer::create<std::vector<Output_> >(num_groups);
151 for (Index_ x = 0; x < l; ++x) {
152 auto ptr = ext->fetch(buffer.data());
155 for (Index_ j = 0; j < otherdim; ++j) {
156 cur_means[group[j]] += ptr[j];
158 group_rss_finish_means(num_groups, group_size, cur_means,
static_cast<Index_
>(s + x), output.mean, opt.mean_placeholder);
161 for (Index_ j = 0; j < otherdim; ++j) {
162 const auto g = group[j];
163 const auto delta = ptr[j] - cur_means[g];
164 cur_rss[g] += delta * delta;
166 for (Group_ g = 0; g < num_groups; ++g) {
167 output.rss[g][s + x] = cur_rss[g];
170 std::fill(cur_means.begin(), cur_means.end(), 0);
171 std::fill(cur_rss.begin(), cur_rss.end(), 0);
173 }, dim, opt.num_threads);
177template<
typename Value_,
typename Index_,
typename Group_,
typename Count_,
typename Output_>
178void group_rss_running_nonempty(
181 const Index_ otherdim,
183 const Group_*
const group,
184 const Group_ num_groups,
185 const Count_*
const group_size,
186 const GroupRssBuffers<Output_>& output,
187 const GroupRssOptions<Output_>& opt
189 const bool do_parallel = opt.num_threads > 1;
190 std::optional<std::vector<std::optional<jiwoo::EquilengthArrays<Output_> > > > all_partial_mean, all_partial_rss;
191 std::optional<std::vector<std::optional<std::vector<Count_> > > > all_partial_count;
194 all_partial_rss.emplace(sanisizer::cast<I<
decltype(all_partial_rss->size())> >(opt.num_threads - 1));
195 all_partial_mean.emplace(sanisizer::cast<I<
decltype(all_partial_mean->size())> >(opt.num_threads));
196 all_partial_count.emplace(sanisizer::cast<I<
decltype(all_partial_count->size())> >(opt.num_threads));
201 for (Group_ g = 0; g < num_groups; ++g) {
202 assert(group_size[g] > 0);
208 for (Group_ g = 0; g < num_groups; ++g) {
209 std::fill_n(output.mean[g], dim, 0);
212 for (Group_ g = 0; g < num_groups; ++g) {
213 std::fill_n(output.rss[g], dim, 0);
218 std::optional<jiwoo::EquilengthArrays<Output_> > cur_mean, cur_rss;
220 Output_*
const * mean_ptrs;
221 Output_*
const * rss_ptrs;
224 mean_ptrs = output.mean.data();
225 rss_ptrs = output.rss.data();
230 rss_ptrs = output.rss.data();
233 sanisizer::cast<I<
decltype(cur_rss->size())> >(num_groups),
234 static_cast<std::size_t
>(dim),
237 rss_ptrs = cur_rss->get();
242 sanisizer::cast<I<
decltype(cur_mean->size())> >(num_groups),
243 static_cast<std::size_t
>(dim),
246 mean_ptrs = cur_mean->get();
249 auto cur_count = sanisizer::create<std::vector<Count_> >(num_groups);
255 auto nonzeros = sanisizer::create<std::vector<std::vector<Index_> > >(num_groups);
256 for (Group_ g = 0; g < num_groups; ++g) {
260 for (Index_ x = 0; x < l; ++x) {
261 auto out = ext->fetch(vbuffer.data(), ibuffer.data());
262 const auto grp = group[s + x];
264 const auto mptr = mean_ptrs[grp];
265 const auto rptr = rss_ptrs[grp];
266 auto& nnz = nonzeros[grp];
267 for (Index_ i = 0; i < out.number; ++i) {
268 const auto d = out.index[i];
273 for (Group_ g = 0; g < num_groups; ++g) {
274 const auto curtotal = cur_count[g];
276 const auto mptr = mean_ptrs[g];
277 const auto rptr = rss_ptrs[g];
278 const auto& nnz = nonzeros[g];
279 for (Index_ d = 0; d < dim; ++d) {
290 for (Index_ x = 0; x < l; ++x) {
291 auto out = ext->fetch(buffer.data());
292 const auto grp = group[s + x];
294 const auto mptr = mean_ptrs[grp];
295 const auto rptr = rss_ptrs[grp];
296 for (Index_ d = 0; d < dim; ++d) {
303 (*all_partial_count)[thread] = std::move(cur_count);
304 (*all_partial_mean)[thread] = std::move(cur_mean);
306 (*all_partial_rss)[thread - 1] = std::move(cur_rss);
309 }, otherdim, opt.num_threads);
313 const auto& ap_mean = *all_partial_mean;
314 const auto& ap_rss = *all_partial_rss;
316 for (Group_ g = 0; g < num_groups; ++g) {
317 const auto cur_output = output.mean[g];
318 const auto cur_global_count = group_size[g];
319 assert(cur_global_count > 0);
320 bool initialized =
false;
322 for (
int u = 0; u < nused; ++u) {
323 const auto cur_count = (*((*all_partial_count)[u]))[g];
324 if (cur_count == 0) {
328 const auto cur_mean = (*(ap_mean[u]))[g];
329 const Output_ mult =
static_cast<Output_
>(cur_count) /
static_cast<Output_
>(cur_global_count);
331 for (Index_ d = 0; d < dim; ++d) {
332 cur_output[d] = cur_mean[d] * mult;
336 for (Index_ d = 0; d < dim; ++d) {
337 cur_output[d] += cur_mean[d] * mult;
346 for (Group_ g = 0; g < num_groups; ++g) {
347 const auto cur_global = output.mean[g];
348 const auto cur_output = output.rss[g];
349 bool initialized =
false;
351 for (
int u = 0; u < nused; ++u) {
352 const auto cur_count = (*((*all_partial_count)[u]))[g];
353 if (cur_count == 0) {
357 const auto cur_mean = (*(ap_mean[u]))[g];
359 for (Index_ d = 0; d < dim; ++d) {
364 const auto cur_rss = (*(ap_rss[u - 1]))[g];
366 for (Index_ d = 0; d < dim; ++d) {
371 for (Index_ d = 0; d < dim; ++d) {
383template<
typename Value_,
typename Index_,
typename Group_,
typename Count_,
typename Output_>
384void group_rss_running(
387 const Group_*
const group,
388 const Group_ num_groups,
389 const Count_*
const group_size,
390 const GroupRssBuffers<Output_>& output,
391 const GroupRssOptions<Output_>& opt
393 const auto dim = (row ? mat.
nrow() : mat.
ncol());
394 const auto otherdim = (row ? mat.
ncol() : mat.
nrow());
396 for (Group_ g = 0; g < num_groups; ++g) {
397 std::fill_n(output.mean[g], dim, opt.mean_placeholder);
398 std::fill_n(output.rss[g], dim, 0);
403 std::size_t num_empty = 0;
404 for (Group_ g = 0; g < num_groups; ++g) {
405 num_empty += (group_size[g] == 0);
409 std::optional<std::vector<Count_> > new_group_size_store;
410 std::optional<GroupRssBuffers<Output_> > new_output_store;
411 std::optional<std::vector<Group_> > new_group_store;
412 Group_ num_non_empty;
413 const Count_* new_group_size;
414 const Group_* new_group;
415 const GroupRssBuffers<Output_>* new_output;
418 num_non_empty = num_groups - num_empty;
419 new_group_size_store.emplace();
420 new_group_size_store->reserve(num_non_empty);
421 new_output_store.emplace();
422 new_output_store->mean.reserve(num_non_empty);
423 new_output_store->rss.reserve(num_non_empty);
425 auto mapping = sanisizer::create<std::vector<std::size_t> >(num_groups);
426 for (Group_ g = 0; g < num_groups; ++g) {
428 mapping[g] = new_group_size_store->size();
429 new_group_size_store->push_back(group_size[g]);
430 new_output_store->mean.push_back(output.mean[g]);
431 new_output_store->rss.push_back(output.rss[g]);
433 std::fill_n(output.mean[g], dim, opt.mean_placeholder);
434 std::fill_n(output.rss[g], dim, 0);
439 for (Index_ i = 0; i < otherdim; ++i) {
440 (*new_group_store)[i] = mapping[group[i]];
443 new_group_size = new_group_size_store->data();
444 new_output = &(*new_output_store);
445 new_group = new_group_store->data();
447 num_non_empty = num_groups;
448 new_group_size = group_size;
449 new_output = &output;
453 group_rss_running_nonempty(
489template<
typename Value_,
typename Index_,
typename Group_,
typename Count_,
typename Output_>
493 const Group_*
const group,
494 const Group_ num_groups,
495 const Count_*
const group_size,
499 assert(sanisizer::is_equal(num_groups, output.
mean.size()));
500 assert(sanisizer::is_equal(num_groups, output.
rss.size()));
502 group_rss_direct(row, mat, group, num_groups, group_size, output, opt);
504 group_rss_running(row, mat, group, num_groups, group_size, output, opt);
527template<
typename Value_,
typename Index_,
typename Group_,
typename Output_>
531 const Group_*
const group,
532 const Group_ num_groups,
536 auto group_size = sanisizer::create<std::vector<Index_> >(num_groups);
537 const auto otherdim = (row ? mat.
ncol() : mat.
nrow());
538 for (Index_ o = 0; o < otherdim; ++o) {
539 group_size[group[o]] += 1;
541 group_rss(row, mat, group, num_groups, group_size.data(), output, opt);
549template<
typename Output_>
556 std::vector<std::vector<Output_> >
mean;
563 std::vector<std::vector<Output_> >
rss;
584template<
typename Output_,
typename Value_,
typename Index_,
typename Group_>
588 const Group_*
const group,
589 const Group_ num_groups,
593 sanisizer::resize(output.
mean, num_groups);
594 sanisizer::resize(output.
rss, num_groups);
597 sanisizer::resize(buffers.
mean, num_groups);
598 sanisizer::resize(buffers.
rss, num_groups);
600 const auto dim = (row ? mat.
nrow() : mat.
ncol());
601 for (Group_ g = 0; g < num_groups; ++g) {
603#ifdef TATAMI_STATS_TEST_DIRTY
607 buffers.
mean[g] = output.
mean[g].data();
609#ifdef TATAMI_STATS_TEST_DIRTY
613 buffers.
rss[g] = output.
rss[g].data();
616 group_rss(row, mat, group, num_groups, 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
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 range(bool row, const tatami::Matrix< Value_, Index_ > &mat, RangeBuffers< Output_ > &output, const RangeOptions< Output_ > &opt)
Definition range.hpp:339
void group_rss(bool row, const tatami::Matrix< Value_, Index_ > &mat, const Group_ *const group, const Group_ num_groups, const Count_ *const group_size, const GroupRssBuffers< Output_ > &output, const GroupRssOptions< Output_ > &opt)
Definition group_rss.hpp:490
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)