Spaces:
Running on Zero
Running on Zero
| # ============================================================================ | |
| # 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 | |
| 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() | |