MAS-AI-0000 commited on
Commit
3ec9229
Β·
verified Β·
1 Parent(s): 4fbc687

Update text_embedder.py

Browse files
Files changed (1) hide show
  1. text_embedder.py +42 -21
text_embedder.py CHANGED
@@ -23,13 +23,18 @@ from __future__ import annotations
23
  import os
24
  import sys
25
  from typing import Optional
26
-
27
  import numpy as np
28
  import torch
29
  import torch.nn.functional as F
30
  from pathlib import Path
31
  from huggingface_hub import snapshot_download
32
 
 
 
 
 
 
33
  # ---------------------------------------------------------------------------
34
  # Make the local 'detree' package importable
35
  # ---------------------------------------------------------------------------
@@ -39,11 +44,11 @@ if _current_dir not in sys.path:
39
 
40
  try:
41
  from detree.model.text_embedding import TextEmbeddingModel
 
42
  except ImportError as _e:
43
- print(f"Warning: could not import TextEmbeddingModel ({_e}). Text embedding will return zeros.")
44
  TextEmbeddingModel = None
45
 
46
-
47
  # ---------------------------------------------------------------------------
48
  # Configuration
49
  # ---------------------------------------------------------------------------
@@ -56,6 +61,8 @@ REPO_ID = "MAS-AI-0000/Authentica"
56
  TEXT_SUBFOLDER = "Lib/Models/Text" # where config.json/model.safetensors live in the repo
57
  EMBEDDING_FILE = "priori1_center10k.pt"
58
  _TEXT_DIR = None
 
 
59
 
60
  try:
61
  # download a local snapshot of just the Text folder and point _TEXT_DIR at it
@@ -78,18 +85,20 @@ except Exception as e:
78
  _model: Optional[object] = None
79
  _tokenizer: Optional[object] = None
80
 
81
-
82
  def _init() -> None:
83
  global _model, _tokenizer
84
 
 
 
85
  if TextEmbeddingModel is None:
86
- print("TextEmbedder: TextEmbeddingModel unavailable β€” embedding disabled.")
87
  return
88
 
89
  if not os.path.exists(_TEXT_DIR):
90
- print(f"TextEmbedder: model directory not found at {_TEXT_DIR!r} β€” embedding disabled.")
91
  return
92
 
 
93
  try:
94
  _model = TextEmbeddingModel(
95
  _TEXT_DIR,
@@ -99,9 +108,10 @@ def _init() -> None:
99
  ).to(DEVICE)
100
  _model.eval()
101
  _tokenizer = _model.tokenizer
102
- print(f"TextEmbedder: model loaded from {_TEXT_DIR!r}")
 
103
  except Exception as exc:
104
- print(f"TextEmbedder: error loading model: {exc}")
105
 
106
 
107
  _init()
@@ -133,23 +143,34 @@ def get_text_embedding(
133
  ``np.ndarray`` of shape ``(1, embedding_dim)`` and dtype float32.
134
  """
135
  if _model is None or _tokenizer is None:
 
136
  return np.zeros((1, 1), dtype=np.float32)
137
 
138
- encoded = _tokenizer(
139
- [text],
140
- return_tensors="pt",
141
- max_length=max_length,
142
- padding="max_length",
143
- truncation=True,
144
- )
145
- encoded = {k: v.to(DEVICE) for k, v in encoded.items()}
 
 
 
 
146
 
147
- # Shape returned by model with hidden_states=True: (batch, num_layers, dim)
148
- embeddings = _model(encoded, hidden_states=True)
149
- embeddings = F.normalize(embeddings, dim=-1) # normalise feature dim
 
150
 
151
- # embeddings: (1, num_layers, dim) β†’ select layer β†’ (1, dim)
152
- selected = embeddings[:, layer, :] # supports negative indexing
 
 
 
 
 
153
 
154
  return selected.cpu().numpy().astype(np.float32)
155
 
 
23
  import os
24
  import sys
25
  from typing import Optional
26
+ import logging
27
  import numpy as np
28
  import torch
29
  import torch.nn.functional as F
30
  from pathlib import Path
31
  from huggingface_hub import snapshot_download
32
 
33
+
34
+
35
+ log = logging.getLogger("text_embedder")
36
+ logging.basicConfig(level=logging.INFO, format="%(levelname)s [%(name)s] %(message)s")
37
+
38
  # ---------------------------------------------------------------------------
39
  # Make the local 'detree' package importable
40
  # ---------------------------------------------------------------------------
 
44
 
45
  try:
46
  from detree.model.text_embedding import TextEmbeddingModel
47
+ log.info("TextEmbeddingModel imported successfully.")
48
  except ImportError as _e:
49
+ log.error(f"Could not import TextEmbeddingModel: {_e}")
50
  TextEmbeddingModel = None
51
 
 
52
  # ---------------------------------------------------------------------------
53
  # Configuration
54
  # ---------------------------------------------------------------------------
 
61
  TEXT_SUBFOLDER = "Lib/Models/Text" # where config.json/model.safetensors live in the repo
62
  EMBEDDING_FILE = "priori1_center10k.pt"
63
  _TEXT_DIR = None
64
+ log.info(f"[config] device={DEVICE!r} max_length={MAX_LENGTH} pooling={POOLING!r}")
65
+
66
 
67
  try:
68
  # download a local snapshot of just the Text folder and point _TEXT_DIR at it
 
85
  _model: Optional[object] = None
86
  _tokenizer: Optional[object] = None
87
 
 
88
  def _init() -> None:
89
  global _model, _tokenizer
90
 
91
+ log.info("_init: starting TextEmbedder initialisation.")
92
+
93
  if TextEmbeddingModel is None:
94
+ log.error("_init: TextEmbeddingModel is None β€” check import error above. Embedding disabled.")
95
  return
96
 
97
  if not os.path.exists(_TEXT_DIR):
98
+ log.error(f"_init: model directory not found at {_TEXT_DIR!r} β€” embedding disabled.")
99
  return
100
 
101
+ log.info(f"_init: loading TextEmbeddingModel from {_TEXT_DIR!r} on device={DEVICE!r} ...")
102
  try:
103
  _model = TextEmbeddingModel(
104
  _TEXT_DIR,
 
108
  ).to(DEVICE)
109
  _model.eval()
110
  _tokenizer = _model.tokenizer
111
+ log.info(f"_init: model loaded OK. tokenizer type={type(_tokenizer).__name__!r}")
112
+ log.info(f"_init: model device={next(_model.parameters()).device}")
113
  except Exception as exc:
114
+ log.exception(f"_init: error loading model: {exc}")
115
 
116
 
117
  _init()
 
143
  ``np.ndarray`` of shape ``(1, embedding_dim)`` and dtype float32.
144
  """
145
  if _model is None or _tokenizer is None:
146
+ log.error("get_text_embedding: model or tokenizer is None β€” returning zeros. Check _init logs.")
147
  return np.zeros((1, 1), dtype=np.float32)
148
 
149
+ log.info(f"get_text_embedding: input text length={len(text)} chars, layer={layer}")
150
+ try:
151
+ encoded = _tokenizer(
152
+ [text],
153
+ return_tensors="pt",
154
+ max_length=max_length,
155
+ padding="max_length",
156
+ truncation=True,
157
+ )
158
+ log.info(f"get_text_embedding: tokenised keys={list(encoded.keys())} "
159
+ f"input_ids shape={encoded['input_ids'].shape}")
160
+ encoded = {k: v.to(DEVICE) for k, v in encoded.items()}
161
 
162
+ # Shape returned by model with hidden_states=True: (batch, num_layers, dim)
163
+ embeddings = _model(encoded, hidden_states=True)
164
+ log.info(f"get_text_embedding: raw embeddings shape={tuple(embeddings.shape)}")
165
+ embeddings = F.normalize(embeddings, dim=-1) # normalise feature dim
166
 
167
+ # embeddings: (1, num_layers, dim) β†’ select layer β†’ (1, dim)
168
+ selected = embeddings[:, layer, :] # supports negative indexing
169
+ log.info(f"get_text_embedding: selected layer={layer} output shape={tuple(selected.shape)} "
170
+ f"norm={selected.norm(dim=-1).item():.4f}")
171
+ except Exception as exc:
172
+ log.exception(f"get_text_embedding: failed during inference: {exc}")
173
+ return np.zeros((1, 1), dtype=np.float32)
174
 
175
  return selected.cpu().numpy().astype(np.float32)
176