Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
57 changes: 17 additions & 40 deletions packages/react-native-executorch/cpp/extensions/nlp/tokenizer.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -67,6 +67,18 @@ TokenizerHostObject::TokenizerHostObject(std::string tokenizerPath)
}
}

std::unique_lock<std::mutex> TokenizerHostObject::tryLockUnique(jsi::Runtime &rt,
std::string_view context) {
std::unique_lock<std::mutex> lock(mutex_, std::try_to_lock);
if (!lock.owns_lock()) {
throw jsi::JSError(rt, std::format("{} is currently in use", context));
}
if (!tokenizer_) {
throw jsi::JSError(rt, std::format("{} has been disposed", context));
}
return lock;
}

jsi::Value TokenizerHostObject::get(jsi::Runtime &rt, const jsi::PropNameID &name) {
auto nameStr = name.utf8(rt);

Expand All @@ -81,14 +93,7 @@ jsi::Value TokenizerHostObject::get(jsi::Runtime &rt, const jsi::PropNameID &nam
throw jsi::JSError(rt, "encode: Usage: encode(text)");
}

std::unique_lock<std::mutex> lock(self->mutex_, std::try_to_lock);
if (!lock.owns_lock()) {
throw jsi::JSError(rt, "encode: Tokenizer is currently in use");
}

if (!self->tokenizer_) {
throw jsi::JSError(rt, "encode: Tokenizer has been disposed");
}
auto lock = self->tryLockUnique(rt, "encode: Tokenizer");

auto text = conversions::asType<std::string>(rt, "encode: text", args[0]);
auto tokens = unwrap(rt, "encode: Failed to encode input",
Expand All @@ -112,14 +117,7 @@ jsi::Value TokenizerHostObject::get(jsi::Runtime &rt, const jsi::PropNameID &nam
skipSpecialTokens = conversions::asType<bool>(rt, "decode: skipSpecialTokens", args[1]);
}

std::unique_lock<std::mutex> lock(self->mutex_, std::try_to_lock);
if (!lock.owns_lock()) {
throw jsi::JSError(rt, "decode: Tokenizer is currently in use");
}

if (!self->tokenizer_) {
throw jsi::JSError(rt, "decode: Tokenizer has been disposed");
}
auto lock = self->tryLockUnique(rt, "decode: Tokenizer");

auto tokens = conversions::asVector<uint64_t>(rt, "decode: tokens", args[0]);

Expand All @@ -142,14 +140,7 @@ jsi::Value TokenizerHostObject::get(jsi::Runtime &rt, const jsi::PropNameID &nam
throw jsi::JSError(rt, "getVocabSize: Usage: getVocabSize()");
}

std::unique_lock<std::mutex> lock(self->mutex_, std::try_to_lock);
if (!lock.owns_lock()) {
throw jsi::JSError(rt, "getVocabSize: Tokenizer is currently in use");
}

if (!self->tokenizer_) {
throw jsi::JSError(rt, "getVocabSize: Tokenizer has been disposed");
}
auto lock = self->tryLockUnique(rt, "getVocabSize: Tokenizer");

return static_cast<double>(self->tokenizer_->vocab_size());
};
Expand All @@ -163,14 +154,7 @@ jsi::Value TokenizerHostObject::get(jsi::Runtime &rt, const jsi::PropNameID &nam
throw jsi::JSError(rt, "idToToken: Usage: idToToken(id)");
}

std::unique_lock<std::mutex> lock(self->mutex_, std::try_to_lock);
if (!lock.owns_lock()) {
throw jsi::JSError(rt, "idToToken: Tokenizer is currently in use");
}

if (!self->tokenizer_) {
throw jsi::JSError(rt, "idToToken: Tokenizer has been disposed");
}
auto lock = self->tryLockUnique(rt, "idToToken: Tokenizer");

auto tokenId = conversions::asType<uint64_t>(rt, "idToToken: id", args[0]);
auto token = unwrap(rt, "idToToken: Failed to convert id to token",
Expand All @@ -188,14 +172,7 @@ jsi::Value TokenizerHostObject::get(jsi::Runtime &rt, const jsi::PropNameID &nam
throw jsi::JSError(rt, "tokenToId: Usage: tokenToId(token)");
}

std::unique_lock<std::mutex> lock(self->mutex_, std::try_to_lock);
if (!lock.owns_lock()) {
throw jsi::JSError(rt, "tokenToId: Tokenizer is currently in use");
}

if (!self->tokenizer_) {
throw jsi::JSError(rt, "tokenToId: Tokenizer has been disposed");
}
auto lock = self->tryLockUnique(rt, "tokenToId: Tokenizer");

auto token = conversions::asType<std::string>(rt, "tokenToId: token", args[0]);
auto tokenId = unwrap(rt, "tokenToId: Failed to convert token to id",
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -3,6 +3,7 @@
#include <memory>
#include <mutex>
#include <string>
#include <string_view>
#include <vector>

#include <jsi/jsi.h>
Expand All @@ -20,6 +21,9 @@ class TokenizerHostObject : public facebook::jsi::HostObject,
std::vector<facebook::jsi::PropNameID> getPropertyNames(facebook::jsi::Runtime &rt) override;

private:
[[nodiscard]] std::unique_lock<std::mutex> tryLockUnique(facebook::jsi::Runtime &rt,
std::string_view context);

std::string tokenizerPath_;
std::unique_ptr<tokenizers::HFTokenizer> tokenizer_;
std::mutex mutex_;
Expand Down