Optimize VAE Decoding Speed and robustness
- Implemented "4D execution mode" in ResnetBlock3D to flatten temporal dimension when convolutions are effectively 2D, reducing reshape overhead. - Updated InflatedCausalConv3d to support direct 4D input processing and use 2D convolution optimization path. - Enhanced InflatedCausalConv3d check_effective_2d to strictly verify stride/dilation/padding compatibility. - Fixed 4D input handling in InflatedCausalConv3d forward pass.
This commit is contained in:
@@ -313,14 +313,7 @@ class ResnetBlock3D(ResnetBlock2D):
|
||||
):
|
||||
hidden_states = input_tensor
|
||||
|
||||
hidden_states = causal_norm_wrapper(self.norm1, hidden_states)
|
||||
hidden_states = retry_on_oom(
|
||||
self.nonlinearity,
|
||||
hidden_states,
|
||||
debug=getattr(self, 'debug', None),
|
||||
operation_name="ResnetBlock3D.nonlinearity"
|
||||
)
|
||||
|
||||
# Handle upsample/downsample first (usually involves 5D operations if temporal)
|
||||
if self.upsample is not None:
|
||||
# upsample_nearest_nhwc fails with large batch sizes.
|
||||
# see https://github.com/huggingface/diffusers/issues/984
|
||||
@@ -333,6 +326,78 @@ class ResnetBlock3D(ResnetBlock2D):
|
||||
input_tensor = self.downsample(input_tensor, memory_state=memory_state)
|
||||
hidden_states = self.downsample(hidden_states, memory_state=memory_state)
|
||||
|
||||
# Check if we can run the main ResNet path in 4D (flattened temporal)
|
||||
# Conditions:
|
||||
# 1. Input is 5D (B, C, T, H, W)
|
||||
# 2. conv1 and conv2 are effectively 2D (spatial only)
|
||||
# 3. conv_shortcut (if present) is effectively 2D
|
||||
# 4. time_embedding_norm is "default" (simple add) - others might need more complex broadcasting logic in 4D
|
||||
|
||||
can_run_4d = (
|
||||
hidden_states.ndim == 5
|
||||
and self.conv1.check_effective_2d()
|
||||
and self.conv2.check_effective_2d()
|
||||
and (self.conv_shortcut is None or self.conv_shortcut.check_effective_2d())
|
||||
and (self.time_embedding_norm == "default" or temb is None)
|
||||
)
|
||||
|
||||
if can_run_4d:
|
||||
B, C, T, H, W = hidden_states.shape
|
||||
|
||||
# Reshape to 4D: (B*T, C, H, W)
|
||||
# Use contiguous to ensure efficient memory layout for conv2d
|
||||
hidden_states = hidden_states.transpose(1, 2).reshape(B * T, C, H, W).contiguous()
|
||||
if input_tensor.ndim == 5:
|
||||
input_tensor = input_tensor.transpose(1, 2).reshape(B * T, C, H, W).contiguous()
|
||||
|
||||
# Prepare temb for 4D
|
||||
temb_4d = None
|
||||
if self.time_emb_proj is not None and temb is not None:
|
||||
if not self.skip_time_act:
|
||||
temb = self.nonlinearity(temb)
|
||||
# temb: (B, C_emb) -> proj -> (B, C_out) -> (B, C_out, 1, 1)
|
||||
# We need (B*T, C_out, 1, 1)
|
||||
temb_proj = self.time_emb_proj(temb)[:, :, None, None]
|
||||
# Expand B to B*T
|
||||
temb_4d = temb_proj.repeat_interleave(T, dim=0)
|
||||
|
||||
# Main path in 4D
|
||||
# causal_norm_wrapper handles 4D input by just calling GroupNorm (no rearrange overhead!)
|
||||
hidden_states = causal_norm_wrapper(self.norm1, hidden_states)
|
||||
hidden_states = self.nonlinearity(hidden_states)
|
||||
|
||||
# InflatedCausalConv3d handles 4D input by dispatching to conv2d (cached check)
|
||||
hidden_states = self.conv1(hidden_states, memory_state=memory_state)
|
||||
|
||||
if temb_4d is not None and self.time_embedding_norm == "default":
|
||||
hidden_states = hidden_states + temb_4d
|
||||
|
||||
hidden_states = causal_norm_wrapper(self.norm2, hidden_states)
|
||||
hidden_states = self.nonlinearity(hidden_states)
|
||||
hidden_states = self.dropout(hidden_states)
|
||||
hidden_states = self.conv2(hidden_states, memory_state=memory_state)
|
||||
|
||||
if self.conv_shortcut is not None:
|
||||
input_tensor = self.conv_shortcut(input_tensor, memory_state=memory_state)
|
||||
|
||||
output_tensor = (input_tensor + hidden_states) / self.output_scale_factor
|
||||
|
||||
# Reshape back to 5D: (B, C, T, H, W)
|
||||
# Output is (B*T, C, H, W)
|
||||
_, C_out, H_out, W_out = output_tensor.shape
|
||||
output_tensor = output_tensor.view(B, T, C_out, H_out, W_out).transpose(1, 2)
|
||||
|
||||
return output_tensor
|
||||
|
||||
# Fallback to standard 5D path
|
||||
hidden_states = causal_norm_wrapper(self.norm1, hidden_states)
|
||||
hidden_states = retry_on_oom(
|
||||
self.nonlinearity,
|
||||
hidden_states,
|
||||
debug=getattr(self, 'debug', None),
|
||||
operation_name="ResnetBlock3D.nonlinearity"
|
||||
)
|
||||
|
||||
hidden_states = self.conv1(hidden_states, memory_state=memory_state)
|
||||
|
||||
if self.time_emb_proj is not None:
|
||||
|
||||
@@ -81,6 +81,63 @@ class InflatedCausalConv3d(Conv3d):
|
||||
def set_memory_device(self, memory_device: _memory_device_t):
|
||||
self.memory_device = memory_device
|
||||
|
||||
def check_effective_2d(self):
|
||||
"""
|
||||
Check if the convolution is effectively 2D (spatial only) AND compatible with 2D optimization path.
|
||||
Returns True if:
|
||||
1. Temporal kernel size is 1 OR weights are tail-inflated (only last temporal slice is non-zero).
|
||||
2. AND stride[0] == 1.
|
||||
3. AND dilation[0] == 1.
|
||||
4. AND padding[0] == 0 (self.padding is modified in __init__, this checks the spatial-only padding).
|
||||
"""
|
||||
if hasattr(self, '_effective_2d'):
|
||||
return self._effective_2d
|
||||
|
||||
# Check stride, dilation, padding first (fast checks)
|
||||
if not (self.stride[0] == 1 and self.dilation[0] == 1 and self.padding[0] == 0):
|
||||
self._effective_2d = False
|
||||
return False
|
||||
|
||||
if self.kernel_size[0] == 1:
|
||||
self._effective_2d = True
|
||||
return True
|
||||
|
||||
# Check for tail inflation (all zeros except last slice in temporal dim)
|
||||
# Weight shape: (Out, In, T, H, W)
|
||||
with torch.no_grad():
|
||||
# Check if all slices except the last one are zero
|
||||
# Use a small threshold for float comparison or exact check for initialized weights
|
||||
if torch.sum(torch.abs(self.weight[:, :, :-1, :, :])) < 1e-6:
|
||||
self._effective_2d = True
|
||||
return True
|
||||
|
||||
self._effective_2d = False
|
||||
return False
|
||||
|
||||
def forward(
|
||||
self,
|
||||
input: Union[Tensor, List[Tensor]],
|
||||
memory_state: MemoryState = MemoryState.UNSET
|
||||
) -> Tensor:
|
||||
assert memory_state != MemoryState.UNSET
|
||||
if memory_state != MemoryState.ACTIVE:
|
||||
self.memory = None
|
||||
|
||||
# Handle 4D input directly (optimization for spatial-only blocks)
|
||||
if torch.is_tensor(input) and input.ndim == 4:
|
||||
# If 4D input is passed, we MUST be in the "effective 2D" path.
|
||||
# Skip extend_head and memory logic, go straight to conv.
|
||||
# _conv_forward handles the actual 2D dispatch.
|
||||
return super().forward(input)
|
||||
|
||||
if (
|
||||
math.isinf(self.memory_limit)
|
||||
and torch.is_tensor(input)
|
||||
and get_sequence_parallel_group() is None
|
||||
):
|
||||
return self.basic_forward(input, memory_state)
|
||||
return self.slicing_forward(input, memory_state)
|
||||
|
||||
def _conv_forward(self, input, weight, bias, *args, **kwargs):
|
||||
"""
|
||||
Override _conv_forward to work around NVIDIA Conv3d memory bug.
|
||||
@@ -92,19 +149,27 @@ class InflatedCausalConv3d(Conv3d):
|
||||
Status is logged at startup in compatibility.py.
|
||||
"""
|
||||
# Optimization: Use fast 2D conv for spatial-only operations (no temporal mixing)
|
||||
# Check: kernel_time=1, stride_time=1, dilation_time=1, padding_time=0
|
||||
if (
|
||||
self.kernel_size[0] == 1
|
||||
and self.stride[0] == 1
|
||||
and self.dilation[0] == 1
|
||||
and self.padding[0] == 0
|
||||
):
|
||||
# Reshape input: (B, C, T, H, W) -> (B*T, C, H, W)
|
||||
B, C, T, H, W = input.shape
|
||||
input_2d = input.permute(0, 2, 1, 3, 4).reshape(B * T, C, H, W)
|
||||
# Check: kernel_time=1 OR tail-inflated weights AND stride/dilation/padding compatibility
|
||||
can_use_2d_path = self.check_effective_2d()
|
||||
|
||||
# Reshape weight: (Out, In, 1, kH, kW) -> (Out, In, kH, kW)
|
||||
weight_2d = weight.squeeze(2)
|
||||
if can_use_2d_path:
|
||||
# Handle 4D input directly (B*T folded into batch)
|
||||
if input.ndim == 4:
|
||||
# Input: (BT, C, H, W)
|
||||
input_2d = input
|
||||
is_4d_input = True
|
||||
else:
|
||||
# Input: (B, C, T, H, W) -> (B*T, C, H, W)
|
||||
B, C, T, H, W = input.shape
|
||||
input_2d = input.permute(0, 2, 1, 3, 4).reshape(B * T, C, H, W)
|
||||
is_4d_input = False
|
||||
|
||||
# Prepare 2D weights
|
||||
if self.kernel_size[0] == 1:
|
||||
weight_2d = weight.squeeze(2)
|
||||
else:
|
||||
# Tail inflation: use the last temporal slice
|
||||
weight_2d = weight[:, :, -1, :, :]
|
||||
|
||||
# Conv2D params
|
||||
stride_2d = self.stride[1:]
|
||||
@@ -116,6 +181,9 @@ class InflatedCausalConv3d(Conv3d):
|
||||
input_2d, weight_2d, bias, stride_2d, padding_2d, dilation_2d, self.groups
|
||||
)
|
||||
|
||||
if is_4d_input:
|
||||
return out_2d
|
||||
|
||||
# Reshape output back: (B*T, Out, H', W') -> (B, Out, T, H', W')
|
||||
_, C_out, H_out, W_out = out_2d.shape
|
||||
out = out_2d.view(B, T, C_out, H_out, W_out).permute(0, 2, 1, 3, 4)
|
||||
@@ -240,19 +308,6 @@ class InflatedCausalConv3d(Conv3d):
|
||||
)
|
||||
return output
|
||||
|
||||
def forward(
|
||||
self,
|
||||
input: Union[Tensor, List[Tensor]],
|
||||
memory_state: MemoryState = MemoryState.UNSET
|
||||
) -> Tensor:
|
||||
assert memory_state != MemoryState.UNSET
|
||||
if memory_state != MemoryState.ACTIVE:
|
||||
self.memory = None
|
||||
if (
|
||||
math.isinf(self.memory_limit)
|
||||
and torch.is_tensor(input)
|
||||
and get_sequence_parallel_group() is None
|
||||
):
|
||||
return self.basic_forward(input, memory_state)
|
||||
return self.slicing_forward(input, memory_state)
|
||||
|
||||
|
||||
Reference in New Issue
Block a user