33template<
typename Output_ =
double>
55template<
typename Output_,
typename Count_>
62 std::vector<Output_*>
mean;
69 std::vector<Output_*>
rss;
82template<
typename Value_,
typename Index_,
typename Group_,
typename Output_,
typename Count_>
86 const Group_*
const group,
87 const Group_ num_groups,
91 const auto dim = (row ? mat.
nrow() : mat.
ncol());
92 const auto otherdim = (row ? mat.
ncol() : mat.
nrow());
95 auto full_group_sizes = sanisizer::create<std::vector<Index_> >(num_groups);
96 for (Index_ i = 0; i < otherdim; ++i) {
97 full_group_sizes[group[i]] += 1;
104 auto cur_means = sanisizer::create<std::vector<Output_> >(num_groups);
105 auto cur_rss = sanisizer::create<std::vector<Output_> >(num_groups);
106 auto cur_non_zeros = sanisizer::create<std::vector<Index_> >(num_groups);
107 auto cur_sizes = sanisizer::create<std::vector<Index_> >(num_groups);
109 for (Index_ x = 0; x < l; ++x) {
110 auto range = ext->fetch(vbuffer.data(), ibuffer.data());
113 for (Index_ i = 0; i <
range.number; ++i) {
114 const auto val =
range.value[i];
115 const auto b = group[
range.index[i]];
116 if (!std::isnan(val)) {
123 for (Group_ g = 0; g < num_groups; ++g) {
124 const auto actual_size = full_group_sizes[g] - cur_sizes[g];
125 cur_sizes[g] = actual_size;
126 output.
count[g][s + x] = actual_size;
128 group_rss_finish_means(num_groups, cur_sizes.data(), cur_means,
static_cast<Index_
>(s + x), output.
mean, opt.
mean_placeholder);
131 for (Index_ i = 0; i <
range.number; ++i) {
132 const auto val =
range.value[i];
133 if (!std::isnan(val)) {
134 const auto g = group[
range.index[i]];
135 const auto delta = val - cur_means[g];
136 cur_rss[g] += delta * delta;
139 for (Group_ g = 0; g < num_groups; ++g) {
140 if (cur_sizes[g] > 0) {
141 const Output_ my_rss = cur_rss[g] + cur_means[g] * cur_means[g] * (cur_sizes[g] - cur_non_zeros[g]);
142 output.
rss[g][s + x] = my_rss;
144 output.
rss[g][s + x] = 0;
148 std::fill(cur_means.begin(), cur_means.end(), 0);
149 std::fill(cur_rss.begin(), cur_rss.end(), 0);
150 std::fill(cur_non_zeros.begin(), cur_non_zeros.end(), 0);
151 std::fill(cur_sizes.begin(), cur_sizes.end(), 0);
159 auto cur_means = sanisizer::create<std::vector<Output_> >(num_groups);
160 auto cur_rss = sanisizer::create<std::vector<Output_> >(num_groups);
161 auto cur_sizes = sanisizer::create<std::vector<Index_> >(num_groups);
163 for (Index_ x = 0; x < l; ++x) {
164 auto ptr = ext->fetch(buffer.data());
167 for (Index_ j = 0; j < otherdim; ++j) {
168 const auto val = ptr[j];
169 if (!std::isnan(val)) {
170 const auto g = group[j];
175 for (Group_ g = 0; g < num_groups; ++g) {
176 output.
count[g][s + x] = cur_sizes[g];
178 group_rss_finish_means(num_groups, cur_sizes.data(), cur_means,
static_cast<Index_
>(s + x), output.
mean, opt.
mean_placeholder);
181 for (Index_ j = 0; j < otherdim; ++j) {
182 const auto val = ptr[j];
183 if (!std::isnan(val)) {
184 const auto g = group[j];
185 const auto delta = val - cur_means[g];
186 cur_rss[g] += delta * delta;
189 for (Group_ g = 0; g < num_groups; ++g) {
190 output.
rss[g][s + x] = cur_rss[g];
193 std::fill(cur_means.begin(), cur_means.end(), 0);
194 std::fill(cur_rss.begin(), cur_rss.end(), 0);
195 std::fill(cur_sizes.begin(), cur_sizes.end(), 0);
201template<
typename Value_,
typename Index_,
typename Group_,
typename Output_,
typename Count_>
202void group_rss_running(
205 const Group_*
const group,
206 const Group_ num_groups,
207 const GroupRssBuffers<Output_, Count_>& output,
208 const GroupRssOptions<Output_>& opt
210 const auto dim = (row ? mat.
nrow() : mat.
ncol());
211 for (Group_ g = 0; g < num_groups; ++g) {
212 std::fill_n(output.rss[g], dim, 0);
213 std::fill_n(output.count[g], dim, 0);
216 const auto otherdim = (row ? mat.
ncol() : mat.
nrow());
218 for (Group_ g = 0; g < num_groups; ++g) {
219 std::fill_n(output.mean[g], dim, opt.mean_placeholder);
223 for (Group_ g = 0; g < num_groups; ++g) {
224 std::fill_n(output.mean[g], dim, 0);
228 const bool do_parallel = opt.num_threads > 1;
229 std::optional<std::vector<std::optional<jiwoo::EquilengthArrays<Output_> > > > all_partial_mean, all_partial_rss;
230 std::optional<std::vector<std::optional<jiwoo::EquilengthArrays<Count_> > > > all_partial_count;
233 all_partial_rss.emplace(sanisizer::cast<I<
decltype(all_partial_rss->size())> >(opt.num_threads - 1));
234 all_partial_mean.emplace(sanisizer::cast<I<
decltype(all_partial_mean->size())> >(opt.num_threads));
235 all_partial_count.emplace(sanisizer::cast<I<
decltype(all_partial_count->size())> >(opt.num_threads));
240 std::optional<jiwoo::EquilengthArrays<Output_> > cur_mean, cur_rss;
241 std::optional<jiwoo::EquilengthArrays<Count_> > cur_count;
243 Output_*
const * mean_ptrs;
244 Output_*
const * rss_ptrs;
245 Count_*
const * count_ptrs;
248 mean_ptrs = output.mean.data();
249 rss_ptrs = output.rss.data();
250 count_ptrs = output.count.data();
255 rss_ptrs = output.rss.data();
258 sanisizer::cast<I<
decltype(cur_rss->size())> >(num_groups),
259 static_cast<std::size_t
>(dim),
262 rss_ptrs = cur_rss->get();
267 sanisizer::cast<I<
decltype(cur_mean->size())> >(num_groups),
268 static_cast<std::size_t
>(dim),
271 mean_ptrs = cur_mean->get();
275 sanisizer::cast<I<
decltype(cur_count->size())> >(num_groups),
276 static_cast<std::size_t
>(dim),
279 count_ptrs = cur_count->get();
286 auto nonzeros = sanisizer::create<std::vector<std::vector<Count_> > >(num_groups);
287 for (Group_ g = 0; g < num_groups; ++g) {
291 auto cur_group_size = sanisizer::create<std::vector<Count_> >(num_groups);
292 for (Index_ x = 0; x < l; ++x) {
293 auto out = ext->fetch(vbuffer.data(), ibuffer.data());
294 const auto grp = group[s + x];
295 ++cur_group_size[grp];
297 const auto mptr = mean_ptrs[grp];
298 const auto rptr = rss_ptrs[grp];
299 const auto cptr = count_ptrs[grp];
300 auto& nnz = nonzeros[grp];
302 for (Index_ i = 0; i < out.number; ++i) {
303 const auto d = out.index[i];
304 const auto val = out.value[i];
305 if (!std::isnan(val)) {
313 for (Group_ g = 0; g < num_groups; ++g) {
314 const auto mptr = mean_ptrs[g];
315 const auto rptr = rss_ptrs[g];
316 const auto cptr = count_ptrs[g];
317 const auto& nnz = nonzeros[g];
318 const auto curtotal = cur_group_size[g];
320 for (Index_ d = 0; d < dim; ++d) {
321 auto& unskipped_total = cptr[d];
322 unskipped_total = curtotal - unskipped_total;
331 for (Index_ x = 0; x < l; ++x) {
332 auto out = ext->fetch(buffer.data());
333 const auto grp = group[s + x];
334 const auto mptr = mean_ptrs[grp];
335 const auto rptr = rss_ptrs[grp];
336 const auto cptr = count_ptrs[grp];
338 for (Index_ d = 0; d < dim; ++d) {
339 const auto val = out[d];
340 if (!std::isnan(val)) {
348 (*all_partial_count)[thread] = std::move(cur_count);
349 (*all_partial_mean)[thread] = std::move(cur_mean);
351 (*all_partial_rss)[thread - 1] = std::move(cur_rss);
354 }, otherdim, opt.num_threads);
358 const auto& ap_mean = *all_partial_mean;
359 const auto& ap_rss = *all_partial_rss;
360 const auto& ap_count = *all_partial_count;
362 for (Group_ g = 0; g < num_groups; ++g) {
363 const auto cur_global_count = output.count[g];
364 for (
int u = 0; u < nused; ++u) {
365 const auto& cur_count = (*(ap_count[u]))[g];
366 for (Index_ d = 0; d < dim; ++d) {
367 cur_global_count[d] += cur_count[d];
373 for (Group_ g = 0; g < num_groups; ++g) {
374 const auto cur_global_count = output.count[g];
375 const auto cur_global_mean = output.mean[g];
377 for (
int u = 0; u < nused; ++u) {
378 const auto& cur_mean = (*(ap_mean[u]))[g];
379 const auto& cur_count = (*(ap_count[u]))[g];
380 for (Index_ d = 0; d < dim; ++d) {
381 if (cur_count[d] > 0) {
382 const auto mult =
static_cast<Output_
>(cur_count[d]) /
static_cast<Output_
>(cur_global_count[d]);
383 cur_global_mean[d] += cur_mean[d] * mult;
391 for (Group_ g = 0; g < num_groups; ++g) {
392 const auto cur_global_mean = output.mean[g];
393 const auto cur_output = output.rss[g];
394 for (
int u = 0; u < nused; ++u) {
395 const auto& cur_mean = (*(ap_mean[u]))[g];
396 const auto& cur_count = (*(ap_count[u]))[g];
398 for (Index_ d = 0; d < dim; ++d) {
402 const auto& cur_rss = (*(ap_rss[u - 1]))[g];
403 for (Index_ d = 0; d < dim; ++d) {
411 for (Group_ g = 0; g < num_groups; ++g) {
412 const auto mptr = output.mean[g];
413 const auto cptr = output.count[g];
414 for (Index_ d = 0; d < dim; ++d) {
416 mptr[d] = opt.mean_placeholder;
445template<
typename Value_,
typename Index_,
typename Group_,
typename Output_,
typename Count_>
449 const Group_*
const group,
450 const Group_ num_groups,
454 assert(sanisizer::is_equal(num_groups, output.
mean.size()));
455 assert(sanisizer::is_equal(num_groups, output.
rss.size()));
456 assert(sanisizer::is_equal(num_groups, output.
count.size()));
458 group_rss_direct(row, mat, group, num_groups, output, opt);
460 group_rss_running(row, mat, group, num_groups, output, opt);
471template<
typename Output_,
typename Count_>
478 std::vector<std::vector<Output_> >
mean;
485 std::vector<std::vector<Output_> >
rss;
492 std::vector<std::vector<Count_> >
count;
515template<
typename Output_,
typename Count_,
typename Value_,
typename Index_,
typename Group_>
519 const Group_*
const group,
520 const Group_ num_groups,
524 sanisizer::resize(output.
mean, num_groups);
525 sanisizer::resize(output.
rss, num_groups);
526 sanisizer::resize(output.
count, num_groups);
529 sanisizer::resize(buffers.
mean, num_groups);
530 sanisizer::resize(buffers.
rss, num_groups);
531 sanisizer::resize(buffers.
count, num_groups);
533 const auto dim = (row ? mat.
nrow() : mat.
ncol());
534 for (Group_ g = 0; g < num_groups; ++g) {
536#ifdef TATAMI_STATS_TEST_DIRTY
540 buffers.
mean[g] = output.
mean[g].data();
543#ifdef TATAMI_STATS_TEST_DIRTY
547 buffers.
rss[g] = output.
rss[g].data();
550#ifdef TATAMI_STATS_TEST_DIRTY
557 group_rss(row, mat, group, num_groups, buffers, opt);