Skip to content

Commit e4dcb4a

Browse files
committed
fuzzy +-345% speedup
1 parent d364394 commit e4dcb4a

1 file changed

Lines changed: 115 additions & 51 deletions

File tree

src/finders/Fuzzy.cpp

Lines changed: 115 additions & 51 deletions
Original file line numberDiff line numberDiff line change
@@ -1,28 +1,104 @@
11
#include "Fuzzy.hpp"
22
#include <algorithm>
33
#include <cmath>
4+
#include <ranges>
45
#include <thread>
5-
#include <unordered_set>
66

77
#include <unistd.h>
88

9-
#include <hyprutils/string/VarList2.hpp>
109
#include <hyprutils/string/String.hpp>
1110

1211
using namespace Hyprutils::String;
1312

13+
namespace {
14+
15+
class CVarListView : public std::ranges::view_interface<CVarListView> {
16+
std::string_view m_str;
17+
char m_sep;
18+
19+
public:
20+
// (removeEmpty = true)
21+
CVarListView(std::string_view str, char sep) : m_str(str), m_sep(sep) {}
22+
23+
class CIterator {
24+
private:
25+
std::string_view m_str;
26+
std::size_t m_index = -1uz;
27+
std::size_t m_count = 0;
28+
char m_sep;
29+
30+
public:
31+
using iterator_category = std::forward_iterator_tag;
32+
using value_type = std::string_view;
33+
using difference_type = std::ptrdiff_t;
34+
using reference = value_type;
35+
using pointer = value_type;
36+
37+
CIterator() = default;
38+
CIterator(std::string_view str, char sep) : m_str(str), m_sep(sep) {
39+
this->operator++();
40+
}
41+
42+
reference operator*() const {
43+
return m_str.substr(m_index, m_count);
44+
}
45+
46+
pointer operator->() const {
47+
return this->operator*();
48+
}
49+
50+
CIterator& operator++() {
51+
m_index += m_count;
52+
do {
53+
m_index += 1;
54+
if (m_index >= m_str.size()) {
55+
m_index = -1uz;
56+
return *this;
57+
}
58+
} while(isSep(m_str[m_index]));
59+
60+
m_count = 1;
61+
62+
while (m_index + m_count < m_str.size() && !isSep(m_str[m_index + m_count]))
63+
m_count += 1;
64+
65+
return *this;
66+
}
67+
68+
CIterator operator++(int) {
69+
auto tmp = *this;
70+
++(*this);
71+
return tmp;
72+
}
73+
74+
bool isSep(char c) const {
75+
return m_sep == ' ' ? std::isspace(c) : m_sep == '/' ? std::isspace(c) || c == m_sep : c == m_sep;
76+
}
77+
78+
bool operator==(const CIterator& other) const {
79+
return m_index == other.m_index;
80+
}
81+
};
82+
83+
using iterator = CIterator;
84+
85+
auto begin() const { return CIterator(m_str, m_sep); }
86+
auto end() const { return CIterator(); }
87+
};
88+
89+
}
90+
1491
static float jaroWinkler(const std::string_view& query, const std::string_view& test) {
1592
const auto LENGTH_A = query.length();
1693
const auto LENGTH_B = test.length();
1794

1895
if (!LENGTH_A && !LENGTH_B)
1996
return 0;
2097

21-
const auto MATCH_DISTANCE = LENGTH_A == 1 && LENGTH_B == 1 ? 0 : ((std::max(LENGTH_A, LENGTH_B) / 2) - 1);
98+
const auto MATCH_DISTANCE = LENGTH_A == 1 && LENGTH_B == 1 ? 0 : ((std::max(LENGTH_A, LENGTH_B) / 2) - 1);
2299

23-
std::vector<bool> matchesA, matchesB;
24-
matchesA.resize(LENGTH_A);
25-
matchesB.resize(LENGTH_B);
100+
bool* matchesA = (bool*)alloca(LENGTH_A * sizeof(bool));
101+
bool* matchesB = (bool*)alloca(LENGTH_B * sizeof(bool));
26102
size_t matches = 0;
27103
for (size_t i = 0; i < LENGTH_A; ++i) {
28104
const size_t start = (i > MATCH_DISTANCE ? i - MATCH_DISTANCE : 0);
@@ -46,7 +122,7 @@ static float jaroWinkler(const std::string_view& query, const std::string_view&
46122
if (!matchesA[i])
47123
continue;
48124

49-
while (k < matchesB.size() && !matchesB[k]) {
125+
while (k < LENGTH_B && !matchesB[k]) {
50126
++k;
51127
}
52128

@@ -80,16 +156,17 @@ constexpr float NO_SALIENT_PENALTY = 0.01F;
80156
constexpr float EXACT_MATCH_SCORE = 2.0F;
81157

82158
//
83-
static float tokenBestMatch(std::string_view qt, std::string_view lastQ, const std::unordered_set<std::string_view>& cset, const std::vector<std::string_view>& cTok) {
159+
static float tokenBestMatch(std::string_view qt, std::string_view lastQ, const CVarListView& cTok) {
84160
if (qt.empty())
85161
return 0.F;
86-
if (cset.contains(qt))
87-
return 1.F;
88162

89163
float best = 0.F;
90164
bool hasExplicitMatch = false; // prefix or substring match
91165

92166
for (auto ct : cTok) {
167+
if (ct == qt)
168+
return 1.F;
169+
93170
// strong prefix match - especially important for the last token (partial typing)
94171
if (ct.starts_with(qt)) {
95172
hasExplicitMatch = true;
@@ -116,63 +193,44 @@ static float tokenBestMatch(std::string_view qt, std::string_view lastQ, const s
116193
return (best - MIN_FUZZY_TO_COUNT) / (1.F - MIN_FUZZY_TO_COUNT);
117194
}
118195

119-
static float scoreCandidate(std::string_view query, std::string_view cand, float freq, char tokenBreak) {
196+
static float scoreCandidate(const std::vector<std::string_view>& qTokens, std::string_view queryLowerTrim, const std::string& query, std::string_view cand, float freq, char tokenBreak) {
120197
const float popFactor = 1.F + (POPULARITY_FACTOR * std::log1p(std::max(0.F, freq)));
121198

122199
// exact matches occupy a reserved band above any achievable fuzzy score, so a popular
123200
// partial match can never outrank them; popularity only orders matches within a band
124-
std::string queryLower{query};
125-
std::ranges::transform(queryLower, queryLower.begin(), ::tolower);
126-
if (trim(queryLower) == trim(std::string{cand}))
201+
if (queryLowerTrim == trim(cand))
127202
return EXACT_MATCH_SCORE + popFactor;
128203

129-
CVarList2 qTokens(std::string{query}, 0, tokenBreak == ' ' ? 's' : tokenBreak, true, false);
130-
CVarList2 cTokens(std::string{cand}, 0, tokenBreak == ' ' ? 's' : tokenBreak, true, false);
131-
132-
std::vector<std::string_view> qTok, cTok;
133-
qTok.reserve(qTokens.size());
134-
cTok.reserve(cTokens.size());
135-
for (const auto& q : qTokens) {
136-
qTok.emplace_back(q);
137-
}
138-
for (const auto& c : cTokens) {
139-
cTok.emplace_back(c);
140-
}
204+
CVarListView cTok(cand, tokenBreak);
141205

142-
if (qTok.empty() || cTok.empty())
206+
if (qTokens.empty() || cTok.begin() == cTok.end())
143207
return 0.F;
144208

145-
std::unordered_set<std::string_view> cset;
146-
cset.reserve(cTok.size());
147-
for (auto t : cTok) {
148-
cset.insert(t);
149-
}
150-
151-
std::string_view lastQ = qTok.back();
209+
std::string_view lastQ = qTokens.back();
152210

153211
// pick salient token as longest
154-
std::string_view salient = qTok[0];
155-
for (auto t : qTok) {
156-
if (t.size() > salient.size())
157-
salient = t;
158-
}
159-
160-
float sum = 0.F;
161-
float minMatch = 1.F;
162-
for (auto qt : qTok) {
163-
float match = tokenBestMatch(qt, lastQ, cset, cTok);
212+
std::string_view salient = qTokens[0];
213+
float salientMatch = tokenBestMatch(qTokens[0], lastQ, cTok);
214+
float sum = 0.F + salientMatch;
215+
float minMatch = std::min(1.F, salientMatch);
216+
for (auto qt : qTokens | std::views::drop(1)) {
217+
float match = tokenBestMatch(qt, lastQ, cTok);
164218
sum += match;
165219
minMatch = std::min(minMatch, match);
220+
221+
if (qt.size() > salient.size()) {
222+
salient = qt;
223+
salientMatch = match;
224+
}
166225
}
167226

168227
// if ANY token matches poorly, penalize heavily
169228
if (minMatch < MIN_TOKEN_MATCH)
170229
return 0.F;
171230

172-
float base = sum / sc<float>(qTok.size()); // normalize it
231+
float base = sum / sc<float>(qTokens.size()); // normalize it
173232

174233
// if salient token doesn't match strongly, kill the score
175-
float salientMatch = tokenBestMatch(salient, lastQ, cset, cTok);
176234
if (salientMatch < MIN_SALIENT_MATCH)
177235
base *= NO_SALIENT_PENALTY;
178236

@@ -188,13 +246,13 @@ struct SScoreData {
188246
size_t idx = 0;
189247
};
190248

191-
static void workerFn(std::vector<SScoreData>& scores, const std::vector<SP<IFinderResult>>& in, const std::string& query, size_t start, size_t end, char tokenBreak) {
249+
static void workerFn(std::vector<SScoreData>& scores, const std::vector<SP<IFinderResult>>& in, const std::vector<std::string_view>& qTokens, const std::string& queryLowerTrim, const std::string& query, size_t start, size_t end, char tokenBreak) {
192250
for (size_t i = start; i < end; ++i) {
193251
auto& ref = scores[i];
194252

195253
float bestScore = 0.F;
196254
for (auto const& candidate : in[i]->fuzzables()) {
197-
auto score = scoreCandidate(query, candidate, in[i]->frequency(), tokenBreak);
255+
auto score = scoreCandidate(qTokens, queryLowerTrim, query, candidate, in[i]->frequency(), tokenBreak);
198256
bestScore = std::max(score, bestScore);
199257
}
200258
ref.score = bestScore;
@@ -234,6 +292,12 @@ static constexpr const decltype(sysconf(0)) MAX_THREADS = 10;
234292

235293
//
236294
std::vector<SP<IFinderResult>> Fuzzy::getNResults(const std::vector<SP<IFinderResult>>& in, const std::string& query, size_t results, char tokenBreak) {
295+
std::string queryLowerTrim{query};
296+
std::ranges::transform(queryLowerTrim, queryLowerTrim.begin(), ::tolower);
297+
queryLowerTrim = trim(queryLowerTrim);
298+
299+
auto qTokens = CVarListView(query, tokenBreak) | std::ranges::to<std::vector>();
300+
237301
std::vector<SScoreData> scores;
238302
scores.resize(in.size());
239303

@@ -251,10 +315,10 @@ std::vector<SP<IFinderResult>> Fuzzy::getNResults(const std::vector<SP<IFinderRe
251315
size_t workElDone = 0, workElPerThread = in.size() / THREADS;
252316
for (long i = 0; i < THREADS; ++i) {
253317
if (i == THREADS - 1) {
254-
workerThreads[i] = std::thread([&, begin = workElDone] { workerFn(scores, in, query, begin, in.size(), tokenBreak); });
318+
workerThreads[i] = std::thread([&, begin = workElDone] { workerFn(scores, in, qTokens, queryLowerTrim, query, begin, in.size(), tokenBreak); });
255319
break;
256320
}
257-
workerThreads[i] = std::thread([&, begin = workElDone, end = workElDone + workElPerThread] { workerFn(scores, in, query, begin, end, tokenBreak); });
321+
workerThreads[i] = std::thread([&, begin = workElDone, end = workElDone + workElPerThread] { workerFn(scores, in, qTokens, queryLowerTrim, query, begin, end, tokenBreak); });
258322

259323
workElDone += workElPerThread;
260324
}
@@ -266,7 +330,7 @@ std::vector<SP<IFinderResult>> Fuzzy::getNResults(const std::vector<SP<IFinderRe
266330

267331
workerThreads.clear();
268332
} else
269-
workerFn(scores, in, query, 0, in.size(), tokenBreak);
333+
workerFn(scores, in, qTokens, queryLowerTrim, query, 0, in.size(), tokenBreak);
270334

271335
return getBestResultsStable(scores, results);
272336
}

0 commit comments

Comments
 (0)