[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:
co-authored by
高鑫
ronnie_zheng
Cursor
parent
80a6014243
commit
90a618e37b
@@ -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(
|
||||||
|
|||||||
Reference in New Issue
Block a user