|
13 | 13 | # See the License for the specific language governing permissions and |
14 | 14 | # limitations under the License. |
15 | 15 |
|
16 | | -"""TAEHV video decoder.""" |
| 16 | +"""TAEHV video decoder and Hunyuan Video 1.5 codec configs.""" |
17 | 17 |
|
18 | 18 | from __future__ import annotations |
19 | 19 |
|
|
26 | 26 | from torch import Tensor |
27 | 27 |
|
28 | 28 | from flashdreams.infra.decoder import DecoderConfig, StreamingVideoDecoder |
| 29 | +from flashdreams.infra.encoder import EncoderConfig, StreamingVideoEncoder |
29 | 30 | from flashdreams.recipes.taehv.checkpoint import ( |
30 | 31 | StateDictTransform, |
31 | 32 | compose, |
32 | 33 | legacy_to_blocks_keys, |
33 | 34 | truncate_oversize_tgrow_weights, |
34 | 35 | ) |
35 | | -from flashdreams.recipes.taehv.impl import TAEHV, TAEHVCache |
| 36 | +from flashdreams.recipes.taehv.impl import TAEHV, TAEHVCache, TAEHVEncoderCache |
36 | 37 |
|
37 | 38 | AVAILABLE_TAEHV_CHECKPOINT_PATHS = { |
38 | 39 | "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", |
39 | 41 | } |
40 | 42 | """Checkpoint paths for the TAEHV decoder.""" |
41 | 43 |
|
@@ -230,9 +232,136 @@ def get_input_temporal_size( |
230 | 232 | return output_temporal_size // r |
231 | 233 |
|
232 | 234 |
|
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 |
235 | 301 |
|
| 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__": |
236 | 365 | config = tyro.cli(TeahvVAEDecoderConfig) |
237 | 366 | model = config.setup() |
238 | 367 | print(model) |
0 commit comments