返回 JoyAI-Echo
inference_wm.py
根目录 / echo_wm / inference_wm.py
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
213 lines PYTHON