Loading...
Searching...
No Matches
cxi.hpp
1// Copyright © 2025 Thibault Vatter
2//
3// This file is part of the wdm library and licensed under the terms of
4// the MIT license. For a copy, see the LICENSE file in the root directory
5// or https://github.com/tnagler/wdm/blob/master/LICENSE.
6
7#pragma once
8
9#include "ranks.hpp"
10#include "utils.hpp"
11#include <memory>
12#include <tuple>
13
14namespace wdm {
15namespace impl {
16
20inline std::vector<int>
21default_tie_seeds()
22{
23 return { 1, 2, 3, 4, 5 };
24}
25
28inline void
29sort_chatterjee_observations(std::vector<double>& x,
30 std::vector<double>& y,
31 std::vector<double>& weights,
32 const std::vector<int>& seeds)
33{
34 std::vector<size_t> order = utils::get_order(x);
35 std::unique_ptr<random::RandomGenerator> tie_generator;
36 for (size_t begin = 0, end; begin < order.size(); begin = end) {
37 end = begin + 1;
38 while (end < order.size() && x[order[end]] == x[order[begin]])
39 ++end;
40 if (end - begin > 1) {
41 if (!tie_generator)
42 tie_generator.reset(new random::RandomGenerator(seeds));
43 std::vector<size_t> tied_order(order.begin() + begin,
44 order.begin() + end);
45 random::shuffle(tied_order, *tie_generator);
46 std::copy(tied_order.begin(), tied_order.end(), order.begin() + begin);
47 }
48 }
49
50 std::vector<double> sorted_x(x.size()), sorted_y(y.size()),
51 sorted_weights(weights.size());
52 for (size_t i = 0; i < order.size(); ++i) {
53 sorted_x[i] = x[order[i]];
54 sorted_y[i] = y[order[i]];
55 sorted_weights[i] = weights[order[i]];
56 }
57 x = sorted_x;
58 y = sorted_y;
59 weights = sorted_weights;
60}
61
62// Conditional null mean and standard deviation for a continuous response.
63inline std::tuple<double, double>
64xi_continuous_inference(const std::vector<double>& probabilities)
65{
66 double edge_weight_sum = 0.0;
67 double null_numerator_mean = 0.0;
68 double squared_edge_weight_sum = 0.0;
69 double adjacent_edge_product_sum = 0.0;
70 double squared_probability_sum = 0.0;
71 double edge_node_product_sum = 0.0;
72
73 for (size_t i = 0; i < probabilities.size(); ++i)
74 squared_probability_sum += probabilities[i] * probabilities[i];
75
76 for (size_t i = 0; i + 1 < probabilities.size(); ++i) {
77 edge_weight_sum += probabilities[i];
78 null_numerator_mean +=
79 probabilities[i] *
80 (1.0 / 3.0 + (probabilities[i] + probabilities[i + 1]) / 6.0);
81 squared_edge_weight_sum += probabilities[i] * probabilities[i];
82 edge_node_product_sum +=
83 probabilities[i] * (probabilities[i] + probabilities[i + 1]);
84 if (i + 2 < probabilities.size())
85 adjacent_edge_product_sum += probabilities[i] * probabilities[i + 1];
86 }
87
88 double null_numerator_variance = squared_edge_weight_sum / 18.0;
89 null_numerator_variance += adjacent_edge_product_sum / 90.0;
90 null_numerator_variance +=
91 edge_weight_sum * edge_weight_sum * squared_probability_sum / 45.0;
92 null_numerator_variance -= edge_weight_sum * edge_node_product_sum / 45.0;
93 if (!std::isfinite(null_numerator_variance) || null_numerator_variance <= 0.0)
94 throw std::runtime_error(
95 "cannot compute the null variance of Chatterjee's xi.");
96
97 return std::make_tuple(3.0 * std::sqrt(null_numerator_variance),
98 1.0 - 3.0 * null_numerator_mean);
99}
100
101// Asymptotic standard deviation for xi with a tied response.
102inline double
103xi_std(const std::vector<double>& r,
104 const std::vector<double>& l,
105 const std::vector<double>& weights = std::vector<double>())
106{
107 double n =
108 (weights.size() > 0) ? utils::sum(weights) : static_cast<double>(r.size());
109
110 // Weighted version
111 std::vector<double> i(r.size());
112 for (size_t k = 0; k < r.size(); ++k)
113 i[k] = k + 1;
114
115 // Sort r and weights together
116 std::vector<size_t> order = utils::get_order(r);
117 std::vector<double> u(r.size()), w(r.size());
118 for (size_t k = 0; k < r.size(); ++k) {
119 u[k] = r[order[k]];
120 w[k] = (weights.size() > 0) ? weights[order[k]] : 1.0;
121 }
122
123 // Weighted cumulative sum
124 std::vector<double> v(r.size());
125 v[0] = u[0] * w[0];
126 for (size_t k = 1; k < r.size(); ++k)
127 v[k] = v[k - 1] + u[k] * w[k];
128
129 double an = 0, bn = 0, cn = 0, dn = 0;
130 for (size_t k = 0; k < r.size(); ++k) {
131 an += (2 * n - 2 * i[k] + 1) * u[k] * u[k] * w[k];
132 cn += (2 * n - 2 * i[k] + 1) * u[k] * w[k];
133 dn += l[k] * (n - l[k]) * ((weights.size() > 0) ? weights[k] : 1.0);
134 }
135 an /= std::pow(n, 4);
136 cn /= std::pow(n, 3);
137 dn /= std::pow(n, 3);
138
139 for (size_t k = 0; k < r.size(); ++k) {
140 double temp = v[k] + (n - i[k]) * u[k] * w[k];
141 bn += temp * temp;
142 }
143 bn /= std::pow(n, 5);
144
145 double tau2 = (an - 2 * bn + cn * cn) / (dn * dn);
146 return std::sqrt(tau2) / std::sqrt(n);
147}
148
166inline std::tuple<double, double, double, double>
167cxi(std::vector<double> x,
168 std::vector<double> y,
169 std::vector<double> weights = std::vector<double>(),
170 bool calculate_std = true,
171 std::string ties_method = "max",
172 std::vector<int> seeds = std::vector<int>(),
173 bool y_continuous = true)
174{
175 utils::check_sizes(x, y, weights);
176
177 if (weights.size() == 0)
178 weights = std::vector<double>(x.size(), 1.0);
179
180 utils::validate_weights(weights);
181 double weight_sum = utils::sum(weights);
182
183 // Zero-mass observations are absent from the weighted empirical measure and
184 // must not create additional edges in predictor order.
185 for (size_t i = weights.size(); i-- > 0;) {
186 if (weights[i] == 0.0) {
187 x.erase(x.begin() + i);
188 y.erase(y.begin() + i);
189 weights.erase(weights.begin() + i);
190 }
191 }
192
193 // Sort in x order and break x ties uniformly without consulting y. An empty
194 // seed vector would draw from std::random_device, making the estimate differ
195 // between calls on the same data; pass seeds to vary the tie ordering.
196 sort_chatterjee_observations(
197 x, y, weights, seeds.empty() ? default_tie_seeds() : seeds);
198
199 std::vector<double> probabilities = weights;
200 for (auto& probability : probabilities)
201 probability /= weight_sum;
202 bool weights_are_unequal = false;
203 for (size_t i = 1; i < probabilities.size(); ++i)
204 weights_are_unequal =
205 weights_are_unequal || probabilities[i] != probabilities[0];
206
207 std::vector<double> ordered_response = y;
208 std::sort(ordered_response.begin(), ordered_response.end());
209 if (ordered_response.front() == ordered_response.back())
210 throw std::runtime_error(
211 "Chatterjee's xi is undefined for a constant response.");
212 bool response_has_ties =
213 std::adjacent_find(ordered_response.begin(), ordered_response.end()) !=
214 ordered_response.end();
215
216 // Weighted empirical distribution at each response.
217 std::vector<double> r = rank0(y, probabilities, ties_method);
218
219 // Weighted empirical survival function at each response.
220 std::vector<double> y_neg(y.size());
221 for (size_t i = 0; i < y.size(); ++i)
222 y_neg[i] = -y[i];
223 std::vector<double> l = rank0(y_neg, probabilities, ties_method);
224
225 // Numerator: base-point weight on edge (i, i + 1).
226 double num = 0.0;
227 for (size_t i = 0; i + 1 < r.size(); ++i)
228 num += probabilities[i] * std::abs(r[i + 1] - r[i]);
229
230 // General weighted-rank denominator, valid for continuous and tied responses.
231 double den = 0.0;
232 for (size_t i = 0; i < l.size(); ++i)
233 den += 2.0 * probabilities[i] * l[i] * (1.0 - l[i]);
234 if (!std::isfinite(den) || den <= 0.0)
235 throw std::runtime_error(
236 "Chatterjee's xi is undefined for a constant response.");
237
238 double xi = 1.0 - num / den;
239
240 if (!calculate_std) {
241 return std::make_tuple(xi,
242 std::numeric_limits<double>::quiet_NaN(),
243 std::numeric_limits<double>::quiet_NaN(),
244 std::numeric_limits<double>::quiet_NaN());
245 } else if (y_continuous && !response_has_ties) {
246 auto inference = xi_continuous_inference(probabilities);
247 return std::make_tuple(
248 xi, std::get<0>(inference), std::get<1>(inference), 1.0 - 3.0 * num);
249 } else {
250 if (weights_are_unequal)
251 throw std::runtime_error(
252 "analytic Chatterjee inference is unavailable for an unequally "
253 "weighted, discrete or tied response.");
254 std::vector<double> raw_r = rank0(y, {}, ties_method);
255 std::vector<double> raw_l = rank0(y_neg, {}, ties_method);
256 return std::make_tuple(xi, xi_std(raw_r, raw_l), 0.0, xi);
257 }
258}
259
260} // namespace impl
261} // namespace wdm
Weighted dependence measures.
Definition wdm.hpp:19