comparison / hy_dev_gen_long.py
Cccccz's picture
Add files using upload-large-folder tool
38a51ff verified
Raw
History Blame Contribute Delete
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()