Skip to content

Commit a3a26ba

Browse files
authored
Merge pull request #886 from ChiahongHong/reduce-binary-search
Avoid repeated negative unigram lookups
2 parents ac23922 + 8156167 commit a3a26ba

3 files changed

Lines changed: 149 additions & 14 deletions

File tree

Source/Engine/gramambular2/reading_grid.cpp

Lines changed: 24 additions & 13 deletions
Original file line numberDiff line numberDiff line change
@@ -54,14 +54,18 @@ bool ReadingGrid::insertReading(const std::string& reading) {
5454
return false;
5555
}
5656

57-
if (!lm_.hasUnigrams(reading)) {
57+
std::string insertedReading = reading;
58+
auto unigrams = lm_.getUnigrams(insertedReading);
59+
if (unigrams.empty()) {
5860
return false;
5961
}
6062

6163
readings_.insert(readings_.begin() + static_cast<ptrdiff_t>(cursor_),
62-
reading);
64+
std::move(insertedReading));
6365
expandGridAt(cursor_);
64-
update();
66+
insert(cursor_,
67+
std::make_shared<Node>(readings_[cursor_], 1, std::move(unigrams)));
68+
update(cursor_, EditType::kInsertion);
6569

6670
// Cursor must only move after update().
6771
++cursor_;
@@ -78,7 +82,7 @@ bool ReadingGrid::deleteReadingBeforeCursor() {
7882
// Cursor must decrement for grid-shrinking and update to work.
7983
--cursor_;
8084
shrinkGridAt(cursor_);
81-
update();
85+
update(cursor_, EditType::kDeletion);
8286
return true;
8387
}
8488

@@ -90,7 +94,7 @@ bool ReadingGrid::deleteReadingAfterCursor() {
9094
readings_.erase(readings_.begin() + static_cast<ptrdiff_t>(cursor_),
9195
readings_.begin() + static_cast<ptrdiff_t>(cursor_ + 1));
9296
shrinkGridAt(cursor_);
93-
update();
97+
update(cursor_, EditType::kDeletion);
9498
return true;
9599
}
96100

@@ -291,7 +295,7 @@ void ReadingGrid::removeAffectedNodes(size_t loc) {
291295
// XXXXX
292296
// XXXXXXXXX
293297
//
294-
if (spans_.empty()) {
298+
if (spans_.empty() || loc == 0) {
295299
return;
296300
}
297301
size_t affectedLength = kMaximumSpanLength - 1;
@@ -323,7 +327,7 @@ std::string ReadingGrid::combineReading(
323327

324328
bool ReadingGrid::hasNodeAt(size_t loc, size_t readingLen,
325329
const std::string& reading) {
326-
if (loc > spans_.size()) {
330+
if (loc >= spans_.size()) {
327331
return false;
328332
}
329333
const NodePtr& n = spans_[loc].nodeOf(readingLen);
@@ -333,14 +337,20 @@ bool ReadingGrid::hasNodeAt(size_t loc, size_t readingLen,
333337
return reading == n->reading();
334338
}
335339

336-
void ReadingGrid::update() {
337-
size_t begin =
338-
(cursor_ <= kMaximumSpanLength) ? 0 : cursor_ - kMaximumSpanLength;
339-
size_t end = cursor_ + kMaximumSpanLength;
340+
void ReadingGrid::update(size_t loc, EditType editType) {
341+
// Spans that do not cross the edit retain their previous lookup result. A
342+
// node means that the lookup succeeded, while a null slot means that the
343+
// same reading was already looked up and did not exist. Only spans that
344+
// include an insertion or cross a deletion boundary need to be queried.
345+
size_t affectedLength = kMaximumSpanLength - 1;
346+
size_t begin = loc <= affectedLength ? 0 : loc - affectedLength;
347+
size_t end = editType == EditType::kInsertion ? loc + 1 : loc;
340348
end = std::min(end, readings_.size());
341349

342350
for (size_t pos = begin; pos < end; pos++) {
343-
for (size_t len = 1; len <= kMaximumSpanLength && pos + len <= end; len++) {
351+
size_t minimumLength = loc - pos + 1;
352+
size_t maximumLength = std::min(kMaximumSpanLength, readings_.size() - pos);
353+
for (size_t len = minimumLength; len <= maximumLength; len++) {
344354
std::string combinedReading =
345355
combineReading(readings_.begin() + static_cast<ptrdiff_t>(pos),
346356
readings_.begin() + static_cast<ptrdiff_t>(pos + len));
@@ -351,7 +361,8 @@ void ReadingGrid::update() {
351361
continue;
352362
}
353363

354-
insert(pos, std::make_shared<Node>(combinedReading, len, unigrams));
364+
insert(pos, std::make_shared<Node>(std::move(combinedReading), len,
365+
std::move(unigrams)));
355366
}
356367
}
357368
}

Source/Engine/gramambular2/reading_grid.h

Lines changed: 6 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -253,14 +253,19 @@ class ReadingGrid {
253253

254254
// Internal methods for maintaining the grid.
255255

256+
enum class EditType {
257+
kInsertion,
258+
kDeletion,
259+
};
260+
256261
void expandGridAt(size_t loc);
257262
void shrinkGridAt(size_t loc);
258263
void removeAffectedNodes(size_t loc);
259264
void insert(size_t loc, const NodePtr& node);
260265
std::string combineReading(std::vector<std::string>::const_iterator begin,
261266
std::vector<std::string>::const_iterator end);
262267
bool hasNodeAt(size_t loc, size_t readingLen, const std::string& reading);
263-
void update();
268+
void update(size_t loc, EditType editType);
264269

265270
// Internal implementation of overrideCandidate, with an optional reading.
266271
bool overrideCandidate(size_t loc, const std::string* reading,

Source/Engine/gramambular2/reading_grid_test.cpp

Lines changed: 119 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -166,6 +166,37 @@ class MockLM : public LanguageModel {
166166
bool hasUnigrams(const std::string&) override { return true; }
167167
};
168168

169+
class SingleReadingCountingLM : public LanguageModel {
170+
public:
171+
std::vector<Unigram> getUnigrams(const std::string& reading) override {
172+
++getUnigramsCount_;
173+
++getUnigramsCounts_[reading];
174+
if (reading.find(ReadingGrid::kDefaultSeparator) != std::string::npos) {
175+
return {};
176+
}
177+
return std::vector<Unigram>{Unigram(reading, -1)};
178+
}
179+
180+
bool hasUnigrams(const std::string& reading) override {
181+
++hasUnigramsCount_;
182+
return reading.find(ReadingGrid::kDefaultSeparator) == std::string::npos;
183+
}
184+
185+
[[nodiscard]] size_t getUnigramsCount() const { return getUnigramsCount_; }
186+
187+
[[nodiscard]] size_t getUnigramsCount(const std::string& reading) const {
188+
auto iter = getUnigramsCounts_.find(reading);
189+
return iter == getUnigramsCounts_.end() ? 0 : iter->second;
190+
}
191+
192+
[[nodiscard]] size_t hasUnigramsCount() const { return hasUnigramsCount_; }
193+
194+
private:
195+
std::map<std::string, size_t> getUnigramsCounts_;
196+
size_t getUnigramsCount_ = 0;
197+
size_t hasUnigramsCount_ = 0;
198+
};
199+
169200
static bool Contains(const std::vector<ReadingGrid::Candidate>& candidates,
170201
const std::string& str) {
171202
return std::any_of(candidates.cbegin(), candidates.cend(),
@@ -269,6 +300,94 @@ TEST(ReadingGridTest, BasicOperations) {
269300
ASSERT_EQ(grid.spans().size(), 0);
270301
}
271302

303+
TEST(ReadingGridTest, InsertReadingQueriesEachCombinationOnce) {
304+
auto lm = std::make_shared<SingleReadingCountingLM>();
305+
ReadingGrid grid(lm);
306+
307+
for (const char* reading : {"a", "b", "c", "d", "e", "f"}) {
308+
ASSERT_TRUE(grid.insertReading(reading));
309+
}
310+
311+
EXPECT_EQ(lm->getUnigramsCount(), 21);
312+
EXPECT_EQ(lm->hasUnigramsCount(), 0);
313+
EXPECT_EQ(lm->getUnigramsCount("a-b"), 1);
314+
EXPECT_EQ(lm->getUnigramsCount("a-b-c-d-e-f"), 1);
315+
EXPECT_EQ(lm->getUnigramsCount("e-f"), 1);
316+
}
317+
318+
TEST(ReadingGridTest, InsertionOnlyQueriesSpansContainingTheEdit) {
319+
auto lm = std::make_shared<SingleReadingCountingLM>();
320+
ReadingGrid grid(lm);
321+
ASSERT_TRUE(grid.insertReading("a"));
322+
ASSERT_TRUE(grid.insertReading("b"));
323+
ASSERT_TRUE(grid.insertReading("c"));
324+
325+
grid.setCursor(1);
326+
ASSERT_TRUE(grid.insertReading("x"));
327+
328+
EXPECT_EQ(lm->getUnigramsCount(), 12);
329+
EXPECT_EQ(lm->getUnigramsCount("b-c"), 1);
330+
EXPECT_EQ(lm->getUnigramsCount("a-x"), 1);
331+
EXPECT_EQ(lm->getUnigramsCount("a-x-b-c"), 1);
332+
EXPECT_EQ(lm->getUnigramsCount("x-b-c"), 1);
333+
}
334+
335+
TEST(ReadingGridTest, InsertionAtBeginningPreservesExistingLookupResults) {
336+
auto lm = std::make_shared<SingleReadingCountingLM>();
337+
ReadingGrid grid(lm);
338+
ASSERT_TRUE(grid.insertReading("b"));
339+
ASSERT_TRUE(grid.insertReading("c"));
340+
341+
grid.setCursor(0);
342+
ASSERT_TRUE(grid.insertReading("a"));
343+
344+
EXPECT_EQ(lm->getUnigramsCount(), 6);
345+
EXPECT_EQ(lm->getUnigramsCount("b-c"), 1);
346+
EXPECT_EQ(lm->getUnigramsCount("a-b"), 1);
347+
EXPECT_EQ(lm->getUnigramsCount("a-b-c"), 1);
348+
}
349+
350+
TEST(ReadingGridTest, DeletionOnlyQueriesSpansCrossingTheEdit) {
351+
auto lm = std::make_shared<SingleReadingCountingLM>();
352+
ReadingGrid grid(lm);
353+
ASSERT_TRUE(grid.insertReading("a"));
354+
ASSERT_TRUE(grid.insertReading("b"));
355+
ASSERT_TRUE(grid.insertReading("c"));
356+
ASSERT_TRUE(grid.insertReading("d"));
357+
358+
grid.setCursor(2);
359+
ASSERT_TRUE(grid.deleteReadingBeforeCursor());
360+
361+
EXPECT_EQ(lm->getUnigramsCount(), 12);
362+
EXPECT_EQ(lm->getUnigramsCount("a-c"), 1);
363+
EXPECT_EQ(lm->getUnigramsCount("a-c-d"), 1);
364+
EXPECT_EQ(lm->getUnigramsCount("c-d"), 1);
365+
366+
grid.setCursor(0);
367+
ASSERT_TRUE(grid.deleteReadingAfterCursor());
368+
EXPECT_EQ(lm->getUnigramsCount(), 12);
369+
ASSERT_EQ(grid.readings(), (std::vector<std::string>{"c", "d"}));
370+
ASSERT_NE(grid.spans()[0].nodeOf(1), nullptr);
371+
EXPECT_EQ(grid.spans()[0].nodeOf(1)->reading(), "c");
372+
373+
grid.setCursor(2);
374+
ASSERT_TRUE(grid.deleteReadingBeforeCursor());
375+
EXPECT_EQ(lm->getUnigramsCount(), 12);
376+
ASSERT_EQ(grid.readings(), (std::vector<std::string>{"c"}));
377+
ASSERT_NE(grid.spans()[0].nodeOf(1), nullptr);
378+
EXPECT_EQ(grid.spans()[0].nodeOf(1)->reading(), "c");
379+
}
380+
381+
TEST(ReadingGridTest, InsertReadingAcceptsReadingOwnedByGrid) {
382+
ReadingGrid grid(std::make_shared<MockLM>());
383+
ASSERT_TRUE(grid.insertReading("a"));
384+
385+
const std::string& aliasedReading = grid.readings()[0];
386+
ASSERT_TRUE(grid.insertReading(aliasedReading));
387+
ASSERT_EQ(grid.readings(), (std::vector<std::string>{"a", "a"}));
388+
ASSERT_EQ(grid.spans()[1].nodeOf(1)->reading(), "a");
389+
}
390+
272391
TEST(ReadingGridTest, InvalidOperations) {
273392
class TestLM : public LanguageModel {
274393
public:

0 commit comments

Comments
 (0)