nin_shortcut is a plain 1x1x1 conv (no causal padding)
This commit is contained in:
parent
ea6ab87a34
commit
0b7217485c
1 changed files with 1 additions and 1 deletions
|
|
@ -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)
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue