如何利用Chinese-Poetry数据集训练高质量生成对抗网络

1次阅读
没有评论

共计 1933 个字符,预计需要花费 5 分钟才能阅读完成。

image.webp

背景与痛点

中文诗歌生成一直是 NLP 领域极具挑战性的任务,主要面临两个核心问题:

如何利用 Chinese-Poetry 数据集训练高质量生成对抗网络

  • 数据稀缺性 :高质量的中文诗歌数据集较少,公开可用的更少。传统方法依赖人工收集整理,耗时耗力。
  • 风格单一性 :现有模型生成的诗歌往往缺乏多样性,容易出现重复模式或生硬拼接的问题。

Chinese-Poetry 数据集收录了唐宋等朝代的经典诗歌,格式统一且质量较高,是训练 GAN 模型的优质语料。但原始数据需要经过专门处理才能用于生成任务。

数据预处理

  1. 数据清洗
  2. 去除非中文字符和标点
  3. 统一诗歌格式(如五言 / 七言)
  4. 过滤掉残缺或重复的诗句

  5. 分词与编码

  6. 使用 jieba 进行分词
  7. 构建词表(建议保留 top 5000 高频词)
  8. 将诗句转换为 token 序列

  9. 韵律处理

  10. 提取平仄信息
  11. 标注韵脚
  12. 构建韵律特征向量

一个简单的预处理示例:

import jieba
from collections import Counter

# 基本分词处理
def preprocess(text):
    words = jieba.lcut(text)
    return [w for w in words if w.strip()]

# 构建词表
def build_vocab(texts, max_size=5000):
    word_counts = Counter()
    for text in texts:
        word_counts.update(preprocess(text))
    return [w for w,_ in word_counts.most_common(max_size)]

模型选型

常见的文本生成 GAN 架构对比:

  • SeqGAN:适合处理序列数据,但训练复杂度高
  • TextGAN:直接优化文本生成指标,但容易模式坍塌
  • MaliGAN:使用重要性采样,稳定性更好

对于诗歌生成任务,推荐使用改进版的 SeqGAN 架构:

  1. 生成器采用 LSTM 网络
  2. 判别器使用 CNN+Attention
  3. 加入韵律约束损失

核心实现

以下是 PyTorch 实现的关键部分:

import torch
import torch.nn as nn

class Generator(nn.Module):
    def __init__(self, vocab_size, embedding_dim, hidden_dim):
        super().__init__()
        self.embedding = nn.Embedding(vocab_size, embedding_dim)
        self.lstm = nn.LSTM(embedding_dim, hidden_dim)
        self.linear = nn.Linear(hidden_dim, vocab_size)

    def forward(self, x):
        emb = self.embedding(x)
        out, _ = self.lstm(emb)
        return self.linear(out)

class Discriminator(nn.Module):
    def __init__(self, vocab_size, embedding_dim):
        super().__init__()
        self.embedding = nn.Embedding(vocab_size, embedding_dim)
        self.conv = nn.Sequential(nn.Conv1d(embedding_dim, 64, 3, padding=1),
            nn.ReLU(),
            nn.MaxPool1d(2)
        )
        self.linear = nn.Linear(64, 1)

    def forward(self, x):
        emb = self.embedding(x).transpose(1, 2)
        features = self.conv(emb).squeeze(2)
        return torch.sigmoid(self.linear(features))

性能优化

关键超参数设置建议:

  • batch_size: 32-128(根据显存调整)
  • 生成器学习率: 0.001
  • 判别器学习率: 0.0001
  • 使用 Adam 优化器
  • 加入梯度裁剪(clipnorm=5)

训练技巧:

  1. 先预训练生成器
  2. 采用课程学习策略
  3. 定期保存检查点

避坑指南

常见问题及解决方案:

  1. 模式坍塌 :增加判别器的能力,加入多样性奖励
  2. 梯度消失 :使用 Wasserstein GAN 或梯度惩罚
  3. 训练不稳定 :调整学习率比例(D:G≈1:4)

效果评估

建议从三个维度评估生成质量:

  1. 流畅性 :BLEU、困惑度
  2. 韵律合规 :平仄和押韵准确率
  3. 意境美感 :人工评分(1- 5 分)

总结与展望

通过合理的数据预处理和模型设计,基于 Chinese-Poetry 数据集可以训练出不错的诗歌生成模型。但目前还存在一些挑战:

  • 如何更好地捕捉诗歌的意境?
  • 能否实现特定风格(如豪放 / 婉约)的定向生成?
  • 如何评估生成诗歌的文学价值?

期待看到更多创新的解决方案。如果你有好的想法,欢迎在评论区分享讨论!

正文完
 0
评论(没有评论)