Skip to content

Commit 00fb7cf

Browse files
committed
Isolate PLSCentroid implementation
1 parent dc4b882 commit 00fb7cf

3 files changed

Lines changed: 286 additions & 1 deletion

File tree

‎cpp/bindings.cpp‎

Lines changed: 3 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -11,6 +11,7 @@
1111
#include "Python.h"
1212
#include "numpy/arrayobject.h"
1313
#include "rf-class-depth.h"
14+
#include "pls-centroid.h"
1415
#include "rf-pca.h"
1516
#include "rf-rp.h"
1617

@@ -55,6 +56,8 @@ static int MLANN_init(mlannIndex *self, PyObject *args) {
5556
self->index = new RFRP(data, n, dim);
5657
else if (strcmp(index_type, "PCA") == 0)
5758
self->index = new RFPCA(data, n, dim);
59+
else if (strcmp(index_type, "PLSCentroid") == 0)
60+
self->index = new PLSCentroid(data, n, dim);
5861
else
5962
self->index = new RFClass(data, n, dim);
6063

‎cpp/pls-centroid.h‎

Lines changed: 282 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,282 @@
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+
};

‎cpp/rf-class-depth.h‎

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -177,7 +177,7 @@ class RFClass : public MLANN {
177177
exact_knn(q, k, elected, out, dist, out_distances);
178178
}
179179

180-
private:
180+
protected:
181181
std::vector<float> log2_tbl;
182182
std::vector<float> t_tbl;
183183

0 commit comments

Comments
 (0)