Spaces:
Running on Zero
Running on Zero
File size: 4,780 Bytes
800ec24 93567bd 800ec24 93567bd 800ec24 93567bd 800ec24 3004afe 800ec24 974817c 800ec24 974817c 800ec24 974817c 800ec24 93567bd 800ec24 7b2c15e 93567bd 800ec24 7b2c15e 93567bd 800ec24 93567bd 974817c 800ec24 974817c 800ec24 | 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 110 111 112 113 114 115 116 117 | # ============================================================================
# SCAIL SAM3 Mask Service (lightweight)
# รับ รูป + วิดีโอ → ทำ SAM3 mask บน GPU (SCAIL-Pose e2e) → คืน
# ref_mask.jpg + rendered_mask_v2.mp4 + rendered_v2.mp4 (driving copy, res ตรงกับ mask)
# ไม่มี checkpoint/generation → Space เบา start ได้ในฟรี ephemeral
# endpoint: /make_masks(image, video, max_persons) -> (ref_mask, driving_mask, rendered, status)
# ============================================================================
import os
import sys
import shutil
import tempfile
import logging
import traceback
import gradio as gr
import spaces
from huggingface_hub import hf_hub_download
logging.basicConfig(level=logging.INFO)
_SCAIL_POSE_DIR = os.path.join(os.path.dirname(os.path.abspath(__file__)), "SCAIL-Pose")
if _SCAIL_POSE_DIR not in sys.path:
sys.path.insert(0, _SCAIL_POSE_DIR)
SAM3_MODEL_PATH = os.getenv(
"SAM3_MODEL", os.path.join(_SCAIL_POSE_DIR, "pretrained_weights", "sam3.pt")
)
_predictor = None
def _ensure_sam3():
"""โหลด sam3.pt จาก facebook/sam3 (gated) — ต้องมี HF_TOKEN ที่ได้สิทธิ์."""
if os.path.exists(SAM3_MODEL_PATH):
return
os.makedirs(os.path.dirname(SAM3_MODEL_PATH), exist_ok=True)
tok = os.getenv("HF_TOKEN") or os.getenv("HUGGING_FACE_HUB_TOKEN")
f = hf_hub_download("facebook/sam3", "sam3.pt", token=tok)
try:
os.symlink(f, SAM3_MODEL_PATH)
except OSError:
shutil.copyfile(f, SAM3_MODEL_PATH)
def _get_predictor():
global _predictor
if _predictor is None:
from ultralytics.models.sam import SAM3VideoSemanticPredictor
overrides = dict(
conf=0.25, task="segment", mode="predict", imgsz=640,
model=SAM3_MODEL_PATH, half=True, save=False, verbose=False,
)
_predictor = SAM3VideoSemanticPredictor(overrides=overrides, new_det_thresh=1.0)
return _predictor
@spaces.GPU(duration=120)
def make_masks(image, video, max_persons=1):
"""รูป + วิดีโอ → SAM3 mask (e2e). คืน (ref_mask.jpg, rendered_mask_v2.mp4, rendered_v2.mp4, status)."""
try:
if image is None or video is None:
return None, None, None, "Missing image or video"
_ensure_sam3()
from NLFPoseExtract.process_animation_aio import process_one
subdir = tempfile.mkdtemp(prefix="mask_")
shutil.copyfile(image, os.path.join(subdir, "ref.png"))
shutil.copyfile(video, os.path.join(subdir, "driving.mp4"))
process_one(
subdir, "driving.mp4", e2e_mode=True, crop_kind=None,
max_persons=int(max_persons), text=["human", "character"],
model_nlf=None, detector=None, predictor=_get_predictor(), image_predictor=None,
)
ref_mask = os.path.join(subdir, "ref_mask.jpg")
drv_mask = os.path.join(subdir, "rendered_mask_v2.mp4")
rendered = os.path.join(subdir, "rendered_v2.mp4")
if not (os.path.exists(ref_mask) and os.path.exists(drv_mask)):
return None, None, None, "Mask generation produced no output"
return ref_mask, drv_mask, rendered, "OK"
except Exception:
logging.exception("make_masks failed")
return None, None, None, traceback.format_exc()
with gr.Blocks(title="SCAIL SAM3 Mask Service") as demo:
gr.Markdown(
"## SCAIL-2 SAM3 Masking Service\n"
"Send a **reference image** + **driving video** → returns colored SAM3 masks "
"(`ref_mask.jpg`, `rendered_mask_v2.mp4`) + the driving copy (`rendered_v2.mp4`, same res as mask). "
"Feed these to `fffiloni/SCAIL-2` `/generate_from_uploads`."
)
with gr.Row():
with gr.Column():
m_image = gr.Image(type="filepath", label="Reference image")
m_video = gr.Video(label="Driving video")
m_persons = gr.Number(value=1, precision=0, label="Max persons")
m_run = gr.Button("Make masks", variant="primary")
m_status = gr.Textbox(label="Status")
with gr.Column():
m_refmask = gr.Image(type="filepath", label="ref_mask")
m_drvmask = gr.Video(label="driving mask (rendered_mask_v2)")
m_rendered = gr.Video(label="rendered (driving copy)")
m_run.click(
make_masks,
inputs=[m_image, m_video, m_persons],
outputs=[m_refmask, m_drvmask, m_rendered, m_status],
api_name="make_masks",
)
if __name__ == "__main__":
try:
_ensure_sam3()
except Exception as _e:
logging.warning("SAM3 weights not ready at startup: %s", _e)
demo.queue(max_size=8).launch()
|