SCAIL-2 / app.py
loveseries's picture
Replace with lightweight SAM3 mask-only service (no checkpoint) (#2)
800ec24
Raw
History Blame Contribute Delete
4.78 kB
# ============================================================================
# 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()