Node updates
This commit is contained in:
@@ -9,8 +9,10 @@ from .nodes import *
|
||||
|
||||
NODE_CLASS_MAPPINGS = {
|
||||
"IG Motion Video Search": IG_MotionVideoSearch,
|
||||
"IG Motion Video Frame": IG_MotionVideoFrame,
|
||||
}
|
||||
|
||||
NODE_DISPLAY_NAME_MAPPINGS = {
|
||||
"IG Motion Video Search": "🐂 IG Motion Video Search",
|
||||
"IG Motion Video Frame": "🖼️ IG Motion Video Frame"
|
||||
}
|
||||
+12
-6
@@ -3,6 +3,10 @@
|
||||
import torch
|
||||
import numpy as np
|
||||
from einops import rearrange
|
||||
import sys
|
||||
import os
|
||||
|
||||
sys.path.insert(0, os.path.abspath(os.path.dirname(__file__)))
|
||||
from dot.models import create_model
|
||||
|
||||
##########################################################################
|
||||
@@ -137,6 +141,8 @@ def flow_to_color(flow_uv, clip_flow=None, convert_to_bgr=False):
|
||||
# DOT args and helper class
|
||||
##########################################################################
|
||||
|
||||
# dot_functions.py
|
||||
|
||||
dot_args = {
|
||||
"model": "dot",
|
||||
"fit_to": 336,
|
||||
@@ -147,13 +153,13 @@ dot_args = {
|
||||
"is_train": False,
|
||||
"worker_idx": 0,
|
||||
"num_workers": 2,
|
||||
"estimator_config": "configs/raft_patch_8.json",
|
||||
"estimator_path": "checkpoints/cvo_raft_patch_8.pth",
|
||||
"estimator_config": os.path.join(os.path.dirname(__file__), "configs/raft_patch_8.json"),
|
||||
"estimator_path": os.path.join(os.path.dirname(__file__), "checkpoints/cvo_raft_patch_8.pth"),
|
||||
"flow_mode": "direct",
|
||||
"refiner_config": "configs/raft_patch_4_alpha.json",
|
||||
"refiner_path": "checkpoints/movi_f_raft_patch_4_alpha.pth",
|
||||
"tracker_config": "configs/cotracker_patch_4_wind_8.json",
|
||||
"tracker_path": "checkpoints/movi_f_cotracker_patch_4_wind_8.pth",
|
||||
"refiner_config": os.path.join(os.path.dirname(__file__), "configs/raft_patch_4_alpha.json"),
|
||||
"refiner_path": os.path.join(os.path.dirname(__file__), "checkpoints/movi_f_raft_patch_4_alpha.pth"),
|
||||
"tracker_config": os.path.join(os.path.dirname(__file__), "configs/cotracker_patch_4_wind_8.json"),
|
||||
"tracker_path": os.path.join(os.path.dirname(__file__), "checkpoints/movi_f_cotracker_patch_4_wind_8.pth"),
|
||||
"sample_mode": "all",
|
||||
"interpolation_version": "torch3d",
|
||||
"inference_mode": "tracks_from_first_to_every_other_frame",
|
||||
|
||||
+4
-2
@@ -37,7 +37,9 @@ def download_checkpoints():
|
||||
Checks for each checkpoint file in the 'checkpoints' directory.
|
||||
If it does not exist locally, downloads it from the provided URL.
|
||||
"""
|
||||
os.makedirs("checkpoints", exist_ok=True)
|
||||
filepath = os.path.dirname(__file__)
|
||||
checkpoint_dir = os.path.join(filepath, "checkpoints")
|
||||
os.makedirs(checkpoint_dir, exist_ok=True)
|
||||
|
||||
# List of (local_filename, url). Adjust as needed for each checkpoint.
|
||||
files_to_download = [
|
||||
@@ -51,7 +53,7 @@ def download_checkpoints():
|
||||
]
|
||||
|
||||
for filename, url in files_to_download:
|
||||
local_path = os.path.join("checkpoints", filename)
|
||||
local_path = os.path.join(checkpoint_dir, filename)
|
||||
|
||||
# Check if file already exists
|
||||
if not os.path.isfile(local_path):
|
||||
|
||||
@@ -8,14 +8,21 @@ import requests
|
||||
from tqdm import tqdm
|
||||
from PIL import Image
|
||||
from pathlib import Path
|
||||
from .database import get_connection # Ensure accessible to ComfyUI
|
||||
import torch.nn.functional as F
|
||||
|
||||
# Make sure this import can find the database.py in your package
|
||||
from .database import get_connection
|
||||
from torchvision import transforms
|
||||
|
||||
# Import your get_dot_frame function
|
||||
from .dot_functions import get_dot_frame
|
||||
from .functions import download_checkpoints
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
# Define URLs for data.bin and embeddings.db on HuggingFace
|
||||
DATA_BIN_URL = "https://huggingface.co/iggy101/MotionVideoSearch/raw/main/index.faiss"
|
||||
EMBEDDINGS_DB_URL = "https://huggingface.co/iggy101/MotionVideoSearch/raw/main/data.sqlite"
|
||||
DATA_BIN_URL = "https://huggingface.co/iggy101/MotionVideoSearch/resolve/main/index.faiss"
|
||||
EMBEDDINGS_DB_URL = "https://huggingface.co/iggy101/MotionVideoSearch/resolve/main/data.sqlite"
|
||||
|
||||
# Directory to store downloaded files
|
||||
DATA_DIR = Path(__file__).parent / "data"
|
||||
@@ -32,6 +39,7 @@ _transform = transforms.Compose([
|
||||
transforms.ToTensor(),
|
||||
])
|
||||
|
||||
|
||||
def download_file(url, dest_path):
|
||||
"""
|
||||
Downloads a file from the specified URL to the destination path with a progress bar.
|
||||
@@ -55,6 +63,7 @@ def download_file(url, dest_path):
|
||||
logger.error(f"Failed to download {url}: {e}")
|
||||
raise
|
||||
|
||||
|
||||
def ensure_file_exists(file_path, url):
|
||||
"""
|
||||
Ensures that the file exists locally, and uses ETag to detect if it needs updating.
|
||||
@@ -95,6 +104,7 @@ def ensure_file_exists(file_path, url):
|
||||
# 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).
|
||||
@@ -123,10 +133,11 @@ def load_model_and_index():
|
||||
|
||||
return _dinov2_vitb14_reg, _faiss_index
|
||||
|
||||
|
||||
class IG_MotionVideoSearch:
|
||||
"""
|
||||
A ComfyUI node that accepts a ComfyUI image and
|
||||
returns the top 5 search results from the FAISS index.
|
||||
returns 5 ranked search results from the FAISS index based on a given starting rank.
|
||||
"""
|
||||
|
||||
@classmethod
|
||||
@@ -134,7 +145,7 @@ class IG_MotionVideoSearch:
|
||||
return {
|
||||
"required": {
|
||||
"image": ("IMAGE",), # ComfyUI "IMAGE" type
|
||||
"top_k": ("INT", {"default": 5, "min": 1, "max": 20, "step": 1}),
|
||||
"starting_rank": ("INT", {"default": 1, "min": 1, "max": 9999999, "step": 1}),
|
||||
}
|
||||
}
|
||||
|
||||
@@ -144,18 +155,19 @@ class IG_MotionVideoSearch:
|
||||
|
||||
CATEGORY = "Motion Video DB" # Appears in ComfyUI under this category in the node menu
|
||||
|
||||
def search(self, image, top_k):
|
||||
def search(self, image, starting_rank):
|
||||
"""
|
||||
Perform the search using the loaded FAISS index and DINOv2 model.
|
||||
|
||||
:param image: A torch.Tensor, shape [batch_size, C, H, W]
|
||||
:param top_k: Number of top results to retrieve
|
||||
:param starting_rank: The starting rank for the search results (defaults to 1).
|
||||
The node will return this rank plus the next 4 lower-ranked results.
|
||||
:return: 5 separate URLs for the search results and a status string with scores
|
||||
"""
|
||||
# Log input details for debugging
|
||||
logger.debug(f"Image type: {type(image)}")
|
||||
logger.debug(f"Image shape: {image.shape}")
|
||||
logger.debug(f"Image dtype: {image.dtype}")
|
||||
logger.debug(f"Starting rank: {starting_rank}")
|
||||
|
||||
# 1. Load model and index if needed
|
||||
model, index = load_model_and_index()
|
||||
@@ -169,7 +181,7 @@ class IG_MotionVideoSearch:
|
||||
logger.warning("Received batch size > 1. Only processing the first image in the batch.")
|
||||
c_img = c_img[0]
|
||||
else:
|
||||
c_img = c_img.squeeze(0) # Remove the batch dimension
|
||||
c_img = c_img.squeeze(0)
|
||||
elif c_img.ndim != 3:
|
||||
raise ValueError(f"Expected image tensor to have 3 or 4 dimensions, got {c_img.ndim}")
|
||||
|
||||
@@ -181,7 +193,7 @@ class IG_MotionVideoSearch:
|
||||
# Convert to PIL image
|
||||
pil_img = Image.fromarray(np_img, mode='RGB')
|
||||
|
||||
# 3. Apply the same resizing logic as your main.py does (multiple of 14, etc.)
|
||||
# 3. Apply the same resizing logic if desired
|
||||
with torch.no_grad():
|
||||
tensor_img = _transform(pil_img).unsqueeze(0).to(device) # shape [1, C, H, W]
|
||||
_, _, h, w = tensor_img.shape
|
||||
@@ -194,23 +206,33 @@ class IG_MotionVideoSearch:
|
||||
# 4. Get the embedding
|
||||
embedding = model(tensor_img).cpu().numpy().astype("float32")
|
||||
|
||||
# 5. Search in FAISS
|
||||
# We want top_k = starting_rank + 4
|
||||
top_k = starting_rank + 4
|
||||
distances, ids = index.search(embedding, top_k)
|
||||
|
||||
# Handle edge cases
|
||||
if ids.size == 0 or (ids.size == 1 and ids[0][0] == -1):
|
||||
return ("No embeddings found in the FAISS index.", "", "", "", "", "No scores available.")
|
||||
return (
|
||||
"No embeddings found in the FAISS index.",
|
||||
"",
|
||||
"",
|
||||
"",
|
||||
"",
|
||||
"No scores available."
|
||||
)
|
||||
|
||||
# Slice out the 5 results we actually want
|
||||
selected_distances = distances[0][starting_rank - 1 : starting_rank - 1 + 5]
|
||||
selected_ids = ids[0][starting_rank - 1 : starting_rank - 1 + 5]
|
||||
|
||||
# 6. Retrieve metadata from SQLite
|
||||
conn = get_connection()
|
||||
cursor = conn.cursor()
|
||||
|
||||
urls = [""] * 5 # Initialize list of 5 URL strings
|
||||
urls = [""] * 5
|
||||
results_str = []
|
||||
|
||||
for rank, (dist, uid) in enumerate(zip(distances[0], ids[0]), start=1):
|
||||
if rank > 5:
|
||||
break
|
||||
for offset, (dist, uid) in enumerate(zip(selected_distances, selected_ids)):
|
||||
rank = starting_rank + offset
|
||||
if uid == -1:
|
||||
results_str.append(f"{rank}. [No valid ID] - Distance: {dist:.4f}")
|
||||
continue
|
||||
@@ -226,12 +248,108 @@ class IG_MotionVideoSearch:
|
||||
)
|
||||
row = cursor.fetchone()
|
||||
if row:
|
||||
urls[rank - 1] = row[0]
|
||||
urls[offset] = row[0]
|
||||
results_str.append(f"{rank}. Distance: {dist:.4f}")
|
||||
else:
|
||||
results_str.append(f"{rank}. [Missing DB row for ID {uid}], Distance: {dist:.4f}")
|
||||
|
||||
conn.close()
|
||||
|
||||
status = "\n".join(results_str)
|
||||
|
||||
return (*urls, status)
|
||||
|
||||
|
||||
###############################################################################
|
||||
# New Node: IG_MotionVideoDotFrame
|
||||
###############################################################################
|
||||
class IG_MotionVideoFrame:
|
||||
"""
|
||||
A ComfyUI node that takes a batch of frames from a video, ensures we have at least
|
||||
24 frames, trims (cuts off) the video to exactly 24 frames, then calls get_dot_frame
|
||||
to produce a colorized motion image. Finally, it returns that as a ComfyUI image.
|
||||
"""
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
"required": {
|
||||
"video_frames": ("IMAGE",),
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("IMAGE",)
|
||||
RETURN_NAMES = ("colored_motion_image",)
|
||||
FUNCTION = "apply"
|
||||
CATEGORY = "Motion Video DB"
|
||||
|
||||
def apply(self, video_frames):
|
||||
"""
|
||||
:param video_frames: A 4D torch.Tensor of shape [N, C, H, W], with N >= 24
|
||||
:return: A single colorized motion image (torch.Tensor of shape [1, C, H, W])
|
||||
"""
|
||||
|
||||
if not isinstance(video_frames, torch.Tensor):
|
||||
raise TypeError("video_frames must be a torch.Tensor.")
|
||||
|
||||
if video_frames.ndim != 4:
|
||||
raise ValueError(
|
||||
f"Expected a 4D tensor [N, C, H, W], but got shape {video_frames.shape}."
|
||||
)
|
||||
|
||||
# AI is convinced that Comfy images are B, C, H, W but they're actually B, H, W, C
|
||||
video_frames = video_frames.permute(0, 3, 1, 2)
|
||||
|
||||
num_frames, channels, height, width = video_frames.shape
|
||||
if num_frames < 24:
|
||||
raise ValueError(
|
||||
f"Video must have at least 24 frames, but got {num_frames}."
|
||||
)
|
||||
|
||||
# 1) Trim the video to 24 frames
|
||||
video_frames = video_frames[:24] # shape: [24, C, H, W]
|
||||
|
||||
# 2) Scale so the shorter side is 'fit_to':
|
||||
fit_to = 336
|
||||
_, _, H, W = video_frames.shape
|
||||
|
||||
# If width < height, we set width to 336 and scale height accordingly.
|
||||
# Otherwise, we set height to 336 and scale width accordingly.
|
||||
if W < H:
|
||||
new_W = fit_to
|
||||
new_H = int(round(H * fit_to / W))
|
||||
else:
|
||||
new_H = fit_to
|
||||
new_W = int(round(W * fit_to / H))
|
||||
|
||||
# Interpolate all frames to this new size
|
||||
video_frames = F.interpolate(
|
||||
video_frames,
|
||||
size=(new_H, new_W),
|
||||
mode='bilinear',
|
||||
align_corners=False
|
||||
)
|
||||
|
||||
# 3) Crop so the height and width are multiples of 8
|
||||
new_H_aligned = (new_H // 8) * 8
|
||||
new_W_aligned = (new_W // 8) * 8
|
||||
|
||||
h_start = (new_H - new_H_aligned) // 2
|
||||
w_start = (new_W - new_W_aligned) // 2
|
||||
|
||||
video_frames = video_frames[
|
||||
:,
|
||||
:,
|
||||
h_start : h_start + new_H_aligned,
|
||||
w_start : w_start + new_W_aligned
|
||||
]
|
||||
|
||||
# get_dot_frame typically expects the shape [N, C, H, W].
|
||||
# This function returns a single frame of shape [C, H, W].
|
||||
download_checkpoints()
|
||||
frame = get_dot_frame(video_frames)
|
||||
|
||||
# Convert to ComfyUI's "IMAGE" format: [1, H, W, C]
|
||||
frame = frame.unsqueeze(0) # shape = [1, C, H, W]
|
||||
frame = frame.permute(0, 2, 3, 1)
|
||||
|
||||
return (frame,)
|
||||
|
||||
@@ -1 +1,2 @@
|
||||
faiss-cpu
|
||||
einshape
|
||||
Reference in New Issue
Block a user