diff --git a/src/h3_blackwell_runtime/vae_encoder.py b/src/h3_blackwell_runtime/vae_encoder.py index 869fa81..a52186a 100644 --- a/src/h3_blackwell_runtime/vae_encoder.py +++ b/src/h3_blackwell_runtime/vae_encoder.py @@ -41,39 +41,29 @@ NIN_LEVELS = frozenset({1, 3, 5}) DOWNSAMPLE_LEVELS = frozenset({0, 1, 2, 3}) -def _causal_conv3d(x, weight, bias, *, kernel_size, stride, spatial_padding): - """Reflect-spatial / causal-temporal 3D conv (matches upstream_vae.CausalConv3d). +def _causal_conv3d(x, weight, bias, *, kernel_size, stride, spatial_padding, temporal_causal): + """Causal 3D conv (matches upstream_vae.CausalConv3d). - Front-only zero temporal padding of ``(kernel_size - 1)`` whenever - ``spatial_padding > 0`` (the reference pads with ``(k//2, k//2, 0)`` in - time), no temporal padding otherwise. Single-frame keyframes go through - the same code path -- the causal front-zero is the reference's - ``autopad="causal_zero"`` behaviour for the encoder. + - ``spatial_padding > 0``: reflect H and W by ``spatial_padding`` on each side. + - ``temporal_causal``: front-zero T by ``kernel_size - 1`` (the reference's + ``causal_padding[0] * 2`` for ``causal_padding[0]=1``, which is every + 3D-causal conv in the H3 VAE). A single-frame input truncates the + temporal taps to the center 1 instead of convolving zero rows. + - If neither applies: no padding at all (reference early-return). """ - # Reference semantics: - # x = F.pad(x, (s_p, s_p, s_p, s_p, 0, 0), mode="reflect") -- spatial - # if x.shape[2] == 1: return super().forward(x, autopad="causal_zero") - # x = F.pad(x, (0, 0, 0, 0, (k-1), 0), mode="constant") -- temporal - # return super().forward(x) if x.shape[2] == 1: - # Keyframe path (matches reference's `autopad="causal_zero"`): apply - # spatial-reflect pad, then truncate the temporal taps to the center 1 - # so the conv acts as a 2D op on the single frame. + # Keyframe path (matches reference's `autopad="causal_zero"`): + # apply spatial-reflect pad (if any), then run an effective 2D conv + # by slicing the kernel to its center temporal tap (T_out stays 1). if spatial_padding > 0: - # F.pad for 5D input (B,C,T,H,W) uses (W_l, W_r, H_l, H_r, T_l, T_r). x = F.pad(x, (spatial_padding, spatial_padding, spatial_padding, spatial_padding, 0, 0), mode="reflect") half = (kernel_size - 1) // 2 kernel_5d = weight[:, :, half:half + 1, :, :] return F.conv3d(x, kernel_5d, bias, (1, stride[1], stride[2]), (0, 0, 0)) if spatial_padding > 0: - return _reflect_and_causal(x, weight, bias, kernel_size, stride, spatial_padding) - return F.conv3d(x, weight, bias, stride, (0, 0, 0)) - - -def _reflect_and_causal(x, weight, bias, kernel_size, stride, spatial_padding): - """Two-step pad (reference CausalConv3d): spatial reflect, temporal front-zero.""" - x = F.pad(x, (spatial_padding, spatial_padding, spatial_padding, spatial_padding, 0, 0), mode="reflect") - x = F.pad(x, (0, 0, 0, 0, kernel_size - 1, 0)) + x = F.pad(x, (spatial_padding, spatial_padding, spatial_padding, spatial_padding, 0, 0), mode="reflect") + if temporal_causal: + x = F.pad(x, (0, 0, 0, 0, kernel_size - 1, 0)) return F.conv3d(x, weight, bias, stride, (0, 0, 0)) @@ -85,29 +75,31 @@ def _group_norm_3d(x, weight, bias): def _resnet(x, p): - residual = x if p["nin"] is None else _causal_conv3d(x, p["nin"][0], p["nin"][1], kernel_size=1, stride=(1, 1, 1), spatial_padding=0) - h = _causal_conv3d(F.silu(_group_norm_3d(x, p["norm1_w"], p["norm1_b"])), p["conv1_w"], p["conv1_b"], kernel_size=3, stride=(1, 1, 1), spatial_padding=1) - h = _causal_conv3d(F.silu(_group_norm_3d(h, p["norm2_w"], p["norm2_b"])), p["conv2_w"], p["conv2_b"], kernel_size=3, stride=(1, 1, 1), spatial_padding=1) + # nin_shortcut uses CausalConv3d(k=1, padding=1) in the reference. + residual = x if p["nin"] is None else _causal_conv3d(x, p["nin"][0], p["nin"][1], kernel_size=1, stride=(1, 1, 1), spatial_padding=1, temporal_causal=True) + h = _causal_conv3d(F.silu(_group_norm_3d(x, p["norm1_w"], p["norm1_b"])), p["conv1_w"], p["conv1_b"], kernel_size=3, stride=(1, 1, 1), spatial_padding=1, temporal_causal=True) + h = _causal_conv3d(F.silu(_group_norm_3d(h, p["norm2_w"], p["norm2_b"])), p["conv2_w"], p["conv2_b"], kernel_size=3, stride=(1, 1, 1), spatial_padding=1, temporal_causal=True) return h.add_(residual) def _downsample(x, p): if p["space"] == 2: # Reference Downsample3D pads H and W by +1 reflect before the conv. - # F.pad for 5D (B,C,T,H,W) uses (W_l, W_r, H_l, H_r, T_l, T_r). x = F.pad(x, (1, 1, 1, 1, 0, 0), mode="reflect") - return _causal_conv3d(x, p["w"], p["b"], kernel_size=3, stride=(p["time"], p["space"], p["space"]), spatial_padding=0) + # Conv uses padding=(1,0,0) -> causal_padding=(1,0,0), so spatial pad=0, + # temporal front-zero is applied. + return _causal_conv3d(x, p["w"], p["b"], kernel_size=3, stride=(p["time"], p["space"], p["space"]), spatial_padding=0, temporal_causal=True) def _encoder_run(x, E): - h = _causal_conv3d(x, E["conv_in"][0], E["conv_in"][1], kernel_size=3, stride=(1, 1, 1), spatial_padding=1) + h = _causal_conv3d(x, E["conv_in"][0], E["conv_in"][1], kernel_size=3, stride=(1, 1, 1), spatial_padding=1, temporal_causal=True) for level in E["down"]: for blk in level["blocks"]: h = _resnet(h, blk) if level["down"] is not None: h = _downsample(h, level["down"]) h = F.silu(_group_norm_3d(h, E["norm_out_w"], E["norm_out_b"])) - return _causal_conv3d(h, E["conv_out"][0], E["conv_out"][1], kernel_size=3, stride=(1, 1, 1), spatial_padding=1) + return _causal_conv3d(h, E["conv_out"][0], E["conv_out"][1], kernel_size=3, stride=(1, 1, 1), spatial_padding=1, temporal_causal=True) class MiniMaxH3VideoVAEEncoder(nn.Module):