Skip to content

Commit 0f17823

Browse files
committed
Add waypoint-1.5-1B support
This patch adds support for waypoint-1.5-1B, a realtime interactive video world model, which is designed to generate on consumer hardware. This patch: 1) extends the existing taehv to support encoder; 2) implements the necessary bits, e.g. DiT, adaptive RMSNorm, etc; 3) supports using JSON file to pass in controls conditions. Signed-off-by: Lin, Peiyong <linpyong@gmail.com>
1 parent 6151302 commit 0f17823

26 files changed

Lines changed: 4122 additions & 36 deletions

File tree

Lines changed: 15 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,15 @@
1+
{
2+
"schema_version": 1,
3+
"actions": [
4+
{"mouse_dx": 0.2, "mouse_dy": 0.2}, {"buttons": [32]}, {}, {}, {}, {"buttons": [1]}, {}, {}, {"buttons": [1, 32]}, {}, {}, {}, {}, {}, {},
5+
{"mouse_dx": 0.2, "mouse_dy": 0.2}, {"buttons": [32]}, {}, {}, {}, {"buttons": [1]}, {}, {}, {"buttons": [1, 32]}, {}, {}, {}, {}, {}, {},
6+
{"mouse_dx": 0.2, "mouse_dy": 0.2}, {"buttons": [32]}, {}, {}, {}, {"buttons": [1]}, {}, {}, {"buttons": [1, 32]}, {}, {}, {}, {}, {}, {},
7+
{"mouse_dx": 0.2, "mouse_dy": 0.2}, {"buttons": [32]}, {}, {}, {}, {"buttons": [1]}, {}, {}, {"buttons": [1, 32]}, {}, {}, {}, {}, {}, {},
8+
{}, {}, {}, {}, {}, {}, {}, {},
9+
{"buttons": [32]}, {"buttons": [32]}, {"buttons": [32]}, {"buttons": [32]}, {"buttons": [32]}, {"buttons": [32]}, {"buttons": [32]}, {"buttons": [32]}, {"buttons": [32]}, {"buttons": [32]},
10+
{"buttons": [65]}, {"buttons": [65]}, {"buttons": [65]}, {"buttons": [65]}, {"buttons": [65]}, {"buttons": [65]}, {"buttons": [65]}, {"buttons": [65]}, {"buttons": [65]}, {"buttons": [65]},
11+
{"buttons": [68]}, {"buttons": [68]}, {"buttons": [68]}, {"buttons": [68]}, {"buttons": [68]}, {"buttons": [68]}, {"buttons": [68]}, {"buttons": [68]}, {"buttons": [68]}, {"buttons": [68]},
12+
{"buttons": [83]}, {"buttons": [83]}, {"buttons": [83]}, {"buttons": [83]}, {"buttons": [83]}, {"buttons": [83]}, {"buttons": [83]}, {"buttons": [83]}, {"buttons": [83]}, {"buttons": [83]},
13+
{}, {}, {}, {}, {}, {}, {}, {}, {}, {}
14+
]
15+
}

‎flashdreams/flashdreams/recipes/taehv/__init__.py‎

Lines changed: 133 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -13,7 +13,7 @@
1313
# See the License for the specific language governing permissions and
1414
# limitations under the License.
1515

16-
"""TAEHV video decoder."""
16+
"""TAEHV video decoder and Hunyuan Video 1.5 codec configs."""
1717

1818
from __future__ import annotations
1919

@@ -26,16 +26,18 @@
2626
from torch import Tensor
2727

2828
from flashdreams.infra.decoder import DecoderConfig, StreamingVideoDecoder
29+
from flashdreams.infra.encoder import EncoderConfig, StreamingVideoEncoder
2930
from flashdreams.recipes.taehv.checkpoint import (
3031
StateDictTransform,
3132
compose,
3233
legacy_to_blocks_keys,
3334
truncate_oversize_tgrow_weights,
3435
)
35-
from flashdreams.recipes.taehv.impl import TAEHV, TAEHVCache
36+
from flashdreams.recipes.taehv.impl import TAEHV, TAEHVCache, TAEHVEncoderCache
3637

3738
AVAILABLE_TAEHV_CHECKPOINT_PATHS = {
3839
"lighttae": "https://huggingface.co/lightx2v/Autoencoders/resolve/main/lighttaew2_1.pth",
40+
"hy1_5": "https://huggingface.co/Overworld-Models/taehv1_5/resolve/main/taehv1_5.pth",
3941
}
4042
"""Checkpoint paths for the TAEHV decoder."""
4143

@@ -230,9 +232,136 @@ def get_input_temporal_size(
230232
return output_temporal_size // r
231233

232234

233-
if __name__ == "__main__":
234-
import tyro
235+
@dataclass(kw_only=True)
236+
class Hy15TAEHVDecoderConfig(DecoderConfig):
237+
"""Config for the Hunyuan Video 1.5 TAEHV decoder."""
238+
239+
_target: Annotated[type, tyro.conf.Suppress] = field(
240+
default_factory=lambda: Hy15TAEHVDecoder
241+
)
242+
243+
checkpoint_path: str = AVAILABLE_TAEHV_CHECKPOINT_PATHS["hy1_5"]
244+
"""Path or URL for the Hunyuan Video 1.5 TAEHV checkpoint."""
245+
246+
state_dict_transform: StateDictTransform | None = legacy_to_blocks_keys
247+
"""Pre-load state-dict remap from the published flat key layout."""
248+
249+
dtype: torch.dtype = torch.bfloat16
250+
"""Network parameter / activation dtype."""
251+
252+
use_cuda_graph: bool = True
253+
"""Wrap the decoder forward in a CUDA graph for replay."""
254+
255+
use_compile: bool = True
256+
"""``torch.compile(mode="max-autotune-no-cudagraphs")``."""
257+
258+
259+
class Hy15TAEHVDecoder(TeahvVAEDecoder):
260+
"""Hunyuan Video 1.5 TAEHV decoder with raw 32-channel latents."""
261+
262+
TEMPORAL_COMPRESSION_RATIO = 4
263+
SPATIAL_COMPRESSION_RATIO = 16
264+
265+
def __init__(self, config: Hy15TAEHVDecoderConfig) -> None:
266+
StreamingVideoDecoder.__init__(self, config)
267+
self.config: Hy15TAEHVDecoderConfig = config
268+
self.need_scaled = False
269+
self.taehv = TAEHV(
270+
checkpoint_path=config.checkpoint_path,
271+
model_type="hy1_5",
272+
use_cuda_graph=config.use_cuda_graph,
273+
use_compile=config.use_compile,
274+
state_dict_transform=config.state_dict_transform,
275+
).to(dtype=config.dtype)
276+
277+
278+
@dataclass(kw_only=True)
279+
class Hy15TAEHVEncoderConfig(EncoderConfig):
280+
"""Config for the Hunyuan Video 1.5 TAEHV encoder."""
281+
282+
_target: Annotated[type, tyro.conf.Suppress] = field(
283+
default_factory=lambda: Hy15TAEHVEncoder
284+
)
285+
286+
checkpoint_path: str = AVAILABLE_TAEHV_CHECKPOINT_PATHS["hy1_5"]
287+
"""Path or URL for the Hunyuan Video 1.5 TAEHV checkpoint."""
288+
289+
state_dict_transform: StateDictTransform | None = legacy_to_blocks_keys
290+
"""Pre-load state-dict remap from the published flat key layout."""
291+
292+
dtype: torch.dtype = torch.bfloat16
293+
"""Network parameter / activation dtype."""
294+
295+
296+
class Hy15TAEHVEncoder(StreamingVideoEncoder[TAEHVEncoderCache]):
297+
"""Hunyuan Video 1.5 TAEHV pixel-video encoder."""
298+
299+
TEMPORAL_COMPRESSION_RATIO = 4
300+
SPATIAL_COMPRESSION_RATIO = 16
235301

302+
def __init__(self, config: Hy15TAEHVEncoderConfig) -> None:
303+
super().__init__(config)
304+
self.config: Hy15TAEHVEncoderConfig = config
305+
self.taehv = TAEHV(
306+
checkpoint_path=config.checkpoint_path,
307+
model_type="hy1_5",
308+
enable_encoder=True,
309+
use_cuda_graph=False,
310+
use_compile=False,
311+
state_dict_transform=config.state_dict_transform,
312+
).to(dtype=config.dtype)
313+
314+
def initialize_autoregressive_cache(self) -> TAEHVEncoderCache:
315+
"""Return an empty causal encoder cache."""
316+
return self.taehv.prepare_encoder_cache()
317+
318+
@torch.no_grad()
319+
def forward(
320+
self,
321+
input: Tensor,
322+
autoregressive_index: int = 0,
323+
cache: TAEHVEncoderCache | None = None,
324+
) -> Tensor:
325+
"""Encode frames in ``[-1, 1]`` to raw Hunyuan Video 1.5 latents."""
326+
if cache is None:
327+
cache = self.initialize_autoregressive_cache()
328+
assert input.ndim >= 4, "Expected input to have shape [..., T, C, H, W]"
329+
330+
*batch_shape, T, C, H, W = input.shape
331+
batch_size = math.prod(batch_shape)
332+
x = input.reshape(batch_size, T, C, H, W).add(1).mul_(0.5)
333+
z = self.taehv.encode(x, cache=cache)
334+
return z.reshape(*batch_shape, *z.shape[1:])
335+
336+
@property
337+
def temporal_compression_ratio(self) -> int:
338+
"""Pixel frames / latent frames for complete causal groups."""
339+
return self.TEMPORAL_COMPRESSION_RATIO
340+
341+
@property
342+
def spatial_compression_ratio(self) -> int:
343+
"""Pixel side / latent side."""
344+
return self.SPATIAL_COMPRESSION_RATIO
345+
346+
def get_output_temporal_size(
347+
self, autoregressive_index: int, input_temporal_size: int
348+
) -> int:
349+
"""Return latent count emitted from complete input frame groups."""
350+
r = self.temporal_compression_ratio
351+
assert input_temporal_size % r == 0, (
352+
f"Hy15 TAEHV encoder input_temporal_size={input_temporal_size} must be "
353+
f"divisible by temporal_compression_ratio={r}."
354+
)
355+
return input_temporal_size // r
356+
357+
def get_input_temporal_size(
358+
self, autoregressive_index: int, output_temporal_size: int
359+
) -> int:
360+
"""Return pixel frame count needed for ``output_temporal_size`` latents."""
361+
return output_temporal_size * self.temporal_compression_ratio
362+
363+
364+
if __name__ == "__main__":
236365
config = tyro.cli(TeahvVAEDecoderConfig)
237366
model = config.setup()
238367
print(model)

‎flashdreams/flashdreams/recipes/taehv/checkpoint.py‎

Lines changed: 14 additions & 14 deletions
Original file line numberDiff line numberDiff line change
@@ -36,22 +36,22 @@
3636
def legacy_to_blocks_keys(
3737
sd: Mapping[str, torch.Tensor],
3838
) -> dict[str, torch.Tensor]:
39-
"""Re-key legacy ``decoder.<i>.*`` weights to ``decoder.blocks.<i>.*``.
39+
"""Re-key flat encoder and decoder weights to their ``blocks`` layouts.
4040
41-
The current :class:`~flashdreams.recipes.taehv.impl.Decoder` wraps
42-
its ``Sequential`` in a ``blocks`` attribute, so older checkpoints
43-
whose keys flatten to ``decoder.<idx>.*`` need rewriting to line up.
44-
Keys already under ``decoder.blocks.`` (and keys outside the
45-
``decoder.`` subtree) pass through unchanged.
41+
The current :class:`~flashdreams.recipes.taehv.impl.Encoder` and
42+
:class:`~flashdreams.recipes.taehv.impl.Decoder` each wrap their
43+
``Sequential`` in a ``blocks`` attribute. Published checkpoints use
44+
flat ``encoder.<idx>.*`` and ``decoder.<idx>.*`` keys. Keys already
45+
under ``*.blocks.`` (and unrelated keys) pass through unchanged.
4646
"""
47-
return {
48-
(
49-
k.replace("decoder.", "decoder.blocks.", 1)
50-
if k.startswith("decoder.") and not k.startswith("decoder.blocks.")
51-
else k
52-
): v
53-
for k, v in sd.items()
54-
}
47+
out: dict[str, torch.Tensor] = {}
48+
for key, value in sd.items():
49+
for prefix in ("encoder.", "decoder."):
50+
if key.startswith(prefix) and not key.startswith(f"{prefix}blocks."):
51+
key = key.replace(prefix, f"{prefix}blocks.", 1)
52+
break
53+
out[key] = value
54+
return out
5555

5656

5757
def truncate_oversize_tgrow_weights(

0 commit comments

Comments
 (0)