nin_shortcut is a plain 1x1x1 conv (no causal padding)

This commit is contained in:
Daniel Maddern 2026-08-19 21:49:40 +07:00
parent ea6ab87a34
commit 0b7217485c

View file

@ -76,7 +76,7 @@ def _group_norm_3d(x, weight, bias):
def _resnet(x, p):
# 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)
residual = x if p["nin"] is None else F.conv3d(x, p["nin"][0], p["nin"][1], (1, 1, 1))
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)