Soumik-404 commited on
Commit
e6ffe6d
·
1 Parent(s): f40136f
Files changed (1) hide show
  1. app/services/embeddings_service.py +20 -10
app/services/embeddings_service.py CHANGED
@@ -5,8 +5,11 @@ import os
5
  from typing import Dict, List, Optional
6
 
7
  import numpy as np
 
 
8
  from PIL import Image
9
  from sentence_transformers import SentenceTransformer
 
10
 
11
  _logger = logging.getLogger(__name__)
12
 
@@ -33,7 +36,8 @@ class EmbeddingService:
33
  self._device = "cpu"
34
 
35
  self._loaded_dimensions: List[int] = []
36
- self._vision_model: Optional[SentenceTransformer] = None
 
37
  self._vision_loaded = False
38
 
39
  def load_model(self, dimension: int) -> None:
@@ -77,12 +81,14 @@ class EmbeddingService:
77
  _logger.info("Patched vision model config: n_inner float -> int")
78
 
79
  _logger.info("Loading vision embedding model from %s", source)
80
- self._vision_model = SentenceTransformer(
 
81
  source,
82
- device=self._device,
83
  trust_remote_code=True,
 
84
  )
85
  self._vision_model.eval()
 
86
  self._vision_loaded = True
87
  _logger.info("Loaded vision embedding model (device=%s)", self._device)
88
 
@@ -99,14 +105,18 @@ class EmbeddingService:
99
  return result.tolist()
100
 
101
  def generate_image_embedding(self, images: List[Image.Image]) -> List[List[float]]:
102
- if not self._vision_loaded or self._vision_model is None:
103
  raise ValueError("Vision model not loaded")
104
- result: np.ndarray = self._vision_model.encode(
105
- images,
106
- convert_to_numpy=True,
107
- show_progress_bar=False,
108
- )
109
- return result.tolist()
 
 
 
 
110
 
111
  @property
112
  def loaded_dimensions(self) -> List[int]:
 
5
  from typing import Dict, List, Optional
6
 
7
  import numpy as np
8
+ import torch
9
+ import torch.nn.functional as F
10
  from PIL import Image
11
  from sentence_transformers import SentenceTransformer
12
+ from transformers import AutoImageProcessor, AutoModel
13
 
14
  _logger = logging.getLogger(__name__)
15
 
 
36
  self._device = "cpu"
37
 
38
  self._loaded_dimensions: List[int] = []
39
+ self._vision_processor: Optional[AutoImageProcessor] = None
40
+ self._vision_model: Optional[AutoModel] = None
41
  self._vision_loaded = False
42
 
43
  def load_model(self, dimension: int) -> None:
 
81
  _logger.info("Patched vision model config: n_inner float -> int")
82
 
83
  _logger.info("Loading vision embedding model from %s", source)
84
+ self._vision_processor = AutoImageProcessor.from_pretrained(source)
85
+ self._vision_model = AutoModel.from_pretrained(
86
  source,
 
87
  trust_remote_code=True,
88
+ _fast_init=False,
89
  )
90
  self._vision_model.eval()
91
+ self._vision_model.to(self._device)
92
  self._vision_loaded = True
93
  _logger.info("Loaded vision embedding model (device=%s)", self._device)
94
 
 
105
  return result.tolist()
106
 
107
  def generate_image_embedding(self, images: List[Image.Image]) -> List[List[float]]:
108
+ if not self._vision_loaded or self._vision_model is None or self._vision_processor is None:
109
  raise ValueError("Vision model not loaded")
110
+ all_embeddings: List[List[float]] = []
111
+ with torch.no_grad():
112
+ for image in images:
113
+ inputs = self._vision_processor(image, return_tensors="pt")
114
+ inputs = {k: v.to(self._device) for k, v in inputs.items()}
115
+ outputs = self._vision_model(**inputs)
116
+ emb = outputs.last_hidden_state[:, 0]
117
+ emb = F.normalize(emb, p=2, dim=1)
118
+ all_embeddings.append(emb.cpu().numpy().flatten().tolist())
119
+ return all_embeddings
120
 
121
  @property
122
  def loaded_dimensions(self) -> List[int]: