Spaces:
Paused
Paused
sync from GitHub: fix: single model instance, OOM retry with batch reduction, better logging
Browse files- demo/qdrant_utils.py +3 -1
- demo/ui/playground.py +8 -7
- demo/ui/upload.py +6 -9
- visual_rag/embedding/visual_embedder.py +60 -25
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 |
-
|
|
|
|
|
|
|
| 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 |
-
|
| 51 |
-
|
| 52 |
-
|
| 53 |
-
|
| 54 |
-
|
| 55 |
-
|
| 56 |
-
|
|
|
|
| 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 |
-
|
| 273 |
-
embedder =
|
| 274 |
-
if
|
| 275 |
-
|
| 276 |
-
|
| 277 |
-
|
| 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 |
-
|
| 623 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 624 |
|
| 625 |
-
|
| 626 |
-
|
| 627 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 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 |
-
|
| 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
|
| 703 |
-
|
| 704 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 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,
|