File size: 3,128 Bytes
38a51ff
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
#!/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()