[NPU][diffusion] add selectable parallel VAE decode strategies (#23248)

Co-authored-by: 高鑫 <gaoxin@gaoxindeMacBook-Pro.local>
Co-authored-by: ronnie_zheng <zl19940307@163.com>
Co-authored-by: Cursor <cursoragent@cursor.com>
This commit is contained in:
GXIN
2026-05-07 21:37:03 +03:00
committed by GitHub
co-authored by 高鑫 ronnie_zheng Cursor
parent 80a6014243
commit 90a618e37b
5 changed files with 428 additions and 292 deletions
@@ -41,6 +41,8 @@ class VAEConfig(ModelConfig):
use_temporal_tiling: bool = True use_temporal_tiling: bool = True
use_parallel_tiling: bool = True use_parallel_tiling: bool = True
use_temporal_scaling_frames: bool = True use_temporal_scaling_frames: bool = True
use_parallel_decode: bool = False
parallel_decode_mode: str = "tiled"
def __post_init__(self): def __post_init__(self):
self.blend_num_frames = ( self.blend_num_frames = (
@@ -137,6 +139,20 @@ class VAEConfig(ModelConfig):
default=VAEConfig.use_parallel_tiling, default=VAEConfig.use_parallel_tiling,
help="Whether to use parallel tiling for VAE", help="Whether to use parallel tiling for VAE",
) )
parser.add_argument(
f"--{prefix}.use-parallel-decode",
action=StoreBoolean,
dest=f"{prefix.replace('-', '_')}.use_parallel_decode",
default=VAEConfig.use_parallel_decode,
help="Whether to use parallel decode for VAE",
)
parser.add_argument(
f"--{prefix}.parallel-decode-mode",
choices=("tiled", "patch", "auto"),
dest=f"{prefix.replace('-', '_')}.parallel_decode_mode",
default=VAEConfig.parallel_decode_mode,
help="Parallel decode mode for VAE",
)
return parser return parser
@@ -0,0 +1,159 @@
# SPDX-License-Identifier: Apache-2.0
from typing import Optional, Tuple
import torch
import torch.nn.functional as F
from sglang.multimodal_gen.runtime.layers.custom_op import CustomOp
from sglang.multimodal_gen.runtime.platforms import current_platform
_is_cuda = current_platform.is_cuda()
if _is_cuda:
from sglang.jit_kernel.diffusion.triton.scale_shift import (
fuse_layernorm_scale_shift_gate_select01_kernel,
fuse_residual_layernorm_scale_shift_gate_select01_kernel,
)
@CustomOp.register("fuse_layernorm_scale_shift_gate_select01")
class FusedLayerNormScaleShiftGateSelect01(CustomOp):
"""Fused layernorm + scale/shift + gate with binary index selection.
CUDA path uses a Triton kernel; other platforms fall back to PyTorch ops.
"""
def forward_cuda(
self,
x: torch.Tensor,
weight: Optional[torch.Tensor],
bias: Optional[torch.Tensor],
scale0: torch.Tensor,
shift0: torch.Tensor,
gate0: torch.Tensor,
scale1: torch.Tensor,
shift1: torch.Tensor,
gate1: torch.Tensor,
index: torch.Tensor,
eps: float,
) -> Tuple[torch.Tensor, torch.Tensor]:
if not x.is_contiguous():
x = x.contiguous()
if not index.is_contiguous():
index = index.contiguous()
return fuse_layernorm_scale_shift_gate_select01_kernel(
x,
weight=weight,
bias=bias,
scale0=scale0.contiguous(),
shift0=shift0.contiguous(),
gate0=gate0.contiguous(),
scale1=scale1.contiguous(),
shift1=shift1.contiguous(),
gate1=gate1.contiguous(),
index=index,
eps=eps,
)
def forward_hip(self, *args, **kwargs):
return self.forward_native(*args, **kwargs)
def forward_native(
self,
x: torch.Tensor,
weight: Optional[torch.Tensor],
bias: Optional[torch.Tensor],
scale0: torch.Tensor,
shift0: torch.Tensor,
gate0: torch.Tensor,
scale1: torch.Tensor,
shift1: torch.Tensor,
gate1: torch.Tensor,
index: torch.Tensor,
eps: float,
) -> Tuple[torch.Tensor, torch.Tensor]:
idx = index.to(dtype=torch.bool).unsqueeze(-1)
shift = torch.where(idx, shift1.unsqueeze(1), shift0.unsqueeze(1))
scale = torch.where(idx, scale1.unsqueeze(1), scale0.unsqueeze(1))
gate = torch.where(idx, gate1.unsqueeze(1), gate0.unsqueeze(1))
x = F.layer_norm(x, (x.shape[-1],), weight=weight, bias=bias, eps=eps)
x = x * (1 + scale) + shift
return x, gate
@CustomOp.register("fuse_residual_layernorm_scale_shift_gate_select01")
class FusedResidualLayerNormScaleShiftGateSelect01(CustomOp):
"""Fused residual + layernorm + scale/shift + gate with binary index selection.
CUDA path uses a Triton kernel; other platforms fall back to PyTorch ops.
"""
def forward_cuda(
self,
x: torch.Tensor,
residual: torch.Tensor,
residual_gate: torch.Tensor,
weight: Optional[torch.Tensor],
bias: Optional[torch.Tensor],
scale0: torch.Tensor,
shift0: torch.Tensor,
gate0: torch.Tensor,
scale1: torch.Tensor,
shift1: torch.Tensor,
gate1: torch.Tensor,
index: torch.Tensor,
eps: float,
) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
if not x.is_contiguous():
x = x.contiguous()
if not index.is_contiguous():
index = index.contiguous()
if not residual.is_contiguous():
residual = residual.contiguous()
if not residual_gate.is_contiguous():
residual_gate = residual_gate.contiguous()
return fuse_residual_layernorm_scale_shift_gate_select01_kernel(
x,
residual=residual,
residual_gate=residual_gate,
weight=weight,
bias=bias,
scale0=scale0.contiguous(),
shift0=shift0.contiguous(),
gate0=gate0.contiguous(),
scale1=scale1.contiguous(),
shift1=shift1.contiguous(),
gate1=gate1.contiguous(),
index=index,
eps=eps,
)
def forward_hip(self, *args, **kwargs):
return self.forward_native(*args, **kwargs)
def forward_native(
self,
x: torch.Tensor,
residual: torch.Tensor,
residual_gate: torch.Tensor,
weight: Optional[torch.Tensor],
bias: Optional[torch.Tensor],
scale0: torch.Tensor,
shift0: torch.Tensor,
gate0: torch.Tensor,
scale1: torch.Tensor,
shift1: torch.Tensor,
gate1: torch.Tensor,
index: torch.Tensor,
eps: float,
) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
idx = index.to(dtype=torch.bool).unsqueeze(-1)
shift = torch.where(idx, shift1.unsqueeze(1), shift0.unsqueeze(1))
scale = torch.where(idx, scale1.unsqueeze(1), scale0.unsqueeze(1))
gate = torch.where(idx, gate1.unsqueeze(1), gate0.unsqueeze(1))
residual_out = residual_gate * x + residual
x = F.layer_norm(
residual_out, (residual_out.shape[-1],), weight=weight, bias=bias, eps=eps
)
x = x * (1 + scale) + shift
return x, residual_out, gate
@@ -14,10 +14,6 @@ from diffusers.models.embeddings import TimestepEmbedding, Timesteps
from diffusers.models.modeling_outputs import Transformer2DModelOutput from diffusers.models.modeling_outputs import Transformer2DModelOutput
from diffusers.models.normalization import AdaLayerNormContinuous from diffusers.models.normalization import AdaLayerNormContinuous
from sglang.jit_kernel.diffusion.triton.scale_shift import (
fuse_layernorm_scale_shift_gate_select01_kernel,
fuse_residual_layernorm_scale_shift_gate_select01_kernel,
)
from sglang.multimodal_gen.configs.models.dits.qwenimage import QwenImageDitConfig from sglang.multimodal_gen.configs.models.dits.qwenimage import QwenImageDitConfig
from sglang.multimodal_gen.runtime.distributed import get_local_torch_device from sglang.multimodal_gen.runtime.distributed import get_local_torch_device
from sglang.multimodal_gen.runtime.distributed.parallel_state import ( from sglang.multimodal_gen.runtime.distributed.parallel_state import (
@@ -25,6 +21,10 @@ from sglang.multimodal_gen.runtime.distributed.parallel_state import (
) )
from sglang.multimodal_gen.runtime.layers.attention import USPAttention from sglang.multimodal_gen.runtime.layers.attention import USPAttention
from sglang.multimodal_gen.runtime.layers.elementwise import MulAdd from sglang.multimodal_gen.runtime.layers.elementwise import MulAdd
from sglang.multimodal_gen.runtime.layers.fused_scale_shift_gate import (
FusedLayerNormScaleShiftGateSelect01,
FusedResidualLayerNormScaleShiftGateSelect01,
)
from sglang.multimodal_gen.runtime.layers.layernorm import ( from sglang.multimodal_gen.runtime.layers.layernorm import (
LayerNormScaleShift, LayerNormScaleShift,
RMSNorm, RMSNorm,
@@ -47,15 +47,11 @@ from sglang.multimodal_gen.runtime.layers.rotary_embedding import (
) )
from sglang.multimodal_gen.runtime.managers.layerwise_offload import OffloadableDiTMixin from sglang.multimodal_gen.runtime.managers.layerwise_offload import OffloadableDiTMixin
from sglang.multimodal_gen.runtime.models.dits.base import CachableDiT from sglang.multimodal_gen.runtime.models.dits.base import CachableDiT
from sglang.multimodal_gen.runtime.platforms import ( from sglang.multimodal_gen.runtime.platforms import AttentionBackendEnum
AttentionBackendEnum,
current_platform,
)
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
logger = init_logger(__name__) # pylint: disable=invalid-name logger = init_logger(__name__) # pylint: disable=invalid-name
try: try:
from nunchaku.models.attention import NunchakuFeedForward # type: ignore[import] from nunchaku.models.attention import NunchakuFeedForward # type: ignore[import]
except Exception: except Exception:
@@ -879,6 +875,10 @@ class QwenImageTransformerBlock(nn.Module):
) )
# Utils # Utils
self.fuse_mul_add = MulAdd() self.fuse_mul_add = MulAdd()
self.fused_ln_ss_gate_select01 = FusedLayerNormScaleShiftGateSelect01()
self.fused_res_ln_ss_gate_select01 = (
FusedResidualLayerNormScaleShiftGateSelect01()
)
nunchaku_enabled = ( nunchaku_enabled = (
quant_config is not None quant_config is not None
@@ -941,9 +941,6 @@ class QwenImageTransformerBlock(nn.Module):
shift, scale, gate = mod_params.chunk(3, dim=-1) shift, scale, gate = mod_params.chunk(3, dim=-1)
if index is not None: if index is not None:
# ROCm currently fails to compile the select01 Triton kernel, so
# keep using the torch.where fallback there.
if x.is_cuda and not current_platform.is_hip():
actual_batch = x.shape[0] actual_batch = x.shape[0]
shift0, shift1 = ( shift0, shift1 = (
shift[:actual_batch], shift[:actual_batch],
@@ -957,70 +954,38 @@ class QwenImageTransformerBlock(nn.Module):
gate[:actual_batch], gate[:actual_batch],
gate[actual_batch : 2 * actual_batch], gate[actual_batch : 2 * actual_batch],
) )
if not x.is_contiguous():
x = x.contiguous()
if not index.is_contiguous():
index = index.contiguous()
if is_scale_residual: if is_scale_residual:
if not residual_x.is_contiguous(): x, residual_out, gate_result = self.fused_res_ln_ss_gate_select01(
residual_x = residual_x.contiguous()
if not gate_x.is_contiguous():
gate_x = gate_x.contiguous()
x, residual_out, gate_result = (
fuse_residual_layernorm_scale_shift_gate_select01_kernel(
x, x,
residual=residual_x, residual_x,
residual_gate=gate_x, gate_x,
weight=getattr(norm_module.norm, "weight", None), getattr(norm_module.norm, "weight", None),
bias=getattr(norm_module.norm, "bias", None), getattr(norm_module.norm, "bias", None),
scale0=scale0.contiguous(), scale0,
shift0=shift0.contiguous(), shift0,
gate0=gate0.contiguous(), gate0,
scale1=scale1.contiguous(), scale1,
shift1=shift1.contiguous(), shift1,
gate1=gate1.contiguous(), gate1,
index=index, index,
eps=norm_module.eps, norm_module.eps,
)
) )
return x, residual_out, gate_result return x, residual_out, gate_result
else: else:
x, gate_result = fuse_layernorm_scale_shift_gate_select01_kernel( x, gate_result = self.fused_ln_ss_gate_select01(
x, x,
weight=getattr(norm_module.norm, "weight", None), getattr(norm_module.norm, "weight", None),
bias=getattr(norm_module.norm, "bias", None), getattr(norm_module.norm, "bias", None),
scale0=scale0.contiguous(), scale0,
shift0=shift0.contiguous(), shift0,
gate0=gate0.contiguous(), gate0,
scale1=scale1.contiguous(), scale1,
shift1=shift1.contiguous(), shift1,
gate1=gate1.contiguous(), gate1,
index=index, index,
eps=norm_module.eps, norm_module.eps,
) )
return x, gate_result return x, gate_result
else:
actual_batch = x.shape[0]
shift0, shift1 = (
shift[:actual_batch],
shift[actual_batch : 2 * actual_batch],
)
scale0, scale1 = (
scale[:actual_batch],
scale[actual_batch : 2 * actual_batch],
)
gate0, gate1 = (
gate[:actual_batch],
gate[actual_batch : 2 * actual_batch],
)
index = index.to(dtype=torch.bool).unsqueeze(-1)
shift_result = torch.where(
index, shift1.unsqueeze(1), shift0.unsqueeze(1)
)
scale_result = torch.where(
index, scale1.unsqueeze(1), scale0.unsqueeze(1)
)
gate_result = torch.where(index, gate1.unsqueeze(1), gate0.unsqueeze(1))
else: else:
shift_result = shift.unsqueeze(1) shift_result = shift.unsqueeze(1)
scale_result = scale.unsqueeze(1) scale_result = scale.unsqueeze(1)
@@ -3,7 +3,6 @@
from typing import Optional, Tuple, Union from typing import Optional, Tuple, Union
import torch import torch
import torch.distributed as dist
import torch.nn as nn import torch.nn as nn
import torch.nn.functional as F import torch.nn.functional as F
from diffusers.models.activations import get_activation from diffusers.models.activations import get_activation
@@ -967,6 +966,25 @@ class AutoencoderKLQwenImage(ParallelTiledVAE):
if self.use_parallel_decode and get_sp_world_size() > 1: if self.use_parallel_decode and get_sp_world_size() > 1:
num_frame = z.shape[2] num_frame = z.shape[2]
num_sample_frames = (num_frame - 1) * self.temporal_compression_ratio + 1 num_sample_frames = (num_frame - 1) * self.temporal_compression_ratio + 1
tile_latent_min_height = (
self.tile_sample_min_height // self.spatial_compression_ratio
)
tile_latent_min_width = (
self.tile_sample_min_width // self.spatial_compression_ratio
)
mode = self.parallel_decode_mode
if mode == "auto":
if (
z.shape[-2] > tile_latent_min_height
or z.shape[-1] > tile_latent_min_width
):
mode = "tiled"
else:
mode = "patch"
if mode == "patch":
decoded = super().parallel_patch_decode(z)[:, :, :num_sample_frames]
else:
decoded = super().parallel_tiled_decode(z)[:, :, :num_sample_frames] decoded = super().parallel_tiled_decode(z)[:, :, :num_sample_frames]
return DecoderOutput(sample=decoded) return DecoderOutput(sample=decoded)
@@ -1056,81 +1074,6 @@ class AutoencoderKLQwenImage(ParallelTiledVAE):
) )
return b return b
def _process_parallel_tiled_outputs(
self,
results: torch.Tensor,
local_dim_metadata: list[torch.Size],
z: torch.Tensor,
world_size: int,
rank: int,
num_t_tiles: int,
num_h_tiles: int,
num_w_tiles: int,
total_spatial_tiles: int,
blend_height: int,
blend_width: int,
) -> torch.Tensor:
local_size = torch.tensor(
[results.size(0)], device=results.device, dtype=torch.int64
)
if rank == 0:
gathered_sizes = [
torch.zeros(1, device=results.device, dtype=torch.int64)
for _ in range(world_size)
]
else:
gathered_sizes = None
dist.gather(local_size, gather_list=gathered_sizes, dst=0)
max_size = 0
if rank == 0:
max_size = max(size.item() for size in gathered_sizes)
max_size_tensor = torch.tensor(
[max_size], device=results.device, dtype=torch.int64
)
dist.broadcast(max_size_tensor, src=0)
max_size = int(max_size_tensor.item())
padded_results = torch.zeros(
max_size, device=results.device, dtype=results.dtype
)
padded_results[: results.size(0)] = results
gathered_dim_metadata = [None] * world_size
dist.all_gather_object(gathered_dim_metadata, local_dim_metadata)
if rank == 0:
gathered_results = [
torch.empty_like(padded_results) for _ in range(world_size)
]
else:
gathered_results = None
dist.gather(padded_results, gather_list=gathered_results, dst=0)
if rank == 0:
gathered_results = torch.stack(gathered_results, dim=0).contiguous()
dec = super()._merge_parallel_tiled_results(
gathered_results,
gathered_dim_metadata,
num_t_tiles,
num_h_tiles,
num_w_tiles,
total_spatial_tiles,
blend_height,
blend_width,
)
shape_tensor = torch.tensor(dec.shape, device=dec.device, dtype=torch.int64)
else:
dec = None
shape_tensor = torch.zeros(5, device=z.device, dtype=torch.int64)
dist.broadcast(shape_tensor, src=0)
if rank != 0:
dec = z.new_empty(tuple(shape_tensor.tolist()))
dist.broadcast(dec, src=0)
return dec
def tiled_encode(self, x: torch.Tensor) -> AutoencoderKLOutput: def tiled_encode(self, x: torch.Tensor) -> AutoencoderKLOutput:
r"""Encode a batch of images using a tiled encoder. r"""Encode a batch of images using a tiled encoder.
@@ -3,8 +3,7 @@
# SPDX-License-Identifier: Apache-2.0 # SPDX-License-Identifier: Apache-2.0
from abc import ABC, abstractmethod from abc import ABC, abstractmethod
from collections.abc import Iterator from math import isqrt, prod
from math import prod
from typing import Optional, cast from typing import Optional, cast
import numpy as np import numpy as np
@@ -32,6 +31,8 @@ class ParallelTiledVAE(ABC, nn.Module):
use_tiling: bool use_tiling: bool
use_temporal_tiling: bool use_temporal_tiling: bool
use_parallel_tiling: bool use_parallel_tiling: bool
use_parallel_decode: bool
parallel_decode_mode: str
def __init__(self, config: VAEConfig, **kwargs) -> None: def __init__(self, config: VAEConfig, **kwargs) -> None:
super().__init__() super().__init__()
@@ -46,6 +47,8 @@ class ParallelTiledVAE(ABC, nn.Module):
self.use_tiling = config.use_tiling self.use_tiling = config.use_tiling
self.use_temporal_tiling = config.use_temporal_tiling self.use_temporal_tiling = config.use_temporal_tiling
self.use_parallel_tiling = config.use_parallel_tiling self.use_parallel_tiling = config.use_parallel_tiling
self.use_parallel_decode = config.use_parallel_decode
self.parallel_decode_mode = config.parallel_decode_mode
@property @property
def device(self): def device(self):
@@ -203,120 +206,6 @@ class ParallelTiledVAE(ABC, nn.Module):
tile_latent_stride_width, tile_latent_stride_width,
) )
def _parallel_data_generator(
self, gathered_results, gathered_dim_metadata
) -> Iterator[tuple[torch.Tensor, int]]:
global_idx = 0
for i, per_rank_metadata in enumerate(gathered_dim_metadata):
_start_shape = 0
for shape in per_rank_metadata:
mul_shape = prod(shape)
yield (
gathered_results[
i, _start_shape : _start_shape + mul_shape
].reshape(shape),
global_idx,
)
_start_shape += mul_shape
global_idx += 1
def _merge_parallel_tiled_results(
self,
gathered_results: torch.Tensor,
gathered_dim_metadata: list[list[torch.Size]],
num_t_tiles: int,
num_h_tiles: int,
num_w_tiles: int,
total_spatial_tiles: int,
blend_height: int,
blend_width: int,
) -> torch.Tensor:
data: list = [
[[[] for _ in range(num_w_tiles)] for _ in range(num_h_tiles)]
for _ in range(num_t_tiles)
]
for current_data, global_idx in self._parallel_data_generator(
gathered_results, gathered_dim_metadata
):
t_idx = global_idx // total_spatial_tiles
spatial_idx = global_idx % total_spatial_tiles
h_idx = spatial_idx // num_w_tiles
w_idx = spatial_idx % num_w_tiles
data[t_idx][h_idx][w_idx] = current_data
result_slices = []
last_slice_data = None
for i, tem_data in enumerate(data):
slice_data = self._merge_spatial_tiles(
tem_data,
blend_height,
blend_width,
self.tile_sample_stride_height,
self.tile_sample_stride_width,
)
if i > 0:
slice_data = self.blend_t(
last_slice_data, slice_data, self.blend_num_frames
)
result_slices.append(
slice_data[:, :, : self.tile_sample_stride_num_frames, :, :]
)
else:
result_slices.append(
slice_data[:, :, : self.tile_sample_stride_num_frames + 1, :, :]
)
last_slice_data = slice_data
return torch.cat(result_slices, dim=2)
def _process_parallel_tiled_outputs(
self,
results: torch.Tensor,
local_dim_metadata: list[torch.Size],
z: torch.Tensor,
world_size: int,
rank: int,
num_t_tiles: int,
num_h_tiles: int,
num_w_tiles: int,
total_spatial_tiles: int,
blend_height: int,
blend_width: int,
) -> torch.Tensor:
local_size = torch.tensor(
[results.size(0)], device=results.device, dtype=torch.int64
)
all_sizes = [
torch.zeros(1, device=results.device, dtype=torch.int64)
for _ in range(world_size)
]
dist.all_gather(all_sizes, local_size)
max_size = max(size.item() for size in all_sizes)
padded_results = torch.zeros(
max_size, device=results.device, dtype=results.dtype
)
padded_results[: results.size(0)] = results
gathered_dim_metadata = [None] * world_size
gathered_results = (
torch.zeros_like(padded_results)
.repeat(world_size, *[1] * len(padded_results.shape))
.contiguous()
)
dist.all_gather_into_tensor(gathered_results, padded_results)
dist.all_gather_object(gathered_dim_metadata, local_dim_metadata)
gathered_dim_metadata = cast(list[list[torch.Size]], gathered_dim_metadata)
return self._merge_parallel_tiled_results(
gathered_results,
gathered_dim_metadata,
num_t_tiles,
num_h_tiles,
num_w_tiles,
total_spatial_tiles,
blend_height,
blend_width,
)
def parallel_tiled_decode(self, z: torch.FloatTensor) -> torch.FloatTensor: def parallel_tiled_decode(self, z: torch.FloatTensor) -> torch.FloatTensor:
""" """
Parallel version of tiled_decode that distributes both temporal and spatial computation across GPUs Parallel version of tiled_decode that distributes both temporal and spatial computation across GPUs
@@ -354,7 +243,6 @@ class ParallelTiledVAE(ABC, nn.Module):
num_w_tiles = (W + tile_latent_stride_width - 1) // tile_latent_stride_width num_w_tiles = (W + tile_latent_stride_width - 1) // tile_latent_stride_width
total_spatial_tiles = num_h_tiles * num_w_tiles total_spatial_tiles = num_h_tiles * num_w_tiles
total_tiles = num_t_tiles * total_spatial_tiles total_tiles = num_t_tiles * total_spatial_tiles
tiles_per_rank = (total_tiles + world_size - 1) // world_size tiles_per_rank = (total_tiles + world_size - 1) // world_size
start_tile_idx = rank * tiles_per_rank start_tile_idx = rank * tiles_per_rank
end_tile_idx = min((rank + 1) * tiles_per_rank, total_tiles) end_tile_idx = min((rank + 1) * tiles_per_rank, total_tiles)
@@ -390,19 +278,184 @@ class ParallelTiledVAE(ABC, nn.Module):
results = z.new_empty((0,), dtype=z.dtype) results = z.new_empty((0,), dtype=z.dtype)
del local_results del local_results
dec = self._process_parallel_tiled_outputs( local_size = torch.tensor(
results, [results.size(0)], device=results.device, dtype=torch.int64
local_dim_metadata, )
z, all_sizes = [
world_size, torch.zeros(1, device=results.device, dtype=torch.int64)
rank, for _ in range(world_size)
num_t_tiles, ]
num_h_tiles, dist.all_gather(all_sizes, local_size)
num_w_tiles, max_size = max(size.item() for size in all_sizes)
total_spatial_tiles,
padded_results = torch.zeros(
max_size, device=results.device, dtype=results.dtype
)
padded_results[: results.size(0)] = results
gathered_dim_metadata = [None] * world_size
gathered_results = (
torch.zeros_like(padded_results)
.repeat(world_size, *[1] * len(padded_results.shape))
.contiguous()
)
dist.all_gather_into_tensor(gathered_results, padded_results)
dist.all_gather_object(gathered_dim_metadata, local_dim_metadata)
gathered_dim_metadata = cast(list[list[torch.Size]], gathered_dim_metadata)
data: list = [
[[[] for _ in range(num_w_tiles)] for _ in range(num_h_tiles)]
for _ in range(num_t_tiles)
]
global_idx = 0
for i, per_rank_metadata in enumerate(gathered_dim_metadata):
start_shape = 0
for shape in per_rank_metadata:
mul_shape = prod(shape)
current_data = gathered_results[
i, start_shape : start_shape + mul_shape
].reshape(shape)
t_idx = global_idx // total_spatial_tiles
spatial_idx = global_idx % total_spatial_tiles
h_idx = spatial_idx // num_w_tiles
w_idx = spatial_idx % num_w_tiles
data[t_idx][h_idx][w_idx] = current_data
start_shape += mul_shape
global_idx += 1
result_slices = []
last_slice_data = None
for i, tem_data in enumerate(data):
slice_data = self._merge_spatial_tiles(
tem_data,
blend_height, blend_height,
blend_width, blend_width,
self.tile_sample_stride_height,
self.tile_sample_stride_width,
) )
if i > 0:
slice_data = self.blend_t(
last_slice_data, slice_data, self.blend_num_frames
)
result_slices.append(
slice_data[:, :, : self.tile_sample_stride_num_frames, :, :]
)
else:
result_slices.append(
slice_data[:, :, : self.tile_sample_stride_num_frames + 1, :, :]
)
last_slice_data = slice_data
return torch.cat(result_slices, dim=2)
def parallel_patch_decode(self, z: torch.FloatTensor) -> torch.FloatTensor:
world_size, rank = get_sp_world_size(), get_sp_parallel_rank()
if world_size <= 1:
return self._decode(z)
tile_latent_min_height = (
self.tile_sample_min_height // self.spatial_compression_ratio
)
tile_latent_min_width = (
self.tile_sample_min_width // self.spatial_compression_ratio
)
tile_latent_stride_height = (
self.tile_sample_stride_height // self.spatial_compression_ratio
)
tile_latent_stride_width = (
self.tile_sample_stride_width // self.spatial_compression_ratio
)
overlap_h = max(0, tile_latent_min_height - tile_latent_stride_height)
overlap_w = max(0, tile_latent_min_width - tile_latent_stride_width)
halo_h = overlap_h // 2
halo_w = overlap_w // 2
_, _, _, latent_h, latent_w = z.shape
scale = self.spatial_compression_ratio
out_h = latent_h * scale
out_w = latent_w * scale
root = isqrt(world_size)
grid_rows, grid_cols = 1, world_size
for rows in range(root, 0, -1):
if world_size % rows == 0:
grid_rows, grid_cols = rows, world_size // rows
break
patch_id = rank
patch_row = patch_id // grid_cols
patch_col = patch_id % grid_cols
h0 = (patch_row * latent_h) // grid_rows
h1 = ((patch_row + 1) * latent_h) // grid_rows
w0 = (patch_col * latent_w) // grid_cols
w1 = ((patch_col + 1) * latent_w) // grid_cols
ext_h0 = max(0, h0 - halo_h)
ext_h1 = min(latent_h, h1 + halo_h)
ext_w0 = max(0, w0 - halo_w)
ext_w1 = min(latent_w, w1 + halo_w)
local_patch = z[:, :, :, ext_h0:ext_h1, ext_w0:ext_w1]
decoded_patch = self._decode(local_patch)
crop_top = (h0 - ext_h0) * scale
crop_bottom = crop_top + (h1 - h0) * scale
crop_left = (w0 - ext_w0) * scale
crop_right = crop_left + (w1 - w0) * scale
decoded_core = decoded_patch[
:, :, :, crop_top:crop_bottom, crop_left:crop_right
].contiguous()
local_result = decoded_core.reshape(-1)
local_dim_metadata = torch.tensor(
decoded_core.shape, device=z.device, dtype=torch.int64
)
local_position = torch.tensor(
[h0 * scale, h1 * scale, w0 * scale, w1 * scale],
device=z.device,
dtype=torch.int64,
)
gathered_positions = [
torch.empty_like(local_position) for _ in range(world_size)
]
dist.all_gather(gathered_positions, local_position)
local_size = torch.tensor(
[local_result.size(0)], device=z.device, dtype=torch.int64
)
gathered_dim_metadata = [
torch.empty_like(local_dim_metadata) for _ in range(world_size)
]
dist.all_gather(gathered_dim_metadata, local_dim_metadata)
all_sizes = [
torch.zeros(1, device=z.device, dtype=torch.int64)
for _ in range(world_size)
]
dist.all_gather(all_sizes, local_size)
max_size = max(size.item() for size in all_sizes)
padded_results = torch.zeros(max_size, device=z.device, dtype=z.dtype)
padded_results[: local_result.size(0)] = local_result
gathered_results = torch.empty(
(world_size, *padded_results.shape),
device=padded_results.device,
dtype=padded_results.dtype,
)
dist.all_gather_into_tensor(gathered_results, padded_results)
dec = z.new_empty(
(
decoded_core.shape[0],
decoded_core.shape[1],
decoded_core.shape[2],
out_h,
out_w,
)
)
for src_rank, positions in enumerate(gathered_positions):
h_start, h_end, w_start, w_end = [int(x.item()) for x in positions]
shape = tuple(int(x.item()) for x in gathered_dim_metadata[src_rank])
patch = gathered_results[src_rank][: prod(shape)].reshape(shape)
dec[:, :, :, h_start:h_end, w_start:w_end] = patch
return dec return dec
def _merge_spatial_tiles( def _merge_spatial_tiles(