You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

基于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分数达标

优化分析与建议

一、模型代码核心错误修复

  1. 解码器掩码强制因果约束:
    移除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)
    
  2. 源序列掩码简化:
    原生Transformer编码器的src_mask用于遮盖序列内部特定位置,而src_key_padding_mask已处理pad token,因此直接将src_mask设为None即可,无需生成全False掩码。
  3. 生成器添加层归一化:
    在生成器的全连接层前加入层归一化,稳定训练过程:
    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))
    
  4. 位置编码验证:
    确保PositionalEncoding实现正确,添加dropout层防止过拟合,同时确认位置编码与嵌入维度完全匹配。

二、训练策略调整

  1. 梯度累积提升有效批次:
    受显存限制无法增大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()
    
  2. 学习率调度优化:
    随机初始化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)
    
  3. 增加训练轮数:
    7万样本+batch4的配置下,3-4轮训练量远不足以让随机初始化Transformer收敛,至少训练10轮,直到验证Loss不再下降再停止。
  4. 加权交叉熵解决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))
    

三、任务适配优化

  1. 标签平滑减少过度自信:
    启用标签平滑(label_smoothing=0.1),降低模型对高频token的过度拟合:
    loss_fn = nn.CrossEntropyLoss(weight=weights.to(device), label_smoothing=0.1)
    
  2. 解码阶段添加结构约束:
    句法树是结构化输出,解码时可加入括号匹配约束:比如维护括号栈,生成闭括号时仅允许匹配当前栈顶的开括号,过滤无效候选token。
  3. 数据增强:
    对输入句子进行同义词替换、语序微调(不改变句法结构),增加训练数据多样性,提升模型泛化能力。

四、对比实验参考

Fairseq的Transformer实现内置了梯度裁剪、权重初始化、学习率调度等训练技巧,而BART作为预训练模型已学习了语言结构信息,因此性能远超随机初始化模型。你的原生实现需要补全上述训练技巧、延长训练时间,才能逐步接近它们的性能。

内容的提问来源于stack exchange,提问作者Crea Teeth

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.07.11 16:54:56