Spaces:
Running on Zero
Running on Zero
| import json | |
| import logging | |
| import os | |
| import re | |
| import shutil | |
| import subprocess | |
| import sys | |
| import tempfile | |
| import traceback | |
| import uuid | |
| import zipfile | |
| from pathlib import Path | |
| from dataclasses import dataclass | |
| import gradio as gr | |
| import spaces | |
| from huggingface_hub import hf_hub_download | |
| from PIL import Image | |
| logging.basicConfig(level=logging.INFO, format="[%(asctime)s] %(levelname)s: %(message)s") | |
| ROOT = Path(__file__).resolve().parent | |
| JOB_ROOT = Path(os.getenv("SCAIL_POSE_JOB_ROOT", str(Path(tempfile.gettempdir()) / "scail_pose_jobs"))) | |
| JOB_ROOT.mkdir(parents=True, exist_ok=True) | |
| SAM3_REPO_ID = os.getenv("SCAIL_POSE_SAM3_REPO_ID", "facebook/sam3") | |
| SAM3_FILENAME = os.getenv("SCAIL_POSE_SAM3_FILENAME", "sam3.pt") | |
| WEIGHTS_DIR = Path(os.getenv("SCAIL_POSE_WEIGHTS_DIR", str(ROOT / "pretrained_weights"))) | |
| SAM3_MODEL_PATH = Path(os.getenv("SCAIL_POSE_SAM3_MODEL", str(WEIGHTS_DIR / SAM3_FILENAME))) | |
| AUTO_DOWNLOAD_SAM3 = os.getenv("SCAIL_POSE_AUTO_DOWNLOAD_SAM3", "1") == "1" | |
| GPU_SIZE = os.getenv("SCAIL_POSE_ZEROGPU_SIZE", "xlarge") | |
| GPU_DURATION = int(os.getenv("SCAIL_POSE_GPU_DURATION", "90")) | |
| IMAGE_EXTS = {".png", ".jpg", ".jpeg", ".webp"} | |
| VIDEO_EXTS = {".mp4", ".mov", ".webm", ".mkv"} | |
| SCAIL_COLORS = [ | |
| (0, 0, 255), | |
| (255, 0, 0), | |
| (0, 255, 0), | |
| (255, 255, 0), | |
| (255, 0, 255), | |
| (0, 255, 255), | |
| ] | |
| _SAM3_VIDEO_PREDICTOR = None | |
| _SAM3_IMAGE_PREDICTOR = None | |
| _SAM3_PREDICTOR_KEY = None | |
| class CharacterRef: | |
| character: str | |
| view: str | |
| path: Path | |
| def _repo_status() -> str: | |
| required = [ | |
| "NLFPoseExtract/process_animation_aio.py", | |
| "NLFPoseExtract/process_replacement.py", | |
| "NLFPoseExtract/v2_helper.py", | |
| "TrackSam3/track.py", | |
| ] | |
| missing = [rel for rel in required if not (ROOT / rel).exists()] | |
| if missing: | |
| return ( | |
| "SCAIL-Pose repo layout was not found. Put this app.py at the root of " | |
| "the SCAIL-Pose checkout.\n\nMissing:\n" + "\n".join(f"- {rel}" for rel in missing) | |
| ) | |
| weight_status = ( | |
| f"SAM3 weights found: {SAM3_MODEL_PATH}" | |
| if SAM3_MODEL_PATH.exists() | |
| else f"SAM3 weights not found yet: {SAM3_MODEL_PATH}" | |
| ) | |
| return ( | |
| "Ready. SCAIL-Pose repo layout detected.\n\n" | |
| f"Job root: {JOB_ROOT}\n" | |
| f"{weight_status}\n\n" | |
| "This Space exports SCAIL-2-compatible input packs. Outputs are temporary." | |
| ) | |
| def _require_repo_layout(): | |
| missing = [] | |
| for rel in ( | |
| "NLFPoseExtract/process_animation_aio.py", | |
| "NLFPoseExtract/process_replacement.py", | |
| "NLFPoseExtract/v2_helper.py", | |
| "TrackSam3/track.py", | |
| ): | |
| if not (ROOT / rel).exists(): | |
| missing.append(rel) | |
| if missing: | |
| raise RuntimeError( | |
| "This app.py must live at the root of the SCAIL-Pose repository. " | |
| f"Missing: {', '.join(missing)}" | |
| ) | |
| def _ensure_sam3_weights() -> Path: | |
| if SAM3_MODEL_PATH.exists(): | |
| return SAM3_MODEL_PATH | |
| if not AUTO_DOWNLOAD_SAM3: | |
| raise RuntimeError( | |
| f"SAM3 weights were not found at {SAM3_MODEL_PATH}. " | |
| "Set SCAIL_POSE_SAM3_MODEL or enable SCAIL_POSE_AUTO_DOWNLOAD_SAM3=1." | |
| ) | |
| WEIGHTS_DIR.mkdir(parents=True, exist_ok=True) | |
| logging.info("Downloading SAM3 weights from %s/%s", SAM3_REPO_ID, SAM3_FILENAME) | |
| downloaded = hf_hub_download( | |
| repo_id=SAM3_REPO_ID, | |
| filename=SAM3_FILENAME, | |
| local_dir=str(WEIGHTS_DIR), | |
| token=os.getenv("HF_TOKEN") or None, | |
| ) | |
| return Path(downloaded) | |
| def _new_job_dir(mode: str) -> Path: | |
| job_dir = JOB_ROOT / f"{mode}_{uuid.uuid4().hex}" | |
| job_dir.mkdir(parents=True, exist_ok=False) | |
| return job_dir | |
| def _as_path(upload, label: str) -> Path: | |
| if upload is None: | |
| raise RuntimeError(f"Missing {label}.") | |
| if isinstance(upload, dict): | |
| upload = upload.get("path") or upload.get("name") | |
| elif hasattr(upload, "path"): | |
| upload = upload.path | |
| elif hasattr(upload, "name") and not isinstance(upload, (str, os.PathLike)): | |
| upload = upload.name | |
| path = Path(upload) | |
| if not path.exists(): | |
| raise RuntimeError(f"{label} does not exist: {path}") | |
| return path | |
| def _save_reference_image(upload, dest: Path) -> Path: | |
| source = _as_path(upload, "reference image") | |
| try: | |
| image = Image.open(source).convert("RGB") | |
| image.save(dest) | |
| except Exception as exc: | |
| raise RuntimeError(f"Could not read reference image: {source}") from exc | |
| return dest | |
| def _copy_video(upload, dest: Path, label: str) -> Path: | |
| source = _as_path(upload, label) | |
| if source.suffix.lower() not in VIDEO_EXTS: | |
| raise RuntimeError(f"{label} should be a video file. Got: {source.name}") | |
| shutil.copy2(source, dest) | |
| return dest | |
| def _copy_reference_to_png(source: Path, dest: Path) -> Path: | |
| dest.parent.mkdir(parents=True, exist_ok=True) | |
| try: | |
| image = Image.open(source).convert("RGB") | |
| image.save(dest) | |
| except Exception as exc: | |
| raise RuntimeError(f"Could not read reference image: {source}") from exc | |
| return dest | |
| def _safe_name(value: str) -> str: | |
| value = value.strip().replace(" ", "_") | |
| value = re.sub(r"[^A-Za-z0-9_.-]+", "_", value) | |
| value = value.strip("._-") | |
| return value or "view" | |
| def _character_sort_key(character: str) -> tuple[int, str]: | |
| match = re.fullmatch(r"character_(\d+)", character) | |
| if match: | |
| return int(match.group(1)), character | |
| return 9999, character | |
| def _parse_character_ref_name(path: Path) -> tuple[str, str] | None: | |
| if path.suffix.lower() not in IMAGE_EXTS: | |
| return None | |
| if path.stem.endswith("_mask"): | |
| return None | |
| parent = path.parent.name | |
| if re.fullmatch(r"character_\d+", parent): | |
| return parent, _safe_name(path.stem) | |
| stem = path.stem | |
| match = re.match(r"^(character_\d+)(?:__|--|_)(.+)$", stem) | |
| if not match: | |
| return None | |
| return match.group(1), _safe_name(match.group(2)) | |
| def _normalize_upload_list(files) -> list[Path]: | |
| if not files: | |
| return [] | |
| if not isinstance(files, list): | |
| files = [files] | |
| return [_as_path(file, "reference file") for file in files] | |
| def _safe_extract_zip(zip_path: str | Path, label: str) -> Path: | |
| zip_path = _as_path(zip_path, label) | |
| if zip_path.suffix.lower() != ".zip": | |
| raise RuntimeError(f"{label} must be a .zip file.") | |
| extract_root = JOB_ROOT / "uploads" / uuid.uuid4().hex | |
| extract_root.mkdir(parents=True, exist_ok=True) | |
| with zipfile.ZipFile(zip_path) as zf: | |
| for item in zf.infolist(): | |
| item_path = Path(item.filename) | |
| if item_path.is_absolute() or ".." in item_path.parts: | |
| raise RuntimeError(f"Unsafe path in zip: {item.filename}") | |
| zf.extractall(extract_root) | |
| visible = [p for p in extract_root.iterdir() if p.name not in {".DS_Store", "__MACOSX"}] | |
| if len(visible) == 1 and visible[0].is_dir(): | |
| return visible[0] | |
| return extract_root | |
| def _collect_character_refs(reference_files, reference_zip) -> list[CharacterRef]: | |
| paths = [] | |
| if reference_zip is not None: | |
| root = _safe_extract_zip(reference_zip, "reference zip") | |
| paths.extend(p for p in root.rglob("*") if p.is_file()) | |
| paths.extend(_normalize_upload_list(reference_files)) | |
| refs = [] | |
| ignored = [] | |
| for path in paths: | |
| parsed = _parse_character_ref_name(path) | |
| if parsed is None: | |
| if path.suffix.lower() in IMAGE_EXTS: | |
| ignored.append(path.name) | |
| continue | |
| character, view = parsed | |
| refs.append(CharacterRef(character=character, view=view, path=path)) | |
| if not refs: | |
| raise RuntimeError( | |
| "No character references found. Use names like `character_0__front.png`, " | |
| "`character_0__back.png`, `character_1__front.png`, or upload a zip with " | |
| "`characters/character_0/front.png`." | |
| ) | |
| seen = set() | |
| duplicates = [] | |
| for ref in refs: | |
| key = (ref.character, ref.view) | |
| if key in seen: | |
| duplicates.append(f"{ref.character}/{ref.view}") | |
| seen.add(key) | |
| if duplicates: | |
| raise RuntimeError("Duplicate character view(s): " + ", ".join(sorted(duplicates))) | |
| if ignored: | |
| logging.info("Ignored reference image(s) without character naming: %s", ", ".join(ignored)) | |
| return sorted(refs, key=lambda ref: (_character_sort_key(ref.character), ref.view.lower())) | |
| def _text_args(text_prompt: str) -> list[str]: | |
| if not text_prompt or not text_prompt.strip(): | |
| return ["human character"] | |
| parts = [part.strip() for part in re.split(r"[\n,]+", text_prompt) if part.strip()] | |
| return parts or ["human character"] | |
| def _run_command(command: list[str], progress=None) -> str: | |
| logging.info("Running command: %s", " ".join(command)) | |
| output_lines = [] | |
| proc = subprocess.Popen( | |
| command, | |
| cwd=str(ROOT), | |
| stdout=subprocess.PIPE, | |
| stderr=subprocess.STDOUT, | |
| text=True, | |
| bufsize=1, | |
| env=os.environ.copy(), | |
| ) | |
| assert proc.stdout is not None | |
| for line in proc.stdout: | |
| line = line.rstrip() | |
| if not line: | |
| continue | |
| logging.info("[SCAIL-Pose] %s", line) | |
| output_lines.append(line) | |
| if progress is not None: | |
| progress(None, desc=line[:120]) | |
| ret = proc.wait() | |
| output = "\n".join(output_lines) | |
| if ret != 0: | |
| raise RuntimeError(f"SCAIL-Pose command failed with exit code {ret}.\n\n{output}") | |
| return output | |
| def _write_metadata(job_dir: Path, mode: str, prompt: str, extra: dict | None = None) -> None: | |
| metadata = { | |
| "mode": mode, | |
| "source": "scail-pose-gradio-pack-builder", | |
| } | |
| if extra: | |
| metadata.update(extra) | |
| (job_dir / "metadata.json").write_text(json.dumps(metadata, indent=2), encoding="utf-8") | |
| (job_dir / "prompt.txt").write_text(prompt or "", encoding="utf-8") | |
| def _zip_pack(job_dir: Path, mode: str) -> Path: | |
| zip_path = job_dir / f"scail2_{mode}_pack.zip" | |
| include = [ | |
| "ref.png", | |
| "ref_mask.png", | |
| "ref_mask.jpg", | |
| "rendered_v2.mp4", | |
| "rendered_mask_v2.mp4", | |
| "replace_mask.mp4", | |
| "prompt.txt", | |
| "metadata.json", | |
| ] | |
| with zipfile.ZipFile(zip_path, "w", compression=zipfile.ZIP_DEFLATED) as zf: | |
| for name in include: | |
| path = job_dir / name | |
| if path.exists(): | |
| zf.write(path, arcname=name) | |
| return zip_path | |
| def _require_outputs(job_dir: Path, names: list[str]) -> None: | |
| missing = [name for name in names if not (job_dir / name).exists()] | |
| if missing: | |
| raise RuntimeError("SCAIL-Pose did not produce expected output(s): " + ", ".join(missing)) | |
| def _load_sam3_predictors(sam3_model: Path): | |
| global _SAM3_VIDEO_PREDICTOR, _SAM3_IMAGE_PREDICTOR, _SAM3_PREDICTOR_KEY | |
| key = str(sam3_model) | |
| if _SAM3_VIDEO_PREDICTOR is not None and _SAM3_PREDICTOR_KEY == key: | |
| return _SAM3_VIDEO_PREDICTOR, _SAM3_IMAGE_PREDICTOR | |
| from ultralytics.models.sam import SAM3SemanticPredictor, SAM3VideoSemanticPredictor | |
| overrides = dict( | |
| conf=0.25, | |
| task="segment", | |
| mode="predict", | |
| imgsz=640, | |
| model=str(sam3_model), | |
| half=True, | |
| save=False, | |
| verbose=False, | |
| ) | |
| _SAM3_VIDEO_PREDICTOR = SAM3VideoSemanticPredictor(overrides=overrides, new_det_thresh=1.0) | |
| _SAM3_IMAGE_PREDICTOR = SAM3SemanticPredictor(overrides=overrides) | |
| _SAM3_PREDICTOR_KEY = key | |
| return _SAM3_VIDEO_PREDICTOR, _SAM3_IMAGE_PREDICTOR | |
| def _write_canonical_pack_zip(pack_root: Path, mode: str) -> Path: | |
| zip_path = pack_root.parent / f"scail2_{mode}_advanced_pack.zip" | |
| with zipfile.ZipFile(zip_path, "w", compression=zipfile.ZIP_DEFLATED) as zf: | |
| for path in sorted(pack_root.rglob("*")): | |
| if path.is_file(): | |
| zf.write(path, arcname=path.relative_to(pack_root).as_posix()) | |
| return zip_path | |
| def _advanced_gallery(pack_root: Path): | |
| items = [] | |
| for image_path in sorted((pack_root / "characters").rglob("*")): | |
| if image_path.is_file() and image_path.suffix.lower() in IMAGE_EXTS: | |
| items.append((str(image_path), image_path.relative_to(pack_root).as_posix())) | |
| return items | |
| def _build_advanced_metadata( | |
| pack_root: Path, | |
| mode: str, | |
| prompt: str, | |
| primary_ref: CharacterRef, | |
| character_names: list[str], | |
| text_prompt: str, | |
| ) -> None: | |
| primary_image = f"characters/{primary_ref.character}/{primary_ref.view}.png" | |
| primary_mask = f"characters/{primary_ref.character}/{primary_ref.view}_mask.png" | |
| metadata = { | |
| "mode": mode, | |
| "source": "scail-pose-gradio-advanced-pack-builder", | |
| "primary": { | |
| "image": primary_image, | |
| "mask": primary_mask, | |
| }, | |
| "driving": { | |
| "video": "rendered_v2.mp4", | |
| "mask_video": "rendered_mask_v2.mp4" if mode == "animation" else "replace_mask.mp4", | |
| }, | |
| "characters": character_names, | |
| "sam3_text": _text_args(text_prompt), | |
| } | |
| (pack_root / "metadata.json").write_text(json.dumps(metadata, indent=2), encoding="utf-8") | |
| (pack_root / "prompt.txt").write_text(prompt or "", encoding="utf-8") | |
| def _prepare_advanced_refs(refs: list[CharacterRef], pack_root: Path) -> list[CharacterRef]: | |
| prepared = [] | |
| for ref in refs: | |
| dest = pack_root / "characters" / ref.character / f"{ref.view}.png" | |
| _copy_reference_to_png(ref.path, dest) | |
| prepared.append(CharacterRef(character=ref.character, view=ref.view, path=dest)) | |
| return prepared | |
| def _animation_command( | |
| job_dir: Path, | |
| sam3_model: Path, | |
| max_persons: int, | |
| text_prompt: str, | |
| crop_mode: str, | |
| ) -> list[str]: | |
| command = [ | |
| sys.executable, | |
| str(ROOT / "NLFPoseExtract" / "process_animation_aio.py"), | |
| "--subdir", | |
| str(job_dir), | |
| "--video_name", | |
| "driving.mp4", | |
| "--e2e_mode", | |
| "--max_persons", | |
| str(int(max_persons)), | |
| "--sam3_model", | |
| str(sam3_model), | |
| "--text", | |
| *_text_args(text_prompt), | |
| ] | |
| if crop_mode == "mask silhouette": | |
| command.append("--crop_e2e_mask") | |
| elif crop_mode == "moving bbox": | |
| command.append("--crop_e2e_bbox") | |
| elif crop_mode == "steady bbox": | |
| command.append("--crop_e2e_steady_bbox") | |
| return command | |
| def _replacement_command( | |
| job_dir: Path, | |
| sam3_model: Path, | |
| text_prompt: str, | |
| matchnearest: bool, | |
| egocentric: bool, | |
| ) -> list[str]: | |
| command = [ | |
| sys.executable, | |
| str(ROOT / "NLFPoseExtract" / "process_replacement.py"), | |
| "--subdir", | |
| str(job_dir), | |
| "--video_name", | |
| "driving.mp4", | |
| "--sam3_model", | |
| str(sam3_model), | |
| "--text", | |
| *_text_args(text_prompt), | |
| ] | |
| if matchnearest: | |
| command.append("--matchnearest") | |
| if egocentric: | |
| command.append("--egocentric") | |
| return command | |
| def build_simple_pack( | |
| mode, | |
| ref_image, | |
| driving_video, | |
| prompt, | |
| sam3_text, | |
| max_persons, | |
| crop_mode, | |
| matchnearest, | |
| egocentric, | |
| progress=gr.Progress(track_tqdm=True), | |
| ): | |
| try: | |
| mode = str(mode or "Animation").lower() | |
| if mode not in {"animation", "replacement"}: | |
| raise RuntimeError(f"Unsupported mode: {mode}") | |
| if mode == "replacement" and matchnearest and egocentric: | |
| raise RuntimeError("matchnearest and egocentric are mutually exclusive.") | |
| progress(0.0, desc="Checking SCAIL-Pose repo") | |
| _require_repo_layout() | |
| progress(0.04, desc="Preparing SAM3 weights") | |
| sam3_model = _ensure_sam3_weights() | |
| if mode == "animation": | |
| job_dir = _new_job_dir("animation") | |
| progress(0.08, desc="Preparing animation inputs") | |
| _save_reference_image(ref_image, job_dir / "ref.png") | |
| _copy_video(driving_video, job_dir / "driving.mp4", "driving video") | |
| _write_metadata( | |
| job_dir, | |
| "animation", | |
| prompt, | |
| { | |
| "driving": { | |
| "video": "rendered_v2.mp4", | |
| "mask_video": "rendered_mask_v2.mp4", | |
| }, | |
| "primary": { | |
| "image": "ref.png", | |
| "mask": "ref_mask.jpg", | |
| }, | |
| "sam3_text": _text_args(sam3_text), | |
| "max_persons": int(max_persons), | |
| "crop_mode": crop_mode, | |
| }, | |
| ) | |
| progress(0.12, desc="Generating animation masks") | |
| logs = _run_command( | |
| _animation_command(job_dir, sam3_model, int(max_persons), sam3_text, crop_mode), | |
| progress=progress, | |
| ) | |
| _require_outputs(job_dir, ["ref_mask.jpg", "rendered_v2.mp4", "rendered_mask_v2.mp4"]) | |
| progress(0.92, desc="Packaging SCAIL-2 input pack") | |
| zip_path = _zip_pack(job_dir, "animation") | |
| progress(1.0, desc="Done") | |
| status = f"Done. Animation pack created at {zip_path}\n\n{logs}" | |
| return ( | |
| str(job_dir / "ref_mask.jpg"), | |
| str(job_dir / "rendered_v2.mp4"), | |
| str(job_dir / "rendered_mask_v2.mp4"), | |
| str(zip_path), | |
| status, | |
| ) | |
| job_dir = _new_job_dir("replacement") | |
| progress(0.08, desc="Preparing replacement inputs") | |
| _save_reference_image(ref_image, job_dir / "ref.png") | |
| _copy_video(driving_video, job_dir / "driving.mp4", "source video") | |
| _write_metadata( | |
| job_dir, | |
| "replacement", | |
| prompt, | |
| { | |
| "driving": { | |
| "video": "rendered_v2.mp4", | |
| "mask_video": "replace_mask.mp4", | |
| }, | |
| "primary": { | |
| "image": "ref.png", | |
| "mask": "ref_mask.png", | |
| }, | |
| "sam3_text": _text_args(sam3_text), | |
| "matchnearest": bool(matchnearest), | |
| "egocentric": bool(egocentric), | |
| }, | |
| ) | |
| progress(0.12, desc="Generating replacement masks") | |
| logs = _run_command( | |
| _replacement_command(job_dir, sam3_model, sam3_text, bool(matchnearest), bool(egocentric)), | |
| progress=progress, | |
| ) | |
| _require_outputs(job_dir, ["ref_mask.png", "rendered_v2.mp4", "replace_mask.mp4"]) | |
| progress(0.92, desc="Packaging SCAIL-2 input pack") | |
| zip_path = _zip_pack(job_dir, "replacement") | |
| progress(1.0, desc="Done") | |
| status = f"Done. Replacement pack created at {zip_path}\n\n{logs}" | |
| return ( | |
| str(job_dir / "ref_mask.png"), | |
| str(job_dir / "rendered_v2.mp4"), | |
| str(job_dir / "replace_mask.mp4"), | |
| str(zip_path), | |
| status, | |
| ) | |
| except Exception: | |
| logging.exception("Simple pack generation failed") | |
| return None, None, None, None, traceback.format_exc() | |
| def build_advanced_character_pack( | |
| reference_files, | |
| reference_zip, | |
| driving_video, | |
| prompt, | |
| sam3_text, | |
| max_subjects, | |
| progress=gr.Progress(track_tqdm=True), | |
| ): | |
| try: | |
| progress(0.0, desc="Checking SCAIL-Pose repo") | |
| _require_repo_layout() | |
| progress(0.04, desc="Preparing SAM3 weights") | |
| sam3_model = _ensure_sam3_weights() | |
| refs = _collect_character_refs(reference_files, reference_zip) | |
| character_names = sorted({ref.character for ref in refs}, key=_character_sort_key) | |
| if len(character_names) > len(SCAIL_COLORS): | |
| raise RuntimeError(f"Too many characters: {len(character_names)}. Limit is {len(SCAIL_COLORS)}.") | |
| max_targets = max(int(max_subjects or len(character_names)), len(character_names)) | |
| max_targets = max(1, min(max_targets, len(SCAIL_COLORS))) | |
| colors = SCAIL_COLORS[:max_targets] | |
| job_dir = _new_job_dir("advanced_animation") | |
| pack_root = job_dir / "scail2_input_pack" | |
| pack_root.mkdir(parents=True, exist_ok=True) | |
| progress(0.08, desc="Preparing references and driving video") | |
| work_driving = job_dir / "driving.mp4" | |
| _copy_video(driving_video, work_driving, "driving video") | |
| shutil.copy2(work_driving, pack_root / "rendered_v2.mp4") | |
| prepared_refs = _prepare_advanced_refs(refs, pack_root) | |
| primary = sorted( | |
| prepared_refs, | |
| key=lambda ref: ( | |
| 0 if ref.character == "character_0" else 1, | |
| 0 if ref.view in {"front", "main", "ref", "reference"} else 1, | |
| _character_sort_key(ref.character), | |
| ref.view, | |
| ), | |
| )[0] | |
| _build_advanced_metadata(pack_root, "animation", prompt, primary, character_names, sam3_text) | |
| progress(0.12, desc="Loading SAM3 predictors") | |
| video_predictor, image_predictor = _load_sam3_predictors(sam3_model) | |
| from TrackSam3.track import get_mask_from_image, get_mask_from_image_via_video, get_mask_from_video | |
| from NLFPoseExtract.v2_helper import save_colored_mask_image, write_colored_mask_video | |
| text = _text_args(sam3_text) | |
| progress(0.20, desc="Tracking subjects in driving video") | |
| drv_masks, drv_colors = get_mask_from_video( | |
| str(work_driving), | |
| video_predictor, | |
| max_targets=max_targets, | |
| sort_by="x", | |
| fixed_colors=colors, | |
| text=text, | |
| ) | |
| if len(drv_masks) == 0: | |
| raise RuntimeError("SAM3 did not detect any valid subject in the driving video.") | |
| import cv2 | |
| cap = cv2.VideoCapture(str(work_driving)) | |
| fps = cap.get(cv2.CAP_PROP_FPS) | |
| cap.release() | |
| fps_int = max(1, int(round(fps or 24))) | |
| write_colored_mask_video( | |
| drv_masks, | |
| drv_colors, | |
| str(pack_root / "rendered_mask_v2.mp4"), | |
| fps_int, | |
| bg_color=(255, 255, 255), | |
| ) | |
| progress(0.52, desc="Generating reference masks") | |
| logs = [ | |
| f"Characters: {', '.join(character_names)}", | |
| f"Reference views: {len(prepared_refs)}", | |
| f"Driving tracks detected: {len(drv_masks)}", | |
| f"Primary: {primary.character}/{primary.view}", | |
| ] | |
| for idx, ref in enumerate(prepared_refs): | |
| character_idx = character_names.index(ref.character) | |
| color = drv_colors[character_idx] if character_idx < len(drv_colors) else SCAIL_COLORS[character_idx] | |
| mask_path = ref.path.with_name(f"{ref.view}_mask.png") | |
| progress(None, desc=f"Masking {ref.character}/{ref.view}") | |
| try: | |
| ref_masks, _ = get_mask_from_image_via_video( | |
| str(ref.path), | |
| video_predictor, | |
| max_targets=1, | |
| sort_by="x", | |
| fixed_colors=[color], | |
| text=text, | |
| ) | |
| except Exception as exc: | |
| logging.warning("Video-based image masking failed for %s, trying image predictor: %s", ref.path, exc) | |
| ref_masks, _ = get_mask_from_image( | |
| str(ref.path), | |
| image_predictor, | |
| max_targets=1, | |
| sort_by="x", | |
| fixed_colors=[color], | |
| text=text, | |
| ) | |
| if len(ref_masks) == 0: | |
| raise RuntimeError(f"SAM3 did not detect a subject in {ref.character}/{ref.view}.") | |
| save_colored_mask_image([ref_masks[0]], [color], str(mask_path), bg_color=(255, 255, 255)) | |
| logs.append(f"- {ref.character}/{ref.view}: {mask_path.relative_to(pack_root).as_posix()}") | |
| progress(0.92, desc="Packaging canonical Advanced Pack") | |
| zip_path = _write_canonical_pack_zip(pack_root, "animation") | |
| gallery = _advanced_gallery(pack_root) | |
| progress(1.0, desc="Done") | |
| status = f"Done. Advanced pack created at {zip_path}\n\n" + "\n".join(logs) | |
| return gallery, str(pack_root / "rendered_v2.mp4"), str(pack_root / "rendered_mask_v2.mp4"), str(zip_path), status | |
| except Exception: | |
| logging.exception("Advanced character pack generation failed") | |
| return [], None, None, None, traceback.format_exc() | |
| def build_two_character_pack( | |
| character_0_ref, | |
| character_1_ref, | |
| driving_video, | |
| prompt, | |
| sam3_text, | |
| progress=gr.Progress(track_tqdm=True), | |
| ): | |
| try: | |
| progress(0.0, desc="Checking SCAIL-Pose repo") | |
| _require_repo_layout() | |
| progress(0.04, desc="Preparing SAM3 weights") | |
| sam3_model = _ensure_sam3_weights() | |
| job_dir = _new_job_dir("two_character_animation") | |
| pack_root = job_dir / "scail2_input_pack" | |
| pack_root.mkdir(parents=True, exist_ok=True) | |
| progress(0.08, desc="Preparing two-character inputs") | |
| work_driving = job_dir / "driving.mp4" | |
| _copy_video(driving_video, work_driving, "driving video") | |
| shutil.copy2(work_driving, pack_root / "rendered_v2.mp4") | |
| prepared_refs = [ | |
| CharacterRef( | |
| character="character_0", | |
| view="front", | |
| path=_copy_reference_to_png(_as_path(character_0_ref, "character 0 reference"), pack_root / "characters" / "character_0" / "front.png"), | |
| ), | |
| CharacterRef( | |
| character="character_1", | |
| view="front", | |
| path=_copy_reference_to_png(_as_path(character_1_ref, "character 1 reference"), pack_root / "characters" / "character_1" / "front.png"), | |
| ), | |
| ] | |
| character_names = ["character_0", "character_1"] | |
| primary = prepared_refs[0] | |
| _build_advanced_metadata(pack_root, "animation", prompt, primary, character_names, sam3_text) | |
| progress(0.12, desc="Loading SAM3 predictors") | |
| video_predictor, image_predictor = _load_sam3_predictors(sam3_model) | |
| from TrackSam3.track import get_mask_from_image, get_mask_from_image_via_video, get_mask_from_video | |
| from NLFPoseExtract.v2_helper import save_colored_mask_image, write_colored_mask_video | |
| text = _text_args(sam3_text) | |
| colors = SCAIL_COLORS[:2] | |
| progress(0.20, desc="Tracking two subjects left-to-right") | |
| drv_masks, drv_colors = get_mask_from_video( | |
| str(work_driving), | |
| video_predictor, | |
| max_targets=2, | |
| sort_by="x", | |
| fixed_colors=colors, | |
| text=text, | |
| ) | |
| if len(drv_masks) < 2: | |
| raise RuntimeError( | |
| f"SAM3 detected {len(drv_masks)} subject(s), but this mode expects two. " | |
| "Try a clearer driving video or use the Advanced Character Pack." | |
| ) | |
| import cv2 | |
| cap = cv2.VideoCapture(str(work_driving)) | |
| fps = cap.get(cv2.CAP_PROP_FPS) | |
| cap.release() | |
| fps_int = max(1, int(round(fps or 24))) | |
| write_colored_mask_video( | |
| drv_masks[:2], | |
| drv_colors[:2], | |
| str(pack_root / "rendered_mask_v2.mp4"), | |
| fps_int, | |
| bg_color=(255, 255, 255), | |
| ) | |
| progress(0.55, desc="Generating reference masks") | |
| logs = [ | |
| "Mode: two-character animation", | |
| "Mapping: left-most driving subject -> character_0; next subject to the right -> character_1", | |
| f"Driving tracks detected: {len(drv_masks)}", | |
| ] | |
| for ref in prepared_refs: | |
| character_idx = character_names.index(ref.character) | |
| color = drv_colors[character_idx] | |
| mask_path = ref.path.with_name(f"{ref.view}_mask.png") | |
| progress(None, desc=f"Masking {ref.character}") | |
| try: | |
| ref_masks, _ = get_mask_from_image_via_video( | |
| str(ref.path), | |
| video_predictor, | |
| max_targets=1, | |
| sort_by="x", | |
| fixed_colors=[color], | |
| text=text, | |
| ) | |
| except Exception as exc: | |
| logging.warning("Video-based image masking failed for %s, trying image predictor: %s", ref.path, exc) | |
| ref_masks, _ = get_mask_from_image( | |
| str(ref.path), | |
| image_predictor, | |
| max_targets=1, | |
| sort_by="x", | |
| fixed_colors=[color], | |
| text=text, | |
| ) | |
| if len(ref_masks) == 0: | |
| raise RuntimeError(f"SAM3 did not detect a subject in {ref.character}.") | |
| save_colored_mask_image([ref_masks[0]], [color], str(mask_path), bg_color=(255, 255, 255)) | |
| logs.append(f"- {ref.character}: {mask_path.relative_to(pack_root).as_posix()}") | |
| progress(0.92, desc="Packaging two-character pack") | |
| zip_path = _write_canonical_pack_zip(pack_root, "two_character_animation") | |
| gallery = _advanced_gallery(pack_root) | |
| progress(1.0, desc="Done") | |
| status = f"Done. Two-character pack created at {zip_path}\n\n" + "\n".join(logs) | |
| return gallery, str(pack_root / "rendered_v2.mp4"), str(pack_root / "rendered_mask_v2.mp4"), str(zip_path), status | |
| except Exception: | |
| logging.exception("Two-character pack generation failed") | |
| return [], None, None, None, traceback.format_exc() | |
| def build_ui(): | |
| with gr.Blocks(title="SCAIL-Pose Pack Builder") as demo: | |
| gr.Markdown( | |
| "# SCAIL-Pose Pack Builder\n" | |
| "Generate SCAIL-2-ready masks and export them as an input pack. " | |
| "Use the downloaded `.zip` in the SCAIL-2 demo Advanced Pack tab." | |
| ) | |
| gr.Textbox(value=_repo_status(), label="Startup status", interactive=False, lines=7) | |
| with gr.Tab("Simple Pack"): | |
| gr.Markdown( | |
| "Create a standard SCAIL-2 pack from one reference image and one video. " | |
| "Choose Animation to follow motion from the driving video, or Replacement to replace " | |
| "a target region in the source video." | |
| ) | |
| simple_mode = gr.Radio( | |
| ["Animation", "Replacement"], | |
| value="Animation", | |
| label="Mode", | |
| ) | |
| with gr.Row(): | |
| simple_ref = gr.Image(type="filepath", label="Reference image") | |
| simple_driving = gr.Video(label="Driving / source video") | |
| simple_prompt = gr.Textbox(label="Prompt for SCAIL-2", lines=3) | |
| simple_sam3_text = gr.Textbox( | |
| value="human character", | |
| label="SAM3 text prompt", | |
| info="Use one or more comma/newline-separated prompts, e.g. `human character` or `woman in red dress, man in blue jacket`.", | |
| ) | |
| with gr.Accordion("Animation options", open=True): | |
| simple_max_persons = gr.Number(value=2, precision=0, label="Max tracked subjects") | |
| simple_crop = gr.Dropdown( | |
| ["none", "mask silhouette", "moving bbox", "steady bbox"], | |
| value="none", | |
| label="Driving crop mode", | |
| ) | |
| with gr.Accordion("Replacement options", open=False): | |
| simple_matchnearest = gr.Checkbox( | |
| value=False, | |
| label="Match nearest target when two people are detected", | |
| ) | |
| simple_egocentric = gr.Checkbox( | |
| value=False, | |
| label="Egocentric: union two detected parts into one actor", | |
| ) | |
| simple_run = gr.Button("Generate pack", variant="primary") | |
| with gr.Row(): | |
| simple_ref_mask = gr.Image(label="Reference mask", interactive=False) | |
| simple_mask_video = gr.Video(label="Mask video") | |
| simple_rendered = gr.Video(label="Rendered / source video") | |
| simple_zip = gr.File(label="Download SCAIL-2 pack") | |
| simple_status = gr.Textbox(label="Run logs", lines=14) | |
| simple_run.click( | |
| build_simple_pack, | |
| inputs=[ | |
| simple_mode, | |
| simple_ref, | |
| simple_driving, | |
| simple_prompt, | |
| simple_sam3_text, | |
| simple_max_persons, | |
| simple_crop, | |
| simple_matchnearest, | |
| simple_egocentric, | |
| ], | |
| outputs=[simple_ref_mask, simple_rendered, simple_mask_video, simple_zip, simple_status], | |
| ) | |
| with gr.Tab("Two Characters"): | |
| gr.Markdown( | |
| "Fast path for the common two-character case. Upload two reference images and one driving video. " | |
| "`Character 0` should match the left-most subject in the driving video; `Character 1` should match " | |
| "the next subject to the right." | |
| ) | |
| with gr.Row(): | |
| two_char0 = gr.Image(type="filepath", label="Character 0 reference") | |
| two_char1 = gr.Image(type="filepath", label="Character 1 reference") | |
| two_driving = gr.Video(label="Driving video") | |
| with gr.Accordion("Optional settings", open=False): | |
| two_prompt = gr.Textbox(label="Prompt for SCAIL-2", lines=3) | |
| two_sam3_text = gr.Textbox( | |
| value="human character", | |
| label="SAM3 text prompt", | |
| info="Use a broad prompt such as `human character`, or comma/newline-separated prompts for distinct subjects.", | |
| ) | |
| two_run = gr.Button("Generate two-character pack", variant="primary") | |
| two_gallery = gr.Gallery( | |
| label="References and generated masks", | |
| columns=4, | |
| height=320, | |
| selected_index=0, | |
| preview=True, | |
| ) | |
| with gr.Row(): | |
| two_rendered = gr.Video(label="Rendered / driving video") | |
| two_mask_video = gr.Video(label="Driving mask video") | |
| two_zip = gr.File(label="Download two-character SCAIL-2 pack") | |
| two_status = gr.Textbox(label="Run logs", lines=14) | |
| two_run.click( | |
| build_two_character_pack, | |
| inputs=[ | |
| two_char0, | |
| two_char1, | |
| two_driving, | |
| two_prompt, | |
| two_sam3_text, | |
| ], | |
| outputs=[two_gallery, two_rendered, two_mask_video, two_zip, two_status], | |
| ) | |
| with gr.Tab("Advanced Character Pack"): | |
| gr.Markdown( | |
| "Create a canonical SCAIL-2 animation pack for multi-reference or multi-character cases. " | |
| "Use this when one character has several views, or when several characters need separate reference slots." | |
| ) | |
| with gr.Accordion("Reference naming", open=True): | |
| gr.Markdown( | |
| "The Advanced Character Pack uses a simple identity convention: " | |
| "`character_0` must match the left-most tracked subject in the driving video, " | |
| "`character_1` the next subject to the right, and so on. " | |
| "Name your reference files according to that left-to-right order.\n\n" | |
| "Example: if the driving video starts with the woman on the left and the man on " | |
| "the right, use `character_0` for the woman and `character_1` for the man.\n\n" | |
| "Upload reference images with names like:\n\n" | |
| "```text\n" | |
| "character_0__front.png\n" | |
| "character_0__back.png\n" | |
| "character_0__closeup.png\n" | |
| "character_1__front.png\n" | |
| "```\n\n" | |
| "Or upload a zip with this structure:\n\n" | |
| "```text\n" | |
| "characters/\n" | |
| " character_0/\n" | |
| " front.png\n" | |
| " back.png\n" | |
| " character_1/\n" | |
| " front.png\n" | |
| "```\n\n" | |
| "All views of the same character should use the same `character_N` prefix. " | |
| "If the detected order is ambiguous or characters cross over, use a simple case first; " | |
| "manual track assignment is not part of this V1." | |
| ) | |
| with gr.Row(): | |
| adv_refs = gr.Files( | |
| label="Reference images", | |
| file_types=[".png", ".jpg", ".jpeg", ".webp"], | |
| type="filepath", | |
| ) | |
| adv_zip_in = gr.File( | |
| label="Reference zip instead of files", | |
| file_types=[".zip"], | |
| type="filepath", | |
| ) | |
| adv_driving = gr.Video(label="Driving video") | |
| adv_prompt = gr.Textbox(label="Prompt for SCAIL-2", lines=3) | |
| with gr.Row(): | |
| adv_sam3_text = gr.Textbox( | |
| value="human character", | |
| label="SAM3 text prompt", | |
| info="Use broad prompts to detect all subjects, or comma/newline-separated prompts for distinct subjects.", | |
| ) | |
| adv_max_subjects = gr.Number(value=2, precision=0, label="Max tracked subjects") | |
| adv_run = gr.Button("Generate advanced character pack", variant="primary") | |
| adv_gallery = gr.Gallery( | |
| label="Parsed references and generated masks", | |
| columns=4, | |
| height=320, | |
| selected_index=0, | |
| preview=True, | |
| ) | |
| with gr.Row(): | |
| adv_rendered = gr.Video(label="Rendered / driving video") | |
| adv_mask_video = gr.Video(label="Driving mask video") | |
| adv_zip = gr.File(label="Download canonical SCAIL-2 Advanced Pack") | |
| adv_status = gr.Textbox(label="Run logs", lines=14) | |
| adv_run.click( | |
| build_advanced_character_pack, | |
| inputs=[ | |
| adv_refs, | |
| adv_zip_in, | |
| adv_driving, | |
| adv_prompt, | |
| adv_sam3_text, | |
| adv_max_subjects, | |
| ], | |
| outputs=[adv_gallery, adv_rendered, adv_mask_video, adv_zip, adv_status], | |
| ) | |
| with gr.Tab("Pack Format"): | |
| gr.Markdown( | |
| "The Simple Pack tab exports flat packs. The Two Characters and Advanced Character Pack tabs " | |
| "export canonical character packs with explicit character folders.\n\n" | |
| "Simple animation pack:\n" | |
| "```text\n" | |
| "ref.png\n" | |
| "ref_mask.jpg\n" | |
| "rendered_v2.mp4\n" | |
| "rendered_mask_v2.mp4\n" | |
| "prompt.txt\n" | |
| "metadata.json\n" | |
| "```\n\n" | |
| "Replacement pack:\n" | |
| "```text\n" | |
| "ref.png\n" | |
| "ref_mask.png\n" | |
| "rendered_v2.mp4\n" | |
| "replace_mask.mp4\n" | |
| "prompt.txt\n" | |
| "metadata.json\n" | |
| "```\n\n" | |
| "Advanced character pack:\n" | |
| "```text\n" | |
| "rendered_v2.mp4\n" | |
| "rendered_mask_v2.mp4\n" | |
| "prompt.txt\n" | |
| "metadata.json\n" | |
| "characters/\n" | |
| " character_0/\n" | |
| " front.png\n" | |
| " front_mask.png\n" | |
| " back.png\n" | |
| " back_mask.png\n" | |
| " character_1/\n" | |
| " front.png\n" | |
| " front_mask.png\n" | |
| "```\n\n" | |
| "For multi-reference, use several views under the same `character_N`. " | |
| "For multi-character, use one folder per identity slot." | |
| ) | |
| return demo | |
| if __name__ == "__main__": | |
| build_ui().queue(max_size=4).launch(show_error=True) |