Skip to content

Commit 4e6bb44

Browse files
authored
Speed up lexical embedding workflow (#536)
1 parent 8feb871 commit 4e6bb44

1 file changed

Lines changed: 52 additions & 33 deletions

File tree

src/pyobo/api/embedding.py

Lines changed: 52 additions & 33 deletions
Original file line numberDiff line numberDiff line change
@@ -3,26 +3,28 @@
33
from __future__ import annotations
44

55
import tempfile
6+
from collections.abc import Mapping
7+
from functools import partial
68
from pathlib import Path
7-
from typing import TYPE_CHECKING, Any, Literal, cast
9+
from typing import TYPE_CHECKING, Any, Literal, TypeAlias, Union, cast
810

911
import bioregistry
1012
import curies
1113
import numpy as np
1214
import pandas as pd
1315
from pystow import get_sentence_transformer
14-
from tqdm import tqdm
16+
from tqdm.contrib.concurrent import process_map
1517
from typing_extensions import Unpack
1618

1719
from pyobo.api.edges import get_edges_df
18-
from pyobo.api.names import get_definition, get_id_name_mapping, get_name
20+
from pyobo.api.names import get_definition, get_id_definition_mapping, get_id_name_mapping, get_name
1921
from pyobo.api.utils import get_version_from_kwargs
2022
from pyobo.constants import GetOntologyKwargs, check_should_force
2123
from pyobo.identifier_utils import wrap_norm_prefix
2224
from pyobo.utils.path import CacheArtifact, get_cache_path
2325

2426
if TYPE_CHECKING:
25-
import sentence_transformers
27+
from sentence_transformers import SentenceTransformer
2628

2729
__all__ = [
2830
"get_graph_embeddings_df",
@@ -37,15 +39,15 @@ def _get_text(
3739
/,
3840
*,
3941
name: str | None = None,
40-
**kwargs: Unpack[GetOntologyKwargs],
4142
) -> str | None:
4243
if name is None:
43-
name = get_name(reference, **kwargs)
44+
name = get_name(reference)
4445
if name is None:
4546
return None
46-
description = get_definition(reference, **kwargs)
47+
description = get_definition(reference)
4748
if description:
4849
name += " " + description
50+
# TODO include synonyms?
4951
return name
5052

5153

@@ -116,20 +118,24 @@ def get_graph_embeddings_df(
116118

117119
EMBEDDING_INDEX_NAME = "luid"
118120
EMBEDDING_DIMENSIONALITY = 384
121+
TransformerHint: TypeAlias = Union[str, "SentenceTransformer", None]
119122

120123

121124
@wrap_norm_prefix
122125
def get_text_embeddings_df(
123126
prefix: str,
124127
*,
125-
model: sentence_transformers.SentenceTransformer | None = None,
128+
model: TransformerHint = None,
129+
encode_kwargs: dict[str, Any] | None = None,
126130
**kwargs: Unpack[GetOntologyKwargs],
127131
) -> pd.DataFrame:
128132
"""Get embeddings for all entities in the resource.
129133
130134
:param prefix: A reference, either as a string or Reference object
131135
:param model: A sentence transformer model. Defaults to ``all-MiniLM-L6-v2`` if not
132136
given.
137+
:param encode_kwargs: Additional keyword arguments to pass to the encoder function
138+
:meth:`sentence_transformers.SentenceTransformer.encode`
133139
:param kwargs: The keyword arguments to forward to ontology getter functions for
134140
names, definitions, and version
135141
@@ -152,27 +158,42 @@ def get_text_embeddings_df(
152158
return df
153159

154160
id_to_name = get_id_name_mapping(prefix, **kwargs)
161+
# no kwargs needed because ontology was loaded above.
162+
id_to_description = get_id_definition_mapping(prefix)
163+
164+
identifiers = list(id_to_name)
165+
texts = process_map(
166+
partial(_id_to_text, id_to_name=id_to_name, id_to_description=id_to_description),
167+
identifiers,
168+
desc=f"[{prefix}] constructing text",
169+
unit_scale=True,
170+
chunksize=1000,
171+
)
155172

156-
luids, texts = [], []
157-
for identifier, name in tqdm(id_to_name.items(), desc=f"[{prefix}] constructing text"):
158-
text = _get_text(curies.ReferenceTuple(prefix, identifier), name=name, **kwargs)
159-
if text is None:
160-
continue
161-
luids.append(identifier)
162-
texts.append(text)
163-
if model is None:
164-
model = get_sentence_transformer()
165-
res = model.encode(texts, show_progress_bar=True)
166-
df = pd.DataFrame(res, index=luids)
173+
model_ = get_sentence_transformer(model)
174+
# TODO update to using MPL
175+
if encode_kwargs is None:
176+
encode_kwargs = {}
177+
encode_kwargs.setdefault("show_progress_bar", True)
178+
res = model_.encode(texts, **encode_kwargs)
179+
df = pd.DataFrame(res, index=identifiers)
167180
df.index.name = EMBEDDING_INDEX_NAME
168181
df.to_csv(path, sep="\t") # index is important here!
169182
return df
170183

171184

185+
def _id_to_text(
186+
identifier: str, id_to_name: Mapping[str, str], id_to_description: Mapping[str, str]
187+
) -> str:
188+
if identifier in id_to_description:
189+
return id_to_name[identifier] + " " + id_to_description[identifier]
190+
return id_to_name[identifier]
191+
192+
172193
def get_text_embedding(
173194
reference: str | curies.Reference | curies.ReferenceTuple,
174195
*,
175-
model: sentence_transformers.SentenceTransformer | None = None,
196+
model: TransformerHint = None,
176197
) -> np.ndarray[tuple[int], np.dtype[np.float64]] | None:
177198
"""Get a text embedding for an entity, or return none if no text is available.
178199
@@ -194,26 +215,25 @@ def get_text_embedding(
194215
.. code-block:: python
195216
196217
import pyobo
197-
from pyobo.api.embedding import get_text_embedding_model
218+
from pystow import get_sentence_transformer
198219
199-
model = get_text_embedding_model()
220+
model = get_sentence_transformer()
200221
embedding = pyobo.get_text_embedding("GO:0000001", model=model)
201222
# [-5.68335280e-02 7.96175096e-03 -3.36112119e-02 2.34440481e-03 ... ]
202223
"""
203224
text = _get_text(reference)
204225
if text is None:
205226
return None
206-
if model is None:
207-
model = get_sentence_transformer()
208-
res = model.encode([text])
227+
model_ = get_sentence_transformer(model)
228+
res = model_.encode([text])
209229
return cast(np.ndarray[tuple[int], np.dtype[np.float64]], res[0])
210230

211231

212232
def get_text_embedding_similarity(
213233
reference_1: str | curies.Reference | curies.ReferenceTuple,
214234
reference_2: str | curies.Reference | curies.ReferenceTuple,
215235
*,
216-
model: sentence_transformers.SentenceTransformer | None = None,
236+
model: TransformerHint = None,
217237
) -> float | None:
218238
"""Get the pairwise similarity.
219239
@@ -237,16 +257,15 @@ def get_text_embedding_similarity(
237257
.. code-block:: python
238258
239259
import pyobo
240-
from pyobo.api.embedding import get_text_embedding_model
260+
from pystow import get_sentence_transformer
241261
242-
model = get_text_embedding_model()
262+
model = get_sentence_transformer()
243263
similarity = pyobo.get_text_embedding_similarity("GO:0000001", "GO:0000004", model=model)
244264
# 0.24702128767967224
245265
"""
246-
if model is None:
247-
model = get_sentence_transformer()
248-
e1 = get_text_embedding(reference_1, model=model)
249-
e2 = get_text_embedding(reference_2, model=model)
266+
model_ = get_sentence_transformer(model)
267+
e1 = get_text_embedding(reference_1, model=model_)
268+
e2 = get_text_embedding(reference_2, model=model_)
250269
if e1 is None or e2 is None:
251270
return None
252-
return cast(float, model.similarity(e1, e2)[0][0].item())
271+
return cast(float, model_.similarity(e1, e2)[0][0].item())

0 commit comments

Comments
 (0)