This commit is contained in:
IDGallagher
2025-01-11 16:38:37 +08:00
parent 8d4a620f22
commit 6ae2d0512a
+40 -10
View File
@@ -14,15 +14,15 @@ from torchvision import transforms
logger = logging.getLogger(__name__)
# Define URLs for data.bin and embeddings.db on HuggingFace
DATA_BIN_URL = "https://huggingface.co/your-username/your-repo/raw/main/data.bin"
EMBEDDINGS_DB_URL = "https://huggingface.co/your-username/your-repo/raw/main/embeddings.db"
DATA_BIN_URL = "https://huggingface.co/iggy101/MotionVideoSearch/raw/main/index.faiss"
EMBEDDINGS_DB_URL = "https://huggingface.co/iggy101/MotionVideoSearch/raw/main/data.sqlite"
# Directory to store downloaded files
DATA_DIR = Path(__file__).parent / "data"
DATA_DIR.mkdir(parents=True, exist_ok=True)
INDEX_PATH = DATA_DIR / "data.bin"
EMBEDDINGS_DB_PATH = DATA_DIR / "embeddings.db"
INDEX_PATH = DATA_DIR / "index.faiss"
EMBEDDINGS_DB_PATH = DATA_DIR / "data.sqlite"
EMBEDDING_DIM = 768
_dinov2_vitb14_reg = None
@@ -57,22 +57,52 @@ def download_file(url, dest_path):
def ensure_file_exists(file_path, url):
"""
Ensures that the file exists locally; downloads it from the URL if it does not.
Ensures that the file exists locally, and uses ETag to detect if it needs updating.
If ETag indicates the remote file changed, download the new version.
"""
local_etag_path = file_path.with_suffix(file_path.suffix + ".etag")
# Attempt to retrieve the remote ETag
remote_etag = None
try:
head_resp = requests.head(url, allow_redirects=True)
head_resp.raise_for_status()
remote_etag = head_resp.headers.get("ETag")
except requests.exceptions.RequestException as e:
logger.warning(f"Failed to retrieve ETag for {url}: {e}")
# If the file does not exist locally, always download
if not file_path.exists():
logger.info(f"{file_path.name} not found. Downloading from HuggingFace...")
download_file(url, file_path)
if remote_etag:
local_etag_path.write_text(remote_etag)
return
# If we have a remote ETag, compare to local ETag
if remote_etag:
if local_etag_path.exists():
local_etag = local_etag_path.read_text().strip()
# If they match, do nothing
if local_etag == remote_etag:
logger.info(f"{file_path.name} is already up to date (ETag match).")
return
# Otherwise, re-download
logger.info(f"{file_path.name} is outdated. Downloading new version from HuggingFace...")
download_file(url, file_path)
local_etag_path.write_text(remote_etag)
else:
logger.info(f"{file_path.name} already exists.")
# If no remote ETag is available, default to the old logic
logger.info(f"{file_path.name} already exists (no ETag to compare).")
def load_model_and_index():
"""
Loads DINOv2 model and the FAISS index (if not already loaded).
Downloads required files if they are missing.
Downloads required files if they are missing or outdated.
"""
global _dinov2_vitb14_reg, _faiss_index
# Ensure that data.bin and embeddings.db are present
# Ensure that data.bin and embeddings.db are present/updated
ensure_file_exists(INDEX_PATH, DATA_BIN_URL)
ensure_file_exists(EMBEDDINGS_DB_PATH, EMBEDDINGS_DB_URL)
@@ -132,7 +162,7 @@ class IG_MotionVideoSearch:
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
# 2. Convert ComfyUI image (torch tensor) into a PIL Image
c_img = image # Directly use the tensor without indexing
c_img = image
if c_img.ndim == 4:
# Assume shape is [batch_size, C, H, W]
if c_img.size(0) > 1:
@@ -204,4 +234,4 @@ class IG_MotionVideoSearch:
conn.close()
status = "\n".join(results_str)
return (*urls, status)
return (*urls, status)