You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何手动实现PyTorch nn.Conv1d?含WaveNet音频场景及dilation=2情况

Understanding nn.Conv1d Forward Pass in WaveNet (dilation=1 and dilation=2)

Great question — let's break this down step by step, starting with the standard convolution (dilation=1) then moving to dilated convolution (dilation=2). I'll also include NumPy implementations for both cases to make the math concrete.


First: Why Output Shape Stays (1,16,8820)?

First, let's confirm the padding setting. PyTorch's nn.Conv1d defaults to padding=0, but here your output sequence length matches the input. The formula for 1D convolution output length is:

output_len = (input_len - kernel_size + 2*padding) // stride + 1

With input_len=8820, kernel_size=3, stride=1 (default), solving for padding to get output_len=8820:

8820 = (8820 - 3 + 2*padding) // 1 + 1 → padding=1

So your Conv1d layer must have padding=1 (either explicitly set or inferred by WaveNet's architecture to maintain sequence length).


Case 1: Dilation=1 (Standard Convolution)

Core Math Logic

Your input is (batch_size, in_channels, seq_len) = (1,16,8820), and the weight tensor is (out_channels, in_channels, kernel_size) = (16,16,3).

For each output channel o (0-15) and each time step t (0-8819) in the output, the value is calculated as:

output[0, o, t] = sum(
    input[0, i, t + k - 1] * weight[o, i, k]
    for i in range(16)  # iterate over input channels
    for k in range(3)   # iterate over kernel elements
)
  • The t + k -1 accounts for padding=1: when t=0, k=0 gives t+k-1=-1 (we use 0 for padded values). When t=8819, k=2 gives t+k-1=8820 (also padded to 0).
  • This uses cross-correlation (PyTorch's default convolution behavior, not strict mathematical convolution).

NumPy Implementation

Here's how to replicate the forward pass with NumPy:

import numpy as np

# Simulate input and weights
input_np = np.random.randn(1, 16, 8820)  # (batch, in_ch, seq)
weight_np = np.random.randn(16, 16, 3)   # (out_ch, in_ch, kernel)
padding = 1

# Step 1: Apply zero padding to input sequence
padded_input = np.pad(input_np, ((0,0), (0,0), (padding, padding)), mode='constant')  # shape (1,16,8822)

# Step 2: Compute output (naive loop version)
output_np = np.zeros((1, 16, 8820))

for o in range(16):
    for t in range(8820):
        # Extract 3-length window from padded input for all input channels
        window = padded_input[0, :, t:t+3]  # shape (16,3)
        # Multiply with weight for output channel o, sum all elements
        output_np[0, o, t] = np.sum(window * weight_np[o, :, :])

# Optional: Faster vectorized version
# Create sliding windows using stride tricks
window_shape = (16, 3)
strides = (padded_input.strides[1], padded_input.strides[2])
sliding_windows = np.lib.stride_tricks.as_strided(
    padded_input[0],
    shape=(8820, *window_shape),
    strides=(padded_input.strides[2], *strides)
)  # shape (8820,16,3)

# Batch matrix multiplication to compute all outputs at once
vectorized_output = np.einsum('oik,tik->to', weight_np, sliding_windows)
vectorized_output = vectorized_output[np.newaxis, ...].transpose(0,2,1)  # shape (1,16,8820)

Case 2: Dilation=2 (Dilated Convolution, WaveNet's Key Feature)

Dilation (or "atrous convolution") spaces out kernel elements, expanding the receptive field without increasing kernel size or computation — this is critical for WaveNet to capture long-range audio patterns efficiently.

Key Changes

  1. Padding Adjustment: To keep output length the same as input, we need padding=2. Using the dilated convolution output length formula:
output_len = (input_len + 2*padding - dilation*(kernel_size-1) - 1) // stride + 1

Solving for output_len=8820, dilation=2, kernel_size=3:

8820 = (8820 + 2*padding - 2*(3-1) -1) //1 +1 → padding=2
  1. Convolution Logic: The kernel elements now skip dilation-1 input steps. For output channel o and time step t, the calculation becomes:
output[0, o, t] = sum(
    input[0, i, t + k*dilation - padding] * weight[o, i, k]
    for i in range(16)
    for k in range(3)
)
  • For t=0, the input positions are 0 +0*2 -2 = -2 (padded 0), 0+1*2-2=0, 0+2*2-2=2.
  • For t=8819, positions are 8819+0*2-2=8817, 8819+1*2-2=8819, 8819+2*2-2=8821 (padded 0).
  • The receptive field per output step jumps from 3 to 1 + (3-1)*2 =5 input steps, letting the model "see" further back in the audio sequence.

NumPy Implementation

dilation = 2
padding = 2

# Step1: Pad input
padded_input_dilated = np.pad(input_np, ((0,0), (0,0), (padding, padding)), mode='constant')  # shape (1,16,8824)

# Step2: Compute output (naive loop version)
output_dilated = np.zeros((1,16,8820))

for o in range(16):
    for t in range(8820):
        # Extract dilated window: positions t, t+2, t+4 in padded input
        window_indices = [t + k*dilation for k in range(3)]
        window = padded_input_dilated[0, :, window_indices]  # shape (16,3)
        output_dilated[0, o, t] = np.sum(window * weight_np[o, :, :])

# Optional: Faster vectorized version
# Create dilated sliding windows using stride tricks
sliding_windows_dilated = np.lib.stride_tricks.as_strided(
    padded_input_dilated[0],
    shape=(8820, 16, 3),
    strides=(padded_input_dilated.strides[2], padded_input_dilated.strides[1], padded_input_dilated.strides[2]*dilation)
)

vectorized_output_dilated = np.einsum('oik,tik->to', weight_np, sliding_windows_dilated)
vectorized_output_dilated = vectorized_output_dilated[np.newaxis, ...].transpose(0,2,1)  # shape (1,16,8820)

内容的提问来源于stack exchange,提问作者Keith

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.05.08 19:47:38