如何手动实现PyTorch nn.Conv1d?含WaveNet音频场景及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 -1accounts forpadding=1: whent=0,k=0givest+k-1=-1(we use 0 for padded values). Whent=8819,k=2givest+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
- 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
- Convolution Logic: The kernel elements now skip
dilation-1input steps. For output channeloand time stept, 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 are0 +0*2 -2 = -2(padded 0),0+1*2-2=0,0+2*2-2=2. - For
t=8819, positions are8819+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 =5input 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

