tatami_stats
Matrix statistics for tatami
Loading...
Searching...
No Matches
group_variance.hpp
Go to the documentation of this file.
1#ifndef TATAMI_STATS_GROUP_VARIANCE_HPP
2#define TATAMI_STATS_GROUP_VARIANCE_HPP
3
4#include <vector>
5#include <algorithm>
6#include <cstddef>
7#include <optional>
8#include <cassert>
9#include <limits>
10#include <cmath>
11
12#include "tatami/tatami.hpp"
13#include "sanisizer/sanisizer.hpp"
14
15#include "group_rss.hpp"
17#include "utils.hpp"
18
25namespace tatami_stats {
26
31template<typename Output_ = double>
57
63template<typename Output_>
70 std::vector<Output_*> mean;
71
77 std::vector<Output_*> variance;
78};
79
100template<typename Value_, typename Index_, typename Group_, typename Count_, typename Output_>
102 bool row,
104 const Group_* const group,
105 const Group_ num_groups,
106 const Count_* const group_size,
107 const GroupVarianceBuffers<Output_>& output,
109) {
110 assert(sanisizer::is_equal(num_groups, output.mean.size()));
111 assert(sanisizer::is_equal(num_groups, output.variance.size()));
112 const auto dim = (row ? mat.nrow() : mat.ncol());
113
114 nanable_ifelse<Value_>(
115 opt.skip_nan,
116
117 [&]() -> void {
118 skip_nan::GroupRssBuffers<Output_, Index_> tmp;
119 tmp.mean = output.mean;
120 tmp.rss = output.variance;
121
122 auto count = sanisizer::create<std::vector<std::vector<Index_> > >(num_groups);
123 tmp.count.reserve(num_groups);
124 for (Group_ g = 0; g < num_groups; ++g) {
125 tatami::resize_container_to_Index_size(count[g], dim);
126 tmp.count.push_back(count[g].data());
127 }
128
130 ropt.num_threads = opt.num_threads;
131 skip_nan::group_rss(row, mat, group, num_groups, tmp, ropt);
132 for (Group_ g = 0; g < num_groups; ++g) {
133 const auto outvar = output.variance[g];
134 const auto curcounts = count[g];
135 for (Index_ d = 0; d < dim; ++d) {
136 if (curcounts[d] <= 1) {
137 outvar[d] = opt.variance_placeholder;
138 } else {
139 outvar[d] /= curcounts[d] - 1;
140 }
141 }
142 }
143 },
144
145 [&]() -> void {
146 GroupRssBuffers<Output_> tmp;
147 tmp.mean = output.mean;
148 tmp.rss = output.variance;
149
150 GroupRssOptions ropt;
151 ropt.num_threads = opt.num_threads;
152 group_rss(row, mat, group, num_groups, group_size, tmp, ropt);
153
154 for (Group_ g = 0; g < num_groups; ++g) {
155 const auto outvar = output.variance[g];
156 const auto gsize = group_size[g];
157 if (gsize <= 1) {
158 std::fill_n(outvar, dim, opt.variance_placeholder);
159 } else {
160 for (Index_ d = 0; d < dim; ++d) {
161 outvar[d] /= gsize - 1;
162 }
163 }
164 }
165 }
166 );
167}
168
187template<typename Value_, typename Index_, typename Group_, typename Output_>
189 bool row,
191 const Group_* const group,
192 const Group_ num_groups,
193 const GroupVarianceBuffers<Output_>& output,
195) {
196 auto group_size = sanisizer::create<std::vector<Index_> >(num_groups);
197 const auto otherdim = (row ? mat.ncol() : mat.nrow());
198 for (Index_ o = 0; o < otherdim; ++o) {
199 group_size[group[o]] += 1;
200 }
201 group_variance(row, mat, group, num_groups, group_size.data(), output, opt);
202}
203
209template<typename Output_>
216 std::vector<std::vector<Output_> > mean;
217
223 std::vector<std::vector<Output_> > variance;
224};
225
244template<typename Output_ = double, typename Value_, typename Index_, typename Group_>
246 bool row,
248 const Group_* const group,
249 const Group_ num_groups,
251) {
253 sanisizer::resize(output.mean, num_groups);
254 sanisizer::resize(output.variance, num_groups);
255
257 sanisizer::resize(buffers.mean, num_groups);
258 sanisizer::resize(buffers.variance, num_groups);
259 const auto dim = (row ? mat.nrow() : mat.ncol());
260
261 for (Group_ g = 0; g < num_groups; ++g) {
263#ifdef TATAMI_STATS_TEST_DIRTY
264 , -1
265#endif
266 );
267 buffers.mean[g] = output.mean[g].data();
269#ifdef TATAMI_STATS_TEST_DIRTY
270 , -1
271#endif
272 );
273 buffers.variance[g] = output.variance[g].data();
274 }
275
276 group_variance(row, mat, group, num_groups, buffers, opt);
277 return output;
278}
279
280}
281
282#endif
virtual Index_ ncol() const=0
virtual Index_ nrow() const=0
Compute group-wise residual sum of squares from a tatami::Matrix.
constexpr Value_ nan_if_available_else_zero()
void group_rss(bool row, const tatami::Matrix< Value_, Index_ > &mat, const Group_ *const group, const Group_ num_groups, const GroupRssBuffers< Output_, Count_ > &output, const GroupRssOptions< Output_ > &opt)
Definition group_rss.hpp:446
Functions to compute statistics from a tatami::Matrix.
Definition count.hpp:20
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 group_variance(bool row, const tatami::Matrix< Value_, Index_ > &mat, const Group_ *const group, const Group_ num_groups, const Count_ *const group_size, const GroupVarianceBuffers< Output_ > &output, const GroupVarianceOptions< Output_ > &opt)
Definition group_variance.hpp:101
void count(const bool row, const tatami::Matrix< Value_, Index_ > &mat, Output_ *const output, Condition_ condition, const CountOptions &opt)
Definition count.hpp:188
void resize_container_to_Index_size(Container_ &container, const Index_ x, Args_ &&... args)
Compute group-wise residual sum of squares while skipping NaNs.
Result buffers for group_variance().
Definition group_variance.hpp:64
std::vector< Output_ * > mean
Definition group_variance.hpp:70
std::vector< Output_ * > variance
Definition group_variance.hpp:77
Options for group_variance().
Definition group_variance.hpp:32
Output_ variance_placeholder
Definition group_variance.hpp:55
Output_ mean_placeholder
Definition group_variance.hpp:49
bool skip_nan
Definition group_variance.hpp:37
int num_threads
Definition group_variance.hpp:43
Results of group_variance().
Definition group_variance.hpp:210
std::vector< std::vector< Output_ > > mean
Definition group_variance.hpp:216
std::vector< std::vector< Output_ > > variance
Definition group_variance.hpp:223
Options for skip_nan::group_rss().
Definition group_rss.hpp:34
int num_threads
Definition group_rss.hpp:39