File size: 3,786 Bytes
fb1ae50
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
"""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