jimstechai commited on
Commit
1ee024e
·
verified ·
1 Parent(s): 5694c5e

Fit embeddings to Vectorize dimension

Browse files
Files changed (2) hide show
  1. __pycache__/app.cpython-313.pyc +0 -0
  2. app.py +27 -5
__pycache__/app.cpython-313.pyc ADDED
Binary file (10 kB). View file
 
app.py CHANGED
@@ -16,6 +16,7 @@ AGENT_TOKEN = os.environ.get("JIMS_RENDER_AGENT_TOKEN") or os.environ.get("JIMS_
16
  MODEL_NAME = os.environ.get("JIMS_EMBEDDING_MODEL", "intfloat/multilingual-e5-small")
17
  ARTIFACT_ID = os.environ.get("JIMS_ACTIVE_ARTIFACT_ID", "hf_space_encoder")
18
  HASH_FALLBACK = os.environ.get("JIMS_EMBEDDING_HASH_FALLBACK_ENABLED", "true").lower() == "true"
 
19
 
20
  model = None
21
  model_error = ""
@@ -45,14 +46,23 @@ def normalize(values: list[float]) -> list[float]:
45
  return [round(value / norm, 8) for value in values]
46
 
47
 
48
- def hash_embed(text: str, dimensions: int = 768) -> list[float]:
 
 
 
 
 
 
 
 
 
49
  vector = [0.0] * dimensions
50
  for token in text.lower().split():
51
  digest = hashlib.sha256(token.encode("utf-8")).digest()
52
  idx = int.from_bytes(digest[:2], "big") % dimensions
53
  sign = 1.0 if digest[2] % 2 == 0 else -1.0
54
  vector[idx] += sign
55
- return normalize(vector)
56
 
57
 
58
  def prefixed(text: str, purpose: str) -> str:
@@ -84,7 +94,7 @@ def embed_texts(texts: list[str], purpose: str) -> tuple[list[list[float]], bool
84
  return [hash_embed(text) for text in texts], True, "hash_fallback"
85
  try:
86
  vectors = loaded.encode([prefixed(text[:16000], purpose) for text in texts], normalize_embeddings=True).tolist()
87
- return [[float(value) for value in vector] for vector in vectors], False, MODEL_NAME
88
  except Exception as exc:
89
  if not HASH_FALLBACK:
90
  raise HTTPException(status_code=500, detail=str(exc)) from exc
@@ -98,7 +108,13 @@ def root() -> dict[str, str]:
98
 
99
  @app.get("/health")
100
  def health() -> dict[str, Any]:
101
- return {"status": "ok", "service": "embedding-service", "model": MODEL_NAME, "model_loaded": model is not None}
 
 
 
 
 
 
102
 
103
 
104
  @app.get("/ready")
@@ -143,7 +159,13 @@ def encode(payload: dict[str, Any]) -> dict[str, Any]:
143
 
144
  @app.get("/v1/artifact/current")
145
  def current_artifact() -> dict[str, Any]:
146
- return {"artifact_id": ARTIFACT_ID, "model": MODEL_NAME, "loaded": model is not None, "error": model_error}
 
 
 
 
 
 
147
 
148
 
149
  @app.post("/v1/reload-artifact", dependencies=[Depends(verify_token)])
 
16
  MODEL_NAME = os.environ.get("JIMS_EMBEDDING_MODEL", "intfloat/multilingual-e5-small")
17
  ARTIFACT_ID = os.environ.get("JIMS_ACTIVE_ARTIFACT_ID", "hf_space_encoder")
18
  HASH_FALLBACK = os.environ.get("JIMS_EMBEDDING_HASH_FALLBACK_ENABLED", "true").lower() == "true"
19
+ TARGET_DIMENSIONS = max(int(os.environ.get("JIMS_EMBEDDING_DIMENSIONS", "768") or "768"), 1)
20
 
21
  model = None
22
  model_error = ""
 
46
  return [round(value / norm, 8) for value in values]
47
 
48
 
49
+ def fit_dimensions(values: list[float]) -> list[float]:
50
+ if len(values) == TARGET_DIMENSIONS:
51
+ return normalize(values)
52
+ if len(values) > TARGET_DIMENSIONS:
53
+ return normalize(values[:TARGET_DIMENSIONS])
54
+ return normalize([*values, *([0.0] * (TARGET_DIMENSIONS - len(values)))])
55
+
56
+
57
+ def hash_embed(text: str) -> list[float]:
58
+ dimensions = TARGET_DIMENSIONS
59
  vector = [0.0] * dimensions
60
  for token in text.lower().split():
61
  digest = hashlib.sha256(token.encode("utf-8")).digest()
62
  idx = int.from_bytes(digest[:2], "big") % dimensions
63
  sign = 1.0 if digest[2] % 2 == 0 else -1.0
64
  vector[idx] += sign
65
+ return fit_dimensions(vector)
66
 
67
 
68
  def prefixed(text: str, purpose: str) -> str:
 
94
  return [hash_embed(text) for text in texts], True, "hash_fallback"
95
  try:
96
  vectors = loaded.encode([prefixed(text[:16000], purpose) for text in texts], normalize_embeddings=True).tolist()
97
+ return [fit_dimensions([float(value) for value in vector]) for vector in vectors], False, MODEL_NAME
98
  except Exception as exc:
99
  if not HASH_FALLBACK:
100
  raise HTTPException(status_code=500, detail=str(exc)) from exc
 
108
 
109
  @app.get("/health")
110
  def health() -> dict[str, Any]:
111
+ return {
112
+ "status": "ok",
113
+ "service": "embedding-service",
114
+ "model": MODEL_NAME,
115
+ "dimension": TARGET_DIMENSIONS,
116
+ "model_loaded": model is not None,
117
+ }
118
 
119
 
120
  @app.get("/ready")
 
159
 
160
  @app.get("/v1/artifact/current")
161
  def current_artifact() -> dict[str, Any]:
162
+ return {
163
+ "artifact_id": ARTIFACT_ID,
164
+ "model": MODEL_NAME,
165
+ "dimension": TARGET_DIMENSIONS,
166
+ "loaded": model is not None,
167
+ "error": model_error,
168
+ }
169
 
170
 
171
  @app.post("/v1/reload-artifact", dependencies=[Depends(verify_token)])