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