video / masks.py
recoilme's picture
Upload folder using huggingface_hub
fb1ae50 verified
Raw
History Blame Contribute Delete
3.79 kB
"""SAM3-based mask generation for SCAIL-2, with disk caching.
Masks are deterministic per input video/reference, so they're cached as .pt
files and reused across generation runs (prompt/seed/steps tuning).
"""
import hashlib
import os
import torch
COMFY_DIR = "/home/ubuntu/ComfyUI"
SAM3_CHECKPOINT = "sam3.1_multiplex_fp16.safetensors"
DEFAULT_CACHE_DIR = os.path.join(os.path.dirname(os.path.abspath(__file__)), "mask_cache")
def _cache_key(paths, extra=""):
h = hashlib.sha256()
for p in paths:
h.update(os.path.abspath(p).encode())
h.update(str(os.path.getmtime(p)).encode())
h.update(extra.encode())
return h.hexdigest()[:16]
def _list_ref_files(path):
if os.path.isdir(path):
return sorted(
os.path.join(path, f) for f in os.listdir(path)
if f.lower().endswith((".png", ".jpg", ".jpeg", ".webp"))
)
return [path]
def _load_sam3():
from nodes import CheckpointLoaderSimple
model, clip, _ = CheckpointLoaderSimple().load_checkpoint(SAM3_CHECKPOINT)
return model, clip
def _unload_sam3(model, clip):
del model, clip
import gc
import comfy.model_management as mm
gc.collect()
mm.soft_empty_cache()
torch.cuda.empty_cache()
def _track(model, cond, images):
from comfy_extras.nodes_sam3 import SAM3_VideoTrack
return SAM3_VideoTrack.execute(
images=images, model=model, conditioning=cond,
detection_threshold=0.5, max_objects=1, detect_interval=1,
).args[0]
def generate_masks(pose_video, ref_images, text="person", replacement_mode=False):
"""Run SAM3 tracking → colored SCAIL-2 masks.
Returns (pose_video_mask, reference_image_mask) as IMAGE tensors [B,H,W,3].
"""
from nodes import CLIPTextEncode
from comfy_extras.nodes_scail import SCAIL2ColoredMask
model, clip = _load_sam3()
try:
with torch.inference_mode(): # no autograd graph — activations freed per frame
cond = CLIPTextEncode().encode(clip=clip, text=text)[0]
print(f" SAM3: tracking pose video ({pose_video.shape[0]} frames)...")
pose_tracks = _track(model, cond, pose_video)
print(f" SAM3: tracking reference ({ref_images.shape[0]} image(s))...")
ref_tracks = _track(model, cond, ref_images)
pose_mask, ref_mask = SCAIL2ColoredMask.execute(
driving_track_data=pose_tracks,
ref_track_data=ref_tracks,
object_indices="",
sort_by="left_to_right",
replacement_mode=replacement_mode,
).args
finally:
_unload_sam3(model, clip)
return pose_mask, ref_mask
def get_masks(pose_path, ref_path, pose_video, ref_images, width, height,
replacement_mode=False, cache_dir=DEFAULT_CACHE_DIR):
"""Load masks from cache or generate+cache them.
Cache key covers input files (path+mtime) and output resolution.
"""
os.makedirs(cache_dir, exist_ok=True)
ref_files = _list_ref_files(ref_path)
key = _cache_key([pose_path] + ref_files, f"{width}x{height}:{replacement_mode}")
pose_cache = os.path.join(cache_dir, f"{key}.pose_mask.pt")
ref_cache = os.path.join(cache_dir, f"{key}.ref_mask.pt")
if os.path.exists(pose_cache) and os.path.exists(ref_cache):
print(f" Masks: loaded from cache ({key})")
return torch.load(pose_cache, weights_only=True), torch.load(ref_cache, weights_only=True)
pose_mask, ref_mask = generate_masks(pose_video, ref_images,
replacement_mode=replacement_mode)
torch.save(pose_mask, pose_cache)
torch.save(ref_mask, ref_cache)
print(f" Masks: generated and cached ({key})")
return pose_mask, ref_mask