使用卷积神经网络实现文本分类的工程实践与模型调优
模型核心架构解析
文本卷积网络(TextCNN)通过一维卷积操作提取局部上下文特征,其数据流向可划分为三个关键模块:
- 词向量映射层:将离散的字符或词汇序列转换为连续的低维稠密张量,输出维度为
[batch_size, sequence_length, embedding_dim]。 - 多尺度特征提取器:并行部署多个不同窗口大小的卷积核(如 2-gram 至 5-gram),经过非线性激活函数(如 ReLU)与全局最大池化操作后,聚合最具判别力的局部语义片段。
- 线性分类头:将多路池化输出的特征向量进行拼接,映射至全连接层,根据预设的类别总数生成预测概率分布。
数据预处理与流水线实现
高质量的模型训练依赖于稳健的数据摄入机制。以下实现涵盖了原始文本解析、词表构建以及 PyTorch 数据集封装,针对大规模语料进行了内存与索引优化。
import os
import torch
import torch.nn as nn
from collections import Counter
from torch.utils.data import Dataset
def load_corpus(filepath: str, limit: int = None) -> tuple[list[str], list[int]]:
"""解析原始语料文件,返回句子列表与对应标签"""
sentences, targets = [], []
with open(filepath, 'r', encoding='utf-8') as fh:
for line in fh:
line = line.strip()
if not line:
continue
parts = line.split('\t')
if len(parts) == 2:
sentences.append(parts[0])
targets.append(int(parts[1]))
if limit is not None:
return sentences[:limit], targets[:limit]
return sentences, targets
def construct_vocabulary_and_embed(raw_sents: list[str], dim: int) -> tuple[dict[str, int], nn.Embedding]:
"""基于词频统计构建词表,并初始化嵌入层"""
special_tokens = {"[PAD]": 0, "[UNK]": 1}
freq_map = Counter()
for sent in raw_sents:
freq_map.update(sent)
vocab = special_tokens.copy()
vocab.update({char: idx + len(special_tokens) for idx, char in enumerate(freq_map)})
embed_layer = nn.Embedding(
num_embeddings=len(vocab),
embedding_dim=dim,
padding_idx=0
)
return vocab, embed_layer
class SentenceDataset(Dataset):
def __init__(self, vocab: dict[str, int], sents: list[str], targets: list[int], max_seq: int):
self.vocab = vocab
self.sents = sents
self.targets = torch.tensor(targets, dtype=torch.long)
self.max_seq = max_seq
def __len__(self) -> int:
return len(self.sents)
def __getitem__(self, idx: int) -> tuple[torch.Tensor, torch.Tensor]:
raw = self.sents[idx]
seq_len = min(len(raw), self.max_seq)
indices = [self.vocab.get(char, 1) for char in raw[:seq_len]]
padding_len = self.max_seq - seq_len
padded_indices = indices + [0] * padding_len
return torch.tensor(padded_indices, dtype=torch.long), self.targets[idx]