PyTorch torch.nn.Conv1d Function
PyTorch torch.nn Reference Manual
torch.nn.Conv1dIs a one-dimensional convolution module in PyTorch.
Mainly used for processing sequence data, such as text, audio, and time series signals.
Function Definition
torch.nn.Conv1d(in_channels, out_channels, kernel_size, stride=1, padding=0, dilation=1, groups=1, bias=True)
Parameter Description
in_channels: Number of input channelsout_channels: Number of output channelskernel_size: Kernel size
Usage Examples
Example 1: Basic Usage
Example
import torch
import torch.nn as nn
# 1D convolution: input channels=3, output channels=64, kernel size=3
conv1d = nn.Conv1d(in_channels=3, out_channels=64, kernel_size=3, padding=1)
# Input: batch=4, channels=3, sequence length=100
x = torch.randn(4, 3, 100)
output = conv1d(x)
print("Input shape:", x.shape)
print("Output shape:", output.shape)
import torch.nn as nn
# 1D convolution: input channels=3, output channels=64, kernel size=3
conv1d = nn.Conv1d(in_channels=3, out_channels=64, kernel_size=3, padding=1)
# Input: batch=4, channels=3, sequence length=100
x = torch.randn(4, 3, 100)
output = conv1d(x)
print("Input shape:", x.shape)
print("Output shape:", output.shape)
Example 2: Text Classification
Example
import torch
import torch.nn as nn
class TextCNN(nn.Module):
def __init__(self, vocab_size, embed_dim, num_classes):
super(TextCNN, self).__init__()
# Embedding layer: (batch, seq) -> (batch, embed, seq)
self.embedding = nn.Embedding(vocab_size, embed_dim)
# Multiple convolution kernels of different sizes
self.convs = nn.ModuleList([
nn.Conv1d(embed_dim, 128, kernel_size=k)
for k in [2, 3, 4, 5]
])
self.fc = nn.Linear(128 * 4, num_classes)
def forward(self, x):
# x: (batch, seq)
x = self.embedding(x) # (batch, seq, embed)
x = x.permute(0, 2, 1) # (batch, embed, seq)
# Convolution + ReLU + GlobalMaxPool
pooled = []
for conv in self.convs:
c = conv(x) # (batch, 128, seq')
c = nn.functional.relu(c)
c = c.max(dim=2)[0] # Global max pooling
pooled.append(c)
x = torch.cat(pooled, dim=1)
return self.fc(x)
# Test
model = TextCNN(vocab_size=10000, embed_dim=128, num_classes=2)
x = torch.randint(0, 10000, (4, 50)) # batch=4, seq=50
output = model(x)
print("Input shape:", x.shape)
print("Output shape:", output.shape)
import torch.nn as nn
class TextCNN(nn.Module):
def __init__(self, vocab_size, embed_dim, num_classes):
super(TextCNN, self).__init__()
# Embedding layer: (batch, seq) -> (batch, embed, seq)
self.embedding = nn.Embedding(vocab_size, embed_dim)
# Multiple convolution kernels of different sizes
self.convs = nn.ModuleList([
nn.Conv1d(embed_dim, 128, kernel_size=k)
for k in [2, 3, 4, 5]
])
self.fc = nn.Linear(128 * 4, num_classes)
def forward(self, x):
# x: (batch, seq)
x = self.embedding(x) # (batch, seq, embed)
x = x.permute(0, 2, 1) # (batch, embed, seq)
# Convolution + ReLU + GlobalMaxPool
pooled = []
for conv in self.convs:
c = conv(x) # (batch, 128, seq')
c = nn.functional.relu(c)
c = c.max(dim=2)[0] # Global max pooling
pooled.append(c)
x = torch.cat(pooled, dim=1)
return self.fc(x)
# Test
model = TextCNN(vocab_size=10000, embed_dim=128, num_classes=2)
x = torch.randint(0, 10000, (4, 50)) # batch=4, seq=50
output = model(x)
print("Input shape:", x.shape)
print("Output shape:", output.shape)
Example 3: Audio Processing
Example
import torch
import torch.nn as nn
# Audio feature extraction
conv_audio = nn.Conv1d(in_channels=1, out_channels=32, kernel_size=5, stride=2)
# Single-channel audio: batch=4, channels=1, sample points=16000
audio = torch.randn(4, 1, 16000)
output = conv_audio(audio)
print("Input shape:", audio.shape)
print("Output shape:", output.shape)
print("Output length:", output.shape[2])
import torch.nn as nn
# Audio feature extraction
conv_audio = nn.Conv1d(in_channels=1, out_channels=32, kernel_size=5, stride=2)
# Single-channel audio: batch=4, channels=1, sample points=16000
audio = torch.randn(4, 1, 16000)
output = conv_audio(audio)
print("Input shape:", audio.shape)
print("Output shape:", output.shape)
print("Output length:", output.shape[2])
Use Cases
- Text Classification: TextCNN
- Audio Processing: Speech feature extraction
- Time Series: Signal filtering
Note: The input tensor shape is (batch, channels, length).
Other Extensions