9 #include "nan_handling.hpp"
28 inline std::vector<double>
29 rank(std::vector<double> x,
30 std::vector<double> weights = std::vector<double>(),
31 std::string ties_method =
"min",
32 std::vector<int> seeds = std::vector<int>())
34 if ((ties_method !=
"min") && (ties_method !=
"average") &&
35 (ties_method !=
"first") && (ties_method !=
"random"))
36 throw std::runtime_error(
37 "ties method must be one of 'min', 'average', 'first', 'random'.");
41 if (weights.size() == 0)
42 weights = std::vector<double>(n, 1.0);
44 if (weights.size() != n) {
45 throw std::runtime_error(
"weights and data must have same size.");
49 std::vector<double> nans;
50 if (utils::any_nan(x)) {
52 for (
size_t i = 0; i < n; i++) {
53 if (std::isnan(x[i])) {
54 x[i] = std::numeric_limits<double>::max();
62 utils::sum(weights) /
static_cast<double>(n - utils::sum(nans));
63 for (
auto& w : weights) {
68 std::vector<size_t> perm = utils::get_order(x);
72 std::unique_ptr<random::RandomGenerator> random_gen;
73 if (ties_method ==
"random")
74 random_gen.reset(
new random::RandomGenerator(seeds));
76 double w_acc = 0.0, w_batch;
77 for (
size_t i = 0, reps; i < n; i += reps) {
81 while ((i + reps < n) && (x[perm[i]] == x[perm[i + reps]]))
82 w_batch += weights[perm[i + reps++]];
85 for (
size_t k = 0; k < reps; ++k)
86 x[perm[i + k]] = w_acc + weights[perm[i]];
89 if ((ties_method ==
"first") || (ties_method ==
"random")) {
92 std::vector<size_t> ord(reps);
93 std::iota(ord.begin(), ord.end(), 0);
94 if (ties_method ==
"random")
95 random::shuffle(ord, *random_gen);
98 for (
size_t k = 0; k < reps; ++k) {
99 ww += weights[perm[i + ord[k]]];
100 x[perm[i + ord[k]]] = w_acc + ww;
102 }
else if (ties_method ==
"average") {
104 for (
size_t k = 0; k < reps; ++k)
105 x[perm[i + k]] += (w_batch - weights[perm[i]]) / 2;
113 if (nans.size() == n) {
114 for (
size_t i = 0; i < x.size(); i++) {
133 inline std::vector<double>
134 rank0(std::vector<double> x,
135 std::vector<double> weights = std::vector<double>(),
136 std::string ties_method =
"min")
138 if ((ties_method !=
"min") && (ties_method !=
"average") &&
139 (ties_method !=
"max"))
140 throw std::runtime_error(
141 "ties_method must be either 'min', 'average', or 'max'.");
145 if (weights.size() == 0)
146 weights = std::vector<double>(n, 1.0);
149 std::vector<size_t> perm = utils::get_order(x);
151 double w_acc = 0.0, w_batch;
152 for (
size_t i = 0, reps; i < n; i += reps) {
156 while ((i + reps < n) && (x[perm[i]] == x[perm[i + reps]]))
157 w_batch += weights[perm[i + reps++]];
160 for (
size_t k = 0; k < reps; ++k)
161 x[perm[i + k]] = w_acc;
167 if ((ties_method ==
"average") && (reps > 1)) {
168 std::vector<double> ww(reps);
169 for (
size_t k = 0; k < reps; ++k)
170 ww[k] = weights[perm[i + k]];
171 double offset = utils::perm_sum(ww, 2) / w_batch;
172 for (
size_t k = 0; k < reps; ++k)
173 x[perm[i + k]] += offset;
174 }
else if (ties_method ==
"max") {
176 for (
size_t k = 0; k < reps; ++k)
177 x[perm[i + k]] = w_acc;
188 inline std::vector<double>
189 bivariate_rank(std::vector<double> x,
190 std::vector<double> y,
191 std::vector<double> weights = std::vector<double>())
193 utils::check_sizes(x, y, weights);
196 std::vector<size_t> perm_x = utils::get_order(x);
197 perm_x = utils::invert_permutation(perm_x);
200 utils::sort_all(x, y, weights);
203 std::vector<size_t> perm_y = utils::get_order(y,
false);
204 perm_y = utils::invert_permutation(perm_y);
207 std::vector<double> counts(y.size(), 0.0);
208 utils::merge_sort_count_per_element(y, weights, counts);
211 std::vector<double> counts_tmp = counts;
212 for (
size_t i = 0; i < counts.size(); i++)
213 counts[i] = counts_tmp[perm_y[perm_x[i]]];
221 median(
const std::vector<double>& x,
222 std::vector<double> weights = std::vector<double>())
224 utils::check_sizes(x, x, weights);
228 auto perm = utils::get_order(x);
231 for (
size_t i = 0; i < n; i++) {
234 w[i] = weights[perm[i]];
239 auto ranks = rank0(xx, w,
"average");
240 if (weights.size() == 0)
241 weights = std::vector<double>(n, 1.0);
242 double rank_avrg = utils::perm_sum(weights, 2) / utils::sum(weights);
246 while (ranks[i] < rank_avrg)
248 if (ranks[i] == rank_avrg)
251 return 0.5 * (xx[i - 1] + xx[i]);
Weighted dependence measures.
Definition: wdm.hpp:19