Skip to content

Commit 37137ef

Browse files
authored
Merge pull request #9 from MementoRC/integrate/task-3-semantic-search
feat: integrate Task-3 enhanced semantic search features
2 parents 967bb6e + ea9eb9a commit 37137ef

8 files changed

Lines changed: 1160 additions & 991 deletions

src/uckn/core/atoms/multi_modal_embeddings.py

Lines changed: 104 additions & 129 deletions
Original file line numberDiff line numberDiff line change
@@ -7,12 +7,61 @@
77

88
import hashlib
99
import logging
10+
import os
1011
import threading
1112
from typing import Any
1213

1314
import numpy as np
1415

15-
from ..ml_environment_manager import get_ml_manager
16+
# Defensive import logic for torch and sentence-transformers
17+
SENTENCE_TRANSFORMERS_AVAILABLE = False
18+
TRANSFORMERS_AVAILABLE = False
19+
SentenceTransformer = None
20+
AutoTokenizer = None
21+
AutoModel = None
22+
torch = None
23+
24+
_DISABLE_TORCH = os.environ.get("UCKN_DISABLE_TORCH", "0") == "1"
25+
26+
if not _DISABLE_TORCH:
27+
# Try importing torch and transformers defensively
28+
try:
29+
try:
30+
import torch
31+
except Exception:
32+
torch = None # type: ignore[assignment]
33+
# Log or print for debugging, but do not raise
34+
else:
35+
try:
36+
from transformers import AutoModel, AutoTokenizer
37+
38+
TRANSFORMERS_AVAILABLE = True
39+
except Exception:
40+
AutoTokenizer = None # type: ignore[assignment]
41+
AutoModel = None # type: ignore[assignment]
42+
TRANSFORMERS_AVAILABLE = False
43+
except Exception:
44+
torch = None
45+
AutoTokenizer = None
46+
AutoModel = None
47+
TRANSFORMERS_AVAILABLE = False
48+
49+
# Try importing sentence-transformers defensively
50+
try:
51+
from sentence_transformers import SentenceTransformer
52+
53+
SENTENCE_TRANSFORMERS_AVAILABLE = True
54+
except Exception:
55+
SentenceTransformer = None
56+
SENTENCE_TRANSFORMERS_AVAILABLE = False
57+
else:
58+
# Torch is disabled by environment variable
59+
torch = None
60+
AutoTokenizer = None
61+
AutoModel = None
62+
SentenceTransformer = None
63+
TRANSFORMERS_AVAILABLE = False
64+
SENTENCE_TRANSFORMERS_AVAILABLE = False
1665

1766

1867
class MultiModalEmbeddings:
@@ -31,25 +80,28 @@ class MultiModalEmbeddings:
3180

3281
def __init__(self, device: str | None = None):
3382
self._logger = logging.getLogger(__name__)
34-
self._ml_manager = get_ml_manager()
35-
36-
# Use ML manager to determine device
37-
self.device = device or self._ml_manager.get_device()
83+
# Defensive: If torch is unavailable, always use cpu
84+
self.device: str = "cpu" # Default value
85+
if (
86+
torch is not None
87+
and hasattr(torch, "cuda")
88+
and callable(getattr(torch.cuda, "is_available", None))
89+
):
90+
self.device = device or ("cuda" if torch.cuda.is_available() else "cpu")
3891
self._lock = threading.Lock()
3992

4093
# Model loading
4194
self.code_tokenizer = None
4295
self.code_model = None
4396
self.text_model = None
4497

45-
# Initialize models based on environment capabilities
46-
if self._ml_manager.should_use_real_ml():
98+
# Only initialize models if not disabled
99+
if not _DISABLE_TORCH:
47100
self._init_code_model()
48101
self._init_text_model()
49102
else:
50-
env_info = self._ml_manager.get_environment_info()
51-
self._logger.info(
52-
f"Using fallback embeddings - Environment: {env_info['environment']}"
103+
self._logger.warning(
104+
"Torch and transformers are disabled by environment variable."
53105
)
54106

55107
# In-memory cache for embeddings
@@ -62,130 +114,56 @@ def is_available(self) -> bool:
62114
Returns:
63115
bool: True if at least one embedding model is available, False otherwise.
64116
"""
65-
# Component is always available - either real ML or fallbacks
66-
caps = self._ml_manager.capabilities
67-
68-
has_real_models = (
69-
caps.sentence_transformers and self.text_model is not None
70-
) or (
71-
caps.transformers
117+
# Component is available if at least one model is initialized
118+
# or if we have the basic dependencies available
119+
has_text_model = SENTENCE_TRANSFORMERS_AVAILABLE and self.text_model is not None
120+
has_code_model = (
121+
TRANSFORMERS_AVAILABLE
72122
and self.code_model is not None
73123
and self.code_tokenizer is not None
74124
)
75125

76-
# Always available: either real models or fallback embeddings
77-
return has_real_models or caps.fallback_embeddings
78-
79-
def _generate_fake_embedding(self, text: str, dim: int = 384) -> list[float]:
80-
"""Generate deterministic fake embedding for testing when ML models unavailable."""
81-
import hashlib
82-
import re
83-
84-
# Extract words for semantic features
85-
words = set(re.findall(r"\w+", text.lower()))
86-
87-
# Create word-based features for first part of embedding
88-
word_features = []
89-
common_words = {
90-
"add",
91-
"sum",
92-
"two",
93-
"numbers",
94-
"values",
95-
"def",
96-
"function",
97-
"class",
98-
"setting",
99-
"config",
100-
"error",
101-
"exception",
102-
"true",
103-
"false",
104-
"return",
105-
"division",
106-
"zero",
107-
"traceback",
108-
"zerodivisionerror",
109-
"by",
110-
}
111-
112-
for common_word in sorted(common_words):
113-
if common_word in words:
114-
word_features.append(1.0)
115-
else:
116-
word_features.append(0.0)
117-
118-
# Pad or truncate to half the dimension
119-
half_dim = dim // 2
120-
while len(word_features) < half_dim:
121-
word_features.append(0.0)
122-
word_features = word_features[:half_dim]
123-
124-
# Create hash-based features for second half
125-
hash_obj = hashlib.md5(text.encode(), usedforsecurity=False)
126-
hash_bytes = hash_obj.digest()
127-
hash_features = []
128-
129-
for i in range(dim - half_dim):
130-
byte_val = hash_bytes[i % len(hash_bytes)]
131-
# Smaller range for hash features to reduce noise
132-
norm_val = (byte_val / 255.0) * 0.2 - 0.1
133-
hash_features.append(norm_val)
134-
135-
# Combine features
136-
embedding = word_features + hash_features
137-
138-
# Normalize to unit vector
139-
norm = sum(x**2 for x in embedding) ** 0.5
140-
if norm > 0:
141-
embedding = [x / norm for x in embedding]
142-
143-
return embedding
126+
# Available if we have at least one working model or basic dependencies
127+
return has_text_model or has_code_model or SENTENCE_TRANSFORMERS_AVAILABLE
144128

145129
def _init_code_model(self):
146-
if not self._ml_manager.capabilities.transformers:
147-
self._logger.debug(
130+
if (
131+
not TRANSFORMERS_AVAILABLE
132+
or AutoTokenizer is None
133+
or AutoModel is None
134+
or torch is None
135+
):
136+
self._logger.warning(
148137
"Transformers not available. Code embedding will fallback to text model."
149138
)
150139
return
151-
152140
try:
153-
self.code_model, self.code_tokenizer = (
154-
self._ml_manager.get_transformers_model(self._CODE_MODEL_NAME)
141+
self.code_tokenizer = AutoTokenizer.from_pretrained(self._CODE_MODEL_NAME)
142+
self.code_model = AutoModel.from_pretrained(self._CODE_MODEL_NAME).to(
143+
self.device
155144
)
156-
if self.code_model and self.code_tokenizer:
157-
self._logger.info(f"Loaded code model: {self._CODE_MODEL_NAME}")
158-
else:
159-
self._logger.warning(
160-
f"Failed to load code model '{self._CODE_MODEL_NAME}'. Falling back to text model."
161-
)
145+
self._logger.info(f"Loaded code model: {self._CODE_MODEL_NAME}")
162146
except Exception as e:
163147
self._logger.warning(
164-
f"Error loading code model '{self._CODE_MODEL_NAME}': {e}. Falling back to text model."
148+
f"Failed to load code model '{self._CODE_MODEL_NAME}': {e}. Falling back to text model."
165149
)
166150
self.code_tokenizer = None
167151
self.code_model = None
168152

169153
def _init_text_model(self):
170-
if not self._ml_manager.capabilities.sentence_transformers:
171-
self._logger.debug(
172-
"SentenceTransformers not available. Text embedding will use fallbacks."
154+
if not SENTENCE_TRANSFORMERS_AVAILABLE or SentenceTransformer is None:
155+
self._logger.warning(
156+
"SentenceTransformers not available. Text embedding will be disabled."
173157
)
174158
return
175-
176159
try:
177-
self.text_model = self._ml_manager.get_sentence_transformer(
178-
self._TEXT_MODEL_NAME
160+
self.text_model = SentenceTransformer(
161+
self._TEXT_MODEL_NAME, device=self.device
179162
)
180-
if self.text_model:
181-
self._logger.info(f"Loaded text model: {self._TEXT_MODEL_NAME}")
182-
else:
183-
self._logger.warning(
184-
f"Failed to load text model '{self._TEXT_MODEL_NAME}'. Using fallbacks."
185-
)
163+
self._logger.info(f"Loaded text model: {self._TEXT_MODEL_NAME}")
186164
except Exception as e:
187-
self._logger.warning(
188-
f"Error loading text model '{self._TEXT_MODEL_NAME}': {e}. Using fallbacks."
165+
self._logger.error(
166+
f"Failed to load text model '{self._TEXT_MODEL_NAME}': {e}"
189167
)
190168
self.text_model = None
191169

@@ -207,17 +185,12 @@ def _embed_code(self, code: str) -> list[float] | None:
207185
cached = self._get_cached_embedding(key)
208186
if cached:
209187
return cached
210-
if (
211-
self.code_model
212-
and self.code_tokenizer
213-
and self._ml_manager.capabilities.torch
214-
):
188+
if self.code_model and self.code_tokenizer and torch is not None:
215189
try:
216190
inputs = self.code_tokenizer(
217191
code, return_tensors="pt", truncation=True, max_length=256
218192
)
219193
inputs = {k: v.to(self.device) for k, v in inputs.items()}
220-
torch = self._ml_manager._get_import("torch")
221194
with torch.no_grad():
222195
outputs = self.code_model(**inputs)
223196
# Use [CLS] token representation
@@ -248,9 +221,7 @@ def _embed_text(self, text: str) -> list[float] | None:
248221
return embedding
249222
except Exception as e:
250223
self._logger.error(f"Text embedding failed: {e}")
251-
252-
# Fallback: Generate deterministic fake embedding for testing
253-
return self._generate_fake_embedding(text)
224+
return None
254225

255226
def _embed_config(self, config: str) -> list[float] | None:
256227
# Simple tokenization: split on newlines, colons, equals, etc.
@@ -284,17 +255,21 @@ def embed(
284255
data_type = data["type"]
285256
data = data["content"]
286257

258+
# Ensure data is a string at this point
259+
if not isinstance(data, str):
260+
self._logger.warning(
261+
f"Expected string data, got {type(data)}. Converting to string."
262+
)
263+
data = str(data)
264+
287265
if data_type == "auto":
288266
# Heuristic: detect type
289-
if isinstance(data, str):
290-
if data.strip().startswith("def ") or data.strip().startswith("class "):
291-
data_type = "code"
292-
elif "=" in data and "\n" in data:
293-
data_type = "config"
294-
elif "Traceback" in data or "Exception" in data:
295-
data_type = "error"
296-
else:
297-
data_type = "text"
267+
if data.strip().startswith("def ") or data.strip().startswith("class "):
268+
data_type = "code"
269+
elif "=" in data and "\n" in data:
270+
data_type = "config"
271+
elif "Traceback" in data or "Exception" in data:
272+
data_type = "error"
298273
else:
299274
data_type = "text"
300275

0 commit comments

Comments
 (0)