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
252 changes: 0 additions & 252 deletions benchmarks/run.py

This file was deleted.

2 changes: 1 addition & 1 deletion tueri/output_scanners/factual_consistency.py
Original file line number Diff line number Diff line change
Expand Up @@ -49,8 +49,8 @@ def __init__(
model=model,
use_onnx=use_onnx,
)
self._model = self._model.to(device())
if not use_onnx:
self._model = self._model.to(device())
self._model.eval()

def scan(self, prompt: str, output: str) -> tuple[str, bool, float]:
Expand Down
6 changes: 6 additions & 0 deletions tueri/transformers_helpers.py
Original file line number Diff line number Diff line change
Expand Up @@ -89,8 +89,11 @@ def get_tokenizer_and_model_for_classification(
model.path,
subfolder=model.subfolder,
revision=model.revision,
torch_dtype="auto",
low_cpu_mem_usage=False,
**model.kwargs,
)
tf_model = tf_model.to(device())
LOGGER.debug("Initialized classification model", model=model, device=device())

return tf_tokenizer, tf_model
Expand Down Expand Up @@ -124,8 +127,11 @@ def get_tokenizer_and_model_for_ner(
model.path,
subfolder=model.subfolder,
revision=model.revision,
torch_dtype="auto",
low_cpu_mem_usage=False,
**model.kwargs,
)
tf_model = tf_model.to(device())
LOGGER.debug("Initialized NER model", model=model, device=device())

return tf_tokenizer, tf_model
Expand Down
4 changes: 2 additions & 2 deletions tueri_api/app/scanner.py
Original file line number Diff line number Diff line change
Expand Up @@ -30,15 +30,15 @@

LOGGER = structlog.getLogger(__name__)

MONGO_URI = os.getenv("MONGO_URI", "mongodb://root:example@localhost:27017/")
MONGO_URL = os.getenv("MONGO_URL", "mongodb://root:example@localhost:27017/")
MONGO_DB, MONGO_COLLECTION = os.getenv("MONGO_DB", "ChatApp"), os.getenv("MONGO_COLLECTION", "TueriScanners")

# Suppress MongoDB heartbeat logs
logging.getLogger("pymongo.topology").setLevel(logging.WARNING)
logging.getLogger("pymongo.serverSelection").setLevel(logging.WARNING)

try:
mongo_client = MongoClient(MONGO_URI, serverSelectionTimeoutMS=5000, heartbeatFrequencyMS=60000)
mongo_client = MongoClient(MONGO_URL, serverSelectionTimeoutMS=5000, heartbeatFrequencyMS=60000)
db = mongo_client[MONGO_DB]
scanners_collection = db[MONGO_COLLECTION]
mongo_client.admin.command("ping")
Expand Down
Loading