|
| 1 | +#pragma once |
| 2 | + |
| 3 | +#include "rf-class-depth.h" |
| 4 | +#include <unordered_set> |
| 5 | + |
| 6 | +// Sparse PLS proposals fitted to projections of neighbor centroids. The |
| 7 | +// original RF PAL objective selects splits; leaf voting and reranking are shared. |
| 8 | +class PLSCentroid : public RFClass { |
| 9 | + public: |
| 10 | + struct Options { |
| 11 | + int sample = 200, support = 16, candidates = 14, sketch_dim = 64; |
| 12 | + uint32_t seed = 17; |
| 13 | + }; |
| 14 | + |
| 15 | + PLSCentroid(const float *data, int n, int d) : PLSCentroid(data, n, d, Options{}) {} |
| 16 | + PLSCentroid(const float *data, int n, int d, Options options_) |
| 17 | + : RFClass(data, n, d), options(options_) { |
| 18 | + if (options.sample < 2 || options.support < 1 || options.candidates < 4 || |
| 19 | + options.sketch_dim < 8 || options.sketch_dim % 2) |
| 20 | + throw std::invalid_argument("Invalid PLSCentroid options"); |
| 21 | + } |
| 22 | + |
| 23 | + void grow(int trees, int depth_, const Eigen::Ref<const UIntRowMatrix> &knn, |
| 24 | + const Eigen::Ref<const RowMatrix> &train, float density_ = -1, int b_ = 1) override { |
| 25 | + if (!empty()) throw std::logic_error("The index has already been grown"); |
| 26 | + if (trees < 1 || train.rows() < 2 || depth_ < 1 || depth_ > std::log2(train.rows()) || |
| 27 | + train.cols() != dim || knn.rows() != train.rows() || knn.cols() < 1 || b_ < 1 || |
| 28 | + knn.maxCoeff() >= uint32_t(n_corpus)) throw std::invalid_argument("Invalid training data"); |
| 29 | + n_trees = trees; depth = depth_; b = b_; |
| 30 | + labels_all.resize(trees); votes_all.resize(trees); forest.resize(trees); |
| 31 | + |
| 32 | + // Targets never depend on test queries and are released after fitting. |
| 33 | + RowMatrix targets = make_targets(knn); |
| 34 | + log2_tbl.resize(train.rows() + 1); t_tbl.resize(train.rows() + 1); |
| 35 | + for (int i = 1; i <= train.rows(); ++i) log2_tbl[i] = std::log2(float(i)); |
| 36 | + for (int i = 0; i <= train.rows(); ++i) t_tbl[i] = i * log2_tbl[i]; |
| 37 | + for (int i = train.rows(); i > 0; --i) t_tbl[i] -= t_tbl[i - 1]; |
| 38 | +#pragma omp parallel for schedule(dynamic, 1) |
| 39 | + for (int tree = 0; tree < trees; ++tree) { |
| 40 | + Scratch scratch; |
| 41 | + scratch.ensure_corpus(n_corpus); |
| 42 | + scratch.generator.seed(options.seed + 104729U * (tree + 1)); |
| 43 | + std::vector<int> rows(train.rows()); |
| 44 | + std::iota(rows.begin(), rows.end(), 0); |
| 45 | + forest[tree].reserve(2 * (1 << std::min(depth, 16)) - 1); |
| 46 | + grow_node(rows.begin(), rows.end(), 0, tree, train, knn, targets, scratch); |
| 47 | + forest[tree].shrink_to_fit(); |
| 48 | + } |
| 49 | + } |
| 50 | + |
| 51 | + void query(const float *data, int k, float threshold, int *out, Distance dist = L2, |
| 52 | + float *distances = nullptr, int *n_elected = nullptr) const override { |
| 53 | + std::vector<uint32_t> elected; |
| 54 | + Eigen::VectorXf votes_total = Eigen::VectorXf::Zero(n_corpus); |
| 55 | + for (int tree = 0; tree < n_trees; ++tree) { |
| 56 | + int node = 0; |
| 57 | + while (forest[tree][node].leaf < 0) { |
| 58 | + const auto &n = forest[tree][node]; |
| 59 | + node = n.normal.project(data) <= n.normal.threshold ? n.left : n.right; |
| 60 | + } |
| 61 | + const int leaf = forest[tree][node].leaf; |
| 62 | + const auto &labels = labels_all[tree][leaf]; |
| 63 | + const auto &votes = votes_all[tree][leaf]; |
| 64 | + for (size_t i = 0; i < labels.size(); ++i) { |
| 65 | + if ((votes_total(labels[i]) += votes[i]) >= threshold) { |
| 66 | + elected.push_back(labels[i]); |
| 67 | + votes_total(labels[i]) = -9999999; |
| 68 | + } |
| 69 | + } |
| 70 | + } |
| 71 | + if (n_elected) *n_elected = elected.size(); |
| 72 | + exact_knn(Eigen::Map<const Eigen::RowVectorXf>(data, dim), k, elected, out, dist, distances); |
| 73 | + } |
| 74 | + |
| 75 | + protected: |
| 76 | + struct Normal { |
| 77 | + std::vector<uint32_t> dims; |
| 78 | + Eigen::VectorXf weights; |
| 79 | + float threshold = 0, gain = 0; |
| 80 | + float project(const float *x) const { |
| 81 | + float value = 0; |
| 82 | + for (size_t i = 0; i < dims.size(); ++i) value += weights[i] * x[dims[i]]; |
| 83 | + return value; |
| 84 | + } |
| 85 | + }; |
| 86 | + struct Node { |
| 87 | + Normal normal; |
| 88 | + int left = -1, right = -1, leaf = -1; |
| 89 | + }; |
| 90 | + struct Scratch : SplitScratch { std::minstd_rand generator; }; |
| 91 | + Options options; |
| 92 | + std::vector<std::vector<Node>> forest; |
| 93 | + |
| 94 | + static std::vector<uint32_t> sample_unique(int n, int k, std::minstd_rand &generator) { |
| 95 | + |
| 96 | + std::vector<uint32_t> reservoir; |
| 97 | + reservoir.reserve(k); |
| 98 | + if (k * 4 < n) { |
| 99 | + // Floyd's algorithm: uniform subset, O(k), no duplicates. |
| 100 | + std::unordered_set<uint32_t> selected; |
| 101 | + for (int i = n - k; i < n; ++i) { |
| 102 | + uint32_t j = std::uniform_int_distribution<int>(0, i)(generator); |
| 103 | + if (!selected.insert(j).second) { selected.insert(i); j = i; } |
| 104 | + reservoir.push_back(j); |
| 105 | + } |
| 106 | + } else { |
| 107 | + reservoir.resize(n); |
| 108 | + std::iota(reservoir.begin(), reservoir.end(), 0); |
| 109 | + for (int i = 0; i < k; ++i) { |
| 110 | + int j = std::uniform_int_distribution<int>(i, n - 1)(generator); |
| 111 | + std::swap(reservoir[i], reservoir[j]); |
| 112 | + } |
| 113 | + reservoir.resize(k); |
| 114 | + } |
| 115 | + |
| 116 | + return reservoir; |
| 117 | + } |
| 118 | + |
| 119 | + Normal random_normal(int support, std::minstd_rand &rng) const { |
| 120 | + Normal normal; |
| 121 | + normal.dims = sample_unique(dim, std::min(dim, support), rng); |
| 122 | + normal.weights.resize(normal.dims.size()); |
| 123 | + std::normal_distribution<float> gaussian; |
| 124 | + for (int i = 0; i < normal.weights.size(); ++i) normal.weights[i] = gaussian(rng); |
| 125 | + normal.weights.normalize(); |
| 126 | + return normal; |
| 127 | + } |
| 128 | + RowMatrix make_targets(const Eigen::Ref<const UIntRowMatrix> &knn) const { |
| 129 | + const int r = options.sketch_dim; |
| 130 | + std::minstd_rand rng(options.seed + 271828U); |
| 131 | + RowMatrix embedding = RowMatrix::Zero(n_corpus, r); |
| 132 | + std::vector<Normal> projections; |
| 133 | + for (int j = 0; j < r; ++j) projections.push_back(random_normal(std::min(16, dim), rng)); |
| 134 | +#pragma omp parallel for |
| 135 | + for (int i = 0; i < n_corpus; ++i) |
| 136 | + for (int j = 0; j < r; ++j) embedding(i,j) = projections[j].project(corpus.row(i).data()); |
| 137 | + RowMatrix targets = RowMatrix::Zero(knn.rows(), r); |
| 138 | +#pragma omp parallel for |
| 139 | + for (int i = 0; i < knn.rows(); ++i) { |
| 140 | + for (int j = 0; j < knn.cols(); ++j) targets.row(i) += embedding.row(knn(i,j)); |
| 141 | + targets.row(i) /= float(knn.cols()); |
| 142 | + } |
| 143 | + return targets; |
| 144 | + } |
| 145 | + |
| 146 | + Normal scan(Normal normal, const std::vector<int> &rows, |
| 147 | + const Eigen::Ref<const RowMatrix> &train, const UIntRowMatrix &labels, |
| 148 | + int n_labels) const { |
| 149 | + const int n = rows.size(), k = labels.cols(); |
| 150 | + std::vector<SplitEntry> order(n); |
| 151 | + for (int i = 0; i < n; ++i) order[i] = {normal.project(train.row(rows[i]).data()), i}; |
| 152 | + miniselect::pdqsort_branchless(order.begin(), order.end(), [](auto &a, auto &c) { return a.key < c.key; }); |
| 153 | + normal.gain = 0; |
| 154 | + |
| 155 | + std::vector<int> counts(n_labels,0); |
| 156 | + std::vector<float> left_ent(n); |
| 157 | + float entropy = 0; |
| 158 | + for (int pos = 0; pos < n; ++pos) { |
| 159 | + for (int j = 0; j < k; ++j) entropy += t_tbl[++counts[labels(order[pos].index,j)]]; |
| 160 | + left_ent[pos] = k * log2_tbl[pos+1] - entropy / float(pos+1); |
| 161 | + } |
| 162 | + const float base = left_ent[n-1]; |
| 163 | + for (int pos = 0; pos < n-1; ++pos) { |
| 164 | + for (int j = 0; j < k; ++j) entropy -= t_tbl[counts[labels(order[pos].index,j)]--]; |
| 165 | + const int remain = n-pos-1; |
| 166 | + if (order[pos].key == order[pos+1].key) continue; |
| 167 | + const float right_ent = k * log2_tbl[remain] - entropy / float(remain); |
| 168 | + const float gain = base - ((pos+1)*(1.f/n)*left_ent[pos] + remain*(1.f/n)*right_ent); |
| 169 | + if (gain > normal.gain + tol) { |
| 170 | + normal.gain = gain; normal.threshold = midpoint(order[pos].key, order[pos+1].key); |
| 171 | + } |
| 172 | + } |
| 173 | + return normal; |
| 174 | + } |
| 175 | + static float midpoint(float a, float c) { |
| 176 | + const float mid = a + (c-a)*.5f; |
| 177 | + return mid < c ? mid : a; // Keep distinct adjacent floats separated. |
| 178 | + } |
| 179 | + UIntRowMatrix compact(const std::vector<int> &rows, const Eigen::Ref<const UIntRowMatrix> &knn, |
| 180 | + int &n_labels, SplitScratch &scratch) const { |
| 181 | + UIntRowMatrix labels(rows.size(), knn.cols()); |
| 182 | + n_labels = 0; |
| 183 | + std::vector<uint32_t> touched; |
| 184 | + for (size_t i = 0; i < rows.size(); ++i) for (int j = 0; j < knn.cols(); ++j) { |
| 185 | + uint32_t id = knn(rows[i],j); |
| 186 | + if (!scratch.votes[id]) { scratch.votes[id] = ++n_labels; touched.push_back(id); } |
| 187 | + labels(i,j) = scratch.votes[id]-1; |
| 188 | + } |
| 189 | + for (auto id : touched) scratch.votes[id] = 0; |
| 190 | + return labels; |
| 191 | + } |
| 192 | + static Eigen::MatrixXf centered_inputs(const std::vector<int> &rows, const std::vector<uint32_t> &dims, |
| 193 | + const Eigen::Ref<const RowMatrix> &train) { |
| 194 | + Eigen::MatrixXf x(rows.size(), dims.size()); |
| 195 | + for (size_t i = 0; i < rows.size(); ++i) for (size_t d = 0; d < dims.size(); ++d) x(i,d) = train(rows[i],dims[d]); |
| 196 | + x.rowwise() -= x.colwise().mean().eval(); |
| 197 | + return x; |
| 198 | + } |
| 199 | + static Eigen::VectorXf leading(const Eigen::MatrixXf &cov, std::minstd_rand &rng) { |
| 200 | + Eigen::VectorXf a(cov.cols()); |
| 201 | + std::normal_distribution<float> gaussian; |
| 202 | + for (int j = 0; j < a.size(); ++j) a[j] = gaussian(rng); |
| 203 | + a.normalize(); |
| 204 | + for (int step = 0; step < 8; ++step) { |
| 205 | + a = (cov*a).eval(); |
| 206 | + if (a.norm() < 1e-12f) break; |
| 207 | + a.normalize(); |
| 208 | + } |
| 209 | + return a; |
| 210 | + } |
| 211 | + std::vector<Normal> proposals(const std::vector<int> &rows, |
| 212 | + const Eigen::Ref<const RowMatrix> &train, const RowMatrix &z, |
| 213 | + std::minstd_rand &rng) { |
| 214 | + std::vector<Normal> pool; |
| 215 | + const int support = std::min(dim, options.support); |
| 216 | + const int axes = std::max(1, options.candidates/4); |
| 217 | + const int randoms = std::max(1, options.candidates/4); |
| 218 | + for (uint32_t d : sample_unique(dim,std::min(dim,axes),rng)) { |
| 219 | + Normal normal; normal.dims = {d}; normal.weights = Eigen::VectorXf::Ones(1); pool.push_back(normal); |
| 220 | + } |
| 221 | + for (int i = 0; i < randoms; ++i) pool.push_back(random_normal(support,rng)); |
| 222 | + Eigen::MatrixXf zc = z.rowwise() - z.colwise().mean(); |
| 223 | + int remaining = options.candidates-pool.size(); |
| 224 | + while (remaining > 0) { |
| 225 | + auto dims = sample_unique(dim,support,rng); |
| 226 | + Eigen::MatrixXf x = centered_inputs(rows,dims,train); |
| 227 | + Eigen::MatrixXf map = x.transpose()*zc / float(rows.size()); |
| 228 | + const int count = std::min(remaining, 3); |
| 229 | + for (int c = 0; c < count; ++c) { |
| 230 | + Eigen::VectorXf a(map.cols()); |
| 231 | + if (c == 0) a = leading((zc.transpose()*zc).eval(),rng); |
| 232 | + else { std::normal_distribution<float> gaussian; for (int j = 0; j < a.size(); ++j) a[j] = gaussian(rng); a.normalize(); } |
| 233 | + Normal normal; normal.dims = dims; normal.weights = map*a; |
| 234 | + if (normal.weights.norm() > 1e-10f) normal.weights.normalize(); |
| 235 | + pool.push_back(std::move(normal)); |
| 236 | + } |
| 237 | + remaining -= count; |
| 238 | + } |
| 239 | + return pool; |
| 240 | + } |
| 241 | + |
| 242 | + int grow_node(std::vector<int>::iterator begin, std::vector<int>::iterator end, int level, int tree, |
| 243 | + const Eigen::Ref<const RowMatrix> &train, const Eigen::Ref<const UIntRowMatrix> &knn, |
| 244 | + const RowMatrix &targets, Scratch &scratch) { |
| 245 | + const int node = forest[tree].size(); |
| 246 | + forest[tree].emplace_back(); |
| 247 | + const int n = end-begin; |
| 248 | + Normal best; |
| 249 | + if (level < depth && n > 1) { |
| 250 | + auto sampled = sample_unique(n,std::min(n,options.sample),scratch.generator); |
| 251 | + std::vector<int> rows; rows.reserve(sampled.size()); |
| 252 | + for (auto i : sampled) rows.push_back(begin[i]); |
| 253 | + RowMatrix z(rows.size(),targets.cols()); |
| 254 | + for (size_t i = 0; i < rows.size(); ++i) z.row(i) = targets.row(rows[i]); |
| 255 | + int n_labels; |
| 256 | + auto labels = compact(rows,knn,n_labels,scratch); |
| 257 | + auto candidates = proposals(rows,train,z,scratch.generator); |
| 258 | + for (auto &candidate : candidates) { |
| 259 | + if (candidate.weights.squaredNorm() < 1e-12f) continue; |
| 260 | + auto evaluated = scan(std::move(candidate),rows,train,labels,n_labels); |
| 261 | + if (evaluated.gain > best.gain + tol) best = std::move(evaluated); |
| 262 | + } |
| 263 | + } |
| 264 | + if (best.dims.empty()) { |
| 265 | + const int leaf = labels_all[tree].size(); |
| 266 | + |
| 267 | + auto votes = count_votes(begin,end,knn,scratch); |
| 268 | + labels_all[tree].push_back(std::move(votes.first)); votes_all[tree].push_back(std::move(votes.second)); |
| 269 | + forest[tree][node].leaf = leaf; |
| 270 | + return node; |
| 271 | + } |
| 272 | + auto mid = std::partition(begin,end,[&](int row) { return best.project(train.row(row).data()) <= best.threshold; }); |
| 273 | + if (mid == begin || mid == end) throw std::logic_error("Hard scan produced an empty full-node child"); |
| 274 | + |
| 275 | + forest[tree][node].normal = std::move(best); |
| 276 | + const int left = grow_node(begin,mid,level+1,tree,train,knn,targets,scratch); |
| 277 | + const int right = grow_node(mid,end,level+1,tree,train,knn,targets,scratch); |
| 278 | + forest[tree][node].left = left; forest[tree][node].right = right; |
| 279 | + return node; |
| 280 | + } |
| 281 | + |
| 282 | +}; |
0 commit comments