Skip to content

Commit f2086c6

Browse files
committed
fix: add cached LMDB environment to support changes in lmdb 2.0.0
Signed-off-by: Luca Foppiano <luca@foppiano.org>
1 parent aed108f commit f2086c6

3 files changed

Lines changed: 20 additions & 6 deletions

File tree

delft/sequenceLabelling/wrapper.py

Lines changed: 2 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -207,7 +207,7 @@ def __init__(
207207

208208
def get_embedding(self, embedding_name, use_cache=True):
209209
"""Return an Embeddings instance for the given name. Override to customize embedding loading."""
210-
return Embeddings(embedding_name, resource_registry=self.registry, use_cache=use_cache)
210+
return Embeddings.get_or_create(embedding_name, resource_registry=self.registry, use_cache=use_cache)
211211

212212
def train(
213213
self,
@@ -807,8 +807,7 @@ def load(self, dir_path="data/models/sequenceLabelling/", weight_file=DEFAULT_WE
807807

808808
if self.model_config.embeddings_name is not None:
809809
# load embeddings
810-
# Do not use cache in 'prediction/production' mode
811-
self.embeddings = self.get_embedding(self.model_config.embeddings_name, use_cache=False)
810+
self.embeddings = self.get_embedding(self.model_config.embeddings_name, use_cache=True)
812811
self.model_config.word_embedding_size = self.embeddings.embed_size
813812
else:
814813
self.embeddings = None

delft/textClassification/wrapper.py

Lines changed: 2 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -155,7 +155,7 @@ def __init__(
155155

156156
def get_embedding(self, embedding_name, use_cache=True):
157157
"""Return an Embeddings instance for the given name. Override to customize embedding loading."""
158-
return Embeddings(embedding_name, resource_registry=self.registry, use_cache=use_cache)
158+
return Embeddings.get_or_create(embedding_name, resource_registry=self.registry, use_cache=use_cache)
159159

160160
def train(self, x_train, y_train, vocab_init=None, incremental=False, callbacks=None):
161161

@@ -565,8 +565,7 @@ def load(self, dir_path="data/models/textClassification/"):
565565

566566
if self.model_config.transformer_name is None:
567567
# load embeddings
568-
# Do not use cache in 'production' mode
569-
self.embeddings = self.get_embedding(self.model_config.embeddings_name, use_cache=False)
568+
self.embeddings = self.get_embedding(self.model_config.embeddings_name, use_cache=True)
570569
self.model_config.word_embedding_size = self.embeddings.embed_size
571570
else:
572571
self.transformer_name = self.model_config.transformer_name

delft/utilities/Embeddings.py

Lines changed: 16 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -38,6 +38,22 @@
3838

3939

4040
class Embeddings(object):
41+
_cache = {} # class-level cache: (name, lmdb_path) -> Embeddings instance
42+
43+
@classmethod
44+
def get_or_create(cls, name, resource_registry=None, lang="en", extension="vec", use_cache=True, load=True):
45+
if use_cache:
46+
lmdb_path = resource_registry.get("embedding-lmdb-path") if resource_registry else None
47+
cache_key = (name, lmdb_path)
48+
if cache_key in cls._cache:
49+
return cls._cache[cache_key]
50+
instance = cls(name, resource_registry=resource_registry, lang=lang, extension=extension,
51+
use_cache=use_cache, load=load)
52+
cls._cache[cache_key] = instance
53+
return instance
54+
return cls(name, resource_registry=resource_registry, lang=lang, extension=extension,
55+
use_cache=use_cache, load=load)
56+
4157
def __init__(self, name, resource_registry=None, lang="en", extension="vec", use_cache=True, load=True):
4258

4359
self.name = name

0 commit comments

Comments
 (0)