Download hy_dev_gen_long.py from Cccccz/comparison: direct link, hf CLI and curl.
- Browser
- Download file 3.13 kB
-
https://huggingface.co/Cccccz/comparison/resolve/main/hy_dev_gen_long.py
- Command line
-
hf download hf://Cccccz/comparison/hy_dev_gen_long.py
-
curl -L -o hy_dev_gen_long.py https://huggingface.co/Cccccz/comparison/resolve/main/hy_dev_gen_long.py
3.13 kB
| #!/usr/bin/env python | |
| """Run HY-WorldPlay-DEV-Predictor's generate_predictor_v4_train_case_eval.py at a | |
| video length other than the 125-frame protocol (long-video rows of the comparison). | |
| HY_NUM_LATENTS (64 = 16 chunks = 253 frames, 128 = 32 chunks = 509 frames) replaces the | |
| generator's hard-coded 32 latents / 125 frames: | |
| * the four bidirectional camera actions are rescaled the way hycache does it | |
| ("w-15,s-16" -> "w-31,s-32" for 64 latents: forward for floor((n-1)/2) latents, | |
| back for the rest, so the video still turns around half-way); | |
| * pose_to_input gets the new latent count and the pipeline call the matching | |
| video_length. | |
| The pipeline's own long-video mechanism (pose-retrieved memory frames, at most 20 | |
| history latents per chunk) is left as is. DisCa's forward drops the ATC-only keywords | |
| the rollout passes (same as hy_dev_gen_disca.py).""" | |
| import inspect, os, sys | |
| DEV = "/local/zoubin/cz/projects/HY-WorldPlay-DEV-Predictor" | |
| sys.path.insert(0, DEV); os.chdir(DEV) | |
| N = int(os.environ["HY_NUM_LATENTS"]) | |
| assert N % 4 == 0 and N > 0, N | |
| FRAMES = (N - 1) * 4 + 1 | |
| import hyvideo.generate as G # noqa: E402 | |
| import predictor_data.prefeature_schema as S # noqa: E402 | |
| from hyvideo.pipelines.worldplay_video_pipeline import HunyuanVideo_1_5_Pipeline # noqa: E402 | |
| import models # noqa: E402 | |
| def rescale(pose): | |
| (fwd, _), (back, _) = [p.strip().split("-") for p in pose.split(",")] | |
| n = N - 1 | |
| return f"{fwd}-{n // 2},{back}-{n - n // 2}" | |
| S.DEFAULT_ACTIONS = tuple((name, rescale(pose)) for name, pose in S.DEFAULT_ACTIONS) | |
| print(f"[hy_dev_gen_long] {N} latents = {FRAMES} frames; actions {S.DEFAULT_ACTIONS}", flush=True) | |
| _pose_to_input = G.pose_to_input | |
| def pose_to_input(pose_data, latent_num, tps=False): | |
| return _pose_to_input(pose_data, N if latent_num == 32 else latent_num, tps) | |
| G.pose_to_input = pose_to_input | |
| _call = HunyuanVideo_1_5_Pipeline.__call__ | |
| def __call__(self, *args, **kw): | |
| if kw.get("video_length") == 125: | |
| kw["video_length"] = FRAMES | |
| return _call(self, *args, **kw) | |
| HunyuanVideo_1_5_Pipeline.__call__ = __call__ | |
| _orig = models.HYWorldPlayPredictorDisCa.forward | |
| _accepted = set(inspect.signature(_orig).parameters) | |
| def forward(self, *args, **kw): | |
| return _orig(self, *args, **{k: v for k, v in kw.items() if k in _accepted}) | |
| models.HYWorldPlayPredictorDisCa.forward = forward | |
| # The generator's completion check (resume + post-encode validation) hard-codes 125 | |
| # decoded frames; load it as a module so the check can be re-pointed at FRAMES. | |
| import av # noqa: E402 | |
| import tools.generate_predictor_v4_train_case_eval as gen # noqa: E402 | |
| def video_is_complete(path): | |
| if not path.is_file() or path.stat().st_size == 0: | |
| return False | |
| try: | |
| with av.open(str(path)) as container: | |
| stream = container.streams.video[0] | |
| if stream.width != 832 or stream.height != 480: | |
| return False | |
| return sum(1 for _ in container.decode(stream)) == FRAMES | |
| except Exception: | |
| return False | |
| gen.video_is_complete = video_is_complete | |
| sys.argv[0] = gen.__file__ | |
| gen.main() | |