Fix causal temporal padding to match reference (2k-1 front zeros when spatial_padding>0)
This commit is contained in:
parent
390135fca6
commit
06fd79a8fd
1 changed files with 14 additions and 15 deletions
|
|
@ -41,24 +41,23 @@ NIN_LEVELS = frozenset({1, 3, 5})
|
|||
DOWNSAMPLE_LEVELS = frozenset({0, 1, 2, 3})
|
||||
|
||||
|
||||
def _causal_front_padding(t_in: int, kernel: int, stride: int) -> int:
|
||||
"""Front-only zero padding: ``kernel-1-(ceil(t_in/stride)-1)*stride``."""
|
||||
t_out = math.ceil(t_in / stride)
|
||||
return max(0, kernel - 1 - (t_out - 1) * stride)
|
||||
|
||||
|
||||
def _causal_conv3d(x, weight, bias, *, kernel_size, stride, spatial_padding):
|
||||
"""Reflect-spatial / causal-temporal 3D conv (single-frame -> truncated taps)."""
|
||||
"""Reflect-spatial / causal-temporal 3D conv (matches upstream_vae.CausalConv3d).
|
||||
|
||||
Front-only zero temporal padding of ``(kernel_size - 1)`` (two taps on each
|
||||
side) whenever ``spatial_padding > 0`` (the reference pads with ``(k//2,
|
||||
k//2, 0)`` in time), no temporal padding otherwise. A single input frame
|
||||
has its temporal taps truncated instead of convolving zero frames
|
||||
(``autopad="causal_zero"``) so a keyframe never leaks zero into the latent.
|
||||
"""
|
||||
t = x.shape[2]
|
||||
if t == 1:
|
||||
if spatial_padding > 0:
|
||||
half = (kernel_size - 1) // 2
|
||||
kernel = weight[:, :, half : half + 2 * spatial_padding + 1]
|
||||
return F.conv3d(x, kernel, bias, (1, stride[1], stride[2]), (0, spatial_padding, spatial_padding))
|
||||
if t == 1 and spatial_padding > 0:
|
||||
half = (kernel_size - 1) // 2
|
||||
kernel = weight[:, :, half : half + 2 * spatial_padding + 1]
|
||||
return F.conv3d(x, kernel, bias, (1, stride[1], stride[2]), (0, spatial_padding, spatial_padding))
|
||||
if spatial_padding > 0:
|
||||
x = F.pad(x, (0, 0, 0, 0, kernel_size - 1, 0))
|
||||
return F.conv3d(x, weight, bias, stride, (0, spatial_padding, spatial_padding))
|
||||
front = _causal_front_padding(t, kernel_size, stride[0])
|
||||
if front > 0:
|
||||
x = F.pad(x, (0, 0, 0, 0, front, 0))
|
||||
return F.conv3d(x, weight, bias, stride, (0, spatial_padding, spatial_padding))
|
||||
|
||||
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue