| 1 | """Public 10-second pure-UCPE image-to-video inference entrypoint.""" |
| 2 | |
| 3 | from __future__ import annotations |
| 4 | |
| 5 | import argparse |
| 6 | import json |
| 7 | import subprocess |
| 8 | import sys |
| 9 | from pathlib import Path |
| 10 | |
| 11 | import torch |
| 12 | import yaml |
| 13 | |
| 14 | ROOT = Path(__file__).resolve().parent |
| 15 | REPO_ROOT = ROOT.parent |
| 16 | for package in ("ltx-core/src", "ltx-pipelines/src"): |
| 17 | sys.path.insert(0, str(ROOT / package)) |
| 18 | |
| 19 | from ltx_core.components.guiders import MultiModalGuiderParams # noqa: E402 |
| 20 | from ltx_pipelines.ti2vid_one_stage import TI2VidOneStagePipeline # noqa: E402 |
| 21 | from ltx_core.model.video_vae.tiling import TilingConfig # noqa: E402 |
| 22 | from ltx_core.model.video_vae.video_vae import get_video_chunks_number # noqa: E402 |
| 23 | from ltx_pipelines.utils.args import ImageConditioningInput # noqa: E402 |
| 24 | from ltx_pipelines.utils.media_io import encode_video # noqa: E402 |
| 25 | |
| 26 | from helpers.action_condition import ( # noqa: E402 |
| 27 | action_config, |
| 28 | build_action_condition, |
| 29 | build_action_trajectory, |
| 30 | ) |
| 31 | from helpers.action_camera import ( # noqa: E402 |
| 32 | DEFAULT_PITCH_LIMIT_DEG, |
| 33 | DEFAULT_ROTATION_SPEED_DEG, |
| 34 | DEFAULT_TRANSLATION_SPEED, |
| 35 | ) |
| 36 | from helpers.action_overlay import overlay_genie_on_video # noqa: E402 |
| 37 | |
| 38 | DEFAULT_CONFIG = ROOT / "configs" / "inference_wm.yaml" |
| 39 | NEGATIVE_PROMPT = ( |
| 40 | "worst quality, inconsistent motion, blurry, jittery, distorted, " |
| 41 | "game UI, video game interface, HUD, heads-up display, menu, status bar, " |
| 42 | "health bar, score, minimap, crosshair, reticle, buttons, icons, subtitles, " |
| 43 | "captions, watermark, logo, text overlay, user interface" |
| 44 | ) |
| 45 | |
| 46 | |
| 47 | def _load_config(path: Path) -> dict: |
| 48 | return yaml.safe_load(path.read_text()) or {} |
| 49 | |
| 50 | |
| 51 | def _override(value, default): |
| 52 | return default if value is None else value |
| 53 | |
| 54 | |
| 55 | def _auto_fov(image: Path, model: str, python_bin: str, width: int, height: int) -> float: |
| 56 | helper = ROOT / "helpers" / "moge_fov.py" |
| 57 | raw = subprocess.run( |
| 58 | [python_bin, str(helper), "--image", str(image), "--model", model, |
| 59 | "--target-width", str(width), "--target-height", str(height)], |
| 60 | check=True, capture_output=True, text=True, |
| 61 | ).stdout.strip() |
| 62 | return float(json.loads(raw)["fov_x_deg"]) |
| 63 | |
| 64 | |
| 65 | def parse_args() -> argparse.Namespace: |
| 66 | parser = argparse.ArgumentParser(description=__doc__) |
| 67 | parser.add_argument("--config", type=Path, default=DEFAULT_CONFIG) |
| 68 | parser.add_argument("--image", type=Path, required=True, help="First-frame image; this entrypoint is I2V-only.") |
| 69 | parser.add_argument("--prompt", default=None) |
| 70 | parser.add_argument("--action-str", required=True) |
| 71 | parser.add_argument("--checkpoint", type=Path, default=None) |
| 72 | parser.add_argument("--gemma-path", type=Path, default=None) |
| 73 | parser.add_argument("--output", type=Path, default=Path("outputs/echo_wm.mp4")) |
| 74 | parser.add_argument("--auto-fov", action="store_true") |
| 75 | parser.add_argument("--moge-model", default="Ruicheng/moge-2-vitl-normal") |
| 76 | parser.add_argument("--moge-python", default=sys.executable) |
| 77 | parser.add_argument("--fov-deg", type=float, default=None) |
| 78 | parser.add_argument("--translation-speed", type=float, default=None) |
| 79 | parser.add_argument("--rotation-speed-deg", type=float, default=None) |
| 80 | parser.add_argument("--pitch-limit-deg", type=float, default=None) |
| 81 | parser.add_argument("--width", type=int, default=None) |
| 82 | parser.add_argument("--height", type=int, default=None) |
| 83 | parser.add_argument("--num-frames", type=int, default=None) |
| 84 | parser.add_argument("--fps", type=float, default=None) |
| 85 | parser.add_argument("--steps", type=int, default=None) |
| 86 | parser.add_argument("--guidance-scale", type=float, default=None) |
| 87 | parser.add_argument("--video-cfg", type=float, default=None, help="Video CFG scale (default: 4.0).") |
| 88 | parser.add_argument("--audio-cfg", type=float, default=None, help="Audio CFG scale (default: 2.0).") |
| 89 | parser.add_argument("--negative-prompt", default=None, help="Negative prompt; overrides config.") |
| 90 | parser.add_argument("--stg-scale", type=float, default=None) |
| 91 | parser.add_argument("--stg-blocks", type=int, nargs="+", default=None) |
| 92 | parser.add_argument("--seed", type=int, default=None) |
| 93 | parser.add_argument("--no-audio", action="store_true") |
| 94 | parser.add_argument( |
| 95 | "--action-overlay", action=argparse.BooleanOptionalAction, default=True, |
| 96 | help="Write a second MP4 with a Genie-style WASD/rotation HUD overlay " |
| 97 | "(default: enabled; disable with --no-action-overlay).", |
| 98 | ) |
| 99 | return parser.parse_args() |
| 100 | |
| 101 | |
| 102 | @torch.inference_mode() |
| 103 | def main() -> None: |
| 104 | args = parse_args() |
| 105 | if not args.image.is_file(): |
| 106 | raise FileNotFoundError(f"I2V first-frame image not found: {args.image}") |
| 107 | cfg = _load_config(args.config) |
| 108 | video_cfg = cfg.get("video", {}) |
| 109 | model_cfg = cfg.get("model", {}) |
| 110 | action_cfg = cfg.get("action", {}) |
| 111 | checkpoint = args.checkpoint or ROOT / model_cfg.get("checkpoint", "checkpoints/echo-wm-base.safetensors") |
| 112 | gemma_path = args.gemma_path or ROOT / model_cfg["gemma_path"] |
| 113 | width = _override(args.width, video_cfg.get("width", 1280)) |
| 114 | height = _override(args.height, video_cfg.get("height", 704)) |
| 115 | num_frames = _override(args.num_frames, video_cfg.get("num_frames", 241)) |
| 116 | fps = _override(args.fps, video_cfg.get("fps", 24.0)) |
| 117 | steps = _override(args.steps, video_cfg.get("steps", 30)) |
| 118 | seed = args.seed if args.seed is not None else video_cfg.get("seed", 42) |
| 119 | legacy_guidance = args.guidance_scale |
| 120 | video_cfg_scale = _override( |
| 121 | args.video_cfg, |
| 122 | legacy_guidance if legacy_guidance is not None else video_cfg.get("video_cfg", 4.0), |
| 123 | ) |
| 124 | audio_cfg_scale = _override( |
| 125 | args.audio_cfg, |
| 126 | legacy_guidance if legacy_guidance is not None else video_cfg.get("audio_cfg", 2.0), |
| 127 | ) |
| 128 | negative_prompt = args.negative_prompt or cfg.get("negative_prompt", NEGATIVE_PROMPT) |
| 129 | stg_scale = _override(args.stg_scale, video_cfg.get("stg_scale", 1.0)) |
| 130 | stg_blocks = _override(args.stg_blocks, video_cfg.get("stg_blocks", [29])) |
| 131 | fov = _override(args.fov_deg, action_cfg.get("fov_deg", 70.0)) |
| 132 | if not args.prompt: |
| 133 | raise ValueError("Provide --prompt using the six-field format in PROMPT_SKILL.md") |
| 134 | prompt = args.prompt |
| 135 | if args.auto_fov: |
| 136 | fov = _auto_fov(args.image, args.moge_model, args.moge_python, width, height) |
| 137 | |
| 138 | action = build_action_condition( |
| 139 | args.action_str, num_frames=num_frames, width=width, height=height, |
| 140 | translation_speed=_override( |
| 141 | args.translation_speed, action_cfg.get("translation_speed", DEFAULT_TRANSLATION_SPEED) |
| 142 | ), |
| 143 | rotation_speed_deg=_override( |
| 144 | args.rotation_speed_deg, action_cfg.get("rotation_speed_deg", DEFAULT_ROTATION_SPEED_DEG) |
| 145 | ), |
| 146 | pitch_limit_deg=_override( |
| 147 | args.pitch_limit_deg, action_cfg.get("pitch_limit_deg", DEFAULT_PITCH_LIMIT_DEG) |
| 148 | ), |
| 149 | fov_deg=fov, device=torch.device("cuda" if torch.cuda.is_available() else "cpu"), fps=fps, |
| 150 | ) |
| 151 | trajectory = None |
| 152 | if args.action_overlay: |
| 153 | trajectory = build_action_trajectory( |
| 154 | args.action_str, |
| 155 | num_frames=num_frames, |
| 156 | translation_speed=_override( |
| 157 | args.translation_speed, action_cfg.get("translation_speed", DEFAULT_TRANSLATION_SPEED) |
| 158 | ), |
| 159 | rotation_speed_deg=_override( |
| 160 | args.rotation_speed_deg, action_cfg.get("rotation_speed_deg", DEFAULT_ROTATION_SPEED_DEG) |
| 161 | ), |
| 162 | pitch_limit_deg=_override( |
| 163 | args.pitch_limit_deg, action_cfg.get("pitch_limit_deg", DEFAULT_PITCH_LIMIT_DEG) |
| 164 | ), |
| 165 | fps=fps, |
| 166 | ) |
| 167 | device = torch.device("cuda" if torch.cuda.is_available() else "cpu") |
| 168 | pipeline = TI2VidOneStagePipeline( |
| 169 | checkpoint_path=str(checkpoint), gemma_root=str(gemma_path), loras=(), device=device, |
| 170 | action_config=action_config(width, height), |
| 171 | ) |
| 172 | video, audio = pipeline( |
| 173 | prompt=prompt, negative_prompt=negative_prompt, seed=seed, height=height, width=width, |
| 174 | num_frames=num_frames, frame_rate=fps, num_inference_steps=steps, |
| 175 | video_guider_params=MultiModalGuiderParams( |
| 176 | cfg_scale=video_cfg_scale, stg_scale=stg_scale, stg_blocks=stg_blocks |
| 177 | ), |
| 178 | audio_guider_params=MultiModalGuiderParams( |
| 179 | cfg_scale=audio_cfg_scale, stg_scale=stg_scale, stg_blocks=stg_blocks |
| 180 | ), |
| 181 | images=[ImageConditioningInput(str(args.image), 0, 1.0)], action_cond=action, |
| 182 | video_tiling_config=TilingConfig.default(), |
| 183 | ) |
| 184 | args.output.parent.mkdir(parents=True, exist_ok=True) |
| 185 | encode_video( |
| 186 | video=video, |
| 187 | fps=int(fps), |
| 188 | audio=None if args.no_audio else audio, |
| 189 | output_path=str(args.output), |
| 190 | video_chunks_number=get_video_chunks_number(num_frames, TilingConfig.default()), |
| 191 | ) |
| 192 | overlay_output = None |
| 193 | if trajectory is not None: |
| 194 | overlay_output = args.output.with_name(f"{args.output.stem}_action{args.output.suffix}") |
| 195 | overlay_genie_on_video(args.output, trajectory, output_path=overlay_output) |
| 196 | metadata = { |
| 197 | "prompt": prompt, "action": args.action_str, "fov_deg": fov, "seed": seed, |
| 198 | "width": width, "height": height, "num_frames": num_frames, "fps": fps, |
| 199 | "action_overlay": bool(args.action_overlay), |
| 200 | "overlay_output": overlay_output.name if overlay_output else None, |
| 201 | "video_cfg": video_cfg_scale, |
| 202 | "audio_cfg": audio_cfg_scale, |
| 203 | "negative_prompt": negative_prompt, |
| 204 | } |
| 205 | args.output.with_suffix(".json").write_text(json.dumps(metadata, indent=2), encoding="utf-8") |
| 206 | print(f"Saved {args.output}") |
| 207 | if overlay_output: |
| 208 | print(f"Saved {overlay_output}") |
| 209 | |
| 210 | |
| 211 | if __name__ == "__main__": |
| 212 | main() |
| 213 |