Yeroyan commited on
Commit
518c15b
·
1 Parent(s): 796678a

sync from GitHub: fix: single model instance, OOM retry with batch reduction, better logging

Browse files
demo/qdrant_utils.py CHANGED
@@ -43,7 +43,9 @@ def init_embedder(model_name: str):
43
  try:
44
  from visual_rag import VisualEmbedder
45
 
46
- return VisualEmbedder(model_name=model_name), None
 
 
47
  except Exception as e:
48
  return None, f"{e}\n\n{traceback.format_exc()}"
49
 
 
43
  try:
44
  from visual_rag import VisualEmbedder
45
 
46
+ embedder = VisualEmbedder(model_name=model_name)
47
+ embedder._load_model()
48
+ return embedder, None
49
  except Exception as e:
50
  return None, f"{e}\n\n{traceback.format_exc()}"
51
 
demo/ui/playground.py CHANGED
@@ -47,13 +47,14 @@ def render_playground_tab():
47
  if not st.session_state.get("model_loaded"):
48
  with st.spinner(f"Loading {model_short}..."):
49
  try:
50
- _ = MultiVectorRetriever(
51
- collection_name=active_collection, model_name=model_name,
52
- qdrant_url=url, qdrant_api_key=api_key,
53
- )
54
- st.session_state["model_loaded"] = True
55
- st.session_state["loaded_model_key"] = cache_key
56
- st.session_state["loaded_model_name"] = model_name
 
57
  except Exception:
58
  st.warning(f"Failed: {model_short}")
59
 
 
47
  if not st.session_state.get("model_loaded"):
48
  with st.spinner(f"Loading {model_short}..."):
49
  try:
50
+ from demo.qdrant_utils import init_embedder
51
+ embedder, err = init_embedder(model_name)
52
+ if err:
53
+ st.warning(f"Failed: {model_short}")
54
+ else:
55
+ st.session_state["model_loaded"] = True
56
+ st.session_state["loaded_model_key"] = cache_key
57
+ st.session_state["loaded_model_name"] = model_name
58
  except Exception:
59
  st.warning(f"Failed: {model_short}")
60
 
demo/ui/upload.py CHANGED
@@ -269,15 +269,12 @@ def process_pdfs(uploaded_files, config):
269
  model_status.info(f"Loading `{model_short}`...")
270
 
271
  output_dtype = np.float16 if vector_dtype == "float16" else np.float32
272
- embedder_key = f"{model_name}::{vector_dtype}"
273
- embedder = None
274
- if st.session_state.get("upload_embedder_key") == embedder_key:
275
- embedder = st.session_state.get("upload_embedder")
276
- if embedder is None:
277
- embedder = VisualEmbedder(model_name=model_name, output_dtype=output_dtype)
278
- embedder._load_model()
279
- st.session_state["upload_embedder_key"] = embedder_key
280
- st.session_state["upload_embedder"] = embedder
281
  model_status.success(f"✅ Model `{model_short}` loaded ({vector_dtype})")
282
 
283
  with phase2:
 
269
  model_status.info(f"Loading `{model_short}`...")
270
 
271
  output_dtype = np.float16 if vector_dtype == "float16" else np.float32
272
+ from demo.qdrant_utils import init_embedder
273
+ embedder, err = init_embedder(model_name)
274
+ if err:
275
+ model_status.error(f"❌ Failed to load model: {err[:100]}")
276
+ return
277
+ embedder.output_dtype = output_dtype
 
 
 
278
  model_status.success(f"✅ Model `{model_short}` loaded ({vector_dtype})")
279
 
280
  with phase2:
visual_rag/embedding/visual_embedder.py CHANGED
@@ -618,13 +618,55 @@ class VisualEmbedder:
618
 
619
  for i in iterator:
620
  batch = images[i : i + batch_size]
 
621
 
622
- with torch.no_grad():
623
- processed = self.processor.process_images(batch).to(self.model.device)
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
624
 
625
- # Extract token info before model forward
626
- if return_token_info:
627
- input_ids = processed["input_ids"]
 
 
 
 
 
 
 
 
 
628
  batch_n_rows = processed.get("n_rows")
629
  batch_n_cols = processed.get("n_cols")
630
  # Qwen2/2.5-VL style grid information (T, H, W)
@@ -681,27 +723,20 @@ class VisualEmbedder:
681
  }
682
  )
683
 
684
- # Generate embeddings
685
- batch_embeddings = self.model(**processed)
686
-
687
- # Extract per-image embeddings
688
- if isinstance(batch_embeddings, torch.Tensor) and batch_embeddings.dim() == 3:
689
- for j in range(batch_embeddings.shape[0]):
690
- embeddings.append(batch_embeddings[j].cpu())
691
- else:
692
- embeddings.extend([e.cpu() for e in batch_embeddings])
693
-
694
- # Memory cleanup
695
- del processed, batch_embeddings
696
- gc.collect()
697
- if torch.cuda.is_available():
698
- torch.cuda.empty_cache()
699
- elif torch.backends.mps.is_available():
700
- torch.mps.empty_cache()
701
 
702
- if return_token_info:
703
- return embeddings, token_infos
704
- return embeddings
 
 
 
 
 
 
 
 
 
705
 
706
  def extract_visual_embedding(
707
  self,
 
618
 
619
  for i in iterator:
620
  batch = images[i : i + batch_size]
621
+ current_batch_size = len(batch)
622
 
623
+ while current_batch_size > 0:
624
+ try:
625
+ sub_batch = batch[:current_batch_size]
626
+ self._embed_image_batch(
627
+ sub_batch, embeddings, token_infos, return_token_info
628
+ )
629
+ if current_batch_size < len(batch):
630
+ remaining = batch[current_batch_size:]
631
+ for single in remaining:
632
+ self._embed_image_batch(
633
+ [single], embeddings, token_infos, return_token_info
634
+ )
635
+ break
636
+ except (torch.cuda.OutOfMemoryError, RuntimeError) as e:
637
+ if "out of memory" not in str(e).lower() and "CUDA" not in str(e):
638
+ raise
639
+ gc.collect()
640
+ if torch.cuda.is_available():
641
+ torch.cuda.empty_cache()
642
+ new_size = max(1, current_batch_size // 2)
643
+ logger.warning(
644
+ f"⚠️ CUDA OOM with batch_size={current_batch_size}, "
645
+ f"retrying with batch_size={new_size}"
646
+ )
647
+ print(
648
+ f"[OOM] Reducing batch size from {current_batch_size} to {new_size}"
649
+ )
650
+ if new_size == current_batch_size:
651
+ raise
652
+ current_batch_size = new_size
653
+
654
+ if return_token_info:
655
+ return embeddings, token_infos
656
+ return embeddings
657
 
658
+ def _embed_image_batch(
659
+ self,
660
+ batch: List[Image.Image],
661
+ embeddings: list,
662
+ token_infos: Optional[list],
663
+ return_token_info: bool,
664
+ ):
665
+ with torch.no_grad():
666
+ processed = self.processor.process_images(batch).to(self.model.device)
667
+
668
+ if return_token_info:
669
+ input_ids = processed["input_ids"]
670
  batch_n_rows = processed.get("n_rows")
671
  batch_n_cols = processed.get("n_cols")
672
  # Qwen2/2.5-VL style grid information (T, H, W)
 
723
  }
724
  )
725
 
726
+ batch_embeddings = self.model(**processed)
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
727
 
728
+ if isinstance(batch_embeddings, torch.Tensor) and batch_embeddings.dim() == 3:
729
+ for j in range(batch_embeddings.shape[0]):
730
+ embeddings.append(batch_embeddings[j].cpu())
731
+ else:
732
+ embeddings.extend([e.cpu() for e in batch_embeddings])
733
+
734
+ del processed, batch_embeddings
735
+ gc.collect()
736
+ if torch.cuda.is_available():
737
+ torch.cuda.empty_cache()
738
+ elif torch.backends.mps.is_available():
739
+ torch.mps.empty_cache()
740
 
741
  def extract_visual_embedding(
742
  self,