33from __future__ import annotations
44
55import tempfile
6+ from collections .abc import Mapping
7+ from functools import partial
68from pathlib import Path
7- from typing import TYPE_CHECKING , Any , Literal , cast
9+ from typing import TYPE_CHECKING , Any , Literal , TypeAlias , Union , cast
810
911import bioregistry
1012import curies
1113import numpy as np
1214import pandas as pd
1315from pystow import get_sentence_transformer
14- from tqdm import tqdm
16+ from tqdm . contrib . concurrent import process_map
1517from typing_extensions import Unpack
1618
1719from 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
1921from pyobo .api .utils import get_version_from_kwargs
2022from pyobo .constants import GetOntologyKwargs , check_should_force
2123from pyobo .identifier_utils import wrap_norm_prefix
2224from pyobo .utils .path import CacheArtifact , get_cache_path
2325
2426if 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
117119EMBEDDING_INDEX_NAME = "luid"
118120EMBEDDING_DIMENSIONALITY = 384
121+ TransformerHint : TypeAlias = Union [str , "SentenceTransformer" , None ]
119122
120123
121124@wrap_norm_prefix
122125def 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+
172193def 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
212232def 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