FL2VA/video_vae/utils.py
| 1 | # SPDX-License-Identifier: Apache-2.0 |
| 2 | # Module helpers for the MiniMax H3 visual VAE (inference-only bundle). |
| 3 | |
| 4 | |
| 5 | def apply_spatial_parallel(module, enabled, chunk_dim=-1): |
| 6 | from .conv import SpatialParallelConv3d |
| 7 | from .norm import FusedGroupNorm3D, SpatialParallelGroupNorm |
| 8 | |
| 9 | if hasattr(module, "set_spatial_parallel"): |
| 10 | module.set_spatial_parallel(enabled) |
| 11 | for m in module.modules(): |
| 12 | if enabled and isinstance(m, FusedGroupNorm3D): |
| 13 | raise NotImplementedError("FusedGroupNorm3D is incompatible with SP") |
| 14 | if isinstance(m, SpatialParallelGroupNorm): |
| 15 | m.spatial_parallel = enabled |
| 16 | elif isinstance(m, SpatialParallelConv3d): |
| 17 | m.spatial_parallel = enabled |
| 18 | m.chunk_dim = chunk_dim |
| 19 | |