基于PyTorch的Transformer Seq2Seq句法分析任务训练难题求助
原生PyTorch Transformer Seq2Seq任务优化求助
任务说明
给定符合特定语法的句子,生成对应句法分析树表示,示例:
- I like apples. -> (S (NP I) (VP like apples))
- John says he likes apples -> (S (NP John) (VP says (S (NP he) (VP likes apples))))
注:括号对应非终结规则(语法成分/POS),每种开闭括号为独立token(如S-opening、S-closing等),文本采用Byte-Pair-Encoding分词。
训练参数
- 损失函数:Cross-Entropy-Loss
- 优化器:AdamW(β₁=0.9, β₂=0.98)
- 学习率:尝试过1e-4、1e-5
- 最大token长度:512
- 批次大小:4(受GPU显存限制)
- 数据集规模:约7万条样本
- 训练轮数:3-4轮
训练结果
- Loss趋于平稳后不再下降,无明显收敛趋势
- F1分数无有效提升,始终处于较低水平
- 评估Loss远高于训练集(评估样本普遍更长且按长度排序)
模型代码
import math import torch import torch.nn as nn from torch.nn import Transformer class Generator(nn.Module): def __init__(self, hidden_size: int, vocab_size: int): super().__init__() self.fc = nn.Linear(hidden_size, vocab_size) def forward(self, x): return self.fc(x) class PositionalEncoding(nn.Module): def __init__(self, d_model: int, max_len: int = 512): super().__init__() position = torch.arange(max_len).unsqueeze(1) div_term = torch.exp(torch.arange(0, d_model, 2) * (-math.log(10000.0) / d_model)) pe = torch.zeros(max_len, 1, d_model) pe[:, 0, 0::2] = torch.sin(position * div_term) pe[:, 0, 1::2] = torch.cos(position * div_term) self.register_buffer('pe', pe) def forward(self, x: torch.Tensor) -> torch.Tensor: return x + self.pe[:x.size(0)] class SimpleTransformer(nn.Module): def __init__(self, vocab_size: int, ntokens=512, d_model=512, num_layers=6, bidirectional=False, device="cpu"): super().__init__() self.d_model = d_model self.src_embed = nn.Embedding(vocab_size, self.d_model) self.tgt_embed = nn.Embedding(vocab_size, self.d_model) self.positional_encoder = PositionalEncoding(d_model=self.d_model, max_len=ntokens) self.model = Transformer(d_model=self.d_model, batch_first=True, num_encoder_layers=num_layers, num_decoder_layers=num_layers) self.bidirectional = bidirectional self.generator = Generator(hidden_size=self.d_model, vocab_size=vocab_size) # 仅全连接层 self.device = device def forward(self, in_ids, l_ids, in_masks, l_masks): in_ids = self.src_embed(in_ids.long()) * math.sqrt(self.d_model) # 按d_model平方根缩放 in_ids = self.positional_encoder(in_ids) l_ids = self.tgt_embed(l_ids.long()) * math.sqrt(self.d_model) l_ids = self.positional_encoder(l_ids) # 生成掩码 src_seq_len = in_ids.size(1) tgt_seq_len = l_ids.size(1) src_mask = torch.zeros(src_seq_len, src_seq_len, device=self.device).type(torch.bool) if not self.bidirectional: tgt_mask = torch.triu(torch.full((tgt_seq_len, tgt_seq_len), float('-inf'), device=self.device), diagonal=1) else: tgt_mask = torch.zeros(tgt_seq_len, tgt_seq_len, device=self.device).type(torch.bool) in_masks = in_masks == 0.0 # 掩码pad token l_masks = l_masks == 0.0 # 掩码pad token out = self.model(src=in_ids, tgt=l_ids, src_mask=src_mask, tgt_mask=tgt_mask, src_key_padding_mask=in_masks, tgt_key_padding_mask=l_masks) return self.generator(out)
已尝试调整
- 编码器/解码器层数:6、3、1
- 模型维度:512(固定)
- 前馈网络维度:2048(默认)
- 位置编码:sin/cos编码、无位置编码
核心问题
无论输入如何,模型始终输出重复序列,无法生成有效句法结构,示例输出:
(s (s (s (s (s (s (s (s (s (s (s (s (s ...
对比实验
- 使用facebook/fairseq训练同任务模型:Loss显著下降,性能远优于原生PyTorch实现
- 使用huggingface预训练bart-base模型:性能最优,验证Loss约0.025,F1分数达标
优化分析与建议
一、模型代码核心错误修复
- 解码器掩码强制因果约束:
移除bidirectional对解码器掩码的控制,Seq2Seq解码器必须使用因果掩码(torch.triu生成的-inf掩码),否则会看见未来token,导致模型崩溃、输出重复。修改后解码器掩码固定为:tgt_mask = torch.triu(torch.full((tgt_seq_len, tgt_seq_len), float('-inf'), device=self.device), diagonal=1) - 源序列掩码简化:
原生Transformer编码器的src_mask用于遮盖序列内部特定位置,而src_key_padding_mask已处理pad token,因此直接将src_mask设为None即可,无需生成全False掩码。 - 生成器添加层归一化:
在生成器的全连接层前加入层归一化,稳定训练过程:class Generator(nn.Module): def __init__(self, hidden_size: int, vocab_size: int): super().__init__() self.norm = nn.LayerNorm(hidden_size) self.fc = nn.Linear(hidden_size, vocab_size) def forward(self, x): return self.fc(self.norm(x)) - 位置编码验证:
确保PositionalEncoding实现正确,添加dropout层防止过拟合,同时确认位置编码与嵌入维度完全匹配。
二、训练策略调整
- 梯度累积提升有效批次:
受显存限制无法增大batch size时,启用梯度累积(如每8步累积一次梯度再更新参数),等效于batch size=32,降低梯度噪声:accumulation_steps = 8 for step, batch in enumerate(dataloader): loss = model(**batch) loss = loss / accumulation_steps loss.backward() if (step + 1) % accumulation_steps == 0: optimizer.step() optimizer.zero_grad() scheduler.step() - 学习率调度优化:
随机初始化Transformer需用学习率预热+余弦退火策略,避免初始学习率过高导致模型发散:from torch.optim.lr_scheduler import CosineAnnealingLR, LambdaLR warmup_steps = 1000 total_steps = len(dataloader) * 10 # 按10轮计算总步数 optimizer = AdamW(model.parameters(), lr=1e-4, betas=(0.9,0.98)) # 预热阶段:前1000步线性提升学习率至1e-4 warmup_scheduler = LambdaLR(optimizer, lr_lambda=lambda step: min(step/warmup_steps, 1.0)) # 预热后余弦退火 scheduler = CosineAnnealingLR(optimizer, T_max=total_steps - warmup_steps) - 增加训练轮数:
7万样本+batch4的配置下,3-4轮训练量远不足以让随机初始化Transformer收敛,至少训练10轮,直到验证Loss不再下降再停止。 - 加权交叉熵解决token分布不平衡:
句法树中开闭括号、非终结符的token频率差异极大,交叉熵会偏向高频token(如导致输出重复的(s),统计目标序列token频率后,给低频token设置更高权重:# 统计token频率 token_counts = torch.bincount(all_target_tokens) weights = 1.0 / token_counts.float() weights = weights / weights.sum() * vocab_size # 归一化权重 loss_fn = nn.CrossEntropyLoss(weight=weights.to(device))
三、任务适配优化
- 标签平滑减少过度自信:
启用标签平滑(label_smoothing=0.1),降低模型对高频token的过度拟合:loss_fn = nn.CrossEntropyLoss(weight=weights.to(device), label_smoothing=0.1) - 解码阶段添加结构约束:
句法树是结构化输出,解码时可加入括号匹配约束:比如维护括号栈,生成闭括号时仅允许匹配当前栈顶的开括号,过滤无效候选token。 - 数据增强:
对输入句子进行同义词替换、语序微调(不改变句法结构),增加训练数据多样性,提升模型泛化能力。
四、对比实验参考
Fairseq的Transformer实现内置了梯度裁剪、权重初始化、学习率调度等训练技巧,而BART作为预训练模型已学习了语言结构信息,因此性能远超随机初始化模型。你的原生实现需要补全上述训练技巧、延长训练时间,才能逐步接近它们的性能。
内容的提问来源于stack exchange,提问作者Crea Teeth
相关产品推荐
相关产品推荐

