biomapper 0.3.2__tar.gz → 0.4.0__tar.gz
This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
- {biomapper-0.3.2 → biomapper-0.4.0}/PKG-INFO +8 -1
- {biomapper-0.3.2 → biomapper-0.4.0}/biomapper/__init__.py +1 -1
- {biomapper-0.3.2 → biomapper-0.4.0}/biomapper/core/protein_metadata_comparison.py +1 -1
- biomapper-0.4.0/biomapper/mapping/clients/__init__.py +1 -0
- {biomapper-0.3.2/biomapper/mapping → biomapper-0.4.0/biomapper/mapping/clients}/chebi_client.py +16 -4
- {biomapper-0.3.2/biomapper/mapping → biomapper-0.4.0/biomapper/mapping/clients}/unichem_client.py +19 -19
- {biomapper-0.3.2/biomapper/mapping → biomapper-0.4.0/biomapper/mapping/clients}/uniprot_focused_mapper.py +21 -19
- biomapper-0.4.0/biomapper/mapping/embeddings/__init__.py +1 -0
- biomapper-0.4.0/biomapper/mapping/embeddings/managers.py +98 -0
- {biomapper-0.3.2 → biomapper-0.4.0}/biomapper/mapping/llm_mapper.py +1 -1
- {biomapper-0.3.2 → biomapper-0.4.0}/biomapper/mapping/metabolite_name_mapper.py +4 -4
- {biomapper-0.3.2 → biomapper-0.4.0}/biomapper/mapping/multi_provider_rag.py +185 -72
- biomapper-0.4.0/biomapper/mapping/rag/__init__.py +160 -0
- biomapper-0.4.0/biomapper/mapping/rag/data/default_prompts.json.template +18 -0
- biomapper-0.4.0/biomapper/mapping/rag/prompts.py +71 -0
- biomapper-0.4.0/biomapper/mapping/rag/store.py +134 -0
- biomapper-0.4.0/biomapper/monitoring/langfuse_tracker.py +138 -0
- biomapper-0.4.0/biomapper/monitoring/metrics.py +72 -0
- biomapper-0.4.0/biomapper/monitoring/traces.py +65 -0
- biomapper-0.4.0/biomapper/schemas/rag_schema.py +77 -0
- biomapper-0.4.0/biomapper/schemas/store_schema.py +12 -0
- biomapper-0.4.0/biomapper/utils/optimization.py +74 -0
- {biomapper-0.3.2 → biomapper-0.4.0}/pyproject.toml +26 -18
- biomapper-0.3.2/biomapper/mapping/rag_mapper.py +0 -207
- {biomapper-0.3.2 → biomapper-0.4.0}/LICENSE +0 -0
- {biomapper-0.3.2 → biomapper-0.4.0}/README.md +0 -0
- {biomapper-0.3.2 → biomapper-0.4.0}/biomapper/core/__init__.py +0 -0
- {biomapper-0.3.2 → biomapper-0.4.0}/biomapper/core/set_analysis.py +0 -0
- {biomapper-0.3.2 → biomapper-0.4.0}/biomapper/mapping/__init__.py +0 -0
- {biomapper-0.3.2/biomapper/mapping → biomapper-0.4.0/biomapper/mapping/clients}/refmet_client.py +0 -0
- {biomapper-0.3.2 → biomapper-0.4.0}/biomapper/py.typed +0 -0
- {biomapper-0.3.2 → biomapper-0.4.0}/biomapper/schemas/__init__.py +0 -0
- {biomapper-0.3.2 → biomapper-0.4.0}/biomapper/schemas/llm_schema.py +0 -0
- {biomapper-0.3.2 → biomapper-0.4.0}/biomapper/schemas/provider_schemas.py +0 -0
- {biomapper-0.3.2 → biomapper-0.4.0}/biomapper/standardization/__init__.py +0 -0
- {biomapper-0.3.2 → biomapper-0.4.0}/biomapper/standardization/ramp_client.py +0 -0
- {biomapper-0.3.2 → biomapper-0.4.0}/biomapper/utils/__init__.py +0 -0
|
@@ -1,6 +1,6 @@
|
|
|
1
1
|
Metadata-Version: 2.1
|
|
2
2
|
Name: biomapper
|
|
3
|
-
Version: 0.
|
|
3
|
+
Version: 0.4.0
|
|
4
4
|
Summary: A unified Python toolkit for biological data harmonization and ontology mapping
|
|
5
5
|
Home-page: https://github.com/arpanauts/biomapper
|
|
6
6
|
License: MIT
|
|
@@ -20,19 +20,26 @@ Classifier: Programming Language :: Python :: 3.13
|
|
|
20
20
|
Classifier: Topic :: Scientific/Engineering :: Bio-Informatics
|
|
21
21
|
Provides-Extra: api
|
|
22
22
|
Provides-Extra: full
|
|
23
|
+
Requires-Dist: chromadb (>=0.6.2,<0.7.0)
|
|
23
24
|
Requires-Dist: cloudpickle (>=3.0.0,<4.0.0)
|
|
25
|
+
Requires-Dist: cryptography (>=43.0.1)
|
|
24
26
|
Requires-Dist: dspy-ai (>=2.1.8,<3.0.0)
|
|
25
27
|
Requires-Dist: langfuse (>=2.57.1,<3.0.0)
|
|
26
28
|
Requires-Dist: libChEBIpy (==1.0.10)
|
|
27
29
|
Requires-Dist: matplotlib (>=3.8.0,<4.0.0)
|
|
30
|
+
Requires-Dist: numpy (<2.0.0)
|
|
28
31
|
Requires-Dist: openai (>=1.14.0,<2.0.0)
|
|
29
32
|
Requires-Dist: pandas (>=2.0.0,<3.0.0)
|
|
30
33
|
Requires-Dist: python-dotenv (>=1.0.1,<2.0.0)
|
|
34
|
+
Requires-Dist: python-multipart (==0.0.10)
|
|
31
35
|
Requires-Dist: pyyaml (>=6.0.1,<7.0.0)
|
|
36
|
+
Requires-Dist: rdkit (>=2023.9.1,<2024.0.0)
|
|
32
37
|
Requires-Dist: requests (>=2.25.1,<3.0.0)
|
|
33
38
|
Requires-Dist: seaborn (>=0.13.0,<0.14.0)
|
|
34
39
|
Requires-Dist: sqlalchemy (>=1.4.0,<2.0.0)
|
|
40
|
+
Requires-Dist: torch (>=2.2.0,<3.0.0)
|
|
35
41
|
Requires-Dist: tqdm (>=4.66.1,<5.0.0)
|
|
42
|
+
Requires-Dist: transformers (>=4.38.2,<5.0.0)
|
|
36
43
|
Requires-Dist: upsetplot (>=0.8.0,<0.9.0)
|
|
37
44
|
Requires-Dist: venn (>=0.1.3,<0.2.0)
|
|
38
45
|
Project-URL: Documentation, https://github.com/arpanauts/biomapper/blob/main/README.md
|
|
@@ -10,7 +10,7 @@ import pandas as pd
|
|
|
10
10
|
from tqdm import tqdm
|
|
11
11
|
import concurrent.futures
|
|
12
12
|
|
|
13
|
-
from ..mapping.uniprot_focused_mapper import UniprotFocusedMapper
|
|
13
|
+
from ..mapping.clients.uniprot_focused_mapper import UniprotFocusedMapper
|
|
14
14
|
|
|
15
15
|
|
|
16
16
|
class ProteinMapping(TypedDict):
|
|
@@ -0,0 +1 @@
|
|
|
1
|
+
"""API clients for various compound and metabolite databases."""
|
{biomapper-0.3.2/biomapper/mapping → biomapper-0.4.0/biomapper/mapping/clients}/chebi_client.py
RENAMED
|
@@ -2,11 +2,23 @@
|
|
|
2
2
|
|
|
3
3
|
import logging
|
|
4
4
|
from dataclasses import dataclass
|
|
5
|
-
from typing import Any, Optional
|
|
5
|
+
from typing import Any, Optional, TYPE_CHECKING
|
|
6
6
|
import requests
|
|
7
7
|
|
|
8
|
-
|
|
9
|
-
from libchebipy import ChebiEntity
|
|
8
|
+
if TYPE_CHECKING:
|
|
9
|
+
from libchebipy import ChebiEntity # type: ignore
|
|
10
|
+
|
|
11
|
+
def chebi_search(query: str) -> list[Any]:
|
|
12
|
+
...
|
|
13
|
+
else:
|
|
14
|
+
try:
|
|
15
|
+
from libchebipy import ChebiEntity, search as chebi_search # type: ignore
|
|
16
|
+
except ImportError:
|
|
17
|
+
ChebiEntity = Any
|
|
18
|
+
|
|
19
|
+
def chebi_search(query: str) -> list[Any]:
|
|
20
|
+
raise ImportError("libchebipy not installed")
|
|
21
|
+
|
|
10
22
|
|
|
11
23
|
logger = logging.getLogger(__name__)
|
|
12
24
|
|
|
@@ -116,7 +128,7 @@ class ChEBIClient:
|
|
|
116
128
|
logger.error("ChEBI entity processing failed: %s", str(e))
|
|
117
129
|
raise ChEBIError(f"Entity processing failed: {str(e)}") from e
|
|
118
130
|
|
|
119
|
-
def get_entity_by_id(self, chebi_id: str) -> ChEBIResult:
|
|
131
|
+
def get_entity_by_id(self, chebi_id: str) -> Optional[ChEBIResult]:
|
|
120
132
|
"""Get ChEBI entity information by ID.
|
|
121
133
|
|
|
122
134
|
Args:
|
{biomapper-0.3.2/biomapper/mapping → biomapper-0.4.0/biomapper/mapping/clients}/unichem_client.py
RENAMED
|
@@ -8,7 +8,7 @@ metabolite information across various chemical databases.
|
|
|
8
8
|
|
|
9
9
|
import logging
|
|
10
10
|
from dataclasses import dataclass
|
|
11
|
-
from typing import
|
|
11
|
+
from typing import Dict, List, Optional, Set, Tuple, Any
|
|
12
12
|
from pathlib import Path
|
|
13
13
|
|
|
14
14
|
import pandas as pd
|
|
@@ -82,7 +82,7 @@ class UniChemClient:
|
|
|
82
82
|
"inchikey": "inchikey",
|
|
83
83
|
}
|
|
84
84
|
|
|
85
|
-
def _get_empty_result(self) ->
|
|
85
|
+
def _get_empty_result(self) -> Dict[str, List[Any]]:
|
|
86
86
|
"""Return an empty result dictionary with all source types."""
|
|
87
87
|
return {
|
|
88
88
|
"chembl_ids": [],
|
|
@@ -94,8 +94,8 @@ class UniChemClient:
|
|
|
94
94
|
}
|
|
95
95
|
|
|
96
96
|
def _process_compound_result(
|
|
97
|
-
self, data:
|
|
98
|
-
) ->
|
|
97
|
+
self, data: List[Dict[str, Any]]
|
|
98
|
+
) -> Dict[str, List[Any]]:
|
|
99
99
|
"""Process compound results into categorized lists of IDs."""
|
|
100
100
|
result = self._get_empty_result()
|
|
101
101
|
|
|
@@ -121,19 +121,19 @@ class UniChemClient:
|
|
|
121
121
|
|
|
122
122
|
return result
|
|
123
123
|
|
|
124
|
-
def get_source_information(self) ->
|
|
124
|
+
def get_source_information(self) -> Dict[str, Any]:
|
|
125
125
|
"""Retrieve information about available data sources."""
|
|
126
126
|
try:
|
|
127
127
|
response = self.session.get(
|
|
128
128
|
f"{self.config.base_url}/sources", timeout=self.config.timeout
|
|
129
129
|
)
|
|
130
130
|
response.raise_for_status()
|
|
131
|
-
json_response:
|
|
131
|
+
json_response: Dict[str, Any] = response.json()
|
|
132
132
|
return json_response
|
|
133
133
|
except requests.RequestException as e:
|
|
134
134
|
raise UniChemError(f"Failed to get source information: {e}") from e
|
|
135
135
|
|
|
136
|
-
def get_structure_search(self, structure: str, search_type: str) ->
|
|
136
|
+
def get_structure_search(self, structure: str, search_type: str) -> Dict[str, Any]:
|
|
137
137
|
"""
|
|
138
138
|
Search for compounds by structure.
|
|
139
139
|
|
|
@@ -146,7 +146,7 @@ class UniChemClient:
|
|
|
146
146
|
|
|
147
147
|
Returns
|
|
148
148
|
-------
|
|
149
|
-
|
|
149
|
+
Dict[str, Any]
|
|
150
150
|
Search results
|
|
151
151
|
"""
|
|
152
152
|
valid_types = {"smiles", "inchi", "inchikey"}
|
|
@@ -159,7 +159,7 @@ class UniChemClient:
|
|
|
159
159
|
timeout=self.config.timeout,
|
|
160
160
|
)
|
|
161
161
|
response.raise_for_status()
|
|
162
|
-
json_response:
|
|
162
|
+
json_response: Dict[str, Any] = response.json()
|
|
163
163
|
return json_response
|
|
164
164
|
except requests.RequestException as e:
|
|
165
165
|
raise UniChemError(f"Structure search failed: {e}") from e
|
|
@@ -203,7 +203,7 @@ class UniChemClient:
|
|
|
203
203
|
|
|
204
204
|
def get_compound_info_by_src_id(
|
|
205
205
|
self, compound_id: str, src_db: str
|
|
206
|
-
) ->
|
|
206
|
+
) -> Dict[str, Any]:
|
|
207
207
|
"""
|
|
208
208
|
Retrieve compound cross-references from UniChem using the REST API.
|
|
209
209
|
"""
|
|
@@ -265,27 +265,27 @@ class UniChemClient:
|
|
|
265
265
|
def map_dataframe(
|
|
266
266
|
self,
|
|
267
267
|
df: pd.DataFrame,
|
|
268
|
-
id_columns:
|
|
269
|
-
target_sources: Optional[
|
|
268
|
+
id_columns: Dict[str, str],
|
|
269
|
+
target_sources: Optional[List[str]] = None,
|
|
270
270
|
prefix_ids: bool = True,
|
|
271
|
-
) ->
|
|
271
|
+
) -> Tuple[pd.DataFrame, MappingResult]:
|
|
272
272
|
"""Map compound IDs in a DataFrame using UniChem.
|
|
273
273
|
|
|
274
274
|
Parameters
|
|
275
275
|
----------
|
|
276
276
|
df : pd.DataFrame
|
|
277
277
|
Input DataFrame containing compound IDs
|
|
278
|
-
id_columns :
|
|
278
|
+
id_columns : Dict[str, str]
|
|
279
279
|
Mapping of source database names to column names
|
|
280
280
|
Example: {'hmdb': 'HMDB_ID', 'pubchem': 'PUBCHEM_ID'}
|
|
281
|
-
target_sources : Optional[
|
|
281
|
+
target_sources : Optional[List[str]], default=None
|
|
282
282
|
List of target databases to map to. If None, defaults to ['chembl', 'drugbank']
|
|
283
283
|
prefix_ids : bool, default=True
|
|
284
284
|
Whether to prefix mapped IDs with database name (e.g., CHEMBL_ID vs chembl)
|
|
285
285
|
|
|
286
286
|
Returns
|
|
287
287
|
-------
|
|
288
|
-
|
|
288
|
+
Tuple[pd.DataFrame, MappingResult]
|
|
289
289
|
Tuple of (mapped DataFrame, mapping statistics)
|
|
290
290
|
"""
|
|
291
291
|
# Validate inputs
|
|
@@ -306,7 +306,7 @@ class UniChemClient:
|
|
|
306
306
|
|
|
307
307
|
mapped_count = 0
|
|
308
308
|
total_count = 0
|
|
309
|
-
mapping_sources:
|
|
309
|
+
mapping_sources: Set[str] = set()
|
|
310
310
|
|
|
311
311
|
# Add INCHI column if not present
|
|
312
312
|
inchi_col_name = "INCHI_ID" if prefix_ids else "inchi"
|
|
@@ -380,8 +380,8 @@ class UniChemClient:
|
|
|
380
380
|
def map_csv(
|
|
381
381
|
self,
|
|
382
382
|
input_path: str | Path,
|
|
383
|
-
id_columns:
|
|
384
|
-
target_sources:
|
|
383
|
+
id_columns: Dict[str, str],
|
|
384
|
+
target_sources: List[str] | None = None,
|
|
385
385
|
output_path: str | Path | None = None,
|
|
386
386
|
prefix_ids: bool = True,
|
|
387
387
|
) -> MappingResult:
|
|
@@ -1,6 +1,8 @@
|
|
|
1
|
+
"""UniProt focused mapper for protein mapping."""
|
|
2
|
+
|
|
1
3
|
import time
|
|
2
4
|
from dataclasses import dataclass
|
|
3
|
-
from typing import
|
|
5
|
+
from typing import Dict, List, Optional, Set, Any
|
|
4
6
|
|
|
5
7
|
import requests
|
|
6
8
|
|
|
@@ -57,7 +59,7 @@ class UniprotFocusedMapper:
|
|
|
57
59
|
session.mount("https://", adapter)
|
|
58
60
|
return session
|
|
59
61
|
|
|
60
|
-
def get_available_mappings(self) ->
|
|
62
|
+
def get_available_mappings(self) -> Dict[str, List[str]]:
|
|
61
63
|
"""Get available mapping categories and target databases.
|
|
62
64
|
|
|
63
65
|
Returns:
|
|
@@ -68,7 +70,7 @@ class UniprotFocusedMapper:
|
|
|
68
70
|
for category, dbs in self.CORE_MAPPINGS.items()
|
|
69
71
|
}
|
|
70
72
|
|
|
71
|
-
def map_id(self, protein_id: str, target_db: str) ->
|
|
73
|
+
def map_id(self, protein_id: str, target_db: str) -> Dict[str, Any]:
|
|
72
74
|
"""Map a UniProt ID to a target database ID.
|
|
73
75
|
|
|
74
76
|
Args:
|
|
@@ -110,7 +112,7 @@ class UniprotFocusedMapper:
|
|
|
110
112
|
|
|
111
113
|
return {}
|
|
112
114
|
|
|
113
|
-
def _submit_job(self, from_db: str, to_db: str, ids:
|
|
115
|
+
def _submit_job(self, from_db: str, to_db: str, ids: List[str]) -> str | None:
|
|
114
116
|
"""Submit a mapping job to the UniProt API.
|
|
115
117
|
|
|
116
118
|
Args:
|
|
@@ -138,7 +140,7 @@ class UniprotFocusedMapper:
|
|
|
138
140
|
except (requests.exceptions.RequestException, ValueError):
|
|
139
141
|
return None
|
|
140
142
|
|
|
141
|
-
def _get_job_results(self, url: str) ->
|
|
143
|
+
def _get_job_results(self, url: str) -> Dict[str, Any]:
|
|
142
144
|
"""Get results from a completed mapping job.
|
|
143
145
|
|
|
144
146
|
Args:
|
|
@@ -160,8 +162,8 @@ class UniprotFocusedMapper:
|
|
|
160
162
|
return {}
|
|
161
163
|
|
|
162
164
|
def _map_to_database(
|
|
163
|
-
self, from_db: str, to_db: str, ids:
|
|
164
|
-
) ->
|
|
165
|
+
self, from_db: str, to_db: str, ids: List[str]
|
|
166
|
+
) -> Dict[str, Any]:
|
|
165
167
|
"""Map identifiers between databases."""
|
|
166
168
|
max_attempts = 20 # Maximum number of polling attempts
|
|
167
169
|
attempt = 0
|
|
@@ -223,7 +225,7 @@ class UniprotFocusedMapper:
|
|
|
223
225
|
return error.response.status_code in {500, 502, 503, 504}
|
|
224
226
|
return False
|
|
225
227
|
|
|
226
|
-
def _retry_request(self, url: str) ->
|
|
228
|
+
def _retry_request(self, url: str) -> Dict[str, Any]:
|
|
227
229
|
"""Retry a failed request with exponential backoff."""
|
|
228
230
|
for attempt in range(self.config.max_retries):
|
|
229
231
|
try:
|
|
@@ -235,8 +237,8 @@ class UniprotFocusedMapper:
|
|
|
235
237
|
continue
|
|
236
238
|
return {} # Return empty dict if all retries fail
|
|
237
239
|
|
|
238
|
-
def _make_request(self, url: str) ->
|
|
239
|
-
response:
|
|
240
|
+
def _make_request(self, url: str) -> Dict[str, Any]:
|
|
241
|
+
response: Dict[str, Any] = {}
|
|
240
242
|
try:
|
|
241
243
|
results_req = requests.get(url, timeout=self.config.timeout)
|
|
242
244
|
results_req.raise_for_status()
|
|
@@ -248,8 +250,8 @@ class UniprotFocusedMapper:
|
|
|
248
250
|
return response
|
|
249
251
|
|
|
250
252
|
def _generate_mappings(
|
|
251
|
-
self, proteins:
|
|
252
|
-
) ->
|
|
253
|
+
self, proteins: Set[str], categories: Optional[List[str]] = None
|
|
254
|
+
) -> Dict[str, Dict[str, List[str]]]:
|
|
253
255
|
"""Generate mappings for a set of proteins with progress tracking.
|
|
254
256
|
|
|
255
257
|
Args:
|
|
@@ -259,7 +261,7 @@ class UniprotFocusedMapper:
|
|
|
259
261
|
Returns:
|
|
260
262
|
Dict mapping protein IDs to their database mappings.
|
|
261
263
|
"""
|
|
262
|
-
mappings:
|
|
264
|
+
mappings: Dict[str, Dict[str, List[str]]] = {}
|
|
263
265
|
|
|
264
266
|
# Process proteins in chunks
|
|
265
267
|
for i in range(0, len(proteins), 100): # Using chunk_size of 100
|
|
@@ -270,8 +272,8 @@ class UniprotFocusedMapper:
|
|
|
270
272
|
return mappings
|
|
271
273
|
|
|
272
274
|
def _process_chunk(
|
|
273
|
-
self, chunk:
|
|
274
|
-
) ->
|
|
275
|
+
self, chunk: List[str], categories: Optional[List[str]] = None
|
|
276
|
+
) -> Dict[str, Dict[str, List[str]]]:
|
|
275
277
|
"""Process a chunk of proteins for mapping.
|
|
276
278
|
|
|
277
279
|
Args:
|
|
@@ -281,12 +283,12 @@ class UniprotFocusedMapper:
|
|
|
281
283
|
Returns:
|
|
282
284
|
Dictionary of mapping results for the chunk
|
|
283
285
|
"""
|
|
284
|
-
chunk_mappings:
|
|
286
|
+
chunk_mappings: Dict[str, Dict[str, List[str]]] = {}
|
|
285
287
|
|
|
286
288
|
for protein in chunk:
|
|
287
289
|
try:
|
|
288
290
|
# Get valid target databases for each category
|
|
289
|
-
target_dbs:
|
|
291
|
+
target_dbs: List[str] = []
|
|
290
292
|
for category in categories or ["Protein/Gene"]:
|
|
291
293
|
if category in self.CORE_MAPPINGS:
|
|
292
294
|
target_dbs.extend(
|
|
@@ -298,11 +300,11 @@ class UniprotFocusedMapper:
|
|
|
298
300
|
)
|
|
299
301
|
|
|
300
302
|
# Map to each target database
|
|
301
|
-
protein_mappings:
|
|
303
|
+
protein_mappings: Dict[str, List[str]] = {}
|
|
302
304
|
for target_db in target_dbs:
|
|
303
305
|
result = self.map_id(protein, target_db)
|
|
304
306
|
if result.get("results"):
|
|
305
|
-
mapped_ids:
|
|
307
|
+
mapped_ids: List[str] = []
|
|
306
308
|
for mapping in result["results"]:
|
|
307
309
|
to_id = mapping.get("to")
|
|
308
310
|
if isinstance(to_id, dict):
|
|
@@ -0,0 +1 @@
|
|
|
1
|
+
"""Embeddings package for managing different embedding models."""
|
|
@@ -0,0 +1,98 @@
|
|
|
1
|
+
"""Embedding model management for compound mapping."""
|
|
2
|
+
from abc import ABC, abstractmethod
|
|
3
|
+
from typing import List, Optional, cast
|
|
4
|
+
import torch
|
|
5
|
+
from transformers import AutoModel, AutoTokenizer # type: ignore
|
|
6
|
+
|
|
7
|
+
|
|
8
|
+
class EmbeddingManager(ABC):
|
|
9
|
+
"""Abstract base class for embedding managers."""
|
|
10
|
+
|
|
11
|
+
@abstractmethod
|
|
12
|
+
def embed_text(self, text: str) -> List[float]:
|
|
13
|
+
"""Embed a single text string."""
|
|
14
|
+
pass
|
|
15
|
+
|
|
16
|
+
@abstractmethod
|
|
17
|
+
def embed_batch(self, texts: List[str]) -> List[List[float]]:
|
|
18
|
+
"""Embed a batch of text strings."""
|
|
19
|
+
pass
|
|
20
|
+
|
|
21
|
+
|
|
22
|
+
class HuggingFaceEmbeddingManager(EmbeddingManager):
|
|
23
|
+
"""Manages embeddings using HuggingFace models."""
|
|
24
|
+
|
|
25
|
+
def __init__(
|
|
26
|
+
self,
|
|
27
|
+
model_name: str = "sentence-transformers/all-mpnet-base-v2",
|
|
28
|
+
device: Optional[str] = None,
|
|
29
|
+
):
|
|
30
|
+
self.model_name = model_name
|
|
31
|
+
self.device = device or ("cuda" if torch.cuda.is_available() else "cpu")
|
|
32
|
+
|
|
33
|
+
self.tokenizer = AutoTokenizer.from_pretrained(model_name)
|
|
34
|
+
self.model = AutoModel.from_pretrained(model_name).to(self.device)
|
|
35
|
+
self.model.eval()
|
|
36
|
+
|
|
37
|
+
def embed_text(self, text: str) -> List[float]:
|
|
38
|
+
"""Embed a single text string."""
|
|
39
|
+
# Tokenize input
|
|
40
|
+
inputs = self.tokenizer(
|
|
41
|
+
text, return_tensors="pt", padding=True, truncation=True, max_length=512
|
|
42
|
+
)
|
|
43
|
+
|
|
44
|
+
# Move tensors to device
|
|
45
|
+
device_inputs = {k: v.to(self.device) for k, v in inputs.items()}
|
|
46
|
+
|
|
47
|
+
with torch.no_grad():
|
|
48
|
+
outputs = self.model(**device_inputs)
|
|
49
|
+
# Use CLS token embedding
|
|
50
|
+
embedding = outputs.last_hidden_state[:, 0, :].cpu().numpy()
|
|
51
|
+
|
|
52
|
+
return cast(List[float], embedding[0].tolist())
|
|
53
|
+
|
|
54
|
+
def embed_batch(self, texts: List[str]) -> List[List[float]]:
|
|
55
|
+
"""Embed a batch of text strings."""
|
|
56
|
+
# Tokenize inputs and get dictionary of tensors
|
|
57
|
+
inputs = self.tokenizer(
|
|
58
|
+
texts, return_tensors="pt", padding=True, truncation=True, max_length=512
|
|
59
|
+
)
|
|
60
|
+
|
|
61
|
+
# Convert inputs to device if it's not already a dictionary
|
|
62
|
+
if not isinstance(inputs, dict):
|
|
63
|
+
inputs = dict(inputs)
|
|
64
|
+
device_inputs = {k: v.to(self.device) for k, v in inputs.items()}
|
|
65
|
+
|
|
66
|
+
with torch.no_grad():
|
|
67
|
+
outputs = self.model(**device_inputs)
|
|
68
|
+
# Use CLS token embeddings for each text in batch
|
|
69
|
+
embeddings = outputs.last_hidden_state[:, 0, :].cpu().numpy()
|
|
70
|
+
|
|
71
|
+
return cast(List[List[float]], embeddings.tolist())
|
|
72
|
+
|
|
73
|
+
|
|
74
|
+
class OpenAIEmbeddingManager(EmbeddingManager):
|
|
75
|
+
"""Manages embeddings using OpenAI's API."""
|
|
76
|
+
|
|
77
|
+
def __init__(
|
|
78
|
+
self, api_key: Optional[str] = None, model: str = "text-embedding-ada-002"
|
|
79
|
+
):
|
|
80
|
+
try:
|
|
81
|
+
import openai
|
|
82
|
+
except ImportError:
|
|
83
|
+
raise ImportError(
|
|
84
|
+
"OpenAI package not found. Install with 'pip install openai'"
|
|
85
|
+
)
|
|
86
|
+
|
|
87
|
+
self.client = openai.OpenAI(api_key=api_key)
|
|
88
|
+
self.model = model
|
|
89
|
+
|
|
90
|
+
def embed_text(self, text: str) -> List[float]:
|
|
91
|
+
"""Embed a single text string."""
|
|
92
|
+
response = self.client.embeddings.create(model=self.model, input=text)
|
|
93
|
+
return response.data[0].embedding
|
|
94
|
+
|
|
95
|
+
def embed_batch(self, texts: List[str]) -> List[List[float]]:
|
|
96
|
+
"""Embed a batch of text strings."""
|
|
97
|
+
response = self.client.embeddings.create(model=self.model, input=texts)
|
|
98
|
+
return [data.embedding for data in response.data]
|
|
@@ -4,15 +4,15 @@ import logging
|
|
|
4
4
|
from dataclasses import dataclass
|
|
5
5
|
from enum import Enum
|
|
6
6
|
from pathlib import Path
|
|
7
|
-
from typing import
|
|
7
|
+
from typing import List, Optional, Callable, Any
|
|
8
8
|
import re
|
|
9
9
|
from collections import defaultdict
|
|
10
10
|
|
|
11
11
|
import pandas as pd
|
|
12
12
|
|
|
13
|
-
from .chebi_client import ChEBIClient
|
|
14
|
-
from .refmet_client import RefMetClient
|
|
15
|
-
from .unichem_client import UniChemClient
|
|
13
|
+
from .clients.chebi_client import ChEBIClient
|
|
14
|
+
from .clients.refmet_client import RefMetClient
|
|
15
|
+
from .clients.unichem_client import UniChemClient
|
|
16
16
|
|
|
17
17
|
logger = logging.getLogger(__name__)
|
|
18
18
|
|